chore: generic fixes (#1187)

This commit is contained in:
delewer
2026-08-13 14:44:09 +03:00
committed by GitHub
parent b8a5634008
commit 0d459995f3
15 changed files with 608 additions and 18 deletions
View File
+116
View File
@@ -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()
+87
View File
@@ -0,0 +1,87 @@
import unittest
from proxy.config import (
CFPROXY_DEFAULT_DOMAINS,
_is_valid_domain,
_normalize_domain_pool,
coerce_domain_list,
parse_dc_ip_list,
)
class ParseDcIpListTest(unittest.TestCase):
def test_parses_multiple_entries(self):
self.assertEqual(
parse_dc_ip_list(['2:149.154.167.220', '4:1.2.3.4']),
{2: '149.154.167.220', 4: '1.2.3.4'},
)
def test_last_entry_wins_for_duplicate_dc(self):
self.assertEqual(parse_dc_ip_list(['2:1.1.1.1', '2:2.2.2.2']), {2: '2.2.2.2'})
def test_rejects_missing_separator(self):
with self.assertRaises(ValueError) as ctx:
parse_dc_ip_list(['2-1.2.3.4'])
self.assertEqual(ctx.exception.kind, 'format')
def test_rejects_short_form_ipv4(self):
for entry in ('2:149.154', '2:1.2.3.4.5', '2:999.1.1.1', '2:abc'):
with self.subTest(entry=entry), self.assertRaises(ValueError) as ctx:
parse_dc_ip_list([entry])
self.assertEqual(ctx.exception.kind, 'invalid')
def test_rejects_non_numeric_dc(self):
with self.assertRaises(ValueError):
parse_dc_ip_list(['x:1.2.3.4'])
class CoerceDomainListTest(unittest.TestCase):
def test_splits_on_common_separators(self):
self.assertEqual(
coerce_domain_list('a.com, b.com; c.com d.com'),
['a.com', 'b.com', 'c.com', 'd.com'],
)
def test_deduplicates_case_insensitively_keeping_first(self):
self.assertEqual(coerce_domain_list(['A.com', 'a.com']), ['A.com'])
def test_flattens_sequences_and_skips_non_strings(self):
self.assertEqual(coerce_domain_list(['a.com b.com', 5, None]), ['a.com', 'b.com'])
def test_returns_empty_for_unsupported_types(self):
self.assertEqual(coerce_domain_list(None), [])
self.assertEqual(coerce_domain_list(42), [])
class DomainValidationTest(unittest.TestCase):
def test_accepts_ordinary_domains(self):
for domain in ('example.com', 'a-b.co.uk', 'x.io'):
self.assertTrue(_is_valid_domain(domain), domain)
def test_rejects_malformed_domains(self):
for domain in ('', 'nodot', '.leading.com', 'trailing.com.',
'-bad.com', 'bad-.com', 'a..com', 'a.1',
'a.' + 'b' * 64, 'a' * 250 + '.com'):
self.assertFalse(_is_valid_domain(domain), domain)
def test_normalize_lowercases_dedupes_and_drops_invalid(self):
self.assertEqual(
_normalize_domain_pool(['B.com ', 'b.com', 'nodot', 'a.com']),
['b.com', 'a.com'],
)
class DefaultDomainsTest(unittest.TestCase):
def test_decoded_defaults_are_valid_domains(self):
self.assertTrue(CFPROXY_DEFAULT_DOMAINS)
for domain in CFPROXY_DEFAULT_DOMAINS:
self.assertTrue(_is_valid_domain(domain), domain)
def test_decoded_defaults_are_unique(self):
self.assertEqual(
len(set(CFPROXY_DEFAULT_DOMAINS)), len(CFPROXY_DEFAULT_DOMAINS)
)
if __name__ == '__main__':
unittest.main()
+133
View File
@@ -0,0 +1,133 @@
import hashlib
import hmac
import os
import struct
import time
import unittest
from proxy.fake_tls import (
CLIENT_RANDOM_LEN,
CLIENT_RANDOM_OFFSET,
SESSION_ID_LEN,
SESSION_ID_OFFSET,
TLS_APPDATA_MAX,
TLS_RECORD_HANDSHAKE,
build_server_hello,
verify_client_hello,
wrap_tls_record,
)
SECRET = bytes.fromhex('00112233445566778899aabbccddeeff')
def _client_hello(secret: bytes = SECRET, timestamp: int = None,
session_id: bytes = None) -> bytes:
if timestamp is None:
timestamp = int(time.time())
if session_id is None:
session_id = os.urandom(SESSION_ID_LEN)
body = bytearray(517)
body[0] = TLS_RECORD_HANDSHAKE
body[1:3] = b'\x03\x01'
struct.pack_into('>H', body, 3, len(body) - 5)
body[5] = 0x01
body[43] = 0x20
body[SESSION_ID_OFFSET:SESSION_ID_OFFSET + SESSION_ID_LEN] = session_id
digest = hmac.new(secret, bytes(body), hashlib.sha256).digest()
client_random = bytearray(digest[:CLIENT_RANDOM_LEN])
ts_bytes = struct.pack('<I', timestamp)
for i in range(4):
client_random[28 + i] = digest[28 + i] ^ ts_bytes[i]
body[CLIENT_RANDOM_OFFSET:CLIENT_RANDOM_OFFSET + CLIENT_RANDOM_LEN] = client_random
return bytes(body)
class VerifyClientHelloTest(unittest.TestCase):
def test_accepts_well_formed_hello(self):
session_id = os.urandom(SESSION_ID_LEN)
now = int(time.time())
result = verify_client_hello(_client_hello(timestamp=now,
session_id=session_id), SECRET)
self.assertIsNotNone(result)
client_random, got_session_id, ts = result
self.assertEqual(len(client_random), CLIENT_RANDOM_LEN)
self.assertEqual(got_session_id, session_id)
self.assertEqual(ts, now)
def test_rejects_wrong_secret(self):
other = bytes.fromhex('ffeeddccbbaa99887766554433221100')
self.assertIsNone(verify_client_hello(_client_hello(), other))
def test_rejects_stale_timestamp(self):
stale = int(time.time()) - 3600
self.assertIsNone(verify_client_hello(_client_hello(timestamp=stale), SECRET))
def test_rejects_tampered_body(self):
hello = bytearray(_client_hello())
hello[300] ^= 0xFF
self.assertIsNone(verify_client_hello(bytes(hello), SECRET))
def test_rejects_short_and_non_handshake_records(self):
self.assertIsNone(verify_client_hello(b'\x16\x03\x01\x00\x10', SECRET))
hello = bytearray(_client_hello())
hello[0] = 0x17
self.assertIsNone(verify_client_hello(bytes(hello), SECRET))
hello = bytearray(_client_hello())
hello[5] = 0x02
self.assertIsNone(verify_client_hello(bytes(hello), SECRET))
class BuildServerHelloTest(unittest.TestCase):
def test_echoes_session_id_and_binds_client_random(self):
session_id = os.urandom(SESSION_ID_LEN)
client_random = os.urandom(CLIENT_RANDOM_LEN)
response = build_server_hello(SECRET, client_random, session_id)
self.assertEqual(response[0], TLS_RECORD_HANDSHAKE)
self.assertEqual(
response[SESSION_ID_OFFSET:SESSION_ID_OFFSET + SESSION_ID_LEN],
session_id,
)
zeroed = bytearray(response)
zeroed[11:11 + 32] = b'\x00' * 32
expected = hmac.new(SECRET, client_random + bytes(zeroed),
hashlib.sha256).digest()
self.assertEqual(response[11:11 + 32], expected)
def test_padding_length_varies_between_calls(self):
sizes = {
len(build_server_hello(SECRET, os.urandom(32), os.urandom(32)))
for _ in range(20)
}
self.assertGreater(len(sizes), 1)
class WrapTlsRecordTest(unittest.TestCase):
def test_short_payload_becomes_one_record(self):
wrapped = wrap_tls_record(b'hello')
self.assertEqual(wrapped, b'\x17\x03\x03\x00\x05hello')
def test_long_payload_is_chunked_to_the_record_limit(self):
payload = os.urandom(TLS_APPDATA_MAX + 100)
wrapped = wrap_tls_record(payload)
offset = 0
chunks = []
while offset < len(wrapped):
length = struct.unpack('>H', wrapped[offset + 3:offset + 5])[0]
self.assertLessEqual(length, TLS_APPDATA_MAX)
chunks.append(wrapped[offset + 5:offset + 5 + length])
offset += 5 + length
self.assertEqual(len(chunks), 2)
self.assertEqual(b''.join(chunks), payload)
def test_empty_payload_produces_no_records(self):
self.assertEqual(wrap_tls_record(b''), b'')
if __name__ == '__main__':
unittest.main()
+130
View File
@@ -0,0 +1,130 @@
import asyncio
import unittest
from proxy.raw_websocket import RawWebSocket, WsHandshakeError, _xor_mask
def _raw_frame(opcode, data, fin=True):
b0 = (0x80 if fin else 0x00) | opcode
n = len(data)
if n < 126:
return bytes([b0, n]) + data
if n < 65536:
return bytes([b0, 126]) + n.to_bytes(2, 'big') + data
return bytes([b0, 127]) + n.to_bytes(8, 'big') + data
class _NullWriter:
def write(self, data):
pass
async def drain(self):
pass
def _recv(chunks, cls=RawWebSocket):
async def _run():
reader = asyncio.StreamReader()
for chunk in chunks:
reader.feed_data(chunk)
reader.feed_eof()
ws = cls(reader, _NullWriter())
return ws, await ws.recv()
return asyncio.run(_run())
class XorMaskTest(unittest.TestCase):
def test_roundtrip(self):
data = bytes(range(256)) * 3
mask = b'\x01\x02\x03\x04'
self.assertEqual(_xor_mask(_xor_mask(data, mask), mask), data)
def test_empty_payload(self):
self.assertEqual(_xor_mask(b'', b'\x01\x02\x03\x04'), b'')
class BuildFrameTest(unittest.TestCase):
def test_short_unmasked_frame(self):
self.assertEqual(
RawWebSocket._build_frame(RawWebSocket.OP_BINARY, b'abc'),
b'\x82\x03abc',
)
def test_extended_length_selects_16bit_header(self):
frame = RawWebSocket._build_frame(RawWebSocket.OP_BINARY, b'x' * 200)
self.assertEqual(frame[:2], b'\x82\x7e')
self.assertEqual(int.from_bytes(frame[2:4], 'big'), 200)
def test_masked_frame_sets_mask_bit_and_is_reversible(self):
payload = b'payload'
frame = RawWebSocket._build_frame(
RawWebSocket.OP_BINARY, payload, mask=True)
self.assertTrue(frame[1] & 0x80)
self.assertEqual(_xor_mask(frame[6:], frame[2:6]), payload)
class RecvTest(unittest.TestCase):
def test_returns_unfragmented_message(self):
_, msg = _recv([_raw_frame(RawWebSocket.OP_BINARY, b'hello')])
self.assertEqual(msg, b'hello')
def test_reassembles_fragmented_message(self):
_, msg = _recv([
_raw_frame(RawWebSocket.OP_BINARY, b'AAA', False),
_raw_frame(RawWebSocket.OP_CONT, b'BBB', False),
_raw_frame(RawWebSocket.OP_CONT, b'CCC', True),
])
self.assertEqual(msg, b'AAABBBCCC')
def test_control_frame_between_fragments_is_skipped(self):
_, msg = _recv([
_raw_frame(RawWebSocket.OP_BINARY, b'AAA', False),
_raw_frame(RawWebSocket.OP_PONG, b''),
_raw_frame(RawWebSocket.OP_CONT, b'BBB', True),
])
self.assertEqual(msg, b'AAABBB')
def test_close_frame_returns_none(self):
ws, msg = _recv([_raw_frame(RawWebSocket.OP_CLOSE, b'\x03\xe8')])
self.assertIsNone(msg)
self.assertTrue(ws._closed)
def test_oversized_frame_is_rejected_before_reading_payload(self):
header = bytes([0x82, 127]) + (1 << 40).to_bytes(8, 'big')
with self.assertRaises(ConnectionError):
_recv([header])
def test_reassembled_message_exceeding_limit_is_rejected(self):
class _Capped(RawWebSocket):
__slots__ = ()
MAX_MESSAGE_LEN = 1500
chunk = b'x' * 1024
with self.assertRaises(ConnectionError):
_recv([
_raw_frame(RawWebSocket.OP_BINARY, chunk, False),
_raw_frame(RawWebSocket.OP_CONT, chunk, False),
], cls=_Capped)
class ParseCloseTest(unittest.TestCase):
def test_known_code_gets_name(self):
code, reason = RawWebSocket._parse_close(b'\x03\xe8bye')
self.assertEqual(code, 1000)
self.assertIn('normal', reason)
def test_empty_payload(self):
self.assertEqual(RawWebSocket._parse_close(b''), (None, ''))
class HandshakeErrorTest(unittest.TestCase):
def test_redirect_status_codes(self):
for code in (301, 302, 303, 307, 308):
self.assertTrue(WsHandshakeError(code, '').is_redirect)
for code in (0, 200, 429, 502):
self.assertFalse(WsHandshakeError(code, '').is_redirect)
if __name__ == '__main__':
unittest.main()
+70
View File
@@ -0,0 +1,70 @@
import unittest
from utils.update_check import _extract_assets, _parse_version_tuple, _version_gt
class ParseVersionTupleTest(unittest.TestCase):
def test_plain_and_prefixed_versions(self):
self.assertEqual(_parse_version_tuple('1.9.1'), (1, 9, 1))
self.assertEqual(_parse_version_tuple('v1.9.1'), (1, 9, 1))
self.assertEqual(_parse_version_tuple(' V2.0 '), (2, 0))
def test_trailing_suffixes_are_truncated_to_digits(self):
self.assertEqual(_parse_version_tuple('1.9.1rc2'), (1, 9, 1))
self.assertEqual(_parse_version_tuple('1.9.1-beta'), (1, 9, 1))
def test_empty_and_non_numeric_segments(self):
self.assertEqual(_parse_version_tuple(''), (0,))
self.assertEqual(_parse_version_tuple(None), (0,))
self.assertEqual(_parse_version_tuple('1.x.3'), (1, 0, 3))
class VersionGtTest(unittest.TestCase):
def test_newer_versions(self):
self.assertTrue(_version_gt('1.9.2', '1.9.1'))
self.assertTrue(_version_gt('1.10.0', '1.9.9'))
self.assertTrue(_version_gt('2.0', '1.9.9'))
def test_equal_and_older_versions(self):
self.assertFalse(_version_gt('1.9.1', '1.9.1'))
self.assertFalse(_version_gt('1.9.1', '1.9.2'))
self.assertFalse(_version_gt('1.9', '1.9.0'))
def test_shorter_version_is_padded_with_zeros(self):
self.assertTrue(_version_gt('1.9.1', '1.9'))
self.assertFalse(_version_gt('1.9', '1.9.1'))
class ExtractAssetsTest(unittest.TestCase):
def test_keeps_name_url_and_digest(self):
data = {'assets': [{
'name': 'TgWsProxy_windows.exe',
'browser_download_url': 'https://example.invalid/a.exe',
'digest': 'sha256:abc',
}]}
self.assertEqual(_extract_assets(data), [{
'name': 'TgWsProxy_windows.exe',
'url': 'https://example.invalid/a.exe',
'digest': 'sha256:abc',
}])
def test_drops_entries_without_name_or_url(self):
data = {'assets': [
{'name': 'a.exe'},
{'browser_download_url': 'https://example.invalid/b.exe'},
]}
self.assertEqual(_extract_assets(data), [])
def test_missing_digest_becomes_empty_string(self):
data = {'assets': [{
'name': 'a.exe', 'browser_download_url': 'https://example.invalid/a.exe',
}]}
self.assertEqual(_extract_assets(data)[0]['digest'], '')
def test_empty_input(self):
self.assertEqual(_extract_assets(None), [])
self.assertEqual(_extract_assets({}), [])
if __name__ == '__main__':
unittest.main()