Loading cryptolyzer/ssh/server.py +154 −19 Original line number Diff line number Diff line Loading @@ -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, Loading @@ -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( Loading Loading @@ -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 Loading @@ -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): Loading Loading @@ -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 Loading @@ -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() Loading @@ -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: Loading cryptolyzer/tls/server.py +29 −0 Original line number Diff line number Diff line Loading @@ -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, Loading @@ -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, Loading Loading @@ -79,6 +82,7 @@ from cryptoparser.tls.subprotocol import ( TlsDistinguishedName, TlsECCurveType, TlsHandshakeCertificateRequest, TlsHandshakeCertificateStatus, TlsHandshakeServerCertificate, TlsHandshakeServerHelloDone, TlsHandshakeServerHello, Loading Loading @@ -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( Loading Loading @@ -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: Loading Loading @@ -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): Loading Loading @@ -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 Loading Loading @@ -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) Loading Loading @@ -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()) Loading test/common/test_analyzer.py +12 −0 Original line number Diff line number Diff line Loading @@ -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 Loading @@ -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.""" Loading test/common/test_transfer.py +43 −33 Original line number Diff line number Diff line Loading @@ -10,7 +10,6 @@ from test.common.classes import ( TestThreadedServerHttpProxy, TestHTTPProxyRequestHandler, ) from test.common.markers import live_dns, live_server import urllib3 Loading @@ -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, _): Loading @@ -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) Loading Loading @@ -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: Loading Loading @@ -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: Loading test/dnsrec/test_client.py +3 −0 Original line number Diff line number Diff line Loading @@ -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 Loading
cryptolyzer/ssh/server.py +154 −19 Original line number Diff line number Diff line Loading @@ -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, Loading @@ -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( Loading Loading @@ -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 Loading @@ -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): Loading Loading @@ -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 Loading @@ -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() Loading @@ -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: Loading
cryptolyzer/tls/server.py +29 −0 Original line number Diff line number Diff line Loading @@ -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, Loading @@ -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, Loading Loading @@ -79,6 +82,7 @@ from cryptoparser.tls.subprotocol import ( TlsDistinguishedName, TlsECCurveType, TlsHandshakeCertificateRequest, TlsHandshakeCertificateStatus, TlsHandshakeServerCertificate, TlsHandshakeServerHelloDone, TlsHandshakeServerHello, Loading Loading @@ -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( Loading Loading @@ -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: Loading Loading @@ -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): Loading Loading @@ -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 Loading Loading @@ -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) Loading Loading @@ -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()) Loading
test/common/test_analyzer.py +12 −0 Original line number Diff line number Diff line Loading @@ -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 Loading @@ -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.""" Loading
test/common/test_transfer.py +43 −33 Original line number Diff line number Diff line Loading @@ -10,7 +10,6 @@ from test.common.classes import ( TestThreadedServerHttpProxy, TestHTTPProxyRequestHandler, ) from test.common.markers import live_dns, live_server import urllib3 Loading @@ -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, _): Loading @@ -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) Loading Loading @@ -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: Loading Loading @@ -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: Loading
test/dnsrec/test_client.py +3 −0 Original line number Diff line number Diff line Loading @@ -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