mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(logging): skip sync success callbacks for internal sub-calls, deflake RAG and Langfuse tests (#45345)
* fix(logging): skip sync success callbacks for internal sub-calls Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(langfuse): assert no retry sleep instead of a wall-clock budget Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test: add types to deflake regression coverage Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * style(tests): avoid unrelated formatting changes Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
01166906fb
commit
df13d59306
3 changed files with 127 additions and 21 deletions
|
|
@ -1895,17 +1895,18 @@ def client(original_function):
|
|||
|
||||
# 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[BoundedLoggingThreadPoolExecutor] = getattr(sys.modules[__name__], "executor")
|
||||
executor.submit(
|
||||
ctx.run,
|
||||
logging_obj.success_handler,
|
||||
result,
|
||||
start_time,
|
||||
end_time,
|
||||
)
|
||||
if not is_internal_call.get():
|
||||
# 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[BoundedLoggingThreadPoolExecutor] = 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:
|
||||
|
|
|
|||
|
|
@ -8,17 +8,19 @@ otherwise record its own duration instead of the call's.
|
|||
import json
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
from base64 import b64encode
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from time import monotonic, sleep
|
||||
from types import MappingProxyType
|
||||
from types import MappingProxyType, SimpleNamespace
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
import opentelemetry.trace as otel_trace
|
||||
import pytest
|
||||
from langfuse import LangfuseOtelSpanAttributes as A
|
||||
from langfuse.api.core import http_client as langfuse_http_client
|
||||
from langfuse.api.core.api_error import ApiError
|
||||
from langfuse.api.core.request_options import RequestOptions
|
||||
from opentelemetry.proto.collector.trace.v1.trace_service_pb2 import ExportTraceServiceRequest
|
||||
|
|
@ -1042,28 +1044,35 @@ def test_auth_check_fails_when_the_keys_reach_no_project():
|
|||
|
||||
|
||||
@pytest.mark.parametrize("status", [500, 503, 429], ids=["http-500", "http-503", "http-429"])
|
||||
def test_auth_check_and_project_id_make_one_round_trip_when_langfuse_is_down(status):
|
||||
def test_auth_check_and_project_id_make_one_round_trip_when_langfuse_is_down(
|
||||
status: int, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
"""Both run on the event loop; the generated client's default retries sleep for seconds, or for Retry-After."""
|
||||
requests: list[httpx.Request] = []
|
||||
requests: Final[list[httpx.Request]] = []
|
||||
sleeps: Final[list[float]] = []
|
||||
monkeypatch.setattr(
|
||||
langfuse_http_client,
|
||||
"time",
|
||||
SimpleNamespace(sleep=sleeps.append, time=time.time),
|
||||
)
|
||||
|
||||
def fail(request: httpx.Request) -> httpx.Response:
|
||||
requests.append(request)
|
||||
return httpx.Response(status, request=request, headers={"retry-after": "20"}, json={"message": "down"})
|
||||
|
||||
client = build_langfuse_client(
|
||||
client: Final = build_langfuse_client(
|
||||
public_key="pk",
|
||||
secret_key="sk",
|
||||
base_url="http://127.0.0.1:1",
|
||||
httpx_client=httpx.Client(transport=httpx.MockTransport(fail)),
|
||||
)
|
||||
|
||||
started = monotonic()
|
||||
failure = client.auth_check()
|
||||
failure: Final = client.auth_check()
|
||||
with pytest.raises(ApiError):
|
||||
client.project_id()
|
||||
assert failure is not None and f"status_code: {status}" in failure.reason
|
||||
assert len(requests) == 2
|
||||
assert monotonic() - started < 0.5
|
||||
assert sleeps == [], "failed auth checks should not sleep for REST retries"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
|
|||
|
|
@ -11,9 +11,12 @@ aquery carries the completion response with real usage and cost.
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
import datetime
|
||||
import json
|
||||
from collections.abc import Callable, Mapping
|
||||
from concurrent.futures import Future
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from typing import Final, cast
|
||||
from unittest.mock import patch
|
||||
|
||||
import httpx
|
||||
|
|
@ -22,8 +25,10 @@ import respx
|
|||
from pydantic import ValidationError
|
||||
|
||||
import litellm
|
||||
import litellm.utils as litellm_utils
|
||||
from litellm._internal_context import is_internal_call
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.litellm_core_utils import litellm_logging
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
from litellm.types.utils import CallTypes, ModelResponse
|
||||
|
||||
|
|
@ -45,11 +50,27 @@ async def _drain_logging_worker() -> None:
|
|||
class RecordingLogger(CustomLogger):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.success_events = []
|
||||
self.success_events: Final[list[dict[str, object]]] = []
|
||||
self.sync_success_events: Final[list[dict[str, object]]] = []
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
async def async_log_success_event(
|
||||
self,
|
||||
kwargs: Mapping[str, object],
|
||||
response_obj: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None:
|
||||
self.success_events.append({"kwargs": kwargs, "response_obj": response_obj})
|
||||
|
||||
def log_success_event(
|
||||
self,
|
||||
kwargs: Mapping[str, object],
|
||||
response_obj: object,
|
||||
start_time: datetime.datetime,
|
||||
end_time: datetime.datetime,
|
||||
) -> None:
|
||||
self.sync_success_events.append({"kwargs": kwargs, "response_obj": response_obj})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("use_router", [False, True])
|
||||
|
|
@ -107,6 +128,81 @@ async def test_aquery_single_billing_event_carries_completion_usage_and_cost(use
|
|||
assert standard_logging_object["response_cost"] > 0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("use_router", [False, True])
|
||||
async def test_aquery_vector_store_search_sub_call_logs_no_sync_success_event_on_the_parent(
|
||||
use_router: bool, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
await _drain_logging_worker()
|
||||
recording_logger: Final = RecordingLogger()
|
||||
|
||||
class InlineExecutor:
|
||||
def submit(
|
||||
self,
|
||||
fn: Callable[..., object],
|
||||
*args: object,
|
||||
**kwargs: object,
|
||||
) -> Future[object]:
|
||||
future: Final = Future[object]()
|
||||
future.set_result(fn(*args, **kwargs))
|
||||
return future
|
||||
|
||||
monkeypatch.setattr(litellm_utils, "executor", InlineExecutor())
|
||||
monkeypatch.setattr(litellm_logging, "executor", InlineExecutor())
|
||||
monkeypatch.setattr(litellm, "callbacks", [recording_logger])
|
||||
|
||||
router_kwargs: Final = (
|
||||
{
|
||||
"router": litellm.Router(
|
||||
model_list=[
|
||||
{
|
||||
"model_name": "gpt-4o-mini",
|
||||
"litellm_params": {"model": "openai/gpt-4o-mini", "api_key": "test-key"},
|
||||
}
|
||||
]
|
||||
)
|
||||
}
|
||||
if use_router
|
||||
else {}
|
||||
)
|
||||
|
||||
response: Final = await litellm.aquery(
|
||||
model="gpt-4o-mini",
|
||||
messages=[{"role": "user", "content": "What is the secret project codename?"}],
|
||||
retrieval_config={"vector_store_id": "vs_test_123", "custom_llm_provider": "openai"},
|
||||
mock_response="The secret project codename is AZURE-FALCON-42.",
|
||||
**router_kwargs,
|
||||
)
|
||||
assert isinstance(response, ModelResponse), "aquery should return its completion response"
|
||||
assert is_internal_call.get() is False, "aquery should restore the internal-call context"
|
||||
await _drain_logging_worker()
|
||||
|
||||
assert len(recording_logger.success_events) == 1, "aquery should run one async success callback"
|
||||
event: Final = recording_logger.success_events[0]
|
||||
response_obj: Final = event["response_obj"]
|
||||
assert isinstance(response_obj, ModelResponse), "the async callback should receive the completion response"
|
||||
assert response_obj.usage.total_tokens > 0, "the async callback response should include completion usage"
|
||||
|
||||
event_kwargs: Final = cast(Mapping[str, object], event["kwargs"])
|
||||
standard_logging_object: Final = cast(Mapping[str, object], event_kwargs["standard_logging_object"])
|
||||
assert cast(str, standard_logging_object["call_type"]) == "aquery", "the async event should be for aquery"
|
||||
assert cast(int, standard_logging_object["total_tokens"]) > 0, (
|
||||
"the async event should include total completion tokens"
|
||||
)
|
||||
assert cast(int, standard_logging_object["prompt_tokens"]) > 0, "the async event should include prompt tokens"
|
||||
assert cast(int, standard_logging_object["completion_tokens"]) > 0, (
|
||||
"the async event should include completion tokens"
|
||||
)
|
||||
assert cast(float, standard_logging_object["response_cost"]) > 0, "the async event should include completion cost"
|
||||
sync_response_types: Final = [
|
||||
type(event["response_obj"]).__name__ for event in recording_logger.sync_success_events
|
||||
]
|
||||
assert recording_logger.sync_success_events == [], (
|
||||
"an internal sub-call must not log on the parent's logging object; "
|
||||
f"sync success event response types: {sync_response_types}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_aquery_response_hidden_params_carry_completion_cost():
|
||||
"""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue