mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix: restore model_id for terminal batch retrieval
Co-authored-by: Ishaan Jaff <ishaan-jaff@users.noreply.github.com>
This commit is contained in:
parent
c89496f378
commit
c144ee56df
2 changed files with 97 additions and 0 deletions
|
|
@ -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 [
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue