SPDX-License-Identifier: AGPL-3.0-only

Bind multicast receivers to the wildcard address and join the group on each
usable IPv4 interface.  OpenBSD associates membership with one interface, so
INADDR_ANY alone is not sufficient on multihomed hosts.  Select the invitation
address as the outgoing multicast interface as well, ensuring the datagram
source matches the address authenticated by XRCP.

Construct the listener directly because network-udp's wildcard path requires
IP_PKTINFO, which OpenBSD does not have; recvfrom is sufficient here because
discovery only needs the peer address.  Use SO_REUSEPORT as well as
SO_REUSEADDR so concurrent listeners can bind the discovery port.

Index: src/Simplex/RemoteControl/Discovery.hs
--- src/Simplex/RemoteControl/Discovery.hs.orig
+++ src/Simplex/RemoteControl/Discovery.hs
@@ -2,10 +2,10 @@
 {-# LANGUAGE DuplicateRecordFields #-}
 {-# LANGUAGE FlexibleContexts #-}
 {-# LANGUAGE GADTs #-}
-{-# LANGUAGE LambdaCase #-}
 {-# LANGUAGE NamedFieldPuns #-}
 {-# LANGUAGE OverloadedStrings #-}
 {-# LANGUAGE PatternSynonyms #-}
+{-# LANGUAGE ScopedTypeVariables #-}
 
 module Simplex.RemoteControl.Discovery
   ( pattern ANY_ADDR_V4,
@@ -23,8 +23,9 @@
 import Control.Monad
 import Data.ByteString (ByteString)
 import Data.Default (def)
-import Data.List (delete, find, partition)
-import Data.Maybe (mapMaybe)
+import Data.Functor (($>))
+import Data.List (delete, find, nub, partition)
+import Data.Maybe (catMaybes, mapMaybe)
 import Data.String (IsString)
 import qualified Data.Text as T
 import Data.Word (Word16)
@@ -37,7 +38,7 @@
 import Simplex.Messaging.Transport.Client (TransportHost (..))
 import Simplex.Messaging.Transport.Server (mkTransportServerConfig, runTransportServerSocket, startTCPServer)
 import Simplex.Messaging.Util (ifM, tshow)
-import Simplex.RemoteControl.Discovery.Multicast (setMembership)
+import Simplex.RemoteControl.Discovery.Multicast (setMembership, setMulticastInterface)
 import Simplex.RemoteControl.Types
 import UnliftIO
 
@@ -100,45 +101,63 @@
           TLS.serverSupported = defaultSupportedParams
         }
 
-withSender :: (UDP.UDPSocket -> IO a) -> IO a
-withSender = bracket (UDP.clientSocket MULTICAST_ADDR_V4 DISCOVERY_PORT False) (UDP.close)
+withSender :: TransportHost -> (UDP.UDPSocket -> IO a) -> IO a
+withSender host = bracket (openSender host) UDP.close
+
+openSender :: TransportHost -> IO UDP.UDPSocket
+openSender (THIPv4 host) =
+  bracketOnError (UDP.clientSocket MULTICAST_ADDR_V4 DISCOVERY_PORT False) UDP.close $ \sock -> do
+    setMulticastInterface (UDP.udpSocket sock) $ N.tupleToHostAddress host
+    pure sock
+openSender _ = ioError $ userError "multicast discovery requires an IPv4 host"
 
 withListener :: TMVar Int -> (UDP.ListenSocket -> IO a) -> IO a
-withListener subscribers = bracket (openListener subscribers) (closeListener subscribers)
+withListener _subscribers action = bracket openListener closeListener $ action . fst
 
-openListener :: TMVar Int -> IO UDP.ListenSocket
-openListener subscribers = do
-  sock <- UDP.serverSocket (MULTICAST_ADDR_V4, read DISCOVERY_PORT)
+openListener :: IO (UDP.ListenSocket, [N.HostAddress])
+openListener = bracketOnError (N.socket N.AF_INET N.Datagram N.defaultProtocol) N.close $ \raw -> do
+  N.setSocketOption raw N.ReuseAddr 1
+  N.setSocketOption raw N.ReusePort 1
+  N.withFdSocket raw N.setCloseOnExecIfNeeded
+  N.bind raw listenerSockAddr4
+  let sock = UDP.ListenSocket raw listenerSockAddr4 False
   logDebug $ "Discovery listener socket: " <> tshow sock
-  let raw = UDP.listenSocket sock
-  -- N.setSocketOption raw N.Broadcast 1
-  joinMulticast subscribers raw (listenerHostAddr4 sock)
-  pure sock
-
-closeListener :: TMVar Int -> UDP.ListenSocket -> IO ()
-closeListener subscribers sock =
-  partMulticast subscribers (UDP.listenSocket sock) (listenerHostAddr4 sock) `finally` UDP.stop sock
-
-joinMulticast :: TMVar Int -> N.Socket -> N.HostAddress -> IO ()
-joinMulticast subscribers sock group = do
-  now <- atomically $ takeTMVar subscribers
-  when (now == 0) $ do
-    setMembership sock group True >>= \case
-      Left e -> atomically (putTMVar subscribers now) >> logError ("setMembership failed " <> tshow e)
-      Right () -> atomically $ putTMVar subscribers (now + 1)
-
-partMulticast :: TMVar Int -> N.Socket -> N.HostAddress -> IO ()
-partMulticast subscribers sock group = do
-  now <- atomically $ takeTMVar subscribers
-  when (now == 1) $
-    setMembership sock group False >>= \case
-      Left e -> atomically (putTMVar subscribers now) >> logError ("setMembership failed " <> tshow e)
-      Right () -> atomically $ putTMVar subscribers (now - 1)
-
-listenerHostAddr4 :: UDP.ListenSocket -> N.HostAddress
-listenerHostAddr4 sock = case UDP.mySockAddr sock of
-  N.SockAddrInet _port host -> host
-  _ -> error "MULTICAST_ADDR_V4 is V4"
+  interfaces <- multicastInterfaces
+  joined <- fmap catMaybes . forM interfaces $ \interface ->
+    (joinMulticast raw multicastHostAddr4 interface $> Just interface)
+      `catch` \(e :: SomeException) -> logError ("Joining multicast interface " <> tshow interface <> " failed: " <> tshow e) $> Nothing
+  when (null joined) $ ioError $ userError "no IPv4 interface accepted the multicast membership"
+  pure (sock, joined)
+
+closeListener :: (UDP.ListenSocket, [N.HostAddress]) -> IO ()
+closeListener (sock, interfaces) =
+  forM_ interfaces (partMulticast raw multicastHostAddr4) `finally` UDP.stop sock
+  where
+    raw = UDP.listenSocket sock
+
+joinMulticast :: N.Socket -> N.HostAddress -> N.HostAddress -> IO ()
+joinMulticast sock group interface = setMembership sock group interface True
+
+partMulticast :: N.Socket -> N.HostAddress -> N.HostAddress -> IO ()
+partMulticast sock group interface =
+  setMembership sock group interface False
+    `catch` \(e :: SomeException) -> logError $ "Leaving multicast interface " <> tshow interface <> " failed: " <> tshow e
+
+multicastInterfaces :: IO [N.HostAddress]
+multicastInterfaces = do
+  interfaces <- nub . mapMaybe interfaceAddress <$> getNetworkInterfaces
+  pure $ nub $ interfaces <> [N.tupleToHostAddress (0, 0, 0, 0)]
+  where
+    interfaceAddress NetworkInterface {ipv4 = IPv4 address} = case N.hostAddressToTuple address of
+      (0, 0, 0, 0) -> Nothing
+      (255, 255, 255, 255) -> Nothing
+      _ -> Just address
+
+listenerSockAddr4 :: N.SockAddr
+listenerSockAddr4 = N.SockAddrInet (read DISCOVERY_PORT) (N.tupleToHostAddress (0, 0, 0, 0))
+
+multicastHostAddr4 :: N.HostAddress
+multicastHostAddr4 = N.tupleToHostAddress (224, 0, 0, 251)
 
 recvAnnounce :: UDP.ListenSocket -> IO (N.SockAddr, ByteString)
 recvAnnounce sock = do
