mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
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:
commit
b66d4e6965
2 changed files with 309 additions and 15 deletions
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue