diff --git a/docs/source/pcapkit/protocols/protocol.rst b/docs/source/pcapkit/protocols/protocol.rst index e8f4a00395..bb8233f447 100644 --- a/docs/source/pcapkit/protocols/protocol.rst +++ b/docs/source/pcapkit/protocols/protocol.rst @@ -61,6 +61,7 @@ utility arguments and methods of specified protocols. .. automethod:: _make_index .. automethod:: _make_payload + .. automethod:: _lookup_next_layer .. automethod:: _decode_next_layer .. automethod:: _import_next_layer diff --git a/docs/source/pcapkit/protocols/transport/transport.rst b/docs/source/pcapkit/protocols/transport/transport.rst index 4bb4e83bba..dbed2373ef 100644 --- a/docs/source/pcapkit/protocols/transport/transport.rst +++ b/docs/source/pcapkit/protocols/transport/transport.rst @@ -19,6 +19,7 @@ which is a base class for transport layer protocols, eg. .. automethod:: register .. automethod:: analyze + .. automethod:: _make_port .. automethod:: _decode_next_layer .. autoattribute:: __layer__ diff --git a/pcapkit/protocols/internet/internet.py b/pcapkit/protocols/internet/internet.py index 2630c66ac7..734c19c090 100644 --- a/pcapkit/protocols/internet/internet.py +++ b/pcapkit/protocols/internet/internet.py @@ -245,10 +245,7 @@ def _import_next_layer(self, proto: 'int', length: 'Optional[int]' = None, *, # elif self._sigterm: from pcapkit.protocols.misc.raw import Raw as protocol # isort: skip # pylint: disable=import-outside-toplevel else: - protocol = self.__proto__[proto] # type: ignore[assignment] - if isinstance(protocol, ModuleDescriptor): - protocol = protocol.klass # type: ignore[unreachable] - self.__proto__[proto] = protocol # update mapping upon import + protocol = self._lookup_next_layer(self.__proto__, proto) next_ = protocol(file_, length, version=version, extension=extension, # type: ignore[abstract] alias=proto, packet=packet, layer=self._exlayer, protocol=self._exproto, diff --git a/pcapkit/protocols/internet/ipv6.py b/pcapkit/protocols/internet/ipv6.py index 1569495142..ea41e5f5cc 100644 --- a/pcapkit/protocols/internet/ipv6.py +++ b/pcapkit/protocols/internet/ipv6.py @@ -31,7 +31,6 @@ from pcapkit.const.ipv6.extension_header import ExtensionHeader as Enum_ExtensionHeader from pcapkit.const.reg.transtype import TransType as Enum_TransType -from pcapkit.corekit.module import ModuleDescriptor from pcapkit.corekit.multidict import OrderedMultiDict from pcapkit.corekit.protochain import ProtoChain from pcapkit.protocols.data.internet.ipv6 import IPv6 as Data_IPv6 @@ -406,10 +405,7 @@ def _import_next_layer(self, proto: 'int', length: 'Optional[int]' = None, *, # from pcapkit.protocols.misc.raw import \ Raw as protocol # isort: skip # pylint: disable=import-outside-toplevel else: - protocol = self.__proto__[proto] # type: ignore[assignment] - if isinstance(protocol, ModuleDescriptor): - protocol = protocol.klass # type: ignore[unreachable] - self.__proto__[proto] = protocol # update mapping upon import + protocol = self._lookup_next_layer(self.__proto__, proto) next_ = protocol(file_, length, version=version, extension=extension, # type: ignore[abstract] alias=proto, packet=packet, layer=self._exlayer, protocol=self._exproto, diff --git a/pcapkit/protocols/protocol.py b/pcapkit/protocols/protocol.py index bf20f34742..81afa58df1 100644 --- a/pcapkit/protocols/protocol.py +++ b/pcapkit/protocols/protocol.py @@ -46,7 +46,7 @@ if TYPE_CHECKING: from enum import IntEnum as StdlibEnum - from typing import IO, Any, DefaultDict, Optional, Type + from typing import IO, Any, Callable, DefaultDict, Optional, Type from aenum import IntEnum as AenumEnum from typing_extensions import Literal, Self @@ -412,10 +412,7 @@ def analyze(cls, proto: 'int', payload: 'bytes', **kwargs: 'Any') -> 'ProtocolBa instance. """ - protocol = cls.__proto__[proto] - if isinstance(protocol, ModuleDescriptor): - protocol = protocol.klass - cls.__proto__[proto] = protocol # update mapping upon import + protocol = cls._lookup_next_layer(cls.__proto__, proto) payload_io = io.BytesIO(payload) try: @@ -1181,6 +1178,11 @@ def _make_index(cls, name: 'str | int | StdlibEnum | AenumEnum', default: 'Optio namespace = cast('dict[int, str]', namespace) index = {v: k for k, v in namespace.items()}[name] else: + # Caught by the handler immediately below and converted, so + # this never escapes -- it is a jump to the shared "name is + # not in namespace" path, not a stdlib exception leaking out + # of the library. A pcapkit exception here would log at + # CRITICAL for something that is handled two lines later. raise KeyError(name) except KeyError as error: if default is None: @@ -1235,6 +1237,50 @@ def _make_payload(cls, data: 'Data') -> 'ProtocolBase': return proto.from_data(data[name]) + @staticmethod + def _lookup_next_layer(registry: 'DefaultDict[int, ModuleDescriptor[ProtocolBase] | Type[ProtocolBase]]', + proto: 'int') -> 'Type[ProtocolBase]': + """Look up the protocol class registered for a next layer code. + + Arguments: + registry: next layer protocol registry, i.e. :attr:`self.__proto__ + `. Passed in rather than read from the + class, so that a caller reaching the registry through an + instance keeps doing so. + proto: next layer protocol index + + Returns: + The class registered for ``proto``, or the fallback ``registry`` + declares -- normally :class:`~pcapkit.protocols.misc.raw.Raw` -- when + ``proto`` is not registered. + + Important: + ``registry`` is a :class:`collections.defaultdict`, so indexing it + with an unregistered code would *insert* that code. It is + class-level -- shared by every instance in the process -- so parsing + a single packet with an unregistered code would grow it, and make + :meth:`self.register ` afterwards report that + code as already registered. The fallback is therefore read from the + default factory rather than through a lookup that records it. + + Resolving a :class:`~pcapkit.corekit.module.ModuleDescriptor` is + still written back, since that is memoisation of an import for a + code that *is* registered rather than a new entry. + + """ + if proto in registry: + protocol = registry[proto] + if isinstance(protocol, ModuleDescriptor): + protocol = protocol.klass + registry[proto] = protocol # update mapping upon import + return protocol + + fallback = cast('Callable[[], ModuleDescriptor[ProtocolBase] | Type[ProtocolBase]]', + registry.default_factory)() + if isinstance(fallback, ModuleDescriptor): + return fallback.klass + return fallback + def _decode_next_layer(self, dict_: '_PT', proto: 'int', length: 'Optional[int]' = None, *, packet: 'Optional[dict[str, Any]]' = None) -> '_PT': r"""Decode next layer protocol. @@ -1298,10 +1344,7 @@ def _import_next_layer(self, proto: 'int', length: 'Optional[int]' = None, *, elif self._sigterm: from pcapkit.protocols.misc.raw import Raw as protocol # isort: skip # pylint: disable=import-outside-toplevel else: - protocol = self.__proto__[proto] # type: ignore[assignment] - if isinstance(protocol, ModuleDescriptor): - protocol = protocol.klass # type: ignore[unreachable] - self.__proto__[proto] = protocol # update mapping upon import + protocol = self._lookup_next_layer(self.__proto__, proto) next_ = protocol(file_, length, alias=proto, packet=packet, layer=self._exlayer, protocol=self._exproto, diff --git a/pcapkit/protocols/schema/link/ethernet.py b/pcapkit/protocols/schema/link/ethernet.py index d93dab9698..9448b56c3e 100644 --- a/pcapkit/protocols/schema/link/ethernet.py +++ b/pcapkit/protocols/schema/link/ethernet.py @@ -8,7 +8,6 @@ from pcapkit.corekit.fields.misc import PayloadField from pcapkit.corekit.fields.numbers import EnumField from pcapkit.corekit.fields.strings import BytesField -from pcapkit.corekit.module import ModuleDescriptor from pcapkit.protocols.schema.schema import Schema, schema_final __all__ = ['Ethernet'] @@ -20,14 +19,35 @@ def callback_payload(self: 'PayloadField', packet: 'dict[str, Any]') -> 'None': - """Callback function for :attr:`Ethernet.payload`.""" + """Callback function for :attr:`Ethernet.payload`. + + Args: + self: Payload field to resolve. + packet: Packet data, whose ``type`` names the next layer. + + Returns: + :obj:`None`; the resolved class is assigned to ``self.protocol``. + + Important: + The lookup goes through :meth:`ProtocolBase._lookup_next_layer + ` rather than + subscripting the registry. :attr:`Ethernet.__proto__ + ` *is* + :attr:`Link.__proto__ `, a + class-level :class:`collections.defaultdict`, so reading a missing key + inserted it -- every unregistered EtherType a parse saw was recorded as + though somebody had registered it, and + :meth:`Link.register ` + afterwards reported it as an overwrite. The helper reads the fallback + without recording it, and memoises a + :class:`~pcapkit.corekit.module.ModuleDescriptor` for a code that really + is registered, which this did not. + + """ from pcapkit.protocols.link.ethernet import Ethernet # pylint: disable=import-outside-toplevel - type_ = packet['type'] - protocol = Ethernet.__proto__[type_] - if isinstance(protocol, ModuleDescriptor): - protocol = protocol.klass - self.protocol = protocol + self.protocol = Ethernet._lookup_next_layer( # pylint: disable=protected-access + Ethernet.__proto__, packet['type']) @schema_final diff --git a/pcapkit/protocols/transport/sctp.py b/pcapkit/protocols/transport/sctp.py index 84ae8a9f3d..cc0c8b6c15 100644 --- a/pcapkit/protocols/transport/sctp.py +++ b/pcapkit/protocols/transport/sctp.py @@ -28,6 +28,7 @@ import struct from typing import TYPE_CHECKING, cast +from pcapkit.const.reg.apptype import TransportProtocol as Enum_TransportProtocol from pcapkit.const.reg.transtype import TransType as Enum_TransType from pcapkit.const.sctp.cause_code import CauseCode as Enum_CauseCode from pcapkit.const.sctp.chunk import Chunk as Enum_Chunk @@ -155,7 +156,6 @@ from pcapkit.protocols.schema.transport.sctp import \ UserInitiatedAbortCause as Schema_UserInitiatedAbortCause from pcapkit.protocols.transport.transport import Transport -from pcapkit.utilities.decorators import beholder from pcapkit.utilities.exceptions import ProtocolError, RegistryError from pcapkit.utilities.warnings import RegistryWarning, warn @@ -578,8 +578,8 @@ def make(self, chunks_value = [] schema = Schema_SCTP( - srcport=srcport, - dstport=dstport, + srcport=self._make_port(srcport, Enum_TransportProtocol.sctp), + dstport=self._make_port(dstport, Enum_TransportProtocol.sctp), vtag=vtag, chksum=b'\x00\x00\x00\x00' if chksum is None else chksum, chunks=chunks_value, @@ -812,66 +812,17 @@ def _decode_next_layer(self, dict_: 'Data_SCTP', proto: 'Optional[int]' = None, ` does for an unregistered transport type. Resolving it to :class:`~pcapkit.protocols.misc.raw.Raw` is - :meth:`self._import_next_layer `'s job. + :meth:`ProtocolBase._import_next_layer + `'s job, + which looks the PPID up through + :meth:`ProtocolBase._lookup_next_layer + ` and so + leaves :attr:`self.__proto__ ` untouched. """ return ProtocolBase._decode_next_layer( # pylint: disable=protected-access self, dict_, proto, length, packet=packet) # type: ignore[arg-type,return-value] - @beholder # type: ignore[arg-type] - def _import_next_layer(self, proto: 'int', length: 'Optional[int]' = None, *, - packet: 'Optional[dict[str, Any]]' = None) -> 'Protocol': - """Import next layer extractor. - - Arguments: - proto: payload protocol identifier (PPID) of the DATA chunk carrying - the payload, or :obj:`None` if the packet carries no user data - length: valid (*non-padding*) length - packet: packet info (passed from :meth:`self.unpack `) - - Returns: - Instance of next layer. - - Important: - This overrides :meth:`ProtocolBase._import_next_layer - ` for one - reason only: to look the PPID up **without** mutating - :attr:`self.__proto__ `. - - The registry is a :class:`collections.defaultdict`, so the base - implementation's ``self.__proto__[proto]`` *inserts* any key it is - handed. Every packet with an unregistered PPID would therefore grow - a registry that is class-level -- shared by every :class:`SCTP` - instance in the process -- and make - :meth:`self.register ` report that PPID as already - registered. The fallback itself is unchanged: an unregistered PPID - still resolves to :class:`~pcapkit.protocols.misc.raw.Raw`, which is - what :attr:`self.__proto__ ` declares as its default. - - """ - if TYPE_CHECKING: - protocol: 'Type[Protocol]' - - file_ = self._get_payload() - if length is None: - length = len(file_) - - if length == 0: - from pcapkit.protocols.misc.null import NoPayload as protocol # isort: skip # pylint: disable=import-outside-toplevel - elif self._sigterm: - from pcapkit.protocols.misc.raw import Raw as protocol # isort: skip # pylint: disable=import-outside-toplevel - elif proto in self.__proto__: - protocol = self.__proto__[proto] # type: ignore[assignment] - if isinstance(protocol, ModuleDescriptor): - protocol = protocol.klass # type: ignore[unreachable] - self.__proto__[proto] = protocol # update mapping upon import - else: - from pcapkit.protocols.misc.raw import Raw as protocol # isort: skip # pylint: disable=import-outside-toplevel - - next_ = protocol(file_, length, alias=proto, packet=packet, # type: ignore[abstract] - layer=self._exlayer, protocol=self._exproto) - return next_ - def _read_sctp_chunks(self) -> 'Chunks': """Read SCTP chunk list. diff --git a/pcapkit/protocols/transport/tcp.py b/pcapkit/protocols/transport/tcp.py index 71d9de2387..5b2cf82584 100644 --- a/pcapkit/protocols/transport/tcp.py +++ b/pcapkit/protocols/transport/tcp.py @@ -44,6 +44,7 @@ import math from typing import TYPE_CHECKING, cast +from pcapkit.const.reg.apptype import TransportProtocol as Enum_TransportProtocol from pcapkit.const.reg.transtype import TransType from pcapkit.const.tcp.checksum import Checksum as Enum_Checksum from pcapkit.const.tcp.flags import Flags as Enum_Flags @@ -549,8 +550,8 @@ def make(self, self._flags = _flag return Schema_TCP( - srcport=srcport, - dstport=dstport, + srcport=self._make_port(srcport, Enum_TransportProtocol.tcp), + dstport=self._make_port(dstport, Enum_TransportProtocol.tcp), seq=seq_no, ack=ack_no, offset={ diff --git a/pcapkit/protocols/transport/transport.py b/pcapkit/protocols/transport/transport.py index ec0629a9a2..fc55b51db2 100644 --- a/pcapkit/protocols/transport/transport.py +++ b/pcapkit/protocols/transport/transport.py @@ -15,10 +15,11 @@ import io from typing import TYPE_CHECKING, Generic +from pcapkit.const.reg.apptype import AppType as Enum_AppType from pcapkit.corekit.module import ModuleDescriptor from pcapkit.protocols.protocol import _PT, _ST from pcapkit.protocols.protocol import ProtocolBase as Protocol -from pcapkit.utilities.exceptions import StructError, UnsupportedCall, stacklevel +from pcapkit.utilities.exceptions import RegistryError, StructError, UnsupportedCall, stacklevel from pcapkit.utilities.logging import DEVMODE, get_logger from pcapkit.utilities.warnings import RegistryWarning, warn @@ -27,6 +28,8 @@ from typing_extensions import Literal + from pcapkit.const.reg.apptype import TransportProtocol as Enum_TransportProtocol + __all__ = ['Transport'] @@ -89,7 +92,7 @@ def register(cls, code: 'int', protocol: 'ModuleDescriptor[Protocol] | Type[Prot if isinstance(protocol, ModuleDescriptor): protocol = protocol.klass if not issubclass(protocol, Protocol): - raise TypeError(f'protocol must be a Protocol subclass, not {protocol!r}') + raise RegistryError(f'protocol must be a Protocol subclass, not {protocol!r}') if code in cls.__proto__: warn(f'port {code} already registered, overwriting', RegistryWarning) cls.__proto__[code] = protocol @@ -108,13 +111,8 @@ def analyze(cls, ports: 'tuple[int, int]', payload: 'bytes', **kwargs: 'Any') -> instance. """ - if ports[0] in cls.__proto__: - protocol = cls.__proto__[ports[0]] - else: - protocol = cls.__proto__[ports[1]] - - if isinstance(protocol, ModuleDescriptor): - protocol = protocol.klass + protocol = cls._lookup_next_layer( + cls.__proto__, ports[0] if ports[0] in cls.__proto__ else ports[1]) payload_io = io.BytesIO(payload) try: @@ -136,6 +134,36 @@ def analyze(cls, ports: 'tuple[int, int]', payload: 'bytes', **kwargs: 'Any') -> # Utilities. ########################################################################## + @staticmethod + def _make_port(port: 'Enum_AppType | int', + proto: 'Enum_TransportProtocol') -> 'Enum_AppType': + """Resolve a port number to its application type. + + Arguments: + port: port number, or the application type itself + proto: transport protocol the port belongs to, which is what + distinguishes e.g. TCP/80 from UDP/80 + + Returns: + The :class:`~pcapkit.const.reg.apptype.AppType` for ``port``. + + Important: + :meth:`self.make ` accepts a bare :obj:`int` for a + port, and the schema field only converts one on the way *out* (in + :meth:`PortEnumField.pre_process + `), + leaving the schema attribute holding whatever it was handed. A + constructed packet therefore reached :meth:`self.read + ` with an :obj:`int` where a parsed one carries an + :class:`~pcapkit.const.reg.apptype.AppType`, and reading ``.port`` + off it raised :exc:`AttributeError`. Normalising here keeps the two + paths agreeing on the type the schema declares. + + """ + if isinstance(port, Enum_AppType): + return port + return Enum_AppType.get(port, proto=proto) + def _decode_next_layer(self, dict_: '_PT', ports: 'tuple[int, int]', length: 'Optional[int]' = None, *, # type: ignore[override] packet: 'Optional[dict[str, Any]]' = None) -> '_PT': # pylint: disable=arguments-renamed """Decode next layer protocol. @@ -153,12 +181,27 @@ def _decode_next_layer(self, dict_: '_PT', ports: 'tuple[int, int]', length: 'Op Returns: Current protocol with next layer extracted. + Important: + The port is forwarded **whether or not it is registered**, since + :meth:`ProtocolBase._import_next_layer + ` passes + it on as ``alias`` and :class:`~pcapkit.protocols.misc.raw.Raw` + records it as :attr:`Data_Raw.protocol + `. Dropping it -- as + this used to, by falling back to :obj:`None` -- anonymised the very + case the field is most useful for: a payload on a port we do not + decode is then indistinguishable from one on port 22. The lower port + is the one carried, for the same reason it is the primary lookup + key. :meth:`SCTP._decode_next_layer + ` and + :meth:`Internet._import_next_layer + ` + already behave this way for an unregistered PPID and transport type. + """ sort_port = sorted(ports) - if sort_port[0] in self.__proto__: - proto = sort_port[0] - elif sort_port[1] in self.__proto__: + if sort_port[0] not in self.__proto__ and sort_port[1] in self.__proto__: proto = sort_port[1] else: - proto = None - return super()._decode_next_layer(dict_, proto, length, packet=packet) # type: ignore[arg-type] + proto = sort_port[0] + return super()._decode_next_layer(dict_, proto, length, packet=packet) diff --git a/pcapkit/protocols/transport/udp.py b/pcapkit/protocols/transport/udp.py index e654131f95..e442e98a88 100644 --- a/pcapkit/protocols/transport/udp.py +++ b/pcapkit/protocols/transport/udp.py @@ -25,6 +25,7 @@ import collections from typing import TYPE_CHECKING +from pcapkit.const.reg.apptype import TransportProtocol as Enum_TransportProtocol from pcapkit.const.reg.transtype import TransType as Enum_TransType from pcapkit.corekit.module import ModuleDescriptor from pcapkit.protocols.data.transport.udp import UDP as Data_UDP @@ -164,8 +165,8 @@ def make(self, """ return Schema_UDP( - srcport=srcport, - dstport=dstport, + srcport=self._make_port(srcport, Enum_TransportProtocol.udp), + dstport=self._make_port(dstport, Enum_TransportProtocol.udp), len=8 + len(payload), checksum=checksum, payload=payload, diff --git a/tests/protocols/internet/test_internet_unit.py b/tests/protocols/internet/test_internet_unit.py index 599f9690f8..9ee9eb280d 100644 --- a/tests/protocols/internet/test_internet_unit.py +++ b/tests/protocols/internet/test_internet_unit.py @@ -160,6 +160,59 @@ def __init__(self, file_: bytes, length: int, **kwargs: object) -> None: self.assertEqual(from_header.file, b'header-payload') self.assertEqual(from_header.length, len(b'header-payload')) + def test_dispatching_an_unregistered_transtype_leaves_the_registry_alone(self) -> None: + """Parsing must not write to the shared ``Internet.__proto__``. + + The registry is a :class:`collections.defaultdict` on a class attribute + shared by every :class:`~pcapkit.protocols.internet.internet.Internet` + subclass instance in the process, so a lookup that inserts the key it + missed turns ordinary parsing into registration. It grows the registry + once per distinct unknown protocol number seen, and it makes + :func:`~pcapkit.foundation.registry.protocols.register_transtype` warn + about an overwrite of something nobody ever registered. + + ICMP is the cheapest demonstration: it has a well-known number that the + default registry does *not* carry, so a 24-byte IPv4 datagram declaring + ``proto=1`` is enough to leak it. + + """ + from pcapkit.const.reg.transtype import TransType + from pcapkit.foundation.registry.protocols import register_transtype + from pcapkit.protocols.internet.internet import Internet + from pcapkit.protocols.internet.ipv4 import IPv4 + from pcapkit.protocols.misc.raw import Raw + + # version/IHL, ToS, total length 24, id, flags/offset, TTL, proto=ICMP, + # checksum, 127.0.0.1 -> 127.0.0.1, then four bytes of payload. + packet = bytes.fromhex('450000180000000040010000' '7f000001' '7f000001') + b'abcd' + self.assertEqual(len(packet), 24) + + registry = Internet.__dict__['__proto__'] + before = set(registry) + self.assertNotIn(TransType.ICMP, before) + + try: + ip = IPv4(packet) + + # The payload is still reachable as Raw, labelled with the protocol + # number it arrived with -- only the registry write is gone. + self.assertIsInstance(ip.payload, Raw) + self.assertEqual(bytes(ip.payload), b'abcd') + self.assertEqual(ip.payload.info.protocol, TransType.ICMP) + self.assertEqual(set(registry), before) + + # A second datagram must behave identically; a class-level leak from + # the first would show up here rather than above. + self.assertIsInstance(IPv4(packet).payload, Raw) + self.assertEqual(set(registry), before) + + # And the number is still registrable without a bogus warning. + with mock.patch('pcapkit.protocols.internet.internet.warn') as warn: + register_transtype(TransType.ICMP, Raw) + self.assertEqual(warn.call_count, 0) + finally: + registry.pop(TransType.ICMP, None) + if __name__ == '__main__': unittest.main() diff --git a/tests/protocols/link/test_link_unit.py b/tests/protocols/link/test_link_unit.py index cbbda097c4..0ae776cf73 100644 --- a/tests/protocols/link/test_link_unit.py +++ b/tests/protocols/link/test_link_unit.py @@ -271,6 +271,23 @@ def test_link_schema_callbacks_resolve_payload_and_auth_fields(self) -> None: registry[custom_type] = ModuleDescriptor('pcapkit.protocols.misc.raw', 'Raw') callback_payload(descriptor_field, {'type': custom_type}) self.assertIs(descriptor_field.protocol, Raw) + + # Resolving a registered descriptor memoises the imported class, + # which the callback used to resolve afresh on every frame. + self.assertIs(registry[custom_type], Raw) + + # An *unregistered* EtherType still resolves to Raw, and must not be + # recorded on the way: the registry is Link.__proto__, shared by + # every Link subclass in the process, so an inserting lookup here + # turns parsing into registration and makes Link.register afterwards + # report an overwrite of something nobody registered. + registry.pop(custom_type, None) + keys = set(registry) + missing_field = PayloadField() + callback_payload(missing_field, {'type': custom_type}) + self.assertIs(missing_field.protocol, Raw) + self.assertNotIn(custom_type, registry) + self.assertEqual(set(registry), keys) finally: if had_original: registry[custom_type] = original diff --git a/tests/protocols/test_protocol_base_unit.py b/tests/protocols/test_protocol_base_unit.py index 298414e929..a00e32aa04 100644 --- a/tests/protocols/test_protocol_base_unit.py +++ b/tests/protocols/test_protocol_base_unit.py @@ -335,6 +335,53 @@ def __index__(cls) -> int: no_stop._exproto = 'raw' self.assertFalse(no_stop._check_term_threshold()) + def test_lookup_next_layer_reads_the_fallback_without_recording_it(self) -> None: + """A missed lookup must not turn into a registration. + + ``__proto__`` is a :class:`collections.defaultdict` on a class + attribute, so ``__proto__[code]`` inserts every code it is handed. The + insertion is worth nothing -- the value is the fallback the factory + would have produced anyway -- and it costs a spurious "already + registered" warning on the next real + :meth:`~pcapkit.protocols.protocol.ProtocolBase.register` call. + + Resolving a :class:`~pcapkit.corekit.module.ModuleDescriptor` for a code + that *is* registered still writes back, since that is memoisation of an + import rather than a new entry. + + """ + DummyProtocol, _, _ = self._make_protocol_class() + from pcapkit.corekit.module import ModuleDescriptor + from pcapkit.protocols.misc.raw import Raw + + DummyProtocol.__proto__ = collections.defaultdict( + lambda: ModuleDescriptor('pcapkit.protocols.misc.raw', 'Raw'), + {7: ModuleDescriptor('pcapkit.protocols.misc.raw', 'Raw')}, + ) + + registry = DummyProtocol.__proto__ + + # A miss resolves to the declared fallback and leaves no trace. + self.assertIs(DummyProtocol._lookup_next_layer(registry, 99), Raw) + self.assertNotIn(99, registry) + self.assertEqual(set(registry), {7}) + + # A hit resolves the descriptor once and keeps the resolved class. + self.assertIs(DummyProtocol._lookup_next_layer(registry, 7), Raw) + self.assertIs(registry[7], Raw) + + # A class registered directly is returned as it is. + registry[8] = Raw + self.assertIs(DummyProtocol._lookup_next_layer(registry, 8), Raw) + + # ``analyze`` and ``_import_next_layer`` both go through the lookup, so + # neither of them records an unregistered code either. + self.assertIsInstance(DummyProtocol.analyze(99, b'body'), Raw) + proto = DummyProtocol(packet=b'abpayload') + proto._sigterm = False + self.assertIsInstance(proto._import_next_layer(99, 3), Raw) + self.assertEqual(set(DummyProtocol.__proto__), {7, 8}) + def test_make_payload_branches(self) -> None: DummyProtocol, DummyData, _ = self._make_protocol_class() from pcapkit.protocols.misc.null import NoPayload diff --git a/tests/protocols/transport/test_sctp_unit.py b/tests/protocols/transport/test_sctp_unit.py index cc73ce801e..e63bad1526 100644 --- a/tests/protocols/transport/test_sctp_unit.py +++ b/tests/protocols/transport/test_sctp_unit.py @@ -161,6 +161,29 @@ def test_common_header_wire_conformance(self) -> None: # A zeroed checksum is not the CRC32c of this packet. self.assertFalse(proto.checksum_valid) + def test_constructed_ports_carry_the_same_type_as_parsed_ones(self) -> None: + """A bare :obj:`int` port is resolved on construction, not left as it is. + + SCTP never tripped over this the way TCP and UDP did -- it keys its next + layer on the DATA chunk's PPID rather than on a port, so it never read + ``srcport.port`` off the schema -- but it did leave a constructed packet + holding an :obj:`int` where a parsed one holds an + :class:`~pcapkit.const.reg.apptype.AppType`. + + """ + from pcapkit.const.reg.apptype import AppType, TransportProtocol + from pcapkit.protocols.transport.sctp import SCTP + + proto = SCTP.__new__(SCTP) + schema = SCTP.make(proto, srcport=9899, dstport=38412, vtag=0x11223344) + + self.assertIsInstance(schema.srcport, AppType) + self.assertEqual(schema.srcport.port, 9899) + self.assertEqual(schema.srcport, + AppType.get(9899, proto=TransportProtocol.sctp)) + self.assertEqual(schema.dstport.port, 38412) + self.assertEqual(schema.pack()[:4], b'\x26\xab\x96\x0c') + def test_checksum_is_the_little_endian_crc32c_of_the_zeroed_packet(self) -> None: import struct diff --git a/tests/protocols/transport/test_tcp_runtime.py b/tests/protocols/transport/test_tcp_runtime.py index 48abc6026f..bd24878a58 100644 --- a/tests/protocols/transport/test_tcp_runtime.py +++ b/tests/protocols/transport/test_tcp_runtime.py @@ -73,6 +73,25 @@ def test_tcp_ipv6_frame_keeps_timestamp_only_options(self) -> None: self.assertEqual(options[2][1].echo, 2559889017) def test_tcp_unregistered_application_payload_falls_back_to_raw(self) -> None: + """An unregistered port is still recorded on the Raw payload. + + ``Data_Raw.protocol`` is "the original enumeration of this protocol", and + it is worth most precisely when the protocol is unknown -- it is then the + only record of what the payload claimed to be. This frame is SSH, which + :mod:`pcapkit` does not decode, so port 22 is what it has to say. It used + to say :obj:`None`, which made the frame indistinguishable from "TCP + payload on a port we cannot name at all", while the equivalent SCTP and + IPv4 cases both named theirs. + + The chain still ends in ``Raw`` rather than in the port's name: the ports + reach :meth:`Transport._decode_next_layer + ` as + plain :obj:`int` (``srcport.port``), not as + :class:`~pcapkit.const.reg.apptype.AppType` members, so + :class:`~pcapkit.protocols.misc.raw.Raw` has no name to label itself + with. + + """ extractor = self._extract('tcp.pcap') frame = extractor.frame[5] tcp = frame.payload.payload.payload @@ -82,7 +101,8 @@ def test_tcp_unregistered_application_payload_falls_back_to_raw(self) -> None: self.assertTrue(tcp.info.flags.ack) self.assertEqual(type(tcp.payload).__name__, 'Raw') self.assertEqual(tcp.payload.name, 'Unknown') - self.assertIsNone(tcp.payload.info.protocol) + self.assertEqual(tcp.info.dstport.port, 22) + self.assertEqual(tcp.payload.info.protocol, 22) self.assertIsNone(tcp.payload.info.error) def test_stream_sample_exposes_no_payload_ack_frame(self) -> None: diff --git a/tests/protocols/transport/test_tcp_udp_unit.py b/tests/protocols/transport/test_tcp_udp_unit.py index 34176c5fd1..0e4e010599 100644 --- a/tests/protocols/transport/test_tcp_udp_unit.py +++ b/tests/protocols/transport/test_tcp_udp_unit.py @@ -1137,6 +1137,43 @@ def test_transport_schema_helpers_cover_selectors_and_port_fields(self) -> None: self.assertEqual(processed.length, 3) self.assertEqual(processed.subtype, MPTCPOption.Reserved_for_Private_Use) + def test_construction_accepts_bare_integer_ports(self) -> None: + """``TCP(srcport=80, dstport=443)`` is the documented construction path. + + :meth:`TCP.make ` advertises + ``Enum_AppType | int``, but the schema field converted an :obj:`int` only + on its way out to :obj:`bytes`, so the schema attribute kept the + :obj:`int` -- and :meth:`TCP.read + `, which reads ``srcport.port`` + to pick the next layer, raised ``AttributeError: 'int' object has no + attribute 'port'``. UDP failed identically; SCTP did not, only because it + keys its next layer on the DATA chunk's PPID and never reads ``.port``. + + """ + from pcapkit.const.reg.apptype import AppType, TransportProtocol + from pcapkit.protocols.transport.tcp import TCP + from pcapkit.protocols.transport.udp import UDP + + tcp = TCP(srcport=80, dstport=443) + self.assertEqual(bytes(tcp)[:4], b'\x00\x50\x01\xbb') + self.assertIsInstance(tcp.info.srcport, AppType) + self.assertEqual(tcp.info.srcport.port, 80) + self.assertEqual(tcp.info.dstport.port, 443) + self.assertEqual(tcp.src.port, 80) + + udp = UDP(srcport=53, dstport=5353, payload=b'data') + self.assertEqual(bytes(udp)[:4], b'\x00\x35\x14\xe9') + self.assertIsInstance(udp.info.dstport, AppType) + self.assertEqual(udp.info.srcport.port, 53) + self.assertEqual(udp.info.dstport.port, 5353) + + # A port with no IANA service name resolves just the same, and an + # AppType passed in is kept as it is rather than round-tripped. + member = AppType.get(80, proto=TransportProtocol.tcp) + unnamed = TCP(srcport=53406, dstport=member) + self.assertEqual(unnamed.info.srcport.port, 53406) + self.assertIs(unnamed.info.dstport, member) + if __name__ == '__main__': unittest.main() diff --git a/tests/protocols/transport/test_transport_unit.py b/tests/protocols/transport/test_transport_unit.py index 928bce27e4..b832b9c006 100644 --- a/tests/protocols/transport/test_transport_unit.py +++ b/tests/protocols/transport/test_transport_unit.py @@ -21,7 +21,7 @@ def test_register_validates_protocols_and_abstract_base(self) -> None: from pcapkit.corekit.module import ModuleDescriptor from pcapkit.protocols.misc.raw import Raw from pcapkit.protocols.transport.transport import Transport - from pcapkit.utilities.exceptions import UnsupportedCall + from pcapkit.utilities.exceptions import RegistryError, UnsupportedCall class DummyTransport(Transport): __proto__ = collections.defaultdict(lambda: Raw, {80: Raw}) @@ -49,6 +49,10 @@ def __index__(cls) -> int: with self.assertRaises(UnsupportedCall): Transport.register(1, Raw) + # RegistryError is (BaseError, TypeError), so the bare TypeError this + # used to raise is still caught by anyone who was catching it. + with self.assertRaises(RegistryError): + DummyTransport.register(81, object) # type: ignore[arg-type] with self.assertRaises(TypeError): DummyTransport.register(81, object) # type: ignore[arg-type] @@ -317,6 +321,77 @@ def __index__(cls) -> int: self.assertEqual(result, 'decoded') decode.assert_called_once_with(data, 2000, 12, packet=None) + def test_decode_next_layer_forwards_the_lower_port_when_neither_is_registered(self) -> None: + """An unregistered port pair still names the port it could not place. + + The port reaches + :meth:`ProtocolBase._import_next_layer + ` as + ``alias``, which is what :class:`~pcapkit.protocols.misc.raw.Raw` records + as ``Data_Raw.protocol``. This used to pass :obj:`None`, so TCP and UDP + anonymised every payload they could not dispatch -- unlike SCTP and IPv4, + which both keep the identifier the packet arrived with. The lower port is + the one carried, matching the lookup's own primary key. + + """ + from pcapkit.protocols.protocol import ProtocolBase + from pcapkit.protocols.transport.transport import Transport + + class DummyTransport(Transport): + __proto__ = collections.defaultdict(lambda: None, {80: object}) + + @property + def name(self) -> str: + return 'Dummy Transport' + + @property + def length(self) -> int: + return 0 + + def read(self, length: int | None = None, **kwargs: object) -> object: + raise NotImplementedError + + def make(self, **kwargs: object) -> object: + raise NotImplementedError + + @classmethod + def __index__(cls) -> int: + return 0 + + transport = object.__new__(DummyTransport) + data = object() + + with mock.patch.object(ProtocolBase, '_decode_next_layer', return_value='decoded') as decode: + result = DummyTransport._decode_next_layer(transport, data, (53406, 22), 12) + + self.assertEqual(result, 'decoded') + decode.assert_called_once_with(data, 22, 12, packet=None) + + # And the lookup left nothing behind, so the ports stay registrable. + self.assertEqual(set(DummyTransport.__proto__), {80}) + + def test_make_port_resolves_an_integer_and_passes_an_apptype_through(self) -> None: + from pcapkit.const.reg.apptype import AppType, TransportProtocol + from pcapkit.protocols.transport.transport import Transport + + resolved = Transport._make_port(80, TransportProtocol.tcp) + self.assertIsInstance(resolved, AppType) + self.assertEqual(resolved.port, 80) + + # The transport protocol has to reach the lookup rather than being + # defaulted away: port 1 is tcpmux over TCP and unassigned over SCTP, so + # the two resolve to different members of the same number. + over_tcp = Transport._make_port(1, TransportProtocol.tcp) + over_sctp = Transport._make_port(1, TransportProtocol.sctp) + self.assertEqual(over_tcp.svc, 'tcpmux') + self.assertEqual(over_sctp.svc, 'unknown') + self.assertIsNot(over_tcp, over_sctp) + self.assertEqual(over_tcp.port, over_sctp.port) + + # An AppType is returned unchanged -- no round trip through its number, + # which would lose the protocol it was resolved for. + self.assertIs(Transport._make_port(over_sctp, TransportProtocol.tcp), over_sctp) + if __name__ == '__main__': unittest.main() diff --git a/tests/protocols/transport/test_udp_runtime.py b/tests/protocols/transport/test_udp_runtime.py index a8a429a719..3e427c37da 100644 --- a/tests/protocols/transport/test_udp_runtime.py +++ b/tests/protocols/transport/test_udp_runtime.py @@ -34,6 +34,10 @@ def test_ipv4_udp_frame_exposes_length_checksum_and_raw_payload(self) -> None: self.assertEqual(udp.info.checksum.hex(), 'ff0b') self.assertEqual(type(udp.payload).__name__, 'Raw') + # Neither port is registered, so the lower of the two is what labels the + # payload -- the same treatment SCTP gives an unregistered PPID. + self.assertEqual(udp.payload.info.protocol, 12345) + def test_ipv6_udp_mdns_frame_exposes_multicast_ports_and_checksum(self) -> None: extractor = self._extract('stream.pcap') frame = extractor.frame[1] @@ -46,6 +50,7 @@ def test_ipv6_udp_mdns_frame_exposes_multicast_ports_and_checksum(self) -> None: self.assertEqual(udp.info.len, 145) self.assertEqual(udp.info.checksum.hex(), '9cb1') self.assertEqual(type(udp.payload).__name__, 'Raw') + self.assertEqual(udp.payload.info.protocol, 5353) if __name__ == '__main__':