mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(responses): honor request cache controls on chat completions bridged to the Responses API (#44676)
* fix(responses): honor request cache controls on chat completions bridged to the Responses API Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(responses): assert spend and cache-hit status on bridged no-cache rows Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): unskip LIT-9196 openai_responses basic translation cases Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(responses): restore azure attribution check on bridged no-cache rows Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(responses): exercise bridged cache controls through a real local cache instead of patched responses Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: kerry <kerry@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
2190003bb8
commit
3242bdfed2
5 changed files with 257 additions and 3 deletions
|
|
@ -194,6 +194,8 @@ class ResponsesToCompletionBridgeHandler:
|
|||
# would raise a duplicate-keyword TypeError.
|
||||
request_data["custom_llm_provider"] = custom_llm_provider
|
||||
request_data["model"] = _restore_routing_prefix(model, custom_llm_provider)
|
||||
if kwargs.get("cache") is not None:
|
||||
request_data["cache"] = kwargs["cache"]
|
||||
result: Final = responses(
|
||||
**request_data,
|
||||
)
|
||||
|
|
@ -289,6 +291,8 @@ class ResponsesToCompletionBridgeHandler:
|
|||
# would raise a duplicate-keyword TypeError.
|
||||
request_data["custom_llm_provider"] = custom_llm_provider
|
||||
request_data["model"] = _restore_routing_prefix(model, custom_llm_provider)
|
||||
if kwargs.get("cache") is not None:
|
||||
request_data["cache"] = kwargs["cache"]
|
||||
result: Final = await aresponses(
|
||||
**request_data,
|
||||
aresponses=True,
|
||||
|
|
|
|||
|
|
@ -5831,6 +5831,7 @@ def completion(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
encoding=_get_encoding(),
|
||||
stream=stream,
|
||||
cache=kwargs.get("cache"),
|
||||
)
|
||||
elif (custom_llm_provider == "openai" and OpenAIGPT5Config.is_model_gpt_5_model(model)) or (
|
||||
custom_llm_provider == "azure"
|
||||
|
|
|
|||
|
|
@ -1,10 +1,12 @@
|
|||
import base64
|
||||
import json
|
||||
import uuid
|
||||
from itertools import chain
|
||||
from itertools import chain, count
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from integration._support.client import Gateway
|
||||
from integration._support.client import Gateway, eventually, string_value
|
||||
from integration._support.database import read_rows
|
||||
from integration._support.wire import Reply, Request, wire_server
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
|
||||
|
|
@ -252,3 +254,167 @@ def test_azure_gpt_6_bridged_stream_returns_text_and_tool_call_on_one_choice(gat
|
|||
assert [(request.method, request.target) for request in wire.drain()] == [
|
||||
("POST", "/openai/responses?api-version=2025-04-01-preview")
|
||||
]
|
||||
|
||||
|
||||
_RESPONSES_TARGET: Final = "/openai/responses?api-version=2025-04-01-preview"
|
||||
|
||||
|
||||
def _responses_json(identity: str) -> bytes:
|
||||
return json.dumps(
|
||||
{
|
||||
"id": identity,
|
||||
"object": "response",
|
||||
"created_at": 1,
|
||||
"status": "completed",
|
||||
"model": "gpt-6-sol",
|
||||
"output": [
|
||||
{
|
||||
"id": f"msg_{identity}",
|
||||
"type": "message",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "hi", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15},
|
||||
}
|
||||
).encode()
|
||||
|
||||
|
||||
def _gpt_6_function_request(model: str, identity: str, **extra: JsonValue) -> dict[str, JsonValue]:
|
||||
return {
|
||||
"model": model,
|
||||
"messages": [{"role": "user", "content": f"What is the weather in Paris? {identity}"}],
|
||||
"tools": [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"description": "Get the weather for a city.",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"city": {"type": "string"}},
|
||||
"required": ["city"],
|
||||
},
|
||||
},
|
||||
}
|
||||
],
|
||||
**extra,
|
||||
}
|
||||
|
||||
|
||||
def test_azure_gpt_6_bridged_no_cache_function_requests_each_reach_provider_and_log_spend(
|
||||
gateway: Gateway,
|
||||
) -> None:
|
||||
identity: Final = f"azure-gpt-6-sol-nocache-{uuid.uuid4().hex}"
|
||||
response_ids: Final = ("resp_first", "resp_second")
|
||||
calls: Final = count()
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "POST"
|
||||
assert request.target == _RESPONSES_TARGET
|
||||
body: Final = _JSON_OBJECT.validate_json(request.body)
|
||||
assert body["model"] == "gpt-6-sol"
|
||||
assert body["input"] == [
|
||||
{
|
||||
"type": "message",
|
||||
"role": "user",
|
||||
"content": [{"type": "input_text", "text": f"What is the weather in Paris? {identity}"}],
|
||||
}
|
||||
]
|
||||
return Reply(body=_responses_json(response_ids[next(calls)]))
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
model="azure/gpt-6-sol",
|
||||
api_base=wire.url,
|
||||
api_key=_API_KEY,
|
||||
api_version="2025-04-01-preview",
|
||||
input_cost_per_token=0.001,
|
||||
output_cost_per_token=0.002,
|
||||
)
|
||||
request: Final = _gpt_6_function_request(model, identity, cache={"no-cache": True})
|
||||
first: Final = gateway.request("POST", "/v1/chat/completions", request)
|
||||
second: Final = gateway.request("POST", "/v1/chat/completions", request)
|
||||
assert first.status_code == 200, first.text
|
||||
assert second.status_code == 200, second.text
|
||||
assert string_value(_JSON_OBJECT.validate_json(first.content)["id"]) == "resp_first", first.text
|
||||
assert string_value(_JSON_OBJECT.validate_json(second.content)["id"]) == "resp_second", second.text
|
||||
assert [(request.method, request.target) for request in wire.drain()] == [
|
||||
("POST", _RESPONSES_TARGET),
|
||||
("POST", _RESPONSES_TARGET),
|
||||
]
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT request_id, status, cache_hit, spend FROM "LiteLLM_SpendLogs" WHERE model_group=%s',
|
||||
(model,),
|
||||
),
|
||||
lambda found: len(found) == 2,
|
||||
seconds=70,
|
||||
)
|
||||
by_response_id: Final = {
|
||||
(decoded := base64.b64decode(string_value(row["request_id"]).removeprefix("resp_")).decode())
|
||||
.rsplit("response_id:", 1)[1]: (decoded, row)
|
||||
for row in rows
|
||||
}
|
||||
for response_id in response_ids:
|
||||
decoded, row = by_response_id[response_id]
|
||||
assert decoded.startswith("litellm:custom_llm_provider:azure;model_id:"), rows
|
||||
assert (
|
||||
string_value(row["status"]),
|
||||
string_value(row["cache_hit"]),
|
||||
float(row["spend"]),
|
||||
) == ("success", "None", pytest.approx(10 * 0.001 + 5 * 0.002)), rows
|
||||
|
||||
|
||||
def test_azure_gpt_6_bridged_function_requests_without_cache_field_still_hit_cache(
|
||||
gateway: Gateway,
|
||||
) -> None:
|
||||
identity: Final = f"azure-gpt-6-sol-cached-{uuid.uuid4().hex}"
|
||||
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.method == "POST"
|
||||
assert request.target == _RESPONSES_TARGET
|
||||
body: Final = _JSON_OBJECT.validate_json(request.body)
|
||||
assert body["model"] == "gpt-6-sol"
|
||||
return Reply(body=_responses_json("resp_cached"))
|
||||
|
||||
with wire_server(respond) as wire, gateway.scenario() as scenario:
|
||||
model: Final = scenario.model(
|
||||
model="azure/gpt-6-sol",
|
||||
api_base=wire.url,
|
||||
api_key=_API_KEY,
|
||||
api_version="2025-04-01-preview",
|
||||
input_cost_per_token=0.001,
|
||||
output_cost_per_token=0.002,
|
||||
)
|
||||
request: Final = _gpt_6_function_request(model, identity)
|
||||
first: Final = gateway.request("POST", "/v1/chat/completions", request)
|
||||
second: Final = gateway.request("POST", "/v1/chat/completions", request)
|
||||
assert first.status_code == 200, first.text
|
||||
assert second.status_code == 200, second.text
|
||||
assert string_value(_JSON_OBJECT.validate_json(first.content)["id"]) == "resp_cached", first.text
|
||||
assert string_value(_JSON_OBJECT.validate_json(second.content)["id"]) == "resp_cached", second.text
|
||||
assert [(request.method, request.target) for request in wire.drain()] == [("POST", _RESPONSES_TARGET)]
|
||||
rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT request_id, status, cache_hit, spend FROM "LiteLLM_SpendLogs" WHERE model_group=%s'
|
||||
" ORDER BY request_id",
|
||||
(model,),
|
||||
),
|
||||
lambda found: len(found) == 2,
|
||||
seconds=70,
|
||||
)
|
||||
priced, cached = rows
|
||||
priced_request: Final = base64.b64decode(
|
||||
string_value(priced["request_id"]).removeprefix("resp_")
|
||||
).decode()
|
||||
assert priced_request.startswith("litellm:custom_llm_provider:azure;model_id:"), rows
|
||||
assert priced_request.endswith(";response_id:resp_cached"), rows
|
||||
assert (
|
||||
priced["status"],
|
||||
priced["cache_hit"],
|
||||
float(priced["spend"]),
|
||||
) == ("success", "None", pytest.approx(10 * 0.001 + 5 * 0.002)), rows
|
||||
assert string_value(cached["request_id"]).startswith("resp_cached_cache_hit"), rows
|
||||
assert (cached["status"], cached["cache_hit"], float(cached["spend"])) == ("success", "True", 0), rows
|
||||
|
|
|
|||
|
|
@ -20,5 +20,4 @@ from integration.translation.runner import run
|
|||
def test_chat_completions_basic_openai_responses(
|
||||
case: TranslationTestCase, gateway: Gateway, provider: SharedProvider
|
||||
) -> None:
|
||||
pytest.skip("BUG: LIT-9196 no-cache ignored on chat completions bridged to responses")
|
||||
run(case, gateway, provider)
|
||||
|
|
|
|||
|
|
@ -1,8 +1,14 @@
|
|||
import itertools
|
||||
from datetime import datetime
|
||||
from typing import Final
|
||||
from unittest.mock import patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
import respx
|
||||
|
||||
import litellm
|
||||
from litellm.caching.caching import Cache
|
||||
from litellm.completion_extras.litellm_responses_transformation.handler import (
|
||||
ResponsesToCompletionBridgeHandler,
|
||||
)
|
||||
|
|
@ -109,3 +115,81 @@ async def test_acompletion_forwards_aws_region_name_to_aresponses():
|
|||
assert result is cached
|
||||
assert _fake_aresponses.kwargs["aws_region_name"] == REGION
|
||||
assert _fake_aresponses.kwargs["custom_llm_provider"] == "bedrock_mantle"
|
||||
|
||||
|
||||
def _responses_api_body(n: int) -> dict:
|
||||
return {
|
||||
"id": f"resp_{n}",
|
||||
"object": "response",
|
||||
"created_at": 1,
|
||||
"status": "completed",
|
||||
"model": "gpt-6-sol",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": f"msg_{n}",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": f"hi {n}", "annotations": []}],
|
||||
}
|
||||
],
|
||||
"usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15},
|
||||
}
|
||||
|
||||
|
||||
_GPT_6_TOOLS: Final = [
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"parameters": {"type": "object", "properties": {"city": {"type": "string"}}},
|
||||
},
|
||||
}
|
||||
]
|
||||
|
||||
_GPT_6_REQUEST: Final = {
|
||||
"model": "azure/gpt-6-sol",
|
||||
"api_base": "https://example.invalid",
|
||||
"api_key": "x",
|
||||
"api_version": "2025-04-01-preview",
|
||||
"messages": [{"role": "user", "content": "weather?"}],
|
||||
"tools": _GPT_6_TOOLS,
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def _bridged_cache_edge(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
|
||||
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
|
||||
ids = itertools.count(1)
|
||||
with respx.mock(assert_all_called=True) as router:
|
||||
route = router.post(url__regex=r"https://example\.invalid/openai/responses.*").mock(
|
||||
side_effect=lambda request: httpx.Response(200, json=_responses_api_body(next(ids)))
|
||||
)
|
||||
yield route
|
||||
|
||||
|
||||
def test_completion_no_cache_reaches_provider_each_time(_bridged_cache_edge):
|
||||
first = litellm.completion(**_GPT_6_REQUEST, cache={"no-cache": True})
|
||||
second = litellm.completion(**_GPT_6_REQUEST, cache={"no-cache": True})
|
||||
|
||||
assert len(_bridged_cache_edge.calls) == 2
|
||||
assert [first.id, second.id] == ["resp_1", "resp_2"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acompletion_no_cache_reaches_provider_each_time(_bridged_cache_edge):
|
||||
first = await litellm.acompletion(**_GPT_6_REQUEST, cache={"no-cache": True})
|
||||
second = await litellm.acompletion(**_GPT_6_REQUEST, cache={"no-cache": True})
|
||||
|
||||
assert len(_bridged_cache_edge.calls) == 2
|
||||
assert [first.id, second.id] == ["resp_1", "resp_2"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_acompletion_without_cache_field_is_served_from_cache(_bridged_cache_edge):
|
||||
first = await litellm.acompletion(**_GPT_6_REQUEST)
|
||||
second = await litellm.acompletion(**_GPT_6_REQUEST)
|
||||
|
||||
assert len(_bridged_cache_edge.calls) == 1
|
||||
assert [first.id, second.id] == ["resp_1", "resp_1"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue