codekingpro/portable-devtools
114k
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
