From 5ac9eec4659275ea6b19f2114657a4f22c12b574 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Fri, 11 Sep 2026 08:06:37 -0700 Subject: [PATCH] refactor(ocr): extract call completion boundary --- litellm/litellm_core_utils/call_completion.py | 229 ++++++++++++++++++ litellm/ocr/main.py | 4 + litellm/rust_bridge/ocr.py | 2 + litellm/utils.py | 215 +++++----------- .../test_call_completion.py | 169 +++++++++++++ 5 files changed, 464 insertions(+), 155 deletions(-) create mode 100644 litellm/litellm_core_utils/call_completion.py create mode 100644 tests/test_litellm/litellm_core_utils/test_call_completion.py diff --git a/litellm/litellm_core_utils/call_completion.py b/litellm/litellm_core_utils/call_completion.py new file mode 100644 index 00000000000..40fb3f77f00 --- /dev/null +++ b/litellm/litellm_core_utils/call_completion.py @@ -0,0 +1,229 @@ +from __future__ import annotations + +import asyncio +import contextvars +import datetime +from collections.abc import Awaitable, Callable, Coroutine +from concurrent.futures import Future +from typing import Protocol + + +class CompletionLogging(Protocol): + def success_handler( + self, + result: object, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: ... + + def failure_handler( + self, + exception: Exception, + traceback_exception: str, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: ... + + def async_success_handler( + self, + result: object, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> Coroutine[object, object, None]: ... + + def async_failure_handler( + self, + exception: Exception, + traceback_exception: str, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> Awaitable[None]: ... + + def handle_sync_success_callbacks_for_async_calls( + self, + result: object, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: ... + + +class CompletionExecutor(Protocol): + def submit( + self, + function: Callable[..., object], + *args: object, + ) -> Future[object]: ... + + +class Completion(Protocol): + def success( + self, + result: object, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: ... + + def failure( + self, + exception: Exception, + traceback_exception: str, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: ... + + async def async_failure( + self, + exception: Exception, + traceback_exception: str, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: ... + + +class PythonCompletion: + def __init__( + self, + logging_obj: CompletionLogging, + executor: CompletionExecutor | None, + *, + async_call: bool, + internal_call: bool, + completion_with_fallbacks: bool, + ) -> None: + self._logging_obj = logging_obj + self._executor = executor + self._async_call = async_call + self._internal_call = internal_call + self._completion_with_fallbacks = completion_with_fallbacks + + def success( + self, + result: object, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: + if not self._async_call: + assert self._executor is not None + context = contextvars.copy_context() + self._executor.submit( + context.run, + self._logging_obj.success_handler, + result, + start_time, + end_time, + ) + return + + if not self._internal_call: + if getattr(self._logging_obj, "_defer_async_logging", False): + + def enqueue_deferred_logging() -> None: + asyncio.create_task(self._dispatch_async_success(result, start_time, end_time)) + + setattr( # noqa: B010 # optional legacy logger field is absent from narrow test doubles + self._logging_obj, + "_enqueue_deferred_logging", + enqueue_deferred_logging, + ) + else: + asyncio.create_task(self._dispatch_async_success(result, start_time, end_time)) + + self._logging_obj.handle_sync_success_callbacks_for_async_calls( + result=result, + start_time=start_time, + end_time=end_time, + ) + + def failure( + self, + exception: Exception, + traceback_exception: str, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: + if self._async_call and self._internal_call: + return + self._logging_obj.failure_handler(exception, traceback_exception, start_time, end_time) + + async def async_failure( + self, + exception: Exception, + traceback_exception: str, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: + if not self._async_call or self._internal_call: + return + await self._logging_obj.async_failure_handler(exception, traceback_exception, start_time, end_time) + + async def _dispatch_async_success( + self, + result: object, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: + if self._completion_with_fallbacks: + return + from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER + + GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue( # pyright: ignore[reportUnknownMemberType] # legacy worker lacks generic coroutine annotations + async_coroutine=self._logging_obj.async_success_handler( + result=result, + start_time=start_time, + end_time=end_time, + ) + ) + self._logging_obj.handle_sync_success_callbacks_for_async_calls( + result=result, + start_time=start_time, + end_time=end_time, + ) + + +class CallCompletion: + def __init__(self, implementation: Completion) -> None: + self._python_implementation = implementation + self._implementation = implementation + self._attached = False + + @property + def python_implementation(self) -> Completion: + return self._python_implementation + + def attach(self, implementation: Completion) -> bool: + if self._attached: + return False + self._implementation = implementation + self._attached = True + return True + + def success( + self, + result: object, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: + self._implementation.success(result, start_time, end_time) + + def failure( + self, + exception: Exception, + traceback_exception: str, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: + self._implementation.failure(exception, traceback_exception, start_time, end_time) + + async def async_failure( + self, + exception: Exception, + traceback_exception: str, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: + await self._implementation.async_failure( + exception, + traceback_exception, + start_time, + end_time, + ) diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 56bfd98895d..f4f968c8d34 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -522,6 +522,7 @@ async def aocr( ) ``` """ + call_completion: Final = kwargs.pop("_litellm_call_completion", None) completion_kwargs: Final[dict[str, object]] = { "model": model, "document": document, @@ -541,6 +542,7 @@ async def aocr( custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, kwargs=kwargs, + call_completion=call_completion, ) try: if rust_enabled() and _rust_ocr_supported(request): @@ -804,6 +806,7 @@ def ocr( print(f"Page {page.index}: {page.markdown}") ``` """ + call_completion: Final = kwargs.pop("_litellm_call_completion", None) completion_kwargs: Final[dict[str, object]] = { "model": model, "document": document, @@ -823,6 +826,7 @@ def ocr( custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, kwargs=kwargs, + call_completion=call_completion, ) try: _is_async: Final = kwargs.pop("aocr", False) is True diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index 89eab71ccba..fbc3756defc 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -52,6 +52,7 @@ class LiteLLMOcrRequest: custom_llm_provider: str | None extra_headers: dict[str, object] | None kwargs: Mapping[str, object] + call_completion: object = None input_sources: Mapping[str, str] | None = None @@ -294,6 +295,7 @@ def _marshal( custom_llm_provider=request.custom_llm_provider, extra_headers=request.extra_headers, kwargs=optional_params, + call_completion=request.call_completion, input_sources=input_sources, ) diff --git a/litellm/utils.py b/litellm/utils.py index bb1bce66d9b..5970ea8876a 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -7,7 +7,6 @@ import ast import asyncio import base64 import binascii -import contextvars import copy import datetime import hashlib @@ -66,7 +65,6 @@ from litellm.constants import ( DEFAULT_EMBEDDING_PARAM_VALUES, DEFAULT_MAX_LRU_CACHE_SIZE, DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT, - DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT, DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET, DEFAULT_TRIM_RATIO, FUNCTION_DEFINITION_TOKEN_COUNT, @@ -81,6 +79,7 @@ from litellm.constants import ( PROVIDERS_THAT_AUTHENTICATE_ON_PROVIDER_INFO, TOOL_CHOICE_OBJECT_TOKEN_COUNT, ) +from litellm.litellm_core_utils.call_completion import CallCompletion, PythonCompletion from litellm.litellm_core_utils.core_helpers import normalize_drop_params from litellm.litellm_core_utils.fallback_generalizations import ( match_capability_generalizations, @@ -279,7 +278,7 @@ except (ImportError, AttributeError, TypeError): # Convert to str (if necessary) claude_json_str = json.dumps(json_data) import importlib.metadata -from collections.abc import AsyncIterator, Callable, Iterable, Iterator, Mapping, Sequence +from collections.abc import Callable, Iterable, Mapping, Sequence from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast, get_args from litellm import utils as litellm_utils @@ -1196,79 +1195,6 @@ def function_setup( raise e -def _dispatch_success_logging( - logging_obj: LiteLLMLoggingObject, - result: object, - start_time: datetime.datetime, - end_time: datetime.datetime, - is_completion_with_fallbacks: bool, - is_litellm_internal_call: bool, -) -> None: - if not is_litellm_internal_call: - if getattr(logging_obj, "_defer_async_logging", False): - - def _enqueue_deferred_logging() -> None: - asyncio.create_task( - _client_async_logging_helper( - logging_obj=logging_obj, - result=result, - start_time=start_time, - end_time=end_time, - is_completion_with_fallbacks=is_completion_with_fallbacks, - ) - ) - - logging_obj._enqueue_deferred_logging = _enqueue_deferred_logging - else: - asyncio.create_task( - _client_async_logging_helper( - logging_obj=logging_obj, - result=result, - start_time=start_time, - end_time=end_time, - is_completion_with_fallbacks=is_completion_with_fallbacks, - ) - ) - - logging_obj.handle_sync_success_callbacks_for_async_calls( - result=result, - start_time=start_time, - end_time=end_time, - ) - - -async def _client_async_logging_helper( - logging_obj: LiteLLMLoggingObject, - result, - start_time, - end_time, - is_completion_with_fallbacks: bool, -): - if ( - is_completion_with_fallbacks is False - ): # don't log the parent event litellm.completion_with_fallbacks as a 'log_success_event', this will lead to double logging the same call - https://github.com/BerriAI/litellm/issues/7477 - print_verbose( - f"Async Wrapper: Completed Call, calling async_success_handler: {logging_obj.async_success_handler}" - ) - ################################################ - # Async Logging Worker - ################################################ - from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER - - GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue( - async_coroutine=logging_obj.async_success_handler(result=result, start_time=start_time, end_time=end_time) - ) - - ################################################ - # Sync Logging Worker - ################################################ - logging_obj.handle_sync_success_callbacks_for_async_calls( - result=result, - start_time=start_time, - end_time=end_time, - ) - - def _get_wrapper_num_retries(kwargs: dict[str, Any], exception: Exception) -> tuple[int | None, dict[str, Any]]: """ Get the number of retries from the kwargs and the retry policy. @@ -1537,6 +1463,7 @@ def client(original_function): start_time: Final = datetime.datetime.now() result = None logging_obj: LiteLLMLoggingObject | None = kwargs.get("litellm_logging_obj", None) + completion: CallCompletion | None = None # only set litellm_call_id if its not in kwargs if "litellm_call_id" not in kwargs: @@ -1552,7 +1479,17 @@ def client(original_function): # Type assertion: logging_obj is guaranteed to be non-None after function_setup assert logging_obj is not None, "logging_obj should not be None after function_setup" + from litellm.litellm_core_utils.thread_pool_executor import executor as logging_executor + completion = CallCompletion( + PythonCompletion( + logging_obj, + logging_executor, + async_call=False, + internal_call=False, + completion_with_fallbacks=False, + ) + ) ## LOAD CREDENTIALS load_credentials_from_list(kwargs) kwargs["litellm_logging_obj"] = logging_obj @@ -1650,7 +1587,12 @@ def client(original_function): except Exception as e: print_verbose(f"Error while checking max token limit: {e}") # MODEL CALL - result = original_function(*args, **kwargs) + invocation_kwargs: Final = ( + {**kwargs, "_litellm_call_completion": completion} + if original_function.__name__ == CallTypes.ocr.value + else kwargs + ) + result = original_function(*args, **invocation_kwargs) end_time = datetime.datetime.now() if _is_streaming_request( kwargs=kwargs, @@ -1704,8 +1646,11 @@ def client(original_function): kwargs=kwargs, ) - _update_response_metadata: Final = getattr(sys.modules[__name__], "update_response_metadata") - _update_response_metadata( + verbose_logger.info("Wrapper: Completed Call, calling success_handler") + completion.success(result, start_time, end_time) + # RETURN RESULT + update_response_metadata = getattr(sys.modules[__name__], "update_response_metadata") + update_response_metadata( result=result, logging_obj=logging_obj, model=model, @@ -1713,21 +1658,6 @@ def client(original_function): start_time=start_time, end_time=end_time, ) - - # LOG SUCCESS - handle streaming success logging in the _next_ object, remove `handle_success` once it's deprecated - verbose_logger.info("Wrapper: Completed Call, calling success_handler") - # Copy the current context to propagate it to the background thread - # This is essential for OpenTelemetry span context propagation - ctx: Final = contextvars.copy_context() - executor: Final = getattr(sys.modules[__name__], "executor") - executor.submit( - ctx.run, - logging_obj.success_handler, - result, - start_time, - end_time, - ) - # RETURN RESULT return result except Exception as e: call_type = original_function.__name__ @@ -1801,11 +1731,14 @@ def client(original_function): end_time = datetime.datetime.now() # LOG FAILURE - handle streaming failure logging in the _next_ object, remove `handle_failure` once it's deprecated - if logging_obj: - logging_obj.failure_handler( - e, traceback_exception, start_time, end_time - ) # DO NOT MAKE THREADED - router retry fallback relies on this! + if completion is not None: + completion.failure(e, traceback_exception, start_time, end_time) + elif logging_obj: + logging_obj.failure_handler(e, traceback_exception, start_time, end_time) raise e + finally: + if completion is not None: + completion.release() @wraps(original_function) async def wrapper_async(*args, **kwargs): @@ -1814,6 +1747,7 @@ def client(original_function): result = None _update_response_metadata: Final = getattr(sys.modules[__name__], "update_response_metadata") logging_obj: LiteLLMLoggingObject | None = kwargs.get("litellm_logging_obj", None) + completion: CallCompletion | None = None LLMCachingHandler: Final = _get_cached_llm_caching_handler() _llm_caching_handler: Final[LLMCachingHandler] = LLMCachingHandler( original_function=original_function, @@ -1839,7 +1773,15 @@ def client(original_function): # Type assertion: logging_obj is guaranteed to be non-None after function_setup assert logging_obj is not None, "logging_obj should not be None after function_setup" - + completion = CallCompletion( + PythonCompletion( + logging_obj, + None, + async_call=True, + internal_call=_is_litellm_internal_call, + completion_with_fallbacks=is_completion_with_fallbacks, + ) + ) modified_kwargs: Final = await async_pre_call_deployment_hook(kwargs, call_type) if modified_kwargs is not None: kwargs = modified_kwargs @@ -1924,7 +1866,12 @@ def client(original_function): # MODEL CALL try: - result = await original_function(*args, **kwargs) + invocation_kwargs: Final = ( + {**kwargs, "_litellm_call_completion": completion} + if original_function.__name__ == CallTypes.aocr.value + else kwargs + ) + result = await original_function(*args, **invocation_kwargs) except Exception as deployment_error: _deployment_call_end_time = datetime.datetime.now() # noqa: DTZ005 # matches the naive datetimes this whole function already times start_time/end_time with try: @@ -1987,20 +1934,13 @@ def client(original_function): args=args, ) + completion.success(result, start_time, end_time) # REBUILD EMBEDDING CACHING if ( isinstance(result, EmbeddingResponse) and _caching_handler_response is not None and _caching_handler_response.final_embedding_cached_response is not None ): - _dispatch_success_logging( - logging_obj=logging_obj, - result=result, - start_time=start_time, - end_time=end_time, - is_completion_with_fallbacks=is_completion_with_fallbacks, - is_litellm_internal_call=_is_litellm_internal_call, - ) return _llm_caching_handler._combine_cached_embedding_response_with_api_result( _caching_handler_response=_caching_handler_response, embedding_response=result, @@ -2016,14 +1956,6 @@ def client(original_function): start_time=start_time, end_time=end_time, ) - _dispatch_success_logging( - logging_obj=logging_obj, - result=result, - start_time=start_time, - end_time=end_time, - is_completion_with_fallbacks=is_completion_with_fallbacks, - is_litellm_internal_call=_is_litellm_internal_call, - ) return result except Exception as e: @@ -2031,17 +1963,12 @@ def client(original_function): # Reuse the timestamp taken right when the deployment call itself failed, before # the failure hook ran, so a slow callback doesn't inflate the reported duration. end_time = _deployment_call_end_time if _deployment_call_end_time is not None else datetime.datetime.now() # noqa: DTZ005 # matches the naive datetimes this whole function already times start_time/end_time with - if logging_obj and not _is_litellm_internal_call: - try: - logging_obj.failure_handler( - e, traceback_exception, start_time, end_time - ) # DO NOT MAKE THREADED - router retry fallback relies on this! - except Exception as e: - raise e - try: - await logging_obj.async_failure_handler(e, traceback_exception, start_time, end_time) - except Exception as e: - raise e + if completion is not None: + completion.failure(e, traceback_exception, start_time, end_time) + await completion.async_failure(e, traceback_exception, start_time, end_time) + elif logging_obj and not _is_litellm_internal_call: + logging_obj.failure_handler(e, traceback_exception, start_time, end_time) + await logging_obj.async_failure_handler(e, traceback_exception, start_time, end_time) call_type = original_function.__name__ num_retries, kwargs = _get_wrapper_num_retries(kwargs=kwargs, exception=e) @@ -2111,6 +2038,8 @@ def client(original_function): raise e finally: + if completion is not None: + completion.release() # Restore trace_id/session_id contextvars to their pre-call value once # this call (in this asyncio Task) is fully done - see # request_correlation_in_logs. Unlike wrapper()'s sync path, it's safe to @@ -7024,26 +6953,7 @@ class TextCompletionStreamWrapper: raise StopAsyncIteration -def mock_stream_usage_chunk(model_response: ModelResponseStream, model: str, prompt_tokens: int) -> ModelResponseStream: - return ModelResponseStream( - id=model_response.id, - choices=[], # mutable-ok: ModelResponseStream only treats a list as explicit choices, a tuple gets a default choice - model=model, - usage=Usage( - prompt_tokens=prompt_tokens, - completion_tokens=DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT, - total_tokens=prompt_tokens + DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT, - ), - ) - - -def mock_completion_streaming_obj( - model_response: ModelResponseStream, - mock_response: str | MockException | ModelResponseStream, - model: str, - n: int | None = None, - prompt_tokens: int | None = None, -) -> Iterator[ModelResponseStream]: +def mock_completion_streaming_obj(model_response, mock_response, model, n: int | None = None): if isinstance(mock_response, litellm.MockException): raise mock_response if isinstance(mock_response, ModelResponseStream): @@ -7063,17 +6973,14 @@ def mock_completion_streaming_obj( _all_choices.append(_streaming_choice) model_response.choices = _all_choices yield model_response - if prompt_tokens is not None: - yield mock_stream_usage_chunk(model_response, model=model, prompt_tokens=prompt_tokens) async def async_mock_completion_streaming_obj( - model_response: ModelResponseStream, + model_response, mock_response: str | MockException | ModelResponseStream, - model: str, + model, n: int | None = None, - prompt_tokens: int | None = None, -) -> AsyncIterator[ModelResponseStream]: +): if isinstance(mock_response, litellm.MockException): raise mock_response if isinstance(mock_response, ModelResponseStream): @@ -7093,8 +7000,6 @@ async def async_mock_completion_streaming_obj( _all_choices.append(_streaming_choice) model_response.choices = _all_choices yield model_response - if prompt_tokens is not None: - yield mock_stream_usage_chunk(model_response, model=model, prompt_tokens=prompt_tokens) ########## Reading Config File ############################ diff --git a/tests/test_litellm/litellm_core_utils/test_call_completion.py b/tests/test_litellm/litellm_core_utils/test_call_completion.py new file mode 100644 index 00000000000..b2c5ca751b9 --- /dev/null +++ b/tests/test_litellm/litellm_core_utils/test_call_completion.py @@ -0,0 +1,169 @@ +import contextvars +import datetime +from collections.abc import Callable, Coroutine +from concurrent.futures import Future +from typing import Final +from unittest.mock import AsyncMock, MagicMock + +import pytest + +from litellm.litellm_core_utils.call_completion import CallCompletion, PythonCompletion + + +class RecordingExecutor: + def __init__(self) -> None: + self.submissions: list[tuple[Callable[..., object], tuple[object, ...]]] = [] + + def submit(self, function: Callable[..., object], *args: object) -> Future[object]: + self.submissions.append((function, args)) + future: Final[Future[object]] = Future() + future.set_result(function(*args)) + return future + + +class RecordingCompletion: + def __init__(self) -> None: + self.successes: list[object] = [] + self.failures: list[Exception] = [] + + def success( + self, + result: object, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: + self.successes.append(result) + + def failure( + self, + exception: Exception, + traceback_exception: str, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: + self.failures.append(exception) + + async def async_failure( + self, + exception: Exception, + traceback_exception: str, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: + self.failures.append(exception) + + +class RecordingLogging: + def __init__(self, marker: contextvars.ContextVar[str], observed: list[tuple[object, str]]) -> None: + self._marker = marker + self._observed = observed + self._defer_async_logging = False + self._enqueue_deferred_logging: Callable[[], None] | None = None + + def success_handler( + self, + result: object, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: + self._observed.append((result, self._marker.get())) + + def failure_handler( + self, + exception: Exception, + traceback_exception: str, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: ... + + async def async_success_handler( + self, + result: object, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: ... + + async def async_failure_handler( + self, + exception: Exception, + traceback_exception: str, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: ... + + def handle_sync_success_callbacks_for_async_calls( + self, + result: object, + start_time: datetime.datetime, + end_time: datetime.datetime, + ) -> None: ... + + +def test_python_completion_preserves_sync_context_and_response_identity() -> None: + marker: Final[contextvars.ContextVar[str]] = contextvars.ContextVar("marker") + marker.set("request-context") + executor: Final = RecordingExecutor() + response: Final = object() + observed: Final[list[tuple[object, str]]] = [] + logging_obj: Final = RecordingLogging(marker, observed) + completion: Final = PythonCompletion( + logging_obj, + executor, + async_call=False, + internal_call=False, + completion_with_fallbacks=False, + ) + now: Final = datetime.datetime.now(datetime.timezone.utc) + + completion.success(response, now, now) + + assert observed == [(response, "request-context")] + assert len(executor.submissions) == 1 + + +@pytest.mark.asyncio +async def test_call_completion_attaches_once_and_forwards_final_objects() -> None: + python_completion: Final = RecordingCompletion() + native_completion: Final = RecordingCompletion() + ignored_completion: Final = RecordingCompletion() + completion: Final = CallCompletion(python_completion) + response: Final = object() + error: Final = ValueError("mapped failure") + now: Final = datetime.datetime.now(datetime.timezone.utc) + + assert completion.python_implementation is python_completion + assert completion.attach(native_completion) + assert not completion.attach(ignored_completion) + + completion.success(response, now, now) + completion.failure(error, "traceback", now, now) + await completion.async_failure(error, "traceback", now, now) + + assert native_completion.successes == [response] + assert native_completion.failures == [error, error] + assert python_completion.successes == [] + assert ignored_completion.successes == [] + + +@pytest.mark.asyncio +async def test_python_completion_retains_deferred_success_arguments(monkeypatch: pytest.MonkeyPatch) -> None: + response: Final = object() + logging_obj: Final = MagicMock() + logging_obj._defer_async_logging = True + logging_obj.async_success_handler = AsyncMock() + completion: Final = PythonCompletion( + logging_obj, + RecordingExecutor(), + async_call=True, + internal_call=False, + completion_with_fallbacks=False, + ) + scheduled: Final[list[Coroutine[object, object, None]]] = [] + monkeypatch.setattr("asyncio.create_task", scheduled.append) + now: Final = datetime.datetime.now(datetime.timezone.utc) + + completion.success(response, now, now) + + logging_obj._enqueue_deferred_logging() + assert len(scheduled) == 1 + scheduled[0].close()