3030
3131from ._record_update import RecordUpdate
3232from ._utils .asyncio import get_running_loop , run_coro_with_timeout
33+ from ._utils .net import (
34+ InterfacesType ,
35+ IPVersion ,
36+ add_multicast_member ,
37+ drop_multicast_member ,
38+ new_respond_socket ,
39+ normalize_interface_choice ,
40+ )
3341from ._utils .time import current_time_millis
3442from .const import _CACHE_CLEANUP_INTERVAL
3543
3846
3947
4048from ._listener import AsyncListener
41- from ._transport import _WrappedTransport , make_wrapped_transport
49+ from ._transport import _strip_zone , _WrappedTransport , make_wrapped_transport
4250
4351_CLOSE_TIMEOUT = 3000 # ms
4452
4553
54+ def _interface_key (interface : str | tuple [tuple [str , int , int ], int ]) -> tuple [str , int ]:
55+ """Return the (address, scope_id) an interface choice maps to, for diffing.
56+
57+ Must produce the same key shape as ``_WrappedTransport.interface_key`` so
58+ the desired set (from ``normalize_interface_choice``) and the current set
59+ (from the bound senders) diff against each other.
60+ """
61+ if isinstance (interface , tuple ):
62+ return (_strip_zone (interface [0 ][0 ]), interface [0 ][2 ])
63+ return (interface , 0 )
64+
65+
4666class AsyncEngine :
4767 """An engine wraps sockets in the event loop."""
4868
4969 __slots__ = (
5070 "_cleanup_timer" ,
5171 "_listen_socket" ,
72+ "_listen_transport" ,
5273 "_respond_sockets" ,
5374 "_setup_task" ,
5475 "loop" ,
@@ -72,6 +93,7 @@ def __init__(
7293 self .senders : list [_WrappedTransport ] = []
7394 self .running_future : asyncio .Future [bool | None ] | None = None
7495 self ._listen_socket = listen_socket
96+ self ._listen_transport : _WrappedTransport | None = None
7597 self ._respond_sockets = respond_sockets
7698 self ._cleanup_timer : asyncio .TimerHandle | None = None
7799 self ._setup_task : asyncio .Task [None ] | None = None
@@ -98,8 +120,6 @@ async def _async_setup(self, loop_thread_ready: threading.Event | None) -> None:
98120
99121 async def _async_create_endpoints (self ) -> None :
100122 """Create endpoints to send and receive."""
101- assert self .loop is not None
102- loop = self .loop
103123 reader_sockets = []
104124 sender_sockets = []
105125 if self ._listen_socket :
@@ -110,22 +130,108 @@ async def _async_create_endpoints(self) -> None:
110130 sender_sockets .append (s )
111131
112132 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 )))
133+ reader = await self ._async_wrap_socket (s , s in sender_sockets )
134+ # The wrap above does not await before returning, so releasing
135+ # the engine's pending handle here keeps ``s`` in exactly one
136+ # place from a concurrent shutdown's point of view.
124137 if s is self ._listen_socket :
138+ # Keep a handle to the shared listen socket so interface
139+ # rescans can add/drop multicast memberships on it.
140+ self ._listen_transport = reader
125141 self ._listen_socket = None
126142 if s in self ._respond_sockets :
127143 self ._respond_sockets .remove (s )
128144
145+ async def _async_wrap_socket (self , sock : socket .socket , is_sender : bool ) -> _WrappedTransport :
146+ """Adopt a socket into a transport, register it, and return the reader wrapper."""
147+ assert self .loop is not None
148+ transport , protocol = await self .loop .create_datagram_endpoint ( # type: ignore[type-var]
149+ lambda : AsyncListener (self .zc ), # type: ignore[arg-type, return-value]
150+ sock = sock ,
151+ )
152+ datagram_transport = cast (asyncio .DatagramTransport , transport )
153+ reader = make_wrapped_transport (datagram_transport )
154+ # No ``await`` between wrapping and registering so a concurrent
155+ # shutdown always sees the transport in exactly one place.
156+ self .protocols .append (cast (AsyncListener , protocol ))
157+ self .readers .append (reader )
158+ if is_sender :
159+ self .senders .append (make_wrapped_transport (datagram_transport ))
160+ return reader
161+
162+ async def async_update_interfaces (
163+ self ,
164+ interfaces : InterfacesType ,
165+ ip_version : IPVersion ,
166+ apple_p2p : bool ,
167+ ) -> bool :
168+ """Reconcile sender/reader sockets to the live interface set.
169+
170+ Adds a per-interface responder socket for each interface that
171+ appeared and tears down the socket for each interface that
172+ disappeared, diffing on the bound address. The shared listen
173+ socket (including the Default single-family dual-use socket) is
174+ never torn down here. Returns whether any responder socket was
175+ added, so the caller can skip re-announcing when nothing appeared.
176+ """
177+ assert self .loop is not None
178+ normalized = normalize_interface_choice (interfaces , ip_version )
179+ desired = {_interface_key (interface ): interface for interface in normalized }
180+ current = {wrapped .interface_key : wrapped for wrapped in self .senders }
181+ listen_transport = self ._listen_transport
182+ listen_socket = listen_transport .sock if listen_transport is not None else None
183+
184+ for bind_address , wrapped in current .items ():
185+ if bind_address in desired :
186+ continue
187+ if listen_transport is not None and wrapped .transport is listen_transport .transport :
188+ # The shared listen / dual-use socket is not a per-interface
189+ # sender; leaving the group or closing it would break receive.
190+ continue
191+ self ._async_close_sender (wrapped , listen_socket )
192+
193+ added = False
194+ for bind_address , interface in desired .items ():
195+ if bind_address in current :
196+ continue
197+ if await self ._async_add_interface (interface , listen_socket , apple_p2p ):
198+ added = True
199+ return added
200+
201+ async def _async_add_interface (
202+ self ,
203+ interface : str | tuple [tuple [str , int , int ], int ],
204+ listen_socket : socket .socket | None ,
205+ apple_p2p : bool ,
206+ ) -> bool :
207+ """Join the multicast group and adopt a responder socket for one interface.
208+
209+ Returns whether a responder socket was actually added.
210+ """
211+ # A unicast instance has no listen socket, so membership is only
212+ # ever managed when ``listen_socket`` is present.
213+ if listen_socket is not None and not add_multicast_member (listen_socket , interface ):
214+ return False
215+ respond_socket = new_respond_socket (interface , apple_p2p = apple_p2p , unicast = self .zc .unicast )
216+ if respond_socket is None :
217+ if listen_socket is not None :
218+ drop_multicast_member (listen_socket , interface )
219+ return False
220+ await self ._async_wrap_socket (respond_socket , is_sender = True )
221+ return True
222+
223+ def _async_close_sender (self , wrapped : _WrappedTransport , listen_socket : socket .socket | None ) -> None :
224+ """Drop a per-interface sender's wrappers/protocol and close its transport."""
225+ transport = wrapped .transport
226+ self .protocols = [
227+ p for p in self .protocols if p .transport is None or p .transport .transport is not transport
228+ ]
229+ self .readers = [w for w in self .readers if w .transport is not transport ]
230+ self .senders = [w for w in self .senders if w .transport is not transport ]
231+ if listen_socket is not None :
232+ drop_multicast_member (listen_socket , wrapped .multicast_interface )
233+ transport .close ()
234+
129235 def _async_cache_cleanup (self ) -> None :
130236 """Periodic cache cleanup."""
131237 now = current_time_millis ()
0 commit comments