diff --git a/litellm/utils.py b/litellm/utils.py index 5970ea8876a..85b20bdd602 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -1658,6 +1658,7 @@ def client(original_function): start_time=start_time, end_time=end_time, ) + completion.release() return result except Exception as e: call_type = original_function.__name__ @@ -1941,12 +1942,14 @@ def client(original_function): and _caching_handler_response is not None and _caching_handler_response.final_embedding_cached_response is not None ): - return _llm_caching_handler._combine_cached_embedding_response_with_api_result( + combined_response: Final = _llm_caching_handler._combine_cached_embedding_response_with_api_result( _caching_handler_response=_caching_handler_response, embedding_response=result, start_time=start_time, end_time=end_time, ) + completion.release() + return combined_response _update_response_metadata( result=result, @@ -1957,6 +1960,7 @@ def client(original_function): end_time=end_time, ) + completion.release() return result except Exception as e: traceback_exception: Final = traceback.format_exc() diff --git a/tests/test_litellm/litellm_core_utils/test_call_completion.py b/tests/test_litellm/litellm_core_utils/test_call_completion.py index 36836f08630..140e482a3c1 100644 --- a/tests/test_litellm/litellm_core_utils/test_call_completion.py +++ b/tests/test_litellm/litellm_core_utils/test_call_completion.py @@ -235,3 +235,33 @@ async def test_async_ocr_wrapper_sends_final_failure_to_attached_completion( assert caught.value is mapped_error assert native_completion.failures == [mapped_error, mapped_error] + + +@pytest.mark.asyncio +async def test_async_ocr_wrapper_retains_completion_until_metadata_finishes( + monkeypatch: pytest.MonkeyPatch, +) -> None: + native_completion: Final = RecordingCompletion() + response: Final = object() + metadata_error: Final = ValueError("metadata failure") + + async def aocr(**kwargs: object) -> object: + completion = kwargs.get("_litellm_call_completion") + assert isinstance(completion, CallCompletion) + assert completion.attach(native_completion) + return response + + def fail_metadata(**kwargs: object) -> None: + raise metadata_error + + monkeypatch.setattr("litellm.utils.function_setup", MagicMock(return_value=(MagicMock(), {}))) + monkeypatch.setattr("litellm.utils.load_credentials_from_list", MagicMock()) + monkeypatch.setattr("litellm.utils.update_response_metadata", fail_metadata) + wrapped: Final = client(aocr) + + with pytest.raises(ValueError, match="metadata failure") as caught: + await wrapped() + + assert caught.value is metadata_error + assert native_completion.successes == [response] + assert native_completion.failures == [metadata_error, metadata_error]