mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(proxy): mark background responses stale_expired when provider returns 404 (#41724)
* fix(proxy): mark background responses stale_expired when provider returns 404 Responses deleted upstream (store=false / ZDR rows dropped after provider retention) now move to stale_expired on the poll instead of being retried every cycle until the 7 day stale cleanup. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(proxy): drop inline comment in check_responses_cost Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): drive _get_response seam in 404 stale_expired tests Assigning AsyncMock on the instance keeps the new tests inside the TQ002/TQ008 test-quality ceilings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): only expire 404 responses polled through a resolved deployment A NotFoundError from the bare SDK fallback can mean a missing or misconfigured deployment, which a config fix inside the staleness window can still recover. Rows whose poll went through their router deployment are the ones the provider actually dropped, so only those move to stale_expired. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): resolve the deployment before fetching so via_router binds once Drops the reinitialised loop flag flagged by review and modernizes the moved annotation. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): drive the real response fetch in the poll 404 regression tests * fix(enterprise): expire background response rows on any provider 404 status The GET path maps the provider's 404 body to litellm.BadRequestError that still carries status_code 404, so an except on NotFoundError never fired and the poll job kept retrying the row every cycle. Key the terminal decision on the status code instead of the class * fix(enterprise): expire a background response only on a 404 naming it and skip a bad deployment entry per job * chore(enterprise): record the poll-cycle bound on the managed-object IN lists * test(integration): cover background response poll retirement on the real proxy --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
parent
d5173e3d70
commit
b7ecb05340
4 changed files with 691 additions and 10 deletions
|
|
@ -6,7 +6,7 @@ same route are non-inference and free.
|
|||
"""
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import TYPE_CHECKING, Dict, Final, Optional, Protocol, cast
|
||||
from typing import TYPE_CHECKING, Dict, Final, Protocol, cast
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
|
@ -47,6 +47,10 @@ def _managed_object_table(prisma_client: "PrismaClient") -> "TableActions[_Manag
|
|||
return ManagedObjectRepository(prisma_client).table
|
||||
|
||||
|
||||
def _is_response_gone_at_provider(error: Exception, provider_response_id: str) -> bool:
|
||||
return getattr(error, "status_code", None) == 404 and provider_response_id in str(error)
|
||||
|
||||
|
||||
class CheckResponsesCost:
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -61,10 +65,15 @@ class CheckResponsesCost:
|
|||
self.prisma_client: PrismaClient = prisma_client
|
||||
self.llm_router: Router = llm_router
|
||||
|
||||
def _resolve_deployment(self, response_id: str) -> bool:
|
||||
model_id: str | None = ResponsesAPIRequestUtils.get_model_id_from_response_id(response_id)
|
||||
return model_id is not None and self.llm_router.get_deployment(model_id=model_id) is not None
|
||||
|
||||
async def _get_response(
|
||||
self,
|
||||
response_id: str,
|
||||
litellm_metadata: Dict[str, str],
|
||||
via_router: bool,
|
||||
) -> ResponsesAPIResponse:
|
||||
"""Fetch the upstream response, using deployment credentials when available.
|
||||
|
||||
|
|
@ -75,8 +84,7 @@ class CheckResponsesCost:
|
|||
sees provider env vars, so it fails for every deployment whose credentials
|
||||
live in the config; the row then never leaves ``queued``.
|
||||
"""
|
||||
model_id: Optional[str] = ResponsesAPIRequestUtils.get_model_id_from_response_id(response_id)
|
||||
if model_id is None or self.llm_router.get_deployment(model_id=model_id) is None:
|
||||
if not via_router:
|
||||
return await litellm.aget_responses(response_id=response_id, litellm_metadata=litellm_metadata)
|
||||
router_response = await self.llm_router.aget_responses(
|
||||
response_id=response_id, litellm_metadata=litellm_metadata
|
||||
|
|
@ -140,6 +148,8 @@ class CheckResponsesCost:
|
|||
- Cost is tracked by the get-responses call, billed because the poll is stamped
|
||||
with BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN
|
||||
- Mark responses in a terminal state as complete in the database
|
||||
- Mark responses the provider no longer has (404 through a resolved
|
||||
deployment) as stale_expired
|
||||
"""
|
||||
try:
|
||||
await self._cleanup_stale_managed_objects()
|
||||
|
|
@ -159,6 +169,7 @@ class CheckResponsesCost:
|
|||
|
||||
verbose_proxy_logger.debug(f"Found {len(jobs)} response jobs to check")
|
||||
completed_jobs: Final[list[_ManagedObjectRow]] = []
|
||||
expired_jobs: Final[list[_ManagedObjectRow]] = []
|
||||
|
||||
for job in jobs:
|
||||
unified_object_id = job.unified_object_id
|
||||
|
|
@ -171,31 +182,48 @@ class CheckResponsesCost:
|
|||
# Get the stored response object to extract model information
|
||||
stored_response = job.file_object
|
||||
model_name = stored_response.get("model", None)
|
||||
|
||||
|
||||
# Decrypt the response ID
|
||||
responses_id_security, _, _ = ResponsesIDSecurity()._decrypt_response_id(unified_object_id)
|
||||
|
||||
|
||||
# Prepare metadata with model information for cost tracking
|
||||
litellm_metadata = {
|
||||
"user_api_key_user_id": job.created_by or "default-user-id",
|
||||
INTERNAL_CALL_ORIGIN_METADATA_KEY: BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN,
|
||||
}
|
||||
|
||||
|
||||
# Add model information if available
|
||||
if model_name:
|
||||
litellm_metadata["model"] = model_name
|
||||
litellm_metadata["model_group"] = model_name # Use same value for model_group
|
||||
|
||||
via_router = self._resolve_deployment(responses_id_security)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning(
|
||||
f"Skipping job {unified_object_id} due to error: {e}"
|
||||
)
|
||||
continue
|
||||
|
||||
provider_response_id = ResponsesAPIRequestUtils.decode_responses_api_response_id(
|
||||
responses_id_security
|
||||
).get("response_id", responses_id_security)
|
||||
try:
|
||||
response = await self._get_response(
|
||||
response_id=responses_id_security,
|
||||
litellm_metadata=litellm_metadata,
|
||||
via_router=via_router,
|
||||
)
|
||||
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"Response {unified_object_id} status: {response.status}, model: {model_name}"
|
||||
)
|
||||
|
||||
|
||||
except Exception as e:
|
||||
if via_router and _is_response_gone_at_provider(e, provider_response_id):
|
||||
verbose_proxy_logger.info(
|
||||
f"Response {unified_object_id} no longer available at provider (404), marking stale_expired: {e}"
|
||||
)
|
||||
expired_jobs.append(job)
|
||||
continue
|
||||
verbose_proxy_logger.warning(
|
||||
f"Skipping job {unified_object_id} due to error: {e}"
|
||||
)
|
||||
|
|
@ -210,6 +238,7 @@ class CheckResponsesCost:
|
|||
# Mark completed jobs in the database
|
||||
if len(completed_jobs) > 0:
|
||||
await _managed_object_table(self.prisma_client).update_many(
|
||||
# bounded-ok: at most MAX_OBJECTS_PER_POLL_CYCLE rows per cycle, the find_many take above
|
||||
where={"id": {"in": [job.id for job in completed_jobs]}},
|
||||
data={"status": "completed"},
|
||||
)
|
||||
|
|
@ -217,3 +246,13 @@ class CheckResponsesCost:
|
|||
f"Marked {len(completed_jobs)} response jobs as completed"
|
||||
)
|
||||
|
||||
if len(expired_jobs) > 0:
|
||||
await _managed_object_table(self.prisma_client).update_many(
|
||||
# bounded-ok: at most MAX_OBJECTS_PER_POLL_CYCLE rows per cycle, the find_many take above
|
||||
where={"id": {"in": [job.id for job in expired_jobs]}},
|
||||
data={"status": "stale_expired"},
|
||||
)
|
||||
verbose_proxy_logger.info(
|
||||
f"Marked {len(expired_jobs)} response jobs as stale_expired"
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
# Grandfathered findings of check_unbounded_in_lists.py: path::scope::kind::subject::occurrence.
|
||||
# Fix a site and delete its line; regenerate with `check_unbounded_in_lists.py --update-baseline`.
|
||||
enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py CheckResponsesCost.check_responses_cost prisma id.in `[job.id for job in completed_jobs]` 0
|
||||
litellm/integrations/shadow_eval_logger.py ShadowEvalLogger._active_jobs prisma job_id.in `[str(record.id) for record in records]` 0
|
||||
litellm/llms/litellm_proxy/skills/handler.py LiteLLMSkillsHandler.list_skills prisma created_by.in `owner_scopes` 0
|
||||
litellm/proxy/_experimental/mcp_server/db.py get_mcp_servers prisma server_id.in `server_ids` 0
|
||||
|
|
|
|||
|
|
@ -0,0 +1,370 @@
|
|||
import itertools
|
||||
import os
|
||||
import uuid
|
||||
from collections.abc import Iterator, Sequence
|
||||
from hashlib import sha256
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from integration._support.client import Gateway, Scenario, eventually, object_value, string_value
|
||||
from integration._support.database import read_rows, scratch_database, write_rows
|
||||
from integration._support.process import graceful_stop_seconds, owned_proxy_process
|
||||
from integration._support.upstream import ScenarioHandle, delete_scenario, register_scenario
|
||||
from integration.cost_calculation.cost_tracking_case import JsonResponse, RoutedResponse
|
||||
from pydantic import JsonValue
|
||||
|
||||
from litellm.constants import MAX_OBJECTS_PER_POLL_CYCLE
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
from litellm.types.utils import BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN
|
||||
|
||||
INPUT_COST_PER_TOKEN: Final = 0.001
|
||||
OUTPUT_COST_PER_TOKEN: Final = 0.002
|
||||
INPUT_TOKENS: Final = 19
|
||||
OUTPUT_TOKENS: Final = 7
|
||||
SCHEDULER_PERIOD_CEILING_SECONDS: Final = 31
|
||||
ONE_POLL_SECONDS: Final = 2 * SCHEDULER_PERIOD_CEILING_SECONDS + 8
|
||||
TWO_POLLS_SECONDS: Final = 4 * SCHEDULER_PERIOD_CEILING_SECONDS + 8
|
||||
OWNED_PROXY_CELL_SECONDS: Final = 2 * graceful_stop_seconds() + 240
|
||||
PROVIDER_ID: Final = "resp_$REQUEST_ID"
|
||||
|
||||
|
||||
def _response_body(status: str) -> dict[str, JsonValue]:
|
||||
completed: Final = status == "completed"
|
||||
return {
|
||||
"id": PROVIDER_ID,
|
||||
"object": "response",
|
||||
"created_at": 1,
|
||||
"status": status,
|
||||
"background": True,
|
||||
"store": False,
|
||||
"error": None,
|
||||
"incomplete_details": None,
|
||||
"model": "gpt-4o-mini",
|
||||
"output": (
|
||||
[
|
||||
{
|
||||
"id": "msg_$REQUEST_ID",
|
||||
"type": "message",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": "pong", "annotations": []}],
|
||||
}
|
||||
]
|
||||
if completed
|
||||
else []
|
||||
),
|
||||
"parallel_tool_calls": True,
|
||||
"temperature": 1.0,
|
||||
"tool_choice": "auto",
|
||||
"tools": [],
|
||||
"top_p": 1.0,
|
||||
"truncation": "disabled",
|
||||
"usage": (
|
||||
{
|
||||
"input_tokens": INPUT_TOKENS,
|
||||
"output_tokens": OUTPUT_TOKENS,
|
||||
"total_tokens": INPUT_TOKENS + OUTPUT_TOKENS,
|
||||
"input_tokens_details": {"cached_tokens": 0},
|
||||
"output_tokens_details": {"reasoning_tokens": 0},
|
||||
}
|
||||
if completed
|
||||
else None
|
||||
),
|
||||
"metadata": {},
|
||||
}
|
||||
|
||||
|
||||
def _json(body: dict[str, JsonValue], status: int = 200) -> JsonResponse:
|
||||
return JsonResponse(content_type="application/json", body=body, status=status)
|
||||
|
||||
|
||||
def _provider_404(message: str) -> JsonResponse:
|
||||
return _json({"error": {"message": message, "type": "invalid_request_error", "param": None, "code": None}}, 404)
|
||||
|
||||
|
||||
def _routes(submission_status: str, retrieve: JsonResponse) -> RoutedResponse:
|
||||
return RoutedResponse(
|
||||
content_type="application/x-routed",
|
||||
routes={"POST /responses": _json(_response_body(submission_status)), f"GET /responses/{PROVIDER_ID}": retrieve},
|
||||
)
|
||||
|
||||
|
||||
def _gone_routes(submission_status: str = "queued") -> RoutedResponse:
|
||||
return _routes(submission_status, _provider_404(f"Response with id '{PROVIDER_ID}' not found."))
|
||||
|
||||
|
||||
def _scenario_id(marker: str) -> str:
|
||||
return f"bg-{marker}-{uuid.uuid4().hex[:12]}"
|
||||
|
||||
|
||||
def _scripted_deployment(scenario: Scenario, marker: str, routes: RoutedResponse) -> tuple[str, ScenarioHandle]:
|
||||
return _deployment_on(scenario, _scenario_id(marker), routes)
|
||||
|
||||
|
||||
def _register_deployment(
|
||||
scenario: Scenario, scenario_id: str, routes: RoutedResponse
|
||||
) -> tuple[str, str, ScenarioHandle]:
|
||||
handle: Final = register_scenario(scenario_id, routes)
|
||||
scenario.cleanups.callback(delete_scenario, handle)
|
||||
created: Final = scenario.gateway.post(
|
||||
"/model/new",
|
||||
{
|
||||
"model_name": f"bg-{sha256(scenario_id.encode()).hexdigest()[:12]}",
|
||||
"litellm_params": {
|
||||
"model": "openai/gpt-4o-mini",
|
||||
"api_key": "sk-scripted-provider",
|
||||
"api_base": handle.api_base(),
|
||||
"input_cost_per_token": INPUT_COST_PER_TOKEN,
|
||||
"output_cost_per_token": OUTPUT_COST_PER_TOKEN,
|
||||
},
|
||||
},
|
||||
)
|
||||
return string_value(created["model_name"]), string_value(object_value(created["model_info"])["id"]), handle
|
||||
|
||||
|
||||
def _deployment_on(scenario: Scenario, scenario_id: str, routes: RoutedResponse) -> tuple[str, ScenarioHandle]:
|
||||
model_name, model_id, handle = _register_deployment(scenario, scenario_id, routes)
|
||||
scenario.cleanups.callback(scenario.delete_model, model_id)
|
||||
return model_name, handle
|
||||
|
||||
|
||||
def _forget_row(unified_id: str, database_url: str | None = None) -> None:
|
||||
write_rows(
|
||||
'DELETE FROM "LiteLLM_ManagedObjectTable" WHERE unified_object_id = %s',
|
||||
(unified_id,),
|
||||
database_url=database_url,
|
||||
)
|
||||
|
||||
|
||||
def _submit_background_response(scenario: Scenario, key: str, model: str, database_url: str | None = None) -> str:
|
||||
response: Final = scenario.gateway.request(
|
||||
"POST", "/v1/responses", {"model": model, "input": "poll me", "background": True, "store": False}, key=key
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
unified_id: Final = string_value(object_value(response.json())["id"])
|
||||
scenario.cleanups.callback(_forget_row, unified_id, database_url)
|
||||
return unified_id
|
||||
|
||||
|
||||
def _row_status(unified_id: str, database_url: str | None = None) -> list[dict[str, JsonValue]]:
|
||||
return read_rows(
|
||||
'SELECT status FROM "LiteLLM_ManagedObjectTable" WHERE unified_object_id = %s',
|
||||
(unified_id,),
|
||||
database_url=database_url,
|
||||
)
|
||||
|
||||
|
||||
def _await_status(unified_id: str, status: str, seconds: float, database_url: str | None = None) -> None:
|
||||
assert eventually(
|
||||
lambda: _row_status(unified_id, database_url), lambda rows: rows == [{"status": status}], seconds=seconds
|
||||
) == [{"status": status}]
|
||||
|
||||
|
||||
def _observed_requests(gateway: Gateway) -> tuple[dict[str, JsonValue], ...]:
|
||||
with httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream:
|
||||
return tuple(map(object_value, upstream.get("/__observations").json()["requests"]))
|
||||
|
||||
|
||||
def _calls_to(requests: Sequence[dict[str, JsonValue]], handle: ScenarioHandle) -> tuple[dict[str, JsonValue], ...]:
|
||||
return tuple(request for request in requests if string_value(request["path"]).startswith(f"/{handle.scenario_id}/"))
|
||||
|
||||
|
||||
def _upstream_log(gateway: Gateway, handle: ScenarioHandle) -> Iterator[tuple[dict[str, JsonValue], ...]]:
|
||||
fresh: Final = (_calls_to(_observed_requests(gateway), handle) for _ in itertools.count())
|
||||
return itertools.accumulate(fresh, lambda seen, calls: (*seen, *calls))
|
||||
|
||||
|
||||
def _polls(log: Sequence[dict[str, JsonValue]], handle: ScenarioHandle) -> tuple[dict[str, JsonValue], ...]:
|
||||
poll_path: Final = f"/{handle.scenario_id}/responses/resp_{handle.scenario_id}"
|
||||
return tuple(call for call in log if call["method"] == "GET" and call["path"] == poll_path)
|
||||
|
||||
|
||||
def _await_polls(gateway: Gateway, handle: ScenarioHandle, at_least: int, seconds: float) -> None:
|
||||
log: Final = _upstream_log(gateway, handle)
|
||||
polls: Final = _polls(
|
||||
eventually(lambda: next(log), lambda seen: len(_polls(seen, handle)) >= at_least, seconds), handle
|
||||
)
|
||||
assert len(polls) >= at_least, polls
|
||||
|
||||
|
||||
def _submissions(log: Sequence[dict[str, JsonValue]], handle: ScenarioHandle) -> tuple[dict[str, JsonValue], ...]:
|
||||
return tuple(
|
||||
call for call in log if call["method"] == "POST" and call["path"] == f"/{handle.scenario_id}/responses"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.timeout(240)
|
||||
@pytest.mark.parametrize("submission_status", ["queued", "in_progress"])
|
||||
def test_a_response_gone_at_the_provider_is_retired_from_polling(gateway: Gateway, submission_status: str) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
key: Final = scenario.key()
|
||||
model, handle = _scripted_deployment(scenario, "gone", _gone_routes(submission_status))
|
||||
log: Final = _upstream_log(gateway, handle)
|
||||
unified_id: Final = _submit_background_response(scenario, key, model)
|
||||
_await_status(unified_id, "stale_expired", ONE_POLL_SECONDS)
|
||||
seen: Final = next(log)
|
||||
assert [call["body"] for call in _submissions(seen, handle)] == [
|
||||
{"model": "gpt-4o-mini", "input": "poll me", "background": True, "store": False}
|
||||
]
|
||||
assert len(_polls(seen, handle)) >= 1, seen
|
||||
caller_view: Final = gateway.request("GET", f"/v1/responses/{unified_id}", key=key)
|
||||
assert caller_view.status_code == 404, caller_view.text
|
||||
assert f"Response with id 'resp_{handle.scenario_id}' not found." in caller_view.text
|
||||
first_clock, _ = _scripted_deployment(scenario, "clock1", _gone_routes())
|
||||
_await_status(_submit_background_response(scenario, key, first_clock), "stale_expired", ONE_POLL_SECONDS)
|
||||
polls_after_first_clock: Final = len(_polls(next(log), handle))
|
||||
second_clock, _ = _scripted_deployment(scenario, "clock2", _gone_routes())
|
||||
_await_status(_submit_background_response(scenario, key, second_clock), "stale_expired", ONE_POLL_SECONDS)
|
||||
assert len(_polls(next(log), handle)) == polls_after_first_clock
|
||||
assert _row_status(unified_id) == [{"status": "stale_expired"}]
|
||||
|
||||
|
||||
@pytest.mark.timeout(180)
|
||||
def test_a_404_that_does_not_name_the_response_keeps_the_row_queued_for_retry(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
model, handle = _scripted_deployment(scenario, "vague404", _routes("queued", _provider_404("Not found.")))
|
||||
unified_id: Final = _submit_background_response(scenario, scenario.key(), model)
|
||||
_await_polls(gateway, handle, 2, TWO_POLLS_SECONDS)
|
||||
assert _row_status(unified_id) == [{"status": "queued"}]
|
||||
|
||||
|
||||
@pytest.mark.timeout(180)
|
||||
def test_a_provider_error_other_than_404_keeps_the_row_queued_for_retry(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
outage: Final = _json({"error": {"message": "The server had an error.", "type": "server_error"}}, 500)
|
||||
model, handle = _scripted_deployment(scenario, "outage", _routes("queued", outage))
|
||||
unified_id: Final = _submit_background_response(scenario, scenario.key(), model)
|
||||
_await_polls(gateway, handle, 2, TWO_POLLS_SECONDS)
|
||||
assert _row_status(unified_id) == [{"status": "queued"}]
|
||||
|
||||
|
||||
def _polled_provider_id(row: dict[str, JsonValue]) -> str:
|
||||
return ResponsesAPIRequestUtils.decode_responses_api_response_id(string_value(row["request_id"]))["response_id"]
|
||||
|
||||
|
||||
def _poll_spend_rows(handle: ScenarioHandle) -> list[dict[str, JsonValue]]:
|
||||
billed: Final = read_rows(
|
||||
'SELECT request_id, call_type, status, prompt_tokens, completion_tokens, spend FROM "LiteLLM_SpendLogs" '
|
||||
"WHERE metadata->>'internal_call_origin' = %s",
|
||||
(BACKGROUND_RESPONSE_COST_POLL_CALL_ORIGIN,),
|
||||
)
|
||||
return [
|
||||
{column: value for column, value in row.items() if column != "request_id"}
|
||||
for row in billed
|
||||
if _polled_provider_id(row) == f"resp_{handle.scenario_id}"
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.timeout(180)
|
||||
def test_a_completed_response_is_marked_completed_and_its_poll_is_billed(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
model, handle = _scripted_deployment(scenario, "done", _routes("queued", _json(_response_body("completed"))))
|
||||
unified_id: Final = _submit_background_response(scenario, scenario.key(), model)
|
||||
_await_status(unified_id, "completed", ONE_POLL_SECONDS)
|
||||
spend_rows: Final = eventually(lambda: _poll_spend_rows(handle), lambda rows: len(rows) >= 1, seconds=70)
|
||||
assert spend_rows[0] == {
|
||||
"call_type": "aget_responses",
|
||||
"status": "success",
|
||||
"prompt_tokens": INPUT_TOKENS,
|
||||
"completion_tokens": OUTPUT_TOKENS,
|
||||
"spend": pytest.approx(INPUT_TOKENS * INPUT_COST_PER_TOKEN + OUTPUT_TOKENS * OUTPUT_COST_PER_TOKEN),
|
||||
}
|
||||
|
||||
|
||||
def _insert_unreadable_row(scenario: Scenario) -> str:
|
||||
unified_id: Final = f"resp_unreadable-{uuid.uuid4().hex[:12]}"
|
||||
write_rows(
|
||||
'INSERT INTO "LiteLLM_ManagedObjectTable" '
|
||||
'("id", "unified_object_id", "model_object_id", "file_object", "file_purpose", "status", "created_at", "updated_at") '
|
||||
"VALUES (%s, %s, %s, '[]'::jsonb, 'response', 'queued', NOW() - INTERVAL '1 hour', NOW())",
|
||||
(str(uuid.uuid4()), unified_id, unified_id),
|
||||
)
|
||||
scenario.cleanups.callback(_forget_row, unified_id)
|
||||
return unified_id
|
||||
|
||||
|
||||
@pytest.mark.timeout(180)
|
||||
def test_a_row_the_poll_cannot_prepare_skips_only_that_row(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
unreadable_id: Final = _insert_unreadable_row(scenario)
|
||||
model, _ = _scripted_deployment(scenario, "gone", _gone_routes())
|
||||
unified_id: Final = _submit_background_response(scenario, scenario.key(), model)
|
||||
_await_status(unified_id, "stale_expired", ONE_POLL_SECONDS)
|
||||
assert _row_status(unreadable_id) == [{"status": "queued"}]
|
||||
|
||||
|
||||
@pytest.mark.timeout(300)
|
||||
def test_responses_gone_at_the_provider_do_not_starve_a_newer_response_out_of_cost_polling(gateway: Gateway) -> None:
|
||||
with gateway.scenario() as scenario:
|
||||
key: Final = scenario.key()
|
||||
gone_ids: Final = tuple(
|
||||
_submit_background_response(
|
||||
scenario, key, _scripted_deployment(scenario, f"gone{index}", _gone_routes())[0]
|
||||
)
|
||||
for index in range(MAX_OBJECTS_PER_POLL_CYCLE)
|
||||
)
|
||||
completable_model, completable = _scripted_deployment(
|
||||
scenario, "done", _routes("queued", _json(_response_body("completed")))
|
||||
)
|
||||
completable_id: Final = _submit_background_response(scenario, key, completable_model)
|
||||
_await_status(completable_id, "completed", TWO_POLLS_SECONDS)
|
||||
assert [_row_status(gone_id) for gone_id in gone_ids] == [
|
||||
[{"status": "stale_expired"}]
|
||||
] * MAX_OBJECTS_PER_POLL_CYCLE
|
||||
spend_rows: Final = eventually(lambda: _poll_spend_rows(completable), lambda rows: len(rows) >= 1, seconds=70)
|
||||
assert spend_rows[0]["spend"] == pytest.approx(
|
||||
INPUT_TOKENS * INPUT_COST_PER_TOKEN + OUTPUT_TOKENS * OUTPUT_COST_PER_TOKEN
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.timeout(OWNED_PROXY_CELL_SECONDS)
|
||||
def test_a_404_on_a_response_whose_deployment_left_the_router_keeps_the_row_queued_for_retry(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
deployment_scenario_id: Final = _scenario_id("left")
|
||||
provider_id: Final = f"resp_{deployment_scenario_id}"
|
||||
env_handle: Final = register_scenario(
|
||||
_scenario_id("envbase"),
|
||||
RoutedResponse(
|
||||
content_type="application/x-routed",
|
||||
routes={f"GET /responses/{provider_id}": _provider_404(f"Response with id '{provider_id}' not found.")},
|
||||
),
|
||||
)
|
||||
try:
|
||||
with scratch_database() as database_url:
|
||||
overrides: Final = {
|
||||
"DATABASE_URL": database_url,
|
||||
"OPENAI_API_BASE": env_handle.api_base(),
|
||||
"OPENAI_API_KEY": "sk-scripted-provider",
|
||||
}
|
||||
with owned_proxy_process(
|
||||
gateway, tmp_path, overrides, remove_environment=("DATABASE_URL_READ_REPLICA",)
|
||||
) as owned:
|
||||
with owned.gateway.scenario() as scenario:
|
||||
outage: Final = _json(
|
||||
{"error": {"message": "The server had an error.", "type": "server_error"}}, 500
|
||||
)
|
||||
model, model_id, _ = _register_deployment(
|
||||
scenario, deployment_scenario_id, _routes("queued", outage)
|
||||
)
|
||||
unified_id: Final = _submit_background_response(scenario, scenario.key(), model, database_url)
|
||||
owned.gateway.post("/model/delete", {"id": model_id})
|
||||
log: Final = _upstream_log(gateway, env_handle)
|
||||
fallback_polls: Final = eventually(
|
||||
lambda: next(log),
|
||||
lambda seen: len(_fallback_polls(seen, env_handle, provider_id)) >= 2,
|
||||
seconds=TWO_POLLS_SECONDS,
|
||||
)
|
||||
assert len(_fallback_polls(fallback_polls, env_handle, provider_id)) >= 2, fallback_polls
|
||||
assert _row_status(unified_id, database_url) == [{"status": "queued"}]
|
||||
finally:
|
||||
delete_scenario(env_handle)
|
||||
|
||||
|
||||
def _fallback_polls(
|
||||
log: Sequence[dict[str, JsonValue]], env_handle: ScenarioHandle, provider_id: str
|
||||
) -> tuple[dict[str, JsonValue], ...]:
|
||||
poll_path: Final = f"/{env_handle.scenario_id}/responses/{provider_id}"
|
||||
return tuple(call for call in log if call["method"] == "GET" and call["path"] == poll_path)
|
||||
|
|
@ -3,7 +3,9 @@ Unit tests for CheckResponsesCost class
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Mapping
|
||||
from datetime import datetime
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -12,6 +14,19 @@ from litellm.constants import MAX_OBJECTS_PER_POLL_CYCLE
|
|||
from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse
|
||||
|
||||
|
||||
class _RecordingManagedObjectTable:
|
||||
def __init__(self, rows: tuple[object, ...]) -> None:
|
||||
self.rows: Final = rows
|
||||
self.updates: tuple[tuple[Mapping[str, object], Mapping[str, object]], ...] = ()
|
||||
|
||||
async def find_many(self, **query: object) -> tuple[object, ...]:
|
||||
return self.rows
|
||||
|
||||
async def update_many(self, where: Mapping[str, object], data: Mapping[str, object]) -> int:
|
||||
self.updates = (*self.updates, (where, data))
|
||||
return len(self.rows)
|
||||
|
||||
|
||||
class TestCheckResponsesCost:
|
||||
"""Test suite for CheckResponsesCost class"""
|
||||
|
||||
|
|
@ -364,6 +379,264 @@ class TestCheckResponsesCost:
|
|||
# Stale cleanup still ran via _expire_stale_rows
|
||||
check_responses_cost_instance._expire_stale_rows.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_responses_cost_marks_404_response_stale_expired(
|
||||
self, check_responses_cost_instance, mock_prisma_client, mock_llm_router
|
||||
):
|
||||
"""A provider 404 on a response polled through its deployment marks the row stale_expired."""
|
||||
import litellm
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
|
||||
encoded_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id(
|
||||
custom_llm_provider="openai",
|
||||
model_id="deployment-404",
|
||||
response_id="resp_upstream_404",
|
||||
)
|
||||
|
||||
mock_job = MagicMock()
|
||||
mock_job.unified_object_id = encoded_response_id
|
||||
mock_job.created_by = "test-user"
|
||||
mock_job.id = "job-404"
|
||||
mock_job.file_object = {"model": "gpt-4o", "id": encoded_response_id}
|
||||
|
||||
mock_llm_router.get_deployment.return_value = {"model_id": "deployment-404"}
|
||||
mock_llm_router.aget_responses = AsyncMock(
|
||||
side_effect=litellm.NotFoundError(
|
||||
message="Response with id 'resp_upstream_404' not found.", model="gpt-5", llm_provider="openai"
|
||||
)
|
||||
)
|
||||
|
||||
table = _RecordingManagedObjectTable(rows=(mock_job,))
|
||||
mock_prisma_client.db.litellm_managedobjecttable = table
|
||||
|
||||
with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_sdk_aget:
|
||||
await check_responses_cost_instance.check_responses_cost()
|
||||
|
||||
mock_sdk_aget.assert_not_called()
|
||||
mock_llm_router.get_deployment.assert_called_once_with(model_id="deployment-404")
|
||||
assert mock_llm_router.aget_responses.call_args.kwargs["response_id"] == encoded_response_id
|
||||
assert table.updates == (({"id": {"in": ["job-404"]}}, {"status": "stale_expired"}),)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_responses_cost_marks_mapped_provider_404_stale_expired(
|
||||
self, check_responses_cost_instance, mock_prisma_client, mock_llm_router
|
||||
):
|
||||
"""The provider's own 404 body, mapped the way the GET path maps it, still expires the row."""
|
||||
import json
|
||||
|
||||
import openai
|
||||
|
||||
import litellm
|
||||
from litellm.llms.openai.common_utils import OpenAIError
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
|
||||
encoded_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id(
|
||||
custom_llm_provider="openai",
|
||||
model_id="deployment-404-mapped",
|
||||
response_id="resp_upstream_404_mapped",
|
||||
)
|
||||
provider_body = json.dumps(
|
||||
{
|
||||
"error": {
|
||||
"message": "Response with id 'resp_upstream_404_mapped' not found.",
|
||||
"type": "invalid_request_error",
|
||||
"param": None,
|
||||
"code": None,
|
||||
}
|
||||
}
|
||||
)
|
||||
with pytest.raises(openai.APIStatusError) as mapped:
|
||||
raise litellm.exception_type(
|
||||
model="gpt-5.5",
|
||||
custom_llm_provider="openai",
|
||||
original_exception=OpenAIError(message=provider_body, status_code=404),
|
||||
completion_kwargs={},
|
||||
extra_kwargs={},
|
||||
)
|
||||
assert mapped.value.status_code == 404
|
||||
|
||||
mock_job = MagicMock()
|
||||
mock_job.unified_object_id = encoded_response_id
|
||||
mock_job.created_by = "test-user"
|
||||
mock_job.id = "job-404-mapped"
|
||||
mock_job.file_object = {"model": "gpt-5.5", "id": encoded_response_id}
|
||||
|
||||
mock_llm_router.get_deployment.return_value = {"model_id": "deployment-404-mapped"}
|
||||
mock_llm_router.aget_responses = AsyncMock(side_effect=mapped.value)
|
||||
|
||||
table = _RecordingManagedObjectTable(rows=(mock_job,))
|
||||
mock_prisma_client.db.litellm_managedobjecttable = table
|
||||
|
||||
with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_sdk_aget:
|
||||
await check_responses_cost_instance.check_responses_cost()
|
||||
|
||||
mock_sdk_aget.assert_not_called()
|
||||
assert table.updates == (({"id": {"in": ["job-404-mapped"]}}, {"status": "stale_expired"}),)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_responses_cost_404_without_router_deployment_keeps_row_for_retry(
|
||||
self, check_responses_cost_instance, mock_prisma_client, mock_llm_router
|
||||
):
|
||||
"""A 404 while polling without a resolved deployment skips the row for retry."""
|
||||
import litellm
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
|
||||
encoded_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id(
|
||||
custom_llm_provider="openai",
|
||||
model_id="deployment-gone",
|
||||
response_id="resp_upstream_gone",
|
||||
)
|
||||
|
||||
mock_job = MagicMock()
|
||||
mock_job.unified_object_id = encoded_response_id
|
||||
mock_job.created_by = "test-user"
|
||||
mock_job.id = "job-404-fallback"
|
||||
mock_job.file_object = {"model": "gpt-4o", "id": encoded_response_id}
|
||||
|
||||
mock_llm_router.get_deployment.return_value = None
|
||||
|
||||
table = _RecordingManagedObjectTable(rows=(mock_job,))
|
||||
mock_prisma_client.db.litellm_managedobjecttable = table
|
||||
|
||||
with patch(
|
||||
"litellm.aget_responses",
|
||||
new_callable=AsyncMock,
|
||||
side_effect=litellm.NotFoundError(
|
||||
message="Response not found", model="gpt-5", llm_provider="openai"
|
||||
),
|
||||
) as mock_sdk_aget:
|
||||
await check_responses_cost_instance.check_responses_cost()
|
||||
|
||||
mock_sdk_aget.assert_awaited_once()
|
||||
mock_llm_router.aget_responses.assert_not_called()
|
||||
assert table.updates == ()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_responses_cost_404_not_naming_the_response_keeps_row_for_retry(
|
||||
self, check_responses_cost_instance, mock_prisma_client, mock_llm_router
|
||||
):
|
||||
"""A 404 that does not name the response (a gateway or a reconfigured deployment) is retried, not expired."""
|
||||
import litellm
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
|
||||
encoded_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id(
|
||||
custom_llm_provider="azure",
|
||||
model_id="deployment-renamed",
|
||||
response_id="resp_upstream_still_there",
|
||||
)
|
||||
|
||||
mock_job = MagicMock()
|
||||
mock_job.unified_object_id = encoded_response_id
|
||||
mock_job.created_by = "test-user"
|
||||
mock_job.id = "job-404-other"
|
||||
mock_job.file_object = {"model": "azure-gpt-5", "id": encoded_response_id}
|
||||
|
||||
mock_llm_router.get_deployment.return_value = {"model_id": "deployment-renamed"}
|
||||
mock_llm_router.aget_responses = AsyncMock(
|
||||
side_effect=litellm.NotFoundError(
|
||||
message="Resource not found", model="gpt-5", llm_provider="azure"
|
||||
)
|
||||
)
|
||||
|
||||
table = _RecordingManagedObjectTable(rows=(mock_job,))
|
||||
mock_prisma_client.db.litellm_managedobjecttable = table
|
||||
|
||||
with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_sdk_aget:
|
||||
await check_responses_cost_instance.check_responses_cost()
|
||||
|
||||
mock_sdk_aget.assert_not_called()
|
||||
assert mock_llm_router.aget_responses.call_args.kwargs["response_id"] == encoded_response_id
|
||||
assert table.updates == ()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_responses_cost_deployment_lookup_error_skips_only_that_job(
|
||||
self, check_responses_cost_instance, mock_prisma_client, mock_llm_router
|
||||
):
|
||||
"""A deployment lookup that raises skips its own job and the cycle still records the next job."""
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
|
||||
broken_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id(
|
||||
custom_llm_provider="openai", model_id="deployment-broken", response_id="resp_upstream_broken"
|
||||
)
|
||||
healthy_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id(
|
||||
custom_llm_provider="openai", model_id="deployment-healthy", response_id="resp_upstream_healthy"
|
||||
)
|
||||
|
||||
broken_job = MagicMock()
|
||||
broken_job.unified_object_id = broken_response_id
|
||||
broken_job.created_by = "test-user"
|
||||
broken_job.id = "job-broken"
|
||||
broken_job.file_object = {"model": "gpt-5.5", "id": broken_response_id}
|
||||
|
||||
healthy_job = MagicMock()
|
||||
healthy_job.unified_object_id = healthy_response_id
|
||||
healthy_job.created_by = "test-user"
|
||||
healthy_job.id = "job-healthy"
|
||||
healthy_job.file_object = {"model": "gpt-5.5", "id": healthy_response_id}
|
||||
|
||||
mock_llm_router.get_deployment.side_effect = [
|
||||
Exception("Model invalid format - <class 'str'>"),
|
||||
{"model_id": "deployment-healthy"},
|
||||
]
|
||||
mock_llm_router.aget_responses = AsyncMock(
|
||||
return_value=ResponsesAPIResponse(
|
||||
id=healthy_response_id,
|
||||
object="response",
|
||||
status="completed",
|
||||
created_at=int(datetime.now().timestamp()),
|
||||
output=[],
|
||||
usage=ResponseAPIUsage(input_tokens=100, output_tokens=50, total_tokens=150),
|
||||
)
|
||||
)
|
||||
|
||||
table = _RecordingManagedObjectTable(rows=(broken_job, healthy_job))
|
||||
mock_prisma_client.db.litellm_managedobjecttable = table
|
||||
|
||||
with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_sdk_aget:
|
||||
await check_responses_cost_instance.check_responses_cost()
|
||||
|
||||
mock_sdk_aget.assert_not_called()
|
||||
mock_llm_router.aget_responses.assert_awaited_once()
|
||||
assert mock_llm_router.aget_responses.call_args.kwargs["response_id"] == healthy_response_id
|
||||
assert table.updates == (({"id": {"in": ["job-healthy"]}}, {"status": "completed"}),)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_responses_cost_non_404_error_keeps_row_for_retry(
|
||||
self, check_responses_cost_instance, mock_prisma_client, mock_llm_router
|
||||
):
|
||||
"""A non-404 provider error through a resolved deployment skips the job so it is retried next cycle."""
|
||||
import litellm
|
||||
from litellm.responses.utils import ResponsesAPIRequestUtils
|
||||
|
||||
encoded_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id(
|
||||
custom_llm_provider="openai",
|
||||
model_id="deployment-500",
|
||||
response_id="resp_upstream_500",
|
||||
)
|
||||
|
||||
mock_job = MagicMock()
|
||||
mock_job.unified_object_id = encoded_response_id
|
||||
mock_job.created_by = "test-user"
|
||||
mock_job.id = "job-500"
|
||||
mock_job.file_object = {"model": "gpt-4o", "id": encoded_response_id}
|
||||
|
||||
mock_llm_router.get_deployment.return_value = {"model_id": "deployment-500"}
|
||||
mock_llm_router.aget_responses = AsyncMock(
|
||||
side_effect=litellm.InternalServerError(
|
||||
message="boom", model="gpt-5", llm_provider="openai"
|
||||
)
|
||||
)
|
||||
|
||||
table = _RecordingManagedObjectTable(rows=(mock_job,))
|
||||
mock_prisma_client.db.litellm_managedobjecttable = table
|
||||
|
||||
with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_sdk_aget:
|
||||
await check_responses_cost_instance.check_responses_cost()
|
||||
|
||||
mock_sdk_aget.assert_not_called()
|
||||
mock_llm_router.aget_responses.assert_awaited_once()
|
||||
assert table.updates == ()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_responses_cost_multiple_jobs(
|
||||
self, check_responses_cost_instance, mock_prisma_client
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue