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()