diff mrjunejune/inference/sdk_provider_integration_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
line wrap: on
line diff
--- /dev/null	Thu Jan 01 00:00:00 1970 +0000
+++ b/mrjunejune/inference/sdk_provider_integration_test.py	Wed Aug 05 09:19:41 2026 -0700
@@ -0,0 +1,166 @@
+import asyncio
+import json
+import os
+import sys
+import tempfile
+import threading
+import time
+import unittest
+from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
+
+from copilot import CopilotClient, RuntimeConnection
+from copilot.rpc import PermissionDecisionReject
+
+CLI_PATH = os.path.abspath(sys.argv.pop(1))
+
+
+class ProviderHandler(BaseHTTPRequestHandler):
+    requests = []
+
+    def do_POST(self) -> None:
+        length = int(self.headers.get("Content-Length", "0"))
+        payload = json.loads(self.rfile.read(length))
+        self.__class__.requests.append((self.path, payload))
+        response_id = f"chatcmpl-{len(self.__class__.requests)}"
+        chunks = [
+            {
+                "id": response_id,
+                "object": "chat.completion.chunk",
+                "created": int(time.time()),
+                "model": "gpt-4",
+                "choices": [
+                    {
+                        "index": 0,
+                        "delta": {"role": "assistant", "content": "hello "},
+                        "finish_reason": None,
+                    }
+                ],
+            },
+            {
+                "id": response_id,
+                "object": "chat.completion.chunk",
+                "created": int(time.time()),
+                "model": "gpt-4",
+                "choices": [
+                    {
+                        "index": 0,
+                        "delta": {"content": "traveler"},
+                        "finish_reason": None,
+                    }
+                ],
+            },
+            {
+                "id": response_id,
+                "object": "chat.completion.chunk",
+                "created": int(time.time()),
+                "model": "gpt-4",
+                "choices": [
+                    {
+                        "index": 0,
+                        "delta": {},
+                        "finish_reason": "stop",
+                    }
+                ],
+            },
+        ]
+        self.send_response(200)
+        self.send_header("Content-Type", "text/event-stream")
+        self.send_header("Cache-Control", "no-cache")
+        self.send_header("Connection", "close")
+        self.end_headers()
+        for chunk in chunks:
+            self.wfile.write(
+                f"data: {json.dumps(chunk, separators=(',', ':'))}\n\n".encode()
+            )
+            self.wfile.flush()
+        self.wfile.write(b"data: [DONE]\n\n")
+        self.wfile.flush()
+        self.close_connection = True
+
+    def log_message(self, _format: str, *_args: object) -> None:
+        return
+
+
+def deny_permission(*_args: object, **_kwargs: object) -> PermissionDecisionReject:
+    return PermissionDecisionReject(feedback="No tools are allowed.")
+
+
+class SdkProviderIntegrationTest(unittest.IsolatedAsyncioTestCase):
+    async def test_two_streamed_turns_use_openai_compatible_provider(self) -> None:
+        self.assertTrue(os.path.isfile(CLI_PATH))
+        ProviderHandler.requests = []
+        server = ThreadingHTTPServer(("127.0.0.1", 0), ProviderHandler)
+        server_thread = threading.Thread(target=server.serve_forever, daemon=True)
+        server_thread.start()
+        events = []
+        idle = asyncio.Event()
+
+        with tempfile.TemporaryDirectory() as base_directory:
+            client = CopilotClient(
+                connection=RuntimeConnection.for_stdio(path=CLI_PATH),
+                base_directory=base_directory,
+                use_logged_in_user=False,
+                log_level="error",
+                mode="empty",
+            )
+            try:
+                await client.start()
+                session = await client.create_session(
+                    session_id="provider-integration",
+                    on_permission_request=deny_permission,
+                    model="gpt-4",
+                    provider={
+                        "type": "openai",
+                        "base_url": (
+                            f"http://127.0.0.1:{server.server_port}/v1"
+                        ),
+                        "wire_api": "completions",
+                        "api_key": "local-test-key",
+                    },
+                    streaming=True,
+                    tools=[],
+                    available_tools=[],
+                    mcp_servers={},
+                    enable_config_discovery=False,
+                    skip_custom_instructions=True,
+                    enable_skills=False,
+                    enable_session_store=True,
+                    on_event=lambda event: (
+                        events.append(event),
+                        idle.set()
+                        if event.type.value == "session.idle"
+                        else None,
+                    ),
+                )
+                for prompt in ("first", "second"):
+                    idle.clear()
+                    await session.send(prompt)
+                    await asyncio.wait_for(idle.wait(), timeout=15)
+                await session.disconnect()
+            finally:
+                await client.stop()
+                server.shutdown()
+                server.server_close()
+                server_thread.join()
+
+        self.assertEqual(len(ProviderHandler.requests), 2)
+        self.assertTrue(
+            all(path.endswith("/chat/completions")
+                for path, _ in ProviderHandler.requests)
+        )
+        deltas = [
+            event.data.delta_content
+            for event in events
+            if event.type.value == "assistant.message_delta"
+        ]
+        self.assertEqual(deltas, ["hello ", "traveler", "hello ", "traveler"])
+        messages = [
+            event.data.content
+            for event in events
+            if event.type.value == "assistant.message"
+        ]
+        self.assertEqual(messages, ["hello traveler", "hello traveler"])
+
+
+if __name__ == "__main__":
+    unittest.main()