Commit 3cc40f54 authored by Szilárd Pfeiffer's avatar Szilárd Pfeiffer
Browse files

Merge branch '93-ike-notify-messages'

Closes: #93
parents 2bea1ae7 101c98d8
Loading
Loading
Loading
Loading
Loading
+133 −0
Original line number Diff line number Diff line
@@ -954,6 +954,48 @@ class Ikev2PayloadNotifyAuthenticationFailed(Ikev2PayloadNotifyNoData):
        return Ikev2NotifyType.AUTHENTICATION_FAILED


class Ikev2NotifyPayloadUseTransportMode(Ikev2PayloadNotifyNoData):
    @classmethod
    def _get_message_type(cls):
        return Ikev2NotifyType.USE_TRANSPORT_MODE


class Ikev2NotifyPayloadHttpCertLookupSupported(Ikev2PayloadNotifyNoData):
    @classmethod
    def _get_message_type(cls):
        return Ikev2NotifyType.HTTP_CERT_LOOKUP_SUPPORTED


class Ikev2NotifyPayloadIkev2FragmentationSupported(Ikev2PayloadNotifyNoData):
    @classmethod
    def _get_message_type(cls):
        return Ikev2NotifyType.IKEV2_FRAGMENTATION_SUPPORTED


class Ikev2NotifyPayloadIntermediateExchangeSupported(Ikev2PayloadNotifyNoData):
    @classmethod
    def _get_message_type(cls):
        return Ikev2NotifyType.INTERMEDIATE_EXCHANGE_SUPPORTED


class Ikev2NotifyPayloadUsePpk(Ikev2PayloadNotifyNoData):
    @classmethod
    def _get_message_type(cls):
        return Ikev2NotifyType.USE_PPK


class Ikev2NotifyPayloadRedirectSupported(Ikev2PayloadNotifyNoData):
    @classmethod
    def _get_message_type(cls):
        return Ikev2NotifyType.REDIRECT_SUPPORTED


class Ikev2NotifyPayloadChildlessIkev2Supported(Ikev2PayloadNotifyNoData):
    @classmethod
    def _get_message_type(cls):
        return Ikev2NotifyType.CHILDLESS_IKEV2_SUPPORTED


@attr.s
class Ikev2PayloadNotifyUnparsed(Ikev2PayloadNotifyBase):
    data: typing.Union[bytes, bytearray] = attr.ib(validator=attr.validators.instance_of((bytes, bytearray)))
@@ -1028,6 +1070,86 @@ class Ikev2NotifyPayloadCookie(Ikev2PayloadNotifyParsedBase):
        composer.compose_raw(self.cookie)


@attr.s
class Ikev2NotifyPayloadSetWindowSize(Ikev2PayloadNotifyParsedBase):
    """Set window size payload notification data parser."""
    window_size: int = attr.ib(validator=[
        attr.validators.instance_of(int),
        attr.validators.in_(range(0, 2 ** 32)),
    ])

    @classmethod
    def _get_message_type(cls):
        return Ikev2NotifyType.SET_WINDOW_SIZE

    @classmethod
    def _parse_data(cls, parser, notification_data_length):
        if notification_data_length != 4:
            raise InvalidValue(notification_data_length, cls, 'notification_data_length')
        parser.parse_numeric('window_size', 4)

    def _compose_data(self, composer):
        composer.compose_numeric(self.window_size, 4)


@attr.s
class Ikev2NotifyPayloadNatDetectionBase(Ikev2PayloadNotifyParsedBase):
    hash_data: typing.Union[bytes, bytearray] = attr.ib(validator=attr.validators.instance_of((bytes, bytearray)))

    @classmethod
    @abc.abstractmethod
    def _get_message_type(cls):
        raise NotImplementedError()

    @classmethod
    def _parse_data(cls, parser, notification_data_length):
        parser.parse_raw('hash_data', notification_data_length)

    def _compose_data(self, composer):
        composer.compose_raw(self.hash_data)


@attr.s
class Ikev2NotifyPayloadNatDetectionSourceIp(Ikev2NotifyPayloadNatDetectionBase):
    """NAT detection source IP payload notification data parser."""

    @classmethod
    def _get_message_type(cls):
        return Ikev2NotifyType.NAT_DETECTION_SOURCE_IP


@attr.s
class Ikev2NotifyPayloadNatDetectionDestinationIp(Ikev2NotifyPayloadNatDetectionBase):
    """NAT detection destination IP payload notification data parser."""

    @classmethod
    def _get_message_type(cls):
        return Ikev2NotifyType.NAT_DETECTION_DESTINATION_IP


@attr.s
class Ikev2NotifyPayloadSignatureHashAlgorithms(Ikev2PayloadNotifyParsedBase):
    """Signature hash algorithms notification (RFC 7427 §4)."""
    hash_algorithms: tuple[int, ...] = attr.ib(
        converter=tuple,
        validator=attr.validators.deep_iterable(attr.validators.instance_of(int)),
    )

    @classmethod
    def _get_message_type(cls):
        return Ikev2NotifyType.SIGNATURE_HASH_ALGORITHMS

    @classmethod
    def _parse_data(cls, parser, notification_data_length):
        if notification_data_length % 2 != 0:
            raise InvalidValue(notification_data_length, cls, 'notification_data_length')
        parser.parse_numeric_array('hash_algorithms', notification_data_length // 2, 2)

    def _compose_data(self, composer):
        for hash_id in self.hash_algorithms:
            composer.compose_numeric(hash_id, 2)


class Ikev2NotifyPayloadVariantBase(VariantParsable):
    @classmethod
    @abc.abstractmethod
@@ -1053,6 +1175,17 @@ class Ikev2NotifyPayloadVariantResponder(Ikev2NotifyPayloadVariantBase):
        return collections.OrderedDict([
            (Ikev2NotifyType.COOKIE, [Ikev2NotifyPayloadCookie, ]),
            (Ikev2NotifyType.INVALID_KE_PAYLOAD, [Ikev2NotifyPayloadInvalidKe, ]),
            (Ikev2NotifyType.SET_WINDOW_SIZE, [Ikev2NotifyPayloadSetWindowSize, ]),
            (Ikev2NotifyType.NAT_DETECTION_SOURCE_IP, [Ikev2NotifyPayloadNatDetectionSourceIp, ]),
            (Ikev2NotifyType.NAT_DETECTION_DESTINATION_IP, [Ikev2NotifyPayloadNatDetectionDestinationIp, ]),
            (Ikev2NotifyType.USE_TRANSPORT_MODE, [Ikev2NotifyPayloadUseTransportMode, ]),
            (Ikev2NotifyType.HTTP_CERT_LOOKUP_SUPPORTED, [Ikev2NotifyPayloadHttpCertLookupSupported, ]),
            (Ikev2NotifyType.SIGNATURE_HASH_ALGORITHMS, [Ikev2NotifyPayloadSignatureHashAlgorithms, ]),
            (Ikev2NotifyType.IKEV2_FRAGMENTATION_SUPPORTED, [Ikev2NotifyPayloadIkev2FragmentationSupported, ]),
            (Ikev2NotifyType.INTERMEDIATE_EXCHANGE_SUPPORTED, [Ikev2NotifyPayloadIntermediateExchangeSupported, ]),
            (Ikev2NotifyType.USE_PPK, [Ikev2NotifyPayloadUsePpk, ]),
            (Ikev2NotifyType.REDIRECT_SUPPORTED, [Ikev2NotifyPayloadRedirectSupported, ]),
            (Ikev2NotifyType.CHILDLESS_IKEV2_SUPPORTED, [Ikev2NotifyPayloadChildlessIkev2Supported, ]),
        ])


+18 −4
Original line number Diff line number Diff line
@@ -93,14 +93,28 @@ class IsakmpMessage(ParsableBase):
        )
    )

    def _collect_payloads_by_type(
        self, payload_type: typing.Union[Ikev1PayloadType, Ikev2PayloadType]
    ) -> list[typing.Union[Ikev1PayloadBase, Ikev2PayloadBase]]:
        return [
            payload for payload in self.payloads
            if payload.get_payload_type() == payload_type
        ]

    def get_payloads_by_type(
        self, payload_type: typing.Union[Ikev1PayloadType, Ikev2PayloadType]
    ) -> list[typing.Union[Ikev1PayloadBase, Ikev2PayloadBase]]:
        return self._collect_payloads_by_type(payload_type)

    def get_payload_by_type(
        self, payload_type: typing.Union[Ikev1PayloadType, Ikev2PayloadType]
    ) -> typing.Union[Ikev1PayloadBase, Ikev2PayloadBase]:
        for payload in self.payloads:
            if payload.get_payload_type() == payload_type:
                return payload

        payloads = self._collect_payloads_by_type(payload_type)
        if not payloads:
            raise KeyError(payload_type)
        if len(payloads) > 1:
            raise IndexError(payload_type)
        return payloads[0]

    @classmethod
    def _parse(cls, parsable):
Compare 5367d32f to 714e10d5
Original line number Diff line number Diff line
Subproject commit 5367d32fc976f5cd9bbedccdcf976905bae9a214
Subproject commit 714e10d550e2b3fe6b9bdff591cbdcb69d5f985a
+66 −1
Original line number Diff line number Diff line
# SPDX-License-Identifier: MPL-2.0

import collections
import unittest

from cryptodatahub.ike.algorithm import (
    Ikev1PayloadType,
    Ikev2PayloadType,
    Ikev2NotifyType,
    Ikev2PayloadType,
    Ikev2ProtocolId,
    Ikev2TransformType,
    Ikev2PseudorandomFunction,
@@ -11,6 +14,7 @@ from cryptodatahub.ike.algorithm import (
from cryptoparser.common.parse import ComposerBinary
from cryptoparser.ike.ikev1 import Ikev1PayloadBase, Ikev1PayloadDoiProtocolSpiBase
from cryptoparser.ike.ikev2 import (
    Ikev2NotifyPayloadNatDetectionBase,
    Ikev2PayloadBase,
    Ikev2PayloadNotifyBase,
    Ikev2PayloadNotifyNoData,
@@ -265,3 +269,64 @@ class TransformTest(Transform):

    def compose(self):
        return self.compose_header(transform_length=0).composed_bytes


class Ikev2NotifyPayloadNatDetectionBaseTest(unittest.TestCase):
    _HASH_DATA = b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f\x10\x11\x12\x13'
    _NOTIFY_TYPE: Ikev2NotifyType
    _PAYLOAD_CLASS: type[Ikev2NotifyPayloadNatDetectionBase]
    _NOTIFY_TYPE_BYTES: bytes

    def setUp(self):
        payload_dict = collections.OrderedDict([
            ('next_payload', b'\x00'),
            ('flags', b'\x00'),
            ('payload_length', b'\x00\x1c'),
            ('protocol_id', b'\x01'),
            ('spi_size', b'\x00'),
            ('notify_type', self._NOTIFY_TYPE_BYTES),
            ('hash_data', self._HASH_DATA),
        ])
        self.payload_bytes = b''.join(payload_dict.values())

        self.nat_payload = self._PAYLOAD_CLASS(
            flags=set(),
            protocol_id=Ikev2ProtocolId.IKE,
            type=self._NOTIFY_TYPE,
            spi=b'',
            hash_data=self._HASH_DATA
        )
        self.nat_payload.next_payload = Ikev2PayloadType.NONE

    def test_get_message_type(self):
        # pylint: disable=protected-access
        self.assertEqual(self._PAYLOAD_CLASS._get_message_type(), self._NOTIFY_TYPE)

    def test_parse(self):
        parsed_payload = self._PAYLOAD_CLASS.parse_exact_size(self.payload_bytes)
        self.assertEqual(parsed_payload.hash_data, self._HASH_DATA)  # pylint: disable=no-member
        self.assertEqual(parsed_payload.type, self._NOTIFY_TYPE)

    def test_compose(self):
        self.assertEqual(self.nat_payload.compose(), self.payload_bytes)

    def test_hash_data_storage(self):
        self.assertEqual(self.nat_payload.hash_data, self._HASH_DATA)  # pylint: disable=no-member

        different_hash = b'\xff\xfe\xfd\xfc\xfb\xfa'
        payload_2 = self._PAYLOAD_CLASS(
            flags=set(),
            protocol_id=Ikev2ProtocolId.IKE,
            type=self._NOTIFY_TYPE,
            spi=b'',
            hash_data=different_hash
        )
        self.assertEqual(payload_2.hash_data, different_hash)  # pylint: disable=no-member

    def test_round_trip_hash_preservation(self):
        composed_bytes = self.nat_payload.compose()
        parsed_payload = self._PAYLOAD_CLASS.parse_exact_size(composed_bytes)

        self.assertEqual(parsed_payload.hash_data, self.nat_payload.hash_data)  # pylint: disable=no-member
        self.assertEqual(parsed_payload.type, self.nat_payload.type)
        self.assertEqual(parsed_payload.spi, self.nat_payload.spi)
+318 −26

File changed.

Preview size limit exceeded, changes collapsed.

Loading