mirror of
https://github.com/jkingsman/Remote-Terminal-for-MeshCore.git
synced 2026-08-07 01:03:34 +02:00
Initial tcp proxy testing
This commit is contained in:
@@ -0,0 +1,204 @@
|
||||
"""Tests for app.tcp_proxy.encoder — binary payload builders."""
|
||||
|
||||
import struct
|
||||
|
||||
from app.tcp_proxy.encoder import (
|
||||
build_contact,
|
||||
build_contact_from_dict,
|
||||
build_device_info,
|
||||
build_self_info,
|
||||
build_self_info_from_runtime,
|
||||
)
|
||||
from app.tcp_proxy.protocol import (
|
||||
PROXY_FW_VER,
|
||||
PROXY_MAX_CHANNELS,
|
||||
PROXY_MAX_CONTACTS_RAW,
|
||||
PUSH_NEW_ADVERT,
|
||||
RESP_CONTACT,
|
||||
RESP_DEVICE_INFO,
|
||||
RESP_SELF_INFO,
|
||||
)
|
||||
|
||||
EXAMPLE_KEY = "ab" * 32 # 64-char hex → 32 bytes
|
||||
|
||||
|
||||
# ── build_contact ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildContact:
|
||||
def test_basic_structure(self):
|
||||
payload = build_contact(EXAMPLE_KEY, name="Alice")
|
||||
assert payload[0] == RESP_CONTACT
|
||||
# public key at bytes 1-32
|
||||
assert payload[1:33] == bytes.fromhex(EXAMPLE_KEY)
|
||||
# total length: 1 + 32 + 1(type) + 1(flags) + 1(path) + 64(path) + 32(name) + 4(adv) + 4(lat) + 4(lon) + 4(lastmod) = 148
|
||||
assert len(payload) == 148
|
||||
|
||||
def test_push_variant(self):
|
||||
payload = build_contact(EXAMPLE_KEY, push=True)
|
||||
assert payload[0] == PUSH_NEW_ADVERT
|
||||
assert len(payload) == 148
|
||||
|
||||
def test_favorite_flag(self):
|
||||
payload = build_contact(EXAMPLE_KEY, favorite=True)
|
||||
flags_byte = payload[34] # byte 1+32+1 = 34
|
||||
assert flags_byte & 0x01 == 1
|
||||
|
||||
def test_not_favorite(self):
|
||||
payload = build_contact(EXAMPLE_KEY, favorite=False)
|
||||
flags_byte = payload[34]
|
||||
assert flags_byte & 0x01 == 0
|
||||
|
||||
def test_flood_path(self):
|
||||
payload = build_contact(EXAMPLE_KEY)
|
||||
path_byte = payload[35] # byte 1+32+1+1 = 35
|
||||
assert path_byte == 0xFF
|
||||
|
||||
def test_direct_path(self):
|
||||
payload = build_contact(
|
||||
EXAMPLE_KEY,
|
||||
direct_path="aabb",
|
||||
direct_path_len=2,
|
||||
direct_path_hash_mode=1,
|
||||
)
|
||||
path_byte = payload[35]
|
||||
# mode=1 → 0x40, hops=2 → 0x02 → packed = 0x42
|
||||
assert path_byte == 0x42
|
||||
|
||||
def test_name_truncated(self):
|
||||
long_name = "A" * 50
|
||||
payload = build_contact(EXAMPLE_KEY, name=long_name)
|
||||
# name field is 32 bytes at offset 100 (1+32+1+1+1+64)
|
||||
name_bytes = payload[100:132]
|
||||
assert name_bytes == b"A" * 32
|
||||
|
||||
def test_lat_lon_encoding(self):
|
||||
payload = build_contact(EXAMPLE_KEY, lat=45.123456, lon=-122.654321)
|
||||
lat_offset = 136 # 1+32+1+1+1+64+32+4 = 136
|
||||
lat = struct.unpack_from("<i", payload, lat_offset)[0]
|
||||
lon = struct.unpack_from("<i", payload, lat_offset + 4)[0]
|
||||
assert abs(lat - 45123456) < 2
|
||||
assert abs(lon - (-122654321)) < 2
|
||||
|
||||
def test_contact_type(self):
|
||||
payload = build_contact(EXAMPLE_KEY, contact_type=2)
|
||||
assert payload[33] == 2 # type byte at offset 1+32
|
||||
|
||||
|
||||
# ── build_contact_from_dict ──────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildContactFromDict:
|
||||
def test_minimal_dict(self):
|
||||
data = {"public_key": EXAMPLE_KEY}
|
||||
payload = build_contact_from_dict(data)
|
||||
assert payload[0] == RESP_CONTACT
|
||||
assert len(payload) == 148
|
||||
|
||||
def test_full_dict(self):
|
||||
data = {
|
||||
"public_key": EXAMPLE_KEY,
|
||||
"type": 1,
|
||||
"favorite": True,
|
||||
"name": "Bob",
|
||||
"direct_path": "ff",
|
||||
"direct_path_len": 1,
|
||||
"direct_path_hash_mode": 0,
|
||||
"last_advert": 1700000000,
|
||||
"lat": 37.7749,
|
||||
"lon": -122.4194,
|
||||
"first_seen": 1699000000,
|
||||
}
|
||||
payload = build_contact_from_dict(data)
|
||||
assert payload[33] == 1 # type
|
||||
assert payload[34] & 0x01 == 1 # favorite
|
||||
|
||||
def test_push_flag(self):
|
||||
data = {"public_key": EXAMPLE_KEY}
|
||||
payload = build_contact_from_dict(data, push=True)
|
||||
assert payload[0] == PUSH_NEW_ADVERT
|
||||
|
||||
|
||||
# ── build_self_info ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildSelfInfo:
|
||||
def test_basic_structure(self):
|
||||
payload = build_self_info()
|
||||
assert payload[0] == RESP_SELF_INFO
|
||||
assert payload[1] == 1 # adv_type = CHAT
|
||||
# minimum length: 1+1+1+1+32+4+4+1+1+1+1+4+4+1+1 + len("RemoteTerm") = 68
|
||||
assert len(payload) >= 58
|
||||
|
||||
def test_name_appended(self):
|
||||
payload = build_self_info(name="TestNode")
|
||||
# name starts at offset 58
|
||||
name_bytes = payload[58:]
|
||||
assert name_bytes == b"TestNode"
|
||||
|
||||
def test_public_key_encoded(self):
|
||||
payload = build_self_info(public_key=EXAMPLE_KEY)
|
||||
assert payload[4:36] == bytes.fromhex(EXAMPLE_KEY)
|
||||
|
||||
def test_radio_params(self):
|
||||
payload = build_self_info(radio_freq=868.0, radio_bw=125.0, radio_sf=12, radio_cr=8)
|
||||
freq = struct.unpack_from("<I", payload, 48)[0]
|
||||
bw = struct.unpack_from("<I", payload, 52)[0]
|
||||
assert freq == 868000
|
||||
assert bw == 125000
|
||||
assert payload[56] == 12 # sf
|
||||
assert payload[57] == 8 # cr
|
||||
|
||||
def test_multi_acks_flag(self):
|
||||
on = build_self_info(multi_acks=True)
|
||||
off = build_self_info(multi_acks=False)
|
||||
assert on[44] == 1
|
||||
assert off[44] == 0
|
||||
|
||||
|
||||
class TestBuildSelfInfoFromRuntime:
|
||||
def test_from_self_info_dict(self):
|
||||
info = {
|
||||
"public_key": EXAMPLE_KEY,
|
||||
"name": "MyRadio",
|
||||
"tx_power": 18,
|
||||
"max_tx_power": 22,
|
||||
"adv_lat": 40.0,
|
||||
"adv_lon": -74.0,
|
||||
"multi_acks": 1,
|
||||
"adv_loc_policy": 1,
|
||||
"radio_freq": 915.0,
|
||||
"radio_bw": 250.0,
|
||||
"radio_sf": 10,
|
||||
"radio_cr": 7,
|
||||
}
|
||||
payload = build_self_info_from_runtime(info)
|
||||
assert payload[0] == RESP_SELF_INFO
|
||||
assert payload[58:] == b"MyRadio"
|
||||
|
||||
def test_missing_fields_use_defaults(self):
|
||||
payload = build_self_info_from_runtime({})
|
||||
assert payload[0] == RESP_SELF_INFO
|
||||
assert payload[58:] == b"RemoteTerm"
|
||||
|
||||
|
||||
# ── build_device_info ────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildDeviceInfo:
|
||||
def test_basic_structure(self):
|
||||
payload = build_device_info()
|
||||
assert payload[0] == RESP_DEVICE_INFO
|
||||
assert payload[1] == PROXY_FW_VER
|
||||
assert payload[2] == PROXY_MAX_CONTACTS_RAW
|
||||
assert payload[3] == PROXY_MAX_CHANNELS
|
||||
|
||||
def test_path_hash_mode(self):
|
||||
payload = build_device_info(path_hash_mode=2)
|
||||
# path_hash_mode is at offset 81 (1+1+1+1+4+12+40+20+1 = 81)
|
||||
assert payload[81] == 2
|
||||
|
||||
def test_expected_length(self):
|
||||
# fw_ver=11 → 1+1+1+1+4+12+40+20+1+1 = 82 bytes
|
||||
payload = build_device_info()
|
||||
assert len(payload) == 82
|
||||
@@ -0,0 +1,365 @@
|
||||
"""Integration tests for the TCP proxy — real asyncio TCP server + client."""
|
||||
|
||||
import asyncio
|
||||
|
||||
import pytest
|
||||
|
||||
from app.tcp_proxy.protocol import (
|
||||
CMD_APP_START,
|
||||
CMD_DEVICE_QUERY,
|
||||
CMD_GET_CHANNEL,
|
||||
CMD_GET_CONTACTS,
|
||||
CMD_GET_DEVICE_TIME,
|
||||
CMD_HAS_CONNECTION,
|
||||
CMD_SET_CHANNEL,
|
||||
CMD_SYNC_NEXT_MESSAGE,
|
||||
FRAME_RX,
|
||||
FRAME_TX,
|
||||
PROXY_FW_VER,
|
||||
PUSH_MSG_WAITING,
|
||||
RESP_CONTACT_END,
|
||||
RESP_CONTACT_START,
|
||||
RESP_CURRENT_TIME,
|
||||
RESP_DEVICE_INFO,
|
||||
RESP_ERR,
|
||||
RESP_NO_MORE_MSGS,
|
||||
RESP_OK,
|
||||
RESP_SELF_INFO,
|
||||
)
|
||||
from app.tcp_proxy.server import dispatch_event, register, unregister
|
||||
from app.tcp_proxy.session import ProxySession
|
||||
|
||||
# ── Helpers ──────────────────────────────────────────────────────────
|
||||
|
||||
EXAMPLE_KEY = "ab" * 32
|
||||
|
||||
|
||||
def _frame_cmd(payload: bytes) -> bytes:
|
||||
"""Wrap a command payload in a 0x3C frame."""
|
||||
return bytes([FRAME_TX]) + len(payload).to_bytes(2, "little") + payload
|
||||
|
||||
|
||||
async def _read_response(reader: asyncio.StreamReader) -> bytes:
|
||||
"""Read one 0x3E-framed response and return the payload."""
|
||||
marker = await reader.readexactly(1)
|
||||
assert marker[0] == FRAME_RX
|
||||
size_bytes = await reader.readexactly(2)
|
||||
size = int.from_bytes(size_bytes, "little")
|
||||
payload = await reader.readexactly(size)
|
||||
return payload
|
||||
|
||||
|
||||
class _ProxyTestHarness:
|
||||
"""Manages a real TCP proxy server for testing."""
|
||||
|
||||
def __init__(self):
|
||||
self._server: asyncio.Server | None = None
|
||||
self.port: int = 0
|
||||
self.sessions: list[ProxySession] = []
|
||||
|
||||
async def start(self):
|
||||
self._server = await asyncio.start_server(self._handle, "127.0.0.1", 0)
|
||||
self.port = self._server.sockets[0].getsockname()[1]
|
||||
|
||||
async def stop(self):
|
||||
for s in self.sessions:
|
||||
try:
|
||||
s.writer.close()
|
||||
except Exception:
|
||||
pass
|
||||
self.sessions.clear()
|
||||
if self._server:
|
||||
self._server.close()
|
||||
await self._server.wait_closed()
|
||||
|
||||
async def _handle(self, reader, writer):
|
||||
session = ProxySession(reader, writer)
|
||||
self.sessions.append(session)
|
||||
register(session)
|
||||
try:
|
||||
await session.run()
|
||||
finally:
|
||||
unregister(session)
|
||||
|
||||
async def connect(self) -> tuple[asyncio.StreamReader, asyncio.StreamWriter]:
|
||||
reader, writer = await asyncio.open_connection("127.0.0.1", self.port)
|
||||
return reader, writer
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def harness():
|
||||
h = _ProxyTestHarness()
|
||||
await h.start()
|
||||
yield h
|
||||
await h.stop()
|
||||
|
||||
|
||||
def _mock_repos_and_runtime():
|
||||
"""Return a context manager that mocks repositories and radio_runtime."""
|
||||
import time
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
contacts = [
|
||||
MagicMock(
|
||||
model_dump=MagicMock(
|
||||
return_value={
|
||||
"public_key": EXAMPLE_KEY,
|
||||
"name": "Alice",
|
||||
"type": 1,
|
||||
"favorite": True,
|
||||
"direct_path": None,
|
||||
"direct_path_len": -1,
|
||||
"direct_path_hash_mode": -1,
|
||||
"last_advert": 0,
|
||||
"lat": 0.0,
|
||||
"lon": 0.0,
|
||||
"first_seen": int(time.time()),
|
||||
}
|
||||
)
|
||||
)
|
||||
]
|
||||
channels = [
|
||||
MagicMock(
|
||||
model_dump=MagicMock(return_value={"key": "cc" * 16, "name": "test", "favorite": True})
|
||||
)
|
||||
]
|
||||
settings_obj = MagicMock(last_message_times={})
|
||||
|
||||
rt = MagicMock()
|
||||
rt.is_connected = True
|
||||
mc = MagicMock()
|
||||
mc.self_info = {
|
||||
"public_key": EXAMPLE_KEY,
|
||||
"name": "TestNode",
|
||||
"tx_power": 20,
|
||||
"max_tx_power": 22,
|
||||
"adv_lat": 0.0,
|
||||
"adv_lon": 0.0,
|
||||
"radio_freq": 915.0,
|
||||
"radio_bw": 250.0,
|
||||
"radio_sf": 10,
|
||||
"radio_cr": 7,
|
||||
}
|
||||
rt.meshcore = mc
|
||||
|
||||
class _Ctx:
|
||||
def __enter__(self_):
|
||||
self_._patches = [
|
||||
patch(
|
||||
"app.repository.ContactRepository.get_favorites",
|
||||
new_callable=AsyncMock,
|
||||
return_value=contacts,
|
||||
),
|
||||
patch(
|
||||
"app.repository.ChannelRepository.get_all",
|
||||
new_callable=AsyncMock,
|
||||
return_value=channels,
|
||||
),
|
||||
patch(
|
||||
"app.repository.AppSettingsRepository.get",
|
||||
new_callable=AsyncMock,
|
||||
return_value=settings_obj,
|
||||
),
|
||||
patch(
|
||||
"app.services.radio_runtime.radio_runtime",
|
||||
rt,
|
||||
),
|
||||
]
|
||||
for p in self_._patches:
|
||||
p.__enter__()
|
||||
return self_
|
||||
|
||||
def __exit__(self_, *args):
|
||||
for p in reversed(self_._patches):
|
||||
p.__exit__(*args)
|
||||
|
||||
return _Ctx()
|
||||
|
||||
|
||||
# ── Tests ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestTcpProxyIntegration:
|
||||
@pytest.mark.asyncio
|
||||
async def test_app_start_returns_self_info(self, harness):
|
||||
reader, writer = await harness.connect()
|
||||
try:
|
||||
with _mock_repos_and_runtime():
|
||||
writer.write(_frame_cmd(bytes([CMD_APP_START]) + b"\x03" + b" " * 6 + b"test"))
|
||||
await writer.drain()
|
||||
resp = await asyncio.wait_for(_read_response(reader), timeout=3)
|
||||
assert resp[0] == RESP_SELF_INFO
|
||||
finally:
|
||||
writer.close()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_device_query_returns_device_info(self, harness):
|
||||
reader, writer = await harness.connect()
|
||||
try:
|
||||
with _mock_repos_and_runtime():
|
||||
# First do APP_START to initialize session state
|
||||
writer.write(_frame_cmd(bytes([CMD_APP_START]) + b"\x03" + b" " * 6 + b"test"))
|
||||
await writer.drain()
|
||||
await asyncio.wait_for(_read_response(reader), timeout=3)
|
||||
|
||||
writer.write(_frame_cmd(bytes([CMD_DEVICE_QUERY, 0x03])))
|
||||
await writer.drain()
|
||||
resp = await asyncio.wait_for(_read_response(reader), timeout=3)
|
||||
assert resp[0] == RESP_DEVICE_INFO
|
||||
assert resp[1] == PROXY_FW_VER
|
||||
finally:
|
||||
writer.close()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_contacts_flow(self, harness):
|
||||
reader, writer = await harness.connect()
|
||||
try:
|
||||
with _mock_repos_and_runtime():
|
||||
writer.write(_frame_cmd(bytes([CMD_GET_CONTACTS])))
|
||||
await writer.drain()
|
||||
|
||||
# Should get CONTACT_START
|
||||
resp1 = await asyncio.wait_for(_read_response(reader), timeout=3)
|
||||
assert resp1[0] == RESP_CONTACT_START
|
||||
count = int.from_bytes(resp1[1:5], "little")
|
||||
assert count == 1
|
||||
|
||||
# One contact
|
||||
resp2 = await asyncio.wait_for(_read_response(reader), timeout=3)
|
||||
assert resp2[0] == 0x03 # RESP_CONTACT
|
||||
|
||||
# CONTACT_END
|
||||
resp3 = await asyncio.wait_for(_read_response(reader), timeout=3)
|
||||
assert resp3[0] == RESP_CONTACT_END
|
||||
finally:
|
||||
writer.close()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_time(self, harness):
|
||||
reader, writer = await harness.connect()
|
||||
try:
|
||||
writer.write(_frame_cmd(bytes([CMD_GET_DEVICE_TIME])))
|
||||
await writer.drain()
|
||||
resp = await asyncio.wait_for(_read_response(reader), timeout=3)
|
||||
assert resp[0] == RESP_CURRENT_TIME
|
||||
assert len(resp) == 5
|
||||
finally:
|
||||
writer.close()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_has_connection(self, harness):
|
||||
reader, writer = await harness.connect()
|
||||
try:
|
||||
with _mock_repos_and_runtime():
|
||||
writer.write(_frame_cmd(bytes([CMD_HAS_CONNECTION])))
|
||||
await writer.drain()
|
||||
resp = await asyncio.wait_for(_read_response(reader), timeout=3)
|
||||
assert resp[0] == RESP_OK
|
||||
val = int.from_bytes(resp[1:5], "little")
|
||||
assert val == 1
|
||||
finally:
|
||||
writer.close()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_channel_returns_error(self, harness):
|
||||
reader, writer = await harness.connect()
|
||||
try:
|
||||
writer.write(_frame_cmd(bytes([CMD_GET_CHANNEL, 5])))
|
||||
await writer.drain()
|
||||
resp = await asyncio.wait_for(_read_response(reader), timeout=3)
|
||||
assert resp[0] == RESP_ERR
|
||||
finally:
|
||||
writer.close()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_set_then_get_channel(self, harness):
|
||||
reader, writer = await harness.connect()
|
||||
try:
|
||||
# SET_CHANNEL: cmd(1) + idx(1) + name(32) + secret(16) = 50
|
||||
name = b"mychan" + b"\x00" * 26 # 32 bytes
|
||||
secret = b"\xdd" * 16
|
||||
cmd = bytes([CMD_SET_CHANNEL, 2]) + name + secret
|
||||
writer.write(_frame_cmd(cmd))
|
||||
await writer.drain()
|
||||
resp = await asyncio.wait_for(_read_response(reader), timeout=3)
|
||||
assert resp[0] == RESP_OK
|
||||
|
||||
# GET_CHANNEL for slot 2
|
||||
writer.write(_frame_cmd(bytes([CMD_GET_CHANNEL, 2])))
|
||||
await writer.drain()
|
||||
resp2 = await asyncio.wait_for(_read_response(reader), timeout=3)
|
||||
assert resp2[0] == 0x12 # RESP_CHANNEL_INFO
|
||||
assert resp2[1] == 2 # idx
|
||||
finally:
|
||||
writer.close()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_sync_next_empty(self, harness):
|
||||
reader, writer = await harness.connect()
|
||||
try:
|
||||
writer.write(_frame_cmd(bytes([CMD_SYNC_NEXT_MESSAGE])))
|
||||
await writer.drain()
|
||||
resp = await asyncio.wait_for(_read_response(reader), timeout=3)
|
||||
assert resp[0] == RESP_NO_MORE_MSGS
|
||||
finally:
|
||||
writer.close()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_event_dispatch_queues_message(self, harness):
|
||||
reader, writer = await harness.connect()
|
||||
try:
|
||||
with _mock_repos_and_runtime():
|
||||
# APP_START to init session
|
||||
writer.write(_frame_cmd(bytes([CMD_APP_START]) + b"\x03" + b" " * 6 + b"test"))
|
||||
await writer.drain()
|
||||
await asyncio.wait_for(_read_response(reader), timeout=3)
|
||||
|
||||
# Set a channel so CHAN messages can be routed
|
||||
name = b"\x00" * 32
|
||||
secret = bytes.fromhex("cc" * 16)
|
||||
writer.write(_frame_cmd(bytes([CMD_SET_CHANNEL, 0]) + name + secret))
|
||||
await writer.drain()
|
||||
await asyncio.wait_for(_read_response(reader), timeout=3)
|
||||
|
||||
# Simulate a broadcast event
|
||||
await dispatch_event(
|
||||
"message",
|
||||
{
|
||||
"type": "CHAN",
|
||||
"outgoing": False,
|
||||
"conversation_key": "cc" * 16,
|
||||
"text": "hello from event",
|
||||
"sender_timestamp": 1700000000,
|
||||
},
|
||||
)
|
||||
|
||||
# Should receive PUSH_MSG_WAITING
|
||||
resp = await asyncio.wait_for(_read_response(reader), timeout=3)
|
||||
assert resp[0] == PUSH_MSG_WAITING
|
||||
|
||||
# Pull the message
|
||||
writer.write(_frame_cmd(bytes([CMD_SYNC_NEXT_MESSAGE])))
|
||||
await writer.drain()
|
||||
msg = await asyncio.wait_for(_read_response(reader), timeout=3)
|
||||
assert msg[0] == 0x11 # RESP_CHANNEL_MSG_RECV_V3
|
||||
finally:
|
||||
writer.close()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_multiple_clients_isolated(self, harness):
|
||||
r1, w1 = await harness.connect()
|
||||
r2, w2 = await harness.connect()
|
||||
try:
|
||||
# Both can get time independently
|
||||
w1.write(_frame_cmd(bytes([CMD_GET_DEVICE_TIME])))
|
||||
w2.write(_frame_cmd(bytes([CMD_GET_DEVICE_TIME])))
|
||||
await w1.drain()
|
||||
await w2.drain()
|
||||
|
||||
resp1 = await asyncio.wait_for(_read_response(r1), timeout=3)
|
||||
resp2 = await asyncio.wait_for(_read_response(r2), timeout=3)
|
||||
assert resp1[0] == RESP_CURRENT_TIME
|
||||
assert resp2[0] == RESP_CURRENT_TIME
|
||||
finally:
|
||||
w1.close()
|
||||
w2.close()
|
||||
@@ -0,0 +1,180 @@
|
||||
"""Tests for app.tcp_proxy.protocol — frame parsing, helpers, constants."""
|
||||
|
||||
from app.tcp_proxy.protocol import (
|
||||
ERR_NOT_FOUND,
|
||||
ERR_UNSUPPORTED,
|
||||
FRAME_RX,
|
||||
FRAME_TX,
|
||||
RESP_ERR,
|
||||
RESP_OK,
|
||||
FrameParser,
|
||||
build_error,
|
||||
build_ok,
|
||||
encode_path_byte,
|
||||
frame_response,
|
||||
pad,
|
||||
)
|
||||
|
||||
# ── frame_response ───────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestFrameResponse:
|
||||
def test_empty_payload(self):
|
||||
result = frame_response(b"")
|
||||
assert result == bytes([FRAME_RX, 0x00, 0x00])
|
||||
|
||||
def test_short_payload(self):
|
||||
result = frame_response(b"\x05\x01")
|
||||
assert result[0] == FRAME_RX
|
||||
size = int.from_bytes(result[1:3], "little")
|
||||
assert size == 2
|
||||
assert result[3:] == b"\x05\x01"
|
||||
|
||||
def test_larger_payload(self):
|
||||
payload = b"\xaa" * 200
|
||||
result = frame_response(payload)
|
||||
assert result[0] == FRAME_RX
|
||||
size = int.from_bytes(result[1:3], "little")
|
||||
assert size == 200
|
||||
assert result[3:] == payload
|
||||
|
||||
|
||||
# ── build_ok / build_error ───────────────────────────────────────────
|
||||
|
||||
|
||||
class TestBuildOk:
|
||||
def test_no_value(self):
|
||||
assert build_ok() == bytes([RESP_OK])
|
||||
|
||||
def test_with_value(self):
|
||||
result = build_ok(42)
|
||||
assert result[0] == RESP_OK
|
||||
assert int.from_bytes(result[1:5], "little") == 42
|
||||
|
||||
def test_zero_value(self):
|
||||
result = build_ok(0)
|
||||
assert len(result) == 5
|
||||
assert int.from_bytes(result[1:5], "little") == 0
|
||||
|
||||
|
||||
class TestBuildError:
|
||||
def test_default_code(self):
|
||||
assert build_error() == bytes([RESP_ERR, ERR_UNSUPPORTED])
|
||||
|
||||
def test_not_found(self):
|
||||
assert build_error(ERR_NOT_FOUND) == bytes([RESP_ERR, ERR_NOT_FOUND])
|
||||
|
||||
|
||||
# ── pad ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestPad:
|
||||
def test_shorter_data(self):
|
||||
result = pad(b"AB", 5)
|
||||
assert result == b"AB\x00\x00\x00"
|
||||
assert len(result) == 5
|
||||
|
||||
def test_exact_data(self):
|
||||
assert pad(b"ABCDE", 5) == b"ABCDE"
|
||||
|
||||
def test_longer_data(self):
|
||||
assert pad(b"ABCDEFGH", 5) == b"ABCDE"
|
||||
|
||||
def test_empty_data(self):
|
||||
assert pad(b"", 3) == b"\x00\x00\x00"
|
||||
|
||||
|
||||
# ── encode_path_byte ────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestEncodePathByte:
|
||||
def test_flood_negative_hop(self):
|
||||
assert encode_path_byte(-1, 0) == 0xFF
|
||||
|
||||
def test_flood_negative_mode(self):
|
||||
assert encode_path_byte(0, -1) == 0xFF
|
||||
|
||||
def test_flood_both_negative(self):
|
||||
assert encode_path_byte(-1, -1) == 0xFF
|
||||
|
||||
def test_zero_hops_mode_zero(self):
|
||||
assert encode_path_byte(0, 0) == 0x00
|
||||
|
||||
def test_three_hops_mode_one(self):
|
||||
# mode=1 → bits 6-7 = 01 → 0x40; hops=3 → 0x03
|
||||
assert encode_path_byte(3, 1) == 0x43
|
||||
|
||||
def test_max_hops_mode_two(self):
|
||||
# mode=2 → bits 6-7 = 10 → 0x80; hops=63 → 0x3F
|
||||
assert encode_path_byte(63, 2) == 0xBF
|
||||
|
||||
|
||||
# ── FrameParser ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestFrameParser:
|
||||
def test_single_complete_frame(self):
|
||||
parser = FrameParser()
|
||||
# 0x3C + 2-byte LE size (3) + 3 bytes payload
|
||||
data = bytes([FRAME_TX, 0x03, 0x00, 0xAA, 0xBB, 0xCC])
|
||||
payloads = parser.feed(data)
|
||||
assert len(payloads) == 1
|
||||
assert payloads[0] == b"\xaa\xbb\xcc"
|
||||
|
||||
def test_two_frames_in_one_chunk(self):
|
||||
parser = FrameParser()
|
||||
frame1 = bytes([FRAME_TX, 0x02, 0x00, 0x01, 0x02])
|
||||
frame2 = bytes([FRAME_TX, 0x01, 0x00, 0xFF])
|
||||
payloads = parser.feed(frame1 + frame2)
|
||||
assert len(payloads) == 2
|
||||
assert payloads[0] == b"\x01\x02"
|
||||
assert payloads[1] == b"\xff"
|
||||
|
||||
def test_split_across_chunks(self):
|
||||
parser = FrameParser()
|
||||
full = bytes([FRAME_TX, 0x04, 0x00, 0x01, 0x02, 0x03, 0x04])
|
||||
# Split in the middle of the payload
|
||||
p1 = parser.feed(full[:5])
|
||||
assert p1 == []
|
||||
p2 = parser.feed(full[5:])
|
||||
assert len(p2) == 1
|
||||
assert p2[0] == b"\x01\x02\x03\x04"
|
||||
|
||||
def test_split_in_header(self):
|
||||
parser = FrameParser()
|
||||
full = bytes([FRAME_TX, 0x01, 0x00, 0xAA])
|
||||
p1 = parser.feed(full[:2]) # marker + first size byte
|
||||
assert p1 == []
|
||||
p2 = parser.feed(full[2:]) # second size byte + payload
|
||||
assert len(p2) == 1
|
||||
assert p2[0] == b"\xaa"
|
||||
|
||||
def test_bad_marker_skipped(self):
|
||||
parser = FrameParser()
|
||||
junk = b"\x00\x00\x00"
|
||||
good = bytes([FRAME_TX, 0x01, 0x00, 0xBB])
|
||||
payloads = parser.feed(junk + good)
|
||||
assert len(payloads) == 1
|
||||
assert payloads[0] == b"\xbb"
|
||||
|
||||
def test_oversized_frame_skipped(self):
|
||||
parser = FrameParser()
|
||||
# Size = 400 (> MAX_FRAME_SIZE=300)
|
||||
bad = bytes([FRAME_TX, 0x90, 0x01])
|
||||
good = bytes([FRAME_TX, 0x01, 0x00, 0xCC])
|
||||
payloads = parser.feed(bad + good)
|
||||
assert len(payloads) == 1
|
||||
assert payloads[0] == b"\xcc"
|
||||
|
||||
def test_empty_feed(self):
|
||||
parser = FrameParser()
|
||||
assert parser.feed(b"") == []
|
||||
|
||||
def test_byte_at_a_time(self):
|
||||
parser = FrameParser()
|
||||
full = bytes([FRAME_TX, 0x02, 0x00, 0xDE, 0xAD])
|
||||
payloads = []
|
||||
for b in full:
|
||||
payloads.extend(parser.feed(bytes([b])))
|
||||
assert len(payloads) == 1
|
||||
assert payloads[0] == b"\xde\xad"
|
||||
@@ -0,0 +1,526 @@
|
||||
"""Tests for app.tcp_proxy.session — ProxySession command handlers."""
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from app.tcp_proxy.protocol import (
|
||||
CMD_APP_START,
|
||||
CMD_DEVICE_QUERY,
|
||||
CMD_GET_BATT_AND_STORAGE,
|
||||
CMD_GET_CHANNEL,
|
||||
CMD_GET_CONTACT_BY_KEY,
|
||||
CMD_GET_CONTACTS,
|
||||
CMD_GET_DEVICE_TIME,
|
||||
CMD_HAS_CONNECTION,
|
||||
CMD_RESET_PATH,
|
||||
CMD_SEND_CHANNEL_TXT_MSG,
|
||||
CMD_SEND_TXT_MSG,
|
||||
CMD_SET_CHANNEL,
|
||||
CMD_SYNC_NEXT_MESSAGE,
|
||||
ERR_NOT_FOUND,
|
||||
PROXY_FW_VER,
|
||||
PUSH_MSG_WAITING,
|
||||
RESP_BATTERY,
|
||||
RESP_CONTACT_END,
|
||||
RESP_CONTACT_START,
|
||||
RESP_CURRENT_TIME,
|
||||
RESP_DEVICE_INFO,
|
||||
RESP_ERR,
|
||||
RESP_MSG_SENT,
|
||||
RESP_NO_MORE_MSGS,
|
||||
RESP_OK,
|
||||
RESP_SELF_INFO,
|
||||
)
|
||||
from app.tcp_proxy.session import ProxySession
|
||||
|
||||
EXAMPLE_KEY = "ab" * 32
|
||||
|
||||
|
||||
# ── Helpers ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _make_session() -> tuple[ProxySession, list[bytes]]:
|
||||
"""Create a ProxySession with a capturing writer."""
|
||||
reader = AsyncMock(spec=asyncio.StreamReader)
|
||||
writer = MagicMock(spec=asyncio.StreamWriter)
|
||||
writer.get_extra_info.return_value = ("127.0.0.1", 12345)
|
||||
|
||||
sent: list[bytes] = []
|
||||
|
||||
def capture_write(data: bytes):
|
||||
sent.append(data)
|
||||
|
||||
writer.write = capture_write
|
||||
writer.drain = AsyncMock()
|
||||
|
||||
session = ProxySession(reader, writer)
|
||||
return session, sent
|
||||
|
||||
|
||||
def _extract_payloads(sent: list[bytes]) -> list[bytes]:
|
||||
"""Extract payloads from framed response bytes."""
|
||||
payloads = []
|
||||
for frame in sent:
|
||||
assert frame[0] == 0x3E
|
||||
size = int.from_bytes(frame[1:3], "little")
|
||||
payloads.append(frame[3 : 3 + size])
|
||||
return payloads
|
||||
|
||||
|
||||
def _make_contact(public_key: str = EXAMPLE_KEY, name: str = "Alice", **kw):
|
||||
return MagicMock(
|
||||
model_dump=MagicMock(
|
||||
return_value={
|
||||
"public_key": public_key,
|
||||
"name": name,
|
||||
"type": 1,
|
||||
"favorite": True,
|
||||
"direct_path": None,
|
||||
"direct_path_len": -1,
|
||||
"direct_path_hash_mode": -1,
|
||||
"last_advert": 0,
|
||||
"lat": 0.0,
|
||||
"lon": 0.0,
|
||||
"first_seen": int(time.time()),
|
||||
**kw,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def _make_channel(key: str = "cc" * 16, name: str = "test", favorite: bool = True):
|
||||
return MagicMock(
|
||||
model_dump=MagicMock(return_value={"key": key, "name": name, "favorite": favorite})
|
||||
)
|
||||
|
||||
|
||||
def _make_settings(last_message_times=None):
|
||||
return MagicMock(last_message_times=last_message_times or {})
|
||||
|
||||
|
||||
def _mock_radio_runtime(connected: bool = True, self_info: dict | None = None):
|
||||
rt = MagicMock()
|
||||
rt.is_connected = connected
|
||||
mc = MagicMock()
|
||||
mc.self_info = self_info or {
|
||||
"public_key": EXAMPLE_KEY,
|
||||
"name": "TestNode",
|
||||
"tx_power": 20,
|
||||
"max_tx_power": 22,
|
||||
"adv_lat": 0.0,
|
||||
"adv_lon": 0.0,
|
||||
"radio_freq": 915.0,
|
||||
"radio_bw": 250.0,
|
||||
"radio_sf": 10,
|
||||
"radio_cr": 7,
|
||||
}
|
||||
rt.meshcore = mc
|
||||
return rt
|
||||
|
||||
|
||||
# ── Tests ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestAppStart:
|
||||
@pytest.mark.asyncio
|
||||
async def test_sends_self_info(self):
|
||||
session, sent = _make_session()
|
||||
contacts = [_make_contact()]
|
||||
channels = [_make_channel()]
|
||||
settings = _make_settings()
|
||||
rt = _mock_radio_runtime()
|
||||
|
||||
with (
|
||||
patch("app.repository.ContactRepository") as cr,
|
||||
patch("app.repository.ChannelRepository") as chr_,
|
||||
patch("app.repository.AppSettingsRepository") as sr,
|
||||
patch("app.services.radio_runtime.radio_runtime", rt),
|
||||
):
|
||||
cr.get_favorites = AsyncMock(return_value=contacts)
|
||||
chr_.get_all = AsyncMock(return_value=channels)
|
||||
sr.get = AsyncMock(return_value=settings)
|
||||
|
||||
await session._cmd_app_start(bytes([CMD_APP_START]))
|
||||
|
||||
payloads = _extract_payloads(sent)
|
||||
assert len(payloads) == 1
|
||||
assert payloads[0][0] == RESP_SELF_INFO
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_populates_contacts_and_channels(self):
|
||||
session, sent = _make_session()
|
||||
contacts = [_make_contact(), _make_contact(public_key="cd" * 32, name="Bob")]
|
||||
channels = [_make_channel(), _make_channel(key="dd" * 16, name="ch2")]
|
||||
settings = _make_settings()
|
||||
rt = _mock_radio_runtime()
|
||||
|
||||
with (
|
||||
patch("app.repository.ContactRepository") as cr,
|
||||
patch("app.repository.ChannelRepository") as chr_,
|
||||
patch("app.repository.AppSettingsRepository") as sr,
|
||||
patch("app.services.radio_runtime.radio_runtime", rt),
|
||||
):
|
||||
cr.get_favorites = AsyncMock(return_value=contacts)
|
||||
chr_.get_all = AsyncMock(return_value=channels)
|
||||
sr.get = AsyncMock(return_value=settings)
|
||||
|
||||
await session._cmd_app_start(bytes([CMD_APP_START]))
|
||||
|
||||
assert len(session.contacts) == 2
|
||||
# Only favorite channels are slotted
|
||||
assert len(session.channel_slots) == 2
|
||||
|
||||
|
||||
class TestDeviceQuery:
|
||||
@pytest.mark.asyncio
|
||||
async def test_sends_device_info(self):
|
||||
session, sent = _make_session()
|
||||
rt = _mock_radio_runtime()
|
||||
|
||||
with patch("app.services.radio_runtime.radio_runtime", rt):
|
||||
await session._cmd_device_query(bytes([CMD_DEVICE_QUERY, 0x03]))
|
||||
|
||||
payloads = _extract_payloads(sent)
|
||||
assert payloads[0][0] == RESP_DEVICE_INFO
|
||||
assert payloads[0][1] == PROXY_FW_VER
|
||||
|
||||
|
||||
class TestGetContacts:
|
||||
@pytest.mark.asyncio
|
||||
async def test_sends_start_contacts_end(self):
|
||||
session, sent = _make_session()
|
||||
contacts = [_make_contact()]
|
||||
|
||||
with patch("app.repository.ContactRepository") as cr:
|
||||
cr.get_favorites = AsyncMock(return_value=contacts)
|
||||
await session._cmd_get_contacts(bytes([CMD_GET_CONTACTS]))
|
||||
|
||||
payloads = _extract_payloads(sent)
|
||||
assert payloads[0][0] == RESP_CONTACT_START
|
||||
count = int.from_bytes(payloads[0][1:5], "little")
|
||||
assert count == 1
|
||||
# Middle payload(s) are contacts
|
||||
assert payloads[-1][0] == RESP_CONTACT_END
|
||||
|
||||
|
||||
class TestGetContactByKey:
|
||||
@pytest.mark.asyncio
|
||||
async def test_found(self):
|
||||
session, sent = _make_session()
|
||||
session.contacts = [
|
||||
{
|
||||
"public_key": EXAMPLE_KEY,
|
||||
"type": 1,
|
||||
"name": "Alice",
|
||||
"favorite": True,
|
||||
"direct_path": None,
|
||||
"direct_path_len": -1,
|
||||
"direct_path_hash_mode": -1,
|
||||
"last_advert": 0,
|
||||
"lat": 0.0,
|
||||
"lon": 0.0,
|
||||
"first_seen": 0,
|
||||
}
|
||||
]
|
||||
|
||||
cmd = bytes([CMD_GET_CONTACT_BY_KEY]) + bytes.fromhex(EXAMPLE_KEY)
|
||||
await session._cmd_get_contact_by_key(cmd)
|
||||
|
||||
payloads = _extract_payloads(sent)
|
||||
assert len(payloads) == 1
|
||||
assert payloads[0][0] == 0x03 # RESP_CONTACT
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_not_found(self):
|
||||
session, sent = _make_session()
|
||||
session.contacts = []
|
||||
|
||||
cmd = bytes([CMD_GET_CONTACT_BY_KEY]) + bytes.fromhex(EXAMPLE_KEY)
|
||||
await session._cmd_get_contact_by_key(cmd)
|
||||
|
||||
payloads = _extract_payloads(sent)
|
||||
assert payloads[0][0] == RESP_ERR
|
||||
assert payloads[0][1] == ERR_NOT_FOUND
|
||||
|
||||
|
||||
class TestGetChannel:
|
||||
@pytest.mark.asyncio
|
||||
async def test_found(self):
|
||||
session, sent = _make_session()
|
||||
key = "cc" * 16
|
||||
session.channel_slots = {0: key}
|
||||
session.channels = [{"key": key, "name": "test"}]
|
||||
|
||||
await session._cmd_get_channel(bytes([CMD_GET_CHANNEL, 0]))
|
||||
|
||||
payloads = _extract_payloads(sent)
|
||||
assert payloads[0][0] == 0x12 # RESP_CHANNEL_INFO
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_slot_returns_error(self):
|
||||
session, sent = _make_session()
|
||||
session.channel_slots = {}
|
||||
|
||||
await session._cmd_get_channel(bytes([CMD_GET_CHANNEL, 5]))
|
||||
|
||||
payloads = _extract_payloads(sent)
|
||||
assert payloads[0][0] == RESP_ERR
|
||||
|
||||
|
||||
class TestSetChannel:
|
||||
@pytest.mark.asyncio
|
||||
async def test_updates_slot_mapping(self):
|
||||
session, sent = _make_session()
|
||||
|
||||
name = b"test" + b"\x00" * 28 # 32 bytes
|
||||
secret = b"\xaa" * 16
|
||||
cmd = bytes([CMD_SET_CHANNEL, 3]) + name + secret
|
||||
await session._cmd_set_channel(cmd)
|
||||
|
||||
assert session.channel_slots[3] == "aa" * 16
|
||||
assert session.key_to_idx["aa" * 16] == 3
|
||||
|
||||
payloads = _extract_payloads(sent)
|
||||
assert payloads[0][0] == RESP_OK
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cleans_stale_mapping(self):
|
||||
session, sent = _make_session()
|
||||
|
||||
# Pre-load slot 0 with key_a
|
||||
session.channel_slots[0] = "aa" * 16
|
||||
session.key_to_idx["aa" * 16] = 0
|
||||
|
||||
# Overwrite slot 0 with key_b
|
||||
name = b"\x00" * 32
|
||||
secret_b = b"\xbb" * 16
|
||||
cmd = bytes([CMD_SET_CHANNEL, 0]) + name + secret_b
|
||||
await session._cmd_set_channel(cmd)
|
||||
|
||||
assert session.channel_slots[0] == "bb" * 16
|
||||
assert "aa" * 16 not in session.key_to_idx
|
||||
|
||||
|
||||
class TestSendDm:
|
||||
@pytest.mark.asyncio
|
||||
async def test_sends_msg_sent_and_ack(self):
|
||||
session, sent = _make_session()
|
||||
session.contacts = [{"public_key": EXAMPLE_KEY}]
|
||||
|
||||
# CMD_SEND_TXT_MSG: cmd(1) + txt_type(1) + attempt(1) + ts(4) + prefix(6) + text
|
||||
prefix = bytes.fromhex(EXAMPLE_KEY[:12])
|
||||
cmd = (
|
||||
bytes([CMD_SEND_TXT_MSG, 0, 0])
|
||||
+ int(time.time()).to_bytes(4, "little")
|
||||
+ prefix
|
||||
+ b"Hello"
|
||||
)
|
||||
|
||||
with patch.object(session, "_do_send_dm", new_callable=AsyncMock):
|
||||
await session._cmd_send_dm(cmd)
|
||||
|
||||
payloads = _extract_payloads(sent)
|
||||
assert payloads[0][0] == RESP_MSG_SENT
|
||||
assert payloads[1][0] == 0x82 # PUSH_ACK
|
||||
# ACK code should match
|
||||
ack_from_sent = payloads[0][2:6]
|
||||
ack_from_push = payloads[1][1:5]
|
||||
assert ack_from_sent == ack_from_push
|
||||
|
||||
|
||||
class TestSendChannel:
|
||||
@pytest.mark.asyncio
|
||||
async def test_sends_ok(self):
|
||||
session, sent = _make_session()
|
||||
key = "cc" * 16
|
||||
session.channel_slots = {0: key}
|
||||
session.channels = [{"key": key, "name": "test"}]
|
||||
|
||||
cmd = (
|
||||
bytes([CMD_SEND_CHANNEL_TXT_MSG, 0, 0])
|
||||
+ int(time.time()).to_bytes(4, "little")
|
||||
+ b"Hello"
|
||||
)
|
||||
|
||||
fake_channel = MagicMock(name="test")
|
||||
with (
|
||||
patch(
|
||||
"app.repository.ChannelRepository.get_by_key",
|
||||
new_callable=AsyncMock,
|
||||
return_value=fake_channel,
|
||||
),
|
||||
patch.object(session, "_do_send_channel", new_callable=AsyncMock),
|
||||
):
|
||||
await session._cmd_send_channel(cmd)
|
||||
|
||||
payloads = _extract_payloads(sent)
|
||||
assert payloads[0][0] == RESP_OK
|
||||
|
||||
|
||||
class TestSimpleCommands:
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_time(self):
|
||||
session, sent = _make_session()
|
||||
await session._cmd_get_time(bytes([CMD_GET_DEVICE_TIME]))
|
||||
payloads = _extract_payloads(sent)
|
||||
assert payloads[0][0] == RESP_CURRENT_TIME
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_battery(self):
|
||||
session, sent = _make_session()
|
||||
await session._cmd_battery(bytes([CMD_GET_BATT_AND_STORAGE]))
|
||||
payloads = _extract_payloads(sent)
|
||||
assert payloads[0][0] == RESP_BATTERY
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_has_connection(self):
|
||||
session, sent = _make_session()
|
||||
rt = _mock_radio_runtime(connected=True)
|
||||
with patch("app.services.radio_runtime.radio_runtime", rt):
|
||||
await session._cmd_has_connection(bytes([CMD_HAS_CONNECTION]))
|
||||
payloads = _extract_payloads(sent)
|
||||
assert payloads[0][0] == RESP_OK
|
||||
val = int.from_bytes(payloads[0][1:5], "little")
|
||||
assert val == 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ok_stub(self):
|
||||
session, sent = _make_session()
|
||||
await session._cmd_ok_stub(bytes([CMD_RESET_PATH]))
|
||||
payloads = _extract_payloads(sent)
|
||||
assert payloads[0][0] == RESP_OK
|
||||
|
||||
|
||||
class TestSyncNext:
|
||||
@pytest.mark.asyncio
|
||||
async def test_empty_queue(self):
|
||||
session, sent = _make_session()
|
||||
await session._cmd_sync_next(bytes([CMD_SYNC_NEXT_MESSAGE]))
|
||||
payloads = _extract_payloads(sent)
|
||||
assert payloads[0][0] == RESP_NO_MORE_MSGS
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dequeues_message(self):
|
||||
session, sent = _make_session()
|
||||
fake_msg = bytes([0x10, 0x00, 0x00, 0x00]) + b"\xaa" * 10
|
||||
session._msg_queue.append(fake_msg)
|
||||
|
||||
await session._cmd_sync_next(bytes([CMD_SYNC_NEXT_MESSAGE]))
|
||||
payloads = _extract_payloads(sent)
|
||||
assert payloads[0] == fake_msg
|
||||
assert len(session._msg_queue) == 0
|
||||
|
||||
|
||||
class TestEventHandlers:
|
||||
@pytest.mark.asyncio
|
||||
async def test_priv_message_queued(self):
|
||||
session, sent = _make_session()
|
||||
data = {
|
||||
"type": "PRIV",
|
||||
"outgoing": False,
|
||||
"conversation_key": EXAMPLE_KEY,
|
||||
"text": "hello",
|
||||
"sender_timestamp": 1700000000,
|
||||
}
|
||||
await session.on_event_message(data)
|
||||
|
||||
assert len(session._msg_queue) == 1
|
||||
payloads = _extract_payloads(sent)
|
||||
assert payloads[0][0] == PUSH_MSG_WAITING
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chan_message_queued(self):
|
||||
session, sent = _make_session()
|
||||
key = "cc" * 16
|
||||
session.key_to_idx = {key: 0}
|
||||
|
||||
data = {
|
||||
"type": "CHAN",
|
||||
"outgoing": False,
|
||||
"conversation_key": key.upper(), # test case normalization
|
||||
"text": "hello",
|
||||
"sender_timestamp": 1700000000,
|
||||
}
|
||||
await session.on_event_message(data)
|
||||
|
||||
assert len(session._msg_queue) == 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_outgoing_message_ignored(self):
|
||||
session, sent = _make_session()
|
||||
data = {"type": "PRIV", "outgoing": True, "conversation_key": EXAMPLE_KEY}
|
||||
await session.on_event_message(data)
|
||||
assert len(session._msg_queue) == 0
|
||||
assert len(sent) == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_chan_unmapped_dropped(self):
|
||||
session, sent = _make_session()
|
||||
session.key_to_idx = {}
|
||||
data = {
|
||||
"type": "CHAN",
|
||||
"outgoing": False,
|
||||
"conversation_key": "ff" * 16,
|
||||
"text": "hello",
|
||||
"sender_timestamp": 0,
|
||||
}
|
||||
await session.on_event_message(data)
|
||||
assert len(session._msg_queue) == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_contact_event_updates_existing_cache(self):
|
||||
session, sent = _make_session()
|
||||
# Contact must already be in favorites cache to receive pushes
|
||||
session.contacts = [
|
||||
{
|
||||
"public_key": EXAMPLE_KEY,
|
||||
"name": "Old",
|
||||
"type": 1,
|
||||
"favorite": True,
|
||||
"direct_path": None,
|
||||
"direct_path_len": -1,
|
||||
"direct_path_hash_mode": -1,
|
||||
"last_advert": 0,
|
||||
"lat": 0.0,
|
||||
"lon": 0.0,
|
||||
"first_seen": 0,
|
||||
}
|
||||
]
|
||||
|
||||
data = {
|
||||
"public_key": EXAMPLE_KEY,
|
||||
"type": 1,
|
||||
"name": "Updated",
|
||||
"favorite": True,
|
||||
"direct_path": None,
|
||||
"direct_path_len": -1,
|
||||
"direct_path_hash_mode": -1,
|
||||
"last_advert": 100,
|
||||
"lat": 0.0,
|
||||
"lon": 0.0,
|
||||
"first_seen": 0,
|
||||
}
|
||||
await session.on_event_contact(data)
|
||||
assert len(session.contacts) == 1
|
||||
assert session.contacts[0]["name"] == "Updated"
|
||||
# Should have sent a PUSH_NEW_ADVERT
|
||||
payloads = _extract_payloads(sent)
|
||||
assert payloads[0][0] == 0x8A # PUSH_NEW_ADVERT
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_contact_event_ignored_for_non_favorites(self):
|
||||
session, sent = _make_session()
|
||||
session.contacts = []
|
||||
|
||||
data = {
|
||||
"public_key": EXAMPLE_KEY,
|
||||
"type": 1,
|
||||
"name": "Stranger",
|
||||
"favorite": False,
|
||||
}
|
||||
await session.on_event_contact(data)
|
||||
assert len(session.contacts) == 0
|
||||
assert len(sent) == 0
|
||||
Reference in New Issue
Block a user