From 884972f9e0a5b859d2560395c9b07e2c7463cc94 Mon Sep 17 00:00:00 2001 From: Jack Kingsman Date: Fri, 27 Feb 2026 16:54:18 -0800 Subject: [PATCH] Add some tests and improve docs --- AGENTS.md | 40 +-- README.md | 2 +- frontend/src/useWebSocket.ts | 8 +- tests/conftest.py | 14 - tests/test_channels_router.py | 257 +++++++++++++++++ tests/test_contacts_router.py | 88 ++++++ tests/test_echo_dedup.py | 166 +++++++++++ tests/test_packet_pipeline.py | 49 ++++ tests/test_packets_router.py | 524 ++++++++++++++++++++++++++++++++++ 9 files changed, 1109 insertions(+), 39 deletions(-) create mode 100644 tests/test_channels_router.py create mode 100644 tests/test_packets_router.py diff --git a/AGENTS.md b/AGENTS.md index 07b13c5..cc4cf18 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -123,8 +123,8 @@ To improve repeater disambiguation in the network visualizer, the backend stores ### Incoming Messages -1. Radio receives message → MeshCore library emits event -2. `event_handlers.py` catches event → stores in database +1. Radio receives raw bytes → `packet_processor.py` parses, decrypts, deduplicates, and stores in database (primary path via `RX_LOG_DATA` event) +2. `event_handlers.py` handles higher-level events (`CONTACT_MSG_RECV`, `ACK`) as a fallback/supplement 3. `ws_manager` broadcasts to connected clients 4. Frontend `useWebSocket` receives → updates React state @@ -270,25 +270,25 @@ All endpoints are prefixed with `/api` (e.g., `/api/health`). | POST | `/api/radio/reconnect` | Manual radio reconnection | | GET | `/api/contacts` | List contacts | | GET | `/api/contacts/repeaters/advert-paths` | List recent unique advert paths for all contacts | -| GET | `/api/contacts/{key}` | Get contact by public key or prefix | -| GET | `/api/contacts/{key}/detail` | Comprehensive contact profile (stats, name history, paths) | -| GET | `/api/contacts/{key}/advert-paths` | List recent unique advert paths for a contact | +| GET | `/api/contacts/{public_key}` | Get contact by public key or prefix | +| GET | `/api/contacts/{public_key}/detail` | Comprehensive contact profile (stats, name history, paths) | +| GET | `/api/contacts/{public_key}/advert-paths` | List recent unique advert paths for a contact | | POST | `/api/contacts` | Create contact (optionally trigger historical DM decrypt) | -| DELETE | `/api/contacts/{key}` | Delete contact | +| DELETE | `/api/contacts/{public_key}` | Delete contact | | POST | `/api/contacts/sync` | Pull from radio | -| POST | `/api/contacts/{key}/add-to-radio` | Push contact to radio | -| POST | `/api/contacts/{key}/remove-from-radio` | Remove contact from radio | -| POST | `/api/contacts/{key}/mark-read` | Mark contact conversation as read | -| POST | `/api/contacts/{key}/command` | Send CLI command to repeater | -| POST | `/api/contacts/{key}/trace` | Trace route to contact | -| POST | `/api/contacts/{key}/repeater/login` | Log in to a repeater | -| POST | `/api/contacts/{key}/repeater/status` | Fetch repeater status telemetry | -| POST | `/api/contacts/{key}/repeater/lpp-telemetry` | Fetch CayenneLPP sensor data | -| POST | `/api/contacts/{key}/repeater/neighbors` | Fetch repeater neighbors | -| POST | `/api/contacts/{key}/repeater/acl` | Fetch repeater ACL | -| POST | `/api/contacts/{key}/repeater/radio-settings` | Fetch radio settings via CLI | -| POST | `/api/contacts/{key}/repeater/advert-intervals` | Fetch advert intervals | -| POST | `/api/contacts/{key}/repeater/owner-info` | Fetch owner info | +| POST | `/api/contacts/{public_key}/add-to-radio` | Push contact to radio | +| POST | `/api/contacts/{public_key}/remove-from-radio` | Remove contact from radio | +| POST | `/api/contacts/{public_key}/mark-read` | Mark contact conversation as read | +| POST | `/api/contacts/{public_key}/command` | Send CLI command to repeater | +| POST | `/api/contacts/{public_key}/trace` | Trace route to contact | +| POST | `/api/contacts/{public_key}/repeater/login` | Log in to a repeater | +| POST | `/api/contacts/{public_key}/repeater/status` | Fetch repeater status telemetry | +| POST | `/api/contacts/{public_key}/repeater/lpp-telemetry` | Fetch CayenneLPP sensor data | +| POST | `/api/contacts/{public_key}/repeater/neighbors` | Fetch repeater neighbors | +| POST | `/api/contacts/{public_key}/repeater/acl` | Fetch repeater ACL | +| POST | `/api/contacts/{public_key}/repeater/radio-settings` | Fetch radio settings via CLI | +| POST | `/api/contacts/{public_key}/repeater/advert-intervals` | Fetch advert intervals | +| POST | `/api/contacts/{public_key}/repeater/owner-info` | Fetch owner info | | GET | `/api/channels` | List channels | | GET | `/api/channels/{key}` | Get channel by key | @@ -343,7 +343,7 @@ All endpoints are prefixed with `/api` (e.g., `/api/health`). Read state (`last_read_at`) is tracked **server-side** for consistency across devices: - Stored as Unix timestamp in `contacts.last_read_at` and `channels.last_read_at` -- Updated via `POST /api/contacts/{key}/mark-read` and `POST /api/channels/{key}/mark-read` +- Updated via `POST /api/contacts/{public_key}/mark-read` and `POST /api/channels/{key}/mark-read` - Bulk update via `POST /api/read-state/mark-all-read` - Aggregated counts via `GET /api/read-state/unreads` (server-side computation) diff --git a/README.md b/README.md index 8e3e61f..c4b2de7 100644 --- a/README.md +++ b/README.md @@ -94,7 +94,7 @@ Access at http://localhost:8000 > **Note:** BLE-in-docker is outside the scope of this README, but the env vars should all still work. -Edit `docker-compose.yaml` to set a serial device for passthrough, or uncomment your transport (serial or TCP).Then: +Edit `docker-compose.yaml` to set a serial device for passthrough, or uncomment your transport (serial or TCP). Then: ```bash docker compose up -d diff --git a/frontend/src/useWebSocket.ts b/frontend/src/useWebSocket.ts index 4a4128a..6477a26 100644 --- a/frontend/src/useWebSocket.ts +++ b/frontend/src/useWebSocket.ts @@ -58,7 +58,7 @@ export function useWebSocket(options: UseWebSocketOptions) { const ws = new WebSocket(wsUrl); ws.onopen = () => { - console.log('WebSocket connected'); + // Connection established (or re-established after disconnect) if (reconnectTimeoutRef.current) { clearTimeout(reconnectTimeoutRef.current); reconnectTimeoutRef.current = null; @@ -70,7 +70,7 @@ export function useWebSocket(options: UseWebSocketOptions) { }; ws.onclose = () => { - console.log('WebSocket disconnected'); + // Connection lost — will auto-reconnect after delay wsRef.current = null; if (!shouldReconnectRef.current) { @@ -82,7 +82,7 @@ export function useWebSocket(options: UseWebSocketOptions) { clearTimeout(reconnectTimeoutRef.current); } reconnectTimeoutRef.current = window.setTimeout(() => { - console.log('Attempting WebSocket reconnect...'); + // Reconnect attempt after disconnect connect(); }, 3000); }; @@ -129,7 +129,7 @@ export function useWebSocket(options: UseWebSocketOptions) { // Heartbeat response, ignore break; default: - console.log('Unknown WebSocket message type:', msg.type); + console.warn('Unknown WebSocket message type:', msg.type); } } catch (e) { console.error('Failed to parse WebSocket message:', e); diff --git a/tests/conftest.py b/tests/conftest.py index c552d0f..99b32cf 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -20,17 +20,3 @@ def cleanup_test_db_dir(): """Clean up temporary pytest DB directory after the test session.""" yield shutil.rmtree(_TEST_DB_DIR, ignore_errors=True) - - -@pytest.fixture -def sample_channel_key(): - """A sample 16-byte channel key for testing.""" - return bytes.fromhex("0123456789abcdef0123456789abcdef") - - -@pytest.fixture -def sample_hashtag_key(): - """A channel key derived from hashtag name '#test'.""" - import hashlib - - return hashlib.sha256(b"#test").digest()[:16] diff --git a/tests/test_channels_router.py b/tests/test_channels_router.py new file mode 100644 index 0000000..13a8149 --- /dev/null +++ b/tests/test_channels_router.py @@ -0,0 +1,257 @@ +"""Tests for the channels router sync endpoint. + +Verifies that POST /api/channels/sync correctly reads channel slots +from the radio and upserts them into the database. +""" + +from contextlib import asynccontextmanager +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest +from meshcore import EventType + +from app.database import Database +from app.radio import radio_manager +from app.repository import ChannelRepository + + +@pytest.fixture +async def test_db(): + """Create an in-memory test database with schema + migrations.""" + import app.repository as repo_module + + db = Database(":memory:") + await db.connect() + + original_db = repo_module.db + repo_module.db = db + + try: + yield db + finally: + repo_module.db = original_db + await db.disconnect() + + +@pytest.fixture(autouse=True) +def _reset_radio_state(): + """Save/restore radio_manager state so tests don't leak.""" + prev = radio_manager._meshcore + prev_lock = radio_manager._operation_lock + yield + radio_manager._meshcore = prev + radio_manager._operation_lock = prev_lock + + +@pytest.fixture +def client(): + """Create an httpx AsyncClient for testing the app.""" + from app.main import app + + transport = httpx.ASGITransport(app=app) + return httpx.AsyncClient(transport=transport, base_url="http://test") + + +def _make_channel_info(name: str, secret: bytes): + """Create a mock channel info response.""" + result = MagicMock() + result.type = EventType.CHANNEL_INFO + result.payload = { + "channel_name": name, + "channel_secret": secret, + } + return result + + +def _make_empty_channel(): + """Create a mock empty channel response.""" + result = MagicMock() + result.type = EventType.CHANNEL_INFO + result.payload = { + "channel_name": "\x00\x00\x00\x00", + "channel_secret": b"", + } + return result + + +def _make_error_response(): + """Create a mock error response (channel slot unused).""" + result = MagicMock() + result.type = EventType.ERROR + result.payload = {} + return result + + +@asynccontextmanager +async def _noop_radio_operation(mc): + """No-op radio_operation context manager that yields mc.""" + yield mc + + +class TestSyncChannelsFromRadio: + """Test POST /api/channels/sync.""" + + @pytest.mark.asyncio + async def test_sync_channels_basic(self, test_db, client): + """Sync creates channels from radio slots.""" + secret_a = bytes.fromhex("0123456789abcdef0123456789abcdef") + secret_b = bytes.fromhex("fedcba9876543210fedcba9876543210") + + mock_mc = MagicMock() + + async def mock_get_channel(idx): + if idx == 0: + return _make_channel_info("#general", secret_a) + if idx == 1: + return _make_channel_info("Private", secret_b) + return _make_empty_channel() + + mock_mc.commands.get_channel = AsyncMock(side_effect=mock_get_channel) + radio_manager._meshcore = mock_mc + + with ( + patch("app.dependencies.radio_manager") as mock_dep_rm, + patch("app.routers.channels.radio_manager") as mock_ch_rm, + ): + mock_dep_rm.is_connected = True + mock_dep_rm.meshcore = mock_mc + mock_ch_rm.radio_operation = lambda desc: _noop_radio_operation(mock_mc) + + response = await client.post("/api/channels/sync?max_channels=5") + + assert response.status_code == 200 + data = response.json() + assert data["synced"] == 2 + + # Verify channels in DB + channels = await ChannelRepository.get_all() + assert len(channels) == 2 + + keys = {ch.key for ch in channels} + assert secret_a.hex().upper() in keys + assert secret_b.hex().upper() in keys + + @pytest.mark.asyncio + async def test_sync_skips_empty_channels(self, test_db, client): + """Empty channel slots are skipped during sync.""" + secret = bytes.fromhex("aabbccddaabbccddaabbccddaabbccdd") + mock_mc = MagicMock() + + async def mock_get_channel(idx): + if idx == 0: + return _make_channel_info("#test", secret) + return _make_empty_channel() + + mock_mc.commands.get_channel = AsyncMock(side_effect=mock_get_channel) + radio_manager._meshcore = mock_mc + + with ( + patch("app.dependencies.radio_manager") as mock_dep_rm, + patch("app.routers.channels.radio_manager") as mock_ch_rm, + ): + mock_dep_rm.is_connected = True + mock_dep_rm.meshcore = mock_mc + mock_ch_rm.radio_operation = lambda desc: _noop_radio_operation(mock_mc) + + response = await client.post("/api/channels/sync?max_channels=5") + + assert response.status_code == 200 + assert response.json()["synced"] == 1 + + @pytest.mark.asyncio + async def test_sync_hashtag_flag(self, test_db, client): + """Channels starting with # are marked as hashtag channels.""" + secret = bytes.fromhex("1122334455667788aabbccddeeff0011") + mock_mc = MagicMock() + + async def mock_get_channel(idx): + if idx == 0: + return _make_channel_info("#hashtag-room", secret) + return _make_empty_channel() + + mock_mc.commands.get_channel = AsyncMock(side_effect=mock_get_channel) + radio_manager._meshcore = mock_mc + + with ( + patch("app.dependencies.radio_manager") as mock_dep_rm, + patch("app.routers.channels.radio_manager") as mock_ch_rm, + ): + mock_dep_rm.is_connected = True + mock_dep_rm.meshcore = mock_mc + mock_ch_rm.radio_operation = lambda desc: _noop_radio_operation(mock_mc) + + response = await client.post("/api/channels/sync?max_channels=3") + + assert response.status_code == 200 + + channel = await ChannelRepository.get_by_key(secret.hex().upper()) + assert channel is not None + assert channel.is_hashtag is True + assert channel.name == "#hashtag-room" + assert channel.on_radio is True + + @pytest.mark.asyncio + async def test_sync_marks_channels_on_radio(self, test_db, client): + """Synced channels have on_radio=True.""" + secret = bytes.fromhex("aabbccddaabbccddaabbccddaabbccdd") + mock_mc = MagicMock() + + async def mock_get_channel(idx): + if idx == 0: + return _make_channel_info("MyChannel", secret) + return _make_empty_channel() + + mock_mc.commands.get_channel = AsyncMock(side_effect=mock_get_channel) + radio_manager._meshcore = mock_mc + + with ( + patch("app.dependencies.radio_manager") as mock_dep_rm, + patch("app.routers.channels.radio_manager") as mock_ch_rm, + ): + mock_dep_rm.is_connected = True + mock_dep_rm.meshcore = mock_mc + mock_ch_rm.radio_operation = lambda desc: _noop_radio_operation(mock_mc) + + await client.post("/api/channels/sync?max_channels=3") + + channel = await ChannelRepository.get_by_key(secret.hex().upper()) + assert channel.on_radio is True + + @pytest.mark.asyncio + async def test_sync_requires_connection(self, test_db, client): + """Sync returns 503 when radio is not connected.""" + with patch("app.dependencies.radio_manager") as mock_rm: + mock_rm.is_connected = False + mock_rm.meshcore = None + + response = await client.post("/api/channels/sync") + + assert response.status_code == 503 + + @pytest.mark.asyncio + async def test_sync_key_normalized_uppercase(self, test_db, client): + """Channel keys are normalized to uppercase hex.""" + secret = bytes.fromhex("aabbccddaabbccddaabbccddaabbccdd") + mock_mc = MagicMock() + + async def mock_get_channel(idx): + if idx == 0: + return _make_channel_info("Test", secret) + return _make_empty_channel() + + mock_mc.commands.get_channel = AsyncMock(side_effect=mock_get_channel) + radio_manager._meshcore = mock_mc + + with ( + patch("app.dependencies.radio_manager") as mock_dep_rm, + patch("app.routers.channels.radio_manager") as mock_ch_rm, + ): + mock_dep_rm.is_connected = True + mock_dep_rm.meshcore = mock_mc + mock_ch_rm.radio_operation = lambda desc: _noop_radio_operation(mock_mc) + + await client.post("/api/channels/sync?max_channels=3") + + channel = await ChannelRepository.get_by_key("AABBCCDDAABBCCDDAABBCCDDAABBCCDD") + assert channel is not None diff --git a/tests/test_contacts_router.py b/tests/test_contacts_router.py index 5b8e1e4..b682200 100644 --- a/tests/test_contacts_router.py +++ b/tests/test_contacts_router.py @@ -546,6 +546,94 @@ class TestSyncContacts: assert messages[0].conversation_key == KEY_A.lower() +class TestCreateContactWithHistorical: + """Test POST /api/contacts with try_historical=true.""" + + @pytest.mark.asyncio + async def test_new_contact_triggers_historical_decrypt(self, test_db, client): + """Creating a new contact with try_historical triggers DM decryption.""" + with patch( + "app.routers.contacts.start_historical_dm_decryption", new_callable=AsyncMock + ) as mock_start: + response = await client.post( + "/api/contacts", + json={"public_key": KEY_A, "name": "Alice", "try_historical": True}, + ) + + assert response.status_code == 200 + assert response.json()["public_key"] == KEY_A + + mock_start.assert_awaited_once() + # Verify correct args: (background_tasks, public_key, name) + call_args = mock_start.call_args + assert call_args[0][1] == KEY_A # public_key + assert call_args[0][2] == "Alice" # display_name + + @pytest.mark.asyncio + async def test_new_contact_without_historical(self, test_db, client): + """Creating a new contact without try_historical does not trigger decryption.""" + with patch( + "app.routers.contacts.start_historical_dm_decryption", new_callable=AsyncMock + ) as mock_start: + response = await client.post( + "/api/contacts", + json={"public_key": KEY_A, "name": "Alice", "try_historical": False}, + ) + + assert response.status_code == 200 + mock_start.assert_not_awaited() + + @pytest.mark.asyncio + async def test_existing_contact_with_historical(self, test_db, client): + """Existing contact with try_historical still triggers decryption.""" + await _insert_contact(KEY_A, "Alice") + + with patch( + "app.routers.contacts.start_historical_dm_decryption", new_callable=AsyncMock + ) as mock_start: + response = await client.post( + "/api/contacts", + json={"public_key": KEY_A, "name": "Alice", "try_historical": True}, + ) + + assert response.status_code == 200 + mock_start.assert_awaited_once() + + @pytest.mark.asyncio + async def test_existing_contact_updates_name_and_decrypts(self, test_db, client): + """Existing contact with try_historical updates name AND triggers decryption.""" + await _insert_contact(KEY_A, "OldName") + + with patch( + "app.routers.contacts.start_historical_dm_decryption", new_callable=AsyncMock + ) as mock_start: + response = await client.post( + "/api/contacts", + json={"public_key": KEY_A, "name": "NewName", "try_historical": True}, + ) + + assert response.status_code == 200 + mock_start.assert_awaited_once() + + # Verify name was also updated + contact = await ContactRepository.get_by_key(KEY_A) + assert contact.name == "NewName" + + @pytest.mark.asyncio + async def test_default_try_historical_is_false(self, test_db, client): + """try_historical defaults to false when not provided.""" + with patch( + "app.routers.contacts.start_historical_dm_decryption", new_callable=AsyncMock + ) as mock_start: + response = await client.post( + "/api/contacts", + json={"public_key": KEY_A, "name": "Alice"}, + ) + + assert response.status_code == 200 + mock_start.assert_not_awaited() + + class TestAddRemoveRadio: """Test add-to-radio and remove-from-radio endpoints.""" diff --git a/tests/test_echo_dedup.py b/tests/test_echo_dedup.py index 6bf97f4..1a0735d 100644 --- a/tests/test_echo_dedup.py +++ b/tests/test_echo_dedup.py @@ -890,3 +890,169 @@ class TestConcurrentDMDedup: msg_type="CHAN", conversation_key=CHANNEL_KEY, limit=10 ) assert len(messages) == 1 + + +class TestMessageAckedBroadcastShape: + """Verify that message_acked broadcasts from _handle_duplicate_message + match the frontend's MessageAckedEvent interface. + + The on_ack handler (event_handlers.py) broadcasts {message_id, ack_count}, + while _handle_duplicate_message broadcasts {message_id, ack_count, paths}. + Both must match what the frontend expects in useWebSocket.ts. + """ + + # Frontend MessageAckedEvent keys (from useWebSocket.ts:113-117) + # The 'paths' key is optional in the TypeScript interface + REQUIRED_KEYS = {"message_id", "ack_count"} + OPTIONAL_KEYS = {"paths"} + + @pytest.mark.asyncio + async def test_outgoing_echo_broadcast_shape(self, test_db, captured_broadcasts): + """Outgoing echo broadcast has all required keys plus paths.""" + from app.packet_processor import create_message_from_decrypted + + msg_id = await MessageRepository.create( + msg_type="CHAN", + text="Sender: Shape test", + conversation_key=CHANNEL_KEY, + sender_timestamp=SENDER_TIMESTAMP, + received_at=SENDER_TIMESTAMP, + outgoing=True, + ) + + pkt_id, _ = await RawPacketRepository.create(b"shape_echo", SENDER_TIMESTAMP + 1) + broadcasts, mock_broadcast = captured_broadcasts + + with patch("app.packet_processor.broadcast_event", mock_broadcast): + await create_message_from_decrypted( + packet_id=pkt_id, + channel_key=CHANNEL_KEY, + sender="Sender", + message_text="Shape test", + timestamp=SENDER_TIMESTAMP, + received_at=SENDER_TIMESTAMP + 1, + path="aabb", + ) + + ack_broadcasts = [b for b in broadcasts if b["type"] == "message_acked"] + assert len(ack_broadcasts) == 1 + + payload = ack_broadcasts[0]["data"] + payload_keys = set(payload.keys()) + + # Must have all required keys + assert payload_keys >= self.REQUIRED_KEYS + # Must only have expected keys + assert payload_keys <= (self.REQUIRED_KEYS | self.OPTIONAL_KEYS) + + # Verify types + assert isinstance(payload["message_id"], int) + assert isinstance(payload["ack_count"], int) + assert payload["message_id"] == msg_id + assert payload["ack_count"] == 1 + + # paths should be a list of dicts with path and received_at keys + assert isinstance(payload["paths"], list) + for p in payload["paths"]: + assert "path" in p + assert "received_at" in p + + @pytest.mark.asyncio + async def test_incoming_echo_broadcast_shape(self, test_db, captured_broadcasts): + """Incoming echo broadcast (with path) has the correct shape.""" + from app.packet_processor import create_message_from_decrypted + + pkt1, _ = await RawPacketRepository.create(b"shape_inc_1", SENDER_TIMESTAMP) + + broadcasts, mock_broadcast = captured_broadcasts + + with patch("app.packet_processor.broadcast_event", mock_broadcast): + await create_message_from_decrypted( + packet_id=pkt1, + channel_key=CHANNEL_KEY, + sender="Other", + message_text="Incoming shape", + timestamp=SENDER_TIMESTAMP, + received_at=SENDER_TIMESTAMP, + path="aa", + ) + + broadcasts.clear() + + pkt2, _ = await RawPacketRepository.create(b"shape_inc_2", SENDER_TIMESTAMP + 1) + + with patch("app.packet_processor.broadcast_event", mock_broadcast): + await create_message_from_decrypted( + packet_id=pkt2, + channel_key=CHANNEL_KEY, + sender="Other", + message_text="Incoming shape", + timestamp=SENDER_TIMESTAMP, + received_at=SENDER_TIMESTAMP + 1, + path="bbcc", + ) + + ack_broadcasts = [b for b in broadcasts if b["type"] == "message_acked"] + assert len(ack_broadcasts) == 1 + + payload = ack_broadcasts[0]["data"] + payload_keys = set(payload.keys()) + + assert payload_keys >= self.REQUIRED_KEYS + assert payload_keys <= (self.REQUIRED_KEYS | self.OPTIONAL_KEYS) + assert payload["ack_count"] == 0 # Not outgoing, no ack increment + + @pytest.mark.asyncio + async def test_dm_echo_broadcast_shape(self, test_db, captured_broadcasts): + """DM duplicate broadcast has the same shape as channel echo.""" + from app.packet_processor import create_dm_message_from_decrypted + + pkt1, _ = await RawPacketRepository.create(b"dm_shape_1", SENDER_TIMESTAMP) + decrypted = DecryptedDirectMessage( + timestamp=SENDER_TIMESTAMP, + flags=0, + message="DM shape test", + dest_hash="fa", + src_hash="a1", + ) + + broadcasts, mock_broadcast = captured_broadcasts + + with patch("app.packet_processor.broadcast_event", mock_broadcast): + msg_id = await create_dm_message_from_decrypted( + packet_id=pkt1, + decrypted=decrypted, + their_public_key=CONTACT_PUB, + our_public_key=OUR_PUB, + received_at=SENDER_TIMESTAMP, + outgoing=True, + path="aabb", + ) + + assert msg_id is not None + broadcasts.clear() + + pkt2, _ = await RawPacketRepository.create(b"dm_shape_2", SENDER_TIMESTAMP + 1) + + with patch("app.packet_processor.broadcast_event", mock_broadcast): + await create_dm_message_from_decrypted( + packet_id=pkt2, + decrypted=decrypted, + their_public_key=CONTACT_PUB, + our_public_key=OUR_PUB, + received_at=SENDER_TIMESTAMP + 1, + outgoing=True, + path="ccddee", + ) + + ack_broadcasts = [b for b in broadcasts if b["type"] == "message_acked"] + assert len(ack_broadcasts) == 1 + + payload = ack_broadcasts[0]["data"] + payload_keys = set(payload.keys()) + + assert payload_keys >= self.REQUIRED_KEYS + assert payload_keys <= (self.REQUIRED_KEYS | self.OPTIONAL_KEYS) + assert isinstance(payload["message_id"], int) + assert isinstance(payload["ack_count"], int) + assert payload["ack_count"] == 1 # Outgoing DM echo increments ack diff --git a/tests/test_packet_pipeline.py b/tests/test_packet_pipeline.py index 80b784c..84e719c 100644 --- a/tests/test_packet_pipeline.py +++ b/tests/test_packet_pipeline.py @@ -387,6 +387,55 @@ class TestAdvertisementPipeline: contact = await ContactRepository.get_by_key(test_pubkey) assert contact.last_path_len == 1 # Still the shorter path + @pytest.mark.asyncio + async def test_advertisement_default_path_len_treated_as_infinity( + self, test_db, captured_broadcasts + ): + """Contact with last_path_len=-1 (unset) is treated as infinite length. + + Any new advertisement should replace the default -1 path since + the code converts -1 to float('inf') for comparison. + """ + from app.packet_processor import _process_advertisement + + test_pubkey = "1234567890abcdef1234567890abcdef1234567890abcdef1234567890abcdef" + await ContactRepository.upsert( + { + "public_key": test_pubkey, + "name": "TestNode", + "type": 1, + "last_seen": 1000, + "last_path_len": -1, # Default unset value + "last_path": None, + } + ) + + from app.decoder import ParsedAdvertisement + + broadcasts, mock_broadcast = captured_broadcasts + + packet_info = MagicMock() + packet_info.path_length = 3 + packet_info.path = bytes.fromhex("aabbcc") + + with patch("app.packet_processor.broadcast_event", mock_broadcast): + with patch("app.packet_processor.parse_advertisement") as mock_parse: + mock_parse.return_value = ParsedAdvertisement( + public_key=test_pubkey, + name="TestNode", + timestamp=1050, + lat=None, + lon=None, + device_role=1, + ) + # Process within 60s window (last_seen=1000, now=1050) + await _process_advertisement(b"", timestamp=1050, packet_info=packet_info) + + # Since -1 is treated as infinity, the new path (len=3) should replace it + contact = await ContactRepository.get_by_key(test_pubkey) + assert contact.last_path_len == 3 + assert contact.last_path == "aabbcc" + @pytest.mark.asyncio async def test_advertisement_replaces_stale_path_outside_window( self, test_db, captured_broadcasts diff --git a/tests/test_packets_router.py b/tests/test_packets_router.py new file mode 100644 index 0000000..a24ae95 --- /dev/null +++ b/tests/test_packets_router.py @@ -0,0 +1,524 @@ +"""Tests for the packets router. + +Covers the historical channel decryption endpoint, background task, +undecrypted count endpoint, and the maintenance endpoint. +""" + +import time +from unittest.mock import patch + +import httpx +import pytest + +from app.database import Database +from app.repository import ChannelRepository, MessageRepository, RawPacketRepository + + +@pytest.fixture +async def test_db(): + """Create an in-memory test database with schema + migrations.""" + import app.repository as repo_module + + db = Database(":memory:") + await db.connect() + + original_db = repo_module.db + repo_module.db = db + + # Also patch the db reference used by the packets router for VACUUM + import app.routers.packets as packets_module + + original_packets_db = packets_module.db + packets_module.db = db + + try: + yield db + finally: + repo_module.db = original_db + packets_module.db = original_packets_db + await db.disconnect() + + +@pytest.fixture +def client(): + """Create an httpx AsyncClient for testing the app.""" + from app.main import app + + transport = httpx.ASGITransport(app=app) + return httpx.AsyncClient(transport=transport, base_url="http://test") + + +async def _insert_raw_packets(count: int, decrypted: bool = False, age_days: int = 0) -> list[int]: + """Insert raw packets and return their IDs.""" + ids = [] + base_ts = int(time.time()) - (age_days * 86400) + for i in range(count): + packet_id, _ = await RawPacketRepository.create( + f"packet_data_{i}_{age_days}_{decrypted}".encode(), base_ts + i + ) + if decrypted: + # Create a message and link it + msg_id = await MessageRepository.create( + msg_type="CHAN", + text=f"decrypted msg {i}", + conversation_key="DEADBEEF" * 4, + sender_timestamp=base_ts + i, + received_at=base_ts + i, + ) + if msg_id is not None: + await RawPacketRepository.mark_decrypted(packet_id, msg_id) + ids.append(packet_id) + return ids + + +class TestUndecryptedCount: + """Test GET /api/packets/undecrypted/count.""" + + @pytest.mark.asyncio + async def test_returns_zero_when_empty(self, test_db, client): + response = await client.get("/api/packets/undecrypted/count") + + assert response.status_code == 200 + assert response.json()["count"] == 0 + + @pytest.mark.asyncio + async def test_counts_only_undecrypted(self, test_db, client): + await _insert_raw_packets(3, decrypted=False) + await _insert_raw_packets(2, decrypted=True) + + response = await client.get("/api/packets/undecrypted/count") + + assert response.status_code == 200 + assert response.json()["count"] == 3 + + +class TestDecryptHistoricalPackets: + """Test POST /api/packets/decrypt/historical.""" + + @pytest.mark.asyncio + async def test_channel_decrypt_with_hex_key(self, test_db, client): + """Channel decryption with a valid hex key starts background task.""" + await _insert_raw_packets(5) + + response = await client.post( + "/api/packets/decrypt/historical", + json={ + "key_type": "channel", + "channel_key": "0123456789abcdef0123456789abcdef", + }, + ) + + assert response.status_code == 200 + data = response.json() + assert data["started"] is True + assert data["total_packets"] == 5 + assert "background" in data["message"].lower() + + @pytest.mark.asyncio + async def test_channel_decrypt_with_hashtag_name(self, test_db, client): + """Channel decryption with a channel name derives key from hash.""" + await _insert_raw_packets(3) + + response = await client.post( + "/api/packets/decrypt/historical", + json={ + "key_type": "channel", + "channel_name": "#general", + }, + ) + + assert response.status_code == 200 + data = response.json() + assert data["started"] is True + assert data["total_packets"] == 3 + + @pytest.mark.asyncio + async def test_channel_decrypt_invalid_hex(self, test_db, client): + """Invalid hex string for channel key returns error.""" + response = await client.post( + "/api/packets/decrypt/historical", + json={ + "key_type": "channel", + "channel_key": "not_valid_hex", + }, + ) + + assert response.status_code == 200 + data = response.json() + assert data["started"] is False + assert "invalid" in data["message"].lower() + + @pytest.mark.asyncio + async def test_channel_decrypt_wrong_key_length(self, test_db, client): + """Channel key with wrong length returns error.""" + response = await client.post( + "/api/packets/decrypt/historical", + json={ + "key_type": "channel", + "channel_key": "aabbccdd", # Only 4 bytes, need 16 + }, + ) + + assert response.status_code == 200 + data = response.json() + assert data["started"] is False + assert "16 bytes" in data["message"] + + @pytest.mark.asyncio + async def test_channel_decrypt_no_key_or_name(self, test_db, client): + """Channel decryption without key or name returns error.""" + response = await client.post( + "/api/packets/decrypt/historical", + json={"key_type": "channel"}, + ) + + assert response.status_code == 200 + data = response.json() + assert data["started"] is False + assert "must provide" in data["message"].lower() + + @pytest.mark.asyncio + async def test_channel_decrypt_no_undecrypted_packets(self, test_db, client): + """Channel decryption with no undecrypted packets returns not started.""" + response = await client.post( + "/api/packets/decrypt/historical", + json={ + "key_type": "channel", + "channel_key": "0123456789abcdef0123456789abcdef", + }, + ) + + assert response.status_code == 200 + data = response.json() + assert data["started"] is False + assert data["total_packets"] == 0 + + @pytest.mark.asyncio + async def test_channel_decrypt_resolves_channel_name(self, test_db, client): + """Channel decryption finds display name from DB when channel exists.""" + key_hex = "0123456789ABCDEF0123456789ABCDEF" + await ChannelRepository.upsert(key=key_hex, name="#test-channel", is_hashtag=True) + await _insert_raw_packets(1) + + response = await client.post( + "/api/packets/decrypt/historical", + json={ + "key_type": "channel", + "channel_key": key_hex.lower(), + }, + ) + + assert response.status_code == 200 + assert response.json()["started"] is True + + @pytest.mark.asyncio + async def test_contact_decrypt_missing_private_key(self, test_db, client): + """Contact decryption without private key returns error.""" + response = await client.post( + "/api/packets/decrypt/historical", + json={ + "key_type": "contact", + "contact_public_key": "aa" * 32, + }, + ) + + assert response.status_code == 200 + data = response.json() + assert data["started"] is False + assert "private_key" in data["message"].lower() + + @pytest.mark.asyncio + async def test_contact_decrypt_missing_contact_key(self, test_db, client): + """Contact decryption without contact public key returns error.""" + response = await client.post( + "/api/packets/decrypt/historical", + json={ + "key_type": "contact", + "private_key": "aa" * 64, + }, + ) + + assert response.status_code == 200 + data = response.json() + assert data["started"] is False + assert "contact_public_key" in data["message"].lower() + + @pytest.mark.asyncio + async def test_contact_decrypt_wrong_private_key_length(self, test_db, client): + """Private key with wrong length returns error.""" + response = await client.post( + "/api/packets/decrypt/historical", + json={ + "key_type": "contact", + "private_key": "aa" * 32, # 32 bytes, need 64 + "contact_public_key": "bb" * 32, + }, + ) + + assert response.status_code == 200 + data = response.json() + assert data["started"] is False + assert "64 bytes" in data["message"] + + @pytest.mark.asyncio + async def test_contact_decrypt_wrong_public_key_length(self, test_db, client): + """Contact public key with wrong length returns error.""" + response = await client.post( + "/api/packets/decrypt/historical", + json={ + "key_type": "contact", + "private_key": "aa" * 64, + "contact_public_key": "bb" * 16, # 16 bytes, need 32 + }, + ) + + assert response.status_code == 200 + data = response.json() + assert data["started"] is False + assert "32 bytes" in data["message"] + + @pytest.mark.asyncio + async def test_contact_decrypt_invalid_hex(self, test_db, client): + """Invalid hex for private key returns error.""" + response = await client.post( + "/api/packets/decrypt/historical", + json={ + "key_type": "contact", + "private_key": "zz" * 64, + "contact_public_key": "bb" * 32, + }, + ) + + assert response.status_code == 200 + data = response.json() + assert data["started"] is False + assert "invalid" in data["message"].lower() + + @pytest.mark.asyncio + async def test_invalid_key_type(self, test_db, client): + """Invalid key_type returns error.""" + response = await client.post( + "/api/packets/decrypt/historical", + json={"key_type": "invalid"}, + ) + + assert response.status_code == 200 + data = response.json() + assert data["started"] is False + assert "key_type" in data["message"].lower() + + +class TestRunHistoricalChannelDecryption: + """Test the _run_historical_channel_decryption background task.""" + + @pytest.mark.asyncio + async def test_decrypts_matching_packets(self, test_db): + """Background task decrypts packets that match the channel key.""" + from app.routers.packets import _run_historical_channel_decryption + + # Insert undecrypted packets + await _insert_raw_packets(3) + channel_key_hex = "AABBCCDDAABBCCDDAABBCCDDAABBCCDD" + channel_key_bytes = bytes.fromhex(channel_key_hex) + + # Each packet must have unique content to avoid message deduplication + call_count = 0 + + def make_unique_result(*_args, **_kwargs): + nonlocal call_count + call_count += 1 + return type( + "DecryptResult", + (), + { + "sender": f"User{call_count}", + "message": f"Hello {call_count}", + "timestamp": 1700000000 + call_count, + }, + )() + + with ( + patch( + "app.routers.packets.try_decrypt_packet_with_channel_key", + side_effect=make_unique_result, + ), + patch( + "app.routers.packets.parse_packet", + return_value=None, + ), + patch("app.routers.packets.broadcast_success") as mock_success, + ): + await _run_historical_channel_decryption(channel_key_bytes, channel_key_hex, "#test") + + mock_success.assert_called_once() + assert "3" in mock_success.call_args[0][1] # "Decrypted 3 messages" + + @pytest.mark.asyncio + async def test_skips_non_matching_packets(self, test_db): + """Background task skips packets that don't match the channel key.""" + from app.routers.packets import _run_historical_channel_decryption + + await _insert_raw_packets(2) + channel_key_hex = "AABBCCDDAABBCCDDAABBCCDDAABBCCDD" + channel_key_bytes = bytes.fromhex(channel_key_hex) + + with ( + patch( + "app.routers.packets.try_decrypt_packet_with_channel_key", + return_value=None, # No match + ), + patch("app.routers.packets.broadcast_success") as mock_success, + ): + await _run_historical_channel_decryption(channel_key_bytes, channel_key_hex, "#test") + + # No success broadcast when nothing was decrypted + mock_success.assert_not_called() + + @pytest.mark.asyncio + async def test_no_packets_returns_early(self, test_db): + """Background task returns early when no undecrypted packets exist.""" + from app.routers.packets import _run_historical_channel_decryption + + channel_key_hex = "AABBCCDDAABBCCDDAABBCCDDAABBCCDD" + channel_key_bytes = bytes.fromhex(channel_key_hex) + + with patch("app.routers.packets.broadcast_success") as mock_success: + await _run_historical_channel_decryption(channel_key_bytes, channel_key_hex) + + mock_success.assert_not_called() + + @pytest.mark.asyncio + async def test_display_name_fallback(self, test_db): + """Uses channel key prefix when no display name is provided.""" + from app.routers.packets import _run_historical_channel_decryption + + await _insert_raw_packets(1) + channel_key_hex = "AABBCCDDAABBCCDDAABBCCDDAABBCCDD" + channel_key_bytes = bytes.fromhex(channel_key_hex) + + mock_result = type( + "DecryptResult", + (), + { + "sender": "User", + "message": "msg", + "timestamp": 1700000000, + }, + )() + + with ( + patch( + "app.routers.packets.try_decrypt_packet_with_channel_key", + return_value=mock_result, + ), + patch("app.routers.packets.parse_packet", return_value=None), + patch("app.routers.packets.broadcast_success") as mock_success, + ): + await _run_historical_channel_decryption( + channel_key_bytes, + channel_key_hex, + None, # No display name + ) + + # Should use key prefix as display name + call_msg = mock_success.call_args[0][0] + assert channel_key_hex[:12] in call_msg + + +class TestMaintenanceEndpoint: + """Test POST /api/packets/maintenance.""" + + @pytest.mark.asyncio + async def test_prune_old_undecrypted(self, test_db, client): + """Prune deletes undecrypted packets older than threshold.""" + await _insert_raw_packets(3, decrypted=False, age_days=30) + await _insert_raw_packets(2, decrypted=False, age_days=0) + + response = await client.post( + "/api/packets/maintenance", + json={"prune_undecrypted_days": 7}, + ) + + assert response.status_code == 200 + data = response.json() + assert data["packets_deleted"] == 3 + + # Verify only recent packets remain + remaining = await RawPacketRepository.get_undecrypted_count() + assert remaining == 2 + + @pytest.mark.asyncio + async def test_purge_linked_raw_packets(self, test_db, client): + """Purge deletes raw packets that are linked to stored messages.""" + await _insert_raw_packets(3, decrypted=True) + await _insert_raw_packets(2, decrypted=False) + + response = await client.post( + "/api/packets/maintenance", + json={"purge_linked_raw_packets": True}, + ) + + assert response.status_code == 200 + data = response.json() + assert data["packets_deleted"] == 3 + + # Undecrypted packets should remain + remaining = await RawPacketRepository.get_undecrypted_count() + assert remaining == 2 + + @pytest.mark.asyncio + async def test_both_prune_and_purge(self, test_db, client): + """Both prune and purge can run in a single request.""" + await _insert_raw_packets(2, decrypted=True) + await _insert_raw_packets(3, decrypted=False, age_days=30) + await _insert_raw_packets(1, decrypted=False, age_days=0) + + response = await client.post( + "/api/packets/maintenance", + json={ + "prune_undecrypted_days": 7, + "purge_linked_raw_packets": True, + }, + ) + + assert response.status_code == 200 + data = response.json() + # 2 linked + 3 old undecrypted = 5 deleted + assert data["packets_deleted"] == 5 + + @pytest.mark.asyncio + async def test_no_options_deletes_nothing(self, test_db, client): + """No options specified means no deletions (only vacuum).""" + await _insert_raw_packets(5) + + response = await client.post( + "/api/packets/maintenance", + json={}, + ) + + assert response.status_code == 200 + data = response.json() + assert data["packets_deleted"] == 0 + + @pytest.mark.asyncio + async def test_vacuum_reports_status(self, test_db, client): + """Maintenance endpoint reports vacuum status.""" + response = await client.post( + "/api/packets/maintenance", + json={}, + ) + + assert response.status_code == 200 + data = response.json() + # vacuumed is a boolean (may be True or False depending on DB state) + assert isinstance(data["vacuumed"], bool) + + @pytest.mark.asyncio + async def test_prune_days_validation(self, test_db, client): + """prune_undecrypted_days must be >= 1.""" + response = await client.post( + "/api/packets/maintenance", + json={"prune_undecrypted_days": 0}, + ) + + assert response.status_code == 422