diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 5b4187846ff..0f22f48bbcf 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -1337,6 +1337,21 @@ class Logging(LiteLLMLoggingBaseClass): additional_args.get("api_base", "") ) + async def async_pre_call(self) -> None: + if self.model_call_details.get("has_logged_async_pre_call"): + return + + self.model_call_details["has_logged_async_pre_call"] = True + self.model_call_details.update(model=self.model, messages=self.messages, log_event_type="pre_api_call") + for callback in litellm._async_input_callback: + try: + if callable(callback): + await callback(self.model_call_details) + except Exception as e: + verbose_logger.exception( + "LiteLLM.LoggingError: [Non-Blocking] Exception occurred while async pre-call logging %s", e + ) + def pre_call(self, input, api_key, model=None, additional_args={}): # Log the exact input to the LLM API try: diff --git a/litellm/utils.py b/litellm/utils.py index d45b29c0f16..06347564324 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -2079,6 +2079,7 @@ def client(original_function): else kwargs ) try: + await logging_obj.async_pre_call() result = await original_function(*args, **call_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 diff --git a/tests/unit/test_utils.py b/tests/unit/test_utils.py index 2c612aa350c..b0c617ddb2e 100644 --- a/tests/unit/test_utils.py +++ b/tests/unit/test_utils.py @@ -5282,6 +5282,112 @@ async def test_wrapper_async_claims_the_budget_reservation_a_supplied_logging_ob assert reservation["callback_bound"] is True +@pytest.mark.asyncio +@pytest.mark.parametrize("registration", ["manager", "input_callback"]) +async def test_wrapper_async_awaits_async_input_callback_before_the_provider_call( + monkeypatch: pytest.MonkeyPatch, registration: str +) -> None: + events: Final[list[tuple[str, str]]] = [] + messages: Final = [{"role": "user", "content": "hello"}] + + async def async_callback(kwargs: dict) -> None: + assert kwargs["model"] == "gpt-4o-mini" + assert kwargs["messages"] == messages + events.append(("async", kwargs["log_event_type"])) + + def sync_callback(kwargs: dict) -> None: + events.append(("sync", kwargs["log_event_type"])) + + if registration == "manager": + litellm.logging_callback_manager.add_litellm_input_callback(async_callback) + else: + monkeypatch.setattr(litellm, "input_callback", [async_callback]) + litellm.logging_callback_manager.add_litellm_input_callback(sync_callback) + + response: Final = await litellm.acompletion(model="gpt-4o-mini", messages=messages, mock_response="ok") + + assert response.choices[0].message.content == "ok" + assert events == [("async", "pre_api_call"), ("sync", "pre_api_call")] + + +@pytest.mark.asyncio +async def test_wrapper_async_continues_when_an_async_input_callback_raises( + monkeypatch: pytest.MonkeyPatch, +) -> None: + events: Final[list[str]] = [] + + async def failing_callback(kwargs: dict) -> None: + events.append("failing") + raise RuntimeError("callback failed") + + async def following_callback(kwargs: dict) -> None: + events.append(kwargs["log_event_type"]) + + monkeypatch.setattr(litellm, "_async_input_callback", [failing_callback, following_callback]) + response: Final = await litellm.acompletion( + model="gpt-4o-mini", messages=[{"role": "user", "content": "hello"}], mock_response="ok" + ) + + assert events == ["failing", "pre_api_call"] + assert response.choices[0].message.content == "ok" + + +@pytest.mark.asyncio +async def test_wrapper_async_skips_async_input_callback_on_cache_hit( + monkeypatch: pytest.MonkeyPatch, +) -> None: + events: Final[list[str]] = [] + + async def async_callback(kwargs: dict) -> None: + events.append("async") + + def sync_callback(kwargs: dict) -> None: + events.append("sync") + + monkeypatch.setattr(litellm, "_async_input_callback", [async_callback]) + monkeypatch.setattr(litellm, "input_callback", [sync_callback]) + monkeypatch.setattr(litellm, "cache", Cache(type="local")) + request: Final = { + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "hello"}], + "mock_response": "ok", + "caching": True, + } + + first: Final = await litellm.acompletion(**request) + await asyncio.gather(*tuple(_PENDING_CACHE_WRITES)) + second: Final = await litellm.acompletion(**request) + + assert first.choices[0].message.content == second.choices[0].message.content == "ok" + assert events == ["async", "sync"] + + +@pytest.mark.asyncio +async def test_wrapper_async_runs_async_input_callback_once_per_logging_object( + monkeypatch: pytest.MonkeyPatch, +) -> None: + events: Final[list[str]] = [] + + 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(), + start_time=datetime.now(), + model="gpt-4o-mini", + messages=[{"role": "user", "content": "hello"}], + litellm_call_id="shared-logging-object", + ) + + 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" + assert events == ["pre_api_call"] + + @pytest.mark.asyncio async def test_wrapper_async_hands_the_budget_reservation_back_when_the_call_fails() -> None: reservation = _budget_reservation()