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