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:
mateo-berri 2026-08-05 18:26:21 -07:00
parent 60d9e6012c
commit eef908d4ad
3 changed files with 80 additions and 8 deletions

View file

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

View file

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

View file

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