mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(batches): register managed output files on batch cancel
update_batch_in_database now fetches the batch row by unified_object_id when the caller omits db_batch_object, so the cancel endpoint attributes newly registered output and error files to the batch owner and returns unified managed ids instead of raw provider ids. Idempotent cancels that do not change the stored status also skip the redundant DB write now. Repair two pre-existing mock tests in test_openai_batches_endpoint.py that asserted values inside lazy percent-format log strings, and give the cancel test's prisma mock an awaitable find_first.
This commit is contained in:
parent
60d9e6012c
commit
eef908d4ad
3 changed files with 80 additions and 8 deletions
|
|
@ -1141,7 +1141,7 @@ async def update_batch_in_database(
|
|||
managed_files_obj: The managed_files proxy hook object
|
||||
prisma_client: Prisma database client
|
||||
verbose_proxy_logger: Logger instance
|
||||
db_batch_object: Optional existing database object (for comparison)
|
||||
db_batch_object: Optional existing database object; fetched by unified_object_id when omitted
|
||||
operation: Description of operation ("update", "cancel", etc.)
|
||||
user_api_key_dict: Optional auth context for creating managed file IDs
|
||||
"""
|
||||
|
|
@ -1154,6 +1154,12 @@ async def update_batch_in_database(
|
|||
if not prisma_client:
|
||||
return
|
||||
|
||||
effective_db_batch_object: Final = (
|
||||
db_batch_object
|
||||
if db_batch_object is not None
|
||||
else await ManagedObjectRepository(prisma_client).table.find_first(where={"unified_object_id": batch_id})
|
||||
)
|
||||
|
||||
# Always normalize the response's file IDs to unified managed IDs
|
||||
# (mutates in place) so the caller returns unified IDs to the user
|
||||
# even when we skip the DB update below for an unchanged status.
|
||||
|
|
@ -1163,16 +1169,17 @@ async def update_batch_in_database(
|
|||
prisma_client=prisma_client,
|
||||
verbose_proxy_logger=verbose_proxy_logger,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
db_batch_object=db_batch_object,
|
||||
db_batch_object=effective_db_batch_object,
|
||||
unified_batch_id=unified_batch_id,
|
||||
)
|
||||
|
||||
# Only update if status has changed (when db_batch_object is provided)
|
||||
if db_batch_object and response.status == db_batch_object.status:
|
||||
if effective_db_batch_object and response.status == effective_db_batch_object.status:
|
||||
return
|
||||
|
||||
if db_batch_object:
|
||||
if effective_db_batch_object:
|
||||
verbose_proxy_logger.info(
|
||||
"Updating batch %s status from %s to %s", batch_id, db_batch_object.status, response.status
|
||||
"Updating batch %s status from %s to %s", batch_id, effective_db_batch_object.status, response.status
|
||||
)
|
||||
else:
|
||||
verbose_proxy_logger.info("Updating batch %s status to %s after %s", batch_id, response.status, operation)
|
||||
|
|
|
|||
|
|
@ -414,7 +414,7 @@ async def test_batch_status_sync_from_provider_to_database():
|
|||
|
||||
# Verify logger was called with status change message
|
||||
mock_logger.info.assert_called()
|
||||
log_message = mock_logger.info.call_args[0][0]
|
||||
log_message = mock_logger.info.call_args[0][0] % mock_logger.info.call_args[0][1:]
|
||||
assert "validating" in log_message
|
||||
assert "completed" in log_message
|
||||
|
||||
|
|
@ -450,6 +450,9 @@ async def test_batch_cancel_updates_database():
|
|||
|
||||
# Mock prisma client
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock()
|
||||
|
||||
# Mock managed_files_obj
|
||||
|
|
@ -482,7 +485,7 @@ async def test_batch_cancel_updates_database():
|
|||
|
||||
# Verify logger was called
|
||||
mock_logger.info.assert_called()
|
||||
log_message = mock_logger.info.call_args[0][0]
|
||||
log_message = mock_logger.info.call_args[0][0] % mock_logger.info.call_args[0][1:]
|
||||
assert "cancel" in log_message.lower()
|
||||
assert "cancelled" in log_message
|
||||
|
||||
|
|
|
|||
|
|
@ -45,9 +45,10 @@ def _build_managed_files_mock(unified_id: str = "file-bWFuYWdlZF9vdXRwdXRfaWQ=")
|
|||
return mock
|
||||
|
||||
|
||||
def _build_prisma_mock():
|
||||
def _build_prisma_mock(db_batch_object=None):
|
||||
mock = MagicMock()
|
||||
mock.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None)
|
||||
mock.db.litellm_managedobjecttable.find_first = AsyncMock(return_value=db_batch_object)
|
||||
mock.db.litellm_managedobjecttable.update = AsyncMock()
|
||||
return mock
|
||||
|
||||
|
|
@ -89,6 +90,67 @@ async def test_update_batch_in_database_stores_unified_output_file_id():
|
|||
assert stored["output_file_id"] != raw_output_file_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cancel_path_registers_output_file_under_batch_owner():
|
||||
unified_id = "file-bWFuYWdlZF9vdXRwdXRfaWQ="
|
||||
db_batch_object = SimpleNamespace(
|
||||
created_by="batch-owner", team_id="batch-team", status="in_progress"
|
||||
)
|
||||
response = _build_batch_response(
|
||||
status="cancelling",
|
||||
output_file_id="file-raw-output",
|
||||
hidden_params={"model_id": "my-model", "model_name": "openai/gpt-4o"},
|
||||
)
|
||||
mock_managed_files = _build_managed_files_mock(unified_id=unified_id)
|
||||
mock_prisma = _build_prisma_mock(db_batch_object=db_batch_object)
|
||||
|
||||
await update_batch_in_database(
|
||||
batch_id="batch_managed_ids_test",
|
||||
unified_batch_id="litellm_proxy;model_id:my-model;llm_batch_id:batch_managed_ids_test",
|
||||
response=response,
|
||||
managed_files_obj=mock_managed_files,
|
||||
prisma_client=mock_prisma,
|
||||
verbose_proxy_logger=MagicMock(),
|
||||
operation="cancel",
|
||||
)
|
||||
|
||||
forwarded_auth = mock_managed_files.store_unified_file_id.call_args.kwargs[
|
||||
"user_api_key_dict"
|
||||
]
|
||||
assert forwarded_auth.user_id == "batch-owner"
|
||||
assert forwarded_auth.team_id == "batch-team"
|
||||
stored = json.loads(
|
||||
mock_prisma.db.litellm_managedobjecttable.update.call_args.kwargs["data"][
|
||||
"file_object"
|
||||
]
|
||||
)
|
||||
assert stored["output_file_id"] == unified_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_batch_derives_model_id_from_unified_batch_id():
|
||||
unified_id = "file-bWFuYWdlZF9vdXRwdXRfaWQ="
|
||||
response = _build_batch_response(output_file_id="file-raw-output", hidden_params={})
|
||||
mock_managed_files = _build_managed_files_mock(unified_id=unified_id)
|
||||
mock_prisma = _build_prisma_mock()
|
||||
|
||||
await update_batch_in_database(
|
||||
batch_id="batch_managed_ids_test",
|
||||
unified_batch_id="litellm_proxy;model_id:model-from-batch-id;llm_batch_id:batch_managed_ids_test",
|
||||
response=response,
|
||||
managed_files_obj=mock_managed_files,
|
||||
prisma_client=mock_prisma,
|
||||
verbose_proxy_logger=MagicMock(),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_id="user-abc"),
|
||||
)
|
||||
|
||||
assert (
|
||||
mock_managed_files.get_unified_output_file_id.call_args.kwargs["model_id"]
|
||||
== "model-from-batch-id"
|
||||
)
|
||||
assert response.output_file_id == unified_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_ensure_batch_response_normalizes_error_file_id():
|
||||
"""Both output_file_id and error_file_id must be normalized to managed IDs."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue