From c144ee56dfae620bb135d779050afc63ef1fe5bb Mon Sep 17 00:00:00 2001 From: Cursor Agent Date: Sun, 22 Mar 2026 17:10:16 +0000 Subject: [PATCH] fix: restore model_id for terminal batch retrieval Co-authored-by: Ishaan Jaff --- litellm/proxy/batches_endpoints/endpoints.py | 15 ++++ .../test_batch_retrieve_input_file_id.py | 82 +++++++++++++++++++ 2 files changed, 97 insertions(+) diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 160e9c23f01..55317f854df 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -406,6 +406,21 @@ async def retrieve_batch( # noqa: PLR0915 verbose_proxy_logger=verbose_proxy_logger, ) + # DB-restored terminal batches can miss hidden_params.model_id. + # Restore it from the unified batch ID so post hooks can re-unify + # output/error file IDs before responding. + if response is not None and unified_batch_id: + current_hidden_params = getattr(response, "_hidden_params", {}) or {} + if not isinstance(current_hidden_params, dict): + current_hidden_params = {} + if not current_hidden_params.get("model_id"): + model_id_from_batch = get_model_id_from_unified_batch_id(unified_batch_id) + if model_id_from_batch: + response._hidden_params = { + **current_hidden_params, + "model_id": model_id_from_batch, + } + # If batch is in a terminal state, return immediately. # Include "complete" (DB-normalized form of "completed"). if response is not None and response.status in [ diff --git a/tests/test_litellm/enterprise/proxy/test_batch_retrieve_input_file_id.py b/tests/test_litellm/enterprise/proxy/test_batch_retrieve_input_file_id.py index 6e9c3c0354b..438e71e073c 100644 --- a/tests/test_litellm/enterprise/proxy/test_batch_retrieve_input_file_id.py +++ b/tests/test_litellm/enterprise/proxy/test_batch_retrieve_input_file_id.py @@ -10,6 +10,7 @@ import base64 import json import pytest +from fastapi import Response from unittest.mock import AsyncMock, MagicMock, patch from litellm.proxy.openai_files_endpoints.common_utils import ( @@ -73,3 +74,84 @@ async def test_should_resolve_raw_input_file_id_to_unified(): f"input_file_id should be unified '{B64_UNIFIED_INPUT_FILE_ID}', " f"got raw '{response.input_file_id}'" ) + + +@pytest.mark.asyncio +async def test_should_restore_model_id_for_terminal_db_batch_before_post_hook(): + """ + When retrieve_batch returns a terminal batch directly from DB, it should + restore hidden_params.model_id from the unified batch ID before running + post_call_success_hook. + """ + from litellm.proxy.batches_endpoints.endpoints import retrieve_batch + from litellm.types.utils import LiteLLMBatch + + terminal_batch = LiteLLMBatch( + id=B64_UNIFIED_BATCH_ID, + completion_window="24h", + created_at=1700000000, + endpoint="/v1/chat/completions", + input_file_id=B64_UNIFIED_INPUT_FILE_ID, + object="batch", + status="completed", + output_file_id="file-raw-provider-output-123", + ) + terminal_batch._hidden_params = {} + + async def _post_call_success_hook(data, user_api_key_dict, response): + assert response._hidden_params.get("model_id") == "model-xyz" + response.output_file_id = "b64-unified-output-file-id" + return response + + mock_proxy_logging_obj = MagicMock( + post_call_success_hook=AsyncMock(side_effect=_post_call_success_hook), + update_request_status=AsyncMock(), + ) + mock_user_api_key_dict = MagicMock() + mock_user_api_key_dict.parent_otel_span = None + mock_user_api_key_dict.allowed_model_region = "" + + mock_request = MagicMock() + mock_request.headers = {} + mock_request.query_params = {} + mock_request.url = MagicMock() + mock_request.url.port = 4000 + mock_request.method = "GET" + mock_request.url.path = f"/v1/batches/{B64_UNIFIED_BATCH_ID}" + + with ( + patch( + "litellm.proxy.batches_endpoints.endpoints.ProxyBaseLLMRequestProcessing.common_processing_pre_call_logic", + new=AsyncMock(return_value=({"batch_id": B64_UNIFIED_BATCH_ID}, MagicMock())), + ), + patch( + "litellm.proxy.batches_endpoints.endpoints.ProxyBaseLLMRequestProcessing.get_custom_headers", + return_value={}, + ), + patch( + "litellm.proxy.batches_endpoints.endpoints.get_batch_from_database", + new=AsyncMock(return_value=(MagicMock(), terminal_batch)), + ), + patch( + "litellm.proxy.batches_endpoints.endpoints.resolve_input_file_id_to_unified", + new=AsyncMock(), + ), + patch( + "litellm.proxy.batches_endpoints.endpoints.resolve_output_file_ids_to_unified", + new=AsyncMock(), + ), + patch("litellm.proxy.proxy_server.general_settings", {}), + patch("litellm.proxy.proxy_server.llm_router", MagicMock()), + patch("litellm.proxy.proxy_server.proxy_config", MagicMock()), + patch("litellm.proxy.proxy_server.version", "1.0.0"), + patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj), + ): + response = await retrieve_batch( + request=mock_request, + fastapi_response=Response(), + user_api_key_dict=mock_user_api_key_dict, + provider=None, + batch_id=B64_UNIFIED_BATCH_ID, + ) + + assert response.output_file_id == "b64-unified-output-file-id"