Merge pull request #35137 from BerriAI/litellm_fix_responses_cost_router_35131

fix(proxy): fetch background responses through the router in CheckResponsesCost
This commit is contained in:
Mateo Wang 2026-08-06 03:26:36 -07:00 committed by GitHub
commit b66d4e6965
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 309 additions and 15 deletions

View file

@ -1,10 +1,10 @@
"""
Polls LiteLLM_ManagedObjectTable to check if the response is complete.
Cost tracking is handled automatically by litellm.aget_responses().
Cost tracking is handled automatically by the get-responses call.
"""
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Dict, Optional, cast
import litellm
from litellm._logging import verbose_proxy_logger
@ -13,11 +13,15 @@ from litellm.constants import (
MAX_OBJECTS_PER_POLL_CYCLE,
STALE_OBJECT_CLEANUP_BATCH_SIZE,
)
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.llms.openai import ResponsesAPIResponse
if TYPE_CHECKING:
from litellm.proxy.utils import PrismaClient, ProxyLogging
from litellm.router import Router
TERMINAL_RESPONSE_STATUSES = frozenset({"completed", "failed", "cancelled", "incomplete"})
class CheckResponsesCost:
def __init__(
@ -33,6 +37,28 @@ class CheckResponsesCost:
self.prisma_client: PrismaClient = prisma_client
self.llm_router: Router = llm_router
async def _get_response(
self,
response_id: str,
litellm_metadata: Dict[str, str],
) -> ResponsesAPIResponse:
"""Fetch the upstream response, using deployment credentials when available.
LiteLLM-encoded response IDs carry the ``model_id`` of the deployment that
served the original request, so routing through ``llm_router`` applies that
deployment's ``api_base`` / ``api_key`` / ``api_version``, exactly like
``GET /v1/responses/{id}`` does. ``litellm.aget_responses`` on its own only
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:
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
)
return cast(ResponsesAPIResponse, router_response)
async def _expire_stale_rows(
self, cutoff: datetime, batch_size: int
) -> int:
@ -87,8 +113,8 @@ class CheckResponsesCost:
Check if background responses are complete and track their cost.
- Get all status="queued" or "in_progress" and file_purpose="response" jobs
- Query the provider to check if response is complete
- Cost is automatically tracked by litellm.aget_responses()
- Mark completed/failed/cancelled responses as complete in the database
- Cost is automatically tracked by the get-responses call
- Mark responses in a terminal state as complete in the database
"""
try:
await self._cleanup_stale_managed_objects()
@ -134,7 +160,7 @@ class CheckResponsesCost:
litellm_metadata["model"] = model_name
litellm_metadata["model_group"] = model_name # Use same value for model_group
response = await litellm.aget_responses(
response = await self._get_response(
response_id=responses_id_security,
litellm_metadata=litellm_metadata,
)
@ -144,21 +170,14 @@ class CheckResponsesCost:
)
except Exception as e:
verbose_proxy_logger.info(
verbose_proxy_logger.warning(
f"Skipping job {unified_object_id} due to error: {e}"
)
continue
# Check if response is in a terminal state
if response.status == "completed":
if response.status in TERMINAL_RESPONSE_STATUSES:
verbose_proxy_logger.info(
f"Response {unified_object_id} is complete. Cost automatically tracked by aget_responses."
)
completed_jobs.append(job)
elif response.status in ["failed", "cancelled"]:
verbose_proxy_logger.info(
f"Response {unified_object_id} has status {response.status}, marking as complete"
f"Response {unified_object_id} has terminal status {response.status}, marking as complete"
)
completed_jobs.append(job)

View file

@ -449,6 +449,281 @@ class TestCheckResponsesCost:
assert "job-3" in completion_call[1]["where"]["id"]["in"]
assert "job-2" not in completion_call[1]["where"]["id"]["in"]
@pytest.mark.asyncio
async def test_encoded_response_id_is_fetched_through_router(
self, check_responses_cost_instance, mock_prisma_client, mock_llm_router
):
"""
Regression test for https://github.com/BerriAI/litellm/issues/35131
A background response created against a deployment whose credentials only
exist in the config (e.g. Azure api_base/api_key) must be fetched through
the router so the deployment credentials are applied. Calling
litellm.aget_responses directly only sees provider env vars, fails, and
leaves the row in "queued" forever.
"""
from litellm.responses.utils import ResponsesAPIRequestUtils
encoded_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id(
custom_llm_provider="azure",
model_id="deployment-abc",
response_id="resp_upstream_123",
)
mock_job = MagicMock()
mock_job.unified_object_id = encoded_response_id
mock_job.created_by = "test-user"
mock_job.id = "job-router"
mock_job.file_object = {"model": "azure-gpt-5", "id": encoded_response_id}
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
)
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
mock_llm_router.aget_responses = AsyncMock(
return_value=ResponsesAPIResponse(
id=encoded_response_id,
object="response",
status="completed",
created_at=int(datetime.now().timestamp()),
output=[],
usage=ResponseAPIUsage(
input_tokens=100, output_tokens=50, total_tokens=150
),
)
)
with patch(
"litellm.aget_responses",
new_callable=AsyncMock,
side_effect=AssertionError(
"must not bypass the router for a deployment-scoped response id"
),
) 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[1]["response_id"]
== encoded_response_id
)
calls = (
mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list
)
assert len(calls) == 1
assert calls[0][1]["data"]["status"] == "completed"
assert calls[0][1]["where"]["id"]["in"] == ["job-router"]
@pytest.mark.asyncio
async def test_encrypted_response_id_is_fetched_through_router(
self, check_responses_cost_instance, mock_prisma_client, mock_llm_router, monkeypatch
):
"""
Rows store the *encrypted* response id when responses id security is on.
After decryption the id still carries the deployment model_id, so the
fetch must go through the router (issue #35131).
"""
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.utils import SpecialEnums
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-test-salt-key-for-response-ids")
encoded_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id(
custom_llm_provider="openai",
model_id="deployment-xyz",
response_id="resp_upstream_456",
)
encrypted_response_id = "resp_" + str(
encrypt_value_helper(
value=SpecialEnums.LITELLM_MANAGED_RESPONSE_API_RESPONSE_ID_COMPLETE_STR.value.format(
encoded_response_id, "test-user", "test-team"
)
)
)
mock_job = MagicMock()
mock_job.unified_object_id = encrypted_response_id
mock_job.created_by = "test-user"
mock_job.id = "job-encrypted"
mock_job.file_object = {"model": "gpt-5", "id": encrypted_response_id}
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
)
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
mock_llm_router.aget_responses = AsyncMock(
return_value=ResponsesAPIResponse(
id=encoded_response_id,
object="response",
status="completed",
created_at=int(datetime.now().timestamp()),
output=[],
usage=None,
)
)
with patch(
"litellm.aget_responses",
new_callable=AsyncMock,
side_effect=AssertionError(
"must not bypass the router for a deployment-scoped response id"
),
) 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[1]["response_id"]
== encoded_response_id
)
calls = (
mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list
)
assert len(calls) == 1
assert calls[0][1]["where"]["id"]["in"] == ["job-encrypted"]
@pytest.mark.asyncio
async def test_response_id_without_model_id_uses_sdk(
self, check_responses_cost_instance, mock_prisma_client, mock_llm_router
):
"""Ids that carry no deployment info can't be routed, so fall back to the SDK."""
mock_job = MagicMock()
mock_job.unified_object_id = "resp_plain_upstream_id"
mock_job.created_by = "test-user"
mock_job.id = "job-plain"
mock_job.file_object = {"model": "gpt-5", "id": "resp_plain_upstream_id"}
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
)
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
mock_llm_router.aget_responses = AsyncMock(
side_effect=AssertionError("router cannot route an id without a model_id")
)
mock_response = ResponsesAPIResponse(
id="resp_plain_upstream_id",
object="response",
status="completed",
created_at=int(datetime.now().timestamp()),
output=[],
usage=None,
)
with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_sdk_aget:
mock_sdk_aget.return_value = mock_response
await check_responses_cost_instance.check_responses_cost()
mock_sdk_aget.assert_called_once()
mock_llm_router.aget_responses.assert_not_called()
@pytest.mark.asyncio
async def test_missing_deployment_falls_back_to_sdk(
self, check_responses_cost_instance, mock_prisma_client, mock_llm_router
):
"""
An encoded id whose deployment was removed from the router must fall back
to the SDK so provider env credentials can still retrieve it, instead of
failing every poll cycle until stale expiration.
"""
from litellm.responses.utils import ResponsesAPIRequestUtils
encoded_response_id = ResponsesAPIRequestUtils._build_responses_api_response_id(
custom_llm_provider="openai",
model_id="deployment-deleted",
response_id="resp_upstream_789",
)
mock_job = MagicMock()
mock_job.unified_object_id = encoded_response_id
mock_job.created_by = "test-user"
mock_job.id = "job-missing-deployment"
mock_job.file_object = {"model": "gpt-5", "id": encoded_response_id}
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
)
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
mock_llm_router.get_deployment = MagicMock(return_value=None)
mock_llm_router.aget_responses = AsyncMock(
side_effect=AssertionError("router has no deployment for this model_id")
)
mock_response = ResponsesAPIResponse(
id=encoded_response_id,
object="response",
status="completed",
created_at=int(datetime.now().timestamp()),
output=[],
usage=None,
)
with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_sdk_aget:
mock_sdk_aget.return_value = mock_response
await check_responses_cost_instance.check_responses_cost()
mock_llm_router.get_deployment.assert_called_once_with(model_id="deployment-deleted")
mock_llm_router.aget_responses.assert_not_called()
mock_sdk_aget.assert_called_once()
assert mock_sdk_aget.call_args[1]["response_id"] == encoded_response_id
calls = (
mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list
)
assert len(calls) == 1
assert calls[0][1]["data"]["status"] == "completed"
assert calls[0][1]["where"]["id"]["in"] == ["job-missing-deployment"]
@pytest.mark.asyncio
async def test_check_responses_cost_with_incomplete_response(
self, check_responses_cost_instance, mock_prisma_client
):
"""'incomplete' is terminal in the Responses API, so the row must not stay queued."""
mock_job = MagicMock()
mock_job.unified_object_id = "resp_test_incomplete"
mock_job.created_by = "test-user"
mock_job.id = "job-incomplete"
mock_job.file_object = {"model": "gpt-5", "id": "resp_test_incomplete"}
mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock(
return_value=[mock_job]
)
mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(
return_value=0
)
mock_response = ResponsesAPIResponse(
id="resp_incomplete",
object="response",
status="incomplete",
created_at=int(datetime.now().timestamp()),
output=[],
usage=None,
)
with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget:
mock_aget.return_value = mock_response
await check_responses_cost_instance.check_responses_cost()
calls = (
mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args_list
)
assert len(calls) == 1
assert calls[0][1]["data"]["status"] == "completed"
assert calls[0][1]["where"]["id"]["in"] == ["job-incomplete"]
@pytest.mark.asyncio
async def test_check_responses_cost_no_model_in_file_object(
self, check_responses_cost_instance, mock_prisma_client