view mrjunejune/inference/sdk_provider_integration_test.py @ 269:de291f396881

install initial production config Copy the ignored repository config into /etc/mrjunejune on first deployment while preserving existing production configuration on later deploys. Co-authored-by: Copilot <[email protected]>
author MrJuneJune <me@mrjunejune.com>
date Fri, 07 Aug 2026 13:08:10 -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()