from __future__ import annotations import json import socket import tempfile import unittest from pathlib import Path from support import make_config from netbox_store_agent.daemon import handle_connection from netbox_store_agent.protocol import response class FakeService: def __init__(self) -> None: self.requests = [] def handle(self, request: object) -> dict[str, object]: self.requests.append(request) return response(200, {"ok": True}) class DaemonFramingTests(unittest.TestCase): def setUp(self) -> None: self.temporary = tempfile.TemporaryDirectory() self.config = make_config(Path(self.temporary.name), require_peers=False) def tearDown(self) -> None: self.temporary.cleanup() def exchange(self, payload: bytes) -> dict[str, object]: server, client = socket.socketpair() try: client.sendall(payload) service = FakeService() handle_connection(server, self.config, service) # type: ignore[arg-type] data = bytearray() while b"\n" not in data: chunk = client.recv(4096) if not chunk: break data.extend(chunk) self.assertEqual(data.count(b"\n"), 1) return json.loads(bytes(data).decode()) finally: client.close() def test_one_request_one_response(self) -> None: result = self.exchange( b'{"protocol_version":1,"method":"GET","path":"/v1/capabilities"}\n' ) self.assertEqual(result["status"], 200) self.assertEqual(set(result), {"protocol_version", "status", "body"}) def test_second_json_line_rejected(self) -> None: result = self.exchange( b'{"protocol_version":1,"method":"GET","path":"/v1/capabilities"}\n{}\n' ) self.assertEqual(result["status"], 400) def test_missing_newline_rejected_when_client_half_closes(self) -> None: server, client = socket.socketpair() try: client.sendall(b"{}") client.shutdown(socket.SHUT_WR) service = FakeService() handle_connection(server, self.config, service) # type: ignore[arg-type] result = json.loads(client.recv(4096).decode()) self.assertEqual(result["status"], 400) finally: client.close() def test_oversized_frame_rejected(self) -> None: result = self.exchange(b"x" * 65536 + b"\n") self.assertEqual(result["status"], 400) if __name__ == "__main__": unittest.main()