mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(logging): run async input callbacks once per provider attempt
This commit is contained in:
parent
6eeb1b1c77
commit
ef80f9da19
4 changed files with 332 additions and 28 deletions
|
|
@ -6,6 +6,8 @@ be settable from user input. Context variables are scoped to the current
|
||||||
asyncio task and cannot be injected via HTTP request bodies.
|
asyncio task and cannot be injected via HTTP request bodies.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import weakref
|
||||||
from collections.abc import Generator
|
from collections.abc import Generator
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from contextvars import ContextVar
|
from contextvars import ContextVar
|
||||||
|
|
@ -23,6 +25,10 @@ _billing_time: Final[ContextVar[datetime | None]] = ContextVar("billing_time", d
|
||||||
|
|
||||||
_post_response: Final[ContextVar[bool]] = ContextVar("post_response", default=False)
|
_post_response: Final[ContextVar[bool]] = ContextVar("post_response", default=False)
|
||||||
|
|
||||||
|
_provider_calls_in_flight: Final[ContextVar[frozenset[tuple[weakref.ref[asyncio.Task[object]], int]]]] = ContextVar(
|
||||||
|
"provider_calls_in_flight", default=frozenset()
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def post_response_phase() -> Generator[None]:
|
def post_response_phase() -> Generator[None]:
|
||||||
|
|
@ -38,6 +44,22 @@ def in_post_response_phase() -> bool:
|
||||||
return _post_response.get()
|
return _post_response.get()
|
||||||
|
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def provider_call_scope(logging_obj: object) -> Generator[bool]:
|
||||||
|
"""Yield whether an outer call in this same task, e.g. a Responses bridge, is already making this logging object's provider call."""
|
||||||
|
task: Final = asyncio.current_task()
|
||||||
|
if task is None:
|
||||||
|
yield False
|
||||||
|
return
|
||||||
|
key: Final = (weakref.ref(task), id(logging_obj))
|
||||||
|
in_flight: Final = _provider_calls_in_flight.get()
|
||||||
|
token: Final = _provider_calls_in_flight.set(in_flight.union((key,)))
|
||||||
|
try:
|
||||||
|
yield key in in_flight
|
||||||
|
finally:
|
||||||
|
_provider_calls_in_flight.reset(token)
|
||||||
|
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def pinned_billing_time(moment: datetime) -> Generator[None]:
|
def pinned_billing_time(moment: datetime) -> Generator[None]:
|
||||||
"""Price every rate lookup inside this block at ``moment`` rather than at each one's own clock read."""
|
"""Price every rate lookup inside this block at ``moment`` rather than at each one's own clock read."""
|
||||||
|
|
|
||||||
|
|
@ -74,6 +74,7 @@ from litellm.litellm_core_utils.core_helpers import (
|
||||||
reconstruct_model_name,
|
reconstruct_model_name,
|
||||||
set_response_cost_in_hidden_params,
|
set_response_cost_in_hidden_params,
|
||||||
)
|
)
|
||||||
|
from litellm.litellm_core_utils.coroutine_checker import coroutine_checker
|
||||||
from litellm.litellm_core_utils.error_normalization import normalize_error
|
from litellm.litellm_core_utils.error_normalization import normalize_error
|
||||||
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
|
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
|
||||||
from litellm.litellm_core_utils.internal_call_metadata import (
|
from litellm.litellm_core_utils.internal_call_metadata import (
|
||||||
|
|
@ -1338,15 +1339,31 @@ class Logging(LiteLLMLoggingBaseClass):
|
||||||
)
|
)
|
||||||
|
|
||||||
async def async_pre_call(self) -> None:
|
async def async_pre_call(self) -> None:
|
||||||
if self.model_call_details.get("has_logged_async_pre_call"):
|
candidates: Final = (
|
||||||
|
*(callback for callback in litellm._async_input_callback if callable(callback)),
|
||||||
|
*(
|
||||||
|
callback
|
||||||
|
for callback in litellm.input_callback
|
||||||
|
if callable(callback) and coroutine_checker.is_async_callable(callback)
|
||||||
|
),
|
||||||
|
)
|
||||||
|
callbacks: Final = tuple(
|
||||||
|
callback
|
||||||
|
for index, callback in enumerate(candidates)
|
||||||
|
if not any(earlier is callback for earlier in candidates[:index])
|
||||||
|
)
|
||||||
|
if not callbacks:
|
||||||
return
|
return
|
||||||
|
|
||||||
self.model_call_details["has_logged_async_pre_call"] = True
|
request: Final = { # mutable-ok: callbacks receive the plain kwargs dict they always got
|
||||||
self.model_call_details.update(model=self.model, messages=self.messages, log_event_type="pre_api_call")
|
**self.model_call_details,
|
||||||
for callback in litellm._async_input_callback:
|
"model": self.model,
|
||||||
|
"messages": self.messages,
|
||||||
|
"log_event_type": "pre_api_call",
|
||||||
|
}
|
||||||
|
for callback in callbacks:
|
||||||
try:
|
try:
|
||||||
if callable(callback):
|
await callback(request)
|
||||||
await callback(self.model_call_details)
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
verbose_logger.exception(
|
verbose_logger.exception(
|
||||||
"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while async pre-call logging %s", e
|
"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while async pre-call logging %s", e
|
||||||
|
|
@ -1460,7 +1477,11 @@ class Logging(LiteLLMLoggingBaseClass):
|
||||||
messages=self.messages,
|
messages=self.messages,
|
||||||
kwargs=self.model_call_details,
|
kwargs=self.model_call_details,
|
||||||
)
|
)
|
||||||
elif callable(callback) and customLogger is not None: # custom logger functions
|
elif (
|
||||||
|
callable(callback)
|
||||||
|
and customLogger is not None
|
||||||
|
and not coroutine_checker.is_async_callable(callback)
|
||||||
|
): # custom logger functions
|
||||||
customLogger.log_input_event(
|
customLogger.log_input_event(
|
||||||
model=self.model,
|
model=self.model,
|
||||||
messages=self.messages,
|
messages=self.messages,
|
||||||
|
|
|
||||||
|
|
@ -52,7 +52,7 @@ import litellm.litellm_core_utils
|
||||||
|
|
||||||
# audio_utils.utils is lazy-loaded - only imported when needed for transcription calls
|
# audio_utils.utils is lazy-loaded - only imported when needed for transcription calls
|
||||||
import litellm.litellm_core_utils.json_validation_rule
|
import litellm.litellm_core_utils.json_validation_rule
|
||||||
from litellm._internal_context import is_internal_call
|
from litellm._internal_context import is_internal_call, provider_call_scope
|
||||||
from litellm._lazy_imports import (
|
from litellm._lazy_imports import (
|
||||||
_get_default_encoding,
|
_get_default_encoding,
|
||||||
_get_modified_max_tokens,
|
_get_modified_max_tokens,
|
||||||
|
|
@ -2079,9 +2079,15 @@ def client(original_function):
|
||||||
else kwargs
|
else kwargs
|
||||||
)
|
)
|
||||||
try:
|
try:
|
||||||
if litellm._async_input_callback:
|
if litellm._async_input_callback or (
|
||||||
await logging_obj.async_pre_call()
|
litellm.input_callback and any(check_coroutine(callback) for callback in litellm.input_callback)
|
||||||
result = await original_function(*args, **call_kwargs)
|
):
|
||||||
|
with provider_call_scope(logging_obj) as nested_provider_call:
|
||||||
|
if not nested_provider_call:
|
||||||
|
await logging_obj.async_pre_call()
|
||||||
|
result = await original_function(*args, **call_kwargs)
|
||||||
|
else:
|
||||||
|
result = await original_function(*args, **call_kwargs)
|
||||||
except Exception as deployment_error:
|
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
|
_deployment_call_end_time = datetime.datetime.now() # noqa: DTZ005 # matches the naive datetimes this whole function already times start_time/end_time with
|
||||||
try:
|
try:
|
||||||
|
|
|
||||||
|
|
@ -2,13 +2,17 @@ import asyncio
|
||||||
import base64
|
import base64
|
||||||
import contextlib
|
import contextlib
|
||||||
import contextvars
|
import contextvars
|
||||||
|
import functools
|
||||||
|
import gc
|
||||||
import io
|
import io
|
||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
import queue
|
import queue
|
||||||
import threading
|
import threading
|
||||||
from collections.abc import Callable, Iterator, Mapping
|
import warnings
|
||||||
|
import weakref
|
||||||
|
from collections.abc import Callable, Coroutine, Iterator, Mapping
|
||||||
from concurrent.futures import Future, ThreadPoolExecutor
|
from concurrent.futures import Future, ThreadPoolExecutor
|
||||||
from datetime import datetime, timedelta, timezone
|
from datetime import datetime, timedelta, timezone
|
||||||
from pathlib import PurePath
|
from pathlib import PurePath
|
||||||
|
|
@ -35,6 +39,7 @@ from litellm.constants import DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT
|
||||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||||
from litellm.integrations.custom_logger import CustomLogger
|
from litellm.integrations.custom_logger import CustomLogger
|
||||||
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
|
from litellm.litellm_core_utils.get_litellm_params import get_litellm_params
|
||||||
|
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||||
from litellm.litellm_core_utils.thread_pool_executor import executor as logging_executor
|
from litellm.litellm_core_utils.thread_pool_executor import executor as logging_executor
|
||||||
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
|
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
|
||||||
from litellm.proxy.utils import is_valid_api_key
|
from litellm.proxy.utils import is_valid_api_key
|
||||||
|
|
@ -5366,32 +5371,282 @@ async def test_wrapper_async_skips_async_input_callback_on_cache_hit(
|
||||||
assert events == ["async", "sync"]
|
assert events == ["async", "sync"]
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
def _async_input_callback_logging_object(
|
||||||
async def test_wrapper_async_runs_async_input_callback_once_per_logging_object(
|
original_function: str, litellm_call_id: str
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
) -> tuple[Logging, Mapping[str, object]]:
|
||||||
) -> None:
|
return litellm.utils.function_setup(
|
||||||
events: Final[list[str]] = []
|
original_function=original_function,
|
||||||
|
|
||||||
async def callback(kwargs: dict) -> None:
|
|
||||||
events.append(kwargs["log_event_type"])
|
|
||||||
|
|
||||||
monkeypatch.setattr(litellm, "_async_input_callback", [callback])
|
|
||||||
logging_obj, kwargs = litellm.utils.function_setup(
|
|
||||||
original_function="acompletion",
|
|
||||||
rules_obj=litellm.utils.Rules(),
|
rules_obj=litellm.utils.Rules(),
|
||||||
start_time=datetime.now(),
|
start_time=datetime.now(),
|
||||||
model="gpt-4o-mini",
|
model="gpt-4o-mini",
|
||||||
messages=[{"role": "user", "content": "hello"}],
|
messages=[{"role": "user", "content": "hello"}],
|
||||||
litellm_call_id="shared-logging-object",
|
litellm_call_id=litellm_call_id,
|
||||||
)
|
)
|
||||||
|
|
||||||
first: Final = await litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="ok")
|
|
||||||
second: Final = await litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="ok")
|
|
||||||
|
|
||||||
assert first.choices[0].message.content == second.choices[0].message.content == "ok"
|
@pytest.mark.asyncio
|
||||||
|
async def test_wrapper_async_reruns_async_input_callback_when_a_failed_attempt_is_retried_on_the_same_logging_object(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
events: Final[list[object]] = []
|
||||||
|
|
||||||
|
async def callback(kwargs: Mapping[str, object]) -> None:
|
||||||
|
events.append(kwargs["log_event_type"])
|
||||||
|
|
||||||
|
monkeypatch.setattr(litellm, "_async_input_callback", [callback])
|
||||||
|
logging_obj, kwargs = _async_input_callback_logging_object("acompletion", "retried-logging-object")
|
||||||
|
rate_limited: Final = litellm.RateLimitError(message="busy", llm_provider="openai", model="gpt-4o-mini")
|
||||||
|
|
||||||
|
with pytest.raises(litellm.RateLimitError):
|
||||||
|
await litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response=rate_limited)
|
||||||
|
retried: Final = await litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="ok")
|
||||||
|
|
||||||
|
assert retried.choices[0].message.content == "ok"
|
||||||
|
assert events == ["pre_api_call", "pre_api_call"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_wrapper_async_runs_async_input_callback_once_when_responses_bridges_to_acompletion(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
call_ids: Final[list[object]] = []
|
||||||
|
|
||||||
|
async def callback(kwargs: Mapping[str, object]) -> None:
|
||||||
|
call_ids.append(kwargs["litellm_call_id"])
|
||||||
|
|
||||||
|
monkeypatch.setattr(litellm, "_async_input_callback", [callback])
|
||||||
|
|
||||||
|
response: Final = await litellm.aresponses(
|
||||||
|
model="openai/gpt-4o-mini",
|
||||||
|
input="hello",
|
||||||
|
use_chat_completions_api=True,
|
||||||
|
mock_response=ModelResponse(choices=[{"message": {"role": "assistant", "content": "bridged"}}]),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.output[0].content[0].text == "bridged"
|
||||||
|
assert len(call_ids) == 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_wrapper_async_awaits_async_input_callback_appended_directly_when_the_logging_object_is_supplied(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
events: Final[list[tuple[str, object]]] = []
|
||||||
|
|
||||||
|
async def async_callback(kwargs: Mapping[str, object]) -> None:
|
||||||
|
events.append(("async", kwargs["log_event_type"]))
|
||||||
|
|
||||||
|
def sync_callback(kwargs: Mapping[str, object]) -> None:
|
||||||
|
events.append(("sync", kwargs["log_event_type"]))
|
||||||
|
|
||||||
|
monkeypatch.setattr("litellm.litellm_core_utils.litellm_logging.customLogger", CustomLogger())
|
||||||
|
logging_obj, kwargs = _async_input_callback_logging_object("acompletion", "supplied-logging-object")
|
||||||
|
monkeypatch.setattr(litellm, "input_callback", [async_callback, sync_callback])
|
||||||
|
|
||||||
|
response: Final = await litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="ok")
|
||||||
|
|
||||||
|
assert response.choices[0].message.content == "ok"
|
||||||
|
assert events == [("async", "pre_api_call"), ("sync", "pre_api_call")]
|
||||||
|
|
||||||
|
|
||||||
|
def test_completion_does_not_leave_an_async_input_callback_unawaited(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||||
|
events: Final[list[object]] = []
|
||||||
|
|
||||||
|
async def async_input_callback(kwargs: Mapping[str, object]) -> None:
|
||||||
|
events.append("async")
|
||||||
|
|
||||||
|
def sync_callback(kwargs: Mapping[str, object]) -> None:
|
||||||
|
events.append("sync")
|
||||||
|
|
||||||
|
monkeypatch.setattr("litellm.litellm_core_utils.litellm_logging.customLogger", CustomLogger())
|
||||||
|
logging_obj, kwargs = _async_input_callback_logging_object("completion", "sync-supplied-logging-object")
|
||||||
|
monkeypatch.setattr(litellm, "input_callback", [async_input_callback, sync_callback])
|
||||||
|
|
||||||
|
with warnings.catch_warnings(record=True) as caught:
|
||||||
|
warnings.simplefilter("always")
|
||||||
|
response: Final = litellm.completion(**kwargs, litellm_logging_obj=logging_obj, mock_response="ok")
|
||||||
|
gc.collect()
|
||||||
|
|
||||||
|
assert response.choices[0].message.content == "ok"
|
||||||
|
assert events == ["sync"]
|
||||||
|
assert [str(w.message) for w in caught if async_input_callback.__name__ in str(w.message)] == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_async_pre_call_hands_callbacks_the_request_without_changing_the_logging_object(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
received: Final[list[tuple[object, object, object]]] = []
|
||||||
|
|
||||||
|
async def callback(kwargs: Mapping[str, object]) -> None:
|
||||||
|
received.append((kwargs["model"], kwargs["messages"], kwargs["log_event_type"]))
|
||||||
|
|
||||||
|
monkeypatch.setattr(litellm, "_async_input_callback", [callback])
|
||||||
|
logging_obj, _ = _async_input_callback_logging_object("acompletion", "untouched-logging-object")
|
||||||
|
before: Final = dict(logging_obj.model_call_details)
|
||||||
|
|
||||||
|
await logging_obj.async_pre_call()
|
||||||
|
|
||||||
|
assert received == [("gpt-4o-mini", [{"role": "user", "content": "hello"}], "pre_api_call")]
|
||||||
|
assert logging_obj.model_call_details == before
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_wrapper_async_runs_async_input_callback_for_concurrent_and_later_calls_sharing_a_logging_object(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
events: Final[list[object]] = []
|
||||||
|
|
||||||
|
async def callback(kwargs: Mapping[str, object]) -> None:
|
||||||
|
events.append(kwargs["log_event_type"])
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
|
||||||
|
monkeypatch.setattr(litellm, "_async_input_callback", [callback])
|
||||||
|
logging_obj, kwargs = _async_input_callback_logging_object("acompletion", "shared-logging-object")
|
||||||
|
|
||||||
|
concurrent: Final = await asyncio.gather(
|
||||||
|
litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="first"),
|
||||||
|
litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="second"),
|
||||||
|
)
|
||||||
|
reused: Final = await litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="third")
|
||||||
|
|
||||||
|
assert [response.choices[0].message.content for response in (*concurrent, reused)] == ["first", "second", "third"]
|
||||||
|
assert events == ["pre_api_call", "pre_api_call", "pre_api_call"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_wrapper_async_runs_async_input_callback_for_a_child_task_reusing_the_logging_object_after_the_parent_call(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
events: Final[list[object]] = []
|
||||||
|
parent_done: Final = asyncio.Event()
|
||||||
|
children: Final[list[asyncio.Task[ModelResponse | CustomStreamWrapper]]] = []
|
||||||
|
logging_obj, kwargs = _async_input_callback_logging_object("acompletion", "child-task-logging-object")
|
||||||
|
|
||||||
|
async def child_call() -> ModelResponse | CustomStreamWrapper:
|
||||||
|
await parent_done.wait()
|
||||||
|
return await litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="child")
|
||||||
|
|
||||||
|
async def callback(kwargs: Mapping[str, object]) -> None:
|
||||||
|
events.append(kwargs["log_event_type"])
|
||||||
|
if not children:
|
||||||
|
children.append(asyncio.create_task(child_call()))
|
||||||
|
|
||||||
|
monkeypatch.setattr(litellm, "_async_input_callback", [callback])
|
||||||
|
|
||||||
|
parent: Final = await litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="parent")
|
||||||
|
parent_done.set()
|
||||||
|
child: Final = await children[0]
|
||||||
|
|
||||||
|
assert (parent.choices[0].message.content, child.choices[0].message.content) == ("parent", "child")
|
||||||
|
assert events == ["pre_api_call", "pre_api_call"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_wrapper_async_does_not_keep_the_parent_task_alive_for_a_child_task_that_outlives_it(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
release_child: Final = asyncio.Event()
|
||||||
|
children: Final[list[asyncio.Task[ModelResponse | CustomStreamWrapper]]] = []
|
||||||
|
logging_obj, kwargs = _async_input_callback_logging_object("acompletion", "outlived-parent-task")
|
||||||
|
|
||||||
|
async def child_call() -> ModelResponse | CustomStreamWrapper:
|
||||||
|
await release_child.wait()
|
||||||
|
return await litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="child")
|
||||||
|
|
||||||
|
async def callback(request: Mapping[str, object]) -> None:
|
||||||
|
if not children:
|
||||||
|
children.append(asyncio.create_task(child_call()))
|
||||||
|
|
||||||
|
async def run_parent() -> tuple[
|
||||||
|
ModelResponse | CustomStreamWrapper, weakref.ref[asyncio.Task[ModelResponse | CustomStreamWrapper]]
|
||||||
|
]:
|
||||||
|
parent_task: Final = asyncio.create_task(
|
||||||
|
litellm.acompletion(**kwargs, litellm_logging_obj=logging_obj, mock_response="parent")
|
||||||
|
)
|
||||||
|
return await parent_task, weakref.ref(parent_task)
|
||||||
|
|
||||||
|
monkeypatch.setattr(litellm, "_async_input_callback", [callback])
|
||||||
|
|
||||||
|
parent, parent_task_ref = await run_parent()
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
gc.collect()
|
||||||
|
|
||||||
|
assert parent_task_ref() is None
|
||||||
|
assert not children[0].done()
|
||||||
|
release_child.set()
|
||||||
|
child: Final = await children[0]
|
||||||
|
assert (parent.choices[0].message.content, child.choices[0].message.content) == ("parent", "child")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_wrapper_async_awaits_registered_async_input_callbacks_that_are_not_coroutine_functions(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
events: Final[list[object]] = []
|
||||||
|
|
||||||
|
async def record(label: str, kwargs: Mapping[str, object]) -> None:
|
||||||
|
events.append(f"{label}:{kwargs['log_event_type']}")
|
||||||
|
|
||||||
|
def returns_coroutine(kwargs: Mapping[str, object]) -> Coroutine[object, object, None]:
|
||||||
|
return record("sync-returning-coroutine", kwargs)
|
||||||
|
|
||||||
|
monkeypatch.setattr(litellm, "_async_input_callback", [functools.partial(record, "partial"), returns_coroutine])
|
||||||
|
|
||||||
|
with warnings.catch_warnings(record=True) as caught:
|
||||||
|
warnings.simplefilter("always")
|
||||||
|
response: Final = await litellm.acompletion(
|
||||||
|
model="gpt-4o-mini", messages=[{"role": "user", "content": "hello"}], mock_response="ok"
|
||||||
|
)
|
||||||
|
gc.collect()
|
||||||
|
|
||||||
|
assert response.choices[0].message.content == "ok"
|
||||||
|
assert events == ["partial:pre_api_call", "sync-returning-coroutine:pre_api_call"]
|
||||||
|
assert [str(w.message) for w in caught if "never awaited" in str(w.message)] == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_wrapper_async_runs_an_async_input_callback_registered_in_both_lists_once(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
events: Final[list[object]] = []
|
||||||
|
|
||||||
|
async def callback(kwargs: Mapping[str, object]) -> None:
|
||||||
|
events.append(kwargs["log_event_type"])
|
||||||
|
|
||||||
|
monkeypatch.setattr(litellm, "_async_input_callback", [callback])
|
||||||
|
monkeypatch.setattr(litellm, "input_callback", [callback])
|
||||||
|
|
||||||
|
response: Final = await litellm.acompletion(
|
||||||
|
model="gpt-4o-mini", messages=[{"role": "user", "content": "hello"}], mock_response="ok"
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.choices[0].message.content == "ok"
|
||||||
assert events == ["pre_api_call"]
|
assert events == ["pre_api_call"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_wrapper_async_does_not_need_async_pre_call_on_a_supplied_logging_object_with_only_sync_input_callbacks(
|
||||||
|
monkeypatch: pytest.MonkeyPatch,
|
||||||
|
) -> None:
|
||||||
|
response: Final = object()
|
||||||
|
logger: Final = MagicMock()
|
||||||
|
del logger.async_pre_call
|
||||||
|
|
||||||
|
async def acompletion(**kwargs: object) -> object:
|
||||||
|
return response
|
||||||
|
|
||||||
|
def sync_callback(kwargs: Mapping[str, object]) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
monkeypatch.setattr(litellm, "_async_input_callback", [])
|
||||||
|
monkeypatch.setattr(litellm, "input_callback", [sync_callback])
|
||||||
|
|
||||||
|
result: Final = await client(acompletion)(model="gpt-4o-mini", litellm_logging_obj=logger)
|
||||||
|
|
||||||
|
assert result is response
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_wrapper_async_hands_the_budget_reservation_back_when_the_call_fails() -> None:
|
async def test_wrapper_async_hands_the_budget_reservation_back_when_the_call_fails() -> None:
|
||||||
reservation = _budget_reservation()
|
reservation = _budget_reservation()
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue