view mrjunejune/inference/sdk_provider_integration_test.py @ 277:1d99147f520c

Merge Qwen services and infinite canvas heads
author MrJuneJune <me@mrjunejune.com>
date Mon, 17 Aug 2026 17:01:40 -0700
parents 1f9877b637e9
children
line wrap: on
line source

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