Commit 413be2ea authored by Szilárd Pfeiffer's avatar Szilárd Pfeiffer
Browse files

Merge remote-tracking branch...

Merge remote-tracking branch 'origin/184-decrease-the-number-of-live-test-to-improve-test-run-stability'

Closes: #184
parents d0167c1a 6327e6af
Loading
Loading
Loading
Loading
+154 −19
Original line number Diff line number Diff line
@@ -3,6 +3,9 @@
import abc
import attr

from cryptodatahub.common.key import PublicKey, PublicKeyParamsRsa
from cryptodatahub.common.parameter import DHParamWellKnown

from cryptodatahub.ssh.algorithm import (
    SshCompressionAlgorithm,
    SshEncryptionAlgorithm,
@@ -12,17 +15,48 @@ from cryptodatahub.ssh.algorithm import (
)

from cryptoparser.common.classes import LanguageTag

from cryptoparser.ssh.record import SshRecordInit
from cryptoparser.ssh.subprotocol import SshMessageCode, SshReasonCode, SshKeyExchangeInit, SshDisconnectMessage
from cryptoparser.common.exception import NotEnoughData

from cryptoparser.ssh.key import SshHostKeyRSA, SshPublicKeyBase
from cryptoparser.ssh.record import SshRecordInit, SshRecordKexDH, SshRecordKexDHGroup
from cryptoparser.ssh.subprotocol import (
    SshDHGroupExchangeGroup,
    SshDHGroupExchangeReply,
    SshDHKeyExchangeReply,
    SshDisconnectMessage,
    SshKeyExchangeInit,
    SshMessageCode,
    SshNewKeys,
    SshReasonCode,
)
from cryptoparser.ssh.version import SshProtocolVersion, SshVersion

from cryptolyzer.common.application import L7ServerBase, L7ServerHandshakeBase, L7ServerConfigurationBase
from cryptolyzer.common.dhparam import get_dh_ephemeral_key_forged, int_to_bytes

from cryptolyzer.ssh.client import SshProtocolMessageDefault
from cryptolyzer.ssh.transfer import SshHandshakeBase


# Public RSA host-key material only: a fixed ~2048-bit odd modulus (public data, never a private key) paired
# with the standard public exponent. It lets the offline mock server present a parseable host key during key
# exchange; the client verifies no signature, so no private key is ever needed.
DEFAULT_SSH_SERVER_HOST_PUBLIC_KEY = SshHostKeyRSA(
    SshHostKeyAlgorithm.SSH_RSA,
    PublicKey.from_params(PublicKeyParamsRsa(
        modulus=int('c0ffee' + 'a5' * 252 + '01', 16),
        public_exponent=65537,
    )),
)
DEFAULT_SSH_SERVER_DH_GROUP_EXCHANGE_GROUPS = (
    DHParamWellKnown.RFC3526_2048_BIT_MODP_GROUP,
    DHParamWellKnown.RFC3526_3072_BIT_MODP_GROUP,
    DHParamWellKnown.RFC3526_4096_BIT_MODP_GROUP,
    DHParamWellKnown.RFC3526_6144_BIT_MODP_GROUP,
    DHParamWellKnown.RFC3526_8192_BIT_MODP_GROUP,
)


@attr.s
class SshServerConfiguration(L7ServerConfigurationBase):  # pylint: disable=too-many-instance-attributes
    protocol_version = attr.ib(
@@ -74,6 +108,23 @@ class SshServerConfiguration(L7ServerConfigurationBase): # pylint: disable=too-
        default=None,
        metadata={'description': 'Maximum number of algorithms allowed per list from remote side (None = unlimited)'}
    )
    key_exchange_reply = attr.ib(
        validator=attr.validators.instance_of(bool),
        default=False,
        metadata={'description': 'Whether to complete the key exchange instead of disconnecting after KEXINIT'}
    )
    host_public_key = attr.ib(
        validator=attr.validators.instance_of(SshPublicKeyBase),
        default=DEFAULT_SSH_SERVER_HOST_PUBLIC_KEY
    )
    dh_group_exchange_groups = attr.ib(
        validator=attr.validators.deep_iterable(member_validator=attr.validators.in_(DHParamWellKnown)),
        default=DEFAULT_SSH_SERVER_DH_GROUP_EXCHANGE_GROUPS
    )
    dh_group_exchange_bounds_tolerated = attr.ib(
        validator=attr.validators.instance_of(bool),
        default=True
    )


@attr.s
@@ -89,6 +140,9 @@ class L7ServerSshBase(L7ServerBase):
        raise NotImplementedError()

    def _get_handshake_class(self):
        if self.configuration is not None and self.configuration.key_exchange_reply:
            return SshServerHandshakeKeyExchange

        return SshServerHandshake

    def _do_handshake(self, last_handshake_message_type):
@@ -140,12 +194,11 @@ class SshServerHandshake(L7ServerHandshakeBase, SshHandshakeBase):
    def _parse_message(self, record):
        return record.packet

    def _process_handshake_message(self, message, last_handshake_message_type):
        self._last_processed_message_type = message.get_message_code()
        self.client_messages[self._last_processed_message_type] = message
    def _disconnect_if_too_many_algorithms(self, message):
        if (self.configuration.max_remote_algorithm_count is None or
                message.get_message_code() != SshMessageCode.KEXINIT):
            return

        if self.configuration.max_remote_algorithm_count is not None:
            if self._last_processed_message_type == SshMessageCode.KEXINIT:
        for attribute in message._get_cipher_attributes():  # pylint: disable=protected-access
            if 'algorithms' not in attribute.name:
                continue
@@ -161,6 +214,12 @@ class SshServerHandshake(L7ServerHandshakeBase, SshHandshakeBase):
                )
                raise StopIteration()

    def _process_handshake_message(self, message, last_handshake_message_type):
        self._last_processed_message_type = message.get_message_code()
        self.client_messages[self._last_processed_message_type] = message

        self._disconnect_if_too_many_algorithms(message)

        if self._last_processed_message_type == last_handshake_message_type:
            self._send_disconnect(SshReasonCode.HOST_NOT_ALLOWED_TO_CONNECT, 'not allowed to connect')
            raise StopIteration()
@@ -184,6 +243,82 @@ class SshServerHandshake(L7ServerHandshakeBase, SshHandshakeBase):
        self.l7_transfer.send(SshRecordInit(SshDisconnectMessage(**kwargs)).compose())


@attr.s
class SshServerHandshakeKeyExchange(SshServerHandshake):
    _RECORD_CLASS_BY_MESSAGE_CODE = {
        SshMessageCode.DH_KEX_INIT: SshRecordKexDH,
        SshMessageCode.DH_GEX_REQUEST: SshRecordKexDHGroup,
        SshMessageCode.DH_GEX_INIT: SshRecordKexDHGroup,
    }
    _FORGED_SIGNATURE = b'\x00' * 4

    _group_exchange_group = attr.ib(init=False, default=None)

    def _parse_record(self):
        if len(self.l7_transfer.buffer) < SshRecordInit.HEADER_SIZE:
            raise NotEnoughData(SshRecordInit.HEADER_SIZE - len(self.l7_transfer.buffer))

        message_code = self.l7_transfer.buffer[5]
        record_class = self._RECORD_CLASS_BY_MESSAGE_CODE.get(message_code, SshRecordInit)
        record = record_class.parse_exact_size(self.l7_transfer.buffer)
        is_handshake = record.packet.get_message_code() != SshMessageCode.DISCONNECT

        return record, len(self.l7_transfer.buffer), is_handshake

    def _process_handshake_message(self, message, last_handshake_message_type):
        message_code = message.get_message_code()
        self._last_processed_message_type = message_code
        self.client_messages[message_code] = message

        if message_code == SshMessageCode.DH_KEX_INIT:
            self._send_key_exchange_reply(SshRecordKexDH, SshDHKeyExchangeReply, self._sorted_groups()[0])
        elif message_code == SshMessageCode.DH_GEX_REQUEST:
            self._send_group_exchange_group(message)
        elif message_code == SshMessageCode.DH_GEX_INIT:
            if self._group_exchange_group is None:
                raise StopIteration()
            self._send_key_exchange_reply(SshRecordKexDHGroup, SshDHGroupExchangeReply, self._group_exchange_group)
        else:
            self._disconnect_if_too_many_algorithms(message)

    def _sorted_groups(self):
        return sorted(self.configuration.dh_group_exchange_groups, key=lambda well_known: well_known.value.key_size)

    def _get_group_exchange_group(self, message):
        sorted_groups = self._sorted_groups()
        if not self.configuration.dh_group_exchange_bounds_tolerated:
            return sorted_groups[0]

        for well_known in sorted_groups:
            if message.gex_min <= well_known.value.key_size <= message.gex_max:
                return well_known

        return None

    def _send_group_exchange_group(self, message):
        well_known = self._get_group_exchange_group(message)
        if well_known is None:
            raise StopIteration()

        self._group_exchange_group = well_known
        parameter_numbers = well_known.value.parameter_numbers
        self.l7_transfer.send(SshRecordKexDHGroup(SshDHGroupExchangeGroup(
            int_to_bytes(parameter_numbers.p, well_known.value.key_size // 8),
            int_to_bytes(parameter_numbers.g, (parameter_numbers.g.bit_length() + 7) // 8),
        )).compose())

    def _send_key_exchange_reply(self, record_class, reply_class, well_known):
        ephemeral_public_key = int_to_bytes(
            get_dh_ephemeral_key_forged(well_known.value.parameter_numbers.p), well_known.value.key_size // 8
        ).lstrip(b'\x00')
        reply = reply_class(
            host_public_key=self.configuration.host_public_key,
            ephemeral_public_key=ephemeral_public_key,
            signature=self._FORGED_SIGNATURE,
        )
        self.l7_transfer.send(record_class(reply).compose() + record_class(SshNewKeys()).compose())


class L7ServerSsh(L7ServerSshBase):
    def __attrs_post_init__(self):
        if self.configuration is None:
+29 −0
Original line number Diff line number Diff line
@@ -9,6 +9,7 @@ from cryptodatahub.common.algorithm import BlockCipher, KeyExchange
from cryptodatahub.common.exception import InvalidValue
from cryptodatahub.common.parameter import DHParameterNumbers, DHParamWellKnown
from cryptodatahub.tls.algorithm import (
    TlsECPointFormat,
    TlsNamedCurve,
    TlsNextProtocolName,
    TlsProtocolName,
@@ -21,7 +22,9 @@ from cryptoparser.common.parse import ComposerBinary
from cryptoparser.common.x509 import SignedCertificateTimestampList

from cryptoparser.tls.extension import (
    TlsCertificateStatusType,
    TlsExtensionApplicationLayerProtocolNegotiation,
    TlsExtensionECPointFormats,
    TlsExtensionEncryptThenMAC,
    TlsExtensionExtendedMasterSecret,
    TlsExtensionKeyShareClientHelloRetry,
@@ -79,6 +82,7 @@ from cryptoparser.tls.subprotocol import (
    TlsDistinguishedName,
    TlsECCurveType,
    TlsHandshakeCertificateRequest,
    TlsHandshakeCertificateStatus,
    TlsHandshakeServerCertificate,
    TlsHandshakeServerHelloDone,
    TlsHandshakeServerHello,
@@ -131,6 +135,10 @@ class TlsServerConfiguration(L7ServerConfigurationBase): # pylint: disable=too-
            attr.validators.deep_iterable(attr.validators.instance_of(bytes))
        )
    )
    certificate_status = attr.ib(
        default=None,
        validator=attr.validators.optional(attr.validators.instance_of(bytes))
    )
    dh_param = attr.ib(
        default=None,
        validator=attr.validators.optional(
@@ -162,6 +170,13 @@ class TlsServerConfiguration(L7ServerConfigurationBase): # pylint: disable=too-
        )
    )
    signed_certificate_timestamps_supported = attr.ib(default=False, validator=attr.validators.instance_of(bool))
    ec_point_formats = attr.ib(
        default=None,
        validator=attr.validators.optional(
            attr.validators.deep_iterable(attr.validators.in_(TlsECPointFormat))
        )
    )
    fallback_scsv_supported = attr.ib(default=False, validator=attr.validators.instance_of(bool))

    def __attrs_post_init__(self):
        if self.min_protocol_version > self.max_protocol_version:
@@ -198,6 +213,9 @@ class TlsServerConfiguration(L7ServerConfigurationBase): # pylint: disable=too-
        if self.certificate_authorities is not None and not self.certificates:
            raise ValueError('certificate_authorities is set but no certificate is configured')

        if self.certificate_status is not None and not self.certificates:
            raise ValueError('certificate_status is set but no certificate is configured')


@attr.s
class L7ServerTlsBase(L7ServerBase):
@@ -395,6 +413,8 @@ class TlsServerHandshake(TlsServer):
            ))
        if self.configuration.signed_certificate_timestamps_supported:
            extensions.append(TlsExtensionSignedCertificateTimestampServer(SignedCertificateTimestampList([])))
        if self.configuration.ec_point_formats is not None:
            extensions.append(TlsExtensionECPointFormats(self.configuration.ec_point_formats))

        return extensions

@@ -465,6 +485,11 @@ class TlsServerHandshake(TlsServer):
                TlsRecord(TlsHandshakeServerCertificate(certificate_chain).compose()).compose()
            )

            if self.configuration.certificate_status is not None:
                self.l7_transfer.send(TlsRecord(TlsHandshakeCertificateStatus(
                    TlsCertificateStatusType.OCSP, self.configuration.certificate_status
                ).compose()).compose())

        if cipher_suite.value.key_exchange in (KeyExchange.DHE, KeyExchange.ADH):
            if self.configuration.dh_param is not None:
                param_bytes = self._compose_dh_param_bytes(self.configuration.dh_param)
@@ -511,6 +536,10 @@ class TlsServerHandshake(TlsServer):

        if message.get_handshake_type() == TlsHandshakeType.CLIENT_HELLO:
            protocol_version = self._check_protocol_version(message)
            if (self.configuration.fallback_scsv_supported and message.fallback_scsv and
                    protocol_version < self.configuration.max_protocol_version):
                self._handle_error(TlsAlertLevel.FATAL, TlsAlertDescription.INAPPROPRIATE_FALLBACK)
                raise StopIteration()
            server_hello = self._prepare_server_hello(message, protocol_version)
            self.l7_transfer.send(TlsRecord(server_hello.compose()).compose())

+12 −0
Original line number Diff line number Diff line
@@ -4,9 +4,12 @@
import unittest
from unittest import mock

import urllib3

from cryptolyzer.common.analyzer import ProtocolHandlerBase
from cryptolyzer.common.transfer import L4TransferSocketParams

from cryptolyzer.dnsrec.dnssec import AnalyzerDnsSec
from cryptolyzer.ike.analyzer import ProtocolHandlerIKEVersionIndependent
from cryptolyzer.ike.versions import AnalyzerVersions

@@ -19,6 +22,15 @@ class TestAnalyzer(unittest.TestCase):
    def test_protocol(self):
        self.assertEqual(ProtocolHandlerIKEVersionIndependent.get_protocol(), 'ike')

    def test_l7_client_from_params_non_ip_fragment(self):
        handler_class = ProtocolHandlerBase.from_protocol('tls')
        uri = urllib3.util.parse_url('tls://localhost:443#not-an-ip')
        l7_client = handler_class._l7_client_from_params(uri, L4TransferSocketParams())
        self.assertEqual(l7_client.address, 'localhost')

    def test_dns_record_default_scheme(self):
        self.assertEqual(AnalyzerDnsSec.get_default_scheme(), 'dns')


class TestAnalyzerThrottle(unittest.TestCase):
    """Tests for throttle functionality in AnalyzerBase._before_probe."""
+43 −33
Original line number Diff line number Diff line
@@ -10,7 +10,6 @@ from test.common.classes import (
    TestThreadedServerHttpProxy,
    TestHTTPProxyRequestHandler,
)
from test.common.markers import live_dns, live_server

import urllib3

@@ -32,44 +31,51 @@ class TestL4ClientTCP(unittest.TestCase):

        return l4_client, result

    @live_dns
    def test_receive_uninitialized(self):
        l4_client = L4ClientTCP('smtp.gmail.com', 587)
        l4_client = L4ClientTCP('localhost', 0)
        with self.assertRaises(NotEnoughData) as context_manager:
            l4_client.receive(1)
        self.assertEqual(context_manager.exception.bytes_needed, 1)

    @live_server
    def test_error_on_close(self):
        address = 'smtp.gmail.com'
        l4_client, _ = self._create_client_and_receive_text(address, 587, 4 + len(address), to_be_closed=False)
        test_http_server = TestThreadedServerHttp('127.0.0.1', 0)
        test_http_server.init_connection()
        test_http_server.start()

        try:
            address = '127.0.0.1'
            port = test_http_server.bind_port

            l4_client = L4ClientTCP(address, port)
            l4_client.init_connection()
            sock = l4_client._socket  # pylint: disable=protected-access
            with mock.patch.object(socket.socket, 'close', side_effect=socket.error):
                l4_client.close()
            sock.close()

        l4_client, _ = self._create_client_and_receive_text(address, 587, 4 + len(address), to_be_closed=False)
            l4_client = L4ClientTCP(address, port)
            l4_client.init_connection()
            sock = l4_client._socket  # pylint: disable=protected-access
            with mock.patch.object(socket.socket, 'close', side_effect=NotImplementedError('not a timeout error')):
                with self.assertRaises(NotImplementedError) as context_manager:
                    l4_client.close()
                self.assertEqual(context_manager.exception.args, ('not a timeout error', ))
            sock.close()
        finally:
            test_http_server.kill()

    @live_dns
    def test_error_connection_refused(self):
        with mock.patch.object(socket, 'create_connection', side_effect=ConnectionRefusedError), \
                self.assertRaises(NetworkError) as context_manager:
            l4_client = L4ClientTCP('badssl.com', 443)
            l4_client = L4ClientTCP('localhost', 0)
            l4_client.init_connection()
        l4_client.close()
        self.assertEqual(context_manager.exception.error, NetworkErrorType.NO_CONNECTION)

    @live_dns
    def test_error_unhandled_exception_rethrown(self):
        with mock.patch.object(socket, 'create_connection', side_effect=NotImplementedError), \
                self.assertRaises(NotImplementedError):
            self._create_client_and_receive_text('badssl.com', 443, 1)
            self._create_client_and_receive_text('localhost', 0, 1)

    @mock.patch.object(TestHTTPProxyRequestHandler, '_get_response_code', return_value=500)
    def test_error_proxy_result_not_http_ok(self, _):
@@ -94,20 +100,26 @@ class TestL4ClientTCP(unittest.TestCase):
        test_http_proxy_server.kill()
        test_http_server.kill()

    @live_server
    def test_receive_until(self):
        address = 'smtp.gmail.com'
        test_http_server = TestThreadedServerHttp('127.0.0.1', 0)
        test_http_server.init_connection()
        test_http_server.start()

        l4_client = L4ClientTCP(address, 587)
        try:
            l4_client = L4ClientTCP('127.0.0.1', test_http_server.bind_port)
            l4_client.init_connection()
            l4_client.send(b'GET / HTTP/1.1\r\nHost: 127.0.0.1\r\n\r\n')

            with self.assertRaises(StopIteration):
                l4_client.receive_until(terminator=b'\r\n', max_line_length=3)
        self.assertEqual(b'220', l4_client.buffer)
            self.assertEqual(b'HTT', l4_client.buffer)
            l4_client.receive_until(terminator=b'\r\n')
            self.assertEqual(l4_client.buffer[-2:], b'\r\n')
        self.assertTrue(l4_client.buffer.decode('ascii').startswith('220 ' + address))
            self.assertTrue(l4_client.buffer.decode('ascii').startswith('HTTP/'))

            l4_client.close()
        finally:
            test_http_server.kill()

    def test_real(self):
        test_http_server = TestThreadedServerHttp('127.0.0.1', 0)
@@ -222,7 +234,6 @@ class TestL4ServerTCP(unittest.TestCase):
        self.assertEqual(context_manager.exception.error, NetworkErrorType.NO_ADDRESS)
        l4_server.close()

    @live_dns
    def test_error_wrong_address(self):
        l4_server = L4ServerTCP('8.8.8.8', 443)
        with self.assertRaises(NetworkError) as context_manager:
@@ -283,7 +294,6 @@ class TestL4ServerUDP(unittest.TestCase):
        self.assertEqual(context_manager.exception.error, NetworkErrorType.NO_ADDRESS)
        l4_server.close()

    @live_dns
    def test_error_wrong_address(self):
        l4_server = L4ServerUDP('8.8.8.8', 443)
        with self.assertRaises(NetworkError) as context_manager:
+3 −0
Original line number Diff line number Diff line
@@ -17,6 +17,9 @@ class TestDnsClient(unittest.TestCase):
            L7ClientDns.from_uri(urllib3.util.parse_url('unknown://mock.site'))
        self.assertEqual(context_manager.exception.args, ('unknown', ))

    def test_supported_schemes(self):
        self.assertEqual(L7ClientDns.get_supported_schemes(), {'dns': L7ClientDns})

    def test_client_dns(self):
        dns_handshake = DnsHandshakeBase(L4TransferSocketParams(timeout=5))
        l7_client = L7ClientDns.from_uri(urllib3.util.parse_url('dns://one.one.one.one'))
Loading