fix: restore model_id for terminal batch retrieval

Co-authored-by: Ishaan Jaff <ishaan-jaff@users.noreply.github.com>
This commit is contained in:
Cursor Agent 2026-03-22 17:10:16 +00:00
parent c89496f378
commit c144ee56df
No known key found for this signature in database
2 changed files with 97 additions and 0 deletions

View file

@ -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 [

View file

@ -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"