Skip to content

Commit 67c69b0

Browse files
committed
Merge branch 'response' into integration
2 parents fd6313f + 0058c03 commit 67c69b0

1 file changed

Lines changed: 61 additions & 24 deletions

File tree

src/zeroconf/_core.py

Lines changed: 61 additions & 24 deletions
Original file line numberDiff line numberDiff line change
@@ -28,7 +28,7 @@
2828
import sys
2929
import threading
3030
from types import TracebackType # noqa # used in type hints
31-
from typing import Awaitable, Dict, List, Optional, Tuple, Type, Union, cast
31+
from typing import Any, Awaitable, Dict, List, Optional, Tuple, Type, Union, cast
3232

3333
from ._cache import DNSCache
3434
from ._dns import DNSQuestion, DNSQuestionType
@@ -105,6 +105,45 @@
105105
_REGISTER_BROADCASTS = 3
106106

107107

108+
class _WrappedTransport:
109+
"""A wrapper for transports."""
110+
111+
__slots__ = (
112+
'transport',
113+
'is_ipv6',
114+
'sock',
115+
'fileno',
116+
'sock_name',
117+
)
118+
119+
def __init__(
120+
self,
121+
transport: asyncio.DatagramTransport,
122+
is_ipv6: bool,
123+
sock: socket.socket,
124+
fileno: int,
125+
sock_name: Any,
126+
) -> None:
127+
"""Initialize the wrapped transport."""
128+
self.transport = transport
129+
self.is_ipv6 = is_ipv6
130+
self.sock = sock
131+
self.fileno = fileno
132+
self.sock_name = sock_name
133+
134+
135+
def _make_wrapped_transport(transport: asyncio.DatagramTransport) -> _WrappedTransport:
136+
"""Make a wrapped transport."""
137+
sock: socket.socket = transport.get_extra_info('socket')
138+
return _WrappedTransport(
139+
transport=transport,
140+
is_ipv6=sock.family == socket.AF_INET6,
141+
sock=sock,
142+
fileno=sock.fileno(),
143+
sock_name=sock.getsockname(),
144+
)
145+
146+
108147
class AsyncEngine:
109148
"""An engine wraps sockets in the event loop."""
110149

@@ -117,8 +156,8 @@ def __init__(
117156
self.loop: Optional[asyncio.AbstractEventLoop] = None
118157
self.zc = zeroconf
119158
self.protocols: List[AsyncListener] = []
120-
self.readers: List[asyncio.DatagramTransport] = []
121-
self.senders: List[asyncio.DatagramTransport] = []
159+
self.readers: List[_WrappedTransport] = []
160+
self.senders: List[_WrappedTransport] = []
122161
self.running_event: Optional[asyncio.Event] = None
123162
self._listen_socket = listen_socket
124163
self._respond_sockets = respond_sockets
@@ -158,9 +197,9 @@ async def _async_create_endpoints(self) -> None:
158197
for s in reader_sockets:
159198
transport, protocol = await loop.create_datagram_endpoint(lambda: AsyncListener(self.zc), sock=s)
160199
self.protocols.append(cast(AsyncListener, protocol))
161-
self.readers.append(cast(asyncio.DatagramTransport, transport))
200+
self.readers.append(_make_wrapped_transport(cast(asyncio.DatagramTransport, transport)))
162201
if s in sender_sockets:
163-
self.senders.append(cast(asyncio.DatagramTransport, transport))
202+
self.senders.append(_make_wrapped_transport(cast(asyncio.DatagramTransport, transport)))
164203

165204
def _async_cache_cleanup(self) -> None:
166205
"""Periodic cache cleanup."""
@@ -186,8 +225,8 @@ def _async_shutdown(self) -> None:
186225
"""Shutdown transports and sockets."""
187226
assert self.running_event is not None
188227
self.running_event.clear()
189-
for transport in itertools.chain(self.senders, self.readers):
190-
transport.close()
228+
for wrapped_transport in itertools.chain(self.senders, self.readers):
229+
wrapped_transport.transport.close()
191230

192231
def close(self) -> None:
193232
"""Close from sync context.
@@ -221,7 +260,7 @@ def __init__(self, zc: 'Zeroconf') -> None:
221260
self.zc = zc
222261
self.data: Optional[bytes] = None
223262
self.last_time: float = 0
224-
self.transport: Optional[asyncio.DatagramTransport] = None
263+
self.transport: Optional[_WrappedTransport] = None
225264
self.sock_description: Optional[str] = None
226265
self._deferred: Dict[str, List[DNSIncoming]] = {}
227266
self._timers: Dict[str, asyncio.TimerHandle] = {}
@@ -309,7 +348,7 @@ def handle_query_or_defer(
309348
msg: DNSIncoming,
310349
addr: str,
311350
port: int,
312-
transport: asyncio.DatagramTransport,
351+
transport: _WrappedTransport,
313352
v6_flow_scope: Union[Tuple[()], Tuple[int, int]] = (),
314353
) -> None:
315354
"""Deal with incoming query packets. Provides a response if
@@ -341,7 +380,7 @@ def _respond_query(
341380
msg: Optional[DNSIncoming],
342381
addr: str,
343382
port: int,
344-
transport: asyncio.DatagramTransport,
383+
transport: _WrappedTransport,
345384
v6_flow_scope: Union[Tuple[()], Tuple[int, int]] = (),
346385
) -> None:
347386
"""Respond to a query and reassemble any truncated deferred packets."""
@@ -362,27 +401,25 @@ def error_received(self, exc: Exception) -> None:
362401
self.log_exception_once(exc, msg_str, exc)
363402

364403
def connection_made(self, transport: asyncio.BaseTransport) -> None:
365-
self.transport = cast(asyncio.DatagramTransport, transport)
366-
sock_name = self.transport.get_extra_info('sockname')
367-
sock_fileno = self.transport.get_extra_info('socket').fileno()
368-
self.sock_description = f"{sock_fileno} ({sock_name})"
404+
wrapped_transport = _make_wrapped_transport(cast(asyncio.DatagramTransport, transport))
405+
self.transport = wrapped_transport
406+
self.sock_description = f"{wrapped_transport.fileno} ({wrapped_transport.sock_name})"
369407

370408
def connection_lost(self, exc: Optional[Exception]) -> None:
371409
"""Handle connection lost."""
372410

373411

374412
def async_send_with_transport(
375413
log_debug: bool,
376-
transport: asyncio.DatagramTransport,
414+
transport: _WrappedTransport,
377415
packet: bytes,
378416
packet_num: int,
379417
out: DNSOutgoing,
380418
addr: Optional[str],
381419
port: int,
382420
v6_flow_scope: Union[Tuple[()], Tuple[int, int]] = (),
383421
) -> None:
384-
s = transport.get_extra_info('socket')
385-
ipv6_socket = s.family == socket.AF_INET6
422+
ipv6_socket = transport.is_ipv6
386423
if addr is None:
387424
real_addr = _MDNS_ADDR6 if ipv6_socket else _MDNS_ADDR
388425
else:
@@ -394,8 +431,8 @@ def async_send_with_transport(
394431
'Sending to (%s, %d) via [socket %s (%s)] (%d bytes #%d) %r as %r...',
395432
real_addr,
396433
port or _MDNS_PORT,
397-
s.fileno(),
398-
transport.get_extra_info('sockname'),
434+
transport.fileno,
435+
transport.sock_name,
399436
len(packet),
400437
packet_num + 1,
401438
out,
@@ -404,9 +441,9 @@ def async_send_with_transport(
404441
# Get flowinfo and scopeid for the IPV6 socket to create a complete IPv6
405442
# address tuple: https://docs.python.org/3.6/library/socket.html#socket-families
406443
if ipv6_socket and not v6_flow_scope:
407-
_, _, sock_flowinfo, sock_scopeid = s.getsockname()
444+
_, _, sock_flowinfo, sock_scopeid = transport.sock_name
408445
v6_flow_scope = (sock_flowinfo, sock_scopeid)
409-
transport.sendto(packet, (real_addr, port or _MDNS_PORT, *v6_flow_scope))
446+
transport.transport.sendto(packet, (real_addr, port or _MDNS_PORT, *v6_flow_scope))
410447

411448

412449
class Zeroconf(QuietLogger):
@@ -832,7 +869,7 @@ def handle_assembled_query(
832869
packets: List[DNSIncoming],
833870
addr: str,
834871
port: int,
835-
transport: asyncio.DatagramTransport,
872+
transport: _WrappedTransport,
836873
v6_flow_scope: Union[Tuple[()], Tuple[int, int]] = (),
837874
) -> None:
838875
"""Respond to a (re)assembled query.
@@ -870,7 +907,7 @@ def send(
870907
addr: Optional[str] = None,
871908
port: int = _MDNS_PORT,
872909
v6_flow_scope: Union[Tuple[()], Tuple[int, int]] = (),
873-
transport: Optional[asyncio.DatagramTransport] = None,
910+
transport: Optional[_WrappedTransport] = None,
874911
) -> None:
875912
"""Sends an outgoing packet threadsafe."""
876913
assert self.loop is not None
@@ -882,7 +919,7 @@ def async_send(
882919
addr: Optional[str] = None,
883920
port: int = _MDNS_PORT,
884921
v6_flow_scope: Union[Tuple[()], Tuple[int, int]] = (),
885-
transport: Optional[asyncio.DatagramTransport] = None,
922+
transport: Optional[_WrappedTransport] = None,
886923
) -> None:
887924
"""Sends an outgoing packet."""
888925
if self.done:

0 commit comments

Comments
 (0)