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