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:
@@ -37,6 +37,7 @@ RUN apt-get update \
|
||||
WORKDIR /app
|
||||
COPY --from=builder /opt/venv /opt/venv
|
||||
COPY proxy ./proxy
|
||||
COPY utils ./utils
|
||||
COPY docs/README.md LICENSE ./
|
||||
|
||||
USER app
|
||||
|
||||
+15
-1
@@ -37,12 +37,26 @@ pip install -e .
|
||||
|
||||
Подробности: `docs/BuildFromSource.md`.
|
||||
|
||||
## Проверки
|
||||
|
||||
Тесты используют только стандартную библиотеку, дополнительные зависимости не нужны:
|
||||
|
||||
```bash
|
||||
python -m unittest discover -s tests -t .
|
||||
```
|
||||
|
||||
Линтер (`ruff` настроен в `pyproject.toml`):
|
||||
|
||||
```bash
|
||||
ruff check .
|
||||
```
|
||||
|
||||
## Pull Request
|
||||
|
||||
Перед открытием PR:
|
||||
|
||||
1. Убедитесь, что изменение решает конкретную проблему.
|
||||
2. Проверьте, что не сломаны существующие сценарии.
|
||||
2. Проверьте, что не сломаны существующие сценарии: запустите тесты и линтер.
|
||||
3. Обновите документацию, если меняется поведение или настройка.
|
||||
|
||||
Небольшие и сфокусированные PR проверяются и принимаются быстрее.
|
||||
|
||||
+15
-1
@@ -37,12 +37,26 @@ Running:
|
||||
|
||||
Details: `docs/BuildFromSource.md`.
|
||||
|
||||
## Checks
|
||||
|
||||
Tests use the standard library only, no extra dependencies:
|
||||
|
||||
```bash
|
||||
python -m unittest discover -s tests -t .
|
||||
```
|
||||
|
||||
Linting (`ruff` is configured in `pyproject.toml`):
|
||||
|
||||
```bash
|
||||
ruff check .
|
||||
```
|
||||
|
||||
## Pull Request
|
||||
|
||||
Before opening a PR:
|
||||
|
||||
1. Make sure your change solves a specific problem.
|
||||
2. Check that existing scenarios aren't broken.
|
||||
2. Check that existing scenarios aren't broken; run the tests and the linter.
|
||||
3. Update documentation if behavior or configuration changes.
|
||||
|
||||
Smaller and focused PRs are reviewed and accepted faster.
|
||||
|
||||
@@ -202,6 +202,7 @@ async def _cfproxy_worker_fallback(reader, writer, relay_init, label,
|
||||
try:
|
||||
ws = await RawWebSocket.connect(worker_domain, worker_domain,
|
||||
timeout=10.0, path=path)
|
||||
break
|
||||
except Exception as exc:
|
||||
cf_worker_pool.report_failure(worker_domain, exc)
|
||||
log.warning("[%s] DC%d%s CF worker %s failed: %s",
|
||||
|
||||
+2
-2
@@ -209,11 +209,11 @@ def parse_dc_ip_list(dc_ip_list: List[str]) -> Dict[int, str]:
|
||||
dc_s, ip_s = entry.split(':', 1)
|
||||
try:
|
||||
dc_n = int(dc_s)
|
||||
_socket.inet_aton(ip_s)
|
||||
_socket.inet_pton(_socket.AF_INET, ip_s)
|
||||
except (ValueError, OSError):
|
||||
err = ValueError(f"Invalid --dc-ip {entry!r}")
|
||||
err.entry = entry
|
||||
err.kind = "invalid"
|
||||
raise err
|
||||
raise err from None
|
||||
dc_redirects[dc_n] = ip_s
|
||||
return dc_redirects
|
||||
|
||||
+1
-1
@@ -296,7 +296,7 @@ class _CfWorkerPool:
|
||||
|
||||
def available_domains(self, worker_domains: List[str]) -> List[str]:
|
||||
now = time.time()
|
||||
domains = list()
|
||||
domains = []
|
||||
for domain in worker_domains:
|
||||
if domain in domains:
|
||||
continue
|
||||
|
||||
+23
-6
@@ -67,18 +67,22 @@ def set_sock_opts(transport, buffer_size):
|
||||
|
||||
|
||||
class RawWebSocket:
|
||||
__slots__ = ('reader', 'writer', '_closed')
|
||||
__slots__ = ('reader', 'writer', '_closed', '_frag')
|
||||
|
||||
OP_CONT = 0x0
|
||||
OP_BINARY = 0x2
|
||||
OP_CLOSE = 0x8
|
||||
OP_PING = 0x9
|
||||
OP_PONG = 0xA
|
||||
|
||||
MAX_MESSAGE_LEN = 16 * 1024 * 1024
|
||||
|
||||
def __init__(self, reader: asyncio.StreamReader,
|
||||
writer: asyncio.StreamWriter):
|
||||
self.reader = reader
|
||||
self.writer = writer
|
||||
self._closed = False
|
||||
self._frag = bytearray()
|
||||
|
||||
@staticmethod
|
||||
async def connect(host: str, domain: str, timeout: float = 10.0,
|
||||
@@ -164,7 +168,7 @@ class RawWebSocket:
|
||||
|
||||
async def recv(self) -> Optional[bytes]:
|
||||
while not self._closed:
|
||||
opcode, payload = await self._read_frame()
|
||||
opcode, payload, fin = await self._read_frame()
|
||||
|
||||
if opcode == self.OP_CLOSE:
|
||||
self._closed = True
|
||||
@@ -192,8 +196,18 @@ class RawWebSocket:
|
||||
if opcode == self.OP_PONG:
|
||||
continue
|
||||
|
||||
if opcode in (0x1, 0x2):
|
||||
if opcode in (self.OP_CONT, 0x1, self.OP_BINARY):
|
||||
if fin and not self._frag:
|
||||
return payload
|
||||
self._frag.extend(payload)
|
||||
if len(self._frag) > self.MAX_MESSAGE_LEN:
|
||||
raise ConnectionError(
|
||||
f"WS message too large: {len(self._frag)} bytes")
|
||||
if not fin:
|
||||
continue
|
||||
message = bytes(self._frag)
|
||||
self._frag.clear()
|
||||
return message
|
||||
continue
|
||||
return None
|
||||
|
||||
@@ -251,17 +265,20 @@ class RawWebSocket:
|
||||
return _st_BBH4s.pack(fb, 0x80 | 126, length, mask_key) + masked
|
||||
return _st_BBQ4s.pack(fb, 0x80 | 127, length, mask_key) + masked
|
||||
|
||||
async def _read_frame(self) -> Tuple[int, bytes]:
|
||||
async def _read_frame(self) -> Tuple[int, bytes, bool]:
|
||||
hdr = await self.reader.readexactly(2)
|
||||
fin = bool(hdr[0] & 0x80)
|
||||
opcode = hdr[0] & 0x0F
|
||||
length = hdr[1] & 0x7F
|
||||
if length == 126:
|
||||
length = _st_H.unpack(await self.reader.readexactly(2))[0]
|
||||
elif length == 127:
|
||||
length = _st_Q.unpack(await self.reader.readexactly(8))[0]
|
||||
if length > self.MAX_MESSAGE_LEN:
|
||||
raise ConnectionError(f"WS frame too large: {length} bytes")
|
||||
if hdr[1] & 0x80:
|
||||
mask_key = await self.reader.readexactly(4)
|
||||
payload = await self.reader.readexactly(length)
|
||||
return opcode, _xor_mask(payload, mask_key)
|
||||
return opcode, _xor_mask(payload, mask_key), fin
|
||||
payload = await self.reader.readexactly(length)
|
||||
return opcode, payload
|
||||
return opcode, payload, fin
|
||||
+8
-1
@@ -73,5 +73,12 @@ packages = ["proxy", "ui", "utils"]
|
||||
[tool.hatch.version]
|
||||
path = "proxy/__init__.py"
|
||||
|
||||
[tool.ruff]
|
||||
target-version = "py38"
|
||||
|
||||
[tool.ruff.lint]
|
||||
ignore = ["F403", "F405"]
|
||||
select = ["E4", "E7", "E9", "F", "B", "C4"]
|
||||
ignore = ["F403", "F405", "B023"]
|
||||
|
||||
[tool.ruff.lint.per-file-ignores]
|
||||
"macos.py" = ["E402"]
|
||||
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
+5
-5
@@ -252,11 +252,11 @@ def _sync_language_combobox(combo: Any, var: Any, cfg_value: str) -> None:
|
||||
|
||||
|
||||
def _entry(ctk, parent, theme, *, var=None, width=0, height=36, radius=10, **kw):
|
||||
opts = dict(
|
||||
font=(theme.ui_font_family, 13), corner_radius=radius,
|
||||
fg_color=theme.bg, border_color=theme.field_border,
|
||||
border_width=1, text_color=theme.text_primary,
|
||||
)
|
||||
opts = {
|
||||
"font": (theme.ui_font_family, 13), "corner_radius": radius,
|
||||
"fg_color": theme.bg, "border_color": theme.field_border,
|
||||
"border_width": 1, "text_color": theme.text_primary,
|
||||
}
|
||||
if var is not None:
|
||||
opts["textvariable"] = var
|
||||
if width:
|
||||
|
||||
Reference in New Issue
Block a user