Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
tls.py2186 linesDownload Raw Back to aioquic
1import datetime
2import ipaddress
3import logging
4import os
5import ssl
6import struct
7from contextlib import contextmanager
8from dataclasses import dataclass, field
9from enum import Enum, IntEnum
10from functools import partial
11from typing import (
12    Any,
13    Callable,
14    Dict,
15    Generator,
16    List,
17    Optional,
18    Sequence,
19    Tuple,
20    TypeVar,
21    Union,
22    cast,
23)
24
25import certifi
26import service_identity
27from cryptography import x509
28from cryptography.exceptions import InvalidSignature
29from cryptography.hazmat.backends import default_backend
30from cryptography.hazmat.primitives import hashes, hmac, serialization
31from cryptography.hazmat.primitives.asymmetric import (
32    dsa,
33    ec,
34    ed448,
35    ed25519,
36    padding,
37    rsa,
38    x448,
39    x25519,
40)
41from cryptography.hazmat.primitives.asymmetric.types import (
42    CertificateIssuerPublicKeyTypes,
43    PrivateKeyTypes,
44)
45from cryptography.hazmat.primitives.kdf.hkdf import HKDFExpand
46from cryptography.hazmat.primitives.serialization import Encoding, PublicFormat
47from OpenSSL import crypto
48
49from .buffer import Buffer, BufferReadError
50
51TLS_VERSION_1_2 = 0x0303
52TLS_VERSION_1_3 = 0x0304
53TLS_VERSION_1_3_DRAFT_28 = 0x7F1C
54TLS_VERSION_1_3_DRAFT_27 = 0x7F1B
55TLS_VERSION_1_3_DRAFT_26 = 0x7F1A
56
57CLIENT_CONTEXT_STRING = b"TLS 1.3, client CertificateVerify"
58SERVER_CONTEXT_STRING = b"TLS 1.3, server CertificateVerify"
59
60T = TypeVar("T")
61
62
63# facilitate mocking for the test suite
64def utcnow() -> datetime.datetime:
65    return datetime.datetime.now(datetime.timezone.utc)
66
67
68class AlertDescription(IntEnum):
69    close_notify = 0
70    unexpected_message = 10
71    bad_record_mac = 20
72    record_overflow = 22
73    handshake_failure = 40
74    bad_certificate = 42
75    unsupported_certificate = 43
76    certificate_revoked = 44
77    certificate_expired = 45
78    certificate_unknown = 46
79    illegal_parameter = 47
80    unknown_ca = 48
81    access_denied = 49
82    decode_error = 50
83    decrypt_error = 51
84    protocol_version = 70
85    insufficient_security = 71
86    internal_error = 80
87    inappropriate_fallback = 86
88    user_canceled = 90
89    missing_extension = 109
90    unsupported_extension = 110
91    unrecognized_name = 112
92    bad_certificate_status_response = 113
93    unknown_psk_identity = 115
94    certificate_required = 116
95    no_application_protocol = 120
96
97
98class Alert(Exception):
99    description: AlertDescription
100
101
102class AlertBadCertificate(Alert):
103    description = AlertDescription.bad_certificate
104
105
106class AlertCertificateExpired(Alert):
107    description = AlertDescription.certificate_expired
108
109
110class AlertDecodeError(Alert):
111    description = AlertDescription.decode_error
112
113
114class AlertDecryptError(Alert):
115    description = AlertDescription.decrypt_error
116
117
118class AlertHandshakeFailure(Alert):
119    description = AlertDescription.handshake_failure
120
121
122class AlertIllegalParameter(Alert):
123    description = AlertDescription.illegal_parameter
124
125
126class AlertInternalError(Alert):
127    description = AlertDescription.internal_error
128
129
130class AlertProtocolVersion(Alert):
131    description = AlertDescription.protocol_version
132
133
134class AlertUnexpectedMessage(Alert):
135    description = AlertDescription.unexpected_message
136
137
138class Direction(Enum):
139    DECRYPT = 0
140    ENCRYPT = 1
141
142
143class Epoch(Enum):
144    INITIAL = 0
145    ZERO_RTT = 1
146    HANDSHAKE = 2
147    ONE_RTT = 3
148
149
150class State(Enum):
151    CLIENT_HANDSHAKE_START = 0
152    CLIENT_EXPECT_SERVER_HELLO = 1
153    CLIENT_EXPECT_ENCRYPTED_EXTENSIONS = 2
154    CLIENT_EXPECT_CERTIFICATE_REQUEST_OR_CERTIFICATE = 3
155    CLIENT_EXPECT_CERTIFICATE = 4
156    CLIENT_EXPECT_CERTIFICATE_VERIFY = 5
157    CLIENT_EXPECT_FINISHED = 6
158    CLIENT_POST_HANDSHAKE = 7
159
160    SERVER_EXPECT_CLIENT_HELLO = 8
161    SERVER_EXPECT_CERTIFICATE = 9
162    SERVER_EXPECT_CERTIFICATE_VERIFY = 10
163    SERVER_EXPECT_FINISHED = 11
164    SERVER_POST_HANDSHAKE = 12
165
166
167def hkdf_label(label: bytes, hash_value: bytes, length: int) -> bytes:
168    full_label = b"tls13 " + label
169    return (
170        struct.pack("!HB", length, len(full_label))
171        + full_label
172        + struct.pack("!B", len(hash_value))
173        + hash_value
174    )
175
176
177def hkdf_expand_label(
178    algorithm: hashes.HashAlgorithm,
179    secret: bytes,
180    label: bytes,
181    hash_value: bytes,
182    length: int,
183) -> bytes:
184    return HKDFExpand(
185        algorithm=algorithm,
186        length=length,
187        info=hkdf_label(label, hash_value, length),
188    ).derive(secret)
189
190
191def hkdf_extract(
192    algorithm: hashes.HashAlgorithm, salt: bytes, key_material: bytes
193) -> bytes:
194    h = hmac.HMAC(salt, algorithm)
195    h.update(key_material)
196    return h.finalize()
197
198
199def load_pem_private_key(
200    data: bytes, password: Optional[bytes] = None
201) -> PrivateKeyTypes:
202    """
203    Load a PEM-encoded private key.
204    """
205    return serialization.load_pem_private_key(data, password=password)
206
207
208def load_pem_x509_certificates(data: bytes) -> List[x509.Certificate]:
209    """
210    Load a chain of PEM-encoded X509 certificates.
211    """
212    boundary = b"-----END CERTIFICATE-----\n"
213    certificates = []
214    for chunk in data.split(boundary):
215        if chunk:
216            certificates.append(x509.load_pem_x509_certificate(chunk + boundary))
217    return certificates
218
219
220def verify_certificate(
221    certificate: x509.Certificate,
222    chain: List[x509.Certificate] = [],
223    server_name: Optional[str] = None,
224    cadata: Optional[bytes] = None,
225    cafile: Optional[str] = None,
226    capath: Optional[str] = None,
227) -> None:
228    # verify dates
229    now = utcnow()
230    if now < certificate.not_valid_before_utc:
231        raise AlertCertificateExpired("Certificate is not valid yet")
232    if now > certificate.not_valid_after_utc:
233        raise AlertCertificateExpired("Certificate is no longer valid")
234
235    # verify subject
236    if server_name is not None:
237        try:
238            ipaddress.ip_address(server_name)
239        except ValueError:
240            is_ip = False
241        else:
242            is_ip = True
243
244        try:
245            if is_ip:
246                service_identity.cryptography.verify_certificate_ip_address(
247                    certificate, server_name
248                )
249            else:
250                service_identity.cryptography.verify_certificate_hostname(
251                    certificate, server_name
252                )
253
254        except (
255            service_identity.CertificateError,
256            service_identity.VerificationError,
257        ) as exc:
258            patterns = service_identity.cryptography.extract_patterns(certificate)
259            if len(patterns) == 0:
260                errmsg = str(exc)
261            elif len(patterns) == 1:
262                errmsg = f"hostname {server_name!r} doesn't match {patterns[0]!r}"
263            else:
264                patterns_repr = ", ".join(repr(pattern) for pattern in patterns)
265                errmsg = (
266                    f"hostname {server_name!r} doesn't match "
267                    f"either of {patterns_repr}"
268                )
269
270            raise AlertBadCertificate(errmsg) from exc
271
272    # load CAs
273    store = crypto.X509Store()
274
275    if cadata is None and cafile is None and capath is None:
276        # Load defaults from certifi.
277        store.load_locations(certifi.where())
278
279    if cadata is not None:
280        for cert in load_pem_x509_certificates(cadata):
281            store.add_cert(crypto.X509.from_cryptography(cert))
282
283    if cafile is not None or capath is not None:
284        store.load_locations(cafile, capath)
285
286    # verify certificate chain
287    store_ctx = crypto.X509StoreContext(
288        store,
289        crypto.X509.from_cryptography(certificate),
290        [crypto.X509.from_cryptography(cert) for cert in chain],
291    )
292    try:
293        store_ctx.verify_certificate()
294    except crypto.X509StoreContextError as exc:
295        raise AlertBadCertificate(exc.args[0])
296
297
298class CipherSuite(IntEnum):
299    AES_128_GCM_SHA256 = 0x1301
300    AES_256_GCM_SHA384 = 0x1302
301    CHACHA20_POLY1305_SHA256 = 0x1303
302    EMPTY_RENEGOTIATION_INFO_SCSV = 0x00FF
303
304
305class CompressionMethod(IntEnum):
306    NULL = 0
307
308
309class ExtensionType(IntEnum):
310    SERVER_NAME = 0
311    STATUS_REQUEST = 5
312    SUPPORTED_GROUPS = 10
313    SIGNATURE_ALGORITHMS = 13
314    ALPN = 16
315    COMPRESS_CERTIFICATE = 27
316    PRE_SHARED_KEY = 41
317    EARLY_DATA = 42
318    SUPPORTED_VERSIONS = 43
319    COOKIE = 44
320    PSK_KEY_EXCHANGE_MODES = 45
321    KEY_SHARE = 51
322    QUIC_TRANSPORT_PARAMETERS = 0x0039
323    QUIC_TRANSPORT_PARAMETERS_DRAFT = 0xFFA5
324    ENCRYPTED_SERVER_NAME = 65486
325
326
327class Group(IntEnum):
328    SECP256R1 = 0x0017
329    SECP384R1 = 0x0018
330    SECP521R1 = 0x0019
331    X25519 = 0x001D
332    X448 = 0x001E
333    GREASE = 0xAAAA
334
335
336class HandshakeType(IntEnum):
337    CLIENT_HELLO = 1
338    SERVER_HELLO = 2
339    NEW_SESSION_TICKET = 4
340    END_OF_EARLY_DATA = 5
341    ENCRYPTED_EXTENSIONS = 8
342    CERTIFICATE = 11
343    CERTIFICATE_REQUEST = 13
344    CERTIFICATE_VERIFY = 15
345    FINISHED = 20
346    KEY_UPDATE = 24
347    COMPRESSED_CERTIFICATE = 25
348    MESSAGE_HASH = 254
349
350
351class NameType(IntEnum):
352    HOST_NAME = 0
353
354
355class PskKeyExchangeMode(IntEnum):
356    PSK_KE = 0
357    PSK_DHE_KE = 1
358
359
360class SignatureAlgorithm(IntEnum):
361    ECDSA_SECP256R1_SHA256 = 0x0403
362    ECDSA_SECP384R1_SHA384 = 0x0503
363    ECDSA_SECP521R1_SHA512 = 0x0603
364    ED25519 = 0x0807
365    ED448 = 0x0808
366    RSA_PKCS1_SHA256 = 0x0401
367    RSA_PKCS1_SHA384 = 0x0501
368    RSA_PKCS1_SHA512 = 0x0601
369    RSA_PSS_PSS_SHA256 = 0x0809
370    RSA_PSS_PSS_SHA384 = 0x080A
371    RSA_PSS_PSS_SHA512 = 0x080B
372    RSA_PSS_RSAE_SHA256 = 0x0804
373    RSA_PSS_RSAE_SHA384 = 0x0805
374    RSA_PSS_RSAE_SHA512 = 0x0806
375
376    # legacy
377    RSA_PKCS1_SHA1 = 0x0201
378    SHA1_DSA = 0x0202
379    ECDSA_SHA1 = 0x0203
380
381
382# BLOCKS
383
384
385@contextmanager
386def pull_block(buf: Buffer, capacity: int) -> Generator:
387    length = int.from_bytes(buf.pull_bytes(capacity), byteorder="big")
388    end = buf.tell() + length
389    yield length
390    if buf.tell() != end:
391        # There was trailing garbage or our parsing was bad.
392        raise AlertDecodeError("extra bytes at the end of a block")
393
394
395@contextmanager
396def push_block(buf: Buffer, capacity: int) -> Generator:
397    """
398    Context manager to push a variable-length block, with `capacity` bytes
399    to write the length.
400    """
401    start = buf.tell() + capacity
402    buf.seek(start)
403    yield
404    end = buf.tell()
405    length = end - start
406    buf.seek(start - capacity)
407    buf.push_bytes(length.to_bytes(capacity, byteorder="big"))
408    buf.seek(end)
409
410
411# LISTS
412
413
414class SkipItem(Exception):
415    "There is nothing to append for this invocation of a pull_list() func"
416
417
418def pull_list(buf: Buffer, capacity: int, func: Callable[[], T]) -> List[T]:
419    """
420    Pull a list of items.
421
422    If the callable raises SkipItem, then iteration continues but nothing
423    is added to the list.
424    """
425    items = []
426    with pull_block(buf, capacity) as length:
427        end = buf.tell() + length
428        while buf.tell() < end:
429            try:
430                items.append(func())
431            except SkipItem:
432                pass
433    return items
434
435
436def push_list(
437    buf: Buffer, capacity: int, func: Callable[[T], None], values: Sequence[T]
438) -> None:
439    """
440    Push a list of items.
441    """
442    with push_block(buf, capacity):
443        for value in values:
444            func(value)
445
446
447def pull_opaque(buf: Buffer, capacity: int) -> bytes:
448    """
449    Pull an opaque value prefixed by a length.
450    """
451    with pull_block(buf, capacity) as length:
452        return buf.pull_bytes(length)
453
454
455def push_opaque(buf: Buffer, capacity: int, value: bytes) -> None:
456    """
457    Push an opaque value prefix by a length.
458    """
459    with push_block(buf, capacity):
460        buf.push_bytes(value)
461
462
463@contextmanager
464def push_extension(buf: Buffer, extension_type: int) -> Generator:
465    buf.push_uint16(extension_type)
466    with push_block(buf, 2):
467        yield
468
469
470# ServerName
471
472
473def pull_server_name(buf: Buffer) -> str:
474    with pull_block(buf, 2):
475        name_type = buf.pull_uint8()
476        if name_type != NameType.HOST_NAME:
477            # We don't know this name_type.
478            raise AlertIllegalParameter(
479                f"ServerName has an unknown name type {name_type}"
480            )
481        return pull_opaque(buf, 2).decode("ascii")
482
483
484def push_server_name(buf: Buffer, server_name: str) -> None:
485    with push_block(buf, 2):
486        buf.push_uint8(NameType.HOST_NAME)
487        push_opaque(buf, 2, server_name.encode("ascii"))
488
489
490# KeyShareEntry
491
492
493KeyShareEntry = Tuple[int, bytes]
494
495
496def pull_key_share(buf: Buffer) -> KeyShareEntry:
497    group = buf.pull_uint16()
498    data = pull_opaque(buf, 2)
499    return (group, data)
500
501
502def push_key_share(buf: Buffer, value: KeyShareEntry) -> None:
503    buf.push_uint16(value[0])
504    push_opaque(buf, 2, value[1])
505
506
507# ALPN
508
509
510def pull_alpn_protocol(buf: Buffer) -> str:
511    try:
512        return pull_opaque(buf, 1).decode("ascii")
513    except UnicodeDecodeError:
514        # We can get arbitrary bytes values for alpns from greasing,
515        # but we expect them to be strings in the rest of the API, so
516        # we ignore them if they don't decode as ASCII.
517        raise SkipItem
518
519
520def push_alpn_protocol(buf: Buffer, protocol: str) -> None:
521    push_opaque(buf, 1, protocol.encode("ascii"))
522
523
524# PRE SHARED KEY
525
526PskIdentity = Tuple[bytes, int]
527
528
529@dataclass
530class OfferedPsks:
531    identities: List[PskIdentity]
532    binders: List[bytes]
533
534
535def pull_psk_identity(buf: Buffer) -> PskIdentity:
536    identity = pull_opaque(buf, 2)
537    obfuscated_ticket_age = buf.pull_uint32()
538    return (identity, obfuscated_ticket_age)
539
540
541def push_psk_identity(buf: Buffer, entry: PskIdentity) -> None:
542    push_opaque(buf, 2, entry[0])
543    buf.push_uint32(entry[1])
544
545
546def pull_psk_binder(buf: Buffer) -> bytes:
547    return pull_opaque(buf, 1)
548
549
550def push_psk_binder(buf: Buffer, binder: bytes) -> None:
551    push_opaque(buf, 1, binder)
552
553
554def pull_offered_psks(buf: Buffer) -> OfferedPsks:
555    return OfferedPsks(
556        identities=pull_list(buf, 2, partial(pull_psk_identity, buf)),
557        binders=pull_list(buf, 2, partial(pull_psk_binder, buf)),
558    )
559
560
561def push_offered_psks(buf: Buffer, pre_shared_key: OfferedPsks) -> None:
562    push_list(
563        buf,
564        2,
565        partial(push_psk_identity, buf),
566        pre_shared_key.identities,
567    )
568    push_list(
569        buf,
570        2,
571        partial(push_psk_binder, buf),
572        pre_shared_key.binders,
573    )
574
575
576# MESSAGES
577
578Extension = Tuple[int, bytes]
579
580
581@dataclass
582class ClientHello:
583    random: bytes
584    legacy_session_id: bytes
585    cipher_suites: List[int]
586    legacy_compression_methods: List[int]
587
588    # extensions
589    alpn_protocols: Optional[List[str]] = None
590    early_data: bool = False
591    key_share: Optional[List[KeyShareEntry]] = None
592    pre_shared_key: Optional[OfferedPsks] = None
593    psk_key_exchange_modes: Optional[List[int]] = None
594    server_name: Optional[str] = None
595    signature_algorithms: Optional[List[int]] = None
596    supported_groups: Optional[List[int]] = None
597    supported_versions: Optional[List[int]] = None
598
599    other_extensions: List[Extension] = field(default_factory=list)
600
601
602def pull_handshake_type(buf: Buffer, expected_type: HandshakeType) -> None:
603    """
604    Pull the message type and assert it is the expected one.
605
606    If it is not, we have a programming error.
607    """
608    message_type = buf.pull_uint8()
609    assert message_type == expected_type
610
611
612def pull_client_hello(buf: Buffer) -> ClientHello:
613    pull_handshake_type(buf, HandshakeType.CLIENT_HELLO)
614    with pull_block(buf, 3):
615        if buf.pull_uint16() != TLS_VERSION_1_2:
616            raise AlertDecodeError("ClientHello version is not 1.2")
617
618        hello = ClientHello(
619            random=buf.pull_bytes(32),
620            legacy_session_id=pull_opaque(buf, 1),
621            cipher_suites=pull_list(buf, 2, buf.pull_uint16),
622            legacy_compression_methods=pull_list(buf, 1, buf.pull_uint8),
623        )
624
625        # extensions
626        after_psk = False
627
628        def pull_extension() -> None:
629            # pre_shared_key MUST be last
630            nonlocal after_psk
631            if after_psk:
632                # the alert is Illegal Parameter per RFC 8446 section 4.2.11.
633                raise AlertIllegalParameter("PreSharedKey is not the last extension")
634
635            extension_type = buf.pull_uint16()
636            extension_length = buf.pull_uint16()
637            if extension_type == ExtensionType.KEY_SHARE:
638                hello.key_share = pull_list(buf, 2, partial(pull_key_share, buf))
639            elif extension_type == ExtensionType.SUPPORTED_VERSIONS:
640                hello.supported_versions = pull_list(buf, 1, buf.pull_uint16)
641            elif extension_type == ExtensionType.SIGNATURE_ALGORITHMS:
642                hello.signature_algorithms = pull_list(buf, 2, buf.pull_uint16)
643            elif extension_type == ExtensionType.SUPPORTED_GROUPS:
644                hello.supported_groups = pull_list(buf, 2, buf.pull_uint16)
645            elif extension_type == ExtensionType.PSK_KEY_EXCHANGE_MODES:
646                hello.psk_key_exchange_modes = pull_list(buf, 1, buf.pull_uint8)
647            elif extension_type == ExtensionType.SERVER_NAME:
648                hello.server_name = pull_server_name(buf)
649            elif extension_type == ExtensionType.ALPN:
650                hello.alpn_protocols = pull_list(
651                    buf, 2, partial(pull_alpn_protocol, buf)
652                )
653            elif extension_type == ExtensionType.EARLY_DATA:
654                hello.early_data = True
655            elif extension_type == ExtensionType.PRE_SHARED_KEY:
656                hello.pre_shared_key = pull_offered_psks(buf)
657                after_psk = True
658            else:
659                hello.other_extensions.append(
660                    (extension_type, buf.pull_bytes(extension_length))
661                )
662
663        pull_list(buf, 2, pull_extension)
664
665    return hello
666
667
668def push_client_hello(buf: Buffer, hello: ClientHello) -> None:
669    buf.push_uint8(HandshakeType.CLIENT_HELLO)
670    with push_block(buf, 3):
671        buf.push_uint16(TLS_VERSION_1_2)
672        buf.push_bytes(hello.random)
673        push_opaque(buf, 1, hello.legacy_session_id)
674        push_list(buf, 2, buf.push_uint16, hello.cipher_suites)
675        push_list(buf, 1, buf.push_uint8, hello.legacy_compression_methods)
676
677        # extensions
678        with push_block(buf, 2):
679            with push_extension(buf, ExtensionType.KEY_SHARE):
680                push_list(buf, 2, partial(push_key_share, buf), hello.key_share)
681
682            with push_extension(buf, ExtensionType.SUPPORTED_VERSIONS):
683                push_list(buf, 1, buf.push_uint16, hello.supported_versions)
684
685            with push_extension(buf, ExtensionType.SIGNATURE_ALGORITHMS):
686                push_list(buf, 2, buf.push_uint16, hello.signature_algorithms)
687
688            with push_extension(buf, ExtensionType.SUPPORTED_GROUPS):
689                push_list(buf, 2, buf.push_uint16, hello.supported_groups)
690
691            if hello.psk_key_exchange_modes is not None:
692                with push_extension(buf, ExtensionType.PSK_KEY_EXCHANGE_MODES):
693                    push_list(buf, 1, buf.push_uint8, hello.psk_key_exchange_modes)
694
695            if hello.server_name is not None:
696                with push_extension(buf, ExtensionType.SERVER_NAME):
697                    push_server_name(buf, hello.server_name)
698
699            if hello.alpn_protocols is not None:
700                with push_extension(buf, ExtensionType.ALPN):
701                    push_list(
702                        buf, 2, partial(push_alpn_protocol, buf), hello.alpn_protocols
703                    )
704
705            for extension_type, extension_value in hello.other_extensions:
706                with push_extension(buf, extension_type):
707                    buf.push_bytes(extension_value)
708
709            if hello.early_data:
710                with push_extension(buf, ExtensionType.EARLY_DATA):
711                    pass
712
713            # pre_shared_key MUST be last
714            if hello.pre_shared_key is not None:
715                with push_extension(buf, ExtensionType.PRE_SHARED_KEY):
716                    push_offered_psks(buf, hello.pre_shared_key)
717
718
719@dataclass
720class ServerHello:
721    random: bytes
722    legacy_session_id: bytes
723    cipher_suite: int
724    compression_method: int
725
726    # extensions
727    key_share: Optional[KeyShareEntry] = None
728    pre_shared_key: Optional[int] = None
729    supported_version: Optional[int] = None
730    other_extensions: List[Tuple[int, bytes]] = field(default_factory=list)
731
732
733def pull_server_hello(buf: Buffer) -> ServerHello:
734    pull_handshake_type(buf, HandshakeType.SERVER_HELLO)
735    with pull_block(buf, 3):
736        if buf.pull_uint16() != TLS_VERSION_1_2:
737            raise AlertDecodeError("ServerHello version is not 1.2")
738
739        hello = ServerHello(
740            random=buf.pull_bytes(32),
741            legacy_session_id=pull_opaque(buf, 1),
742            cipher_suite=buf.pull_uint16(),
743            compression_method=buf.pull_uint8(),
744        )
745
746        # extensions
747        def pull_extension() -> None:
748            extension_type = buf.pull_uint16()
749            extension_length = buf.pull_uint16()
750            if extension_type == ExtensionType.SUPPORTED_VERSIONS:
751                hello.supported_version = buf.pull_uint16()
752            elif extension_type == ExtensionType.KEY_SHARE:
753                hello.key_share = pull_key_share(buf)
754            elif extension_type == ExtensionType.PRE_SHARED_KEY:
755                hello.pre_shared_key = buf.pull_uint16()
756            else:
757                hello.other_extensions.append(
758                    (extension_type, buf.pull_bytes(extension_length))
759                )
760
761        pull_list(buf, 2, pull_extension)
762
763    return hello
764
765
766def push_server_hello(buf: Buffer, hello: ServerHello) -> None:
767    buf.push_uint8(HandshakeType.SERVER_HELLO)
768    with push_block(buf, 3):
769        buf.push_uint16(TLS_VERSION_1_2)
770        buf.push_bytes(hello.random)
771
772        push_opaque(buf, 1, hello.legacy_session_id)
773        buf.push_uint16(hello.cipher_suite)
774        buf.push_uint8(hello.compression_method)
775
776        # extensions
777        with push_block(buf, 2):
778            if hello.supported_version is not None:
779                with push_extension(buf, ExtensionType.SUPPORTED_VERSIONS):
780                    buf.push_uint16(hello.supported_version)
781
782            if hello.key_share is not None:
783                with push_extension(buf, ExtensionType.KEY_SHARE):
784                    push_key_share(buf, hello.key_share)
785
786            if hello.pre_shared_key is not None:
787                with push_extension(buf, ExtensionType.PRE_SHARED_KEY):
788                    buf.push_uint16(hello.pre_shared_key)
789
790            for extension_type, extension_value in hello.other_extensions:
791                with push_extension(buf, extension_type):
792                    buf.push_bytes(extension_value)
793
794
795@dataclass
796class NewSessionTicket:
797    ticket_lifetime: int = 0
798    ticket_age_add: int = 0
799    ticket_nonce: bytes = b""
800    ticket: bytes = b""
801
802    # extensions
803    max_early_data_size: Optional[int] = None
804    other_extensions: List[Tuple[int, bytes]] = field(default_factory=list)
805
806
807def pull_new_session_ticket(buf: Buffer) -> NewSessionTicket:
808    new_session_ticket = NewSessionTicket()
809
810    pull_handshake_type(buf, HandshakeType.NEW_SESSION_TICKET)
811    with pull_block(buf, 3):
812        new_session_ticket.ticket_lifetime = buf.pull_uint32()
813        new_session_ticket.ticket_age_add = buf.pull_uint32()
814        new_session_ticket.ticket_nonce = pull_opaque(buf, 1)
815        new_session_ticket.ticket = pull_opaque(buf, 2)
816
817        def pull_extension() -> None:
818            extension_type = buf.pull_uint16()
819            extension_length = buf.pull_uint16()
820            if extension_type == ExtensionType.EARLY_DATA:
821                new_session_ticket.max_early_data_size = buf.pull_uint32()
822            else:
823                new_session_ticket.other_extensions.append(
824                    (extension_type, buf.pull_bytes(extension_length))
825                )
826
827        pull_list(buf, 2, pull_extension)
828
829    return new_session_ticket
830
831
832def push_new_session_ticket(buf: Buffer, new_session_ticket: NewSessionTicket) -> None:
833    buf.push_uint8(HandshakeType.NEW_SESSION_TICKET)
834    with push_block(buf, 3):
835        buf.push_uint32(new_session_ticket.ticket_lifetime)
836        buf.push_uint32(new_session_ticket.ticket_age_add)
837        push_opaque(buf, 1, new_session_ticket.ticket_nonce)
838        push_opaque(buf, 2, new_session_ticket.ticket)
839
840        with push_block(buf, 2):
841            if new_session_ticket.max_early_data_size is not None:
842                with push_extension(buf, ExtensionType.EARLY_DATA):
843                    buf.push_uint32(new_session_ticket.max_early_data_size)
844
845            for extension_type, extension_value in new_session_ticket.other_extensions:
846                with push_extension(buf, extension_type):
847                    buf.push_bytes(extension_value)
848
849
850@dataclass
851class EncryptedExtensions:
852    alpn_protocol: Optional[str] = None
853    early_data: bool = False
854
855    other_extensions: List[Tuple[int, bytes]] = field(default_factory=list)
856
857
858def pull_encrypted_extensions(buf: Buffer) -> EncryptedExtensions:
859    extensions = EncryptedExtensions()
860
861    pull_handshake_type(buf, HandshakeType.ENCRYPTED_EXTENSIONS)
862    with pull_block(buf, 3):
863
864        def pull_extension() -> None:
865            extension_type = buf.pull_uint16()
866            extension_length = buf.pull_uint16()
867            if extension_type == ExtensionType.ALPN:
868                extensions.alpn_protocol = pull_list(
869                    buf, 2, partial(pull_alpn_protocol, buf)
870                )[0]
871            elif extension_type == ExtensionType.EARLY_DATA:
872                extensions.early_data = True
873            else:
874                extensions.other_extensions.append(
875                    (extension_type, buf.pull_bytes(extension_length))
876                )
877
878        pull_list(buf, 2, pull_extension)
879
880    return extensions
881
882
883def push_encrypted_extensions(buf: Buffer, extensions: EncryptedExtensions) -> None:
884    buf.push_uint8(HandshakeType.ENCRYPTED_EXTENSIONS)
885    with push_block(buf, 3):
886        with push_block(buf, 2):
887            if extensions.alpn_protocol is not None:
888                with push_extension(buf, ExtensionType.ALPN):
889                    push_list(
890                        buf,
891                        2,
892                        partial(push_alpn_protocol, buf),
893                        [extensions.alpn_protocol],
894                    )
895
896            if extensions.early_data:
897                with push_extension(buf, ExtensionType.EARLY_DATA):
898                    pass
899
900            for extension_type, extension_value in extensions.other_extensions:
901                with push_extension(buf, extension_type):
902                    buf.push_bytes(extension_value)
903
904
905CertificateEntry = Tuple[bytes, bytes]
906
907
908@dataclass
909class Certificate:
910    request_context: bytes = b""
911    certificates: List[CertificateEntry] = field(default_factory=list)
912
913
914def pull_certificate(buf: Buffer) -> Certificate:
915    certificate = Certificate()
916
917    pull_handshake_type(buf, HandshakeType.CERTIFICATE)
918    with pull_block(buf, 3):
919        certificate.request_context = pull_opaque(buf, 1)
920
921        def pull_certificate_entry(buf: Buffer) -> CertificateEntry:
922            data = pull_opaque(buf, 3)
923            extensions = pull_opaque(buf, 2)
924            return (data, extensions)
925
926        certificate.certificates = pull_list(
927            buf, 3, partial(pull_certificate_entry, buf)
928        )
929
930    return certificate
931
932
933def push_certificate(buf: Buffer, certificate: Certificate) -> None:
934    buf.push_uint8(HandshakeType.CERTIFICATE)
935    with push_block(buf, 3):
936        push_opaque(buf, 1, certificate.request_context)
937
938        def push_certificate_entry(buf: Buffer, entry: CertificateEntry) -> None:
939            push_opaque(buf, 3, entry[0])
940            push_opaque(buf, 2, entry[1])
941
942        push_list(
943            buf, 3, partial(push_certificate_entry, buf), certificate.certificates
944        )
945
946
947@dataclass
948class CertificateRequest:
949    request_context: bytes = b""
950    signature_algorithms: Optional[List[int]] = None
951    other_extensions: List[Tuple[int, bytes]] = field(default_factory=list)
952
953
954def pull_certificate_request(buf: Buffer) -> CertificateRequest:
955    certificate_request = CertificateRequest()
956
957    pull_handshake_type(buf, HandshakeType.CERTIFICATE_REQUEST)
958    with pull_block(buf, 3):
959        certificate_request.request_context = pull_opaque(buf, 1)
960
961        def pull_extension() -> None:
962            extension_type = buf.pull_uint16()
963            extension_length = buf.pull_uint16()
964            if extension_type == ExtensionType.SIGNATURE_ALGORITHMS:
965                certificate_request.signature_algorithms = pull_list(
966                    buf, 2, buf.pull_uint16
967                )
968            else:
969                certificate_request.other_extensions.append(
970                    (extension_type, buf.pull_bytes(extension_length))
971                )
972
973        pull_list(buf, 2, pull_extension)
974
975    return certificate_request
976
977
978def push_certificate_request(
979    buf: Buffer, certificate_request: CertificateRequest
980) -> None:
981    buf.push_uint8(HandshakeType.CERTIFICATE_REQUEST)
982    with push_block(buf, 3):
983        push_opaque(buf, 1, certificate_request.request_context)
984
985        with push_block(buf, 2):
986            with push_extension(buf, ExtensionType.SIGNATURE_ALGORITHMS):
987                push_list(
988                    buf, 2, buf.push_uint16, certificate_request.signature_algorithms
989                )
990
991            for extension_type, extension_value in certificate_request.other_extensions:
992                with push_extension(buf, extension_type):
993                    buf.push_bytes(extension_value)
994
995
996@dataclass
997class CertificateVerify:
998    algorithm: int
999    signature: bytes
1000
1001
1002def pull_certificate_verify(buf: Buffer) -> CertificateVerify:
1003    pull_handshake_type(buf, HandshakeType.CERTIFICATE_VERIFY)
1004    with pull_block(buf, 3):
1005        algorithm = buf.pull_uint16()
1006        signature = pull_opaque(buf, 2)
1007
1008    return CertificateVerify(algorithm=algorithm, signature=signature)
1009
1010
1011def push_certificate_verify(buf: Buffer, verify: CertificateVerify) -> None:
1012    buf.push_uint8(HandshakeType.CERTIFICATE_VERIFY)
1013    with push_block(buf, 3):
1014        buf.push_uint16(verify.algorithm)
1015        push_opaque(buf, 2, verify.signature)
1016
1017
1018@dataclass
1019class Finished:
1020    verify_data: bytes = b""
1021
1022
1023def pull_finished(buf: Buffer) -> Finished:
1024    finished = Finished()
1025
1026    pull_handshake_type(buf, HandshakeType.FINISHED)
1027    finished.verify_data = pull_opaque(buf, 3)
1028
1029    return finished
1030
1031
1032def push_finished(buf: Buffer, finished: Finished) -> None:
1033    buf.push_uint8(HandshakeType.FINISHED)
1034    push_opaque(buf, 3, finished.verify_data)
1035
1036
1037# CONTEXT
1038
1039
1040class KeySchedule:
1041    def __init__(self, cipher_suite: CipherSuite):
1042        self.algorithm = cipher_suite_hash(cipher_suite)
1043        self.cipher_suite = cipher_suite
1044        self.generation = 0
1045        self.hash = hashes.Hash(self.algorithm)
1046        self.hash_empty_value = self.hash.copy().finalize()
1047        self.secret = bytes(self.algorithm.digest_size)
1048
1049    def certificate_verify_data(self, context_string: bytes) -> bytes:
1050        return b" " * 64 + context_string + b"\x00" + self.hash.copy().finalize()
1051
1052    def finished_verify_data(self, secret: bytes) -> bytes:
1053        hmac_key = hkdf_expand_label(
1054            algorithm=self.algorithm,
1055            secret=secret,
1056            label=b"finished",
1057            hash_value=b"",
1058            length=self.algorithm.digest_size,
1059        )
1060
1061        h = hmac.HMAC(hmac_key, algorithm=self.algorithm)
1062        h.update(self.hash.copy().finalize())
1063        return h.finalize()
1064
1065    def derive_secret(self, label: bytes) -> bytes:
1066        return hkdf_expand_label(
1067            algorithm=self.algorithm,
1068            secret=self.secret,
1069            label=label,
1070            hash_value=self.hash.copy().finalize(),
1071            length=self.algorithm.digest_size,
1072        )
1073
1074    def extract(self, key_material: Optional[bytes] = None) -> None:
1075        if key_material is None:
1076            key_material = bytes(self.algorithm.digest_size)
1077
1078        if self.generation:
1079            self.secret = hkdf_expand_label(
1080                algorithm=self.algorithm,
1081                secret=self.secret,
1082                label=b"derived",
1083                hash_value=self.hash_empty_value,
1084                length=self.algorithm.digest_size,
1085            )
1086
1087        self.generation += 1
1088        self.secret = hkdf_extract(
1089            algorithm=self.algorithm, salt=self.secret, key_material=key_material
1090        )
1091
1092    def update_hash(self, data: bytes) -> None:
1093        self.hash.update(data)
1094
1095
1096class KeyScheduleProxy:
1097    def __init__(self, cipher_suites: List[CipherSuite]):
1098        self.__schedules = dict(map(lambda c: (c, KeySchedule(c)), cipher_suites))
1099
1100    def extract(self, key_material: Optional[bytes] = None) -> None:
1101        for k in self.__schedules.values():
1102            k.extract(key_material)
1103
1104    def select(self, cipher_suite: CipherSuite) -> KeySchedule:
1105        return self.__schedules[cipher_suite]
1106
1107    def update_hash(self, data: bytes) -> None:
1108        for k in self.__schedules.values():
1109            k.update_hash(data)
1110
1111
1112CIPHER_SUITES: Dict = {
1113    CipherSuite.AES_128_GCM_SHA256: hashes.SHA256,
1114    CipherSuite.AES_256_GCM_SHA384: hashes.SHA384,
1115    CipherSuite.CHACHA20_POLY1305_SHA256: hashes.SHA256,
1116}
1117
1118SIGNATURE_ALGORITHMS: Dict = {
1119    SignatureAlgorithm.ECDSA_SECP256R1_SHA256: (None, hashes.SHA256),
1120    SignatureAlgorithm.ECDSA_SECP384R1_SHA384: (None, hashes.SHA384),
1121    SignatureAlgorithm.ECDSA_SECP521R1_SHA512: (None, hashes.SHA512),
1122    SignatureAlgorithm.RSA_PKCS1_SHA1: (padding.PKCS1v15, hashes.SHA1),
1123    SignatureAlgorithm.RSA_PKCS1_SHA256: (padding.PKCS1v15, hashes.SHA256),
1124    SignatureAlgorithm.RSA_PKCS1_SHA384: (padding.PKCS1v15, hashes.SHA384),
1125    SignatureAlgorithm.RSA_PKCS1_SHA512: (padding.PKCS1v15, hashes.SHA512),
1126    SignatureAlgorithm.RSA_PSS_RSAE_SHA256: (padding.PSS, hashes.SHA256),
1127    SignatureAlgorithm.RSA_PSS_RSAE_SHA384: (padding.PSS, hashes.SHA384),
1128    SignatureAlgorithm.RSA_PSS_RSAE_SHA512: (padding.PSS, hashes.SHA512),
1129}
1130
1131GROUP_TO_CURVE: Dict = {
1132    Group.SECP256R1: ec.SECP256R1,
1133    Group.SECP384R1: ec.SECP384R1,
1134    Group.SECP521R1: ec.SECP521R1,
1135}
1136CURVE_TO_GROUP = dict((v, k) for k, v in GROUP_TO_CURVE.items())
1137
1138
1139def cipher_suite_hash(cipher_suite: CipherSuite) -> hashes.HashAlgorithm:
1140    return CIPHER_SUITES[cipher_suite]()
1141
1142
1143def decode_public_key(
1144    key_share: KeyShareEntry,
1145) -> Union[ec.EllipticCurvePublicKey, x25519.X25519PublicKey, x448.X448PublicKey, None]:
1146    if key_share[0] == Group.X25519:
1147        return x25519.X25519PublicKey.from_public_bytes(key_share[1])
1148    elif key_share[0] == Group.X448:
1149        return x448.X448PublicKey.from_public_bytes(key_share[1])
1150    elif key_share[0] in GROUP_TO_CURVE:
1151        return ec.EllipticCurvePublicKey.from_encoded_point(
1152            GROUP_TO_CURVE[key_share[0]](), key_share[1]
1153        )
1154    else:
1155        return None
1156
1157
1158def encode_public_key(
1159    public_key: Union[
1160        ec.EllipticCurvePublicKey, x25519.X25519PublicKey, x448.X448PublicKey
1161    ],
1162) -> KeyShareEntry:
1163    if isinstance(public_key, x25519.X25519PublicKey):
1164        return (Group.X25519, public_key.public_bytes(Encoding.Raw, PublicFormat.Raw))
1165    elif isinstance(public_key, x448.X448PublicKey):
1166        return (Group.X448, public_key.public_bytes(Encoding.Raw, PublicFormat.Raw))
1167    return (
1168        CURVE_TO_GROUP[public_key.curve.__class__],
1169        public_key.public_bytes(Encoding.X962, PublicFormat.UncompressedPoint),
1170    )
1171
1172
1173def negotiate(
1174    supported: List[T], offered: Optional[List[Any]], exc: Optional[Alert] = None
1175) -> T:
1176    if offered is not None:
1177        for c in supported:
1178            if c in offered:
1179                return c
1180
1181    if exc is not None:
1182        raise exc
1183    return None
1184
1185
1186def signature_algorithm_params(signature_algorithm: int) -> Tuple:
1187    if signature_algorithm in (SignatureAlgorithm.ED25519, SignatureAlgorithm.ED448):
1188        return tuple()
1189
1190    padding_cls, algorithm_cls = SIGNATURE_ALGORITHMS[signature_algorithm]
1191    algorithm = algorithm_cls()
1192    if padding_cls is None:
1193        return (ec.ECDSA(algorithm),)
1194    elif padding_cls == padding.PSS:
1195        padding_obj = padding_cls(
1196            mgf=padding.MGF1(algorithm), salt_length=algorithm.digest_size
1197        )
1198    else:
1199        padding_obj = padding_cls()
1200    return padding_obj, algorithm

Showing the first 1,200 of 2186 lines. Download the file for the rest.

codekingpro/portable-devtools · Team Ai