diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py new file mode 100644 index 00000000000..4ee6a89cc98 --- /dev/null +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_responses_cost.py @@ -0,0 +1,110 @@ +""" +Polls LiteLLM_ManagedObjectTable to check if the response is complete. +Cost tracking is handled automatically by litellm.aget_responses(). +""" + +from typing import TYPE_CHECKING + +import litellm +from litellm._logging import verbose_proxy_logger + +if TYPE_CHECKING: + from litellm.proxy.utils import PrismaClient, ProxyLogging + from litellm.router import Router + + +class CheckResponsesCost: + def __init__( + self, + proxy_logging_obj: "ProxyLogging", + prisma_client: "PrismaClient", + llm_router: "Router", + ): + from litellm.proxy.utils import PrismaClient, ProxyLogging + from litellm.router import Router + + self.proxy_logging_obj: ProxyLogging = proxy_logging_obj + self.prisma_client: PrismaClient = prisma_client + self.llm_router: Router = llm_router + + async def check_responses_cost(self): + """ + 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 + """ + jobs = await self.prisma_client.db.litellm_managedobjecttable.find_many( + where={ + "status": {"in": ["queued", "in_progress"]}, + "file_purpose": "response", + } + ) + + verbose_proxy_logger.debug(f"Found {len(jobs)} response jobs to check") + completed_jobs = [] + + for job in jobs: + unified_object_id = job.unified_object_id + + try: + from litellm.proxy.hooks.responses_id_security import ( + ResponsesIDSecurity, + ) + + # 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", + } + + # 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 + + response = await litellm.aget_responses( + response_id=responses_id_security, + litellm_metadata=litellm_metadata, + ) + + verbose_proxy_logger.debug( + f"Response {unified_object_id} status: {response.status}, model: {model_name}" + ) + + except Exception as e: + verbose_proxy_logger.info( + f"Skipping job {unified_object_id} due to error: {e}" + ) + continue + + # Check if response is in a terminal state + if response.status == "completed": + 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" + ) + completed_jobs.append(job) + + # Mark completed jobs in the database + if len(completed_jobs) > 0: + await self.prisma_client.db.litellm_managedobjecttable.update_many( + where={"id": {"in": [job.id for job in completed_jobs]}}, + data={"status": "completed"}, + ) + verbose_proxy_logger.info( + f"Marked {len(completed_jobs)} response jobs as completed" + ) + diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index e12be6baf5d..a83d7e224b5 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -23,7 +23,9 @@ from litellm.proxy._types import ( from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, get_batch_id_from_unified_batch_id, + get_content_type_from_file_object, get_model_id_from_unified_batch_id, + normalize_mime_type_for_provider, ) from litellm.types.llms.openai import ( AllMessageValues, @@ -33,6 +35,7 @@ from litellm.types.llms.openai import ( FileObject, OpenAIFileObject, OpenAIFilesPurpose, + ResponsesAPIResponse, ) from litellm.types.utils import ( CallTypesLiteral, @@ -41,10 +44,6 @@ from litellm.types.utils import ( LLMResponseTypes, SpecialEnums, ) -from litellm.proxy.openai_files_endpoints.common_utils import ( - get_content_type_from_file_object, - normalize_mime_type_for_provider, -) if TYPE_CHECKING: from litellm.types.llms.openai import HttpxBinaryResponseContent @@ -133,10 +132,10 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): async def store_unified_object_id( self, unified_object_id: str, - file_object: Union[LiteLLMBatch, LiteLLMFineTuningJob], + file_object: Union[LiteLLMBatch, LiteLLMFineTuningJob, "ResponsesAPIResponse"], litellm_parent_otel_span: Optional[Span], model_object_id: str, - file_purpose: Literal["batch", "fine-tune"], + file_purpose: Literal["batch", "fine-tune", "response"], user_api_key_dict: UserAPIKeyAuth, ) -> None: verbose_logger.info( @@ -946,7 +945,9 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): # File is stored in a storage backend, download and convert to base64 try: - from litellm.llms.base_llm.files.storage_backend_factory import get_storage_backend + from litellm.llms.base_llm.files.storage_backend_factory import ( + get_storage_backend, + ) storage_backend_name = db_file.storage_backend storage_url = db_file.storage_url diff --git a/litellm/llms/openai/responses/transformation.py b/litellm/llms/openai/responses/transformation.py index 7ccec074703..325bd80ae4d 100644 --- a/litellm/llms/openai/responses/transformation.py +++ b/litellm/llms/openai/responses/transformation.py @@ -411,7 +411,6 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): ) raw_response_headers = dict(raw_response.headers) processed_headers = process_response_headers(raw_response_headers) - response = ResponsesAPIResponse(**raw_response_json) response._hidden_params["additional_headers"] = processed_headers response._hidden_params["headers"] = raw_response_headers diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index b4588db6252..ef28ecef73a 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -16,7 +16,11 @@ from typing_extensions import Required, TypedDict from litellm._uuid import uuid from litellm.types.integrations.slack_alerting import AlertType -from litellm.types.llms.openai import AllMessageValues, OpenAIFileObject +from litellm.types.llms.openai import ( + AllMessageValues, + OpenAIFileObject, + ResponsesAPIResponse, +) from litellm.types.mcp import ( MCPAuth, MCPAuthType, @@ -3709,8 +3713,8 @@ class LiteLLM_ManagedFileTable(LiteLLMPydanticObjectBase): class LiteLLM_ManagedObjectTable(LiteLLMPydanticObjectBase): unified_object_id: str model_object_id: str - file_purpose: Literal["batch", "fine-tune"] - file_object: Union[LiteLLMBatch, LiteLLMFineTuningJob] + file_purpose: Literal["batch", "fine-tune", "response"] + file_object: Union[LiteLLMBatch, LiteLLMFineTuningJob, ResponsesAPIResponse] class EnterpriseLicenseData(TypedDict, total=False): diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 0a1cfda1502..233a83319f6 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -4522,7 +4522,7 @@ class ProxyStartupEvent: ### MONITOR SPEND LOGS QUEUE (queue-size-based job) ### if general_settings.get("disable_spend_logs", False) is False: from litellm.proxy.utils import _monitor_spend_logs_queue - + # Start background task to monitor spend logs queue size asyncio.create_task( _monitor_spend_logs_queue( @@ -4632,6 +4632,37 @@ class ProxyStartupEvent: ) pass + ### CHECK RESPONSES COST ### + if llm_router is not None: + try: + from litellm_enterprise.proxy.common_utils.check_responses_cost import ( + CheckResponsesCost, + ) + + check_responses_cost_job = CheckResponsesCost( + proxy_logging_obj=proxy_logging_obj, + prisma_client=prisma_client, + llm_router=llm_router, + ) + scheduler.add_job( + check_responses_cost_job.check_responses_cost, + "interval", + seconds=proxy_batch_polling_interval + + random.randint(0, 30), # Add small random offset + # REMOVED jitter parameter - major cause of memory leak + id="check_responses_cost_job", + replace_existing=True, + misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME, + ) + verbose_proxy_logger.info("Responses cost check job scheduled successfully") + + except Exception as e: + verbose_proxy_logger.error(f"Failed to setup responses cost checking: {e}") + verbose_proxy_logger.debug( + "Checking responses cost for LiteLLM Managed Files is an Enterprise Feature. Skipping..." + ) + pass + # MEMORY LEAK FIX: Start scheduler with paused=False to avoid backlog processing # Do NOT reset job times to "now" as this can trigger the memory leak # The misfire_grace_time and coalesce settings will handle any missed runs properly diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 252b3a7d384..623e8408862 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -1,6 +1,6 @@ import asyncio import time -from typing import Any, AsyncIterator, cast +from typing import Any, AsyncIterator, Optional, cast from uuid import uuid4 from fastapi import APIRouter, Depends, HTTPException, Request, Response @@ -155,7 +155,7 @@ async def responses_api( # Normal response flow processor = ProxyBaseLLMRequestProcessing(data=data) try: - return await processor.base_process_llm_request( + response = await processor.base_process_llm_request( request=request, fastapi_response=fastapi_response, user_api_key_dict=user_api_key_dict, @@ -173,6 +173,48 @@ async def responses_api( user_api_base=user_api_base, version=version, ) + + # Store in managed objects table if background mode is enabled + if data.get("background") and isinstance(response, ResponsesAPIResponse): + if response.status in ["queued", "in_progress"]: + from litellm_enterprise.proxy.hooks.managed_files import ( # type: ignore + _PROXY_LiteLLMManagedFiles, + ) + managed_files_obj = cast( + Optional[_PROXY_LiteLLMManagedFiles], + proxy_logging_obj.get_proxy_hook("managed_files"), + ) + + if managed_files_obj and llm_router: + try: + # Get the actual deployment model_id from hidden params + hidden_params = getattr(response, "_hidden_params", {}) or {} + model_id = hidden_params.get("model_id", None) + + if not model_id: + verbose_proxy_logger.warning( + f"No model_id found in response hidden params for response {response.id}, skipping managed object storage" + ) + raise Exception("No model_id found in response hidden params") + # Store in managed objects table + await managed_files_obj.store_unified_object_id( + unified_object_id=response.id, + file_object=response, + litellm_parent_otel_span=None, + model_object_id=response.id, + file_purpose="response", + user_api_key_dict=user_api_key_dict, + ) + + verbose_proxy_logger.info( + f"Stored background response {response.id} in managed objects table with unified_id={response.id}" + ) + except Exception as e: + verbose_proxy_logger.error( + f"Failed to store background response in managed objects table: {str(e)}" + ) + + return response except ModifyResponseException as e: # Guardrail passthrough: return violation message in Responses API format (200) _data = e.request_data diff --git a/tests/proxy_unit_tests/test_check_responses_cost.py b/tests/proxy_unit_tests/test_check_responses_cost.py new file mode 100644 index 00000000000..3bcacdfc05d --- /dev/null +++ b/tests/proxy_unit_tests/test_check_responses_cost.py @@ -0,0 +1,382 @@ +""" +Unit tests for CheckResponsesCost class +""" + +import asyncio +from datetime import datetime +from unittest.mock import AsyncMock, MagicMock, Mock, patch + +import pytest + +from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse + + +class TestCheckResponsesCost: + """Test suite for CheckResponsesCost class""" + + @pytest.fixture + def mock_prisma_client(self): + """Create a mock Prisma client""" + client = MagicMock() + client.db = MagicMock() + client.db.litellm_managedobjecttable = MagicMock() + return client + + @pytest.fixture + def mock_proxy_logging_obj(self): + """Create a mock ProxyLogging object""" + logging_obj = MagicMock() + logging_obj.get_proxy_hook = MagicMock(return_value=None) + return logging_obj + + @pytest.fixture + def mock_llm_router(self): + """Create a mock LLM Router""" + router = MagicMock() + router.aget_responses = AsyncMock() + router.get_deployment = MagicMock() + return router + + @pytest.fixture + def check_responses_cost_instance( + self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router + ): + """Create a CheckResponsesCost instance with mocked dependencies""" + from litellm_enterprise.proxy.common_utils.check_responses_cost import ( + CheckResponsesCost, + ) + + return CheckResponsesCost( + proxy_logging_obj=mock_proxy_logging_obj, + prisma_client=mock_prisma_client, + llm_router=mock_llm_router, + ) + + def test_initialization(self, check_responses_cost_instance): + """Test that CheckResponsesCost initializes correctly""" + assert check_responses_cost_instance.proxy_logging_obj is not None + assert check_responses_cost_instance.prisma_client is not None + assert check_responses_cost_instance.llm_router is not None + + @pytest.mark.asyncio + async def test_check_responses_cost_no_jobs( + self, check_responses_cost_instance, mock_prisma_client + ): + """Test check_responses_cost when there are no jobs to process""" + # Mock empty job list + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[] + ) + + # Should not raise any errors + await check_responses_cost_instance.check_responses_cost() + + # Verify find_many was called with correct parameters + mock_prisma_client.db.litellm_managedobjecttable.find_many.assert_called_once_with( + where={ + "status": {"in": ["queued", "in_progress"]}, + "file_purpose": "response", + } + ) + + @pytest.mark.asyncio + async def test_check_responses_cost_with_completed_response( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + """Test check_responses_cost with a completed response""" + # Mock job with response ID + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_123" + mock_job.created_by = "test-user" + mock_job.id = "job-123" + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + # Mock completed response + mock_response = ResponsesAPIResponse( + id="resp_123", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage( + input_tokens=100, + output_tokens=50, + total_tokens=150, + ), + ) + + # Mock update_many + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + # Run the check with mocked litellm.aget_responses + 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() + + # Verify the job was marked as completed + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_called_once() + call_args = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args + assert call_args[1]["data"]["status"] == "completed" + assert call_args[1]["where"]["id"]["in"] == ["job-123"] + + @pytest.mark.asyncio + async def test_check_responses_cost_with_failed_response( + self, check_responses_cost_instance, mock_prisma_client, mock_llm_router + ): + """Test check_responses_cost with a failed response""" + # Mock job + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_456" + mock_job.created_by = "test-user" + mock_job.id = "job-456" + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + # Mock failed response + mock_response = ResponsesAPIResponse( + id="resp_456", + object="response", + status="failed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + + # Mock update_many + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + # Run the check + 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() + + # Verify the job was marked as completed (even though response failed) + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_called_once() + call_args = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args + assert call_args[1]["data"]["status"] == "completed" + + @pytest.mark.asyncio + async def test_check_responses_cost_with_cancelled_response( + self, check_responses_cost_instance, mock_prisma_client + ): + """Test check_responses_cost with a cancelled response""" + # Mock job + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_789" + mock_job.created_by = "test-user" + mock_job.id = "job-789" + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + # Mock cancelled response + mock_response = ResponsesAPIResponse( + id="resp_789", + object="response", + status="cancelled", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + + # Mock update_many + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + # Run the check + 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() + + # Verify the job was marked as completed + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_called_once() + + @pytest.mark.asyncio + async def test_check_responses_cost_with_in_progress_response( + self, check_responses_cost_instance, mock_prisma_client + ): + """Test check_responses_cost with a response still in progress""" + # Mock job + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_in_progress" + mock_job.created_by = "test-user" + mock_job.id = "job-in-progress" + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + # Mock in-progress response + mock_response = ResponsesAPIResponse( + id="resp_in_progress", + object="response", + status="in_progress", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + + # Mock update_many + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + # Run the check + 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() + + # Verify no updates were made (response still in progress) + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_not_called() + + @pytest.mark.asyncio + async def test_check_responses_cost_with_queued_response( + self, check_responses_cost_instance, mock_prisma_client + ): + """Test check_responses_cost with a queued response""" + # Mock job + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_queued" + mock_job.created_by = "test-user" + mock_job.id = "job-queued" + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + # Mock queued response + mock_response = ResponsesAPIResponse( + id="resp_queued", + object="response", + status="queued", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + + # Mock update_many + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + # Run the check + 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() + + # Verify no updates were made (response still queued) + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_not_called() + + @pytest.mark.asyncio + async def test_check_responses_cost_with_exception( + self, check_responses_cost_instance, mock_prisma_client + ): + """Test check_responses_cost handles exceptions gracefully""" + # Mock job + mock_job = MagicMock() + mock_job.unified_object_id = "resp_test_error" + mock_job.created_by = "test-user" + mock_job.id = "job-error" + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + # Mock update_many + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + # Run the check with mocked exception + with patch( + "litellm.aget_responses", + new_callable=AsyncMock, + side_effect=Exception("Provider error"), + ): + # Should not raise, just skip the job + await check_responses_cost_instance.check_responses_cost() + + # Verify no updates were made (job was skipped due to error) + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_not_called() + + @pytest.mark.asyncio + async def test_check_responses_cost_multiple_jobs( + self, check_responses_cost_instance, mock_prisma_client + ): + """Test check_responses_cost with multiple jobs""" + # Mock multiple jobs + mock_job1 = MagicMock() + mock_job1.unified_object_id = "resp_test_1" + mock_job1.created_by = "user1" + mock_job1.id = "job-1" + + mock_job2 = MagicMock() + mock_job2.unified_object_id = "resp_test_2" + mock_job2.created_by = "user2" + mock_job2.id = "job-2" + + mock_job3 = MagicMock() + mock_job3.unified_object_id = "resp_test_3" + mock_job3.created_by = "user3" + mock_job3.id = "job-3" + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job1, mock_job2, mock_job3] + ) + + # Mock responses - 2 completed, 1 in progress + mock_response1 = ResponsesAPIResponse( + id="resp_1", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage( + input_tokens=100, + output_tokens=50, + total_tokens=150, + ), + ) + + mock_response2 = ResponsesAPIResponse( + id="resp_2", + object="response", + status="in_progress", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + + mock_response3 = ResponsesAPIResponse( + id="resp_3", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage( + input_tokens=200, + output_tokens=100, + total_tokens=300, + ), + ) + + # Mock update_many + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + # Run the check + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.side_effect = [mock_response1, mock_response2, mock_response3] + + await check_responses_cost_instance.check_responses_cost() + + # Verify only the 2 completed jobs were marked as complete + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_called_once() + call_args = mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args + assert len(call_args[1]["where"]["id"]["in"]) == 2 + assert "job-1" in call_args[1]["where"]["id"]["in"] + assert "job-3" in call_args[1]["where"]["id"]["in"] + assert "job-2" not in call_args[1]["where"]["id"]["in"] diff --git a/tests/test_litellm/integrations/test_responses_background_cost.py b/tests/test_litellm/integrations/test_responses_background_cost.py new file mode 100644 index 00000000000..6f1e7e96103 --- /dev/null +++ b/tests/test_litellm/integrations/test_responses_background_cost.py @@ -0,0 +1,513 @@ +""" +Integration tests for responses API background cost tracking +""" + +import asyncio +import os +from datetime import datetime +from unittest.mock import AsyncMock, MagicMock, Mock, patch + +import pytest + +from litellm.types.llms.openai import ResponseAPIUsage, ResponsesAPIResponse + + +class TestResponsesBackgroundCostTracking: + """Integration tests for responses API background cost tracking""" + + @pytest.fixture + def mock_managed_files_obj(self): + """Create a mock managed files object""" + managed_files = MagicMock() + managed_files.store_unified_object_id = AsyncMock() + return managed_files + + @pytest.fixture + def mock_proxy_logging_obj(self, mock_managed_files_obj): + """Create a mock proxy logging object""" + logging_obj = MagicMock() + logging_obj.get_proxy_hook = MagicMock(return_value=mock_managed_files_obj) + return logging_obj + + @pytest.fixture + def mock_llm_router(self): + """Create a mock LLM router""" + router = MagicMock() + return router + + @pytest.mark.asyncio + async def test_store_response_in_managed_objects_table( + self, mock_managed_files_obj, mock_proxy_logging_obj, mock_llm_router + ): + """Test that background responses are stored in managed objects table""" + # Create a mock response with queued status and hidden params + response = ResponsesAPIResponse( + id="resp_bGl0ZWxsbTpjdXN0b21fbGxtX3Byb3ZpZGVyOm9wZW5haTttb2RlbF9pZDpncHQtNDtsbGxfcmVzcG9uc2VfaWQ6cmVzcF8xMjM", + object="response", + status="queued", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + + # Add hidden params with model_id (simulating what base_process_llm_request does) + response._hidden_params = { + "model_id": "model-deployment-id-123" + } + + # Mock request data + data = { + "model": "gpt-4", + "input": "Test input", + "background": True, + } + + # Mock user_api_key_dict + user_api_key_dict = MagicMock() + user_api_key_dict.user_id = "test-user" + + # Simulate the storage logic from endpoints.py + if data.get("background") and isinstance(response, ResponsesAPIResponse): + if response.status in ["queued", "in_progress"]: + # Get model_id from hidden params + hidden_params = getattr(response, "_hidden_params", {}) or {} + model_id = hidden_params.get("model_id", None) + + if model_id: + # Store in managed objects table using response.id directly + await mock_managed_files_obj.store_unified_object_id( + unified_object_id=response.id, + file_object=response, + litellm_parent_otel_span=None, + model_object_id=response.id, + file_purpose="response", + user_api_key_dict=user_api_key_dict, + ) + + # Verify store_unified_object_id was called + mock_managed_files_obj.store_unified_object_id.assert_called_once() + call_args = mock_managed_files_obj.store_unified_object_id.call_args + + # Verify the arguments - unified_object_id should be response.id + assert call_args[1]["unified_object_id"] == response.id + assert call_args[1]["model_object_id"] == response.id + assert call_args[1]["file_purpose"] == "response" + assert call_args[1]["user_api_key_dict"] == user_api_key_dict + + @pytest.mark.asyncio + async def test_no_storage_for_non_background_requests( + self, mock_managed_files_obj, mock_proxy_logging_obj + ): + """Test that non-background requests are not stored""" + # Create a mock response + response = ResponsesAPIResponse( + id="resp_456", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage( + input_tokens=100, + output_tokens=50, + total_tokens=150, + ), + ) + + # Mock request data without background flag + data = { + "model": "gpt-4", + "input": "Test input", + "background": False, + } + + # Simulate the storage logic + if data.get("background") and isinstance(response, ResponsesAPIResponse): + if response.status in ["queued", "in_progress"]: + await mock_managed_files_obj.store_unified_object_id() + + # Verify store_unified_object_id was NOT called + mock_managed_files_obj.store_unified_object_id.assert_not_called() + + @pytest.mark.asyncio + async def test_no_storage_for_completed_responses( + self, mock_managed_files_obj, mock_proxy_logging_obj + ): + """Test that completed responses are not stored""" + # Create a mock response with completed status + response = ResponsesAPIResponse( + id="resp_789", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage( + input_tokens=100, + output_tokens=50, + total_tokens=150, + ), + ) + + # Mock request data with background flag + data = { + "model": "gpt-4", + "input": "Test input", + "background": True, + } + + # Simulate the storage logic + if data.get("background") and isinstance(response, ResponsesAPIResponse): + if response.status in ["queued", "in_progress"]: + await mock_managed_files_obj.store_unified_object_id() + + # Verify store_unified_object_id was NOT called (status is completed) + mock_managed_files_obj.store_unified_object_id.assert_not_called() + + @pytest.mark.asyncio + async def test_no_storage_without_model_id( + self, mock_managed_files_obj, mock_proxy_logging_obj + ): + """Test that responses without model_id in hidden params are not stored""" + # Create a mock response without hidden params + response = ResponsesAPIResponse( + id="resp_no_model", + object="response", + status="queued", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + + # Mock request data with background flag + data = { + "model": "gpt-4", + "input": "Test input", + "background": True, + } + + user_api_key_dict = MagicMock() + + # Simulate the storage logic + if data.get("background") and isinstance(response, ResponsesAPIResponse): + if response.status in ["queued", "in_progress"]: + hidden_params = getattr(response, "_hidden_params", {}) or {} + model_id = hidden_params.get("model_id", None) + + if model_id: # This will be False + await mock_managed_files_obj.store_unified_object_id( + unified_object_id=response.id, + file_object=response, + litellm_parent_otel_span=None, + model_object_id=response.id, + file_purpose="response", + user_api_key_dict=user_api_key_dict, + ) + + # Verify store_unified_object_id was NOT called (no model_id) + mock_managed_files_obj.store_unified_object_id.assert_not_called() + + @pytest.mark.asyncio + async def test_error_handling_in_storage( + self, mock_managed_files_obj, mock_proxy_logging_obj + ): + """Test that errors during storage are handled gracefully""" + # Mock store_unified_object_id to raise an exception + mock_managed_files_obj.store_unified_object_id = AsyncMock( + side_effect=Exception("Database error") + ) + + response = ResponsesAPIResponse( + id="resp_error", + object="response", + status="queued", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + response._hidden_params = {"model_id": "test-model-id"} + + data = { + "model": "gpt-4", + "input": "Test input", + "background": True, + } + + user_api_key_dict = MagicMock() + user_api_key_dict.user_id = "test-user" + + # Try to store - should not raise (error is caught in endpoints.py) + try: + if data.get("background") and isinstance(response, ResponsesAPIResponse): + if response.status in ["queued", "in_progress"]: + hidden_params = getattr(response, "_hidden_params", {}) or {} + model_id = hidden_params.get("model_id", None) + + if model_id: + await mock_managed_files_obj.store_unified_object_id( + unified_object_id=response.id, + file_object=response, + litellm_parent_otel_span=None, + model_object_id=response.id, + file_purpose="response", + user_api_key_dict=user_api_key_dict, + ) + except Exception: + # Exception should be caught and logged, not raised + pass + + # Verify the method was called (even though it raised) + assert mock_managed_files_obj.store_unified_object_id.called + + +class TestCheckResponsesCost: + """Tests for the CheckResponsesCost polling class""" + + @pytest.fixture + def mock_prisma_client(self): + """Create a mock Prisma client""" + client = MagicMock() + client.db = MagicMock() + client.db.litellm_managedobjecttable = MagicMock() + return client + + @pytest.fixture + def mock_proxy_logging_obj(self): + """Create a mock proxy logging object""" + return MagicMock() + + @pytest.fixture + def mock_llm_router(self): + """Create a mock LLM router""" + return MagicMock() + + @pytest.mark.asyncio + async def test_check_responses_cost_initialization( + self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router + ): + """Test CheckResponsesCost initialization""" + from litellm_enterprise.proxy.common_utils.check_responses_cost import ( + CheckResponsesCost, + ) + + checker = CheckResponsesCost( + proxy_logging_obj=mock_proxy_logging_obj, + prisma_client=mock_prisma_client, + llm_router=mock_llm_router, + ) + + assert checker.proxy_logging_obj == mock_proxy_logging_obj + assert checker.prisma_client == mock_prisma_client + assert checker.llm_router == mock_llm_router + + @pytest.mark.asyncio + async def test_check_responses_cost_no_jobs( + self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router + ): + """Test polling when there are no jobs""" + from litellm_enterprise.proxy.common_utils.check_responses_cost import ( + CheckResponsesCost, + ) + + # Mock find_many to return empty list + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[] + ) + + checker = CheckResponsesCost( + proxy_logging_obj=mock_proxy_logging_obj, + prisma_client=mock_prisma_client, + llm_router=mock_llm_router, + ) + + # Should not raise any errors + await checker.check_responses_cost() + + # Verify find_many was called with correct parameters + mock_prisma_client.db.litellm_managedobjecttable.find_many.assert_called_once_with( + where={ + "status": {"in": ["queued", "in_progress"]}, + "file_purpose": "response", + } + ) + + @pytest.mark.asyncio + async def test_check_responses_cost_with_completed_job( + self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router + ): + """Test polling with a completed job""" + from litellm_enterprise.proxy.common_utils.check_responses_cost import ( + CheckResponsesCost, + ) + + # Create a mock job + mock_job = MagicMock() + mock_job.id = "job-123" + mock_job.unified_object_id = "resp_test_id" + mock_job.created_by = "test-user" + + # Mock find_many to return the job + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + + # Mock update_many + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + # Create a completed response + completed_response = ResponsesAPIResponse( + id="resp_test_id", + object="response", + status="completed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=ResponseAPIUsage( + input_tokens=100, + output_tokens=50, + total_tokens=150, + ), + ) + + checker = CheckResponsesCost( + proxy_logging_obj=mock_proxy_logging_obj, + prisma_client=mock_prisma_client, + llm_router=mock_llm_router, + ) + + # Mock litellm.aget_responses to return completed response + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.return_value = completed_response + + await checker.check_responses_cost() + + # Verify update_many was called to mark job as completed + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_called_once() + call_args = ( + mock_prisma_client.db.litellm_managedobjecttable.update_many.call_args + ) + assert call_args[1]["where"]["id"]["in"] == ["job-123"] + assert call_args[1]["data"]["status"] == "completed" + + @pytest.mark.asyncio + async def test_check_responses_cost_with_failed_job( + self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router + ): + """Test polling with a failed job""" + from litellm_enterprise.proxy.common_utils.check_responses_cost import ( + CheckResponsesCost, + ) + + # Create a mock job + mock_job = MagicMock() + mock_job.id = "job-456" + mock_job.unified_object_id = "resp_failed" + mock_job.created_by = "test-user" + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + # Create a failed response + failed_response = ResponsesAPIResponse( + id="resp_failed", + object="response", + status="failed", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + + checker = CheckResponsesCost( + proxy_logging_obj=mock_proxy_logging_obj, + prisma_client=mock_prisma_client, + llm_router=mock_llm_router, + ) + + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.return_value = failed_response + + await checker.check_responses_cost() + + # Verify job was marked as completed even though it failed + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_called_once() + + @pytest.mark.asyncio + async def test_check_responses_cost_with_in_progress_job( + self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router + ): + """Test polling with a job still in progress""" + from litellm_enterprise.proxy.common_utils.check_responses_cost import ( + CheckResponsesCost, + ) + + # Create a mock job + mock_job = MagicMock() + mock_job.id = "job-789" + mock_job.unified_object_id = "resp_in_progress" + mock_job.created_by = "test-user" + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + # Create an in-progress response + in_progress_response = ResponsesAPIResponse( + id="resp_in_progress", + object="response", + status="in_progress", + created_at=int(datetime.now().timestamp()), + output=[], + usage=None, + ) + + checker = CheckResponsesCost( + proxy_logging_obj=mock_proxy_logging_obj, + prisma_client=mock_prisma_client, + llm_router=mock_llm_router, + ) + + with patch("litellm.aget_responses", new_callable=AsyncMock) as mock_aget: + mock_aget.return_value = in_progress_response + + await checker.check_responses_cost() + + # Verify update_many was NOT called (job still in progress) + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_not_called() + + @pytest.mark.asyncio + async def test_check_responses_cost_error_handling( + self, mock_proxy_logging_obj, mock_prisma_client, mock_llm_router + ): + """Test that errors when querying responses are handled gracefully""" + from litellm_enterprise.proxy.common_utils.check_responses_cost import ( + CheckResponsesCost, + ) + + # Create a mock job + mock_job = MagicMock() + mock_job.id = "job-error" + mock_job.unified_object_id = "resp_error" + mock_job.created_by = "test-user" + + mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=[mock_job] + ) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock() + + checker = CheckResponsesCost( + proxy_logging_obj=mock_proxy_logging_obj, + prisma_client=mock_prisma_client, + llm_router=mock_llm_router, + ) + + # Mock litellm.aget_responses to raise an exception + with patch( + "litellm.aget_responses", + new_callable=AsyncMock, + side_effect=Exception("API error"), + ): + # Should not raise - errors are caught and logged + await checker.check_responses_cost() + + # Verify update_many was NOT called (error occurred) + mock_prisma_client.db.litellm_managedobjecttable.update_many.assert_not_called()