mirror of
https://github.com/Flowseal/tg-ws-proxy.git
synced 2026-09-05 18:16:11 +03:00
chore: generic fixes (#1187)
This commit is contained in:
@@ -0,0 +1,116 @@
|
||||
import os
|
||||
import unittest
|
||||
|
||||
from proxy._aes import Cipher, algorithms, modes
|
||||
from proxy.bridge import MsgSplitter
|
||||
from proxy.utils import (
|
||||
PROTO_ABRIDGED_INT,
|
||||
PROTO_INTERMEDIATE_INT,
|
||||
PROTO_PADDED_INTERMEDIATE_INT,
|
||||
)
|
||||
|
||||
|
||||
def _relay_init() -> bytes:
|
||||
return os.urandom(64)
|
||||
|
||||
|
||||
def _encryptor(relay_init: bytes):
|
||||
enc = Cipher(
|
||||
algorithms.AES(relay_init[8:40]), modes.CTR(relay_init[40:56])
|
||||
).encryptor()
|
||||
enc.update(b'\x00' * 64)
|
||||
return enc
|
||||
|
||||
|
||||
def _abridged(payload: bytes) -> bytes:
|
||||
words = len(payload) // 4
|
||||
if words < 0x7F:
|
||||
return bytes([words]) + payload
|
||||
return b'\x7f' + words.to_bytes(3, 'little') + payload
|
||||
|
||||
|
||||
def _intermediate(payload: bytes) -> bytes:
|
||||
return len(payload).to_bytes(4, 'little') + payload
|
||||
|
||||
|
||||
class MsgSplitterTest(unittest.TestCase):
|
||||
def _split(self, proto_int, packets, chunk_sizes=None):
|
||||
relay_init = _relay_init()
|
||||
splitter = MsgSplitter(relay_init, proto_int)
|
||||
enc = _encryptor(relay_init)
|
||||
stream = enc.update(b''.join(packets))
|
||||
|
||||
chunks = []
|
||||
if chunk_sizes is None:
|
||||
chunks = [stream]
|
||||
else:
|
||||
offset = 0
|
||||
for size in chunk_sizes:
|
||||
chunks.append(stream[offset:offset + size])
|
||||
offset += size
|
||||
if offset < len(stream):
|
||||
chunks.append(stream[offset:])
|
||||
|
||||
parts = []
|
||||
for chunk in chunks:
|
||||
parts.extend(splitter.split(chunk))
|
||||
return splitter, stream, parts
|
||||
|
||||
def test_abridged_stream_splits_into_packets(self):
|
||||
packets = [_abridged(b'a' * 4), _abridged(b'b' * 16), _abridged(b'c' * 40)]
|
||||
_, stream, parts = self._split(PROTO_ABRIDGED_INT, packets)
|
||||
self.assertEqual(len(parts), 3)
|
||||
self.assertEqual(b''.join(parts), stream)
|
||||
self.assertEqual([len(p) for p in parts], [5, 17, 41])
|
||||
|
||||
def test_intermediate_stream_splits_into_packets(self):
|
||||
packets = [_intermediate(b'a' * 8), _intermediate(b'b' * 12)]
|
||||
_, stream, parts = self._split(PROTO_INTERMEDIATE_INT, packets)
|
||||
self.assertEqual(len(parts), 2)
|
||||
self.assertEqual(b''.join(parts), stream)
|
||||
|
||||
def test_padded_intermediate_uses_intermediate_framing(self):
|
||||
packets = [_intermediate(b'z' * 20)]
|
||||
_, stream, parts = self._split(PROTO_PADDED_INTERMEDIATE_INT, packets)
|
||||
self.assertEqual(parts, [stream])
|
||||
|
||||
def test_partial_packet_is_buffered_until_complete(self):
|
||||
packets = [_abridged(b'a' * 20)]
|
||||
_, stream, parts = self._split(
|
||||
PROTO_ABRIDGED_INT, packets, chunk_sizes=[1] * (len(packets[0]) - 1)
|
||||
)
|
||||
self.assertEqual(parts, [stream])
|
||||
|
||||
def test_split_preserves_stream_across_arbitrary_chunking(self):
|
||||
packets = [_intermediate(bytes([i]) * 16) for i in range(8)]
|
||||
_, stream, parts = self._split(
|
||||
PROTO_INTERMEDIATE_INT, packets, chunk_sizes=[7, 3, 50, 11]
|
||||
)
|
||||
self.assertEqual(b''.join(parts), stream)
|
||||
self.assertEqual(len(parts), 8)
|
||||
|
||||
def test_empty_chunk_yields_nothing(self):
|
||||
splitter = MsgSplitter(_relay_init(), PROTO_INTERMEDIATE_INT)
|
||||
self.assertEqual(splitter.split(b''), [])
|
||||
|
||||
def test_zero_length_packet_disables_splitting(self):
|
||||
relay_init = _relay_init()
|
||||
splitter = MsgSplitter(relay_init, PROTO_INTERMEDIATE_INT)
|
||||
enc = _encryptor(relay_init)
|
||||
stream = enc.update((0).to_bytes(4, 'little') + b'tail')
|
||||
parts = splitter.split(stream)
|
||||
self.assertEqual(parts, [stream])
|
||||
self.assertEqual(splitter.split(b'raw'), [b'raw'])
|
||||
|
||||
def test_flush_returns_buffered_tail_once(self):
|
||||
relay_init = _relay_init()
|
||||
splitter = MsgSplitter(relay_init, PROTO_INTERMEDIATE_INT)
|
||||
enc = _encryptor(relay_init)
|
||||
partial = enc.update(_intermediate(b'x' * 32)[:10])
|
||||
self.assertEqual(splitter.split(partial), [])
|
||||
self.assertEqual(splitter.flush(), [partial])
|
||||
self.assertEqual(splitter.flush(), [])
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user