2828import threading
2929from typing import TYPE_CHECKING , cast
3030
31+ from ._logger import log
3132from ._record_update import RecordUpdate
3233from ._utils .asyncio import get_running_loop , run_coro_with_timeout
34+ from ._utils .net import (
35+ InterfacesType ,
36+ IPVersion ,
37+ add_multicast_member ,
38+ drop_multicast_member ,
39+ new_respond_socket ,
40+ normalize_interface_choice ,
41+ )
3342from ._utils .time import current_time_millis
3443from .const import _CACHE_CLEANUP_INTERVAL
3544
3847
3948
4049from ._listener import AsyncListener
41- from ._transport import _WrappedTransport , make_wrapped_transport
50+ from ._transport import _strip_zone , _WrappedTransport , make_wrapped_transport
4251
4352_CLOSE_TIMEOUT = 3000 # ms
4453
4554
55+ def _interface_key (interface : str | tuple [tuple [str , int , int ], int ]) -> tuple [str , int ]:
56+ """Return the (address, scope_id) an interface choice maps to, for diffing.
57+
58+ Must produce the same key shape as ``_WrappedTransport.interface_key`` so
59+ the desired set (from ``normalize_interface_choice``) and the current set
60+ (from the bound senders) diff against each other.
61+ """
62+ if isinstance (interface , tuple ):
63+ return (_strip_zone (interface [0 ][0 ]), interface [0 ][2 ])
64+ return (interface , 0 )
65+
66+
4667class AsyncEngine :
4768 """An engine wraps sockets in the event loop."""
4869
4970 __slots__ = (
5071 "_cleanup_timer" ,
5172 "_listen_socket" ,
73+ "_listen_transport" ,
5274 "_respond_sockets" ,
5375 "_setup_task" ,
5476 "loop" ,
@@ -72,6 +94,7 @@ def __init__(
7294 self .senders : list [_WrappedTransport ] = []
7395 self .running_future : asyncio .Future [bool | None ] | None = None
7496 self ._listen_socket = listen_socket
97+ self ._listen_transport : _WrappedTransport | None = None
7598 self ._respond_sockets = respond_sockets
7699 self ._cleanup_timer : asyncio .TimerHandle | None = None
77100 self ._setup_task : asyncio .Task [None ] | None = None
@@ -98,8 +121,6 @@ async def _async_setup(self, loop_thread_ready: threading.Event | None) -> None:
98121
99122 async def _async_create_endpoints (self ) -> None :
100123 """Create endpoints to send and receive."""
101- assert self .loop is not None
102- loop = self .loop
103124 reader_sockets = []
104125 sender_sockets = []
105126 if self ._listen_socket :
@@ -110,22 +131,110 @@ async def _async_create_endpoints(self) -> None:
110131 sender_sockets .append (s )
111132
112133 for s in reader_sockets :
113- transport , protocol = await loop .create_datagram_endpoint ( # type: ignore[type-var]
114- lambda : AsyncListener (self .zc ), # type: ignore[arg-type, return-value]
115- sock = s ,
116- )
117- # Register the wrapped transport before releasing the engine's
118- # handle so a concurrent shutdown always sees ``s`` in exactly
119- # one place; do not add an ``await`` between these two steps.
120- self .protocols .append (cast (AsyncListener , protocol ))
121- self .readers .append (make_wrapped_transport (cast (asyncio .DatagramTransport , transport )))
122- if s in sender_sockets :
123- self .senders .append (make_wrapped_transport (cast (asyncio .DatagramTransport , transport )))
134+ reader = await self ._async_wrap_socket (s , s in sender_sockets )
135+ # The wrap above does not await before returning, so releasing
136+ # the engine's pending handle here keeps ``s`` in exactly one
137+ # place from a concurrent shutdown's point of view.
124138 if s is self ._listen_socket :
139+ # Keep a handle to the shared listen socket so interface
140+ # rescans can add/drop multicast memberships on it.
141+ self ._listen_transport = reader
125142 self ._listen_socket = None
126143 if s in self ._respond_sockets :
127144 self ._respond_sockets .remove (s )
128145
146+ async def _async_wrap_socket (self , sock : socket .socket , is_sender : bool ) -> _WrappedTransport :
147+ """Adopt a socket into a transport, register it, and return the reader wrapper."""
148+ assert self .loop is not None
149+ transport , protocol = await self .loop .create_datagram_endpoint ( # type: ignore[type-var]
150+ lambda : AsyncListener (self .zc ), # type: ignore[arg-type, return-value]
151+ sock = sock ,
152+ )
153+ datagram_transport = cast (asyncio .DatagramTransport , transport )
154+ reader = make_wrapped_transport (datagram_transport )
155+ # No ``await`` between wrapping and registering so a concurrent
156+ # shutdown always sees the transport in exactly one place.
157+ self .protocols .append (cast (AsyncListener , protocol ))
158+ self .readers .append (reader )
159+ if is_sender :
160+ self .senders .append (make_wrapped_transport (datagram_transport ))
161+ return reader
162+
163+ async def async_update_interfaces (
164+ self ,
165+ interfaces : InterfacesType ,
166+ ip_version : IPVersion ,
167+ apple_p2p : bool ,
168+ ) -> bool :
169+ """Reconcile sender/reader sockets to the live interface set.
170+
171+ Adds a per-interface responder socket for each interface that
172+ appeared and tears down the socket for each interface that
173+ disappeared, diffing on the bound address. The shared listen
174+ socket (including the Default single-family dual-use socket) is
175+ never torn down here. Returns whether any responder socket was
176+ added, so the caller can skip re-announcing when nothing appeared.
177+ """
178+ assert self .loop is not None
179+ normalized = normalize_interface_choice (interfaces , ip_version )
180+ desired = {_interface_key (interface ): interface for interface in normalized }
181+ current = {wrapped .interface_key : wrapped for wrapped in self .senders }
182+ listen_transport = self ._listen_transport
183+ listen_socket = listen_transport .sock if listen_transport is not None else None
184+
185+ for bind_address , wrapped in current .items ():
186+ if bind_address in desired :
187+ continue
188+ if listen_transport is not None and wrapped .transport is listen_transport .transport :
189+ # The shared listen / dual-use socket is not a per-interface
190+ # sender; leaving the group or closing it would break receive.
191+ continue
192+ self ._async_close_sender (wrapped , listen_socket )
193+
194+ added = False
195+ for bind_address , interface in desired .items ():
196+ if bind_address in current :
197+ continue
198+ if await self ._async_add_interface (interface , listen_socket , apple_p2p ):
199+ added = True
200+ return added
201+
202+ async def _async_add_interface (
203+ self ,
204+ interface : str | tuple [tuple [str , int , int ], int ],
205+ listen_socket : socket .socket | None ,
206+ apple_p2p : bool ,
207+ ) -> bool :
208+ """Join the multicast group and adopt a responder socket for one interface.
209+
210+ Returns whether a responder socket was actually added.
211+ """
212+ # A unicast instance has no listen socket, so membership is only
213+ # ever managed when ``listen_socket`` is present.
214+ if listen_socket is not None and not add_multicast_member (listen_socket , interface ):
215+ log .debug ("Interface %r not added: could not join multicast group" , interface )
216+ return False
217+ respond_socket = new_respond_socket (interface , apple_p2p = apple_p2p , unicast = self .zc .unicast )
218+ if respond_socket is None :
219+ if listen_socket is not None :
220+ drop_multicast_member (listen_socket , interface )
221+ log .debug ("Interface %r not added: no responder socket" , interface )
222+ return False
223+ await self ._async_wrap_socket (respond_socket , is_sender = True )
224+ return True
225+
226+ def _async_close_sender (self , wrapped : _WrappedTransport , listen_socket : socket .socket | None ) -> None :
227+ """Drop a per-interface sender's wrappers/protocol and close its transport."""
228+ transport = wrapped .transport
229+ self .protocols = [
230+ p for p in self .protocols if p .transport is None or p .transport .transport is not transport
231+ ]
232+ self .readers = [w for w in self .readers if w .transport is not transport ]
233+ self .senders = [w for w in self .senders if w .transport is not transport ]
234+ if listen_socket is not None :
235+ drop_multicast_member (listen_socket , wrapped .multicast_interface )
236+ transport .close ()
237+
129238 def _async_cache_cleanup (self ) -> None :
130239 """Periodic cache cleanup."""
131240 now = current_time_millis ()
0 commit comments