89 lines
3.4 KiB
Python
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']
|