SPDX-License-Identifier: AGPL-3.0-only

Multicast membership is per socket and per interface.  Let callers name the
interface for each membership, provide IP_MULTICAST_IF for the sender, and
report setsockopt failures with their actual errno.

Index: src/Simplex/RemoteControl/Discovery/Multicast.hsc
--- src/Simplex/RemoteControl/Discovery/Multicast.hsc.orig
+++ src/Simplex/RemoteControl/Discovery/Multicast.hsc
@@ -1,8 +1,10 @@
 module Simplex.RemoteControl.Discovery.Multicast
-  ( setMembership
+  ( setMembership,
+    setMulticastInterface,
   ) where
 
 import Foreign (Ptr, allocaBytes, castPtr, pokeByteOff)
+import Foreign.C.Error (throwErrnoIfMinus1_)
 import Foreign.C.Types (CInt (..))
 import Network.Socket
 
@@ -10,20 +12,27 @@
 
 {- | Toggle multicast group membership.
 
-NB: Group membership is per-host, not per-process. A socket is only used to access system interface for groups.
+Group membership is per socket, so every listening socket has to join and
+leave the group independently.
 -}
-setMembership :: Socket -> HostAddress -> Bool -> IO (Either CInt ())
-setMembership sock group membership = allocaBytes #{size struct ip_mreq} $ \mReqPtr -> do
+setMembership :: Socket -> HostAddress -> HostAddress -> Bool -> IO ()
+setMembership sock group interface membership = allocaBytes #{size struct ip_mreq} $ \mReqPtr -> do
   #{poke struct ip_mreq, imr_multiaddr} mReqPtr group
-  #{poke struct ip_mreq, imr_interface} mReqPtr (0 :: HostAddress) -- attempt to contact the group on ANY interface
+  #{poke struct ip_mreq, imr_interface} mReqPtr interface
   withFdSocket sock $ \fd -> do
-    rc <- c_setsockopt fd c_IPPROTO_IP flag (castPtr mReqPtr) (#{size struct ip_mreq})
-    if rc == 0
-      then pure $ Right ()
-      else pure $ Left rc
+    throwErrnoIfMinus1_ "setsockopt(multicast membership)" $
+      c_setsockopt fd c_IPPROTO_IP flag (castPtr mReqPtr) (#{size struct ip_mreq})
   where
     flag = if membership then c_IP_ADD_MEMBERSHIP else c_IP_DROP_MEMBERSHIP
 
+-- | Select the IPv4 interface used for outgoing multicast datagrams.
+setMulticastInterface :: Socket -> HostAddress -> IO ()
+setMulticastInterface sock interface = allocaBytes #{size struct in_addr} $ \interfacePtr -> do
+  pokeByteOff interfacePtr 0 interface
+  withFdSocket sock $ \fd ->
+    throwErrnoIfMinus1_ "setsockopt(multicast interface)" $
+      c_setsockopt fd c_IPPROTO_IP c_IP_MULTICAST_IF (castPtr interfacePtr) (#{size struct in_addr})
+
 #ifdef mingw32_HOST_OS
 
 foreign import stdcall unsafe "setsockopt"
@@ -33,6 +42,9 @@
 c_IP_ADD_MEMBERSHIP  = 12
 c_IP_DROP_MEMBERSHIP = 13
 
+c_IP_MULTICAST_IF :: CInt
+c_IP_MULTICAST_IF = 9
+
 #else
 
 foreign import ccall unsafe "setsockopt"
@@ -42,6 +54,9 @@
 c_IP_ADD_MEMBERSHIP  = #const IP_ADD_MEMBERSHIP
 c_IP_DROP_MEMBERSHIP = #const IP_DROP_MEMBERSHIP
 
+c_IP_MULTICAST_IF :: CInt
+c_IP_MULTICAST_IF = #const IP_MULTICAST_IF
+
 #endif
 
 c_IPPROTO_IP :: CInt
