Mercurial
diff mrjunejune/inference/copilot_sidecar_test.py @ 260:1f9877b637e9
Add Copilot-powered cyberpunk JRPG chat
Integrate the production JRPG chat with Seobeo streaming, Deita persistence, and a Bazel-managed Copilot SDK and LiteLLM inference stack.
Co-authored-by: Copilot <[email protected]>
| author | MrJuneJune <mrjunejune@users.noreply.github.com> |
|---|---|
| date | Wed, 05 Aug 2026 09:19:41 -0700 |
| parents | |
| children | 056790c4fb0d |
line wrap: on
line diff
--- /dev/null Thu Jan 01 00:00:00 1970 +0000 +++ b/mrjunejune/inference/copilot_sidecar_test.py Wed Aug 05 09:19:41 2026 -0700 @@ -0,0 +1,447 @@ +import asyncio +import types +import unittest +from dataclasses import replace + +from mrjunejune.inference.copilot_sidecar import Sidecar, SidecarConfig + + +def event(event_type, **data): + return types.SimpleNamespace( + type=types.SimpleNamespace(value=event_type), + data=types.SimpleNamespace(**data), + ) + + +class FakeSession: + def __init__(self, session_id, behavior="complete"): + self.session_id = session_id + self.behavior = behavior + self.handlers = [] + self.prompts = [] + self.abort_calls = 0 + self.disconnect_calls = 0 + self.disconnect_error = False + + def on(self, handler): + self.handlers.append(handler) + + def unsubscribe(): + if handler in self.handlers: + self.handlers.remove(handler) + + return unsubscribe + + def emit(self, value): + for handler in list(self.handlers): + handler(value) + + async def send(self, prompt): + self.prompts.append(prompt) + message_id = f"message-{len(self.prompts)}" + if self.behavior == "complete": + self.emit( + event( + "assistant.message_delta", + delta_content=f"{prompt}-delta", + message_id=message_id, + ) + ) + self.emit( + event( + "assistant.message", + content=f"{prompt}-answer", + message_id=message_id, + model="test-model", + ) + ) + self.emit( + event( + "assistant.usage", + model="test-model", + input_tokens=3, + output_tokens=5, + ) + ) + self.emit(event("session.idle", aborted=False)) + elif self.behavior == "error": + self.emit( + event( + "session.error", + error_type="provider", + error_code="upstream_error", + message="upstream failed", + status_code=502, + ) + ) + return message_id + + async def abort(self): + self.abort_calls += 1 + if self.behavior == "abort_error": + raise RuntimeError("abort failed") + self.emit(event("session.idle", aborted=True)) + + async def disconnect(self): + self.disconnect_calls += 1 + if self.disconnect_error: + raise RuntimeError("disconnect failed") + + +class FakeClient: + def __init__(self, behavior="complete"): + self.behavior = behavior + self.sessions = {} + self.create_calls = [] + self.resume_calls = [] + self.delete_calls = [] + self.started = False + self.stopped = False + self.resume_started = None + self.resume_release = None + + async def start(self): + self.started = True + + async def stop(self): + self.stopped = True + + async def resume_session(self, session_id, **kwargs): + self.resume_calls.append((session_id, kwargs)) + if self.resume_started is not None: + self.resume_started.set() + if self.resume_release is not None: + await self.resume_release.wait() + if session_id not in self.sessions: + raise LookupError(session_id) + return self.sessions[session_id] + + async def create_session(self, session_id, **kwargs): + self.create_calls.append((session_id, kwargs)) + session = FakeSession(session_id, self.behavior) + self.sessions[session_id] = session + return session + + async def delete_session(self, session_id): + self.delete_calls.append(session_id) + self.sessions.pop(session_id, None) + + +class SidecarTest(unittest.IsolatedAsyncioTestCase): + async def asyncSetUp(self): + self.output = [] + + async def capture(payload): + self.output.append(payload) + + self.capture = capture + self.config = SidecarConfig( + base_url="http://litellm.invalid/v1", + model="test-model", + wire_api="responses", + base_directory="/not-used-by-fake", + ) + + async def make_sidecar(self, behavior="complete"): + client = FakeClient(behavior) + sidecar = Sidecar(client, self.config, self.capture) + await sidecar.start() + return sidecar, client + + async def test_health_and_invalid_protocol(self): + sidecar, _ = await self.make_sidecar() + await sidecar.announce_ready() + await sidecar.dispatch({"command": "health", "request_id": "health-1"}) + await sidecar.dispatch({"command": "unknown", "request_id": "bad-1"}) + + self.assertEqual(self.output[0]["type"], "ready") + self.assertIsNone(self.output[0]["request_id"]) + self.assertEqual(self.output[1]["request_id"], "health-1") + self.assertEqual( + [item["type"] for item in self.output[-2:]], + ["turn.error", "turn.done"], + ) + + async def test_multi_turn_reuses_one_session_and_configures_sdk(self): + sidecar, client = await self.make_sidecar() + for request_id, prompt in (("r1", "first"), ("r2", "second")): + await sidecar.dispatch( + { + "command": "turn.start", + "request_id": request_id, + "conversation_id": "conversation-a", + "prompt": prompt, + } + ) + await sidecar.drain_events() + + self.assertEqual(len(client.create_calls), 1) + self.assertEqual(len(client.resume_calls), 1) + session = client.sessions["conversation-a"] + self.assertEqual(session.prompts, ["first", "second"]) + options = client.create_calls[0][1] + self.assertEqual(options["provider"]["type"], "openai") + self.assertEqual( + options["provider"]["base_url"], "http://litellm.invalid/v1" + ) + self.assertEqual(options["provider"]["wire_api"], "responses") + self.assertEqual(options["model"], "test-model") + self.assertEqual(options["available_tools"], []) + self.assertTrue(options["streaming"]) + self.assertEqual( + [item["request_id"] for item in self.output if item["type"] == "turn.done"], + ["r1", "r2"], + ) + + async def test_concurrent_conversations_keep_correlation(self): + sidecar, _ = await self.make_sidecar() + await asyncio.gather( + sidecar.dispatch( + { + "command": "turn.start", + "request_id": "left-request", + "conversation_id": "left", + "prompt": "left", + } + ), + sidecar.dispatch( + { + "command": "turn.start", + "request_id": "right-request", + "conversation_id": "right", + "prompt": "right", + } + ), + ) + await sidecar.drain_events() + + correlated = { + (item["request_id"], item["conversation_id"]) + for item in self.output + if item["type"] in ("assistant.delta", "assistant.completed") + } + self.assertEqual( + correlated, + {("left-request", "left"), ("right-request", "right")}, + ) + + async def test_abort_finishes_active_turn(self): + sidecar, client = await self.make_sidecar("pending") + await sidecar.dispatch( + { + "command": "turn.start", + "request_id": "turn-request", + "conversation_id": "abort-me", + "prompt": "wait", + } + ) + await sidecar.dispatch( + { + "command": "turn.abort", + "request_id": "abort-request", + "conversation_id": "abort-me", + } + ) + await sidecar.drain_events() + + self.assertEqual(client.sessions["abort-me"].abort_calls, 1) + done = [ + item + for item in self.output + if item["type"] == "turn.done" + and item["request_id"] == "turn-request" + ] + self.assertEqual(done[0]["aborted"], True) + abort_accept = [ + item + for item in self.output + if item["type"] == "turn.accepted" + and item["request_id"] == "abort-request" + ] + self.assertEqual(abort_accept[0]["target_request_id"], "turn-request") + + async def test_session_error_is_terminal(self): + sidecar, _ = await self.make_sidecar("error") + await sidecar.dispatch( + { + "command": "turn.start", + "request_id": "error-request", + "conversation_id": "error-conversation", + "prompt": "fail", + } + ) + await sidecar.drain_events() + + errors = [item for item in self.output if item["type"] == "turn.error"] + done = [item for item in self.output if item["type"] == "turn.done"] + self.assertEqual(errors[0]["error"]["code"], "upstream_error") + self.assertEqual(errors[0]["error"]["status_code"], 502) + self.assertEqual(len(done), 1) + self.assertTrue(done[0]["failed"]) + + async def test_abort_error_terminates_abort_request(self): + sidecar, _ = await self.make_sidecar("abort_error") + await sidecar.dispatch( + { + "command": "turn.start", + "request_id": "active-request", + "conversation_id": "abort-error", + "prompt": "wait", + } + ) + await sidecar.dispatch( + { + "command": "turn.abort", + "request_id": "failed-abort", + "conversation_id": "abort-error", + } + ) + + abort_events = [ + item for item in self.output if item["request_id"] == "failed-abort" + ] + self.assertEqual( + [item["type"] for item in abort_events], + ["turn.accepted", "turn.error", "turn.done"], + ) + self.assertTrue(abort_events[-1]["failed"]) + + async def test_delete_and_shutdown_release_resources(self): + sidecar, client = await self.make_sidecar("pending") + await sidecar.dispatch( + { + "command": "turn.start", + "request_id": "active", + "conversation_id": "delete-me", + "prompt": "wait", + } + ) + session = client.sessions["delete-me"] + await sidecar.dispatch( + { + "command": "conversation.delete", + "request_id": "delete-request", + "conversation_id": "delete-me", + } + ) + await sidecar.dispatch( + { + "command": "shutdown", + "request_id": "shutdown-request", + "conversation_id": None, + } + ) + + self.assertEqual(session.disconnect_calls, 1) + self.assertEqual(client.delete_calls, ["delete-me"]) + self.assertTrue(client.stopped) + self.assertTrue(sidecar.shutting_down) + shutdown = [ + item + for item in self.output + if item["request_id"] == "shutdown-request" + ] + self.assertEqual(shutdown[0]["type"], "turn.done") + + async def test_idle_and_overflow_sessions_are_evicted(self): + self.config = replace( + self.config, + idle_timeout_seconds=3600, + max_sessions=1, + ) + sidecar, client = await self.make_sidecar() + for conversation_id in ("old", "new"): + await sidecar.dispatch( + { + "command": "turn.start", + "request_id": f"request-{conversation_id}", + "conversation_id": conversation_id, + "prompt": conversation_id, + } + ) + await sidecar.drain_events() + await asyncio.sleep(0.01) + + old_session = client.sessions["old"] + old_session.disconnect_error = True + await sidecar.evict_idle_sessions() + self.assertEqual(old_session.disconnect_calls, 1) + self.assertNotIn("old", sidecar._conversations) + self.assertIn("new", sidecar._conversations) + + sidecar._conversations["new"].last_used -= 4000 + await sidecar.evict_idle_sessions() + self.assertEqual(client.sessions["new"].disconnect_calls, 1) + self.assertEqual(sidecar._conversations, {}) + await sidecar.dispatch( + { + "command": "shutdown", + "request_id": "shutdown-eviction", + "conversation_id": None, + } + ) + + async def test_shutdown_waits_for_inflight_session_creation(self): + client = FakeClient("pending") + client.resume_started = asyncio.Event() + client.resume_release = asyncio.Event() + sidecar = Sidecar(client, self.config, self.capture) + await sidecar.start() + + start_task = asyncio.create_task( + sidecar.dispatch( + { + "command": "turn.start", + "request_id": "starting", + "conversation_id": "race", + "prompt": "wait", + } + ) + ) + await client.resume_started.wait() + shutdown_task = asyncio.create_task( + sidecar.dispatch( + { + "command": "shutdown", + "request_id": "shutdown", + "conversation_id": None, + } + ) + ) + await asyncio.sleep(0) + self.assertTrue(sidecar.shutting_down) + client.resume_release.set() + await asyncio.gather(start_task, shutdown_task) + + self.assertTrue(client.stopped) + self.assertEqual(client.sessions["race"].disconnect_calls, 1) + await sidecar.dispatch( + { + "command": "turn.start", + "request_id": "too-late", + "conversation_id": "late", + "prompt": "no", + } + ) + late_error = [ + item + for item in self.output + if item["request_id"] == "too-late" and item["type"] == "turn.error" + ] + self.assertEqual(late_error[0]["error"]["code"], "shutting_down") + + async def test_command_gates_are_released(self): + sidecar, _ = await self.make_sidecar() + for index in range(20): + await sidecar.dispatch( + { + "command": "unknown", + "request_id": f"unknown-{index}", + "conversation_id": f"conversation-{index}", + } + ) + self.assertEqual(sidecar._conversation_gates, {}) + + +if __name__ == "__main__": + unittest.main()