Files
matrix-screen-controller/核桃派软件源代码/tests/test_mobile_protocol.py
T

89 lines
3.4 KiB
Python

import json
from pathlib import Path
import pytest
from app.mobile.protocol import Handshake, ProtocolError, Reassembler, fragments, json_object
from cryptography.hazmat.primitives.asymmetric import ec
@pytest.mark.parametrize("mtu", [23, 64, 247, 517])
def test_fragments_roundtrip_and_replay(mtu):
data = ("奇妙小屏幕" * 200).encode()
parts = list(fragments(3, 0, data, mtu))
receiver = Reassembler()
for part in parts[:-1]:
assert receiver.accept(part) is None
assert receiver.accept(parts[-1]) == (3, data)
with pytest.raises(ProtocolError):
receiver.accept(parts[0])
def test_fragment_reordering_and_deadline():
parts = list(fragments(1, 0, b"hello world", 23))
with pytest.raises(ProtocolError):
Reassembler().accept(parts[1])
now = [0.0]
receiver = Reassembler(lambda: now[0])
receiver.accept(parts[0])
now[0] = 15
with pytest.raises(ProtocolError, match="TIMEOUT"):
receiver.accept(parts[1])
def pair():
client, server = Handshake(), Handshake()
return client.finish(server.hello, server=False), server.finish(client.hello, server=True)
def test_bidirectional_encryption_and_replay():
client, server = pair()
wire = client.encrypt('你好'.encode())
assert server.decrypt(wire) == '你好'.encode()
assert client.decrypt(server.encrypt(b'reply')) == b'reply'
with pytest.raises(ProtocolError):
server.decrypt(wire)
with pytest.raises(ProtocolError):
server.decrypt(client.encrypt(b'next'))
def test_tamper_direction_and_connection_isolation():
client, server = pair()
wire = client.encrypt(b'command')
with pytest.raises(ProtocolError):
client.decrypt(wire)
with pytest.raises(ProtocolError):
server.decrypt(wire[:-1] + bytes([wire[-1] ^ 1]))
_, other = pair()
with pytest.raises(ProtocolError):
other.decrypt(wire)
def test_invalid_hello_and_json():
handshake = Handshake()
hello = json.loads(handshake.hello)
hello['protocol_major'] = 2
with pytest.raises(ProtocolError, match='UNSUPPORTED_VERSION'):
handshake.finish(json.dumps(hello).encode(), server=True)
for raw in (b'{"a":1,"a":2}', b'{"a":NaN}', b'[]', b'\xff'):
with pytest.raises(ProtocolError):
json_object(raw)
def test_shared_golden_vector():
path = Path(__file__).resolve().parents[2] / '移动端相关内容/安卓app/安卓程序源代码/sharedCore/src/jvmTest/resources/protocol-v1.json'
vector = json.loads(path.read_text(encoding='utf-8'))
def raw(key):
return bytes.fromhex(vector[key])
client = Handshake(private_key=ec.derive_private_key(int(vector['client_scalar']), ec.SECP256R1()), random_bytes=raw('client_random'))
server = Handshake(private_key=ec.derive_private_key(int(vector['server_scalar']), ec.SECP256R1()), random_bytes=raw('server_random'))
assert client.hello == raw('client_hello')
assert server.hello == raw('server_hello')
cc, sc = client.finish(server.hello, server=False), server.finish(client.hello, server=True)
assert cc.transcript == raw('transcript')
assert cc.encrypt(raw('request')) == raw('client_record')
assert sc.encrypt(raw('response')) == raw('server_record')
assert sc.decrypt(raw('client_record')) == raw('request')
assert cc.decrypt(raw('server_record')) == raw('response')
assert [p.hex() for p in fragments(3, 0, raw('client_record'), 23)] == vector['fragments_mtu23']