mirror of
https://github.com/BerriAI/litellm.git
synced 2026-08-28 05:25:59 +00:00
* feat(logging): add opt-in session_id/trace_id correlation to JSON log records via contextvars Adds two ContextVar instances (session_id_var, trace_id_var) to litellm/_logging.py and two setter functions (set_session_id, set_trace_id). Logging.__init__() now calls both setters after assigning litellm_trace_id so every JSON log record emitted within the async request context carries trace_id and, when provided, session_id — enabling log correlation in Loki, CloudWatch Logs Insights, and other structured-log sinks without any changes to individual log call sites. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com> * fix(logging): guard session_id/trace_id injection against overwriting caller-supplied extra fields * fix(logging): always reset session_id_var to empty string when no session_id provided * feat: gate request correlation IDs in logs behind request_correlation_in_logs flag * refactor: move correlation ID injection into CorrelationContextFilter * feat(logging): extend request_correlation_in_logs to plaintext logs and StandardLoggingPayload Plaintext log lines (json_logs off) now get the same trace_id/session_id suffix as JSON logs via a new CorrelationPlainFormatter, so the flag has a visible effect regardless of log format. StandardLoggingPayload gets a new independent session_id field, populated from litellm_session_id. trace_id's existing session_id-first fallback is preserved when request_correlation_in_logs is off; with the flag on, an explicit litellm_trace_id now takes priority over litellm_session_id so the two fields carry genuinely independent values. * fix(logging): restore correlation context after nested calls; sanitize correlation ids Addresses two review findings on this PR. CorrelationContextFilter's trace_id/session_id contextvars were set on every Logging.__init__ but never reset, so a nested LiteLLM call sharing the same asyncio Task as an outer request (e.g. a guardrail's own LLM-as-judge call, an MCP sampling call) would leave the outer request's subsequent log lines stamped with the nested call's ids instead of its own. set_trace_id/ set_session_id now return their contextvars.Token, and Logging stores them and resets both once its own success/failure handler actually completes, via a new idempotent _restore_correlation_context() called from all four terminal handlers. set_trace_id/set_session_id also now strip control characters and bound length before storing a caller-controlled trace_id/session_id, since these values can originate from request input (litellm_session_id, x-litellm- trace-id) and get interpolated into plain-text log lines - without this, a caller could embed \r/\n or escape sequences to forge fake log entries. * fix(logging): restore correlation context after nested calls, not before The previous commit called _restore_correlation_context() as the first line of each terminal handler, before that handler's own callback dispatch loop runs. That's backwards: a nested LiteLLM call triggered from within a callback (e.g. a guardrail's own LLM-as-judge call) would then capture the *already-reset* value as its own pre-call baseline, and its own reset would restore to that instead of the true outer value - verified live to still leak. success_handler/async_success_handler/failure_handler/async_failure_handler are now thin wrappers: the original bodies move to _success_handler_body/etc, called inside a try/finally that restores context only once the full body - including any nested calls its own callback dispatch triggers - has actually finished, mirroring proper stack-scoped nesting semantics. * test(logging): cover async_failure_handler's correlation-context restore Codecov flagged the new async_failure_handler wrapper (try/finally around _async_failure_handler_body) as uncovered - the method had no direct test at all before this PR's refactor split it into a wrapper. Adds a test that awaits it directly and asserts both that async_log_failure_event still fires and that _restore_correlation_context() puts the pre-call trace_id/session_id back. * fix(logging): restore correlation context by value, not by contextvars.Token veria-ai correctly flagged that contextvars.Token.reset() only works in the exact Context it was created in, and litellm's async success path (and streaming failure path) dispatch async_success_handler/async_failure_handler via asyncio.create_task and the global logging worker - a different Context than Logging.__init__ ran in. reset_trace_id/reset_session_id silently swallowed the resulting ValueError, so the restore was a no-op for exactly those paths. Verified independently: reproduced the raw contextvars behavior, then confirmed litellm's async success dispatch really does go through asyncio.create_task + GLOBAL_LOGGING_WORKER (litellm/utils.py). Logging now captures the pre-call *value* (not a Token) and restores via a plain set_trace_id()/set_session_id() call, which works regardless of which Task/Context calls it. reset_trace_id/reset_session_id are removed as dead/unreliable code. Added a regression test that spawns __init__ and the restore in different asyncio Tasks - confirmed it fails against the prior Token-based commit and passes here. * fix(logging): restore correlation context in the originating task too Greptile's re-review correctly identified a remaining gap: for a successful acompletion(), async_success_handler is dispatched via asyncio.create_task + the global logging worker into a *different* Task than the one wrapper_async/Logging.__init__ ran in. The prior fix (43c164a) only restored the handler's own (detached, throwaway) Task - it never touched the originating request Task, which keeps this call's trace_id/ session_id set for the rest of its own execution (e.g. nested calls made via the same Task). wrapper()/wrapper_async() in litellm/utils.py now restore the originating Task's correlation context in a finally block once the whole call is done, regardless of what detached logging tasks it spawned along the way. Since the wrapped body rebinds its own `kwargs` local via function_setup(), sharing the dict object doesn't work here; a small mutable holder carries the constructed Logging instance back out to the outer wrapper instead. _restore_correlation_context() is no longer guarded against repeat calls: with value-based (not Token-based) restoration, each distinct Task that calls it needs its own restore to take effect in that Task's own view of the contextvars, so multiple calls (once per Task involved in an attempt) are required, not just tolerated. Added a regression test using mock_response to exercise the real success dispatch path (asyncio.create_task + GLOBAL_LOGGING_WORKER) without a live provider call, asserting the *test's own* (originating) task context is restored after the call - this is exactly the case Greptile flagged and the prior commit didn't cover. * fix(logging): restore correlation context when function_setup itself fails Greptile's 4th finding: if function_setup() constructs Logging() (whose __init__ already mutates trace_id_var/session_id_var) and then raises before returning - e.g. update_environment_variables() throws - the caller's wrapper()/wrapper_async() never receives a logging_obj reference, so its own restore-on-finally never fires. The correlation ids leak into every subsequent log line on that thread/task with no way to clear them. function_setup()'s own except block now restores the context itself in that case, using whatever logging_obj it managed to construct before failing (locals().get(), safe against the earlier failure modes where logging_obj was never assigned at all). Added a regression test that monkeypatches Logging.update_environment_variables to raise after construction, confirmed it fails without this fix (the leaked ids show up directly in the raised exception's own log line) and passes with it. Broader sweep (test_utils.py, test_router.py, test_main_module_header.py, streaming handler tests, plus all logging-specific tests): 722 passed. * fix(logging): don't assume every litellm_logging_obj is a real Logging instance CI caught a real regression from the last commit: tests/test_litellm/llms/xai/test_xai_key_fallback.py injects a minimal FakeLogging stand-in (only implementing update_from_kwargs) as litellm_logging_obj for a narrow realtime-config unit test, bypassing the real Logging class entirely. wrapper()/ wrapper_async()'s finally block and function_setup()'s except block both unconditionally called _restore_correlation_context() on whatever ended up in the holder, which doesn't exist on that stand-in. _restore_correlation_context is new plumbing specific to this PR's feature, not part of any pre-existing stand-in's expected interface, so callers of it can't assume every object playing the litellm_logging_obj role implements it. Added _restore_correlation_context_if_supported(), a small getattr-guarded helper, and used it at all three call sites. * fix(logging): don't restore context too early on setup failure or streaming Two more findings from Greptile's 5th review round. 1. function_setup()'s except block restored correlation context *after* logging the "Error in function_setup" exception, so that diagnostic log line itself was stamped with the doomed call's ids instead of the outer ids - misleading, since the failed call never produces anything else to attribute those ids to. Restore now happens before the log call. 2. wrapper()/wrapper_async() restored the originating task's context as soon as a streaming call returned, before the caller ever starts iterating the CustomStreamWrapper it just got back. Any log lines emitted while iterating (in the same thread/task) incorrectly showed the pre-call ids instead of this call's own ones. The wrapper finally block now skips the restore when the return value is a stream wrapper, deferring to the terminal handler that already fires once the stream is actually assembled/exhausted. Both verified with tests that fail against the prior commit and pass against this one. Broader sweep unchanged at 829 passing. * fix(logging): best-effort correlation cleanup on abandoned streams Greptile's 7th finding: if a caller returns a streaming response and never fully consumes it - stops iterating early, drops the reference, cancels it - the terminal handler that normally restores the originating task's trace_id/session_id never fires, since it only runs once the stream is actually assembled/exhausted. The ids leak into every subsequent log line in that thread/task with no bound. There's no reliable Python hook for "this was abandoned without being closed" - CustomStreamWrapper has no close()/__aexit__/context-manager convention today, and the only automatic option is __del__, whose timing is inherently unpredictable (delayed by cyclic GC, not guaranteed at interpreter shutdown, can run on a different thread). This is a best-effort safety net, not a guarantee, and is documented as such in the docstring. Testing this via real garbage collection proved unreliable in practice: per-chunk logging submits work to a thread pool executor whose worker thread transiently holds its own bound-method reference to the wrapper until that task completes, so refcount doesn't hit zero on a deterministic schedule even with polling. Tests call __del__ directly instead - a plain method, safe to invoke early - which exercises exactly the restore logic real garbage collection would eventually trigger, plus a case confirming a broken logging_obj can never make __del__ raise. * fix(logging): restore consumer's context at every real stream exit point Two more findings from this round. Veria AI: even a *fully consumed* stream never restored the actual consuming thread/task's correlation context. The terminal success dispatch (dispatch_success_handlers via asyncio.create_task for async, or success_handler via the shared executor for sync) only restores whatever detached context it runs in - never the caller's own thread/task that's running the for/async for loop. Same root cause as the wrapper-level fix two rounds ago, just missed for the streaming-completion path. Greptile: explicit aclose() (client disconnect, router fallback aborting a partial stream) closed the underlying stream without restoring correlation context either, since request wrappers intentionally skip restoration for returned streams and no terminal handler runs on this path. Added CustomStreamWrapper._restore_consumer_correlation_context(), called from every point control genuinely returns to the consumer: the final raise StopIteration/StopAsyncIteration on natural exhaustion (both sync branches, both async branches), _handle_stream_fallback_error (the shared choke point for all three failure-raising call sites), and aclose(). __del__ now delegates to the same helper instead of duplicating it. Verified with tests extending the existing streaming-exhaustion cases to assert the consuming context is restored after the loop completes (fails against the prior commit, passes now), plus a dedicated aclose() test. Broader sweep: 832 passing. * fix(logging): don't let a delayed __del__ finalizer clobber a newer active call If an abandoned stream's __del__ fires late (after cyclic GC delay), a different call may have already taken over the correlation contextvars in the same Task/thread. Restoring unconditionally would stomp that active call's trace_id/session_id with the abandoned stream's stale pre-call snapshot. __del__ now only restores when the contextvars still hold the ids this call itself set. * fix(logging): compare sanitized ids in the __del__ ownership guard set_trace_id()/set_session_id() sanitize (strip control chars, bound length) before storing, so the contextvar's value can differ from the raw litellm_trace_id/litellm_session_id. The __del__ ownership guard was comparing against the raw values, so a caller-supplied id containing control characters or exceeding 256 chars would never match, permanently skipping cleanup. Capture what set_trace_id()/set_session_id() actually stored and compare against that instead. * fix(logging): restore consumer context on the synthesized finish_reason chunk Both __next__ and _finalize_completed_stream() have a branch that fires when the underlying stream ends without ever emitting an explicit finish_reason chunk: they synthesize one via finish_reason_handler() and return it. A consumer that stops as soon as it sees finish_reason - a common pattern - never calls __next__()/__anext__() again, so the existing restore in the sent_last_chunk-is-True StopIteration branch never runs for them. The underlying stream is already exhausted at this point regardless of whether the caller keeps iterating, so restoring here is safe. * fix(logging): don't restore correlation context before the caller receives the final chunk The previous fix (5147c69186) restored context immediately before returning the synthesized finish_reason chunk from __next__/_finalize_completed_stream, reasoning that completion_stream was already exhausted. But that chunk is still this call's own data, and the caller's own application-level log statements processing it run in the same synchronous frame right after the return - restoring first made those lines carry the wrong (outer) ids, exactly what wrapper()/wrapper_async() deliberately avoid by not restoring while a stream is being iterated. Revert to not restoring there. A caller that keeps iterating still gets a correct, deterministic restore on its very next __next__()/__anext__() call (completion_stream is exhausted, so that immediately re-raises StopIteration/StopAsyncIteration through the already-restoring branch). A caller that stops right after finish_reason relies on aclose() or the best-effort __del__ guard, same as any other stream the caller doesn't fully exhaust. * refactor(logging): hoist a safely-hoistable function-body import to module top CorrelationContextFilter.filter()'s `import litellm` was a function-body import; verified it can move to module top without a circular-import failure (litellm/__init__.py already imports from litellm._logging before setting request_correlation_in_logs, but a bare `import litellm` only binds the already-in-sys.modules module object - the attribute itself isn't read until filter() actually runs, by which point litellm is fully initialized). * test(logging): move correlation tests into their conventionally-mapped files tests/test_litellm/ mirrors litellm/ in a parallel path. Correlation tests for the Logging class (litellm_logging.py), function_setup/wrapper_async (utils.py), and CustomStreamWrapper (streaming_handler.py) had all landed in test_logging.py, which only maps to litellm/_logging.py itself. Moving each group to its correctly-mapped file: test_litellm_logging.py (Logging class init/restore), test_utils.py (function_setup, wrapper_async), and test_streaming_handler.py (CustomStreamWrapper) in the next commit. test_logging.py keeps only what actually exercises _logging.py's own contextvars/filters/formatters/sanitization. No behavior change - same assertions, same coverage, just relocated. * fix(logging): restore correlation context unconditionally in wrapper()'s sync path Blocking finding from review: a caller-visible correlation feature was silently misattributing one request's logs to a different, unrelated one on the sync/threaded path. wrapper()/wrapper_async() both left trace_id/session_id "open" across a stream's entire iteration so the caller's own log lines while consuming it would carry the right ids. That's safe for wrapper_async(): each async call gets its own asyncio Task with its own copy of the contextvars, and Tasks are never recycled across requests, so a leftover value can only ever affect that one already-abandoned Task. It is not safe for wrapper() (sync): a plain OS thread has no such per-call isolation, and a thread pool's worker threads *are* recycled across unrelated requests. If a sync stream was abandoned (client disconnect, early break, an uncaught exception) without ever being exhausted or closed, nothing restored its contextvars, and a pool could later hand that same thread to a completely different call, which would inherit the abandoned request's ids as its own "pre-call" baseline and then restore back to that poison when it finished - permanently misattributing every subsequent log line on that thread, including its own, to the abandoned request. Strengthening the __del__ finalizer already added for this can't fix it: finalizer timing is exactly what a permanently-reused thread can't rely on. wrapper() now restores unconditionally in its own finally, before a sync stream is ever handed back to the caller. The trade-off: a sync stream consumer's own application-level log statements while iterating no longer automatically carry this call's ids (litellm's own internal per-chunk logging is unaffected, since it's dispatched separately). That's an acceptable cost for eliminating a silent cross-request misattribution bug. wrapper_async() keeps the existing conditional (skip-if-streaming) behavior, justified by the Task-isolation argument above; CustomStreamWrapper's __del__/aclose()/next-iteration restore machinery remains meaningful and necessary there. This also simplifies wrapper()/wrapper_async() back toward their original shape: both previously used a mutable-dict-holder split into a separate _body function to smuggle logging_obj/result out to an outer finally, working around function_setup() rebinding its own local `kwargs`. That restructuring is no longer needed - `logging_obj` (and, for wrapper_async(), `result`) were already function-level locals in scope for a plain try/finally; three of wrapper_async()'s retry-return statements now assign through `result` first so it accurately reflects what's actually returned even on a retry path. Regression test: test_abandoned_sync_stream_does_not_contaminate_a_later_call_on_the_same_thread in test_streaming_handler.py reproduces the exact reported scenario with a real single-worker ThreadPoolExecutor - confirmed it fails with the prior (skip-restore-on-stream) wrapper() and passes with this fix. * refactor(logging): use Mapping instead of bare dict for read-only params _get_standard_logging_payload_trace_id/_session_id only read litellm_params (.get() calls, no mutation) - annotate it as Mapping[str, Any] rather than a bare mutable dict, per the repo's no-mutable-collection-in-annotation rule. * fix(logging): scope request_correlation_in_logs to the async/proxy path only Blocking review finding: wrapper() (the sync entry point) used the same skip-restore-on-stream design as wrapper_async(), but a plain OS thread has no per-call context isolation the way an asyncio Task does, and a thread pool's worker threads are recycled across unrelated requests - an abandoned sync stream could leave its ids stuck on a thread a pool later hands to a completely different request, misattributing that request's logs. A fix existed and was tested (restore unconditionally in wrapper()'s own finally), but it doesn't benefit this feature's primary consumer - the proxy only ever calls the async entry point - and carries sync-specific complexity this PR doesn't need. Scope the feature to async only instead: Logging.__init__() takes a new supports_correlation_logging parameter (default True), threaded down from a new function_setup(..., is_async_call: bool = True) parameter. wrapper() is the one caller that passes is_async_call=False; every other function_setup() call site (wrapper_async(), the router, and proxy/MCP-internal call sites) is already async and keeps the default. With supports_correlation_logging=False, Logging.__init__() never calls set_trace_id()/set_session_id() at all, so a sync call has nothing to leak in the first place. wrapper() reverts to its pre-review shape with no correlation-specific code at all. StandardLoggingPayload's own trace_id/session_id fields are unaffected either way - they're a deterministic per-call read of self.litellm_trace_id/self.litellm_session_id, not ambient contextvar state, so they were never exposed to the cross-request bug. Full sync/direct-SDK support (stamping + its own safe-restore mechanism) is deferred to a follow-up PR; the fix and its regression test already exist in this branch's history at commit9f3a20f4b2and can be resurrected there. Tests: replaced the two wrapper()-level tests with ones proving the new invariant (sync calls, streaming and non-streaming, never touch trace_id_var/session_id_var even when the caller explicitly passes litellm_trace_id/litellm_session_id), and added a direct unit test for the supports_correlation_logging=False gate on Logging.__init__ itself. Verified live: a real proxy (Postgres-backed, real OpenAI calls) shows clean trace_id/session_id isolation across two concurrent sessions with no cross-contamination; a standalone script confirms real sync SDK calls against a real model never touch the correlation contextvars. * feat(logging): fall back to W3C traceparent/baggage for trace_id/session_id request_correlation_in_logs previously only resolved trace_id/session_id from litellm-specific sources: x-litellm-trace-id/x-litellm-session-id headers, a generic x-<vendor>-session-id header, or Anthropic-style metadata.user_id. If none were present, trace_id fell back to an auto-generated UUID unrelated to anything else, and session_id stayed empty - even when the caller already had real distributed-tracing instrumentation sending the actual industry-standard headers for this. Add a fallback to the W3C Trace Context traceparent header (trace-id component) and W3C Baggage header (session.id entry), so a request already carrying real OpenTelemetry trace context correlates litellm's own logs with the same trace in the caller's observability backend (Datadog, Honeycomb, Tempo, etc.) instead of getting an unrelated generated id. Precedence is unchanged for existing sources: explicit litellm headers and the Anthropic metadata path both still win over this new fallback, which only fires when neither found anything. trace_id and session_id are resolved independently here (unlike the existing chain_id mechanism, which uses one shared value for both), since traceparent and baggage are semantically distinct W3C concepts. New helpers _trace_id_from_traceparent/_session_id_from_baggage in litellm_pre_call_utils.py parse the header formats directly (no new dependency - both are simple fixed-width/delimited strings), wired into LiteLLMProxyRequestSetup.add_litellm_metadata_from_request_headers() only when the corresponding litellm_trace_id/litellm_session_id key isn't already set by the existing paths. Verified live against a real proxy: a bare traceparent header produces a log trace_id exactly matching its trace-id component; a traceparent alongside an explicit x-litellm-trace-id header (different value) produces a log showing the explicit header's value, proving precedence. * fix(logging): reserve trace_id/session_id in JsonFormatter against message-content spoofing JsonFormatter merges keys parsed from the message body before applying extra record attributes, and the extra-attributes loop skips a key that's already present. A caller-controlled log message that happens to parse as JSON/dict with a "trace_id"/"session_id" key (e.g. the proxy logging a raw request-header dict) could therefore make the JSON record carry the attacker-supplied value instead of the real correlation context set via CorrelationContextFilter. trace_id/session_id are now applied from the LogRecord's own attributes after message-content parsing, unconditionally overwriting anything the message body claimed for those two keys. * style(logging): fix import order (ruff I001) in _logging.py and litellm_logging.py - _logging.py: import litellm belongs after the stdlib from-imports, grouped with the other litellm.* imports, not before them. - litellm_logging.py: the refactor to Mapping introduced a second, separate `from collections.abc import Mapping` instead of merging it into the existing `from collections.abc import Callable` import. Caught by the strict-rule budget gate (ruff-strict-budget.json caps I001 at 0 new violations); both auto-fixed with `ruff check --fix --select I001`. * style(logging): freeze mutable-collection constructions flagged by LIT002 Five sites in this PR's diff built a mutable list/dict literal instead of a frozen value: a plain list of optional strings in CorrelationPlainFormatter, a `kwargs or {}` fallback, a `metadata or {}` fallback, two `[...]` candidate orderings, and a `dict(headers)` copy feeding a dict comprehension. Each is build-once/read-only, so this rewrites them as tuples, MappingProxyType, or a plain conditional `.get()` instead of seeding then reading a fresh mutable collection - no behavior change, confirmed by the existing test suite. Caught by the type-discipline budget gate (LIT002 capped at 0 new violations). * fix(logging): reserve trace_id/session_id even when no correlation context is active Live-proxy verification surfaced a gap in the earlier message-content-spoofing fix (7f390a57fc): that fix only overwrites trace_id/session_id from the LogRecord's own attribute, so it does nothing for a log line emitted before CorrelationContextFilter has stamped anything on this record (e.g. the "Request Headers" debug line, which fires before Logging.__init__() runs for the request). On such a record, a caller-supplied header literally named trace_id/session_id still got promoted into the JSON output via the embedded JSON/dict-repr parser, since there was no genuine value to protect. Fixed at the source: trace_id/session_id are now excluded unconditionally from the message-content-parsing promotion step, not just superseded afterward. Verified live against a real proxy - the exact adversarial request (headers literally named trace_id/session_id) no longer leaks into any JSON log record. Added a regression test for this no-active-context variant specifically, confirmed it fails against the prior commit and passes now. Also fixes an unrelated basedpyright regression from an earlier rebase's conflict resolution: litellm/utils.py's `logging_obj` was incorrectly re-annotated `Final` at its second assignment in function_setup() (it's first declared `None` a few lines earlier), which basedpyright correctly rejects. * fix(proxy): stop logging the raw W3C baggage session_id value _session_id_from_baggage() extracts the caller-controlled session.id entry verbatim - it isn't sanitized until set_session_id() runs later in Logging.__init__(). The debug log line for this extraction interpolated the raw value directly, so a caller could embed terminal control characters or ANSI escape sequences that forge/alter plaintext log output for anyone tailing the proxy's logs. Verified live: a baggage header with an embedded ANSI escape reached the terminal as a real, unescaped control sequence before this fix. Drops the value from the log line entirely (the extraction succeeding is enough signal on its own) rather than sanitizing-then-logging, matching veria-ai's suggestion. Added a regression test using caplog that fails against the prior commit and passes now. * fix(logging): restore consumer context only after stream-failure exception mapping _map_anthropic_exception/_map_aleph_alpha_exception synchronously log a debug diagnostic (the raw status code) as part of exception_type()'s mapping. _handle_stream_fallback_error restored the consumer's outer correlation context before calling exception_type(), so that diagnostic log line carried the outer (or empty) trace_id/session_id instead of the failing stream's own - flagged by Greptile. Moved the restore to run after mapping completes, matching the same restore-after-not-before pattern already applied elsewhere in this file for success/finish_reason handling. Added a regression test that captures the correlation context live during a mocked exception_type() call; fails against the prior commit, passes now. * fix(logging): restore consumer context only after aclose()'s stream close completes aclose() restored the consumer's outer correlation context as its first statement, before awaiting the underlying provider stream's own aclose()/ close(). If that close attempt raises, the except branch's debug diagnostic ran under the already-restored outer context instead of the closing stream's own trace_id/session_id - flagged by Greptile, same restore-too-early pattern as the stream-failure fix inf1cf9589d6. Moved the restore to the end of aclose(), after the close attempt (and its diagnostic logging) completes. Added a regression test with a fake stream whose aclose() raises, capturing the correlation context live during the diagnostic log call; fails against the prior commit, passes now. * style(logging): satisfy new strict-lint budgets introduced upstream (Final, ANN401, S110, TRY300, kwargs typing) Rebasing onto litellm_internal_staging pulled in 116 upstream commits that introduced/tightened several lint gates this PR's own code now trips: - LIT010 (every local/module-level variable must be Final): added Final annotations across _logging.py, litellm_logging.py, streaming_handler.py, litellm_pre_call_utils.py, and utils.py. Where a name is genuinely reassigned (logging_obj: starts None, later set to the real object) or branch-assigned, either restructured into a single ternary expression (ordered_candidates) or suppressed with `# rebind-ok: <reason>` matching this repo's documented escape hatch. - LIT011 (parameter mutation): suppressed the two new `data[key] = value` writes in litellm_pre_call_utils.py with `# rebind-ok`, matching the unsuppressed precedent already used for every other `data[...]` write in the same function - `data` is an intentional out-param there. - ANN001/ANN003/ANN202 (missing parameter/return type annotations): fully typed success_handler/_success_handler_body, their async twins, and failure_handler/_failure_handler_body/async variants in litellm_logging.py, plus function_setup in utils.py (added Rules to its existing TYPE_CHECKING block for the rules_obj: Rules annotation). - ANN401 (explicit Any disallowed): suppressed with `# noqa: ANN401` on the handful of genuinely-heterogeneous result/*args/**kwargs parameters, since ordinary suppression is this repo's documented path. - S110 (try/except/pass): added to the existing BLE001 noqa on the one best-effort correlation-cleanup try/except this PR added. - TRY300 (return inside try): moved two `return result` statements into `else:` blocks in the retry-fallback paths this PR's own diff touched. - reportPrivateUsage (basedpyright): renamed the two new StandardLoggingPayloadSetup static methods (get_standard_logging_payload_ trace_id/session_id) to drop their leading underscore, since they're genuinely called from a sibling module-level function in the same file. No behavior change - confirmed by the full existing test suite (819 passed) plus all four lint gates (ruff format, ruff-strict, type-discipline, basedpyright) passing clean. * fix(lint): stop RUF100 flagging noqa suppressions the strict gate needs CI's plain "ruff check" job uses the default ruff.toml, a narrower config than ruff-strict.toml (used only by the strict-rule budget gate). ANN401 and S110 aren't enabled in the default config, so RUF100 (unused-noqa) flagged the `# noqa: ANN401`/`# noqa: ...,S110` suppressions this PR added as pointless under that config, even though they're genuinely needed under ruff-strict.toml. - ANN401: added to ruff.toml's existing `lint.external` list (same mechanism already used for C901/TID251, enforced by the strict gate but not by this config) - these Any usages are genuinely dynamic/forwarded, so the suppression itself is correct and just needed registering. - S110: fixed the underlying code instead of registering another external code - the try/except/pass in CustomStreamWrapper._restore_consumer_correlation_context now logs at debug level on failure (matching the existing best-effort-cleanup pattern in _record_partial_usage_for_failure elsewhere in this file), which satisfies S110's own suggestion directly and needs no suppression at all. Verified against both ruff.toml and ruff-strict.toml directly, plus all three other gates (ruff format, type-discipline, basedpyright) and the full test suite (821 passed). * fix(lint): scope the ANN401 exemption to file level instead of a repo-wide noqa Ruff has no per-line-scoped way to register a noqa code across configs (that requires the default ruff.toml's lint.external list, which is repo-wide in scope even though the noqa itself is per-line). Since ruff does support file-level exemptions via per-file-ignores, and ANN401 only needed exempting in exactly two files, moved the exemption there instead: - ruff-strict.toml: added [lint.per-file-ignores] disabling ANN401 for litellm_logging.py and utils.py specifically, with a comment explaining why (heterogeneous response/forwarded-args parameters with no fitting concrete type - already verified by trying CostResponseTypes and hitting a real basedpyright mismatch). - ruff.toml: reverted the ANN401 entry from lint.external - no longer needed, since there's no `# noqa: ANN401` left anywhere for RUF100 to second-guess. - Removed the now-redundant `# noqa: ANN401` from the 10 affected parameters in both files, keeping the existing kwargs-ok reasons and adding a short inline comment on the `result`/`*args` lines pointing at the ruff-strict.toml exemption for context. Verified against both configs directly (ANN401 clean under ruff-strict.toml for these files, RUF100 clean under the default config), all four gates (ruff format, ruff-strict, type-discipline, basedpyright), and the full test suite (821 passed). * fix(logging): redact credential-shaped trace_id/session_id before stamping log records CorrelationContextFilter stamps trace_id/session_id onto a LogRecord after SecretRedactionFilter has already run, so a caller-controlled value (e.g. via x-litellm-trace-id or a W3C baggage header) that happens to look like a real credential reached JSON and plaintext logs unredacted. Apply the same credential redaction already used elsewhere in this module at _sanitize_correlation_id(), the single choke point both set_trace_id() and set_session_id() route through, so every caller-facing entry point is covered without depending on filter ordering. * fix(logging): restore correlation context when a stream's max-duration timeout fires CustomStreamWrapper.__anext__() called _check_max_streaming_duration() before entering its try block, so the litellm.Timeout it raises bypassed the except Exception -> _handle_stream_fallback_error path entirely, leaking the timed-out stream's own trace_id/session_id into whatever the consumer's task logs next. Move the check inside the try so it flows through the same restoration path every other stream failure already uses. * test(streaming): make dispatch_failure_handlers mock awaitable for the async max-duration test Moving _check_max_streaming_duration() inside __anext__()'s try block (prior commit) means a max-duration Timeout now dispatches failure handlers through the same path every other stream failure already uses, instead of bypassing it entirely. dispatch_failure_handlers is async on the real Logging class; the test's plain MagicMock logging_obj made asyncio.create_task() choke on a non-coroutine return value once that path actually got exercised. --------- Co-authored-by: Deepanshu <deepanshu.lulla@alpha-sense.com> Co-authored-by: Claude Sonnet 4.6 <noreply@anthropic.com>
1272 lines
44 KiB
Python
1272 lines
44 KiB
Python
"""
|
|
Unit tests for StandardLoggingPayloadSetup
|
|
"""
|
|
|
|
import json
|
|
import os
|
|
import sys
|
|
from datetime import datetime
|
|
from unittest.mock import AsyncMock
|
|
|
|
sys.path.insert(
|
|
0, os.path.abspath("../..")
|
|
) # Adds the parent directory to the system-path
|
|
from datetime import datetime as dt_object
|
|
import time
|
|
import pytest
|
|
import litellm
|
|
from litellm.types.utils import (
|
|
StandardLoggingPayload,
|
|
Usage,
|
|
StandardLoggingMetadata,
|
|
StandardLoggingModelInformation,
|
|
StandardLoggingHiddenParams,
|
|
)
|
|
from create_mock_standard_logging_payload import (
|
|
create_standard_logging_payload,
|
|
create_standard_logging_payload_with_long_content,
|
|
)
|
|
from litellm.litellm_core_utils.litellm_logging import (
|
|
StandardLoggingPayloadSetup,
|
|
)
|
|
|
|
from litellm.integrations.custom_logger import CustomLogger
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"response_obj,expected_values",
|
|
[
|
|
# Test None input
|
|
(None, (0, 0, 0)),
|
|
# Test empty dict
|
|
({}, (0, 0, 0)),
|
|
# Test valid usage dict
|
|
(
|
|
{
|
|
"usage": {
|
|
"prompt_tokens": 10,
|
|
"completion_tokens": 20,
|
|
"total_tokens": 30,
|
|
}
|
|
},
|
|
(10, 20, 30),
|
|
),
|
|
# Test with litellm.Usage object
|
|
(
|
|
{"usage": Usage(prompt_tokens=15, completion_tokens=25, total_tokens=40)},
|
|
(15, 25, 40),
|
|
),
|
|
# Test invalid usage type
|
|
({"usage": "invalid"}, (0, 0, 0)),
|
|
# Test None usage
|
|
({"usage": None}, (0, 0, 0)),
|
|
],
|
|
)
|
|
def test_get_usage(response_obj, expected_values):
|
|
"""
|
|
Make sure values returned from get_usage are always integers
|
|
"""
|
|
|
|
usage = StandardLoggingPayloadSetup.get_usage_from_response_obj(response_obj)
|
|
|
|
# Check types
|
|
assert isinstance(usage.prompt_tokens, int)
|
|
assert isinstance(usage.completion_tokens, int)
|
|
assert isinstance(usage.total_tokens, int)
|
|
|
|
# Check values
|
|
assert usage.prompt_tokens == expected_values[0]
|
|
assert usage.completion_tokens == expected_values[1]
|
|
assert usage.total_tokens == expected_values[2]
|
|
|
|
|
|
def test_get_usage_from_image_generation_response():
|
|
"""
|
|
Test that image generation usage (with input_tokens/output_tokens format)
|
|
is correctly transformed to standard usage format with image_tokens preserved.
|
|
|
|
Note: get_usage_from_response_obj() is used by multiple endpoints including
|
|
/images/generations and Response API (/responses), both of which use the
|
|
input_tokens/output_tokens format instead of prompt_tokens/completion_tokens.
|
|
|
|
This tests the fix for the bug where image_tokens were being lost during
|
|
spend log creation for /images/generations endpoint.
|
|
"""
|
|
# Simulating image generation response usage from OpenAI
|
|
response_obj = {
|
|
"usage": {
|
|
"input_tokens": 13,
|
|
"output_tokens": 372,
|
|
"total_tokens": 385,
|
|
"input_tokens_details": {
|
|
"image_tokens": 0,
|
|
"text_tokens": 13,
|
|
},
|
|
"output_tokens_details": {
|
|
"image_tokens": 272,
|
|
"text_tokens": 100,
|
|
},
|
|
}
|
|
}
|
|
|
|
usage = StandardLoggingPayloadSetup.get_usage_from_response_obj(response_obj)
|
|
|
|
# Check basic token counts are mapped correctly
|
|
assert usage.prompt_tokens == 13
|
|
assert usage.completion_tokens == 372
|
|
assert usage.total_tokens == 385
|
|
|
|
# Check that prompt_tokens_details contains image_tokens and text_tokens
|
|
assert usage.prompt_tokens_details is not None
|
|
assert usage.prompt_tokens_details.image_tokens == 0
|
|
assert usage.prompt_tokens_details.text_tokens == 13
|
|
|
|
# Check that completion_tokens_details contains image_tokens and text_tokens
|
|
assert usage.completion_tokens_details is not None
|
|
assert usage.completion_tokens_details.image_tokens == 272
|
|
assert usage.completion_tokens_details.text_tokens == 100
|
|
|
|
|
|
def test_get_additional_headers():
|
|
additional_headers = {
|
|
"x-ratelimit-limit-requests": "2000",
|
|
"x-ratelimit-remaining-requests": "1999",
|
|
"x-ratelimit-limit-tokens": "160000",
|
|
"x-ratelimit-remaining-tokens": "160000",
|
|
"llm_provider-date": "Tue, 29 Oct 2024 23:57:37 GMT",
|
|
"llm_provider-content-type": "application/json",
|
|
"llm_provider-transfer-encoding": "chunked",
|
|
"llm_provider-connection": "keep-alive",
|
|
"llm_provider-anthropic-ratelimit-requests-limit": "2000",
|
|
"llm_provider-anthropic-ratelimit-requests-remaining": "1999",
|
|
"llm_provider-anthropic-ratelimit-requests-reset": "2024-10-29T23:57:40Z",
|
|
"llm_provider-anthropic-ratelimit-tokens-limit": "160000",
|
|
"llm_provider-anthropic-ratelimit-tokens-remaining": "160000",
|
|
"llm_provider-anthropic-ratelimit-tokens-reset": "2024-10-29T23:57:36Z",
|
|
"llm_provider-request-id": "req_01F6CycZZPSHKRCCctcS1Vto",
|
|
"llm_provider-via": "1.1 google",
|
|
"llm_provider-cf-cache-status": "DYNAMIC",
|
|
"llm_provider-x-robots-tag": "none",
|
|
"llm_provider-server": "cloudflare",
|
|
"llm_provider-cf-ray": "8da71bdbc9b57abb-SJC",
|
|
"llm_provider-content-encoding": "gzip",
|
|
"llm_provider-x-ratelimit-limit-requests": "2000",
|
|
"llm_provider-x-ratelimit-remaining-requests": "1999",
|
|
"llm_provider-x-ratelimit-limit-tokens": "160000",
|
|
"llm_provider-x-ratelimit-remaining-tokens": "160000",
|
|
}
|
|
additional_logging_headers = StandardLoggingPayloadSetup.get_additional_headers(
|
|
additional_headers
|
|
)
|
|
# Typed rate-limit fields are coerced to int
|
|
assert additional_logging_headers is not None
|
|
assert additional_logging_headers.get("x_ratelimit_limit_requests") == 2000
|
|
assert additional_logging_headers.get("x_ratelimit_remaining_requests") == 1999
|
|
assert additional_logging_headers.get("x_ratelimit_limit_tokens") == 160000
|
|
assert additional_logging_headers.get("x_ratelimit_remaining_tokens") == 160000
|
|
# Provider-specific headers are preserved verbatim (not dropped)
|
|
assert (
|
|
additional_logging_headers.get("llm_provider-request-id")
|
|
== "req_01F6CycZZPSHKRCCctcS1Vto"
|
|
)
|
|
assert (
|
|
additional_logging_headers.get(
|
|
"llm_provider-anthropic-ratelimit-requests-reset"
|
|
)
|
|
== "2024-10-29T23:57:40Z"
|
|
)
|
|
|
|
|
|
def all_fields_present(standard_logging_metadata: StandardLoggingMetadata):
|
|
for field in StandardLoggingMetadata.__annotations__.keys():
|
|
assert field in standard_logging_metadata
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"metadata_key, metadata_value",
|
|
[
|
|
("user_api_key_alias", "test_alias"),
|
|
("user_api_key_hash", "test_hash"),
|
|
("user_api_key_team_id", "test_team_id"),
|
|
("user_api_key_user_id", "test_user_id"),
|
|
("user_api_key_team_alias", "test_team_alias"),
|
|
("user_api_key_spend", 10.50),
|
|
("spend_logs_metadata", {"key": "value"}),
|
|
("requester_ip_address", "127.0.0.1"),
|
|
("requester_metadata", {"user_agent": "test_agent"}),
|
|
],
|
|
)
|
|
def test_get_standard_logging_metadata(metadata_key, metadata_value):
|
|
"""
|
|
Test that the get_standard_logging_metadata function correctly sets the metadata fields.
|
|
All fields in StandardLoggingMetadata should ALWAYS be present.
|
|
"""
|
|
metadata = {metadata_key: metadata_value}
|
|
standard_logging_metadata = (
|
|
StandardLoggingPayloadSetup.get_standard_logging_metadata(metadata)
|
|
)
|
|
|
|
print("standard_logging_metadata", standard_logging_metadata)
|
|
|
|
# Assert that all fields in StandardLoggingMetadata are present
|
|
all_fields_present(standard_logging_metadata)
|
|
|
|
# Assert that the specific metadata field is set correctly
|
|
assert standard_logging_metadata[metadata_key] == metadata_value
|
|
|
|
|
|
def test_get_standard_logging_metadata_user_api_key_hash():
|
|
valid_hash = "a" * 64 # 64 character string
|
|
metadata = {"user_api_key": valid_hash}
|
|
result = StandardLoggingPayloadSetup.get_standard_logging_metadata(metadata)
|
|
assert result["user_api_key_hash"] == valid_hash
|
|
|
|
|
|
def test_get_standard_logging_metadata_invalid_user_api_key():
|
|
invalid_hash = "not_a_valid_hash"
|
|
metadata = {"user_api_key": invalid_hash}
|
|
result = StandardLoggingPayloadSetup.get_standard_logging_metadata(metadata)
|
|
all_fields_present(result)
|
|
assert result["user_api_key_hash"] is None
|
|
|
|
|
|
def test_get_standard_logging_metadata_non_string_user_api_key():
|
|
"""Non-string user_api_key should not be set as user_api_key_hash."""
|
|
metadata = {"user_api_key": 12345}
|
|
result = StandardLoggingPayloadSetup.get_standard_logging_metadata(metadata)
|
|
all_fields_present(result)
|
|
assert result["user_api_key_hash"] is None
|
|
|
|
|
|
def test_get_standard_logging_metadata_none_user_api_key():
|
|
"""None user_api_key should not be set as user_api_key_hash."""
|
|
metadata = {"user_api_key": None}
|
|
result = StandardLoggingPayloadSetup.get_standard_logging_metadata(metadata)
|
|
all_fields_present(result)
|
|
assert result["user_api_key_hash"] is None
|
|
|
|
|
|
def test_get_standard_logging_metadata_invalid_keys():
|
|
metadata = {
|
|
"user_api_key_alias": "test_alias",
|
|
"invalid_key": "should_be_ignored",
|
|
"another_invalid_key": 123,
|
|
}
|
|
result = StandardLoggingPayloadSetup.get_standard_logging_metadata(metadata)
|
|
all_fields_present(result)
|
|
assert result["user_api_key_alias"] == "test_alias"
|
|
assert "invalid_key" not in result
|
|
assert "another_invalid_key" not in result
|
|
|
|
|
|
def test_cleanup_timestamps():
|
|
"""Test cleanup_timestamps with different input types"""
|
|
# Test with datetime objects
|
|
now = dt_object.now()
|
|
start = now
|
|
end = now
|
|
completion = now
|
|
|
|
result = StandardLoggingPayloadSetup.cleanup_timestamps(start, end, completion)
|
|
|
|
assert all(isinstance(x, float) for x in result)
|
|
assert len(result) == 3
|
|
|
|
# Test with float timestamps
|
|
start_float = time.time()
|
|
end_float = start_float + 1
|
|
completion_float = end_float
|
|
|
|
result = StandardLoggingPayloadSetup.cleanup_timestamps(
|
|
start_float, end_float, completion_float
|
|
)
|
|
|
|
assert all(isinstance(x, float) for x in result)
|
|
assert result[0] == start_float
|
|
assert result[1] == end_float
|
|
assert result[2] == completion_float
|
|
|
|
# Test with mixed types
|
|
result = StandardLoggingPayloadSetup.cleanup_timestamps(
|
|
start_float, end, completion_float
|
|
)
|
|
assert all(isinstance(x, float) for x in result)
|
|
|
|
# Test invalid input
|
|
with pytest.raises(ValueError):
|
|
StandardLoggingPayloadSetup.cleanup_timestamps(
|
|
"invalid", end_float, completion_float
|
|
)
|
|
|
|
|
|
def test_get_model_cost_information():
|
|
"""Test get_model_cost_information with different inputs"""
|
|
# Test with None values
|
|
result = StandardLoggingPayloadSetup.get_model_cost_information(
|
|
base_model=None,
|
|
custom_pricing=None,
|
|
custom_llm_provider=None,
|
|
init_response_obj={},
|
|
)
|
|
assert result["model_map_key"] == ""
|
|
assert result["model_map_value"] is None # this was not found in model cost map
|
|
# assert all fields in StandardLoggingModelInformation are present
|
|
assert all(
|
|
field in result for field in StandardLoggingModelInformation.__annotations__
|
|
)
|
|
|
|
# Test with valid model
|
|
result = StandardLoggingPayloadSetup.get_model_cost_information(
|
|
base_model="gpt-5-mini",
|
|
custom_pricing=False,
|
|
custom_llm_provider="openai",
|
|
init_response_obj={},
|
|
)
|
|
litellm_info_gpt_3_5_turbo_model_map_value = litellm.get_model_info(
|
|
model="gpt-5-mini", custom_llm_provider="openai"
|
|
)
|
|
print("result", result)
|
|
assert result["model_map_key"] == "gpt-5-mini"
|
|
assert result["model_map_value"] is not None
|
|
assert result["model_map_value"] == litellm_info_gpt_3_5_turbo_model_map_value
|
|
# assert all fields in StandardLoggingModelInformation are present
|
|
assert all(
|
|
field in result for field in StandardLoggingModelInformation.__annotations__
|
|
)
|
|
|
|
|
|
def test_get_model_cost_information_custom_pricing_uses_base_model():
|
|
result = StandardLoggingPayloadSetup.get_model_cost_information(
|
|
base_model="bedrock/invoke/global.anthropic.claude-opus-4-6-v1",
|
|
custom_pricing=True,
|
|
custom_llm_provider="bedrock",
|
|
init_response_obj={"model": "invoke_test_claude"},
|
|
)
|
|
assert result["model_map_value"] is not None
|
|
assert result["model_map_key"] != "invoke_test_claude"
|
|
|
|
|
|
def test_standard_logging_payload_uses_deployment_when_no_base_model():
|
|
"""metadata["deployment"] is used for cost-map lookup when base_model is not set."""
|
|
from datetime import datetime
|
|
|
|
from litellm.litellm_core_utils.litellm_logging import (
|
|
Logging,
|
|
get_standard_logging_object_payload,
|
|
)
|
|
|
|
logging_obj = Logging(
|
|
model="invoke_test_claude",
|
|
messages=[{"role": "user", "content": "hi"}],
|
|
stream=False,
|
|
call_type="completion",
|
|
start_time=datetime.now(),
|
|
litellm_call_id="test-deploy-fallback",
|
|
function_id="test-fn",
|
|
)
|
|
|
|
kwargs = {
|
|
"model": "invoke_test_claude",
|
|
"messages": [{"role": "user", "content": "hi"}],
|
|
"custom_llm_provider": "bedrock",
|
|
"litellm_params": {
|
|
"metadata": {
|
|
"deployment": "bedrock/invoke/global.anthropic.claude-opus-4-6-v1",
|
|
},
|
|
},
|
|
}
|
|
mock_response = {
|
|
"id": "chatcmpl-deploy-test",
|
|
"object": "chat.completion",
|
|
"model": "invoke_test_claude",
|
|
"usage": {"prompt_tokens": 5, "completion_tokens": 10, "total_tokens": 15},
|
|
"choices": [
|
|
{
|
|
"index": 0,
|
|
"message": {"role": "assistant", "content": "hello"},
|
|
"finish_reason": "stop",
|
|
}
|
|
],
|
|
}
|
|
|
|
payload = get_standard_logging_object_payload(
|
|
kwargs=kwargs,
|
|
init_response_obj=mock_response,
|
|
start_time=datetime.now(),
|
|
end_time=datetime.now(),
|
|
logging_obj=logging_obj,
|
|
status="success",
|
|
)
|
|
|
|
assert payload is not None
|
|
assert payload["model_map_information"]["model_map_value"] is not None
|
|
assert payload["model_map_information"]["model_map_key"] != "invoke_test_claude"
|
|
|
|
|
|
def test_get_hidden_params():
|
|
"""Test get_hidden_params with different inputs"""
|
|
# Test with None
|
|
result = StandardLoggingPayloadSetup.get_hidden_params(None)
|
|
assert result["model_id"] is None
|
|
assert result["cache_key"] is None
|
|
assert result["api_base"] is None
|
|
assert result["response_cost"] is None
|
|
assert result["additional_headers"] is None
|
|
|
|
# assert all fields in StandardLoggingHiddenParams are present
|
|
assert all(field in result for field in StandardLoggingHiddenParams.__annotations__)
|
|
|
|
# Test with valid params
|
|
hidden_params = {
|
|
"model_id": "test-model",
|
|
"cache_key": "test-cache",
|
|
"api_base": "https://api.test.com",
|
|
"response_cost": 0.001,
|
|
"additional_headers": {
|
|
"x-ratelimit-limit-requests": "2000",
|
|
"x-ratelimit-remaining-requests": "1999",
|
|
},
|
|
}
|
|
result = StandardLoggingPayloadSetup.get_hidden_params(hidden_params)
|
|
assert result["model_id"] == "test-model"
|
|
assert result["cache_key"] == "test-cache"
|
|
assert result["api_base"] == "https://api.test.com"
|
|
assert result["response_cost"] == 0.001
|
|
assert result["additional_headers"] is not None
|
|
assert result["additional_headers"]["x_ratelimit_limit_requests"] == 2000
|
|
# assert all fields in StandardLoggingHiddenParams are present
|
|
assert all(field in result for field in StandardLoggingHiddenParams.__annotations__)
|
|
|
|
|
|
def test_get_final_response_obj():
|
|
"""Test get_final_response_obj with different input types and redaction scenarios"""
|
|
# Test with direct response_obj
|
|
response_obj = {"choices": [{"message": {"content": "test content"}}]}
|
|
result = StandardLoggingPayloadSetup.get_final_response_obj(
|
|
response_obj=response_obj, init_response_obj=None, kwargs={}
|
|
)
|
|
assert result == response_obj
|
|
|
|
# Test redaction when litellm.turn_off_message_logging is True
|
|
litellm.turn_off_message_logging = True
|
|
try:
|
|
model_response = litellm.ModelResponse(
|
|
choices=[
|
|
litellm.Choices(message=litellm.Message(content="sensitive content"))
|
|
]
|
|
)
|
|
kwargs = {"messages": [{"role": "user", "content": "original message"}]}
|
|
result = StandardLoggingPayloadSetup.get_final_response_obj(
|
|
response_obj=model_response, init_response_obj=model_response, kwargs=kwargs
|
|
)
|
|
|
|
print("result", result)
|
|
print("type(result)", type(result))
|
|
# Verify response message content was redacted
|
|
assert result["choices"][0]["message"]["content"] == "redacted-by-litellm"
|
|
# Verify that redaction occurred in kwargs
|
|
assert kwargs["messages"][0]["content"] == "redacted-by-litellm"
|
|
finally:
|
|
# Reset litellm.turn_off_message_logging to its original value
|
|
litellm.turn_off_message_logging = False
|
|
|
|
|
|
def testget_standard_logging_payload_trace_id():
|
|
"""Test get_standard_logging_payload_trace_id with different input scenarios"""
|
|
# Test case 1: When litellm_trace_id is provided in litellm_params
|
|
from unittest.mock import MagicMock
|
|
|
|
# Create a mock Logging object
|
|
mock_logging_obj = MagicMock()
|
|
mock_logging_obj.litellm_trace_id = "default-trace-id"
|
|
|
|
# Test when litellm_trace_id is in litellm_params
|
|
litellm_params = {"litellm_trace_id": "dynamic-trace-id"}
|
|
result = StandardLoggingPayloadSetup.get_standard_logging_payload_trace_id(
|
|
logging_obj=mock_logging_obj, litellm_params=litellm_params
|
|
)
|
|
assert result == "dynamic-trace-id"
|
|
|
|
# Test case 2: When litellm_trace_id is not provided in litellm_params
|
|
litellm_params = {}
|
|
result = StandardLoggingPayloadSetup.get_standard_logging_payload_trace_id(
|
|
logging_obj=mock_logging_obj, litellm_params=litellm_params
|
|
)
|
|
assert result == "default-trace-id"
|
|
|
|
# Test case 3: When litellm_params is None
|
|
result = StandardLoggingPayloadSetup.get_standard_logging_payload_trace_id(
|
|
logging_obj=mock_logging_obj, litellm_params={}
|
|
)
|
|
assert result == "default-trace-id"
|
|
|
|
# Test case 4: When litellm_trace_id in params is not a string
|
|
litellm_params = {"litellm_trace_id": 12345}
|
|
result = StandardLoggingPayloadSetup.get_standard_logging_payload_trace_id(
|
|
logging_obj=mock_logging_obj, litellm_params=litellm_params
|
|
)
|
|
assert result == "12345"
|
|
assert isinstance(result, str)
|
|
|
|
|
|
def testget_standard_logging_payload_trace_id_prioritizes_trace_id_when_flag_on(monkeypatch):
|
|
"""With request_correlation_in_logs on, an explicit litellm_trace_id wins over litellm_session_id."""
|
|
from unittest.mock import MagicMock
|
|
|
|
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
|
mock_logging_obj = MagicMock()
|
|
mock_logging_obj.litellm_trace_id = "default-trace-id"
|
|
|
|
litellm_params = {"litellm_trace_id": "the-trace-id", "litellm_session_id": "the-session-id"}
|
|
result = StandardLoggingPayloadSetup.get_standard_logging_payload_trace_id(
|
|
logging_obj=mock_logging_obj, litellm_params=litellm_params
|
|
)
|
|
assert result == "the-trace-id"
|
|
|
|
|
|
def testget_standard_logging_payload_trace_id_prioritizes_session_id_when_flag_off(monkeypatch):
|
|
"""With request_correlation_in_logs off (default), legacy behavior is preserved:
|
|
litellm_session_id still wins over litellm_trace_id."""
|
|
from unittest.mock import MagicMock
|
|
|
|
monkeypatch.setattr(litellm, "request_correlation_in_logs", False)
|
|
mock_logging_obj = MagicMock()
|
|
mock_logging_obj.litellm_trace_id = "default-trace-id"
|
|
|
|
litellm_params = {"litellm_trace_id": "the-trace-id", "litellm_session_id": "the-session-id"}
|
|
result = StandardLoggingPayloadSetup.get_standard_logging_payload_trace_id(
|
|
logging_obj=mock_logging_obj, litellm_params=litellm_params
|
|
)
|
|
assert result == "the-session-id"
|
|
|
|
|
|
def testget_standard_logging_payload_session_id_when_flag_on(monkeypatch):
|
|
"""Test get_standard_logging_payload_session_id with different input scenarios, flag enabled"""
|
|
from unittest.mock import MagicMock
|
|
|
|
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
|
|
mock_logging_obj = MagicMock()
|
|
mock_logging_obj.litellm_session_id = ""
|
|
|
|
# Test case 1: litellm_session_id provided directly in litellm_params
|
|
litellm_params = {"litellm_session_id": "dynamic-session-id"}
|
|
result = StandardLoggingPayloadSetup.get_standard_logging_payload_session_id(
|
|
logging_obj=mock_logging_obj, litellm_params=litellm_params
|
|
)
|
|
assert result == "dynamic-session-id"
|
|
|
|
# Test case 2: falls back to metadata.session_id when not in litellm_params directly
|
|
litellm_params = {"metadata": {"session_id": "metadata-session-id"}}
|
|
result = StandardLoggingPayloadSetup.get_standard_logging_payload_session_id(
|
|
logging_obj=mock_logging_obj, litellm_params=litellm_params
|
|
)
|
|
assert result == "metadata-session-id"
|
|
|
|
# Test case 3: falls back to logging_obj.litellm_session_id when nothing else is set
|
|
mock_logging_obj.litellm_session_id = "obj-session-id"
|
|
result = StandardLoggingPayloadSetup.get_standard_logging_payload_session_id(
|
|
logging_obj=mock_logging_obj, litellm_params={}
|
|
)
|
|
assert result == "obj-session-id"
|
|
|
|
# Test case 4: empty string when no session id was supplied anywhere
|
|
mock_logging_obj.litellm_session_id = ""
|
|
result = StandardLoggingPayloadSetup.get_standard_logging_payload_session_id(
|
|
logging_obj=mock_logging_obj, litellm_params={}
|
|
)
|
|
assert result == ""
|
|
|
|
# Test case 5: non-string session id in params is coerced to str
|
|
litellm_params = {"litellm_session_id": 98765}
|
|
result = StandardLoggingPayloadSetup.get_standard_logging_payload_session_id(
|
|
logging_obj=mock_logging_obj, litellm_params=litellm_params
|
|
)
|
|
assert result == "98765"
|
|
assert isinstance(result, str)
|
|
|
|
# Test case 6: trace_id and session_id are independent - passing only a trace id
|
|
# must not populate session_id
|
|
litellm_params = {"litellm_trace_id": "some-trace-id"}
|
|
result = StandardLoggingPayloadSetup.get_standard_logging_payload_session_id(
|
|
logging_obj=mock_logging_obj, litellm_params=litellm_params
|
|
)
|
|
assert result == ""
|
|
|
|
|
|
def testget_standard_logging_payload_session_id_empty_when_flag_off(monkeypatch):
|
|
"""When request_correlation_in_logs is off (default), session_id is always empty,
|
|
even if litellm_session_id was explicitly supplied - preserves the pre-existing
|
|
StandardLoggingPayload shape for callers who haven't opted in."""
|
|
from unittest.mock import MagicMock
|
|
|
|
monkeypatch.setattr(litellm, "request_correlation_in_logs", False)
|
|
mock_logging_obj = MagicMock()
|
|
mock_logging_obj.litellm_session_id = "obj-session-id"
|
|
|
|
litellm_params = {"litellm_session_id": "dynamic-session-id"}
|
|
result = StandardLoggingPayloadSetup.get_standard_logging_payload_session_id(
|
|
logging_obj=mock_logging_obj, litellm_params=litellm_params
|
|
)
|
|
assert result == ""
|
|
|
|
|
|
def test_truncate_standard_logging_payload():
|
|
"""
|
|
1. original messages, response, and error_str should NOT BE MODIFIED, since these are from kwargs
|
|
2. the `messages`, `response`, and `error_str` in new standard_logging_payload should be truncated
|
|
"""
|
|
_custom_logger = CustomLogger()
|
|
standard_logging_payload: StandardLoggingPayload = (
|
|
create_standard_logging_payload_with_long_content()
|
|
)
|
|
original_messages = standard_logging_payload["messages"]
|
|
len_original_messages = len(str(original_messages))
|
|
original_response = standard_logging_payload["response"]
|
|
len_original_response = len(str(original_response))
|
|
original_error_str = standard_logging_payload["error_str"]
|
|
len_original_error_str = len(str(original_error_str))
|
|
|
|
_custom_logger.truncate_standard_logging_payload_content(standard_logging_payload)
|
|
|
|
# Original messages, response, and error_str should NOT BE MODIFIED
|
|
assert standard_logging_payload["messages"] != original_messages
|
|
assert standard_logging_payload["response"] != original_response
|
|
assert standard_logging_payload["error_str"] != original_error_str
|
|
assert len_original_messages == len(str(original_messages))
|
|
assert len_original_response == len(str(original_response))
|
|
assert len_original_error_str == len(str(original_error_str))
|
|
|
|
print(
|
|
"logged standard_logging_payload",
|
|
json.dumps(standard_logging_payload, indent=2),
|
|
)
|
|
|
|
# Logged messages, response, and error_str should be truncated
|
|
# assert len of messages is less than 10_500
|
|
assert len(str(standard_logging_payload["messages"])) < 10_500
|
|
# assert len of response is less than 10_500
|
|
assert len(str(standard_logging_payload["response"])) < 10_500
|
|
# assert len of error_str is less than 10_500
|
|
assert len(str(standard_logging_payload["error_str"])) < 10_500
|
|
|
|
|
|
def test_strip_trailing_slash():
|
|
common_api_base = "https://api.test.com"
|
|
assert (
|
|
StandardLoggingPayloadSetup.strip_trailing_slash(common_api_base + "/")
|
|
== common_api_base
|
|
)
|
|
assert (
|
|
StandardLoggingPayloadSetup.strip_trailing_slash(common_api_base)
|
|
== common_api_base
|
|
)
|
|
|
|
|
|
def test_get_error_information():
|
|
"""Test get_error_information with different types of exceptions"""
|
|
|
|
# Test with None
|
|
result = StandardLoggingPayloadSetup.get_error_information(None)
|
|
print("error_information", json.dumps(result, indent=2))
|
|
assert result["error_code"] == ""
|
|
assert result["error_class"] == ""
|
|
assert result["llm_provider"] == ""
|
|
|
|
# Test with a basic Exception
|
|
basic_exception = Exception("Test error")
|
|
result = StandardLoggingPayloadSetup.get_error_information(basic_exception)
|
|
print("error_information", json.dumps(result, indent=2))
|
|
assert result["error_code"] == ""
|
|
assert result["error_class"] == "Exception"
|
|
assert result["llm_provider"] == ""
|
|
|
|
# Test with litellm exception from provider
|
|
litellm_exception = litellm.exceptions.RateLimitError(
|
|
message="Test error",
|
|
llm_provider="openai",
|
|
model="gpt-5-mini",
|
|
response=None,
|
|
litellm_debug_info=None,
|
|
max_retries=None,
|
|
num_retries=None,
|
|
)
|
|
result = StandardLoggingPayloadSetup.get_error_information(litellm_exception)
|
|
print("error_information", json.dumps(result, indent=2))
|
|
assert result["error_code"] == "429"
|
|
assert result["error_class"] == "RateLimitError"
|
|
assert result["llm_provider"] == "openai"
|
|
assert result["error_message"] == "litellm.RateLimitError: Test error"
|
|
|
|
|
|
def test_get_response_time():
|
|
"""Test get_response_time with different streaming scenarios"""
|
|
# Test case 1: Non-streaming response
|
|
start_time = 1000.0
|
|
end_time = 1005.0
|
|
completion_start_time = 1003.0
|
|
stream = False
|
|
|
|
response_time = StandardLoggingPayloadSetup.get_response_time(
|
|
start_time_float=start_time,
|
|
end_time_float=end_time,
|
|
completion_start_time_float=completion_start_time,
|
|
stream=stream,
|
|
)
|
|
|
|
# For non-streaming, should return end_time - start_time
|
|
assert response_time == 5.0
|
|
|
|
# Test case 2: Streaming response
|
|
start_time = 1000.0
|
|
end_time = 1010.0
|
|
completion_start_time = 1002.0
|
|
stream = True
|
|
|
|
response_time = StandardLoggingPayloadSetup.get_response_time(
|
|
start_time_float=start_time,
|
|
end_time_float=end_time,
|
|
completion_start_time_float=completion_start_time,
|
|
stream=stream,
|
|
)
|
|
|
|
# For streaming, should return completion_start_time - start_time
|
|
assert response_time == 2.0
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"metadata, expected_requester_metadata",
|
|
[
|
|
({"metadata": {"test": "test2"}}, {"test": "test2"}),
|
|
({"metadata": {"test": "test2"}, "model_id": "test-model"}, {"test": "test2"}),
|
|
(
|
|
{
|
|
"metadata": {
|
|
"test": "test2",
|
|
},
|
|
"model_id": "test-model",
|
|
"requester_metadata": {"test": "test2"},
|
|
},
|
|
{"test": "test2"},
|
|
),
|
|
],
|
|
)
|
|
def test_standard_logging_metadata_requester_metadata(
|
|
metadata, expected_requester_metadata
|
|
):
|
|
result = StandardLoggingPayloadSetup.get_standard_logging_metadata(metadata)
|
|
assert result["requester_metadata"] == expected_requester_metadata
|
|
|
|
|
|
def test_cost_breakdown_in_standard_logging_payload():
|
|
"""
|
|
Test that cost breakdown fields are properly included in StandardLoggingPayload.
|
|
Tests input_cost, output_cost, tool_usage_cost, and total_cost fields.
|
|
"""
|
|
from litellm.litellm_core_utils.litellm_logging import (
|
|
get_standard_logging_object_payload,
|
|
Logging,
|
|
)
|
|
from litellm.types.utils import Usage
|
|
from datetime import datetime
|
|
import time
|
|
|
|
# Create a mock logging object with cost breakdown
|
|
logging_obj = Logging(
|
|
model="gpt-5.5",
|
|
messages=[{"role": "user", "content": "Hello"}],
|
|
stream=False,
|
|
call_type="completion",
|
|
start_time=datetime.now(),
|
|
litellm_call_id="test-123",
|
|
function_id="test-function",
|
|
)
|
|
|
|
# Simulate cost breakdown being stored during cost calculation
|
|
logging_obj.set_cost_breakdown(
|
|
input_cost=0.001,
|
|
output_cost=0.002,
|
|
total_cost=0.0035,
|
|
cost_for_built_in_tools_cost_usd_dollar=0.0005,
|
|
)
|
|
|
|
# Mock response object
|
|
mock_response = {
|
|
"id": "chatcmpl-123",
|
|
"object": "chat.completion",
|
|
"model": "gpt-5.5",
|
|
"usage": {
|
|
"prompt_tokens": 10,
|
|
"completion_tokens": 20,
|
|
"total_tokens": 30,
|
|
},
|
|
"choices": [
|
|
{
|
|
"index": 0,
|
|
"message": {
|
|
"role": "assistant",
|
|
"content": "Hello! How can I help you today?",
|
|
},
|
|
"finish_reason": "stop",
|
|
}
|
|
],
|
|
}
|
|
|
|
# Create kwargs
|
|
kwargs = {
|
|
"model": "gpt-5.5",
|
|
"messages": [{"role": "user", "content": "Hello"}],
|
|
"response_cost": 0.0035,
|
|
"custom_llm_provider": "openai",
|
|
}
|
|
|
|
start_time = datetime.now()
|
|
end_time = datetime.now()
|
|
|
|
# Get the standard logging payload
|
|
payload = get_standard_logging_object_payload(
|
|
kwargs=kwargs,
|
|
init_response_obj=mock_response,
|
|
start_time=start_time,
|
|
end_time=end_time,
|
|
logging_obj=logging_obj,
|
|
status="success",
|
|
)
|
|
|
|
# Verify the cost breakdown field is present
|
|
assert payload is not None
|
|
assert payload["cost_breakdown"] is not None
|
|
assert payload["cost_breakdown"]["input_cost"] == 0.001
|
|
assert payload["cost_breakdown"]["output_cost"] == 0.002
|
|
assert payload["cost_breakdown"]["tool_usage_cost"] == 0.0005
|
|
assert payload["cost_breakdown"]["total_cost"] == 0.0035
|
|
assert payload["response_cost"] == 0.0035
|
|
|
|
print("✅ Cost breakdown test passed!")
|
|
|
|
|
|
def test_cost_breakdown_missing_in_standard_logging_payload():
|
|
"""
|
|
Test that cost breakdown field is None when not available (e.g., for embedding calls)
|
|
"""
|
|
from litellm.litellm_core_utils.litellm_logging import (
|
|
get_standard_logging_object_payload,
|
|
Logging,
|
|
)
|
|
from datetime import datetime
|
|
|
|
# Create a mock logging object without cost breakdown
|
|
logging_obj = Logging(
|
|
model="gpt-5.5",
|
|
messages=[{"role": "user", "content": "Hello"}],
|
|
stream=False,
|
|
call_type="embedding", # Non-completion call type
|
|
start_time=datetime.now(),
|
|
litellm_call_id="test-123",
|
|
function_id="test-function",
|
|
)
|
|
|
|
# No cost breakdown stored
|
|
|
|
# Mock response object
|
|
mock_response = {
|
|
"object": "list",
|
|
"data": [{"embedding": [0.1, 0.2, 0.3]}],
|
|
"model": "text-embedding-3-small",
|
|
"usage": {"prompt_tokens": 10, "total_tokens": 10},
|
|
}
|
|
|
|
kwargs = {
|
|
"model": "text-embedding-3-small",
|
|
"input": ["Hello"],
|
|
"response_cost": 0.0001,
|
|
"custom_llm_provider": "openai",
|
|
}
|
|
|
|
start_time = datetime.now()
|
|
end_time = datetime.now()
|
|
|
|
# Get the standard logging payload
|
|
payload = get_standard_logging_object_payload(
|
|
kwargs=kwargs,
|
|
init_response_obj=mock_response,
|
|
start_time=start_time,
|
|
end_time=end_time,
|
|
logging_obj=logging_obj,
|
|
status="success",
|
|
)
|
|
|
|
# Verify the cost breakdown field is None for non-completion calls
|
|
assert payload is not None
|
|
assert payload["cost_breakdown"] is None
|
|
assert payload["response_cost"] == 0.0001
|
|
|
|
print("✅ Cost breakdown missing test passed!")
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"use_combined_usage_object",
|
|
[False, True],
|
|
ids=["normal_usage_dict", "combined_usage_object"],
|
|
)
|
|
def test_usage_dict_roundtrip_in_payload(use_combined_usage_object):
|
|
"""
|
|
Regression test: verify that usage data flows correctly through
|
|
get_standard_logging_object_payload without unnecessary Pydantic round-trips.
|
|
|
|
Checks:
|
|
- usage_object in StandardLoggingMetadata is a plain dict with correct token values
|
|
- prompt_tokens, completion_tokens, total_tokens on the payload match the usage dict
|
|
- Works for both normal usage dict path and combined_usage_object (realtime API) path
|
|
"""
|
|
from litellm.litellm_core_utils.litellm_logging import (
|
|
get_standard_logging_object_payload,
|
|
Logging,
|
|
)
|
|
from datetime import datetime
|
|
|
|
logging_obj = Logging(
|
|
model="gpt-5.5",
|
|
messages=[{"role": "user", "content": "Hi"}],
|
|
stream=False,
|
|
call_type="completion",
|
|
start_time=datetime.now(),
|
|
litellm_call_id="test-usage-roundtrip",
|
|
function_id="test-fn",
|
|
)
|
|
|
|
mock_response = {
|
|
"id": "chatcmpl-usage-test",
|
|
"object": "chat.completion",
|
|
"model": "gpt-5.5",
|
|
"usage": {
|
|
"prompt_tokens": 42,
|
|
"completion_tokens": 58,
|
|
"total_tokens": 100,
|
|
},
|
|
"choices": [
|
|
{
|
|
"index": 0,
|
|
"message": {"role": "assistant", "content": "Hello!"},
|
|
"finish_reason": "stop",
|
|
}
|
|
],
|
|
}
|
|
|
|
kwargs = {
|
|
"model": "gpt-5.5",
|
|
"messages": [{"role": "user", "content": "Hi"}],
|
|
"response_cost": 0.01,
|
|
"custom_llm_provider": "openai",
|
|
}
|
|
|
|
if use_combined_usage_object:
|
|
kwargs["combined_usage_object"] = Usage(
|
|
prompt_tokens=42, completion_tokens=58, total_tokens=100
|
|
)
|
|
|
|
start_time = datetime.now()
|
|
end_time = datetime.now()
|
|
|
|
payload = get_standard_logging_object_payload(
|
|
kwargs=kwargs,
|
|
init_response_obj=mock_response,
|
|
start_time=start_time,
|
|
end_time=end_time,
|
|
logging_obj=logging_obj,
|
|
status="success",
|
|
)
|
|
|
|
assert payload is not None
|
|
|
|
# Top-level token fields must match
|
|
assert payload["prompt_tokens"] == 42
|
|
assert payload["completion_tokens"] == 58
|
|
assert payload["total_tokens"] == 100
|
|
|
|
# usage_object in metadata must be a plain dict (not a Pydantic model)
|
|
usage_obj = payload["metadata"]["usage_object"]
|
|
assert isinstance(usage_obj, dict)
|
|
assert usage_obj["prompt_tokens"] == 42
|
|
assert usage_obj["completion_tokens"] == 58
|
|
assert usage_obj["total_tokens"] == 100
|
|
|
|
|
|
def test_standard_logging_payload_uses_actual_model_for_azure_router():
|
|
from litellm.litellm_core_utils.litellm_logging import (
|
|
Logging,
|
|
get_standard_logging_object_payload,
|
|
)
|
|
|
|
logging_obj = Logging(
|
|
model="azure_ai/model-router",
|
|
messages=[{"role": "user", "content": "Hello"}],
|
|
stream=False,
|
|
call_type="completion",
|
|
start_time=datetime.now(),
|
|
litellm_call_id="test-azure-router-opt-in",
|
|
function_id="test-fn",
|
|
)
|
|
|
|
kwargs = {
|
|
"model": "azure_ai/model-router",
|
|
"messages": [{"role": "user", "content": "Hello"}],
|
|
"response_cost": 0.00001,
|
|
"custom_llm_provider": "azure_ai",
|
|
}
|
|
mock_response = {
|
|
"id": "chatcmpl-azure-router-opt-in",
|
|
"object": "chat.completion",
|
|
"model": "azure_ai/gpt-5-nano-2025-08-07",
|
|
"usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
|
|
"choices": [
|
|
{
|
|
"index": 0,
|
|
"message": {"role": "assistant", "content": "hello"},
|
|
"finish_reason": "stop",
|
|
}
|
|
],
|
|
}
|
|
|
|
payload = get_standard_logging_object_payload(
|
|
kwargs=kwargs,
|
|
init_response_obj=mock_response,
|
|
start_time=datetime.now(),
|
|
end_time=datetime.now(),
|
|
logging_obj=logging_obj,
|
|
status="success",
|
|
)
|
|
assert payload is not None
|
|
assert payload["model"] == "azure_ai/gpt-5-nano-2025-08-07"
|
|
|
|
|
|
def test_standard_logging_payload_uses_actual_model_for_azure_router_with_underscore():
|
|
from litellm.litellm_core_utils.litellm_logging import (
|
|
Logging,
|
|
get_standard_logging_object_payload,
|
|
)
|
|
|
|
logging_obj = Logging(
|
|
model="azure_ai/model_router",
|
|
messages=[{"role": "user", "content": "Hello"}],
|
|
stream=False,
|
|
call_type="completion",
|
|
start_time=datetime.now(),
|
|
litellm_call_id="test-azure-router-underscore",
|
|
function_id="test-fn",
|
|
)
|
|
|
|
kwargs = {
|
|
"model": "azure_ai/model_router",
|
|
"messages": [{"role": "user", "content": "Hello"}],
|
|
"response_cost": 0.00001,
|
|
"custom_llm_provider": "azure_ai",
|
|
}
|
|
mock_response = {
|
|
"id": "chatcmpl-azure-router-underscore",
|
|
"object": "chat.completion",
|
|
"model": "azure_ai/gpt-5-nano-2025-08-07",
|
|
"usage": {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
|
|
"choices": [
|
|
{
|
|
"index": 0,
|
|
"message": {"role": "assistant", "content": "hello"},
|
|
"finish_reason": "stop",
|
|
}
|
|
],
|
|
}
|
|
|
|
payload = get_standard_logging_object_payload(
|
|
kwargs=kwargs,
|
|
init_response_obj=mock_response,
|
|
start_time=datetime.now(),
|
|
end_time=datetime.now(),
|
|
logging_obj=logging_obj,
|
|
status="success",
|
|
)
|
|
assert payload is not None
|
|
assert payload["model"] == "azure_ai/gpt-5-nano-2025-08-07"
|
|
|
|
|
|
def test_merge_litellm_metadata_basic():
|
|
"""
|
|
Test that merge_litellm_metadata correctly merges metadata and litellm_metadata.
|
|
User API key fields (from metadata) should take precedence over model-related fields (from litellm_metadata).
|
|
"""
|
|
litellm_params = {
|
|
"metadata": {
|
|
"user_api_key": "test-key-123",
|
|
"user_api_key_user_id": "user-456",
|
|
"user_api_key_team_id": "team-789",
|
|
},
|
|
"litellm_metadata": {
|
|
"model_group": "gpt-4-group",
|
|
"model_info": {"id": "model-123"},
|
|
"tags": ["tag1", "tag2"],
|
|
},
|
|
}
|
|
|
|
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
|
|
|
|
# Check that user API key fields are present
|
|
assert result["user_api_key"] == "test-key-123"
|
|
assert result["user_api_key_user_id"] == "user-456"
|
|
assert result["user_api_key_team_id"] == "team-789"
|
|
|
|
# Check that model-related fields are present
|
|
assert result["model_group"] == "gpt-4-group"
|
|
assert result["model_info"] == {"id": "model-123"}
|
|
assert result["tags"] == ["tag1", "tag2"]
|
|
|
|
|
|
def test_merge_litellm_metadata_precedence():
|
|
"""
|
|
Test that metadata fields take precedence over litellm_metadata when there are conflicts.
|
|
"""
|
|
litellm_params = {
|
|
"metadata": {
|
|
"tags": ["user-tag1", "user-tag2"],
|
|
"custom_field": "from_metadata",
|
|
},
|
|
"litellm_metadata": {
|
|
"tags": ["model-tag1", "model-tag2"], # This should NOT overwrite
|
|
"custom_field": "from_litellm_metadata", # This should NOT overwrite
|
|
"model_group": "gpt-4-group", # This should be included
|
|
},
|
|
}
|
|
|
|
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
|
|
|
|
# metadata values should take precedence
|
|
assert result["tags"] == ["user-tag1", "user-tag2"]
|
|
assert result["custom_field"] == "from_metadata"
|
|
|
|
# litellm_metadata values should only be included if not in metadata
|
|
assert result["model_group"] == "gpt-4-group"
|
|
|
|
|
|
def test_merge_litellm_metadata_skip_non_serializable():
|
|
"""
|
|
Test that non-serializable objects like UserAPIKeyAuth are skipped.
|
|
"""
|
|
from litellm.proxy._types import UserAPIKeyAuth
|
|
|
|
user_api_key_auth = UserAPIKeyAuth(
|
|
api_key="test-key",
|
|
user_id="test-user",
|
|
team_id="test-team",
|
|
)
|
|
|
|
litellm_params = {
|
|
"metadata": {
|
|
"user_api_key": "test-key-123",
|
|
"user_api_key_auth": user_api_key_auth, # This should be skipped
|
|
"safe_field": "safe_value",
|
|
},
|
|
"litellm_metadata": {
|
|
"model_group": "gpt-4-group",
|
|
},
|
|
}
|
|
|
|
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
|
|
|
|
# user_api_key_auth should be skipped
|
|
assert "user_api_key_auth" not in result
|
|
|
|
# Other fields should be present
|
|
assert result["user_api_key"] == "test-key-123"
|
|
assert result["safe_field"] == "safe_value"
|
|
assert result["model_group"] == "gpt-4-group"
|
|
|
|
|
|
def test_merge_litellm_metadata_empty_params():
|
|
"""
|
|
Test that merge_litellm_metadata handles empty or missing metadata gracefully.
|
|
"""
|
|
# Test with empty litellm_params
|
|
result = StandardLoggingPayloadSetup.merge_litellm_metadata({})
|
|
assert result == {}
|
|
|
|
# Test with only metadata
|
|
litellm_params = {
|
|
"metadata": {
|
|
"user_api_key": "test-key",
|
|
}
|
|
}
|
|
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
|
|
assert result == {"user_api_key": "test-key"}
|
|
|
|
# Test with only litellm_metadata
|
|
litellm_params = {
|
|
"litellm_metadata": {
|
|
"model_group": "gpt-4-group",
|
|
}
|
|
}
|
|
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
|
|
assert result == {"model_group": "gpt-4-group"}
|
|
|
|
# Test with None values
|
|
litellm_params = {
|
|
"metadata": None,
|
|
"litellm_metadata": None,
|
|
}
|
|
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
|
|
assert result == {}
|
|
|
|
|
|
def test_merge_litellm_metadata_bedrock_passthrough_scenario():
|
|
"""
|
|
Test merge_litellm_metadata in a Bedrock passthrough scenario where both
|
|
user API key metadata and model metadata need to be merged.
|
|
|
|
This is the specific scenario that was fixed - bedrock passthrough requests
|
|
should include complete user authentication metadata in logging.
|
|
"""
|
|
litellm_params = {
|
|
"metadata": {
|
|
# User API key fields from authentication
|
|
"user_api_key": "sk-bedrock-test-key-123",
|
|
"user_api_key_hash": "hashed-key-123",
|
|
"user_api_key_user_id": "bedrock-user-456",
|
|
"user_api_key_team_id": "bedrock-team-789",
|
|
"user_api_key_org_id": "bedrock-org-101",
|
|
"user_api_key_alias": "bedrock-key-alias",
|
|
"user_api_key_team_alias": "bedrock-team-alias",
|
|
"user_api_key_end_user_id": "end-user-123",
|
|
"user_api_key_request_route": "/bedrock/model/invoke",
|
|
},
|
|
"litellm_metadata": {
|
|
# Model-related fields from Bedrock configuration
|
|
"model_group": "bedrock-claude-group",
|
|
"model_info": {
|
|
"id": "anthropic.claude-3-sonnet",
|
|
"mode": "chat",
|
|
},
|
|
"aws_region_name": "us-east-1",
|
|
"tags": ["production", "bedrock"],
|
|
},
|
|
}
|
|
|
|
result = StandardLoggingPayloadSetup.merge_litellm_metadata(litellm_params)
|
|
|
|
# Verify all user API key fields are present
|
|
assert result["user_api_key"] == "sk-bedrock-test-key-123"
|
|
assert result["user_api_key_hash"] == "hashed-key-123"
|
|
assert result["user_api_key_user_id"] == "bedrock-user-456"
|
|
assert result["user_api_key_team_id"] == "bedrock-team-789"
|
|
assert result["user_api_key_org_id"] == "bedrock-org-101"
|
|
assert result["user_api_key_alias"] == "bedrock-key-alias"
|
|
assert result["user_api_key_team_alias"] == "bedrock-team-alias"
|
|
assert result["user_api_key_end_user_id"] == "end-user-123"
|
|
assert result["user_api_key_request_route"] == "/bedrock/model/invoke"
|
|
|
|
# Verify all model-related fields are present
|
|
assert result["model_group"] == "bedrock-claude-group"
|
|
assert result["model_info"] == {
|
|
"id": "anthropic.claude-3-sonnet",
|
|
"mode": "chat",
|
|
}
|
|
assert result["aws_region_name"] == "us-east-1"
|
|
assert result["tags"] == ["production", "bedrock"]
|
|
|
|
# Verify total number of fields (9 user fields + 4 model fields = 13)
|
|
assert len(result) == 13
|