diff --git a/tests/integration/caching/test_semantic_cache_tool_turns.py b/tests/integration/caching/test_semantic_cache_tool_turns.py index 5d6e6a956a1..ac80b806c40 100644 --- a/tests/integration/caching/test_semantic_cache_tool_turns.py +++ b/tests/integration/caching/test_semantic_cache_tool_turns.py @@ -1,14 +1,19 @@ +import asyncio import hashlib import json import math import threading import uuid -from collections.abc import Iterator, Mapping, Sequence +from collections.abc import Callable, Iterator, Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor from dataclasses import dataclass, field from pathlib import Path from typing import Final from urllib.parse import urlsplit +import anthropic +import httpx +import openai import pytest import yaml from integration._support.client import Gateway, eventually, gateway_from_environment @@ -46,6 +51,8 @@ class _Peer: points: list[Mapping[str, JsonValue]] = field(default_factory=list) # mutable-ok: the Qdrant collection embedded: list[str] = field(default_factory=list) # mutable-ok: every prompt the proxy embedded answered: list[str] = field(default_factory=list) # mutable-ok: every completion the provider served + qdrant_down: threading.Event = field(default_factory=threading.Event) + refused_writes: list[str] = field(default_factory=list) # mutable-ok: upserts refused while Qdrant is down def stored(self) -> int: with self.lock: @@ -55,6 +62,11 @@ class _Peer: path: Final = urlsplit(request.target).path body: Final = _JSON_OBJECT.validate_json(request.body) if request.body else {} collection: Final = f"/qdrant/collections/{_COLLECTION}" + if path.startswith("/qdrant/") and self.qdrant_down.is_set(): + if path == f"{collection}/points": + with self.lock: + self.refused_writes.append(path) + return Reply(status=503, body=b'{"status":{"error":"qdrant unavailable"}}') if request.method == "GET" and path == f"{collection}/exists": return self._json({"result": {"exists": False}, "status": "ok"}) if request.method in {"GET", "PUT"} and path in {collection, f"{collection}/index"}: @@ -279,3 +291,254 @@ def test_openai_agent_loop_tool_calls_on_chat_completions_are_not_served_a_cache assert len(set(answers[:2])) == 2, f"a different tool call replayed an earlier cached answer: {answers}" assert answers[:2] == semantic_proxy.peer.answered[-2:], semantic_proxy.peer.answered assert answers[2] == answers[0], f"a repeated tool call missed the cache: {answers}" + + +def _tool_turn(task: str, calls: Sequence[tuple[str, str]], results: Sequence[tuple[str, str]]) -> list[JsonValue]: + return [ + {"role": "user", "content": task}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": call_id, "type": "function", "function": {"name": "read_file", "arguments": arguments}} + for call_id, arguments in calls + ], + }, + *({"role": "tool", "tool_call_id": call_id, "content": output} for call_id, output in results), + ] + + +def _claude_turn(task: str, calls: Sequence[tuple[str, str]], results: Sequence[tuple[str, str]]) -> list[JsonValue]: + return [ + {"role": "user", "content": task}, + { + "role": "assistant", + "content": [ + {"type": "tool_use", "id": call_id, "name": "Read", "input": {"file_path": file_path}} + for call_id, file_path in calls + ], + }, + { + "role": "user", + "content": [ + {"type": "tool_result", "tool_use_id": call_id, "content": output} for call_id, output in results + ], + }, + ] + + +_TWO_READS: Final = (("c1", '{"path":"a.py"}'), ("c2", '{"path":"b.py"}')) +_TWO_CLAUDE_READS: Final = (("toolu_1", "a.py"), ("toolu_2", "b.py")) + + +@pytest.mark.parametrize( + ("route", "model", "build", "calls", "first", "second"), + [ + pytest.param( + "/v1/chat/completions", + _CHAT_MODEL, + _tool_turn, + _TWO_READS, + (("c1", "x = 1"), ("c2", "y = 2")), + (("c2", "x = 1"), ("c1", "y = 2")), + id="chat-parallel-results-answer-the-other-call", + ), + pytest.param( + "/v1/messages", + _CLAUDE_MODEL, + _claude_turn, + _TWO_CLAUDE_READS, + (("toolu_1", "x = 1"), ("toolu_2", "y = 2")), + (("toolu_2", "x = 1"), ("toolu_1", "y = 2")), + id="messages-parallel-results-answer-the-other-call", + ), + pytest.param( + "/v1/chat/completions", + _CHAT_MODEL, + _tool_turn, + _TWO_READS[:1], + (("c1", "x = 1"),), + (("c9", "x = 1"),), + id="chat-result-for-an-unknown-call", + ), + pytest.param( + "/v1/chat/completions", + _CHAT_MODEL, + _tool_turn, + _TWO_READS, + (("c1", '"},{"result_of_call":2,"output":"y = 2'),), + (("c1", ""), ("c2", "y = 2")), + id="chat-output-shaped-like-a-second-result", + ), + pytest.param( + "/v1/chat/completions", + _CHAT_MODEL, + _tool_turn, + _TWO_READS[:1], + (("c1", ""),), + (("c1", "x" * 5000),), + id="chat-empty-and-5kb-output", + ), + ], +) +def test_tool_turns_that_differ_only_in_their_results_do_not_share_a_cached_answer( + semantic_proxy: _SemanticProxy, + route: str, + model: str, + build: Callable[[str, Sequence[tuple[str, str]], Sequence[tuple[str, str]]], list[JsonValue]], + calls: Sequence[tuple[str, str]], + first: Sequence[tuple[str, str]], + second: Sequence[tuple[str, str]], +) -> None: + task: Final = f"read both files {uuid.uuid4().hex}" + + answers: Final = _send_turns( + semantic_proxy, route, model, (build(task, calls, first), build(task, calls, second), build(task, calls, first)) + ) + + assert answers[0] != answers[1], f"a different tool result replayed an earlier cached answer: {answers}" + assert answers[:2] == semantic_proxy.peer.answered[-2:], semantic_proxy.peer.answered + assert answers[2] == answers[0], f"a repeated tool turn missed the cache: {answers}" + + +def _openai_sdk(proxy: _SemanticProxy, turn: list[JsonValue]) -> str: + sdk: Final = openai.OpenAI( + base_url=str(proxy.gateway.client.base_url) + "/v1", + api_key=proxy.gateway.key, + http_client=httpx.Client(trust_env=False, timeout=15), + ) + completion: Final = sdk.chat.completions.create(model=_CHAT_MODEL, messages=turn) # pyright: ignore[reportArgumentType, reportCallIssue] # JSON turns are validated by the proxy, not the SDK types + return str(completion.choices[0].message.content) + + +def _async_openai_sdk(proxy: _SemanticProxy, turn: list[JsonValue]) -> str: + async def create() -> str: + async with httpx.AsyncClient(trust_env=False, timeout=15) as http_client: + sdk: Final = openai.AsyncOpenAI( + base_url=str(proxy.gateway.client.base_url) + "/v1", api_key=proxy.gateway.key, http_client=http_client + ) + completion: Final = await sdk.chat.completions.create(model=_CHAT_MODEL, messages=turn) # pyright: ignore[reportArgumentType, reportCallIssue] # JSON turns are validated by the proxy, not the SDK types + return str(completion.choices[0].message.content) + + return asyncio.run(create()) + + +def _anthropic_text(message: anthropic.types.Message) -> str: + block: Final = message.content[0] + assert isinstance(block, anthropic.types.TextBlock), message + return block.text + + +def _anthropic_sdk(proxy: _SemanticProxy, turn: list[JsonValue]) -> str: + sdk: Final = anthropic.Anthropic( + base_url=str(proxy.gateway.client.base_url), + api_key=proxy.gateway.key, + http_client=httpx.Client(trust_env=False, timeout=15), + ) + return _anthropic_text(sdk.messages.create(model=_CLAUDE_MODEL, max_tokens=16, messages=turn)) # pyright: ignore[reportArgumentType] # JSON turns are validated by the proxy, not the SDK types + + +def _async_anthropic_sdk(proxy: _SemanticProxy, turn: list[JsonValue]) -> str: + async def create() -> str: + async with httpx.AsyncClient(trust_env=False, timeout=15) as http_client: + sdk: Final = anthropic.AsyncAnthropic( + base_url=str(proxy.gateway.client.base_url), api_key=proxy.gateway.key, http_client=http_client + ) + message: Final = await sdk.messages.create(model=_CLAUDE_MODEL, max_tokens=16, messages=turn) # pyright: ignore[reportArgumentType] # JSON turns are validated by the proxy, not the SDK types + return _anthropic_text(message) + + return asyncio.run(create()) + + +@pytest.mark.parametrize( + ("send", "build"), + [ + pytest.param(_openai_sdk, _tool_turn, id="openai-sdk"), + pytest.param(_async_openai_sdk, _tool_turn, id="async-openai-sdk"), + pytest.param(_anthropic_sdk, _claude_turn, id="anthropic-sdk"), + pytest.param(_async_anthropic_sdk, _claude_turn, id="async-anthropic-sdk"), + ], +) +def test_sdk_clients_get_a_fresh_answer_for_each_tool_turn_and_a_hit_for_a_repeat( + semantic_proxy: _SemanticProxy, + send: Callable[[_SemanticProxy, list[JsonValue]], str], + build: Callable[[str, Sequence[tuple[str, str]], Sequence[tuple[str, str]]], list[JsonValue]], +) -> None: + task: Final = f"read the file {uuid.uuid4().hex}" + calls: Final = _TWO_READS[:1] if build is _tool_turn else _TWO_CLAUDE_READS[:1] + call_id: Final = calls[0][0] + alpha: Final = build(task, calls, ((call_id, "x = 1"),)) + beta: Final = build(task, calls, ((call_id, "PermissionError"),)) + + def answer(turn: list[JsonValue]) -> str: + stored_before: Final = semantic_proxy.peer.stored() + answered_before: Final = len(semantic_proxy.peer.answered) + text: Final = send(semantic_proxy, turn) + if len(semantic_proxy.peer.answered) > answered_before: + eventually(semantic_proxy.peer.stored, lambda count: count > stored_before) + return text + + answers: Final = [answer(turn) for turn in (alpha, beta, alpha)] + + assert answers[0] != answers[1], f"a different tool result replayed an earlier cached answer: {answers}" + assert answers[:2] == semantic_proxy.peer.answered[-2:], semantic_proxy.peer.answered + assert answers[2] == answers[0], f"a repeated tool turn missed the cache: {answers}" + + +def test_concurrent_agent_loops_each_get_their_own_answer_and_their_own_hit(semantic_proxy: _SemanticProxy) -> None: + task: Final = f"read the file {uuid.uuid4().hex}" + turns: Final = tuple(_tool_turn(task, _TWO_READS[:1], (("c1", f"line {index}"),)) for index in range(8)) + + def send(turn: list[JsonValue]) -> str: + response: Final = semantic_proxy.gateway.request( + "POST", "/v1/chat/completions", {"model": _CHAT_MODEL, "messages": turn} + ) + assert response.status_code == 200, response.text + choice: Final = _JSON_OBJECT.validate_python(_list(_JSON_OBJECT.validate_json(response.content)["choices"])[0]) + return str(_JSON_OBJECT.validate_python(choice["message"])["content"]) + + stored_before: Final = semantic_proxy.peer.stored() + with ThreadPoolExecutor(max_workers=len(turns)) as pool: + first: Final = tuple(pool.map(send, turns)) + eventually(semantic_proxy.peer.stored, lambda count: count >= stored_before + len(turns)) + answered_before: Final = len(semantic_proxy.peer.answered) + with ThreadPoolExecutor(max_workers=len(turns)) as pool: + second: Final = tuple(pool.map(send, turns)) + + assert len(set(first)) == len(turns), f"concurrent tool turns shared a cached answer: {first}" + assert second == first, f"a repeated tool turn got another turn's answer: {first} {second}" + assert len(semantic_proxy.peer.answered) == answered_before, "a repeated tool turn missed the cache" + + +def test_tool_turns_are_answered_while_qdrant_is_down_and_cached_once_it_recovers( + semantic_proxy: _SemanticProxy, +) -> None: + task: Final = f"read the file {uuid.uuid4().hex}" + turn: Final = _tool_turn(task, _TWO_READS[:1], (("c1", "x = 1"),)) + body: Final[dict[str, JsonValue]] = {"model": _CHAT_MODEL, "messages": turn} + + def send() -> httpx.Response: + return semantic_proxy.gateway.request("POST", "/v1/chat/completions", body) + + refused_before: Final = len(semantic_proxy.peer.refused_writes) + answered_before: Final = len(semantic_proxy.peer.answered) + semantic_proxy.peer.qdrant_down.set() + try: + outage: Final = (send(), send()) + eventually(lambda: len(semantic_proxy.peer.refused_writes), lambda count: count >= refused_before + 2) + finally: + semantic_proxy.peer.qdrant_down.clear() + answered_during_outage: Final = len(semantic_proxy.peer.answered) - answered_before + stored_before: Final = semantic_proxy.peer.stored() + recovered: Final = send() + eventually(semantic_proxy.peer.stored, lambda count: count > stored_before) + answered_after_recovery: Final = len(semantic_proxy.peer.answered) + repeat: Final = send() + + assert [response.status_code for response in outage] == [200, 200], [response.text for response in outage] + assert answered_during_outage == 2, "an outage request was not sent to the provider" + assert recovered.status_code == 200, recovered.text + assert answered_after_recovery - answered_before == 3, "the first request after recovery hit a stale entry" + assert repeat.status_code == 200, repeat.text + assert repeat.json()["choices"] == recovered.json()["choices"] + assert len(semantic_proxy.peer.answered) == answered_after_recovery, "the repeat after recovery missed the cache"