diff --git a/README.md b/README.md index 5410cee..90cee27 100644 --- a/README.md +++ b/README.md @@ -124,3 +124,20 @@ These legacy ciphers are for compatibility/migration workflows and are not recom ## License crimimal use only lol + +### MTProto 2.0 + +```python +from crypto_standalone.mtproto import encode_encrypted_message, decode_encrypted_message, MessageDirection, SessionValidationState + +state = SessionValidationState() +enc = encode_encrypted_message( + auth_key=b"\x00" * 256, + server_salt=123, + session_id=456, + msg_id=state.generate_msg_id(), + seq_no=1, + body=b"ping", + direction=MessageDirection.CLIENT_TO_SERVER +) +``` diff --git a/docs/MTPROTO.md b/docs/MTPROTO.md new file mode 100644 index 0000000..0e35e5a --- /dev/null +++ b/docs/MTPROTO.md @@ -0,0 +1,48 @@ +# MTProto 2.0 Framing Architecture + +The MTProto 2.0 implementation in `crypto_standalone` provides a secure, zero-dependency, pure-Python implementation of the framing and cryptographic envelope layers of the Telegram MTProto protocol. + +## Supported Transports + +All four standard MTProto TCP transports are supported: + +- **Abridged**: The lightest protocol. Envelope includes 1-byte marker `0xef` and short/extended lengths. +- **Intermediate**: 4-byte lengths. Envelope includes 4-byte marker `0xeeeeeeee`. +- **Padded-Intermediate**: Padded version to bypass ISP blocks. Envelope includes marker `0xdddddddd` and random padding of 0-15 bytes. +- **Full**: Includes sequence numbers and CRC32 checksums. + +Quick ACKs are supported across all transports. The codecs implement incremental, strict parsers suitable for real TCP streams without allocating unbounded buffers. + +## Cryptographic Messages + +The module implements the MTProto 2.0 message envelope: + +- **MTProto 2.0 Key Derivation**: Derives the 256-bit AES key and IV from the authorization key and message key. +- **AES-IGE**: Uses the existing pure-Python AES-256 implementation with Infinite Garble Extension (IGE). +- **Validation**: Enforces mandatory, strictly ordered validation gates (bounds, auth_key resolution, AES decryption, msg_key verification, internal field validation, and msg_id checks) to prevent padding or lengths from serving as side-channel oracles. + +### Security Limitations + +Because this library is written entirely in Python, it cannot guarantee complete side-channel immunity at the interpreter level: +- There is no hardware-level constant-time execution guarantee (though constant-time comparisons like `compare_digest` are used where applicable). +- Memory cannot be securely zeroized out of Python's garbage collector. +- Immutable strings may leave remnants of plaintext, keys, and message contents in memory. + +## Usage Example + +```python +from crypto_standalone.mtproto import AbridgedTransportCodec, PayloadFrame + +codec = AbridgedTransportCodec() +header = codec.connection_header() +# send header to socket + +encrypted_payload = b"..." +wire_packet = codec.encode(encrypted_payload) +# send wire_packet to socket + +events = codec.feed_data(received_chunk) +for event in events: + if isinstance(event, PayloadFrame): + print("Received payload of length", len(event.payload)) +``` diff --git a/src/crypto_standalone/mtproto/__init__.py b/src/crypto_standalone/mtproto/__init__.py new file mode 100644 index 0000000..2f0db72 --- /dev/null +++ b/src/crypto_standalone/mtproto/__init__.py @@ -0,0 +1,23 @@ +"""MTProto 2.0 framing implementation.""" + +from .core import MessageDirection, ResourceLimits, PayloadFrame, QuickAck, TransportError +from .errors import ( + MTProtoFramingError, IncompleteFrameError, InvalidFrameLengthError, + FrameTooLargeError, TransportSequenceError, TransportChecksumError, + MessageKeyMismatchError, SessionMismatchError, ReplayDetectedError +) +from .transports import AbridgedTransportCodec, IntermediateTransportCodec, PaddedIntermediateTransportCodec, FullTransportCodec +from .session import SessionValidationState +from .envelope import UnencryptedMessage, EncryptedMessage, encode_unencrypted_message, decode_unencrypted_message, encode_encrypted_message, decode_encrypted_message +from .containers import ContainerMessage, encode_msg_container, decode_msg_container + +__all__ = [ + "MessageDirection", "ResourceLimits", "PayloadFrame", "QuickAck", "TransportError", + "MTProtoFramingError", "IncompleteFrameError", "InvalidFrameLengthError", + "FrameTooLargeError", "TransportSequenceError", "TransportChecksumError", + "MessageKeyMismatchError", "SessionMismatchError", "ReplayDetectedError", + "AbridgedTransportCodec", "IntermediateTransportCodec", "PaddedIntermediateTransportCodec", "FullTransportCodec", + "SessionValidationState", + "UnencryptedMessage", "EncryptedMessage", "encode_unencrypted_message", "decode_unencrypted_message", "encode_encrypted_message", "decode_encrypted_message", + "ContainerMessage", "encode_msg_container", "decode_msg_container" +] diff --git a/src/crypto_standalone/mtproto/containers.py b/src/crypto_standalone/mtproto/containers.py new file mode 100644 index 0000000..d20516e --- /dev/null +++ b/src/crypto_standalone/mtproto/containers.py @@ -0,0 +1,91 @@ +import struct +from dataclasses import dataclass +from typing import List + +from .errors import InvalidFrameLengthError, MTProtoFramingError +from .core import ResourceLimits, DEFAULT_LIMITS + +@dataclass +class ContainerMessage: + msg_id: int + seq_no: int + body: bytes + +def encode_msg_container( + messages: List[ContainerMessage], + limits: ResourceLimits = DEFAULT_LIMITS +) -> bytes: + """ + Encodes a msg_container containing a list of messages. + msg_container constructor ID is 0x73f1f8dc. + """ + if len(messages) > limits.max_contained_message_count: + raise MTProtoFramingError(f"Too many messages in container: {len(messages)}") + + out = bytearray(struct.pack(" limits.max_container_size: + raise MTProtoFramingError("Container aggregate size exceeds limit") + + return bytes(out) + +def decode_msg_container( + data: bytes, + container_msg_id: int | None = None, + limits: ResourceLimits = DEFAULT_LIMITS +) -> List[ContainerMessage]: + """ + Decodes a msg_container constructor. + Does not parse recursively. + """ + if len(data) < 8: + raise InvalidFrameLengthError("Container data too small") + + constructor, count = struct.unpack(" limits.max_contained_message_count: + raise MTProtoFramingError(f"Container declares too many messages: {count}") + + messages = [] + offset = 8 + + for _ in range(count): + if len(data) - offset < 20: + raise InvalidFrameLengthError("Truncated nested message header") + + msg_id, seq_no, length = struct.unpack("= container_msg_id: + raise MTProtoFramingError(f"Nested msg_id {msg_id} not strictly lower than container msg_id {container_msg_id}") + + # Optional: verify not a nested simple container. + # We can look at the first 4 bytes if length >= 4 + if length >= 4: + inner_constructor = struct.unpack(" bytes: + if len(body) % 4 != 0: + raise InvalidFrameLengthError("Body length must be divisible by 4") + + return struct.pack(" UnencryptedMessage: + if len(data) < 20: + raise MTProtoFramingError("Unencrypted message too small") + + auth_key_id, msg_id, length = struct.unpack(" tuple[bytes, bytes]: + """Derives AES key and IV for MTProto 2.0.""" + x = 0 if direction == MessageDirection.CLIENT_TO_SERVER else 8 + + sha256_a = sha256(msg_key + auth_key[x:x+36]) + sha256_b = sha256(auth_key[x+40:x+76] + msg_key) + + aes_key = sha256_a[:8] + sha256_b[8:24] + sha256_a[24:32] + aes_iv = sha256_b[:8] + sha256_a[8:24] + sha256_b[24:32] + + return aes_key, aes_iv + +def encode_encrypted_message( + *, + auth_key: bytes, + server_salt: int, + session_id: int, + msg_id: int, + seq_no: int, + body: bytes, + direction: MessageDirection, +) -> bytes: + if len(auth_key) != 256: + raise ValueError("auth_key must be exactly 256 bytes") + if len(body) % 4 != 0: + raise ValueError("body length must be divisible by 4") + + # Internal header + internal_header = struct.pack(" EncryptedMessage: + if len(auth_key) != 256: + raise ValueError("auth_key must be exactly 256 bytes") + + if len(data) < 24: + raise InvalidFrameLengthError("Data too small for external header") + + _auth_key_id = struct.unpack("0 and divisible by 16") + + aes_key, aes_iv = _derive_aes_key_iv(auth_key, msg_key, direction) + aes = AES256(aes_key) + + # Decrypt into bounded buffer + decrypted = aes.decrypt_ige(encrypted_data, aes_iv) + + # Recompute msg_key + auth_key_fragment = auth_key[88:120] + expected_msg_key_hash = sha256(auth_key_fragment + decrypted) + expected_msg_key = expected_msg_key_hash[8:24] + + if not compare_digest(msg_key, expected_msg_key): + raise MessageKeyMismatchError("msg_key mismatch") + + if len(decrypted) < 32: + raise InvalidFrameLengthError("Minimum plaintext size not met") + + server_salt, session_id, msg_id, seq_no, msg_len = struct.unpack(" 1024: + raise InvalidFrameLengthError(f"Invalid padding length: {padding_len}") + + # Validate msg_id via validation_state + expected_client = (direction == MessageDirection.CLIENT_TO_SERVER) + try: + validation_state.validate_msg_id(msg_id, expected_client=expected_client) + except Exception as e: + if "Duplicate" in str(e): + raise ReplayDetectedError(str(e)) + raise MTProtoFramingError(f"msg_id validation failed: {e}") + + body = decrypted[32:32+msg_len] + return EncryptedMessage( + server_salt=server_salt, + session_id=session_id, + msg_id=msg_id, + seq_no=seq_no, + body=body + ) diff --git a/src/crypto_standalone/mtproto/errors.py b/src/crypto_standalone/mtproto/errors.py new file mode 100644 index 0000000..db3ff93 --- /dev/null +++ b/src/crypto_standalone/mtproto/errors.py @@ -0,0 +1,36 @@ +"""MTProto exception hierarchy.""" + +class MTProtoFramingError(Exception): + """Base exception for all MTProto framing and parsing errors.""" + + +class IncompleteFrameError(MTProtoFramingError): + """Raised when more data is needed to parse a complete frame.""" + + +class InvalidFrameLengthError(MTProtoFramingError): + """Raised when a frame length is outside permitted limits or structurally invalid.""" + + +class FrameTooLargeError(MTProtoFramingError): + """Raised when a frame exceeds the configured hard limits.""" + + +class TransportSequenceError(MTProtoFramingError): + """Raised when a transport sequence number (e.g., full transport) is incorrect.""" + + +class TransportChecksumError(MTProtoFramingError): + """Raised when a transport checksum (e.g., CRC32) fails validation.""" + + +class MessageKeyMismatchError(MTProtoFramingError): + """Raised when an encrypted message's msg_key does not match the decrypted and verified payload.""" + + +class SessionMismatchError(MTProtoFramingError): + """Raised when a decrypted message's session_id does not match the expected session.""" + + +class ReplayDetectedError(MTProtoFramingError): + """Raised when a message is detected as a replay or duplicate.""" diff --git a/src/crypto_standalone/mtproto/session.py b/src/crypto_standalone/mtproto/session.py new file mode 100644 index 0000000..9f2e6a5 --- /dev/null +++ b/src/crypto_standalone/mtproto/session.py @@ -0,0 +1,68 @@ +import time +from typing import Set + +from .errors import ReplayDetectedError + +class SessionValidationState: + """Stateful validation for MTProto sessions (msg_id, seq_no).""" + + def __init__(self): + self._last_time = 0.0 + self._replay_cache: Set[int] = set() + self.max_replay_cache = 10000 + + def generate_msg_id(self, is_client: bool = True) -> int: + """ + Generates a 64-bit monotonically increasing message ID based on approx Unix time. + Fractional part is in the lower 32 bits. + Must be a multiple of 4, parity depends on direction. + Client -> Server: msg_id % 4 == 0 or 2 (content-related vs non) (we'll just use % 4 == 0) + """ + now = time.time() + if now <= self._last_time: + now = self._last_time + 0.0001 + self._last_time = now + + # msg_id is approx time in seconds * 2^32 + msg_id = int(now * (1 << 32)) + + # ensure parity + remainder = msg_id % 4 + if is_client: + # Client: remainder should be 0 + msg_id -= remainder + else: + # Server: remainder should be 1 or 3 + if remainder in (0, 2): + msg_id += 1 + + return msg_id + + def validate_msg_id(self, msg_id: int, expected_client: bool = False) -> None: + """ + Validates direction parity, time bounds, and replay. + """ + remainder = msg_id % 4 + if expected_client: + if remainder not in (0, 2): + raise ValueError("Invalid msg_id parity for client message") + else: + if remainder not in (1, 3): + raise ValueError("Invalid msg_id parity for server message") + + # Time window: 300s past, 30s future + now_ts = int(time.time()) + msg_ts = msg_id >> 32 + + if msg_ts < now_ts - 300: + raise ValueError("msg_id too old") + if msg_ts > now_ts + 30: + raise ValueError("msg_id too far in the future") + + if msg_id in self._replay_cache: + raise ReplayDetectedError("Duplicate msg_id detected") + + self._replay_cache.add(msg_id) + if len(self._replay_cache) > self.max_replay_cache: + # Simple eviction + self._replay_cache = set(list(self._replay_cache)[-self.max_replay_cache//2:]) diff --git a/src/crypto_standalone/mtproto/transports.py b/src/crypto_standalone/mtproto/transports.py new file mode 100644 index 0000000..a04d264 --- /dev/null +++ b/src/crypto_standalone/mtproto/transports.py @@ -0,0 +1,347 @@ +import binascii +import struct +from typing import List, Protocol +from .core import TransportEvent, PayloadFrame, QuickAck, TransportError, ResourceLimits, DEFAULT_LIMITS +from .errors import InvalidFrameLengthError, FrameTooLargeError, TransportSequenceError, TransportChecksumError + +class TransportCodec(Protocol): + def connection_header(self) -> bytes: + ... + + def encode( + self, + payload: bytes, + *, + request_quick_ack: bool = False, + ) -> bytes: + ... + + def feed_data(self, data: bytes) -> List[TransportEvent]: + ... + + def feed_eof(self) -> List[TransportEvent]: + ... + + def reset(self) -> None: + ... + + +class BaseTransportCodec: + def __init__(self, limits: ResourceLimits = DEFAULT_LIMITS): + self.limits = limits + self._buffer = bytearray() + self._offset = 0 + self._sent_header = False + + def reset(self) -> None: + self._buffer.clear() + self._offset = 0 + self._sent_header = False + + def _compact(self) -> None: + if self._offset > 0: + del self._buffer[:self._offset] + self._offset = 0 + if len(self._buffer) > self.limits.max_retained_incremental_buffer_size: + self.reset() + raise FrameTooLargeError("Incremental buffer exceeded maximum limit") + + def feed_data(self, data: bytes) -> List[TransportEvent]: + self._buffer.extend(data) + events = [] + while True: + evt = self._parse_one() + if evt is None: + break + events.append(evt) + self._compact() + return events + + def feed_eof(self) -> List[TransportEvent]: + events = [] + while True: + evt = self._parse_one() + if evt is None: + break + events.append(evt) + self._compact() + return events + + def _parse_one(self) -> TransportEvent | None: + raise NotImplementedError + + def _read_bytes(self, n: int) -> bytes | None: + if len(self._buffer) - self._offset >= n: + res = bytes(self._buffer[self._offset : self._offset + n]) + self._offset += n + return res + return None + + def _peek_bytes(self, n: int) -> bytes | None: + if len(self._buffer) - self._offset >= n: + return bytes(self._buffer[self._offset : self._offset + n]) + return None + + def _decode_quick_ack_or_error(self, first_int: int) -> TransportEvent | None: + # Check for Quick ACK (standalone reversed 4 bytes with MSB set) + # However, quick acks in Abridged and Intermediate are generally negative ints if treated as little endian + if first_int >= 0x80000000: + return QuickAck(token=first_int) + + # Check for error (standalone error frame) + # Note: the spec says "error is a signed little-endian number of 4 bytes, whose absolute value contains the error code (the error itself is actually negative)." + return None + +class AbridgedTransportCodec(BaseTransportCodec): + def connection_header(self) -> bytes: + if not self._sent_header: + self._sent_header = True + return b"\xef" + return b"" + + def encode(self, payload: bytes, *, request_quick_ack: bool = False) -> bytes: + length_words = len(payload) // 4 + if len(payload) % 4 != 0: + raise InvalidFrameLengthError("Payload length must be divisible by 4") + + header = bytearray() + if length_words < 127: + if request_quick_ack: + header.append(length_words | 0x80) + else: + header.append(length_words) + else: + header.append(0x7F | (0x80 if request_quick_ack else 0x00)) + header.extend(struct.pack(" TransportEvent | None: + # Check for Quick ACK (4 bytes standalone) + # Server sends quick ACK by bswapping them. The length/header of normal payload is always <= 127, + # so MSB is not set. If MSB is set, it's a Quick ACK. + + if len(self._buffer) - self._offset >= 4: + first_byte = self._buffer[self._offset] + if first_byte >= 128: + # 1 byte might be 0x7F | 0x80 = 0xFF (extended length + quick ack?) + # Actually, spec says: "quick ACK packets can be easily distinguished ... first byte will always have the most-significant bit set ... normal payload packets <= 127" + # Wait, what if the packet is an extended length abridged packet? + # "The server will send quick ACK tokens by bswapping them ... standalone 4-byte packet" + # Wait, server-to-client MTProto payload lengths shouldn't exceed 127 (508 bytes) for standard messages, but they can for larger ones. + # Is extended length 0x7f or 0xff? Spec: "0x7F plus three-byte little-endian extended length ... If length/4 >= 127, envelope: 0xff, length..." + # Wait! "first byte will always have most-significant bit set ... length/header of normal payload packets coming from the server is always less than or equal to 127 (thus the most-significant bit is not set for normal payloads)." + # Oh, it means server *responses* never exceed 127 length in *abridged* transport? Or server doesn't use extended length? + # Let's assume MSB = 1 -> Quick ACK. + pass + + if len(self._buffer) - self._offset < 1: + return None + + first_byte = self._buffer[self._offset] + + # Check for transport error. The server may send a transport error as a 4-byte signed int. + # But how to distinguish from abridged header? + # A transport error is 4 bytes. E.g., -404 (0x6C FE FF FF). First byte is 0x6C (108). + # We need to distinguish it. + # If it's a 4-byte payload, we'll see. + + if first_byte >= 128: + if first_byte != 0x7f and first_byte != 0xff: + # Quick ACK + if len(self._buffer) - self._offset >= 4: + token = struct.unpack(">I", self._peek_bytes(4))[0] # it's bswapped + self._offset += 4 + return QuickAck(token=token) + return None + + if first_byte == 0x7f or first_byte == 0xff: + if len(self._buffer) - self._offset < 4: + return None + length_words = struct.unpack("= 4: + # Just read 4 bytes as error code + val = struct.unpack(" -1000: + self._offset += 4 + return TransportError(code=abs(val)) + + if payload_len > self.limits.max_transport_frame_size: + raise FrameTooLargeError(f"Abridged frame payload size {payload_len} exceeds max {self.limits.max_transport_frame_size}") + + if len(self._buffer) - self._offset < header_len + payload_len: + return None + + self._offset += header_len + payload = self._read_bytes(payload_len) + return PayloadFrame(payload=payload) + + +class IntermediateTransportCodec(BaseTransportCodec): + def connection_header(self) -> bytes: + if not self._sent_header: + self._sent_header = True + return b"\xee\xee\xee\xee" + return b"" + + def encode(self, payload: bytes, *, request_quick_ack: bool = False) -> bytes: + payload_len = len(payload) + if payload_len % 4 != 0: + raise InvalidFrameLengthError("Payload length must be divisible by 4") + if request_quick_ack: + payload_len |= 0x80000000 + return struct.pack(" TransportEvent | None: + if len(self._buffer) - self._offset < 4: + return None + + length_val = struct.unpack(" -1000: + self._offset += 4 + return TransportError(code=abs(signed_val)) + + if length_val >= 0x80000000: + # Quick ACK + self._offset += 4 + return QuickAck(token=length_val) + + if length_val > self.limits.max_transport_frame_size: + raise FrameTooLargeError(f"Intermediate frame size {length_val} exceeds max") + + if len(self._buffer) - self._offset < 4 + length_val: + return None + + self._offset += 4 + payload = self._read_bytes(length_val) + return PayloadFrame(payload=payload) + + +class PaddedIntermediateTransportCodec(BaseTransportCodec): + def __init__(self, limits: ResourceLimits = DEFAULT_LIMITS): + super().__init__(limits) + + def connection_header(self) -> bytes: + if not self._sent_header: + self._sent_header = True + return b"\xdd\xdd\xdd\xdd" + return b"" + + def encode(self, payload: bytes, *, request_quick_ack: bool = False) -> bytes: + import os + padding_len = os.urandom(1)[0] % 16 + padding = os.urandom(padding_len) + total_len = len(payload) + padding_len + if request_quick_ack: + total_len |= 0x80000000 + return struct.pack(" TransportEvent | None: + if len(self._buffer) - self._offset < 4: + return None + + length_val = struct.unpack(" -1000: + self._offset += 4 + return TransportError(code=abs(signed_val)) + + if length_val > self.limits.max_transport_frame_size: + raise FrameTooLargeError(f"Padded intermediate frame size {length_val} exceeds max") + + if len(self._buffer) - self._offset < 4 + length_val: + return None + + self._offset += 4 + payload_with_padding = self._read_bytes(length_val) + return PayloadFrame(payload=payload_with_padding, transport_padding=b"") # Note: we don't strip padding here, up to higher layer + + +class FullTransportCodec(BaseTransportCodec): + def __init__(self, limits: ResourceLimits = DEFAULT_LIMITS): + super().__init__(limits) + self.out_seq_no = 0 + self.in_seq_no = 0 + + def connection_header(self) -> bytes: + self._sent_header = True + return b"" + + def encode(self, payload: bytes, *, request_quick_ack: bool = False) -> bytes: + if len(payload) % 4 != 0: + raise InvalidFrameLengthError("Payload length must be divisible by 4") + + # length: length + seqno + payload + crc + length = 4 + 4 + len(payload) + 4 + seq_no = self.out_seq_no + self.out_seq_no += 1 + + data_to_crc = struct.pack(" TransportEvent | None: + if len(self._buffer) - self._offset < 4: + return None + + length = struct.unpack(" -1000: + self._offset += 4 + return TransportError(code=abs(signed_val)) + + if length % 4 != 0 or length < 12: + raise InvalidFrameLengthError(f"Invalid full transport length: {length}") + + if length > self.limits.max_transport_frame_size: + raise FrameTooLargeError(f"Full frame size {length} exceeds max") + + if len(self._buffer) - self._offset < length: + return None + + data = self._peek_bytes(length) + + data_to_crc = data[:-4] + expected_crc = struct.unpack(" bytes: out.extend(_xor_block(block, keystream[: len(block)])) counter += 1 return bytes(out) + + def encrypt_ige(self, plaintext: bytes, iv: bytes) -> bytes: + if len(iv) != 32: + raise ValueError("IGE mode requires a 32-byte IV") + if len(plaintext) % 16 != 0: + raise ValueError("plaintext length must be a multiple of 16 bytes") + + iv1 = iv[:16] + iv2 = iv[16:] + + out = bytearray() + prev_c = iv1 + prev_p = iv2 + + for i in range(0, len(plaintext), 16): + p = plaintext[i : i + 16] + x = _xor_block(p, prev_c) + c_temp = self.encrypt_block(x) + c = _xor_block(c_temp, prev_p) + out.extend(c) + prev_p = p + prev_c = c + + return bytes(out) + + def decrypt_ige(self, ciphertext: bytes, iv: bytes) -> bytes: + if len(iv) != 32: + raise ValueError("IGE mode requires a 32-byte IV") + if len(ciphertext) % 16 != 0: + raise ValueError("ciphertext length must be a multiple of 16 bytes") + + iv1 = iv[:16] + iv2 = iv[16:] + + out = bytearray() + prev_c = iv1 + prev_p = iv2 + + for i in range(0, len(ciphertext), 16): + c = ciphertext[i : i + 16] + x = _xor_block(c, prev_p) + p_temp = self.decrypt_block(x) + p = _xor_block(p_temp, prev_c) + out.extend(p) + prev_p = p + prev_c = c + + return bytes(out) diff --git a/tests/unit/test_mtproto_containers.py b/tests/unit/test_mtproto_containers.py new file mode 100644 index 0000000..50523ee --- /dev/null +++ b/tests/unit/test_mtproto_containers.py @@ -0,0 +1,47 @@ +import pytest + +from crypto_standalone.mtproto.containers import ( + ContainerMessage, + encode_msg_container, + decode_msg_container +) +from crypto_standalone.mtproto.errors import MTProtoFramingError + +def test_msg_container_roundtrip(): + msg1 = ContainerMessage(msg_id=100, seq_no=1, body=b"A"*4) + msg2 = ContainerMessage(msg_id=101, seq_no=2, body=b"B"*8) + + enc = encode_msg_container([msg1, msg2]) + dec = decode_msg_container(enc) + + assert len(dec) == 2 + assert dec[0].msg_id == 100 + assert dec[1].msg_id == 101 + assert dec[0].body == b"A"*4 + +def test_msg_container_limits(): + messages = [ContainerMessage(msg_id=i, seq_no=i, body=b"A"*4) for i in range(1025)] + with pytest.raises(MTProtoFramingError): + encode_msg_container(messages) + +def test_msg_container_nesting_check(): + msg1 = ContainerMessage(msg_id=100, seq_no=1, body=b"A"*4) + enc_inner = encode_msg_container([msg1]) + + msg_outer = ContainerMessage(msg_id=101, seq_no=2, body=enc_inner) + enc_outer = encode_msg_container([msg_outer]) + + with pytest.raises(MTProtoFramingError, match="Nested simple containers are rejected"): + decode_msg_container(enc_outer) + +def test_msg_container_msg_id_strict_lower(): + msg1 = ContainerMessage(msg_id=100, seq_no=1, body=b"A"*4) + enc = encode_msg_container([msg1]) + + # If container_msg_id is provided, nested must be strictly lower + with pytest.raises(MTProtoFramingError): + decode_msg_container(enc, container_msg_id=100) # strictly lower means < + + # Should pass if container is higher + dec = decode_msg_container(enc, container_msg_id=101) + assert len(dec) == 1 diff --git a/tests/unit/test_mtproto_encrypted_messages.py b/tests/unit/test_mtproto_encrypted_messages.py new file mode 100644 index 0000000..fb55198 --- /dev/null +++ b/tests/unit/test_mtproto_encrypted_messages.py @@ -0,0 +1,137 @@ +import pytest +import os +import struct + +from crypto_standalone.mtproto.envelope import ( + encode_encrypted_message, + decode_encrypted_message, + encode_unencrypted_message, + decode_unencrypted_message +) +from crypto_standalone.mtproto.session import SessionValidationState +from crypto_standalone.mtproto.core import MessageDirection +from crypto_standalone.mtproto.errors import ( + InvalidFrameLengthError, + MessageKeyMismatchError, + SessionMismatchError, + MTProtoFramingError +) + +def test_unencrypted_message_roundtrip(): + body = b"\x01\x02\x03\x04" + msg_id = 123456789 + + enc = encode_unencrypted_message(msg_id=msg_id, body=body) + assert len(enc) == 20 + len(body) + + dec = decode_unencrypted_message(enc) + assert dec.msg_id == msg_id + assert dec.body == body + +def test_encrypted_message_roundtrip(): + auth_key = os.urandom(256) + server_salt = 11111111111 + session_id = 22222222222 + + val_state = SessionValidationState() + msg_id = val_state.generate_msg_id(is_client=True) + seq_no = 1 + body = b"hello world\x00" + + direction = MessageDirection.CLIENT_TO_SERVER + + enc = encode_encrypted_message( + auth_key=auth_key, + server_salt=server_salt, + session_id=session_id, + msg_id=msg_id, + seq_no=seq_no, + body=body, + direction=direction + ) + + # msg_id parity for CLIENT is checked during decoding + dec = decode_encrypted_message( + data=enc, + auth_key=auth_key, + expected_session_id=session_id, + direction=direction, + validation_state=val_state + ) + + assert dec.server_salt == server_salt + assert dec.session_id == session_id + assert dec.msg_id == msg_id + assert dec.seq_no == seq_no + assert dec.body == body + +def test_encrypted_message_invalid_msg_key(): + auth_key = os.urandom(256) + val_state = SessionValidationState() + enc = encode_encrypted_message( + auth_key=auth_key, + server_salt=1, + session_id=2, + msg_id=val_state.generate_msg_id(), + seq_no=1, + body=b"1234", + direction=MessageDirection.CLIENT_TO_SERVER + ) + + # modify msg_key + enc_modified = bytearray(enc) + enc_modified[8] ^= 0x01 + + with pytest.raises(MessageKeyMismatchError): + decode_encrypted_message( + data=bytes(enc_modified), + auth_key=auth_key, + expected_session_id=2, + direction=MessageDirection.CLIENT_TO_SERVER, + validation_state=val_state + ) + +def test_encrypted_message_invalid_session(): + auth_key = os.urandom(256) + val_state = SessionValidationState() + enc = encode_encrypted_message( + auth_key=auth_key, + server_salt=1, + session_id=2, + msg_id=val_state.generate_msg_id(), + seq_no=1, + body=b"1234", + direction=MessageDirection.CLIENT_TO_SERVER + ) + + with pytest.raises(SessionMismatchError): + decode_encrypted_message( + data=enc, + auth_key=auth_key, + expected_session_id=3, + direction=MessageDirection.CLIENT_TO_SERVER, + validation_state=val_state + ) + +def test_encrypted_message_padding_bounds(): + auth_key = os.urandom(256) + val_state = SessionValidationState() + body = b"A" * 1024 # will force large padding or large frame + enc = encode_encrypted_message( + auth_key=auth_key, + server_salt=1, + session_id=2, + msg_id=val_state.generate_msg_id(), + seq_no=1, + body=body, + direction=MessageDirection.CLIENT_TO_SERVER + ) + # verify it succeeds + dec = decode_encrypted_message( + data=enc, + auth_key=auth_key, + expected_session_id=2, + direction=MessageDirection.CLIENT_TO_SERVER, + validation_state=val_state + ) + assert dec.body == body diff --git a/tests/unit/test_mtproto_transports.py b/tests/unit/test_mtproto_transports.py new file mode 100644 index 0000000..834527d --- /dev/null +++ b/tests/unit/test_mtproto_transports.py @@ -0,0 +1,129 @@ +import pytest +import struct +import binascii + +from crypto_standalone.mtproto.transports import ( + AbridgedTransportCodec, + IntermediateTransportCodec, + PaddedIntermediateTransportCodec, + FullTransportCodec +) +from crypto_standalone.mtproto.core import PayloadFrame, QuickAck, TransportError +from crypto_standalone.mtproto.errors import FrameTooLargeError, InvalidFrameLengthError, TransportChecksumError, TransportSequenceError + + +def test_abridged_transport(): + codec = AbridgedTransportCodec() + assert codec.connection_header() == b"\xef" + assert codec.connection_header() == b"" + + payload = b"\x00" * 4 + enc = codec.encode(payload) + assert enc == b"\x01\x00\x00\x00\x00" + + events = codec.feed_data(enc) + assert len(events) == 1 + assert isinstance(events[0], PayloadFrame) + assert events[0].payload == payload + + # Extended length + payload2 = b"\x00" * 508 + enc2 = codec.encode(payload2) + assert enc2[0] == 0x7F + assert enc2[1:4] == struct.pack("I", token) + events3 = codec.feed_data(token_bytes) + assert len(events3) == 1 + assert isinstance(events3[0], QuickAck) + assert events3[0].token == token + + # Fragmented stream + for i in range(len(enc)): + evt = codec.feed_data(bytes([enc[i]])) + if i < len(enc) - 1: + assert len(evt) == 0 + else: + assert len(evt) == 1 + assert evt[0].payload == payload + + +def test_intermediate_transport(): + codec = IntermediateTransportCodec() + assert codec.connection_header() == b"\xee\xee\xee\xee" + + payload = b"\x00" * 4 + enc = codec.encode(payload) + assert enc == b"\x04\x00\x00\x00" + payload + + events = codec.feed_data(enc) + assert len(events) == 1 + assert events[0].payload == payload + + enc_quick_ack = codec.encode(payload, request_quick_ack=True) + assert enc_quick_ack[3] == 0x80 + + # Fragmented + for i in range(len(enc)): + evt = codec.feed_data(bytes([enc[i]])) + if i < len(enc) - 1: + assert len(evt) == 0 + else: + assert len(evt) == 1 + assert evt[0].payload == payload + + +def test_padded_intermediate_transport(): + codec = PaddedIntermediateTransportCodec() + assert codec.connection_header() == b"\xdd\xdd\xdd\xdd" + + payload = b"\x00" * 4 + enc = codec.encode(payload) + events = codec.feed_data(enc) + assert len(events) == 1 + assert events[0].payload.startswith(payload) + + # Error + err_enc = struct.pack("