diff --git a/tests/fixtures_e2e_cancellation.py b/tests/fixtures_e2e_cancellation.py index 203463d24..7aedb4981 100644 --- a/tests/fixtures_e2e_cancellation.py +++ b/tests/fixtures_e2e_cancellation.py @@ -120,15 +120,33 @@ class StubModelServer: outer.calls.append(body) if outer.latency_sec: time.sleep(outer.latency_sec) - return self._send(outer._completion(body, len(outer.calls))) + return self._send(outer._completion(body, len(outer.calls)), stream=bool(body.get("stream"))) - def _send(self, payload): - data = json.dumps(payload).encode("utf-8") + def _send(self, payload, *, stream=False): + content_type = "application/json" + if stream: + content_type = "text/event-stream" + choices = [] + for choice in payload["choices"]: + delta = dict(choice["message"]) + if delta.get("tool_calls"): + delta["tool_calls"] = [ + dict(call, index=index) for index, call in enumerate(delta["tool_calls"]) + ] + choices.append({"index": choice["index"], "delta": delta, + "finish_reason": choice["finish_reason"]}) + common = {"id": payload["id"], "model": payload["model"], "object": "chat.completion.chunk"} + frames = [{**common, "choices": choices}, {**common, "choices": [], "usage": payload["usage"]}] + data = ("".join("data: " + json.dumps(frame) + "\n\n" for frame in frames) + + "data: [DONE]\n\n").encode("utf-8") + else: + data = json.dumps(payload).encode("utf-8") self.send_response(200) - self.send_header("Content-Type", "application/json") + self.send_header("Content-Type", content_type) self.send_header("Content-Length", str(len(data))) self.end_headers() self.wfile.write(data) + self.wfile.flush() def log_message(self, *_args): return diff --git a/tests/test_loopback_model_stream.py b/tests/test_loopback_model_stream.py index 62bc29454..b058a6081 100644 --- a/tests/test_loopback_model_stream.py +++ b/tests/test_loopback_model_stream.py @@ -8,6 +8,7 @@ import pytest from openai import OpenAI from ouroboros.llm_stream import consume_stream +from tests.fixtures_e2e_cancellation import StubModelServer from tests.system_e2e.harness import ScriptedStubModel @@ -41,3 +42,33 @@ def test_loopback_model_round_trips_through_the_real_sdk(stream, tool_call): assert result["usage"]["prompt_tokens"] == 10 assert result["usage"]["completion_tokens"] == 5 assert result["usage"]["total_tokens"] == 15 + + +@pytest.mark.serial +@pytest.mark.parametrize("stream", [False, True]) +@pytest.mark.parametrize("tool_call", [False, True]) +def test_cancellation_model_round_trips_through_the_real_sdk(stream, tool_call): + tools = [{"type": "function", "function": {"name": "list_files", "parameters": { + "type": "object", "properties": {"path": {"type": "string"}}, + "required": ["path"], "additionalProperties": False, + }}}] + with StubModelServer(mode="keepalive" if tool_call else "finish") as model: + with OpenAI(base_url=model.base_url, api_key="keyless-fixture", max_retries=0) as client: + response = client.chat.completions.create( + model="mock-model", messages=[{"role": "user", "content": "Execute the fixture."}], + tools=tools, stream=stream, + **({"stream_options": {"include_usage": True}} if stream else {}), + ) + result = consume_stream(response).model_dump() if stream else response.model_dump() + assert len(model.calls) == 1 and model.calls[0]["stream"] is stream + message = result["choices"][0]["message"] + if tool_call: + call = message["tool_calls"][0] + assert call["id"] == "call_1" and call["type"] == "function" + assert call["function"]["name"] == "list_files" + assert json.loads(call["function"]["arguments"]) == {"path": "."} + else: + assert message["content"] == "Done." + assert result["usage"]["prompt_tokens"] == 10 + assert result["usage"]["completion_tokens"] == 5 + assert result["usage"]["total_tokens"] == 15