diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index 37b51e1d3af..290045a5a87 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -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) diff --git a/tests/openai_endpoints_tests/test_openai_batches_endpoint.py b/tests/openai_endpoints_tests/test_openai_batches_endpoint.py index c6f4128f2c5..db8f75cf640 100644 --- a/tests/openai_endpoints_tests/test_openai_batches_endpoint.py +++ b/tests/openai_endpoints_tests/test_openai_batches_endpoint.py @@ -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 diff --git a/tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py b/tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py index d8669960674..74139fa9238 100644 --- a/tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py +++ b/tests/test_litellm/enterprise/proxy/test_batch_update_db_managed_output_file_id.py @@ -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."""