litellm/tests/test_litellm/test_utils.py
Deepanshu Lulla 9ce96c2d34
feat(logging): add opt-in session_id and trace_id correlation to JSON log records via contextvars (#34418)
* 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 commit 9f3a20f4b2 and 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 in f1cf9589d6.

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>
2026-08-10 10:40:13 -07:00

5256 lines
201 KiB
Python

import json
import logging
import os
import sys
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from jsonschema import validate
sys.path.insert(
0, os.path.abspath("../..")
) # Adds the parent directory to the system path
import litellm
from litellm._logging import (
CorrelationContextFilter,
JsonFormatter,
session_id_var,
trace_id_var,
verbose_logger,
)
from litellm.proxy.utils import is_valid_api_key
from litellm.types.utils import (
CallTypes,
Delta,
LlmProviders,
ModelResponseStream,
PromptTokensDetailsWrapper,
StreamingChoices,
Usage,
)
from litellm.utils import (
ProviderConfigManager,
TextCompletionStreamWrapper,
_check_provider_match,
_is_streaming_request,
get_api_key,
get_llm_provider,
get_optional_params_image_gen,
get_prompt_cache_min_tokens,
is_cached_message,
is_prompt_caching_valid_prompt,
)
# Adds the parent directory to the system path
def test_usage_openai_cache_write_tokens_populates_both_names():
"""OpenAI reports cache-write tokens as prompt_tokens_details.cache_write_tokens.
The Usage constructor must expose it under both cache_write_tokens (canonical,
OpenAI naming) and cache_creation_tokens (legacy, Anthropic naming)."""
usage = Usage(
prompt_tokens=1000,
completion_tokens=10,
total_tokens=1010,
prompt_tokens_details={"cached_tokens": 0, "cache_write_tokens": 800},
)
assert usage.prompt_tokens_details.cache_write_tokens == 800
assert usage.prompt_tokens_details.cache_creation_tokens == 800
def test_usage_anthropic_cache_creation_maps_to_cache_write_tokens():
"""Anthropic/Bedrock report the top-level cache_creation_input_tokens field.
It must be normalized onto the OpenAI cache_write_tokens name as well as the
legacy cache_creation_tokens name."""
usage = Usage(
prompt_tokens=500,
completion_tokens=50,
total_tokens=550,
cache_creation_input_tokens=300,
cache_read_input_tokens=120,
)
assert usage.prompt_tokens_details.cache_write_tokens == 300
assert usage.prompt_tokens_details.cache_creation_tokens == 300
assert usage.prompt_tokens_details.cached_tokens == 120
def test_prompt_tokens_details_no_cache_write_tokens_when_absent():
"""A read-only cache hit (no cache write) must not surface cache-write fields."""
details = PromptTokensDetailsWrapper(cached_tokens=800)
assert details.cached_tokens == 800
assert not hasattr(details, "cache_write_tokens")
assert not hasattr(details, "cache_creation_tokens")
def test_prompt_tokens_details_cache_write_creation_stay_in_sync_on_assignment():
"""Assigning either name after construction must mirror to the other, so a
caller that sets only one field can't leave the pair silently out of sync."""
details = PromptTokensDetailsWrapper(cache_write_tokens=100)
assert details.cache_write_tokens == details.cache_creation_tokens == 100
details.cache_write_tokens = 250
assert details.cache_write_tokens == details.cache_creation_tokens == 250
details.cache_creation_tokens = 375
assert details.cache_write_tokens == details.cache_creation_tokens == 375
@pytest.fixture
def local_model_cost_map(monkeypatch):
original_model_cost = litellm.model_cost
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
litellm.model_cost = litellm.get_model_cost_map(url="")
litellm.get_model_info.cache_clear()
try:
yield
finally:
litellm.model_cost = original_model_cost
litellm.get_model_info.cache_clear()
def test_get_model_info_surfaces_supports_adaptive_thinking(local_model_cost_map):
"""supports_adaptive_thinking must flow through get_model_info like every other
capability flag: both from an explicit cost-map entry and from a
fallback-generalization rule for an unmapped model. Regression: the field shipped
in the JSON but was never declared on ModelInfo nor copied during construction, so
get_model_info (and _supports_factory) silently dropped it for any provider-prefixed
or unmapped name."""
explicit = litellm.get_model_info(model="claude-opus-4-8")
assert explicit["supports_adaptive_thinking"] is True
generalized = litellm.get_model_info(
model="claude-opus-4-9", custom_llm_provider="anthropic"
)
assert generalized["supports_adaptive_thinking"] is True
def test_check_provider_match_azure_ai_allows_openai_and_azure():
"""
Test that azure_ai provider can match openai and azure models.
This is needed for Azure Model Router which can route to OpenAI models.
"""
# azure_ai should match openai models
assert (
_check_provider_match(
model_info={"litellm_provider": "openai"}, custom_llm_provider="azure_ai"
)
is True
)
# azure_ai should match azure models
assert (
_check_provider_match(
model_info={"litellm_provider": "azure"}, custom_llm_provider="azure_ai"
)
is True
)
# azure_ai should NOT match other providers
assert (
_check_provider_match(
model_info={"litellm_provider": "anthropic"}, custom_llm_provider="azure_ai"
)
is False
)
def test_check_provider_match_github_allows_upstream_provider_metadata():
"""
Test that github provider can match upstream provider metadata.
GitHub Models can provide models from multiple providers.
"""
assert (
_check_provider_match(
model_info={"litellm_provider": "openai"},
custom_llm_provider="github",
)
is True
)
assert (
_check_provider_match(
model_info={"litellm_provider": "github"},
custom_llm_provider="github",
)
is True
)
assert (
_check_provider_match(
model_info={"litellm_provider": "anthropic"},
custom_llm_provider="github",
)
is True
)
def test_supports_function_calling_github_openai_alias():
assert litellm.utils.supports_function_calling(model="github/gpt-4o-mini") is True
assert (
litellm.utils.supports_function_calling(
model="gpt-4o-mini", custom_llm_provider="github"
)
is True
)
def test_supports_function_calling_github_anthropic_alias():
assert (
litellm.utils.supports_function_calling(
model="github/claude-3-7-sonnet-20250219"
)
is True
)
def test_supports_function_calling_deepinfra_llama():
"""Test that deepinfra Llama models correctly report function calling support.
Regression test for https://github.com/BerriAI/litellm/issues/22619
"""
assert (
litellm.utils.supports_function_calling(
model="deepinfra/meta-llama/Llama-3.3-70B-Instruct-Turbo"
)
is True
)
def test_supports_function_calling_unknown_github_alias_returns_false():
assert (
litellm.utils.supports_function_calling(
model="github/non-existent-model-for-capability-check"
)
is False
)
def test_get_optional_params_image_gen():
from litellm.llms.azure.image_generation import AzureGPTImageGenerationConfig
provider_config = AzureGPTImageGenerationConfig()
optional_params = get_optional_params_image_gen(
model="gpt-image-1",
response_format="b64_json",
n=3,
custom_llm_provider="azure",
drop_params=True,
provider_config=provider_config,
)
assert optional_params is not None
assert "response_format" not in optional_params
assert optional_params["n"] == 3
def test_get_optional_params_image_gen_vertex_ai_size():
"""Test that Vertex AI image generation properly handles size parameter and maps it to aspectRatio"""
# Test with various size parameters
test_cases = [
("1024x1024", "1:1"), # Square aspect ratio
("256x256", "1:1"), # Square aspect ratio
("512x512", "1:1"), # Square aspect ratio
("1792x1024", "16:9"), # Landscape aspect ratio
("1024x1792", "9:16"), # Portrait aspect ratio
("unsupported", "1:1"), # Default to square for unsupported sizes
]
for size_input, expected_aspect_ratio in test_cases:
optional_params = get_optional_params_image_gen(
model="vertex_ai/imagegeneration@006",
size=size_input,
n=2,
custom_llm_provider="vertex_ai",
drop_params=True,
)
assert optional_params is not None
assert optional_params["aspectRatio"] == expected_aspect_ratio
assert optional_params["sampleCount"] == 2
assert "size" not in optional_params # size should be converted to aspectRatio
# Test without size parameter
optional_params = get_optional_params_image_gen(
model="vertex_ai/imagegeneration@006",
n=1,
custom_llm_provider="vertex_ai",
drop_params=True,
)
assert optional_params is not None
assert (
"aspectRatio" not in optional_params
) # aspectRatio should not be set if size is not provided
assert optional_params["sampleCount"] == 1
def test_get_optional_params_image_gen_filters_empty_values():
optional_params = get_optional_params_image_gen(
model="gpt-image-1",
custom_llm_provider="openai",
extra_body={},
)
assert optional_params == {}
def test_gpt_image_provider_detection_covers_existing_family():
for image_model in ("gpt-image-1", "gpt-image-1-mini", "gpt-image-1.5"):
model, custom_llm_provider, _, _ = litellm.get_llm_provider(model=image_model)
assert model == image_model
assert custom_llm_provider == "openai"
def test_gpt_image_2_provider_and_model_info(local_model_cost_map):
model, custom_llm_provider, _, _ = litellm.get_llm_provider(model="gpt-image-2")
assert model == "gpt-image-2"
assert custom_llm_provider == "openai"
model_info = litellm.get_model_info(model="gpt-image-2")
assert model_info["litellm_provider"] == "openai"
assert model_info["mode"] == "image_generation"
assert model_info["input_cost_per_token"] == 5e-06
assert model_info["input_cost_per_image_token"] == 8e-06
assert model_info["output_cost_per_token"] == 1e-05
assert model_info["output_cost_per_image_token"] == 3e-05
assert (
"/v1/images/generations"
in litellm.model_cost["gpt-image-2"]["supported_endpoints"]
)
assert (
"/v1/images/edits" in litellm.model_cost["gpt-image-2"]["supported_endpoints"]
)
assert model_info["supports_vision"] is True
assert model_info["supports_pdf_input"] is True
def test_gpt_image_2_snapshot_model_info(local_model_cost_map):
model, custom_llm_provider, _, _ = litellm.get_llm_provider(
model="gpt-image-2-2026-04-21"
)
assert model == "gpt-image-2-2026-04-21"
assert custom_llm_provider == "openai"
model_info = litellm.get_model_info(model="gpt-image-2-2026-04-21")
assert model_info["litellm_provider"] == "openai"
assert model_info["mode"] == "image_generation"
assert model_info["output_cost_per_image_token"] == 3e-05
def test_azure_gpt_image_2_model_info(local_model_cost_map):
model, custom_llm_provider, _, _ = litellm.get_llm_provider(
model="azure/gpt-image-2"
)
assert model == "gpt-image-2"
assert custom_llm_provider == "azure"
model_info = litellm.get_model_info(
model="gpt-image-2", custom_llm_provider="azure"
)
assert model_info["litellm_provider"] == "azure"
assert model_info["mode"] == "image_generation"
assert model_info["input_cost_per_token"] == 5e-06
assert model_info["input_cost_per_image_token"] == 8e-06
assert model_info["output_cost_per_token"] == 1e-05
assert model_info["output_cost_per_image_token"] == 3e-05
def test_all_model_configs():
from litellm.llms.vertex_ai.vertex_ai_partner_models.ai21.transformation import (
VertexAIAi21Config,
)
from litellm.llms.vertex_ai.vertex_ai_partner_models.llama3.transformation import (
VertexAILlama3Config,
)
assert (
"max_completion_tokens"
in VertexAILlama3Config().get_supported_openai_params(model="llama3")
)
assert VertexAILlama3Config().map_openai_params(
{"max_completion_tokens": 10}, {}, "llama3", drop_params=False
) == {"max_tokens": 10}
assert "max_completion_tokens" in VertexAIAi21Config().get_supported_openai_params(
model="jamba-1.5-mini@001"
)
assert VertexAIAi21Config().map_openai_params(
{"max_completion_tokens": 10}, {}, "jamba-1.5-mini@001", drop_params=False
) == {"max_tokens": 10}
from litellm.llms.fireworks_ai.chat.transformation import FireworksAIConfig
assert "max_completion_tokens" in FireworksAIConfig().get_supported_openai_params(
model="llama3"
)
assert FireworksAIConfig().map_openai_params(
model="llama3",
non_default_params={"max_completion_tokens": 10},
optional_params={},
drop_params=False,
) == {"max_tokens": 10}
from litellm.llms.nvidia_nim.chat.transformation import NvidiaNimConfig
assert "max_completion_tokens" in NvidiaNimConfig().get_supported_openai_params(
model="llama3"
)
assert NvidiaNimConfig().map_openai_params(
model="llama3",
non_default_params={"max_completion_tokens": 10},
optional_params={},
drop_params=False,
) == {"max_tokens": 10}
from litellm.llms.ollama.chat.transformation import OllamaChatConfig
assert "max_completion_tokens" in OllamaChatConfig().get_supported_openai_params(
model="llama3"
)
assert OllamaChatConfig().map_openai_params(
model="llama3",
non_default_params={"max_completion_tokens": 10},
optional_params={},
drop_params=False,
) == {"num_predict": 10}
from litellm.llms.predibase.chat.transformation import PredibaseConfig
assert "max_completion_tokens" in PredibaseConfig().get_supported_openai_params(
model="llama3"
)
assert PredibaseConfig().map_openai_params(
model="llama3",
non_default_params={"max_completion_tokens": 10},
optional_params={},
drop_params=False,
) == {"max_new_tokens": 10}
from litellm.llms.codestral.completion.transformation import (
CodestralTextCompletionConfig,
)
assert (
"max_completion_tokens"
in CodestralTextCompletionConfig().get_supported_openai_params(model="llama3")
)
assert CodestralTextCompletionConfig().map_openai_params(
model="llama3",
non_default_params={"max_completion_tokens": 10},
optional_params={},
drop_params=False,
) == {"max_tokens": 10}
from litellm.llms.volcengine.chat.transformation import (
VolcEngineChatConfig as VolcEngineConfig,
)
assert "max_completion_tokens" in VolcEngineConfig().get_supported_openai_params(
model="llama3"
)
assert VolcEngineConfig().map_openai_params(
model="llama3",
non_default_params={"max_completion_tokens": 10},
optional_params={},
drop_params=False,
) == {"max_tokens": 10}
from litellm.llms.ai21.chat.transformation import AI21ChatConfig
assert "max_completion_tokens" in AI21ChatConfig().get_supported_openai_params(
"jamba-1.5-mini@001"
)
assert AI21ChatConfig().map_openai_params(
model="jamba-1.5-mini@001",
non_default_params={"max_completion_tokens": 10},
optional_params={},
drop_params=False,
) == {"max_tokens": 10}
from litellm.llms.azure.chat.gpt_transformation import AzureOpenAIConfig
assert "max_completion_tokens" in AzureOpenAIConfig().get_supported_openai_params(
model="gpt-3.5-turbo"
)
assert AzureOpenAIConfig().map_openai_params(
model="gpt-3.5-turbo",
non_default_params={"max_completion_tokens": 10},
optional_params={},
api_version="2022-12-01",
drop_params=False,
) == {"max_completion_tokens": 10}
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
assert (
"max_completion_tokens"
in AmazonConverseConfig().get_supported_openai_params(
model="anthropic.claude-3-sonnet-20240229-v1:0"
)
)
assert AmazonConverseConfig().map_openai_params(
model="anthropic.claude-3-sonnet-20240229-v1:0",
non_default_params={"max_completion_tokens": 10},
optional_params={},
drop_params=False,
) == {"maxTokens": 10}
from litellm.llms.codestral.completion.transformation import (
CodestralTextCompletionConfig,
)
assert (
"max_completion_tokens"
in CodestralTextCompletionConfig().get_supported_openai_params(model="llama3")
)
assert CodestralTextCompletionConfig().map_openai_params(
model="llama3",
non_default_params={"max_completion_tokens": 10},
optional_params={},
drop_params=False,
) == {"max_tokens": 10}
from litellm import AmazonAnthropicClaudeConfig, AmazonAnthropicConfig
assert (
"max_completion_tokens"
in AmazonAnthropicClaudeConfig().get_supported_openai_params(
model="anthropic.claude-3-sonnet-20240229-v1:0"
)
)
assert AmazonAnthropicClaudeConfig().map_openai_params(
non_default_params={"max_completion_tokens": 10},
optional_params={},
model="anthropic.claude-3-sonnet-20240229-v1:0",
drop_params=False,
) == {"max_tokens": 10}
assert (
"max_completion_tokens"
in AmazonAnthropicConfig().get_supported_openai_params(model="")
)
assert AmazonAnthropicConfig().map_openai_params(
non_default_params={"max_completion_tokens": 10},
optional_params={},
model="",
drop_params=False,
) == {"max_tokens_to_sample": 10}
from litellm.llms.databricks.chat.transformation import DatabricksConfig
assert "max_completion_tokens" in DatabricksConfig().get_supported_openai_params()
assert DatabricksConfig().map_openai_params(
model="databricks/llama-3-70b-instruct",
drop_params=False,
non_default_params={"max_completion_tokens": 10},
optional_params={},
) == {"max_tokens": 10}
from litellm.llms.vertex_ai.vertex_ai_partner_models.anthropic.transformation import (
VertexAIAnthropicConfig,
)
assert (
"max_completion_tokens"
in VertexAIAnthropicConfig().get_supported_openai_params(
model="claude-sonnet-4-6"
)
)
assert VertexAIAnthropicConfig().map_openai_params(
non_default_params={"max_completion_tokens": 10},
optional_params={},
model="claude-sonnet-4-6",
drop_params=False,
) == {"max_tokens": 10}
from litellm.llms.gemini.chat.transformation import GoogleAIStudioGeminiConfig
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
assert "max_completion_tokens" in VertexGeminiConfig().get_supported_openai_params(
model="gemini-1.0-pro"
)
assert VertexGeminiConfig().map_openai_params(
model="gemini-1.0-pro",
non_default_params={"max_completion_tokens": 10},
optional_params={},
drop_params=False,
) == {"max_output_tokens": 10}
assert (
"max_completion_tokens"
in GoogleAIStudioGeminiConfig().get_supported_openai_params(
model="gemini-1.0-pro"
)
)
assert GoogleAIStudioGeminiConfig().map_openai_params(
model="gemini-1.0-pro",
non_default_params={"max_completion_tokens": 10},
optional_params={},
drop_params=False,
) == {"max_output_tokens": 10}
assert "max_completion_tokens" in VertexGeminiConfig().get_supported_openai_params(
model="gemini-1.0-pro"
)
assert VertexGeminiConfig().map_openai_params(
model="gemini-1.0-pro",
non_default_params={"max_completion_tokens": 10},
optional_params={},
drop_params=False,
) == {"max_output_tokens": 10}
def test_anthropic_web_search_in_model_info():
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
supported_models = [
"anthropic/claude-4-sonnet-20250514",
"anthropic/claude-sonnet-4-5-20250929",
]
for model in supported_models:
from litellm.utils import get_model_info
model_info = get_model_info(model)
assert model_info is not None
assert (
model_info["supports_web_search"] is True
), f"Model {model} should support web search"
assert (
model_info["search_context_cost_per_query"] is not None
), f"Model {model} should have a search context cost per query"
def test_cohere_embedding_optional_params():
from litellm import get_optional_params_embeddings
optional_params = get_optional_params_embeddings(
model="embed-v4.0",
custom_llm_provider="cohere",
input="Hello, world!",
input_type="search_query",
dimensions=512,
)
assert optional_params is not None
def validate_model_cost_values(model_data, exceptions=None):
"""
Validates that cost values in model data do not exceed 1.
Args:
model_data (dict): The model data dictionary
exceptions (list, optional): List of model IDs that are allowed to have costs > 1
Returns:
tuple: (is_valid, violations) where is_valid is a boolean and violations is a list of error messages
"""
if exceptions is None:
exceptions = []
violations = []
# Define all cost-related fields to check
cost_fields = [
"input_cost_per_token",
"output_cost_per_token",
"input_cost_per_character",
"output_cost_per_character",
"input_cost_per_image",
"output_cost_per_image",
"input_cost_per_pixel",
"output_cost_per_pixel",
"input_cost_per_second",
"output_cost_per_second",
"output_cost_per_second_1080p",
"input_cost_per_query",
"input_cost_per_request",
"input_cost_per_audio_token",
"output_cost_per_audio_token",
"output_cost_per_image_token",
"input_cost_per_video_token",
"output_cost_per_video_token",
"input_cost_per_audio_per_second",
"input_cost_per_video_per_second",
"input_cost_per_token_above_128k_tokens",
"output_cost_per_token_above_128k_tokens",
"input_cost_per_token_above_200k_tokens",
"output_cost_per_token_above_200k_tokens",
"input_cost_per_token_above_272k_tokens",
"output_cost_per_token_above_272k_tokens",
"input_cost_per_character_above_128k_tokens",
"output_cost_per_character_above_128k_tokens",
"input_cost_per_image_above_128k_tokens",
"input_cost_per_video_per_second_above_8s_interval",
"input_cost_per_video_per_second_above_15s_interval",
"input_cost_per_video_per_second_above_128k_tokens",
"input_cost_per_token_batches",
"output_cost_per_token_batches",
"input_cost_per_token_cache_hit",
"cache_creation_input_token_cost",
"cache_creation_input_audio_token_cost",
"cache_read_input_token_cost",
"cache_read_input_audio_token_cost",
"input_dbu_cost_per_token",
"output_db_cost_per_token",
"output_dbu_cost_per_token",
"output_cost_per_reasoning_token",
"citation_cost_per_token",
]
# Also check nested cost fields
nested_cost_fields = [
"search_context_cost_per_query",
]
for model_id, model_info in model_data.items():
# Skip if this model is in exceptions
if model_id in exceptions:
continue
# Check direct cost fields
for field in cost_fields:
if field in model_info and model_info[field] is not None:
cost_value = model_info[field]
# Convert string values to float if needed
if isinstance(cost_value, str):
try:
cost_value = float(cost_value)
except (ValueError, TypeError):
# Skip if we can't convert to float
continue
if isinstance(cost_value, (int, float)) and cost_value > 1:
violations.append(
f"Model '{model_id}' has {field} = {cost_value} which exceeds 1"
)
# Check nested cost fields
for field in nested_cost_fields:
if field in model_info and model_info[field] is not None:
nested_costs = model_info[field]
if isinstance(nested_costs, dict):
for nested_field, nested_value in nested_costs.items():
# Convert string values to float if needed
if isinstance(nested_value, str):
try:
nested_value = float(nested_value)
except (ValueError, TypeError):
# Skip if we can't convert to float
continue
if isinstance(nested_value, (int, float)) and nested_value > 1:
violations.append(
f"Model '{model_id}' has {field}.{nested_field} = {nested_value} which exceeds 1"
)
return len(violations) == 0, violations
def test_aaamodel_prices_and_context_window_json_is_valid():
"""
Validates the `model_prices_and_context_window.json` file.
If this test fails after you update the json, you need to update the schema or correct the change you made.
"""
INTENDED_SCHEMA = {
"type": "object",
"additionalProperties": {
"type": "object",
"properties": {
"supports_computer_use": {"type": "boolean"},
"cache_creation_input_audio_token_cost": {"type": "number"},
"cache_creation_input_token_cost": {"type": "number"},
"cache_creation_input_token_cost_above_1hr": {"type": "number"},
"cache_creation_input_token_cost_above_200k_tokens": {"type": "number"},
"cache_creation_input_token_cost_above_272k_tokens": {"type": "number"},
"cache_creation_input_token_cost_above_272k_tokens_flex": {
"type": "number"
},
"cache_creation_input_token_cost_flex": {"type": "number"},
"cache_creation_input_token_cost_priority": {"type": "number"},
"cache_read_input_token_cost": {"type": "number"},
"cache_read_input_token_cost_above_200k_tokens": {"type": "number"},
"cache_read_input_token_cost_above_272k_tokens": {"type": "number"},
"cache_read_input_token_cost_above_272k_tokens_flex": {
"type": "number"
},
"cache_read_input_token_cost_above_512k_tokens": {"type": "number"},
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": {
"type": "number"
},
"cache_read_input_audio_token_cost": {"type": "number"},
"audio_transcription_config": {"type": "string"},
"deprecation_date": {"type": "string"},
"input_cost_per_audio_per_second": {"type": "number"},
"input_cost_per_audio_per_second_above_128k_tokens": {"type": "number"},
"input_cost_per_audio_token": {"type": "number"},
"input_cost_per_image_token": {"type": "number"},
"input_cost_per_character": {"type": "number"},
"input_cost_per_character_above_128k_tokens": {"type": "number"},
"input_cost_per_image": {"type": "number"},
"input_cost_per_image_above_128k_tokens": {"type": "number"},
"input_cost_per_image_token": {"type": "number"},
"input_cost_per_video_token": {"type": "number"},
"input_cost_per_token_above_200k_tokens": {"type": "number"},
"input_cost_per_token_above_256k_tokens": {"type": "number"},
"input_cost_per_token_above_272k_tokens": {"type": "number"},
"input_cost_per_token_above_512k_tokens": {"type": "number"},
"cache_read_input_token_cost_flex": {"type": "number"},
"cache_read_input_token_cost_priority": {"type": "number"},
"cache_read_input_token_cost_above_200k_tokens_priority": {
"type": "number"
},
"cache_read_input_token_cost_above_272k_tokens_priority": {
"type": "number"
},
"input_cost_per_token_flex": {"type": "number"},
"input_cost_per_token_priority": {"type": "number"},
"input_cost_per_token_above_200k_tokens_priority": {"type": "number"},
"input_cost_per_token_above_272k_tokens_priority": {"type": "number"},
"input_cost_per_token_above_272k_tokens_flex": {"type": "number"},
"input_cost_per_audio_token_priority": {"type": "number"},
"output_cost_per_token_flex": {"type": "number"},
"output_cost_per_token_priority": {"type": "number"},
"output_cost_per_token_above_200k_tokens_priority": {"type": "number"},
"output_cost_per_token_above_272k_tokens_priority": {"type": "number"},
"output_cost_per_token_above_272k_tokens_flex": {"type": "number"},
"regional_processing_uplift_multiplier_eu": {"type": "number"},
"regional_processing_uplift_multiplier_us": {"type": "number"},
"input_cost_per_pixel": {"type": "number"},
"input_cost_per_query": {"type": "number"},
"input_cost_per_request": {"type": "number"},
"input_cost_per_second": {"type": "number"},
"input_cost_per_token": {"type": "number"},
"input_cost_per_token_above_128k_tokens": {"type": "number"},
"input_cost_per_token_batches": {"type": "number"},
"input_cost_per_token_cache_hit": {"type": "number"},
"input_cost_per_video_per_second": {"type": "number"},
"input_cost_per_video_per_second_above_8s_interval": {"type": "number"},
"input_cost_per_video_per_second_above_15s_interval": {
"type": "number"
},
"input_cost_per_video_per_second_above_128k_tokens": {"type": "number"},
"input_dbu_cost_per_token": {"type": "number"},
"annotation_cost_per_page": {"type": "number"},
"ocr_cost_per_page": {"type": "number"},
"ocr_cost_per_credit": {"type": "number"},
"code_interpreter_cost_per_session": {"type": "number"},
"inference_geo": {"type": "string"},
"litellm_provider": {"type": "string"},
"max_input_tokens": {"type": "number"},
"max_output_tokens": {"type": "number"},
"max_tokens": {"type": "number"},
"metadata": {"type": "object"},
"provider_specific_entry": {"type": "object"},
"mode": {
"type": "string",
"enum": [
"audio_speech",
"audio_transcription",
"chat",
"completion",
"container",
"image_edit",
"embedding",
"image_generation",
"video_generation",
"moderation",
"rerank",
"realtime",
"responses",
"ocr",
"search",
"vector_store",
],
},
"output_cost_per_audio_token": {"type": "number"},
"output_cost_per_character": {"type": "number"},
"output_cost_per_character_above_128k_tokens": {"type": "number"},
"output_cost_per_image": {"type": "number"},
"output_cost_per_image_token": {"type": "number"},
"output_cost_per_video_token": {"type": "number"},
"output_cost_per_pixel": {"type": "number"},
"output_cost_per_second": {"type": "number"},
"output_cost_per_second_1080p": {"type": "number"},
"output_cost_per_token": {"type": "number"},
"output_cost_per_token_above_128k_tokens": {"type": "number"},
"output_cost_per_token_above_200k_tokens": {"type": "number"},
"output_cost_per_token_above_256k_tokens": {"type": "number"},
"output_cost_per_token_above_272k_tokens": {"type": "number"},
"output_cost_per_token_above_512k_tokens": {"type": "number"},
"output_cost_per_token_batches": {"type": "number"},
"output_cost_per_reasoning_token": {"type": "number"},
"output_cost_per_video_per_second": {"type": "number"},
"output_db_cost_per_token": {"type": "number"},
"output_dbu_cost_per_token": {"type": "number"},
"output_vector_size": {"type": "number"},
"rpd": {"type": "number"},
"rpm": {"type": "number"},
"source": {"type": "string"},
"comment": {"type": "string"},
"supports_assistant_prefill": {"type": "boolean"},
"supports_audio_input": {"type": "boolean"},
"supports_audio_output": {"type": "boolean"},
"gemini_native_audio": {"type": "boolean"},
"gemini_audio_only_live": {"type": "boolean"},
"supports_embedding_image_input": {"type": "boolean"},
"supports_function_calling": {"type": "boolean"},
"supports_image_input": {"type": "boolean"},
"supports_nova_canvas_image_edit": {"type": "boolean"},
"supports_parallel_function_calling": {"type": "boolean"},
"supports_parallel_tool_use_config": {"type": "boolean"},
"supports_pdf_input": {"type": "boolean"},
"prompt_cache_min_tokens": {"type": "number"},
"supports_prompt_caching": {"type": "boolean"},
"supports_response_schema": {"type": "boolean"},
"supports_system_messages": {"type": "boolean"},
"supports_tool_choice": {"type": "boolean"},
"supports_video_input": {"type": "boolean"},
"supports_vision": {"type": "boolean"},
"supports_web_search": {"type": "boolean"},
"supports_url_context": {"type": "boolean"},
"supports_multimodal": {"type": "boolean"},
"uses_embed_content": {"type": "boolean"},
"supports_reasoning": {"type": "boolean"},
"supports_minimal_reasoning_effort": {"type": "boolean"},
"supports_low_reasoning_effort": {"type": "boolean"},
"supports_none_reasoning_effort": {"type": "boolean"},
"supports_xhigh_reasoning_effort": {"type": "boolean"},
"supports_max_reasoning_effort": {"type": "boolean"},
"supports_adaptive_thinking": {"type": "boolean"},
"supports_mid_conversation_system": {"type": "boolean"},
"supports_sampling_params": {"type": "boolean"},
"supports_output_config": {"type": "boolean"},
"supports_speed": {"type": "boolean"},
"bedrock_output_config_effort_ceiling": {
"type": "string",
"enum": ["low", "medium", "high", "max", "xhigh"],
},
"bedrock_converse_supports_strict_tools": {"type": "boolean"},
"tpm": {"type": "number"},
"provider_specific_entry": {"type": "object"},
"supported_endpoints": {
"type": "array",
"items": {
"type": "string",
"enum": [
"/v1/responses",
"/v1/embeddings",
"/v1/chat/completions",
"/v1/completions",
"/v1/messages",
"/v1/images/generations",
"/v1/realtime",
"/v1/realtime/transcription_sessions",
"/v1/images/variations",
"/v1/images/edits",
"/v1/batch",
"/v1/audio/transcriptions",
"/v1/audio/speech",
"/v1/ocr",
"/vertex_ai/live",
],
},
},
"supported_regions": {
"type": "array",
"items": {
"type": "string",
},
},
"search_context_cost_per_query": {
"type": "object",
"properties": {
"search_context_size_low": {"type": "number"},
"search_context_size_medium": {"type": "number"},
"search_context_size_high": {"type": "number"},
},
"additionalProperties": False,
},
"web_search_billing_unit": {
"type": "string",
"enum": ["per_prompt", "per_query"],
},
"citation_cost_per_token": {"type": "number"},
"supported_modalities": {
"type": "array",
"items": {
"type": "string",
"enum": ["text", "audio", "image", "video"],
},
},
"supported_output_modalities": {
"type": "array",
"items": {
"type": "string",
"enum": ["text", "image", "audio", "code", "video"],
},
},
"supports_native_streaming": {"type": "boolean"},
"supports_image_size": {"type": "boolean"},
"supports_native_structured_output": {"type": "boolean"},
"use_openai_responses_path": {"type": "boolean"},
"tiered_pricing": {
"type": "array",
"items": {
"type": "object",
"properties": {
"range": {
"type": "array",
"items": {"type": "number"},
"minItems": 2,
"maxItems": 2,
},
"input_cost_per_token": {"type": "number"},
"output_cost_per_token": {"type": "number"},
"cache_read_input_token_cost": {"type": "number"},
"output_cost_per_reasoning_token": {"type": "number"},
"max_results_range": {
"type": "array",
"items": {"type": "number"},
"minItems": 2,
"maxItems": 2,
},
"input_cost_per_query": {"type": "number"},
},
"additionalProperties": False,
},
},
},
"additionalProperties": False,
},
}
prod_json = os.path.join(
os.path.dirname(__file__), "..", "..", "model_prices_and_context_window.json"
)
with open(prod_json, "r") as model_prices_file:
actual_json = json.load(model_prices_file)
assert isinstance(actual_json, dict)
actual_json.pop(
"sample_spec", None
) # remove the sample, whose schema is inconsistent with the real data
actual_json.pop(
"fallback_generalizations", None
) # reserved meta key, not a model entry
# Validate schema
validate(actual_json, INTENDED_SCHEMA)
# Validate cost values
# Define exceptions for models that are allowed to have costs > 1
# Add model IDs here if they legitimately have costs > 1
exceptions = [
# Add any model IDs that should be exempt from the cost validation
# Example: "expensive-model-id",
]
is_valid, violations = validate_model_cost_values(actual_json, exceptions)
if not is_valid:
error_message = "Cost validation failed:\n" + "\n".join(violations)
error_message += "\n\nTo add exceptions, add the model ID to the 'exceptions' list in the test function."
raise AssertionError(error_message)
def test_max_tokens_consistency():
"""
Test that max_tokens == max_output_tokens for all models.
According to the spec in model_prices_and_context_window.json:
- max_tokens is a LEGACY parameter
- It should be set to max_output_tokens if the provider specifies it
This test ensures consistency across all model definitions.
"""
import json
from pathlib import Path
# Load the model configuration
config_path = (
Path(__file__).parent.parent.parent / "model_prices_and_context_window.json"
)
with open(config_path, "r") as f:
models = json.load(f)
inconsistencies = []
for model_name, config in models.items():
# Skip the sample_spec
if model_name == "sample_spec":
continue
# Check if both max_tokens and max_output_tokens exist
if isinstance(config, dict):
max_tokens = config.get("max_tokens")
max_output_tokens = config.get("max_output_tokens")
# Only validate if both exist
if max_tokens is not None and max_output_tokens is not None:
if max_tokens != max_output_tokens:
inconsistencies.append(
{
"model": model_name,
"max_tokens": max_tokens,
"max_output_tokens": max_output_tokens,
}
)
if inconsistencies:
error_msg = f"\n\n❌ Found {len(inconsistencies)} models with max_tokens != max_output_tokens:\n\n"
for item in inconsistencies[:10]: # Show first 10
error_msg += f" {item['model']}: max_tokens={item['max_tokens']}, max_output_tokens={item['max_output_tokens']}\n"
if len(inconsistencies) > 10:
error_msg += f"\n ... and {len(inconsistencies) - 10} more\n"
error_msg += "\nTo fix these inconsistencies, run: poetry run python fix_max_tokens_inconsistencies.py"
raise AssertionError(error_msg)
def test_get_model_info_gemini():
"""
Tests if ALL gemini models have 'tpm' and 'rpm' in the model info
"""
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
model_map = litellm.model_cost
for model, info in model_map.items():
if (
model.startswith("gemini/")
and not "gemma" in model
and not "learnlm" in model
and not "imagen" in model
and not "veo" in model
and not "lyria" in model
and not "robotics" in model
):
assert info.get("tpm") is not None, f"{model} does not have tpm"
assert info.get("rpm") is not None, f"{model} does not have rpm"
def test_get_model_info_bedrock_regional_inference_profile_pricing(local_model_cost_map):
"""Regression LIT-4056: with the bedrock/ routing prefix (plain, converse/, or
invoke/), the exact regional cost-map entry must win over the region-stripped
base entry, matching the unprefixed control form."""
regional = litellm.model_cost["au.anthropic.claude-opus-4-8"]
base = litellm.model_cost["anthropic.claude-opus-4-8"]
assert regional["input_cost_per_token"] > base["input_cost_per_token"]
for model in (
"bedrock/au.anthropic.claude-opus-4-8",
"bedrock/converse/au.anthropic.claude-opus-4-8",
"bedrock/invoke/au.anthropic.claude-opus-4-8",
):
info = litellm.get_model_info(model=model)
assert info["key"] == "au.anthropic.claude-opus-4-8", model
assert info["input_cost_per_token"] == regional["input_cost_per_token"], model
assert info["output_cost_per_token"] == regional["output_cost_per_token"], model
control = litellm.get_model_info(model="au.anthropic.claude-opus-4-8", custom_llm_provider="bedrock")
assert control["key"] == "au.anthropic.claude-opus-4-8"
def test_get_model_info_bedrock_regional_profile_without_entry_falls_back_to_base(local_model_cost_map):
"""A regional profile with no dedicated cost-map entry must still resolve to its
region-stripped base entry."""
assert "apac.anthropic.claude-opus-4-8" not in litellm.model_cost
info = litellm.get_model_info(model="bedrock/apac.anthropic.claude-opus-4-8")
assert info["key"] == "anthropic.claude-opus-4-8"
def test_get_model_info_bedrock_double_provider_prefix_resolves(local_model_cost_map):
"""A doubled bedrock/ prefix routes at runtime via strip_bedrock_routing_prefix,
so model info must resolve it to the same entry the request actually bills as."""
info = litellm.get_model_info(model="bedrock/bedrock/us.anthropic.claude-sonnet-4-6")
assert info["key"] == "us.anthropic.claude-sonnet-4-6"
def test_openai_models_in_model_info():
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="")
model_map = litellm.model_cost
violated_models = []
for model, info in model_map.items():
if (
info.get("litellm_provider") == "openai"
and info.get("supports_vision") is True
):
if info.get("supports_pdf_input") is not True:
violated_models.append(model)
assert (
len(violated_models) == 0
), f"The following models should support pdf input: {violated_models}"
def test_supports_tool_choice_simple_tests():
"""
simple sanity checks
"""
assert litellm.utils.supports_tool_choice(model="gpt-4o") == True
assert (
litellm.utils.supports_tool_choice(
model="bedrock/anthropic.claude-3-sonnet-20240229-v1:0"
)
== True
)
assert (
litellm.utils.supports_tool_choice(
model="anthropic.claude-3-sonnet-20240229-v1:0"
)
is True
)
assert (
litellm.utils.supports_tool_choice(
model="anthropic.claude-3-sonnet-20240229-v1:0",
custom_llm_provider="bedrock_converse",
)
is True
)
assert (
litellm.utils.supports_tool_choice(model="us.amazon.nova-micro-v1:0") is False
)
assert (
litellm.utils.supports_tool_choice(model="bedrock/us.amazon.nova-micro-v1:0")
is False
)
assert (
litellm.utils.supports_tool_choice(
model="us.amazon.nova-micro-v1:0", custom_llm_provider="bedrock_converse"
)
is False
)
assert litellm.utils.supports_tool_choice(model="perplexity/sonar") is False
def test_check_provider_match():
"""
Test the _check_provider_match function for various provider scenarios
"""
# Test bedrock and bedrock_converse cases
model_info = {"litellm_provider": "bedrock"}
assert litellm.utils._check_provider_match(model_info, "bedrock") is True
assert litellm.utils._check_provider_match(model_info, "bedrock_converse") is True
# Test bedrock_converse provider
model_info = {"litellm_provider": "bedrock_converse"}
assert litellm.utils._check_provider_match(model_info, "bedrock") is True
assert litellm.utils._check_provider_match(model_info, "bedrock_converse") is True
# Test non-matching provider
model_info = {"litellm_provider": "bedrock"}
assert litellm.utils._check_provider_match(model_info, "openai") is False
def test_check_provider_match_none_value_matches_any_provider():
"""
A ``litellm_provider`` of None must be treated the same as a missing
key: both mean "no provider constraint" and should match any
``custom_llm_provider``.
Regression test for https://github.com/BerriAI/litellm/issues/28336.
Before the fix, ``register_model`` persisted ``litellm_provider: None``
via ``get_model_info`` for deployments registered without a provider
(e.g. ``Router.add_deployment``), which caused ``_check_provider_match``
to drop custom pricing intermittently.
"""
# Missing key already returned True; None must behave identically.
assert litellm.utils._check_provider_match({}, "openai") is True
assert (
litellm.utils._check_provider_match({"litellm_provider": None}, "openai")
is True
)
assert (
litellm.utils._check_provider_match({"litellm_provider": None}, "anthropic")
is True
)
# When custom_llm_provider is also None nothing constrains the match.
assert litellm.utils._check_provider_match({"litellm_provider": None}, None) is True
def test_get_provider_rerank_config():
"""
Test the get_provider_rerank_config function for various providers
"""
from litellm import HostedVLLMRerankConfig
from litellm.utils import LlmProviders, ProviderConfigManager
# Test for hosted_vllm provider
config = ProviderConfigManager.get_provider_rerank_config(
"my_model", LlmProviders.HOSTED_VLLM, "http://localhost", []
)
assert isinstance(config, HostedVLLMRerankConfig)
# Models that should be skipped during testing
OLD_PROVIDERS = ["aleph_alpha", "palm"]
SKIP_MODELS = [
"azure/mistral",
"azure/command-r",
"jamba",
"deepinfra",
"mistral.",
]
# Bedrock models to block - organized by type
BEDROCK_REGIONS = ["ap-northeast-1", "eu-central-1", "us-east-1", "us-west-2"]
BEDROCK_COMMITMENTS = ["1-month-commitment", "6-month-commitment"]
BEDROCK_MODELS = {
"anthropic.claude-v1",
"anthropic.claude-v2",
"anthropic.claude-v2:1",
"anthropic.claude-instant-v1",
}
# Generate block_list dynamically
block_list = set()
for region in BEDROCK_REGIONS:
for commitment in BEDROCK_COMMITMENTS:
for model in BEDROCK_MODELS:
block_list.add(f"bedrock/{region}/{commitment}/{model}")
block_list.add(f"bedrock/{region}/{model}")
# Add Cohere models
for commitment in BEDROCK_COMMITMENTS:
block_list.add(f"bedrock/*/{commitment}/cohere.command-text-v14")
block_list.add(f"bedrock/*/{commitment}/cohere.command-light-text-v14")
print("block_list", block_list)
def test_supports_computer_use_utility():
"""
Tests the litellm.utils.supports_computer_use utility function.
"""
from litellm.utils import supports_computer_use
# Ensure LITELLM_LOCAL_MODEL_COST_MAP is set for consistent test behavior,
# as supports_computer_use relies on get_model_info.
# This also requires litellm.model_cost to be populated.
original_env_var = os.getenv("LITELLM_LOCAL_MODEL_COST_MAP")
original_model_cost = getattr(litellm, "model_cost", None)
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
litellm.model_cost = litellm.get_model_cost_map(url="") # Load with local/backup
try:
# Test a model known to support computer_use from backup JSON
supports_cu_anthropic = supports_computer_use(
model="anthropic/claude-4-sonnet-20250514"
)
assert supports_cu_anthropic is True
# Test a model known not to have the flag or set to false (defaults to False via get_model_info)
supports_cu_gpt = supports_computer_use(model="gpt-3.5-turbo")
assert supports_cu_gpt is False
finally:
# Restore original environment and model_cost to avoid side effects
if original_env_var is None:
del os.environ["LITELLM_LOCAL_MODEL_COST_MAP"]
else:
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = original_env_var
if original_model_cost is not None:
litellm.model_cost = original_model_cost
elif hasattr(litellm, "model_cost"):
delattr(litellm, "model_cost")
def test_get_model_info_shows_supports_computer_use():
"""
Tests if 'supports_computer_use' is correctly retrieved by get_model_info.
We'll use 'claude-4-sonnet-20250514' as it's configured
in the backup JSON to have supports_computer_use: True.
"""
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
# Ensure litellm.model_cost is loaded, relying on the backup mechanism if primary fails
# as per previous debugging.
litellm.model_cost = litellm.get_model_cost_map(url="")
# This model should have 'supports_computer_use': True in the backup JSON
model_known_to_support_computer_use = "claude-4-sonnet-20250514"
info = litellm.get_model_info(model_known_to_support_computer_use)
print(f"Info for {model_known_to_support_computer_use}: {info}")
# After the fix in utils.py, this should now be present and True
assert info.get("supports_computer_use") is True
# Optionally, test a model known NOT to support it, or where it's undefined (should default to False)
# For example, if "gpt-3.5-turbo" doesn't have it defined, it should be False.
model_known_not_to_support_computer_use = "gpt-3.5-turbo"
info_gpt = litellm.get_model_info(model_known_not_to_support_computer_use)
print(f"Info for {model_known_not_to_support_computer_use}: {info_gpt}")
assert (
info_gpt.get("supports_computer_use") is None
) # Expecting None due to the default in ModelInfoBase
@pytest.mark.parametrize(
"model, custom_llm_provider",
[
("gpt-3.5-turbo", "openai"),
("anthropic.claude-sonnet-4-5-20250929-v1:0", "bedrock"),
("gemini-2.5-pro", "vertex_ai"),
],
)
def test_pre_process_non_default_params(model, custom_llm_provider):
from pydantic import BaseModel
from litellm.utils import ProviderConfigManager, pre_process_non_default_params
provider_config = ProviderConfigManager.get_provider_chat_config(
model=model, provider=LlmProviders(custom_llm_provider)
)
class ResponseFormat(BaseModel):
x: str
y: str
passed_params = {
"model": "gpt-3.5-turbo",
"response_format": ResponseFormat,
}
special_params = {}
processed_non_default_params = pre_process_non_default_params(
model=model,
passed_params=passed_params,
special_params=special_params,
custom_llm_provider=custom_llm_provider,
additional_drop_params=None,
provider_config=provider_config,
)
print(processed_non_default_params)
# Vertex AI / Gemini uses Pydantic's model_json_schema() which doesn't
# include additionalProperties: False (Gemini rejects it). Other
# providers use OpenAI's to_strict_json_schema() which does.
expected_schema = {
"properties": {
"x": {"title": "X", "type": "string"},
"y": {"title": "Y", "type": "string"},
},
"required": ["x", "y"],
"title": "ResponseFormat",
"type": "object",
}
if custom_llm_provider not in ("vertex_ai", "vertex_ai_beta", "gemini"):
expected_schema["additionalProperties"] = False
assert processed_non_default_params == {
"response_format": {
"type": "json_schema",
"json_schema": {
"schema": expected_schema,
"name": "ResponseFormat",
"strict": True,
},
}
}
@pytest.mark.parametrize(
"custom_llm_provider, expected",
[
("vertex_ai", True),
("vertex_ai_beta", True),
("gdc", True),
("openai", False),
("bedrock", False),
("not_a_real_provider", False),
],
)
def test_provider_supports_vertex_params(custom_llm_provider, expected):
from litellm.utils import _provider_supports_vertex_params
assert _provider_supports_vertex_params(custom_llm_provider) is expected
@pytest.mark.parametrize(
"model, custom_llm_provider, should_keep",
[
("gemini-2.5-pro", "vertex_ai", True),
("gemini-2.5-pro", "vertex_ai_beta", True),
("gdc/gemini-2.5-flash", "gdc", True),
("gpt-4o", "openai", False),
],
)
def test_vertex_params_not_stripped_for_vertex_family(
model, custom_llm_provider, should_keep
):
optional_params = litellm.utils.get_optional_params(
model=model,
custom_llm_provider=custom_llm_provider,
vertex_project="my-project",
vertex_location="us-central1",
)
assert ("vertex_project" in optional_params) is should_keep
assert ("vertex_location" in optional_params) is should_keep
if should_keep:
assert optional_params["vertex_project"] == "my-project"
assert optional_params["vertex_location"] == "us-central1"
from litellm.utils import supports_function_calling
class TestProxyFunctionCalling:
"""Test class for proxy function calling capabilities."""
@pytest.fixture(autouse=True)
def reset_mock_cache(self):
"""Reset model cache before each test."""
from litellm.utils import _model_cache
_model_cache.flush_cache()
@pytest.mark.parametrize(
"direct_model,proxy_model,expected_result",
[
# OpenAI models
("gpt-3.5-turbo", "litellm_proxy/gpt-3.5-turbo", True),
("gpt-4", "litellm_proxy/gpt-4", True),
("gpt-4o", "litellm_proxy/gpt-4o", True),
("gpt-4o-mini", "litellm_proxy/gpt-4o-mini", True),
("gpt-4-turbo", "litellm_proxy/gpt-4-turbo", True),
("gpt-4-1106-preview", "litellm_proxy/gpt-4-1106-preview", True),
# Azure OpenAI models
("azure/gpt-4", "litellm_proxy/azure/gpt-4", True),
("azure/gpt-3.5-turbo", "litellm_proxy/azure/gpt-3.5-turbo", True),
(
"azure/gpt-4-1106-preview",
"litellm_proxy/azure/gpt-4-1106-preview",
True,
),
# Anthropic models (Claude supports function calling)
(
"claude-sonnet-4-6",
"litellm_proxy/claude-sonnet-4-6",
True,
),
# Google models
("gemini-2.5-pro", "litellm_proxy/gemini-2.5-pro", True),
("gemini/gemini-2.5-pro", "litellm_proxy/gemini/gemini-2.5-pro", True),
("gemini/gemini-2.5-flash", "litellm_proxy/gemini/gemini-2.5-flash", True),
# Groq models (mixed support)
("groq/gemma-7b-it", "litellm_proxy/groq/gemma-7b-it", True),
(
"groq/llama-3.3-70b-versatile",
"litellm_proxy/groq/llama-3.3-70b-versatile",
True,
),
# Cohere models (generally don't support function calling)
("command-nightly", "litellm_proxy/command-nightly", False),
],
)
def test_proxy_function_calling_support_consistency(
self, direct_model, proxy_model, expected_result
):
"""Test that proxy models have the same function calling support as their direct counterparts."""
direct_result = supports_function_calling(direct_model)
proxy_result = supports_function_calling(proxy_model)
# Both should match the expected result
assert (
direct_result == expected_result
), f"Direct model {direct_model} should return {expected_result}"
assert (
proxy_result == expected_result
), f"Proxy model {proxy_model} should return {expected_result}"
# Direct and proxy should be consistent
assert (
direct_result == proxy_result
), f"Mismatch: {direct_model}={direct_result} vs {proxy_model}={proxy_result}"
@pytest.mark.parametrize(
"proxy_model_name,underlying_model,expected_proxy_result",
[
# Custom model names that cannot be resolved without proxy configuration context
# These will return False because LiteLLM cannot determine the underlying model
(
"litellm_proxy/bedrock-claude-3-haiku",
"bedrock/anthropic.claude-3-haiku-20240307-v1:0",
False,
),
(
"litellm_proxy/bedrock-claude-3-sonnet",
"bedrock/anthropic.claude-3-sonnet-20240229-v1:0",
False,
),
(
"litellm_proxy/bedrock-claude-3-opus",
"bedrock/anthropic.claude-sonnet-4-5-20250929-v1:0",
False,
),
(
"litellm_proxy/bedrock-claude-instant",
"bedrock/anthropic.claude-instant-v1",
False,
),
(
"litellm_proxy/bedrock-titan-text",
"bedrock/amazon.titan-text-express-v1",
False,
),
# Azure with custom deployment names (cannot be resolved)
("litellm_proxy/my-gpt4-deployment", "azure/gpt-4", False),
("litellm_proxy/production-gpt35", "azure/gpt-3.5-turbo", False),
("litellm_proxy/dev-gpt4o", "azure/gpt-4o", False),
# Custom OpenAI deployments (cannot be resolved)
("litellm_proxy/company-gpt4", "gpt-4", False),
("litellm_proxy/internal-gpt35", "gpt-3.5-turbo", False),
# Vertex AI with custom names (cannot be resolved)
("litellm_proxy/vertex-gemini-pro", "vertex_ai/gemini-1.5-pro", False),
("litellm_proxy/vertex-gemini-flash", "vertex_ai/gemini-1.5-flash", False),
# Anthropic with custom names (cannot be resolved)
("litellm_proxy/claude-prod", "anthropic/claude-3-sonnet-20240229", False),
("litellm_proxy/claude-dev", "anthropic/claude-3-haiku-20240307", False),
# Groq with custom names (cannot be resolved)
("litellm_proxy/fast-llama", "groq/llama-3.1-8b-instant", False),
("litellm_proxy/groq-gemma", "groq/gemma-7b-it", False),
# Cohere with custom names (cannot be resolved)
("litellm_proxy/cohere-command", "cohere/command-r", False),
("litellm_proxy/cohere-command-plus", "cohere/command-r-plus", False),
# Together AI with custom names (cannot be resolved)
(
"litellm_proxy/together-llama",
"together_ai/meta-llama/Llama-2-70b-chat-hf",
False,
),
(
"litellm_proxy/together-mistral",
"together_ai/mistralai/Mistral-7B-Instruct-v0.1",
False,
),
# Ollama with custom names (cannot be resolved)
("litellm_proxy/local-llama", "ollama/llama2", False),
("litellm_proxy/local-mistral", "ollama/mistral", False),
],
)
def test_proxy_custom_model_names_without_config(
self, proxy_model_name, underlying_model, expected_proxy_result
):
"""
Test proxy models with custom model names that differ from underlying models.
Without proxy configuration context, LiteLLM cannot resolve custom model names
to their underlying models, so these will return False.
This demonstrates the limitation and documents the expected behavior.
"""
# Test the underlying model directly first to establish what it SHOULD return
try:
underlying_result = supports_function_calling(underlying_model)
print(
f"Underlying model {underlying_model} supports function calling: {underlying_result}"
)
except Exception as e:
print(f"Warning: Could not test underlying model {underlying_model}: {e}")
# Test the proxy model - this will return False due to lack of configuration context
proxy_result = supports_function_calling(proxy_model_name)
assert (
proxy_result == expected_proxy_result
), f"Proxy model {proxy_model_name} should return {expected_proxy_result} (without config context)"
def test_proxy_model_resolution_with_custom_names_documentation(self):
"""
Document the behavior and limitation for custom proxy model names.
This test demonstrates:
1. The current limitation with custom model names
2. How the proxy server would handle this in production
3. The expected behavior for both scenarios
"""
# Case 1: Custom model name that cannot be resolved
custom_model = "litellm_proxy/my-custom-claude"
result = supports_function_calling(custom_model)
assert (
result is False
), "Custom model names return False without proxy config context"
# Case 2: Model name that can be resolved (matches pattern)
resolvable_model = "litellm_proxy/claude-sonnet-4-5-20250929"
result = supports_function_calling(resolvable_model)
assert result is True, "Resolvable model names work with fallback logic"
# Documentation notes:
print("""
PROXY MODEL RESOLUTION BEHAVIOR:
✅ WORKS (with current fallback logic):
- litellm_proxy/gpt-4
- litellm_proxy/claude-sonnet-4-5-20250929
- litellm_proxy/anthropic/claude-3-haiku-20240307
❌ DOESN'T WORK (requires proxy server config):
- litellm_proxy/my-custom-gpt4
- litellm_proxy/bedrock-claude-3-haiku
- litellm_proxy/production-model
💡 SOLUTION: Use LiteLLM proxy server with proper model_list configuration
that maps custom names to underlying models.
""")
@pytest.mark.parametrize(
"proxy_model_with_hints,expected_result",
[
# These are proxy models where we can infer the underlying model from the name
("litellm_proxy/gpt-4-with-functions", True), # Hints at GPT-4
("litellm_proxy/claude-3-haiku-prod", True), # Hints at Claude 3 Haiku
(
"litellm_proxy/bedrock-anthropic-claude-3-sonnet",
True,
), # Hints at Bedrock Claude 3 Sonnet
],
)
def test_proxy_models_with_naming_hints(
self, proxy_model_with_hints, expected_result
):
"""
Test proxy models with names that provide hints about the underlying model.
Note: These will currently fail because the hint-based resolution isn't implemented yet,
but they demonstrate what could be possible with enhanced model name inference.
"""
# This test documents potential future enhancement
proxy_result = supports_function_calling(proxy_model_with_hints)
# Currently these will return False, but we document the expected behavior
# In the future, we could implement smarter model name inference
print(
f"Model {proxy_model_with_hints}: current={proxy_result}, desired={expected_result}"
)
# For now, we expect False (current behavior), but document the limitation
assert (
proxy_result is False
), f"Current limitation: {proxy_model_with_hints} returns False without inference"
@pytest.mark.parametrize(
"proxy_model,expected_result",
[
# Test specific proxy models that should support function calling
("litellm_proxy/gpt-3.5-turbo", True),
("litellm_proxy/gpt-4", True),
("litellm_proxy/gpt-4o", True),
("litellm_proxy/claude-sonnet-4-6", True),
("litellm_proxy/gemini/gemini-2.5-pro", True),
# Test proxy models that should not support function calling
("litellm_proxy/command-nightly", False),
("litellm_proxy/anthropic.claude-instant-v1", False),
],
)
def test_proxy_only_function_calling_support(self, proxy_model, expected_result):
"""
Test proxy models independently to ensure they report correct function calling support.
This test focuses on proxy models without comparing to direct models,
useful for cases where we only care about the proxy behavior.
"""
try:
result = supports_function_calling(model=proxy_model)
assert (
result == expected_result
), f"Proxy model {proxy_model} returned {result}, expected {expected_result}"
except Exception as e:
pytest.fail(f"Error testing proxy model {proxy_model}: {e}")
def test_litellm_utils_supports_function_calling_import(self):
"""Test that supports_function_calling can be imported from litellm.utils."""
try:
from litellm.utils import supports_function_calling
assert callable(supports_function_calling)
except ImportError as e:
pytest.fail(f"Failed to import supports_function_calling: {e}")
def test_litellm_supports_function_calling_import(self):
"""Test that supports_function_calling can be imported from litellm directly."""
try:
import litellm
assert hasattr(litellm, "supports_function_calling")
assert callable(litellm.supports_function_calling)
except Exception as e:
pytest.fail(f"Failed to access litellm.supports_function_calling: {e}")
@pytest.mark.parametrize(
"model_name",
[
"litellm_proxy/gpt-3.5-turbo",
"litellm_proxy/gpt-4",
"litellm_proxy/claude-sonnet-4-6",
"litellm_proxy/gemini/gemini-2.5-pro",
],
)
def test_proxy_model_with_custom_llm_provider_none(self, model_name):
"""
Test proxy models with custom_llm_provider=None parameter.
This tests the supports_function_calling function with the custom_llm_provider
parameter explicitly set to None, which is a common usage pattern.
"""
try:
result = supports_function_calling(
model=model_name, custom_llm_provider=None
)
# All the models in this test should support function calling
assert (
result is True
), f"Model {model_name} should support function calling but returned {result}"
except Exception as e:
pytest.fail(
f"Error testing {model_name} with custom_llm_provider=None: {e}"
)
def test_edge_cases_and_malformed_proxy_models(self):
"""Test edge cases and malformed proxy model names."""
test_cases = [
("litellm_proxy/", False), # Empty model name after proxy prefix
("litellm_proxy", False), # Just the proxy prefix without slash
("litellm_proxy//gpt-3.5-turbo", False), # Double slash
("litellm_proxy/nonexistent-model", False), # Non-existent model
]
for model_name, expected_result in test_cases:
try:
result = supports_function_calling(model=model_name)
# For malformed models, we expect False or the function to handle gracefully
assert (
result == expected_result
), f"Edge case {model_name} returned {result}, expected {expected_result}"
except Exception:
# It's acceptable for malformed model names to raise exceptions
# rather than returning False, as long as they're handled gracefully
pass
def test_proxy_model_resolution_demonstration(self):
"""
Demonstration test showing the current issue with proxy model resolution.
This test documents the current behavior and can be used to verify
when the issue is fixed.
"""
direct_model = "gpt-3.5-turbo"
proxy_model = "litellm_proxy/gpt-3.5-turbo"
direct_result = supports_function_calling(model=direct_model)
proxy_result = supports_function_calling(model=proxy_model)
print(f"\nDemonstration of proxy model resolution:")
print(
f"Direct model '{direct_model}' supports function calling: {direct_result}"
)
print(f"Proxy model '{proxy_model}' supports function calling: {proxy_result}")
# This assertion will currently fail due to the bug
# When the bug is fixed, this test should pass
if direct_result != proxy_result:
pytest.skip(
f"Known issue: Proxy model resolution inconsistency. "
f"Direct: {direct_result}, Proxy: {proxy_result}. "
f"This test will pass when the issue is resolved."
)
assert direct_result == proxy_result, (
f"Proxy model resolution issue: {direct_model} -> {direct_result}, "
f"{proxy_model} -> {proxy_result}"
)
@pytest.mark.parametrize(
"proxy_model_name,underlying_bedrock_model,expected_proxy_result,description",
[
# Bedrock Converse API mappings - these are the real-world scenarios
(
"litellm_proxy/bedrock-claude-3-haiku",
"bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0",
False,
"Bedrock Claude 3 Haiku via Converse API",
),
(
"litellm_proxy/bedrock-claude-3-sonnet",
"bedrock/converse/anthropic.claude-3-sonnet-20240229-v1:0",
False,
"Bedrock Claude 3 Sonnet via Converse API",
),
(
"litellm_proxy/bedrock-claude-3-opus",
"bedrock/converse/anthropic.claude-sonnet-4-5-20250929-v1:0",
False,
"Bedrock Claude 3 Opus via Converse API",
),
(
"litellm_proxy/bedrock-claude-3-5-sonnet",
"bedrock/converse/anthropic.claude-haiku-4-5-20251001-v1:0",
False,
"Bedrock Claude 3.5 Sonnet via Converse API",
),
# Bedrock Legacy API mappings (non-converse)
(
"litellm_proxy/bedrock-claude-instant",
"bedrock/anthropic.claude-instant-v1",
False,
"Bedrock Claude Instant Legacy API",
),
(
"litellm_proxy/bedrock-claude-v2",
"bedrock/anthropic.claude-v2",
False,
"Bedrock Claude v2 Legacy API",
),
(
"litellm_proxy/bedrock-claude-v2-1",
"bedrock/anthropic.claude-v2:1",
False,
"Bedrock Claude v2.1 Legacy API",
),
# Bedrock other model providers via Converse API
(
"litellm_proxy/bedrock-titan-text",
"bedrock/converse/amazon.titan-text-express-v1",
False,
"Bedrock Titan Text Express via Converse API",
),
(
"litellm_proxy/bedrock-titan-text-premier",
"bedrock/converse/amazon.titan-text-premier-v1:0",
False,
"Bedrock Titan Text Premier via Converse API",
),
(
"litellm_proxy/bedrock-llama3-8b",
"bedrock/converse/meta.llama3-8b-instruct-v1:0",
False,
"Bedrock Llama 3 8B via Converse API",
),
(
"litellm_proxy/bedrock-llama3-70b",
"bedrock/converse/meta.llama3-70b-instruct-v1:0",
False,
"Bedrock Llama 3 70B via Converse API",
),
(
"litellm_proxy/bedrock-mistral-7b",
"bedrock/converse/mistral.mistral-7b-instruct-v0:2",
False,
"Bedrock Mistral 7B via Converse API",
),
(
"litellm_proxy/bedrock-mistral-8x7b",
"bedrock/converse/mistral.mixtral-8x7b-instruct-v0:1",
False,
"Bedrock Mistral 8x7B via Converse API",
),
(
"litellm_proxy/bedrock-mistral-large",
"bedrock/converse/mistral.mistral-large-2402-v1:0",
False,
"Bedrock Mistral Large via Converse API",
),
# Company-specific naming patterns (real-world examples)
(
"litellm_proxy/prod-claude-haiku",
"bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0",
False,
"Production Claude Haiku",
),
(
"litellm_proxy/dev-claude-sonnet",
"bedrock/converse/anthropic.claude-3-sonnet-20240229-v1:0",
False,
"Development Claude Sonnet",
),
(
"litellm_proxy/staging-claude-opus",
"bedrock/converse/anthropic.claude-sonnet-4-5-20250929-v1:0",
False,
"Staging Claude Opus",
),
(
"litellm_proxy/cost-optimized-claude",
"bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0",
False,
"Cost-optimized Claude deployment",
),
(
"litellm_proxy/high-performance-claude",
"bedrock/converse/anthropic.claude-sonnet-4-5-20250929-v1:0",
False,
"High-performance Claude deployment",
),
# Regional deployment examples
(
"litellm_proxy/us-east-claude",
"bedrock/converse/anthropic.claude-3-sonnet-20240229-v1:0",
False,
"US East Claude deployment",
),
(
"litellm_proxy/eu-west-claude",
"bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0",
False,
"EU West Claude deployment",
),
(
"litellm_proxy/ap-south-llama",
"bedrock/converse/meta.llama3-70b-instruct-v1:0",
False,
"Asia Pacific Llama deployment",
),
],
)
def test_bedrock_converse_api_proxy_mappings(
self,
proxy_model_name,
underlying_bedrock_model,
expected_proxy_result,
description,
):
"""
Test real-world Bedrock Converse API proxy model mappings.
This test covers the specific scenario where proxy model names like
'bedrock-claude-3-haiku' map to underlying Bedrock Converse API models like
'bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0'.
These mappings are typically defined in proxy server configuration files
and cannot be resolved by LiteLLM without that context.
"""
print(f"\nTesting: {description}")
print(f" Proxy model: {proxy_model_name}")
print(f" Underlying model: {underlying_bedrock_model}")
# Test the underlying model directly to verify it supports function calling
try:
underlying_result = supports_function_calling(underlying_bedrock_model)
print(f" Underlying model function calling support: {underlying_result}")
# Most Bedrock Converse API models with Anthropic Claude should support function calling
if "anthropic.claude-3" in underlying_bedrock_model:
assert (
underlying_result is True
), f"Claude 3 models should support function calling: {underlying_bedrock_model}"
except Exception as e:
print(
f" Warning: Could not test underlying model {underlying_bedrock_model}: {e}"
)
# Test the proxy model - should return False due to lack of configuration context
proxy_result = supports_function_calling(proxy_model_name)
print(f" Proxy model function calling support: {proxy_result}")
assert proxy_result == expected_proxy_result, (
f"Proxy model {proxy_model_name} should return {expected_proxy_result} "
f"(without config context). Description: {description}"
)
def test_real_world_proxy_config_documentation(self):
"""
Document how real-world proxy configurations would handle model mappings.
This test provides documentation on how the proxy server configuration
would typically map custom model names to underlying models.
"""
print("""
REAL-WORLD PROXY SERVER CONFIGURATION EXAMPLE:
===============================================
In a proxy_server_config.yaml file, you would define:
model_list:
- model_name: bedrock-claude-3-haiku
litellm_params:
model: bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
aws_region_name: us-east-1
- model_name: bedrock-claude-3-sonnet
litellm_params:
model: bedrock/converse/anthropic.claude-3-sonnet-20240229-v1:0
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
aws_region_name: us-east-1
- model_name: prod-claude-haiku
litellm_params:
model: bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0
aws_access_key_id: os.environ/PROD_AWS_ACCESS_KEY_ID
aws_secret_access_key: os.environ/PROD_AWS_SECRET_ACCESS_KEY
aws_region_name: us-west-2
FUNCTION CALLING WITH PROXY SERVER:
===================================
When using the proxy server with this configuration:
1. Client calls: supports_function_calling("bedrock-claude-3-haiku")
2. Proxy server resolves to: bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0
3. LiteLLM evaluates the underlying model's capabilities
4. Returns: True (because Claude 3 Haiku supports function calling)
Without the proxy server configuration context, LiteLLM cannot resolve
the custom model name and returns False.
BEDROCK CONVERSE API BENEFITS:
==============================
The Bedrock Converse API provides:
- Standardized function calling interface across providers
- Better tool use capabilities compared to legacy APIs
- Consistent request/response format
- Enhanced streaming support for function calls
""")
# Verify that direct underlying models work as expected
bedrock_models = [
"bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0",
"bedrock/converse/anthropic.claude-3-sonnet-20240229-v1:0",
"bedrock/converse/anthropic.claude-sonnet-4-5-20250929-v1:0",
]
for model in bedrock_models:
try:
result = supports_function_calling(model)
print(f"Direct test - {model}: {result}")
# Claude 3 models should support function calling
assert (
result is True
), f"Claude 3 model should support function calling: {model}"
except Exception as e:
print(f"Could not test {model}: {e}")
@pytest.mark.parametrize(
"proxy_model_name,underlying_bedrock_model,expected_proxy_result,description",
[
# Bedrock Converse API mappings - these are the real-world scenarios
(
"litellm_proxy/bedrock-claude-3-haiku",
"bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0",
False,
"Bedrock Claude 3 Haiku via Converse API",
),
(
"litellm_proxy/bedrock-claude-3-sonnet",
"bedrock/converse/anthropic.claude-3-sonnet-20240229-v1:0",
False,
"Bedrock Claude 3 Sonnet via Converse API",
),
(
"litellm_proxy/bedrock-claude-3-opus",
"bedrock/converse/anthropic.claude-sonnet-4-5-20250929-v1:0",
False,
"Bedrock Claude 3 Opus via Converse API",
),
(
"litellm_proxy/bedrock-claude-3-5-sonnet",
"bedrock/converse/anthropic.claude-haiku-4-5-20251001-v1:0",
False,
"Bedrock Claude 3.5 Sonnet via Converse API",
),
# Bedrock Legacy API mappings (non-converse)
(
"litellm_proxy/bedrock-claude-instant",
"bedrock/anthropic.claude-instant-v1",
False,
"Bedrock Claude Instant Legacy API",
),
(
"litellm_proxy/bedrock-claude-v2",
"bedrock/anthropic.claude-v2",
False,
"Bedrock Claude v2 Legacy API",
),
(
"litellm_proxy/bedrock-claude-v2-1",
"bedrock/anthropic.claude-v2:1",
False,
"Bedrock Claude v2.1 Legacy API",
),
# Bedrock other model providers via Converse API
(
"litellm_proxy/bedrock-titan-text",
"bedrock/converse/amazon.titan-text-express-v1",
False,
"Bedrock Titan Text Express via Converse API",
),
(
"litellm_proxy/bedrock-titan-text-premier",
"bedrock/converse/amazon.titan-text-premier-v1:0",
False,
"Bedrock Titan Text Premier via Converse API",
),
(
"litellm_proxy/bedrock-llama3-8b",
"bedrock/converse/meta.llama3-8b-instruct-v1:0",
False,
"Bedrock Llama 3 8B via Converse API",
),
(
"litellm_proxy/bedrock-llama3-70b",
"bedrock/converse/meta.llama3-70b-instruct-v1:0",
False,
"Bedrock Llama 3 70B via Converse API",
),
(
"litellm_proxy/bedrock-mistral-7b",
"bedrock/converse/mistral.mistral-7b-instruct-v0:2",
False,
"Bedrock Mistral 7B via Converse API",
),
(
"litellm_proxy/bedrock-mistral-8x7b",
"bedrock/converse/mistral.mixtral-8x7b-instruct-v0:1",
False,
"Bedrock Mistral 8x7B via Converse API",
),
(
"litellm_proxy/bedrock-mistral-large",
"bedrock/converse/mistral.mistral-large-2402-v1:0",
False,
"Bedrock Mistral Large via Converse API",
),
# Company-specific naming patterns (real-world examples)
(
"litellm_proxy/prod-claude-haiku",
"bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0",
False,
"Production Claude Haiku",
),
(
"litellm_proxy/dev-claude-sonnet",
"bedrock/converse/anthropic.claude-3-sonnet-20240229-v1:0",
False,
"Development Claude Sonnet",
),
(
"litellm_proxy/staging-claude-opus",
"bedrock/converse/anthropic.claude-sonnet-4-5-20250929-v1:0",
False,
"Staging Claude Opus",
),
(
"litellm_proxy/cost-optimized-claude",
"bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0",
False,
"Cost-optimized Claude deployment",
),
(
"litellm_proxy/high-performance-claude",
"bedrock/converse/anthropic.claude-sonnet-4-5-20250929-v1:0",
False,
"High-performance Claude deployment",
),
# Regional deployment examples
(
"litellm_proxy/us-east-claude",
"bedrock/converse/anthropic.claude-3-sonnet-20240229-v1:0",
False,
"US East Claude deployment",
),
(
"litellm_proxy/eu-west-claude",
"bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0",
False,
"EU West Claude deployment",
),
(
"litellm_proxy/ap-south-llama",
"bedrock/converse/meta.llama3-70b-instruct-v1:0",
False,
"Asia Pacific Llama deployment",
),
],
)
def test_bedrock_converse_api_proxy_mappings(
self,
proxy_model_name,
underlying_bedrock_model,
expected_proxy_result,
description,
):
"""
Test real-world Bedrock Converse API proxy model mappings.
This test covers the specific scenario where proxy model names like
'bedrock-claude-3-haiku' map to underlying Bedrock Converse API models like
'bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0'.
These mappings are typically defined in proxy server configuration files
and cannot be resolved by LiteLLM without that context.
"""
print(f"\nTesting: {description}")
print(f" Proxy model: {proxy_model_name}")
print(f" Underlying model: {underlying_bedrock_model}")
# Test the underlying model directly to verify it supports function calling
try:
underlying_result = supports_function_calling(underlying_bedrock_model)
print(f" Underlying model function calling support: {underlying_result}")
# Most Bedrock Converse API models with Anthropic Claude should support function calling
if "anthropic.claude-3" in underlying_bedrock_model:
assert (
underlying_result is True
), f"Claude 3 models should support function calling: {underlying_bedrock_model}"
except Exception as e:
print(
f" Warning: Could not test underlying model {underlying_bedrock_model}: {e}"
)
# Test the proxy model - should return False due to lack of configuration context
proxy_result = supports_function_calling(proxy_model_name)
print(f" Proxy model function calling support: {proxy_result}")
assert proxy_result == expected_proxy_result, (
f"Proxy model {proxy_model_name} should return {expected_proxy_result} "
f"(without config context). Description: {description}"
)
def test_real_world_proxy_config_documentation(self):
"""
Document how real-world proxy configurations would handle model mappings.
This test provides documentation on how the proxy server configuration
would typically map custom model names to underlying models.
"""
print("""
REAL-WORLD PROXY SERVER CONFIGURATION EXAMPLE:
===============================================
In a proxy_server_config.yaml file, you would define:
model_list:
- model_name: bedrock-claude-3-haiku
litellm_params:
model: bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
aws_region_name: us-east-1
- model_name: bedrock-claude-3-sonnet
litellm_params:
model: bedrock/converse/anthropic.claude-3-sonnet-20240229-v1:0
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
aws_region_name: us-east-1
- model_name: prod-claude-haiku
litellm_params:
model: bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0
aws_access_key_id: os.environ/PROD_AWS_ACCESS_KEY_ID
aws_secret_access_key: os.environ/PROD_AWS_SECRET_ACCESS_KEY
aws_region_name: us-west-2
FUNCTION CALLING WITH PROXY SERVER:
===================================
When using the proxy server with this configuration:
1. Client calls: supports_function_calling("bedrock-claude-3-haiku")
2. Proxy server resolves to: bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0
3. LiteLLM evaluates the underlying model's capabilities
4. Returns: True (because Claude 3 Haiku supports function calling)
Without the proxy server configuration context, LiteLLM cannot resolve
the custom model name and returns False.
BEDROCK CONVERSE API BENEFITS:
==============================
The Bedrock Converse API provides:
- Standardized function calling interface across providers
- Better tool use capabilities compared to legacy APIs
- Consistent request/response format
- Enhanced streaming support for function calls
""")
# Verify that direct underlying models work as expected
bedrock_models = [
"bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0",
"bedrock/converse/anthropic.claude-3-sonnet-20240229-v1:0",
"bedrock/converse/anthropic.claude-sonnet-4-5-20250929-v1:0",
]
for model in bedrock_models:
try:
result = supports_function_calling(model)
print(f"Direct test - {model}: {result}")
# Claude 3 models should support function calling
assert (
result is True
), f"Claude 3 model should support function calling: {model}"
except Exception as e:
print(f"Could not test {model}: {e}")
@pytest.mark.parametrize(
"proxy_model_name,underlying_bedrock_model,expected_proxy_result,description",
[
# Bedrock Converse API mappings - these are the real-world scenarios
(
"litellm_proxy/bedrock-claude-3-haiku",
"bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0",
False,
"Bedrock Claude 3 Haiku via Converse API",
),
(
"litellm_proxy/bedrock-claude-3-sonnet",
"bedrock/converse/anthropic.claude-3-sonnet-20240229-v1:0",
False,
"Bedrock Claude 3 Sonnet via Converse API",
),
(
"litellm_proxy/bedrock-claude-3-opus",
"bedrock/converse/anthropic.claude-sonnet-4-5-20250929-v1:0",
False,
"Bedrock Claude 3 Opus via Converse API",
),
(
"litellm_proxy/bedrock-claude-3-5-sonnet",
"bedrock/converse/anthropic.claude-haiku-4-5-20251001-v1:0",
False,
"Bedrock Claude 3.5 Sonnet via Converse API",
),
# Bedrock Legacy API mappings (non-converse)
(
"litellm_proxy/bedrock-claude-instant",
"bedrock/anthropic.claude-instant-v1",
False,
"Bedrock Claude Instant Legacy API",
),
(
"litellm_proxy/bedrock-claude-v2",
"bedrock/anthropic.claude-v2",
False,
"Bedrock Claude v2 Legacy API",
),
(
"litellm_proxy/bedrock-claude-v2-1",
"bedrock/anthropic.claude-v2:1",
False,
"Bedrock Claude v2.1 Legacy API",
),
# Bedrock other model providers via Converse API
(
"litellm_proxy/bedrock-titan-text",
"bedrock/converse/amazon.titan-text-express-v1",
False,
"Bedrock Titan Text Express via Converse API",
),
(
"litellm_proxy/bedrock-titan-text-premier",
"bedrock/converse/amazon.titan-text-premier-v1:0",
False,
"Bedrock Titan Text Premier via Converse API",
),
(
"litellm_proxy/bedrock-llama3-8b",
"bedrock/converse/meta.llama3-8b-instruct-v1:0",
False,
"Bedrock Llama 3 8B via Converse API",
),
(
"litellm_proxy/bedrock-llama3-70b",
"bedrock/converse/meta.llama3-70b-instruct-v1:0",
False,
"Bedrock Llama 3 70B via Converse API",
),
(
"litellm_proxy/bedrock-mistral-7b",
"bedrock/converse/mistral.mistral-7b-instruct-v0:2",
False,
"Bedrock Mistral 7B via Converse API",
),
(
"litellm_proxy/bedrock-mistral-8x7b",
"bedrock/converse/mistral.mixtral-8x7b-instruct-v0:1",
False,
"Bedrock Mistral 8x7B via Converse API",
),
(
"litellm_proxy/bedrock-mistral-large",
"bedrock/converse/mistral.mistral-large-2402-v1:0",
False,
"Bedrock Mistral Large via Converse API",
),
# Company-specific naming patterns (real-world examples)
(
"litellm_proxy/prod-claude-haiku",
"bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0",
False,
"Production Claude Haiku",
),
(
"litellm_proxy/dev-claude-sonnet",
"bedrock/converse/anthropic.claude-3-sonnet-20240229-v1:0",
False,
"Development Claude Sonnet",
),
(
"litellm_proxy/staging-claude-opus",
"bedrock/converse/anthropic.claude-sonnet-4-5-20250929-v1:0",
False,
"Staging Claude Opus",
),
(
"litellm_proxy/cost-optimized-claude",
"bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0",
False,
"Cost-optimized Claude deployment",
),
(
"litellm_proxy/high-performance-claude",
"bedrock/converse/anthropic.claude-sonnet-4-5-20250929-v1:0",
False,
"High-performance Claude deployment",
),
# Regional deployment examples
(
"litellm_proxy/us-east-claude",
"bedrock/converse/anthropic.claude-3-sonnet-20240229-v1:0",
False,
"US East Claude deployment",
),
(
"litellm_proxy/eu-west-claude",
"bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0",
False,
"EU West Claude deployment",
),
(
"litellm_proxy/ap-south-llama",
"bedrock/converse/meta.llama3-70b-instruct-v1:0",
False,
"Asia Pacific Llama deployment",
),
],
)
def test_bedrock_converse_api_proxy_mappings(
self,
proxy_model_name,
underlying_bedrock_model,
expected_proxy_result,
description,
):
"""
Test real-world Bedrock Converse API proxy model mappings.
This test covers the specific scenario where proxy model names like
'bedrock-claude-3-haiku' map to underlying Bedrock Converse API models like
'bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0'.
These mappings are typically defined in proxy server configuration files
and cannot be resolved by LiteLLM without that context.
"""
print(f"\nTesting: {description}")
print(f" Proxy model: {proxy_model_name}")
print(f" Underlying model: {underlying_bedrock_model}")
# Test the underlying model directly to verify it supports function calling
try:
underlying_result = supports_function_calling(underlying_bedrock_model)
print(f" Underlying model function calling support: {underlying_result}")
# Most Bedrock Converse API models with Anthropic Claude should support function calling
if "anthropic.claude-3" in underlying_bedrock_model:
assert (
underlying_result is True
), f"Claude 3 models should support function calling: {underlying_bedrock_model}"
except Exception as e:
print(
f" Warning: Could not test underlying model {underlying_bedrock_model}: {e}"
)
# Test the proxy model - should return False due to lack of configuration context
proxy_result = supports_function_calling(proxy_model_name)
print(f" Proxy model function calling support: {proxy_result}")
assert proxy_result == expected_proxy_result, (
f"Proxy model {proxy_model_name} should return {expected_proxy_result} "
f"(without config context). Description: {description}"
)
def test_real_world_proxy_config_documentation(self):
"""
Document how real-world proxy configurations would handle model mappings.
This test provides documentation on how the proxy server configuration
would typically map custom model names to underlying models.
"""
print("""
REAL-WORLD PROXY SERVER CONFIGURATION EXAMPLE:
===============================================
In a proxy_server_config.yaml file, you would define:
model_list:
- model_name: bedrock-claude-3-haiku
litellm_params:
model: bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
aws_region_name: us-east-1
- model_name: bedrock-claude-3-sonnet
litellm_params:
model: bedrock/converse/anthropic.claude-3-sonnet-20240229-v1:0
aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
aws_region_name: us-east-1
- model_name: prod-claude-haiku
litellm_params:
model: bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0
aws_access_key_id: os.environ/PROD_AWS_ACCESS_KEY_ID
aws_secret_access_key: os.environ/PROD_AWS_SECRET_ACCESS_KEY
aws_region_name: us-west-2
FUNCTION CALLING WITH PROXY SERVER:
===================================
When using the proxy server with this configuration:
1. Client calls: supports_function_calling("bedrock-claude-3-haiku")
2. Proxy server resolves to: bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0
3. LiteLLM evaluates the underlying model's capabilities
4. Returns: True (because Claude 3 Haiku supports function calling)
Without the proxy server configuration context, LiteLLM cannot resolve
the custom model name and returns False.
BEDROCK CONVERSE API BENEFITS:
==============================
The Bedrock Converse API provides:
- Standardized function calling interface across providers
- Better tool use capabilities compared to legacy APIs
- Consistent request/response format
- Enhanced streaming support for function calls
""")
# Verify that direct underlying models work as expected
bedrock_models = [
"bedrock/converse/anthropic.claude-3-haiku-20240307-v1:0",
"bedrock/converse/anthropic.claude-3-sonnet-20240229-v1:0",
"bedrock/converse/anthropic.claude-sonnet-4-5-20250929-v1:0",
]
for model in bedrock_models:
try:
result = supports_function_calling(model)
print(f"Direct test - {model}: {result}")
# Claude 3 models should support function calling
assert (
result is True
), f"Claude 3 model should support function calling: {model}"
except Exception as e:
print(f"Could not test {model}: {e}")
def test_register_model_with_scientific_notation():
"""
Test that the register_model function can handle scientific notation in the model name.
"""
import uuid
# Use a truly unique model name with uuid to avoid conflicts when tests run in parallel
test_model_name = f"test-scientific-notation-model-{uuid.uuid4().hex[:12]}"
# Clear LRU caches that might have stale data
from litellm.utils import (
_invalidate_model_cost_lowercase_map,
)
_invalidate_model_cost_lowercase_map()
model_cost_dict = {
test_model_name: {
"max_tokens": 8192,
"input_cost_per_token": "3e-07",
"output_cost_per_token": "6e-07",
"litellm_provider": "openai",
"mode": "chat",
},
}
litellm.register_model(model_cost_dict)
registered_model = litellm.model_cost[test_model_name]
print(registered_model)
assert registered_model["input_cost_per_token"] == 3e-07
assert registered_model["output_cost_per_token"] == 6e-07
assert registered_model["litellm_provider"] == "openai"
assert registered_model["mode"] == "chat"
# Clean up after test
if test_model_name in litellm.model_cost:
del litellm.model_cost[test_model_name]
_invalidate_model_cost_lowercase_map()
def test_register_model_openrouter_without_slash():
"""
Test that register_model handles openrouter models without '/' in the name.
Fixes https://github.com/BerriAI/litellm/issues/18936
Previously, the code did `split_string[1]` which would fail with IndexError
when the model name didn't contain '/'. Now it uses `split_string[-1]` which
always works.
"""
# Clear any existing entries
litellm.openrouter_models.discard("my-custom-alias")
litellm.openrouter_models.discard("gpt-4")
litellm.openrouter_models.discard("openai/gpt-4")
# Test 1: Model name without '/' (this was the bug - would raise IndexError)
litellm.register_model(
{
"my-custom-alias": {
"max_tokens": 8192,
"input_cost_per_token": 0.00001,
"output_cost_per_token": 0.00002,
"litellm_provider": "openrouter",
"mode": "chat",
},
}
)
assert "my-custom-alias" in litellm.openrouter_models
# Test 2: Model name with single '/' (openrouter/model format)
litellm.register_model(
{
"openrouter/gpt-4": {
"max_tokens": 8192,
"input_cost_per_token": 0.00001,
"output_cost_per_token": 0.00002,
"litellm_provider": "openrouter",
"mode": "chat",
},
}
)
assert "gpt-4" in litellm.openrouter_models
# Test 3: Model name with double '/' (openrouter/provider/model format)
litellm.register_model(
{
"openrouter/openai/gpt-4-turbo": {
"max_tokens": 8192,
"input_cost_per_token": 0.00001,
"output_cost_per_token": 0.00002,
"litellm_provider": "openrouter",
"mode": "chat",
},
}
)
assert "openai/gpt-4-turbo" in litellm.openrouter_models
def test_reasoning_content_preserved_in_text_completion_wrapper():
"""Ensure reasoning_content is copied from delta to text_choices."""
chunk = ModelResponseStream(
id="test-id",
created=1234567890,
model="test-model",
object="chat.completion.chunk",
choices=[
StreamingChoices(
finish_reason=None,
index=0,
delta=Delta(
content="Some answer text",
role="assistant",
reasoning_content="Here's my chain of thought...",
),
)
],
)
wrapper = TextCompletionStreamWrapper(
completion_stream=None, # Not used in convert_to_text_completion_object
model="test-model",
stream_options=None,
)
transformed = wrapper.convert_to_text_completion_object(chunk)
assert "choices" in transformed
assert len(transformed["choices"]) == 1
choice = transformed["choices"][0]
assert choice["text"] == "Some answer text"
assert choice["reasoning_content"] == "Here's my chain of thought..."
def test_anthropic_claude_4_invoke_chat_provider_config():
"""Test that the Anthropic Claude 4 Invoke chat provider config is correct."""
from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import (
AmazonAnthropicClaudeConfig,
)
from litellm.utils import ProviderConfigManager
config = ProviderConfigManager.get_provider_chat_config(
model="invoke/us.anthropic.claude-sonnet-4-20250514-v1:0",
provider=LlmProviders.BEDROCK,
)
print(config)
assert isinstance(config, AmazonAnthropicClaudeConfig)
def test_bedrock_application_inference_profile():
model = "arn:aws:bedrock:us-east-2:<AWS-ACCOUNT-ID>:inference-profile/us.anthropic.claude-3-5-haiku-20241022-v1:0"
from pydantic import BaseModel
from litellm import completion
from litellm.utils import supports_tool_choice
result = supports_tool_choice(model, custom_llm_provider="bedrock")
result_2 = supports_tool_choice(model, custom_llm_provider="bedrock_converse")
print(result)
assert result == result_2
assert result is True
def test_image_response_utils():
"""Test that the image response utils are correct."""
from litellm.utils import ImageResponse
result = {
"created": None,
"data": [
{
"b64_json": "/9j/.../2Q==",
"revised_prompt": None,
"url": None,
"timings": {"inference": 0.9612685777246952},
"index": 0,
}
],
"id": "91559891cxxx-PDX",
"model": "black-forest-labs/FLUX.1-schnell-Free",
"object": "list",
"hidden_params": {"additional_headers": {}},
}
image_response = ImageResponse(**result)
def test_is_valid_api_key():
import hashlib
# Valid sk- keys
assert is_valid_api_key("sk-abc123")
assert is_valid_api_key("sk-ABC_123-xyz")
# Valid hashed key (64 hex chars)
assert is_valid_api_key("a" * 64)
assert is_valid_api_key("0123456789abcdef" * 4) # 16*4 = 64
# Real SHA-256 hash
real_hash = hashlib.sha256(b"my_secret_key").hexdigest()
assert len(real_hash) == 64
assert is_valid_api_key(real_hash)
# Invalid: too short
assert not is_valid_api_key("sk-")
assert not is_valid_api_key("")
# Invalid: too long
assert not is_valid_api_key("sk-" + "a" * 200)
# Invalid: wrong prefix
assert not is_valid_api_key("pk-abc123")
# Invalid: wrong chars in sk- key
assert not is_valid_api_key("sk-abc$%#@!")
# Invalid: not a string
assert not is_valid_api_key(None)
assert not is_valid_api_key(12345)
# Invalid: wrong length for hash
assert not is_valid_api_key("a" * 63)
assert not is_valid_api_key("a" * 65)
def test_block_key_hashing_logic():
"""
Test that block_key() function only hashes keys that start with "sk-"
"""
import hashlib
from litellm.proxy.utils import hash_token
# Test cases: (input_key, should_be_hashed, expected_output)
test_cases = [
("sk-1234567890abcdef", True, hash_token("sk-1234567890abcdef")),
("sk-test-key", True, hash_token("sk-test-key")),
("abc123", False, "abc123"), # Should not be hashed
("hashed_key_123", False, "hashed_key_123"), # Should not be hashed
("", False, ""), # Empty string should not be hashed
("sk-", True, hash_token("sk-")), # Edge case: just "sk-"
]
for input_key, should_be_hashed, expected_output in test_cases:
# Simulate the logic from block_key() function
if input_key.startswith("sk-"):
hashed_token = hash_token(token=input_key)
else:
hashed_token = input_key
assert hashed_token == expected_output, f"Failed for input: {input_key}"
# Additional verification: if it should be hashed, verify it's actually a hash
if should_be_hashed:
# SHA-256 hashes are 64 characters long and contain only hex digits
assert (
len(hashed_token) == 64
), f"Hash length should be 64, got {len(hashed_token)} for {input_key}"
assert all(
c in "0123456789abcdef" for c in hashed_token
), f"Hash should contain only hex digits for {input_key}"
else:
# If not hashed, it should be the original string
assert (
hashed_token == input_key
), f"Non-hashed key should remain unchanged: {input_key}"
print("✅ All block_key hashing logic tests passed!")
def test_generate_gcp_iam_access_token():
"""
Test the _generate_gcp_iam_access_token function with mocked GCP IAM client.
"""
from unittest.mock import Mock, patch
service_account = "projects/-/serviceAccounts/test@project.iam.gserviceaccount.com"
expected_token = "test-access-token-12345"
# Mock the GCP IAM client and its response
mock_response = Mock()
mock_response.access_token = expected_token
mock_client = Mock()
mock_client.generate_access_token.return_value = mock_response
# Mock the iam_credentials_v1 module
mock_iam_credentials_v1 = Mock()
mock_iam_credentials_v1.IAMCredentialsClient = Mock(return_value=mock_client)
mock_iam_credentials_v1.GenerateAccessTokenRequest = Mock()
# Test successful token generation by mocking sys.modules
with patch.dict(
"sys.modules", {"google.cloud.iam_credentials_v1": mock_iam_credentials_v1}
):
from litellm._redis import _generate_gcp_iam_access_token
result = _generate_gcp_iam_access_token(service_account)
assert result == expected_token
mock_iam_credentials_v1.IAMCredentialsClient.assert_called_once()
mock_client.generate_access_token.assert_called_once()
# Verify the request was created with correct parameters
mock_iam_credentials_v1.GenerateAccessTokenRequest.assert_called_once_with(
name=service_account,
scope=["https://www.googleapis.com/auth/cloud-platform"],
)
def test_generate_gcp_iam_access_token_import_error():
"""
Test that _generate_gcp_iam_access_token raises ImportError when google-cloud-iam is not available.
"""
# Import the function first, before mocking
from litellm._redis import _generate_gcp_iam_access_token
# Mock the import to fail when the function tries to import google.cloud.iam_credentials_v1
original_import = __builtins__["__import__"]
def mock_import(name, *args, **kwargs):
if name == "google.cloud.iam_credentials_v1":
raise ImportError("No module named 'google.cloud.iam_credentials_v1'")
return original_import(name, *args, **kwargs)
with patch("builtins.__import__", side_effect=mock_import):
with pytest.raises(ImportError) as exc_info:
_generate_gcp_iam_access_token("test-service-account")
assert "google-cloud-iam is required" in str(exc_info.value)
assert "pip install google-cloud-iam" in str(exc_info.value)
def test_generate_azure_ad_redis_token():
"""Test _generate_azure_ad_redis_token with mocked Azure credential."""
from unittest.mock import Mock, patch
expected_token = "azure-access-token-12345"
mock_token = Mock()
mock_token.token = expected_token
mock_credential = Mock()
mock_credential.get_token.return_value = mock_token
mock_azure_identity = Mock()
mock_azure_identity.DefaultAzureCredential = Mock(return_value=mock_credential)
mock_azure_identity.ClientSecretCredential = Mock()
mock_azure_identity.ManagedIdentityCredential = Mock()
with patch.dict(
"sys.modules", {"azure.identity": mock_azure_identity, "azure": Mock()}
):
from litellm._redis import _generate_azure_ad_redis_token
result = _generate_azure_ad_redis_token()
assert result == expected_token
mock_credential.get_token.assert_called_once_with(
"https://redis.azure.com/.default"
)
def test_generate_azure_ad_redis_token_service_principal():
"""Test _generate_azure_ad_redis_token with service principal credentials."""
from unittest.mock import Mock, patch
expected_token = "sp-access-token-67890"
mock_token = Mock()
mock_token.token = expected_token
mock_credential = Mock()
mock_credential.get_token.return_value = mock_token
mock_client_secret_credential = Mock(return_value=mock_credential)
mock_azure_identity = Mock()
mock_azure_identity.DefaultAzureCredential = Mock()
mock_azure_identity.ClientSecretCredential = mock_client_secret_credential
mock_azure_identity.ManagedIdentityCredential = Mock()
with patch.dict(
"sys.modules", {"azure.identity": mock_azure_identity, "azure": Mock()}
):
from litellm._redis import _generate_azure_ad_redis_token
result = _generate_azure_ad_redis_token(
azure_client_id="test-client-id",
azure_tenant_id="test-tenant-id",
azure_client_secret="test-secret",
)
assert result == expected_token
mock_client_secret_credential.assert_called_once_with(
client_id="test-client-id",
tenant_id="test-tenant-id",
client_secret="test-secret",
)
def test_generate_azure_ad_redis_token_import_error():
"""Test that _generate_azure_ad_redis_token raises ImportError when azure-identity is missing."""
from unittest.mock import patch
from litellm._redis import _generate_azure_ad_redis_token
with patch.dict("sys.modules", {"azure.identity": None}):
with pytest.raises(ImportError) as exc_info:
_generate_azure_ad_redis_token()
assert "azure-identity is required" in str(exc_info.value)
def test_redis_client_logic_azure_ad_auth():
"""Test that _get_redis_client_logic sets up Azure AD auth when REDIS_AZURE_AD_TOKEN=true.
Mocks ``azure.identity`` via ``sys.modules`` so the test does not require
the real ``azure-identity`` package to be installed in the CI environment.
"""
from unittest.mock import Mock, patch
mock_credential = Mock()
mock_azure_identity = Mock()
mock_azure_identity.DefaultAzureCredential = Mock(return_value=mock_credential)
mock_azure_identity.ClientSecretCredential = Mock(return_value=mock_credential)
mock_azure_identity.ManagedIdentityCredential = Mock(return_value=mock_credential)
with patch.dict(
"sys.modules", {"azure.identity": mock_azure_identity, "azure": Mock()}
):
from litellm._redis import _get_redis_client_logic
redis_kwargs = _get_redis_client_logic(
host="myredis.redis.cache.windows.net",
port="6380",
azure_redis_ad_token="true",
ssl=True,
)
assert "redis_connect_func" in redis_kwargs
# Marker for async paths to detect Azure AD auth
assert hasattr(redis_kwargs["redis_connect_func"], "_azure_redis_ad_token")
assert redis_kwargs["redis_connect_func"]._azure_redis_ad_token is True
# Live credential object (not raw secret) is exposed for async paths
assert hasattr(redis_kwargs["redis_connect_func"], "_azure_credential")
# Raw credentials must NOT be exposed on the function
assert not hasattr(redis_kwargs["redis_connect_func"], "_azure_client_secret")
assert not hasattr(redis_kwargs["redis_connect_func"], "_azure_client_id")
assert not hasattr(redis_kwargs["redis_connect_func"], "_azure_tenant_id")
# Azure-specific kwargs should be removed from the dict passed to Redis
assert "azure_redis_ad_token" not in redis_kwargs
assert "azure_client_id" not in redis_kwargs
if __name__ == "__main__":
# Allow running this test file directly for debugging
pytest.main([__file__, "-v"])
def test_model_info_for_vertex_ai_deepseek_model():
model_info = litellm.get_model_info(
model="vertex_ai/deepseek-ai/deepseek-r1-0528-maas"
)
assert model_info is not None
assert model_info["litellm_provider"] == "vertex_ai-deepseek_models"
assert model_info["mode"] == "chat"
assert model_info["input_cost_per_token"] is not None
assert model_info["output_cost_per_token"] is not None
print("vertex deepseek model info", model_info)
def test_model_info_for_openrouter_kimi_k2_5():
"""
Test that openrouter/moonshotai/kimi-k2.5 model info is correctly configured
in model_prices_and_context_window.json.
Model properties from OpenRouter API:
- context_length: 262144
- pricing: prompt=$0.0000006, completion=$0.000003, input_cache_read=$0.0000001
- modality: text+image->text (supports vision)
- supports: tool_choice, tools (function calling)
"""
import json
from pathlib import Path
# Load directly from the local JSON file
json_path = Path(__file__).parents[2] / "model_prices_and_context_window.json"
with open(json_path) as f:
model_cost = json.load(f)
model_info = model_cost.get("openrouter/moonshotai/kimi-k2.5")
assert (
model_info is not None
), "Model not found in model_prices_and_context_window.json"
assert model_info["litellm_provider"] == "openrouter"
assert model_info["mode"] == "chat"
# Verify context window
assert model_info["max_input_tokens"] == 262144
assert model_info["max_output_tokens"] == 262144
assert model_info["max_tokens"] == 262144
# Verify pricing
assert model_info["input_cost_per_token"] == 6e-07
assert model_info["output_cost_per_token"] == 3e-06
assert model_info["cache_read_input_token_cost"] == 1e-07
# Verify capabilities
assert model_info["supports_vision"] is True
assert model_info["supports_function_calling"] is True
assert model_info["supports_tool_choice"] is True
print("openrouter kimi-k2.5 model info", model_info)
def test_gemini_embedding_2_ga_in_cost_map():
"""GA and Vertex preview gemini-embedding-2 entries align with multimodal unit pricing."""
import json
from pathlib import Path
json_path = Path(__file__).parents[2] / "model_prices_and_context_window.json"
with open(json_path) as f:
model_cost = json.load(f)
for key, provider in (
("gemini/gemini-embedding-2", "gemini"),
("vertex_ai/gemini-embedding-2", "vertex_ai"),
("vertex_ai/gemini-embedding-2-preview", "vertex_ai"),
("gemini-embedding-2", "vertex_ai-embedding-models"),
):
info = model_cost.get(key)
assert (
info is not None
), f"{key} missing from model_prices_and_context_window.json"
assert info["litellm_provider"] == provider
assert info.get("mode") == "embedding"
assert info.get("supports_multimodal") is True
assert info.get("input_cost_per_token") == 2e-07
assert info.get("input_cost_per_image") == 0.00012
assert info.get("input_cost_per_audio_per_second") == 0.00016
assert info.get("input_cost_per_video_per_second") == 0.00079
if provider in ("vertex_ai-embedding-models", "vertex_ai"):
assert (
info.get("uses_embed_content") is True
), f"{key} must have uses_embed_content=true for correct Vertex AI routing"
def test_gemini_lyria_3_preview_models_in_cost_map():
import json
from pathlib import Path
json_path = Path(__file__).parents[2] / "model_prices_and_context_window.json"
with open(json_path) as f:
model_cost = json.load(f)
clip = model_cost.get("gemini/lyria-3-clip-preview")
pro = model_cost.get("gemini/lyria-3-pro-preview")
assert clip is not None and pro is not None
assert clip["litellm_provider"] == "gemini" and pro["litellm_provider"] == "gemini"
assert clip["max_input_tokens"] == 131072 == pro["max_input_tokens"]
assert clip["output_cost_per_image"] == 0.04
def test_model_info_for_fireworks_short_form_models():
"""
Test that fireworks_ai short-form model entries (fireworks_ai/<model>)
are correctly configured in model_prices_and_context_window.json.
These entries enable cost attribution for models called via short-form
names (e.g., fireworks_ai/glm-4p7 instead of
fireworks_ai/accounts/fireworks/models/glm-4p7).
"""
import json
from pathlib import Path
json_path = Path(__file__).parents[2] / "model_prices_and_context_window.json"
with open(json_path) as f:
model_cost = json.load(f)
# glm-4p7: short-form and long-form
for key in [
"fireworks_ai/glm-4p7",
"fireworks_ai/accounts/fireworks/models/glm-4p7",
]:
info = model_cost.get(key)
assert (
info is not None
), f"{key} not found in model_prices_and_context_window.json"
assert info["litellm_provider"] == "fireworks_ai"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == 6e-07
assert info["output_cost_per_token"] == 2.2e-06
assert info["max_input_tokens"] == 202800
assert info["supports_reasoning"] is True
# minimax-m2p1: short-form and long-form
for key in [
"fireworks_ai/minimax-m2p1",
"fireworks_ai/accounts/fireworks/models/minimax-m2p1",
]:
info = model_cost.get(key)
assert (
info is not None
), f"{key} not found in model_prices_and_context_window.json"
assert info["litellm_provider"] == "fireworks_ai"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == 3e-07
assert info["output_cost_per_token"] == 1.2e-06
assert info["max_input_tokens"] == 204800
# kimi-k2p5: short-form only (long-form already existed)
info = model_cost.get("fireworks_ai/kimi-k2p5")
assert (
info is not None
), "fireworks_ai/kimi-k2p5 not found in model_prices_and_context_window.json"
assert info["litellm_provider"] == "fireworks_ai"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == 6e-07
assert info["output_cost_per_token"] == 3e-06
assert info["max_input_tokens"] == 262144
class TestGetValidModelsWithCLI:
"""Test get_valid_models function as used in CLI token usage"""
def test_get_valid_models_with_cli_pattern(self):
"""Test get_valid_models with litellm_proxy provider and CLI token pattern"""
# Mock the HTTP request that get_valid_models makes to the proxy
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = {
"data": [
{"id": "gpt-3.5-turbo", "object": "model"},
{"id": "gpt-4", "object": "model"},
{"id": "litellm_proxy/gemini/gemini-2.5-flash", "object": "model"},
{"id": "claude-3-sonnet", "object": "model"},
]
}
with patch.object(
litellm.module_level_client, "get", return_value=mock_response
) as mock_get:
# Test the exact pattern used in cli_token_usage.py
result = litellm.get_valid_models(
check_provider_endpoint=True,
custom_llm_provider="litellm_proxy",
api_key="sk-test-cli-key-123",
api_base="http://localhost:4000/",
)
# Verify the function returns a list of model names
assert isinstance(result, list)
assert len(result) == 4
# All models get prefixed with "litellm_proxy/" by the get_models method
assert "litellm_proxy/gpt-3.5-turbo" in result
assert "litellm_proxy/gpt-4" in result
# Note: This model already had the prefix, so it gets double-prefixed
assert "litellm_proxy/litellm_proxy/gemini/gemini-2.5-flash" in result
assert "litellm_proxy/claude-3-sonnet" in result
# Verify the HTTP request was made with correct parameters
mock_get.assert_called_once()
_, call_kwargs = mock_get.call_args
# Check that the request was made to the correct endpoint
assert call_kwargs["url"].startswith("http://localhost:4000/")
assert call_kwargs["url"].endswith("/v1/models")
# Check that the API key was included in headers
assert "headers" in call_kwargs
headers = call_kwargs["headers"]
assert headers.get("Authorization") == "Bearer sk-test-cli-key-123"
class TestIsCachedMessage:
"""Test is_cached_message function for context caching detection.
Fixes GitHub issue #17821 - TypeError when content is string instead of list.
"""
def test_string_content_returns_false(self):
"""String content should return False without crashing."""
message = {"role": "user", "content": "Hello world"}
assert is_cached_message(message) is False
def test_none_content_returns_false(self):
"""None content should return False."""
message = {"role": "user", "content": None}
assert is_cached_message(message) is False
def test_missing_content_returns_false(self):
"""Message without content key should return False."""
message = {"role": "user"}
assert is_cached_message(message) is False
def test_list_content_without_cache_control_returns_false(self):
"""List content without cache_control should return False."""
message = {"role": "user", "content": [{"type": "text", "text": "Hello"}]}
assert is_cached_message(message) is False
def test_list_content_with_cache_control_returns_true(self):
"""List content with cache_control ephemeral should return True."""
message = {
"role": "user",
"content": [
{
"type": "text",
"text": "Hello",
"cache_control": {"type": "ephemeral"},
}
],
}
assert is_cached_message(message) is True
def test_list_with_non_dict_items_skips_them(self):
"""List content with non-dict items should skip them gracefully."""
message = {
"role": "user",
"content": ["string_item", 123, {"type": "text", "text": "Hello"}],
}
assert is_cached_message(message) is False
def test_list_with_mixed_items_finds_cached(self):
"""Mixed content list should find cached item."""
message = {
"role": "user",
"content": [
"string_item",
{"type": "image", "url": "..."},
{
"type": "text",
"text": "cached",
"cache_control": {"type": "ephemeral"},
},
],
}
assert is_cached_message(message) is True
def test_wrong_cache_control_type_returns_false(self):
"""Non-ephemeral cache_control type should return False."""
message = {
"role": "user",
"content": [
{
"type": "text",
"text": "Hello",
"cache_control": {"type": "permanent"},
}
],
}
assert is_cached_message(message) is False
def test_empty_list_content_returns_false(self):
"""Empty list content should return False."""
message = {"role": "user", "content": []}
assert is_cached_message(message) is False
def test_message_level_cache_control_returns_true(self):
"""Message with string content and message-level cache_control should return True.
This is the format injected by the cache_control_injection_points hook
when the message content is a string (common for system messages).
Fixes GitHub issue #18519 - Gemini models ignoring cache_control_injection_points.
"""
message = {
"role": "system",
"content": "You are a helpful assistant.",
"cache_control": {"type": "ephemeral"},
}
assert is_cached_message(message) is True
def test_message_level_cache_control_wrong_type_returns_false(self):
"""Message-level cache_control with non-ephemeral type should return False."""
message = {
"role": "system",
"content": "You are a helpful assistant.",
"cache_control": {"type": "permanent"},
}
assert is_cached_message(message) is False
def test_message_level_cache_control_non_dict_returns_false(self):
"""Message-level cache_control that's not a dict should return False."""
message = {
"role": "system",
"content": "You are a helpful assistant.",
"cache_control": "ephemeral",
}
assert is_cached_message(message) is False
@pytest.mark.asyncio
class TestProxyLoggingBudgetAlerts:
"""Test budget_alerts method in ProxyLogging class."""
async def test_budget_alerts_when_alerting_is_none(self):
"""Test that budget_alerts returns early when alerting is None."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging.alerting = None
proxy_logging.slack_alerting_instance = AsyncMock()
proxy_logging.email_logging_instance = AsyncMock()
user_info = MagicMock()
# Should return without calling any alerting instances
await proxy_logging.budget_alerts(type="user_budget", user_info=user_info)
# Verify no calls were made
proxy_logging.slack_alerting_instance.budget_alerts.assert_not_called()
proxy_logging.email_logging_instance.budget_alerts.assert_not_called()
async def test_budget_alerts_with_slack_only(self):
"""Test that budget_alerts calls slack_alerting_instance when slack is in alerting."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging.alerting = ["slack"]
proxy_logging.slack_alerting_instance = AsyncMock()
user_info = MagicMock()
await proxy_logging.budget_alerts(type="token_budget", user_info=user_info)
proxy_logging.slack_alerting_instance.budget_alerts.assert_called_once_with(
type="token_budget", user_info=user_info
)
async def test_budget_alerts_with_email_only(self):
"""Test that budget_alerts calls email_logging_instance when email is in alerting."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging.alerting = ["email"]
proxy_logging.email_logging_instance = AsyncMock()
user_info = MagicMock()
await proxy_logging.budget_alerts(type="team_budget", user_info=user_info)
proxy_logging.email_logging_instance.budget_alerts.assert_called_once_with(
type="team_budget", user_info=user_info
)
async def test_budget_alerts_with_email_when_instance_is_none(self):
"""Test that budget_alerts does not call email_logging_instance when it is None."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging.alerting = ["email"]
proxy_logging.email_logging_instance = None
user_info = MagicMock()
# Should not raise an error
await proxy_logging.budget_alerts(
type="organization_budget", user_info=user_info
)
async def test_budget_alerts_with_both_slack_and_email(self):
"""Test that budget_alerts calls both slack and email instances when both are in alerting."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging.alerting = ["slack", "email"]
proxy_logging.slack_alerting_instance = AsyncMock()
proxy_logging.email_logging_instance = AsyncMock()
user_info = MagicMock()
await proxy_logging.budget_alerts(type="proxy_budget", user_info=user_info)
proxy_logging.slack_alerting_instance.budget_alerts.assert_called_once_with(
type="proxy_budget", user_info=user_info
)
proxy_logging.email_logging_instance.budget_alerts.assert_called_once_with(
type="proxy_budget", user_info=user_info
)
@pytest.mark.parametrize(
"alert_type",
[
"token_budget",
"user_budget",
"soft_budget",
"team_budget",
"organization_budget",
"proxy_budget",
"projected_limit_exceeded",
],
)
async def test_budget_alerts_with_all_alert_types(self, alert_type):
"""Test that budget_alerts works with all supported alert types."""
from litellm.caching.caching import DualCache
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging.alerting = ["slack", "email"]
proxy_logging.slack_alerting_instance = AsyncMock()
proxy_logging.email_logging_instance = AsyncMock()
user_info = MagicMock()
await proxy_logging.budget_alerts(type=alert_type, user_info=user_info)
proxy_logging.slack_alerting_instance.budget_alerts.assert_called_once_with(
type=alert_type, user_info=user_info
)
proxy_logging.email_logging_instance.budget_alerts.assert_called_once_with(
type=alert_type, user_info=user_info
)
async def test_budget_alerts_soft_budget_with_alert_emails_bypasses_alerting_none(
self,
):
"""
Test that soft_budget alerts with alert_emails bypass the alerting=None check
and send emails even when alerting is None.
This tests the new logic that allows team-specific soft budget email alerts
via metadata.soft_budget_alerting_emails to work even when global alerting is disabled.
"""
from litellm.caching.caching import DualCache
from litellm.proxy._types import CallInfo, Litellm_EntityType
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging.alerting = None # Global alerting is disabled
proxy_logging.slack_alerting_instance = AsyncMock()
proxy_logging.email_logging_instance = AsyncMock()
# Create CallInfo with alert_emails set (simulating team metadata extraction)
user_info = CallInfo(
token="test-token",
spend=100.0,
soft_budget=50.0,
user_id="test-user",
team_id="test-team",
team_alias="test-team-alias",
event_group=Litellm_EntityType.TEAM,
alert_emails=["team1@example.com", "team2@example.com"],
)
# Should send email even though alerting is None (because of alert_emails)
await proxy_logging.budget_alerts(type="soft_budget", user_info=user_info)
# Verify slack was NOT called (alerting is None)
proxy_logging.slack_alerting_instance.budget_alerts.assert_not_called()
# Verify email WAS called (bypasses alerting=None check)
proxy_logging.email_logging_instance.budget_alerts.assert_called_once_with(
type="soft_budget", user_info=user_info
)
async def test_budget_alerts_soft_budget_without_alert_emails_respects_alerting_none(
self,
):
"""
Test that soft_budget alerts WITHOUT alert_emails still respect alerting=None
and do not send emails when alerting is None.
"""
from litellm.caching.caching import DualCache
from litellm.proxy._types import CallInfo, Litellm_EntityType
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging.alerting = None
proxy_logging.slack_alerting_instance = AsyncMock()
proxy_logging.email_logging_instance = AsyncMock()
# Create CallInfo WITHOUT alert_emails
user_info = CallInfo(
token="test-token",
spend=100.0,
soft_budget=50.0,
user_id="test-user",
team_id="test-team",
team_alias="test-team-alias",
event_group=Litellm_EntityType.TEAM,
alert_emails=None, # No alert emails
)
# Should NOT send email (alerting is None and no alert_emails)
await proxy_logging.budget_alerts(type="soft_budget", user_info=user_info)
# Verify no calls were made
proxy_logging.slack_alerting_instance.budget_alerts.assert_not_called()
proxy_logging.email_logging_instance.budget_alerts.assert_not_called()
async def test_budget_alerts_soft_budget_with_empty_alert_emails_respects_alerting_none(
self,
):
"""
Test that soft_budget alerts with empty alert_emails list still respect alerting=None.
"""
from litellm.caching.caching import DualCache
from litellm.proxy._types import CallInfo, Litellm_EntityType
from litellm.proxy.utils import ProxyLogging
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
proxy_logging.alerting = None
proxy_logging.slack_alerting_instance = AsyncMock()
proxy_logging.email_logging_instance = AsyncMock()
# Create CallInfo with empty alert_emails list
user_info = CallInfo(
token="test-token",
spend=100.0,
soft_budget=50.0,
user_id="test-user",
team_id="test-team",
team_alias="test-team-alias",
event_group=Litellm_EntityType.TEAM,
alert_emails=[], # Empty list
)
# Should NOT send email (alert_emails is empty)
await proxy_logging.budget_alerts(type="soft_budget", user_info=user_info)
# Verify no calls were made
proxy_logging.slack_alerting_instance.budget_alerts.assert_not_called()
proxy_logging.email_logging_instance.budget_alerts.assert_not_called()
def test_azure_ai_claude_provider_config():
"""Test that Azure AI Claude models return AzureAnthropicConfig for proper tool transformation."""
from litellm import AzureAIStudioConfig, AzureAnthropicConfig
from litellm.utils import ProviderConfigManager
# Claude models should return AzureAnthropicConfig
config = ProviderConfigManager.get_provider_chat_config(
model="claude-sonnet-4-5",
provider=LlmProviders.AZURE_AI,
)
assert isinstance(config, AzureAnthropicConfig)
# Test case-insensitive matching
config = ProviderConfigManager.get_provider_chat_config(
model="Claude-Opus-4",
provider=LlmProviders.AZURE_AI,
)
assert isinstance(config, AzureAnthropicConfig)
# Non-Claude models should return AzureAIStudioConfig
config = ProviderConfigManager.get_provider_chat_config(
model="mistral-large",
provider=LlmProviders.AZURE_AI,
)
assert isinstance(config, AzureAIStudioConfig)
# Tests for thinking blocks helper functions
# Related to issue: https://github.com/BerriAI/litellm/issues/18926
def test_any_assistant_message_has_thinking_blocks_with_thinking():
"""Test that function returns True when any assistant message has thinking_blocks."""
from litellm.utils import any_assistant_message_has_thinking_blocks
messages = [
{"role": "user", "content": "Hello"},
{
"role": "assistant",
"thinking_blocks": [{"type": "thinking", "thinking": "Let me think..."}],
"tool_calls": [{"id": "123", "function": {"name": "test"}}],
},
{"role": "tool", "tool_call_id": "123", "content": "result"},
{
"role": "assistant",
"tool_calls": [{"id": "456", "function": {"name": "test2"}}],
# No thinking_blocks here - Claude sometimes doesn't include them
},
]
assert any_assistant_message_has_thinking_blocks(messages) is True
def test_any_assistant_message_has_thinking_blocks_without_thinking():
"""Test that function returns False when no assistant message has thinking_blocks."""
from litellm.utils import any_assistant_message_has_thinking_blocks
messages = [
{"role": "user", "content": "Hello"},
{
"role": "assistant",
"tool_calls": [{"id": "123", "function": {"name": "test"}}],
},
{"role": "tool", "tool_call_id": "123", "content": "result"},
]
assert any_assistant_message_has_thinking_blocks(messages) is False
def test_any_assistant_message_has_thinking_blocks_empty_list():
"""Test that function returns False when thinking_blocks is an empty list."""
from litellm.utils import any_assistant_message_has_thinking_blocks
messages = [
{"role": "user", "content": "Hello"},
{
"role": "assistant",
"thinking_blocks": [], # Empty list
"tool_calls": [{"id": "123", "function": {"name": "test"}}],
},
]
assert any_assistant_message_has_thinking_blocks(messages) is False
def test_last_assistant_with_tool_calls_has_no_thinking_blocks_issue_18926():
"""
Test the scenario from issue #18926 where:
- First assistant message HAS thinking_blocks
- Second assistant message has NO thinking_blocks
The old logic would drop thinking because the LAST tool_call message
has no thinking_blocks, but this breaks because the first message
still has thinking blocks in the conversation.
"""
from litellm.utils import (
any_assistant_message_has_thinking_blocks,
last_assistant_with_tool_calls_has_no_thinking_blocks,
)
messages = [
{"role": "user", "content": "Build a feature"},
{
"role": "assistant",
"thinking_blocks": [
{"type": "thinking", "thinking": "Let me analyze the requirements..."}
],
"tool_calls": [
{
"id": "toolu_1",
"function": {"name": "file_editor", "arguments": "{}"},
}
],
},
{
"role": "tool",
"tool_call_id": "toolu_1",
"content": "File contents here...",
},
{
"role": "assistant",
# NO thinking_blocks - Claude sometimes doesn't include them
"content": [{"type": "text", "text": "Let me explore more..."}],
"tool_calls": [
{
"id": "toolu_2",
"function": {"name": "file_editor", "arguments": "{}"},
}
],
},
]
# Last assistant with tool_calls has no thinking_blocks
assert last_assistant_with_tool_calls_has_no_thinking_blocks(messages) is True
# But ANY assistant message has thinking_blocks
assert any_assistant_message_has_thinking_blocks(messages) is True
# So we should NOT drop thinking - the combination tells us thinking is in use
# The fix uses both checks: only drop if last has none AND no message has any
should_drop_thinking = last_assistant_with_tool_calls_has_no_thinking_blocks(
messages
) and not any_assistant_message_has_thinking_blocks(messages)
assert should_drop_thinking is False
class TestAdditionalDropParamsForNonOpenAIProviders:
"""
Test additional_drop_params functionality for non-OpenAI providers.
Fixes https://github.com/BerriAI/litellm/issues/19225
The bug was that additional_drop_params only filtered params for OpenAI/Azure
providers, but not for other providers like Bedrock. This caused OpenAI-specific
params like prompt_cache_key to be passed to Bedrock, resulting in errors.
"""
def test_additional_drop_params_filters_for_bedrock(self):
"""
Test that additional_drop_params correctly filters params for Bedrock provider.
Before the fix, prompt_cache_key would be passed through to Bedrock even when
specified in additional_drop_params, causing:
'BedrockException - {"message":"The model returned the following errors:
prompt_cache_key: Extra inputs are not permitted"}'
"""
from litellm.utils import add_provider_specific_params_to_optional_params
optional_params = {}
passed_params = {
"prompt_cache_key": "test_key_123",
"temperature": 0.7,
"model": "bedrock/anthropic.claude-v2",
}
openai_params = ["temperature", "max_tokens", "top_p", "model"]
result = add_provider_specific_params_to_optional_params(
optional_params=optional_params,
passed_params=passed_params,
custom_llm_provider="bedrock",
openai_params=openai_params,
additional_drop_params=["prompt_cache_key"],
)
# prompt_cache_key should be filtered out
assert "prompt_cache_key" not in result
# temperature should still be there (it's in openai_params, not filtered)
# Note: temperature is in openai_params so it won't be added by this function
# The function only adds params NOT in openai_params
def test_additional_drop_params_filters_multiple_params_for_non_openai(self):
"""Test filtering multiple params for non-OpenAI providers."""
from litellm.utils import add_provider_specific_params_to_optional_params
optional_params = {}
passed_params = {
"prompt_cache_key": "test_key",
"some_openai_only_param": "value1",
"another_openai_param": "value2",
"keep_this_param": "keep_me",
}
openai_params = ["temperature", "max_tokens"]
result = add_provider_specific_params_to_optional_params(
optional_params=optional_params,
passed_params=passed_params,
custom_llm_provider="anthropic",
openai_params=openai_params,
additional_drop_params=["prompt_cache_key", "some_openai_only_param"],
)
# Filtered params should not be present
assert "prompt_cache_key" not in result
assert "some_openai_only_param" not in result
# Non-filtered params should be present
assert result.get("another_openai_param") == "value2"
assert result.get("keep_this_param") == "keep_me"
def test_additional_drop_params_none_keeps_all_params(self):
"""Test that when additional_drop_params is None, all params are kept."""
from litellm.utils import add_provider_specific_params_to_optional_params
optional_params = {}
passed_params = {
"prompt_cache_key": "test_key",
"custom_param": "value",
}
openai_params = ["temperature"]
result = add_provider_specific_params_to_optional_params(
optional_params=optional_params,
passed_params=passed_params,
custom_llm_provider="bedrock",
openai_params=openai_params,
additional_drop_params=None,
)
# All params should be present when additional_drop_params is None
assert result.get("prompt_cache_key") == "test_key"
assert result.get("custom_param") == "value"
def test_additional_drop_params_empty_list_keeps_all_params(self):
"""Test that when additional_drop_params is empty list, all params are kept."""
from litellm.utils import add_provider_specific_params_to_optional_params
optional_params = {}
passed_params = {
"prompt_cache_key": "test_key",
"custom_param": "value",
}
openai_params = ["temperature"]
result = add_provider_specific_params_to_optional_params(
optional_params=optional_params,
passed_params=passed_params,
custom_llm_provider="bedrock",
openai_params=openai_params,
additional_drop_params=[],
)
# All params should be present when additional_drop_params is empty
assert result.get("prompt_cache_key") == "test_key"
assert result.get("custom_param") == "value"
class TestDropParamsWithPromptCacheKey:
"""
Test that drop_params: true correctly drops prompt_cache_key for non-OpenAI providers.
Fixes https://github.com/BerriAI/litellm/issues/19225
prompt_cache_key is an OpenAI-specific parameter that should be automatically
dropped when using providers like Bedrock that don't support it.
"""
def test_prompt_cache_key_in_default_params(self):
"""Verify prompt_cache_key is now in DEFAULT_CHAT_COMPLETION_PARAM_VALUES."""
from litellm.constants import DEFAULT_CHAT_COMPLETION_PARAM_VALUES
assert "prompt_cache_key" in DEFAULT_CHAT_COMPLETION_PARAM_VALUES
assert "prompt_cache_retention" in DEFAULT_CHAT_COMPLETION_PARAM_VALUES
def test_drop_params_removes_prompt_cache_key_for_bedrock(self):
"""
Test that get_optional_params with drop_params=True removes prompt_cache_key
for Bedrock provider since it's not in Bedrock's supported params.
"""
from litellm.utils import get_optional_params
# Call get_optional_params for Bedrock with prompt_cache_key
# drop_params=True should remove it since Bedrock doesn't support it
result = get_optional_params(
model="anthropic.claude-3-sonnet-20240229-v1:0",
custom_llm_provider="bedrock",
prompt_cache_key="test_cache_key",
temperature=0.7,
drop_params=True,
)
# prompt_cache_key should be dropped for Bedrock
assert "prompt_cache_key" not in result
# temperature should remain (it's supported by Bedrock)
assert result.get("temperature") == 0.7
class TestGetOptionalParamsDeepSeek:
"""Tests that deepseek provider uses DeepSeekChatConfig for parameter mapping."""
def test_deepseek_supports_thinking_param(self):
"""
Verify that get_optional_params for deepseek accepts the 'thinking' param,
which is only supported by DeepSeekChatConfig, not OpenAIConfig.
"""
from litellm.utils import get_optional_params
result = get_optional_params(
model="deepseek-reasoner",
custom_llm_provider="deepseek",
thinking={"type": "enabled"},
)
assert result.get("thinking") == {"type": "enabled"}
def test_deepseek_supports_reasoning_effort_param(self):
"""
Verify that get_optional_params for deepseek accepts 'reasoning_effort',
which is only supported by DeepSeekChatConfig, not OpenAIConfig.
"""
from litellm.utils import get_optional_params
result = get_optional_params(
model="deepseek-reasoner",
custom_llm_provider="deepseek",
reasoning_effort="high",
)
assert result.get("thinking") == {"type": "enabled"}
def test_deepseek_thinking_strips_budget_tokens(self):
"""
DeepSeekChatConfig strips budget_tokens from thinking param.
This would not happen with OpenAIConfig.
"""
from litellm.utils import get_optional_params
result = get_optional_params(
model="deepseek-reasoner",
custom_llm_provider="deepseek",
thinking={"type": "enabled", "budget_tokens": 5000},
)
assert "budget_tokens" not in result.get("thinking", {})
assert result.get("thinking") == {"type": "enabled"}
class TestIsStreamingRequest:
def test_stream_true_in_kwargs(self):
assert (
_is_streaming_request(kwargs={"stream": True}, call_type="acompletion")
is True
)
def test_stream_false_in_kwargs(self):
assert (
_is_streaming_request(kwargs={"stream": False}, call_type="acompletion")
is False
)
def test_no_stream_in_kwargs(self):
assert _is_streaming_request(kwargs={}, call_type="acompletion") is False
def test_generate_content_stream_string(self):
assert (
_is_streaming_request(
kwargs={}, call_type=CallTypes.generate_content_stream.value
)
is True
)
def test_agenerate_content_stream_string(self):
assert (
_is_streaming_request(
kwargs={}, call_type=CallTypes.agenerate_content_stream.value
)
is True
)
def test_generate_content_stream_enum(self):
assert (
_is_streaming_request(
kwargs={}, call_type=CallTypes.generate_content_stream
)
is True
)
def test_agenerate_content_stream_enum(self):
assert (
_is_streaming_request(
kwargs={}, call_type=CallTypes.agenerate_content_stream
)
is True
)
def test_non_streaming_call_type_string(self):
assert _is_streaming_request(kwargs={}, call_type="acompletion") is False
def test_non_streaming_call_type_enum(self):
assert (
_is_streaming_request(kwargs={}, call_type=CallTypes.acompletion) is False
)
def test_stream_true_overrides_non_streaming_call_type(self):
assert (
_is_streaming_request(
kwargs={"stream": True}, call_type=CallTypes.acompletion
)
is True
)
class TestCallbackAsyncSyncSeparation:
"""Test that LoggingCallbackManager auto-routes async callbacks to async lists."""
def setup_method(self):
"""Reset callback lists before each test."""
litellm.input_callback = []
litellm.success_callback = []
litellm.failure_callback = []
litellm._async_input_callback = []
litellm._async_success_callback = []
litellm._async_failure_callback = []
def test_async_success_callback_routed_to_async_list(self):
async def my_async_cb(*args, **kwargs):
pass
litellm.logging_callback_manager.add_litellm_success_callback(my_async_cb)
assert my_async_cb in litellm._async_success_callback
assert my_async_cb not in litellm.success_callback
def test_sync_success_callback_stays_in_sync_list(self):
def my_sync_cb(*args, **kwargs):
pass
litellm.logging_callback_manager.add_litellm_success_callback(my_sync_cb)
assert my_sync_cb in litellm.success_callback
assert my_sync_cb not in litellm._async_success_callback
def test_string_callback_stays_in_sync_list(self):
litellm.logging_callback_manager.add_litellm_success_callback("langfuse")
assert "langfuse" in litellm.success_callback
assert "langfuse" not in litellm._async_success_callback
def test_async_failure_callback_routed_to_async_list(self):
async def my_async_cb(*args, **kwargs):
pass
litellm.logging_callback_manager.add_litellm_failure_callback(my_async_cb)
assert my_async_cb in litellm._async_failure_callback
assert my_async_cb not in litellm.failure_callback
def test_sync_failure_callback_stays_in_sync_list(self):
def my_sync_cb(*args, **kwargs):
pass
litellm.logging_callback_manager.add_litellm_failure_callback(my_sync_cb)
assert my_sync_cb in litellm.failure_callback
assert my_sync_cb not in litellm._async_failure_callback
def test_dynamodb_routed_to_async_success(self):
litellm.logging_callback_manager.add_litellm_success_callback("dynamodb")
assert "dynamodb" in litellm._async_success_callback
assert "dynamodb" not in litellm.success_callback
def test_openmeter_routed_to_async_success(self):
litellm.logging_callback_manager.add_litellm_success_callback("openmeter")
assert "openmeter" in litellm._async_success_callback
assert "openmeter" not in litellm.success_callback
def test_async_input_callback_routed_to_async_list(self):
async def my_async_cb(*args, **kwargs):
pass
litellm.logging_callback_manager.add_litellm_input_callback(my_async_cb)
assert my_async_cb in litellm._async_input_callback
assert my_async_cb not in litellm.input_callback
def test_sync_input_callback_stays_in_sync_list(self):
def my_sync_cb(*args, **kwargs):
pass
litellm.logging_callback_manager.add_litellm_input_callback(my_sync_cb)
assert my_sync_cb in litellm.input_callback
assert my_sync_cb not in litellm._async_input_callback
class TestMetadataNoneHandling:
"""
Test that metadata=None in kwargs doesn't cause TypeError.
When metadata key exists with value None (e.g., from Azure OpenAI streaming),
dict.get("metadata", {}) returns None (key exists, so default is ignored).
The fix uses (kwargs.get("metadata") or {}) which handles both missing key
and explicit None value.
Related: #20871
"""
def test_metadata_none_get_previous_models(self):
"""kwargs.get("metadata") or {} should return {} when metadata is None."""
kwargs = {"metadata": None}
previous_models = (kwargs.get("metadata") or {}).get("previous_models", None)
assert previous_models is None
def test_metadata_none_model_group_check(self):
"""'model_group' in (kwargs.get("metadata") or {}) should not raise TypeError."""
kwargs = {"metadata": None}
_is_litellm_router_call = "model_group" in (kwargs.get("metadata") or {})
assert _is_litellm_router_call is False
def test_metadata_missing_key(self):
"""Should work when metadata key is completely absent."""
kwargs = {}
previous_models = (kwargs.get("metadata") or {}).get("previous_models", None)
assert previous_models is None
def test_metadata_present_with_values(self):
"""Should work when metadata has actual values."""
kwargs = {"metadata": {"previous_models": ["model1"], "model_group": "test"}}
previous_models = (kwargs.get("metadata") or {}).get("previous_models", None)
assert previous_models == ["model1"]
_is_litellm_router_call = "model_group" in (kwargs.get("metadata") or {})
assert _is_litellm_router_call is True
def test_metadata_none_causes_error_with_old_pattern(self):
"""Demonstrate the bug: dict.get('metadata', {}) returns None when key exists with None value."""
kwargs = {"metadata": None}
# Old pattern: kwargs.get("metadata", {}) returns None because key exists
result = kwargs.get("metadata", {})
assert result is None # This is the root cause of the bug
# Attempting to use .get() on None raises AttributeError or TypeError
with pytest.raises((TypeError, AttributeError)):
kwargs.get("metadata", {}).get("previous_models", None)
# Attempting 'in' on None raises TypeError
with pytest.raises(TypeError):
"model_group" in kwargs.get("metadata", {})
def test_litellm_params_metadata_none(self):
"""litellm_params.get("metadata") or {} should handle None value."""
litellm_params = {"metadata": None}
metadata = litellm_params.get("metadata") or {}
assert metadata == {}
class TestValidateAndFixThinkingParam:
"""Tests for validate_and_fix_thinking_param."""
def test_none_returns_none(self):
from litellm.utils import validate_and_fix_thinking_param
assert validate_and_fix_thinking_param(thinking=None) is None
def test_already_snake_case(self):
from litellm.utils import validate_and_fix_thinking_param
thinking = {"type": "enabled", "budget_tokens": 32000}
result = validate_and_fix_thinking_param(thinking=thinking)
assert result == {"type": "enabled", "budget_tokens": 32000}
def test_camel_case_normalized(self):
from litellm.utils import validate_and_fix_thinking_param
thinking = {"type": "enabled", "budgetTokens": 32000}
result = validate_and_fix_thinking_param(thinking=thinking)
assert result == {"type": "enabled", "budget_tokens": 32000}
assert "budgetTokens" not in result
def test_both_keys_snake_case_wins(self):
from litellm.utils import validate_and_fix_thinking_param
thinking = {"type": "enabled", "budget_tokens": 10000, "budgetTokens": 50000}
result = validate_and_fix_thinking_param(thinking=thinking)
assert result == {"type": "enabled", "budget_tokens": 10000}
assert "budgetTokens" not in result
def test_original_dict_not_mutated(self):
from litellm.utils import validate_and_fix_thinking_param
thinking = {"type": "enabled", "budgetTokens": 32000}
validate_and_fix_thinking_param(thinking=thinking)
assert "budgetTokens" in thinking
assert "budget_tokens" not in thinking
def test_deepseek_v4_models_in_cost_map():
"""
Test that deepseek-v4-flash and deepseek-v4-pro entries are correctly
configured in model_prices_and_context_window.json.
Prices sourced from https://api-docs.deepseek.com/quick_start/pricing:
- deepseek-v4-flash: $0.14/M input, $0.28/M output
- deepseek-v4-pro: $0.435/M input, $0.87/M output (75% discounted active price)
Closes https://github.com/BerriAI/litellm/issues/26709
"""
import json
from pathlib import Path
json_path = Path(__file__).parents[2] / "model_prices_and_context_window.json"
with open(json_path) as f:
model_cost = json.load(f)
# --- bare model names ---
for key, expected_input, expected_output, expected_cache in [
("deepseek-v4-flash", 1.4e-07, 2.8e-07, 2.8e-09),
("deepseek-v4-pro", 4.35e-07, 8.7e-07, 3.625e-09),
]:
info = model_cost.get(key)
assert info is not None, f"{key} missing from model_prices_and_context_window.json"
assert info["litellm_provider"] == "deepseek"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == expected_input
assert info["output_cost_per_token"] == expected_output
assert info["cache_read_input_token_cost"] == expected_cache
assert info["max_input_tokens"] == 1_000_000
assert info["supports_function_calling"] is True
assert info["supports_tool_choice"] is True
# --- provider-prefixed names ---
for key, expected_input, expected_output, expected_cache in [
("deepseek/deepseek-v4-flash", 1.4e-07, 2.8e-07, 2.8e-09),
("deepseek/deepseek-v4-pro", 4.35e-07, 8.7e-07, 3.625e-09),
]:
info = model_cost.get(key)
assert info is not None, f"{key} missing from model_prices_and_context_window.json"
assert info["litellm_provider"] == "deepseek"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == expected_input
assert info["output_cost_per_token"] == expected_output
assert info["cache_read_input_token_cost"] == expected_cache
assert info["supports_function_calling"] is True
assert info["supports_tool_choice"] is True
def test_deepseek_v4_models_in_backup_cost_map():
"""
Test that deepseek-v4-flash and deepseek-v4-pro entries are correctly
configured in litellm/model_prices_and_context_window_backup.json.
"""
import json
from pathlib import Path
json_path = Path(__file__).parents[2] / "litellm" / "model_prices_and_context_window_backup.json"
with open(json_path) as f:
model_cost = json.load(f)
# --- bare model names ---
for key, expected_input, expected_output, expected_cache in [
("deepseek-v4-flash", 1.4e-07, 2.8e-07, 2.8e-09),
("deepseek-v4-pro", 4.35e-07, 8.7e-07, 3.625e-09),
]:
info = model_cost.get(key)
assert info is not None, f"{key} missing from backup JSON"
assert info["litellm_provider"] == "deepseek"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == expected_input
assert info["output_cost_per_token"] == expected_output
assert info["cache_read_input_token_cost"] == expected_cache
assert info["max_input_tokens"] == 1_000_000
# --- provider-prefixed names ---
for key, expected_input, expected_output, expected_cache in [
("deepseek/deepseek-v4-flash", 1.4e-07, 2.8e-07, 2.8e-09),
("deepseek/deepseek-v4-pro", 4.35e-07, 8.7e-07, 3.625e-09),
]:
info = model_cost.get(key)
assert info is not None, f"{key} missing from backup JSON"
assert info["litellm_provider"] == "deepseek"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == expected_input
assert info["output_cost_per_token"] == expected_output
assert info["cache_read_input_token_cost"] == expected_cache
_FIREWORKS_MODELS = [
(
"accounts/fireworks/models/glm-5p2",
1.4e-06,
4.4e-06,
1.4e-07,
1048576,
131072,
False,
True,
),
(
"accounts/fireworks/models/glm-5p1",
1.4e-06,
4.4e-06,
2.6e-07,
202800,
131072,
False,
True,
),
(
"accounts/fireworks/routers/glm-5p1-fast",
2.8e-06,
8.8e-06,
5.2e-07,
202800,
131072,
False,
True,
),
(
"accounts/fireworks/models/qwen3p7-plus",
4e-07,
1.6e-06,
8e-08,
262144,
65536,
True,
True,
),
(
"accounts/fireworks/models/minimax-m3",
3e-07,
1.2e-06,
6e-08,
512000,
512000,
True,
True,
),
(
"accounts/fireworks/models/minimax-m2p7",
3e-07,
1.2e-06,
6e-08,
196608,
196608,
False,
True,
),
(
"accounts/fireworks/models/kimi-k2p7-code",
9.5e-07,
4e-06,
1.9e-07,
262144,
32768,
True,
True,
),
(
"accounts/fireworks/routers/kimi-k2p7-code-fast",
1.9e-06,
8e-06,
3.8e-07,
262144,
32768,
True,
True,
),
(
"accounts/fireworks/models/kimi-k2p6",
9.5e-07,
4e-06,
1.6e-07,
262144,
32768,
True,
True,
),
(
"accounts/fireworks/routers/kimi-k2p6-fast",
2e-06,
8e-06,
3e-07,
262144,
32768,
True,
True,
),
(
"accounts/fireworks/models/gpt-oss-120b",
1.5e-07,
6e-07,
1.5e-08,
131072,
32768,
False,
True,
),
(
"accounts/fireworks/models/gpt-oss-20b",
7e-08,
3e-07,
3.5e-08,
131072,
32768,
False,
True,
),
(
"accounts/fireworks/models/deepseek-v4-pro",
1.74e-06,
3.48e-06,
1.45e-07,
1048576,
384000,
False,
True,
),
(
"accounts/fireworks/models/deepseek-v4-flash",
1.4e-07,
2.8e-07,
2.8e-08,
1048576,
384000,
False,
True,
),
]
_FIREWORKS_SHORT_FORMS = [
"glm-5p2",
"glm-5p1",
"qwen3p7-plus",
"minimax-m3",
"minimax-m2p7",
"kimi-k2p7-code",
"kimi-k2p6",
"gpt-oss-120b",
"gpt-oss-20b",
"deepseek-v4-pro",
"deepseek-v4-flash",
]
_FIREWORKS_ROUTER_SHORT_FORMS = [
"glm-5p1-fast",
"kimi-k2p6-fast",
"kimi-k2p7-code-fast",
]
def _assert_fireworks_entry(
model_cost,
model_path,
expected_input,
expected_output,
expected_cache,
expected_max_input,
expected_max_output,
expected_vision,
expected_reasoning,
):
info = model_cost.get(f"fireworks_ai/{model_path}")
assert info is not None, f"fireworks_ai/{model_path} missing from model cost map"
assert info["litellm_provider"] == "fireworks_ai"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == expected_input
assert info["output_cost_per_token"] == expected_output
assert info["cache_read_input_token_cost"] == expected_cache
assert info["max_input_tokens"] == expected_max_input
assert info["max_output_tokens"] == expected_max_output
assert info["max_tokens"] == expected_max_output
assert info["supports_function_calling"] is True
assert info["supports_tool_choice"] is True
assert info["supports_reasoning"] is expected_reasoning
assert info["supports_response_schema"] is True
assert info["supports_vision"] is expected_vision
def test_fireworks_models_in_cost_map():
import json
from pathlib import Path
json_path = Path(__file__).parents[2] / "model_prices_and_context_window.json"
with open(json_path) as f:
model_cost = json.load(f)
for entry in _FIREWORKS_MODELS:
_assert_fireworks_entry(model_cost, *entry)
for short in _FIREWORKS_SHORT_FORMS:
long_key = f"fireworks_ai/accounts/fireworks/models/{short}"
short_key = f"fireworks_ai/{short}"
assert model_cost.get(short_key) == model_cost.get(
long_key
), f"short-form {short_key} does not match long-form {long_key}"
for short in _FIREWORKS_ROUTER_SHORT_FORMS:
long_key = f"fireworks_ai/accounts/fireworks/routers/{short}"
short_key = f"fireworks_ai/{short}"
assert model_cost.get(short_key) == model_cost.get(
long_key
), f"short-form {short_key} does not match long-form {long_key}"
def test_fireworks_models_in_backup_cost_map():
import json
from pathlib import Path
json_path = (
Path(__file__).parents[2]
/ "litellm"
/ "model_prices_and_context_window_backup.json"
)
with open(json_path) as f:
model_cost = json.load(f)
for entry in _FIREWORKS_MODELS:
_assert_fireworks_entry(model_cost, *entry)
for short in _FIREWORKS_SHORT_FORMS:
long_key = f"fireworks_ai/accounts/fireworks/models/{short}"
short_key = f"fireworks_ai/{short}"
assert model_cost.get(short_key) == model_cost.get(
long_key
), f"short-form {short_key} does not match long-form {long_key}"
for short in _FIREWORKS_ROUTER_SHORT_FORMS:
long_key = f"fireworks_ai/accounts/fireworks/routers/{short}"
short_key = f"fireworks_ai/{short}"
assert model_cost.get(short_key) == model_cost.get(
long_key
), f"short-form {short_key} does not match long-form {long_key}"
class TestBedrockBaseModelLabelKeepsTools:
"""Regression for #29618: a Bedrock deployment whose ``base_model`` is a friendly
label must not silently drop ``tools``/``tool_choice`` under ``drop_params``."""
TOOLS = [
{
"type": "function",
"function": {
"name": "get_weather",
"parameters": {
"type": "object",
"properties": {"city": {"type": "string"}},
},
},
}
]
def test_base_model_label_keeps_tools_with_drop_params(self):
from litellm.utils import get_optional_params
result = get_optional_params(
model="eu.anthropic.claude-haiku-4-5-20251001-v1:0",
custom_llm_provider="bedrock",
base_model="claude-haiku-4-5",
tools=self.TOOLS,
tool_choice="auto",
drop_params=True,
)
assert "tools" in result
assert "tool_choice" in result
def test_base_model_label_alone_drops_tools(self):
"""Without the real model id the label resolves to no tool support, so passing
the label as ``model`` is exactly what dropped tools before the fix."""
from litellm.utils import get_optional_params
result = get_optional_params(
model="claude-haiku-4-5",
custom_llm_provider="bedrock",
tools=self.TOOLS,
tool_choice="auto",
drop_params=True,
)
assert "tools" not in result
def test_aws_bedrock_project_id_excluded_from_bedrock_optional_params():
"""`aws_bedrock_project_id` is sent as a bedrock-mantle request header, so it
must never reach optional_params (and from there the request body), while
other aws_* params keep flowing for boto3 auth."""
from litellm.utils import get_optional_params
result = get_optional_params(
model="mantle/anthropic.claude-mythos-preview",
custom_llm_provider="bedrock",
max_tokens=10,
aws_bedrock_project_id="proj_abc123def456",
aws_region_name="us-east-1",
)
assert "aws_bedrock_project_id" not in result
assert result["aws_region_name"] == "us-east-1"
class TestGetOptionalParamsTencent:
"""Tests that tencent provider uses TencentChatConfig for parameter mapping."""
def test_tencent_supports_thinking_param(self):
"""Verify get_optional_params for tencent accepts the 'thinking' param."""
from unittest.mock import patch
from litellm.utils import get_optional_params
with patch(
"litellm.llms.tencent.chat.transformation.supports_reasoning",
return_value=True,
):
result = get_optional_params(
model="tencent/deepseek-v4-pro",
custom_llm_provider="tencent",
thinking={"type": "enabled"},
)
assert result.get("thinking") == {"type": "enabled"}
def test_tencent_supports_reasoning_effort(self):
"""Verify get_optional_params for tencent converts reasoning_effort to thinking."""
from unittest.mock import patch
from litellm.utils import get_optional_params
with patch(
"litellm.llms.tencent.chat.transformation.supports_reasoning",
return_value=True,
):
result = get_optional_params(
model="tencent/deepseek-v4-pro",
custom_llm_provider="tencent",
reasoning_effort="medium",
)
assert result.get("thinking") == {"type": "enabled"}
def test_tencent_supported_params_includes_thinking_and_reasoning_effort(self):
"""Verify get_supported_openai_params for tencent includes custom params."""
from unittest.mock import patch
from litellm.litellm_core_utils.get_supported_openai_params import (
get_supported_openai_params,
)
with patch(
"litellm.llms.tencent.chat.transformation.supports_reasoning",
return_value=True,
):
params = get_supported_openai_params(
model="tencent/deepseek-v4-pro",
custom_llm_provider="tencent",
)
assert "thinking" in params
assert "reasoning_effort" in params
def test_tencent_messages_config_routing(self):
"""Verify ProviderConfigManager routes tencent to TencentAnthropicMessagesConfig."""
import litellm
from litellm.llms.tencent.messages.transformation import (
TencentAnthropicMessagesConfig,
)
from litellm.utils import ProviderConfigManager
config = ProviderConfigManager.get_provider_anthropic_messages_config(
model="deepseek-v4-pro",
provider=litellm.LlmProviders.TENCENT,
)
assert isinstance(config, TencentAnthropicMessagesConfig)
assert config.custom_llm_provider == "tencent"
class TestValidateEnvironmentTencent:
"""Tests that validate_environment resolves TENCENT_API_KEY for the tencent provider."""
def test_reports_key_present(self):
with patch.dict(os.environ, {"TENCENT_API_KEY": "sk-tencent"}):
result = litellm.validate_environment(model="tencent/deepseek-v4-pro")
assert result["keys_in_environment"] is True
assert result["missing_keys"] == []
def test_reports_key_missing(self):
with patch.dict(os.environ, {}, clear=True):
result = litellm.validate_environment(model="tencent/deepseek-v4-pro")
assert result["keys_in_environment"] is False
assert "TENCENT_API_KEY" in result["missing_keys"]
class TestVertexEmbeddingEncodingFormat:
"""vertex_ai/gemini embeddings must accept encoding_format="float" — it's
the OpenAI SDK default and float lists are exactly what the vertex API
returns. Other values keep the unsupported-param behavior (drop with
drop_params, raise otherwise). Issue #33173."""
def test_encoding_format_float_is_accepted_and_dropped(self):
optional_params = litellm.utils.get_optional_params_embeddings(
model="gemini-embedding-001",
encoding_format="float",
custom_llm_provider="vertex_ai",
)
assert "encoding_format" not in optional_params
def test_encoding_format_float_accepted_for_gemini_provider(self):
optional_params = litellm.utils.get_optional_params_embeddings(
model="gemini-embedding-001",
encoding_format="float",
custom_llm_provider="gemini",
)
assert "encoding_format" not in optional_params
def test_encoding_format_base64_still_rejected_without_drop_params(self):
with pytest.raises(Exception) as excinfo:
litellm.utils.get_optional_params_embeddings(
model="gemini-embedding-001",
encoding_format="base64",
custom_llm_provider="vertex_ai",
)
assert "encoding_format" in str(excinfo.value)
def test_encoding_format_base64_dropped_with_drop_params(self):
optional_params = litellm.utils.get_optional_params_embeddings(
model="gemini-embedding-001",
encoding_format="base64",
custom_llm_provider="vertex_ai",
drop_params=True,
)
assert "encoding_format" not in optional_params
def test_dimensions_still_mapped(self):
optional_params = litellm.utils.get_optional_params_embeddings(
model="gemini-embedding-001",
encoding_format="float",
dimensions=256,
custom_llm_provider="vertex_ai",
)
assert optional_params.get("outputDimensionality") == 256
@pytest.mark.parametrize(
"model",
[
"vertex_ai/gemini-2.5-flash-image",
"vertex_ai/gemini-3-pro-image",
"vertex_ai/gemini-3-pro-image-preview",
"vertex_ai/gemini-3.1-flash-image",
"vertex_ai/gemini-3.1-flash-image-preview",
"gemini/gemini-2.5-flash-image",
"gemini/gemini-3-pro-image",
"gemini/gemini-3-pro-image-preview",
"gemini/gemini-3.1-flash-image",
"gemini/gemini-3.1-flash-image-preview",
],
)
def test_gemini_image_models_do_not_support_reasoning(
model: str, local_model_cost_map: None
) -> None:
assert model in litellm.model_cost, (
f"{model} is missing from the local model cost map. "
"Add its entry to litellm/model_prices_and_context_window_backup.json."
)
assert litellm.supports_reasoning(model) is False, (
f"{model} incorrectly classified as reasoning-capable. "
"Add 'supports_reasoning: false' to its model_cost entry."
)
PROMPT_CACHE_MESSAGES = [{"role": "user", "content": "the quick brown fox jumps over the lazy dog " * 155}]
@pytest.mark.parametrize(
"model, expected_min_tokens",
[
("claude-opus-4-6", 4096),
("claude-opus-4-7", 2048),
("claude-opus-4-8", 1024),
("claude-fable-5", 512),
],
)
def test_get_prompt_cache_min_tokens_resolves_per_model(
model: str, expected_min_tokens: int, local_model_cost_map: None
) -> None:
"""The smallest cacheable prefix is a per-model property, read from the cost map's
prompt_cache_min_tokens. Anthropic's minimum spans 512..4096 across models and moves in both
directions across releases, so a single global constant is wrong for every model but one."""
assert get_prompt_cache_min_tokens(model=model) == expected_min_tokens
def test_get_prompt_cache_min_tokens_differs_per_platform_for_same_model(local_model_cost_map: None) -> None:
"""The same model can carry a different minimum per platform, so the threshold must come from
the platform's own cost-map entry rather than being derived from the model family name."""
assert get_prompt_cache_min_tokens(model="claude-fable-5") == 512
assert get_prompt_cache_min_tokens(model="anthropic.claude-fable-5") == 1024
assert get_prompt_cache_min_tokens(model="claude-fable-5") != get_prompt_cache_min_tokens(
model="anthropic.claude-fable-5"
)
def test_get_prompt_cache_min_tokens_unmapped_model_falls_back_to_default(local_model_cost_map: None) -> None:
"""get_model_info raises for a model it has no entry for. The resolver must swallow that and
fall back to the default, otherwise the raise reaches callers that would read it as
"not cacheable" -- turning an unknown model into a silently uncacheable one."""
assert get_prompt_cache_min_tokens(model="totally-unknown-model-xyz") == 1024
def test_is_prompt_caching_valid_prompt_uses_per_model_minimum(local_model_cost_map: None) -> None:
"""Regression: a prompt between two models' minimums is cacheable on one and not the other.
A 1403-token prompt clears claude-opus-4-8's 1024 minimum but not claude-opus-4-6's 4096, so
the flat-1024 check reported claude-opus-4-6 as cacheable and the cache write was rejected
upstream. Both assertions must live together: is_prompt_caching_valid_prompt returns False on
any internal error, so the True case is what proves the False case isn't a swallowed exception."""
token_count = litellm.token_counter(
model="claude-opus-4-6", messages=PROMPT_CACHE_MESSAGES, use_default_image_token_count=True
)
assert 1024 <= token_count < 4096, (
f"prompt drifted to {token_count} tokens; it must sit between claude-opus-4-8's 1024 minimum "
"and claude-opus-4-6's 4096 minimum for this test to distinguish them"
)
assert is_prompt_caching_valid_prompt(model="claude-opus-4-6", messages=PROMPT_CACHE_MESSAGES) is False
assert is_prompt_caching_valid_prompt(model="claude-opus-4-8", messages=PROMPT_CACHE_MESSAGES) is True
def test_is_prompt_caching_valid_prompt_explicit_min_token_count_overrides_model(local_model_cost_map: None) -> None:
"""An explicit min_token_count wins over the model-resolved value in both directions. Callers
holding only a model-group alias resolve the threshold themselves and pass it, because an alias
resolves to nothing here and would silently fall back to the default."""
assert (
is_prompt_caching_valid_prompt(model="claude-opus-4-6", messages=PROMPT_CACHE_MESSAGES, min_token_count=512)
is True
)
assert (
is_prompt_caching_valid_prompt(model="claude-opus-4-8", messages=PROMPT_CACHE_MESSAGES, min_token_count=8192)
is False
)
def test_custom_logger_guards_ignore_subclass_instances(monkeypatch: pytest.MonkeyPatch) -> None:
"""Regression LIT-4392: the success/failure existence guards used isinstance, so a user
subclass of a built-in logger already promoted into the callback lists made the guard
report the built-in itself as registered and the configured logger was silently skipped.
The exact-class assertions must hold alongside the subclass assertions: the guards still
have to dedup a second instance of the same class, only a subclass must stop matching."""
from litellm.integrations.custom_logger import CustomLogger
from litellm.utils import (
_custom_logger_class_exists_in_failure_callbacks,
_custom_logger_class_exists_in_success_callbacks,
)
class BuiltinLogger(CustomLogger):
pass
class UserSubclassLogger(BuiltinLogger):
pass
builtin_instance = BuiltinLogger()
monkeypatch.setattr(litellm, "success_callback", [UserSubclassLogger()])
monkeypatch.setattr(litellm, "failure_callback", [UserSubclassLogger()])
monkeypatch.setattr(litellm, "_async_success_callback", [])
monkeypatch.setattr(litellm, "_async_failure_callback", [])
assert _custom_logger_class_exists_in_success_callbacks(builtin_instance) is False
assert _custom_logger_class_exists_in_failure_callbacks(builtin_instance) is False
monkeypatch.setattr(litellm, "success_callback", [BuiltinLogger()])
monkeypatch.setattr(litellm, "failure_callback", [BuiltinLogger()])
assert _custom_logger_class_exists_in_success_callbacks(builtin_instance) is True
assert _custom_logger_class_exists_in_failure_callbacks(builtin_instance) is True
@pytest.mark.asyncio
async def test_s3_v2_success_callback_registers_alongside_user_subclass(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Regression LIT-4392: with a user S3Logger subclass registered via litellm_settings.callbacks
and success_callback ["s3_v2"], the built-in s3_v2 logger was never added and S3 logs were
silently dropped while requests kept returning 200."""
from litellm.integrations.s3_v2 import S3Logger
from litellm.utils import _add_custom_logger_callback_to_specific_event
class UserS3Logger(S3Logger):
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
pass
user_logger = UserS3Logger()
monkeypatch.setattr(litellm, "success_callback", [user_logger, "s3_v2"])
monkeypatch.setattr(litellm, "_async_success_callback", [user_logger])
monkeypatch.setattr(litellm, "failure_callback", [])
monkeypatch.setattr(litellm, "_async_failure_callback", [])
_add_custom_logger_callback_to_specific_event("s3_v2", "success")
assert any(type(cb) is S3Logger for cb in litellm.success_callback)
assert any(type(cb) is S3Logger for cb in litellm._async_success_callback)
assert "s3_v2" not in litellm.success_callback
assert user_logger in litellm.success_callback
@pytest.mark.asyncio
async def test_builtin_string_callback_registers_when_subclass_already_active(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Regression LIT-4392, litellm.callbacks path: the inline dedup in function_setup also
matched subclass instances, so a built-in name in litellm.callbacks was dropped whenever a
user subclass was already promoted into _async_success_callback."""
from litellm.integrations.s3_v2 import S3Logger
class UserS3Logger(S3Logger):
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
pass
user_logger = UserS3Logger()
monkeypatch.setattr(litellm, "callbacks", ["s3_v2"])
monkeypatch.setattr(litellm, "input_callback", [])
monkeypatch.setattr(litellm, "success_callback", [user_logger])
monkeypatch.setattr(litellm, "failure_callback", [])
monkeypatch.setattr(litellm, "_async_success_callback", [user_logger])
monkeypatch.setattr(litellm, "_async_failure_callback", [])
await litellm.acompletion(
model="gpt-5.6",
messages=[{"role": "user", "content": "hi"}],
mock_response="ok",
)
assert any(type(cb) is S3Logger for cb in litellm._async_success_callback)
def test_reapply_runtime_registrations_replays_register_model_overrides(monkeypatch):
"""
register_model is the documented way to override pricing for a model. A
price-data reload swaps litellm.model_cost for a freshly fetched catalog,
so without replaying those registrations the override is silently lost and
the model reverts to upstream pricing.
"""
from litellm import utils as litellm_utils
from litellm.utils import (
_invalidate_model_cost_lowercase_map,
reapply_runtime_model_cost_registrations,
)
monkeypatch.setattr(
litellm_utils,
"_runtime_registered_model_cost",
dict(litellm_utils._runtime_registered_model_cost),
)
# Only the recorded half is under test here; the live-router rebuild is covered
# in test_router_model_cost_isolation.py. Routers built by earlier tests in this
# process stay in the weak set until they are collected, so leaving the callback
# installed would make this depend on when that happens.
monkeypatch.setattr(litellm_utils._LiveDeploymentReplay, "callback", None)
saved_model_cost = litellm.model_cost
try:
litellm.register_model(
model_cost={
"openai/gpt-4o": {
"litellm_provider": "openai",
"mode": "chat",
"input_cost_per_token": 0.000123,
}
}
)
litellm.model_cost = {
"openai/gpt-4o": {
"litellm_provider": "openai",
"mode": "chat",
"input_cost_per_token": 0.000999,
"max_input_tokens": 4242,
}
}
_invalidate_model_cost_lowercase_map()
reapply_runtime_model_cost_registrations()
assert litellm.model_cost["openai/gpt-4o"]["input_cost_per_token"] == 0.000123
assert litellm.model_cost["openai/gpt-4o"]["max_input_tokens"] == 4242
finally:
litellm.model_cost = saved_model_cost
_invalidate_model_cost_lowercase_map()
def test_reapply_runtime_registrations_drops_request_scoped_registrations(monkeypatch):
"""
Per-request custom pricing describes one call, so it must not be re-asserted
over every future catalog. Replaying it would let a one-off price outlive
the catalog generation it was applied to and silently beat fresh upstream
pricing forever, while a durable override registered alongside it survives.
"""
from litellm import utils as litellm_utils
from litellm.utils import (
_invalidate_model_cost_lowercase_map,
reapply_runtime_model_cost_registrations,
)
monkeypatch.setattr(
litellm_utils,
"_runtime_registered_model_cost",
dict(litellm_utils._runtime_registered_model_cost),
)
saved_model_cost = litellm.model_cost
try:
litellm.register_model(
model_cost={"openai/gpt-4o": {"litellm_provider": "openai", "input_cost_per_token": 0.000111}},
persist_across_reloads=True,
)
litellm.register_model(
model_cost={"openai/gpt-4o-mini": {"litellm_provider": "openai", "input_cost_per_token": 0.000222}},
persist_across_reloads=False,
)
litellm.model_cost = {
"openai/gpt-4o": {"litellm_provider": "openai", "input_cost_per_token": 0.000999},
"openai/gpt-4o-mini": {"litellm_provider": "openai", "input_cost_per_token": 0.000888},
}
_invalidate_model_cost_lowercase_map()
reapply_runtime_model_cost_registrations()
assert litellm.model_cost["openai/gpt-4o"]["input_cost_per_token"] == 0.000111
assert litellm.model_cost["openai/gpt-4o-mini"]["input_cost_per_token"] == 0.000888
finally:
litellm.model_cost = saved_model_cost
_invalidate_model_cost_lowercase_map()
def test_ai21_api_key_is_resolved_from_the_documented_env_var(monkeypatch: pytest.MonkeyPatch) -> None:
"""The ai21 branch resolved a misspelled env var, so the name every other ai21 code path
reads, and the only name documented, was ignored."""
monkeypatch.setattr(litellm, "api_key", None)
monkeypatch.setattr(litellm, "ai21_key", None)
monkeypatch.delenv("AI211_API_KEY", raising=False)
monkeypatch.setenv("AI21_API_KEY", "sk-ai21-resolved-from-env")
assert get_api_key(llm_provider="ai21", dynamic_api_key=None) == "sk-ai21-resolved-from-env"
class _JsonCapture(logging.Handler):
def __init__(self):
super().__init__()
self.formatter = JsonFormatter()
self.records: list[dict] = []
self.addFilter(CorrelationContextFilter())
def emit(self, record):
self.records.append(json.loads(self.formatter.format(record)))
def _make_capture_logger(name: str) -> tuple[logging.Logger, _JsonCapture]:
lg = logging.getLogger(name)
cap = _JsonCapture()
lg.addHandler(cap)
lg.setLevel(logging.DEBUG)
return lg, cap
@pytest.mark.asyncio
async def test_wrapper_async_restores_originating_task_context_after_success(monkeypatch):
"""A successful acompletion() dispatches async_success_handler via
asyncio.create_task + the global logging worker - a different Task than the
one running acompletion() itself (this test's own task). That handler's own
restore only fixes up the detached child task it runs in; wrapper_async's own
finally block (in litellm/utils.py) must separately restore the *originating*
task's trace_id/session_id, since nothing else does.
"""
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
trace_id_var.set("outer-trace-wrapper-test")
session_id_var.set("outer-session-wrapper-test")
try:
await litellm.acompletion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "hi"}],
mock_response="Hello there!",
litellm_session_id="mock-call-session",
num_retries=0,
)
assert trace_id_var.get() == "outer-trace-wrapper-test"
assert session_id_var.get() == "outer-session-wrapper-test"
finally:
trace_id_var.set("")
session_id_var.set("")
def test_function_setup_failure_after_logging_construction_restores_context(monkeypatch):
"""If function_setup() constructs Logging() (which already mutated
trace_id_var/session_id_var in __init__) but then raises before returning,
the caller's wrapper() never gets a logging_obj reference to restore from.
function_setup()'s own except block must restore the correlation context
itself in that case, or it leaks into every subsequent log line in this
thread/task until something unrelated happens to reset it."""
from litellm.litellm_core_utils.litellm_logging import Logging
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
def _boom(self, *args, **kwargs):
raise RuntimeError("simulated failure after Logging() construction")
monkeypatch.setattr(Logging, "update_environment_variables", _boom)
trace_id_var.set("pre-setup-failure-trace")
session_id_var.set("pre-setup-failure-session")
try:
with pytest.raises(RuntimeError, match="simulated failure"):
litellm.completion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "hi"}],
mock_response="Hello there!",
litellm_session_id="doomed-call-session",
num_retries=0,
)
assert trace_id_var.get() == "pre-setup-failure-trace"
assert session_id_var.get() == "pre-setup-failure-session"
finally:
trace_id_var.set("")
session_id_var.set("")
def test_function_setup_failure_log_line_shows_outer_not_doomed_ids(monkeypatch):
"""The 'Error in function_setup' diagnostic log line itself must be stamped
with the outer/pre-call correlation ids, not the doomed call's own ids -
restoring context must happen *before* logging the exception, not after,
since the failed call never produces a usable logging object for anything
else to be attributed to."""
from litellm.litellm_core_utils.litellm_logging import Logging
monkeypatch.setattr(litellm, "request_correlation_in_logs", True)
def _boom(self, *args, **kwargs):
raise RuntimeError("simulated failure after Logging() construction")
monkeypatch.setattr(Logging, "update_environment_variables", _boom)
lg, cap = _make_capture_logger("test.function_setup_failure_log_order")
# verbose_logger is a distinct, module-level logger from our throwaway one -
# temporarily attach the same capture handler so we see its own emitted record.
verbose_logger.addHandler(cap)
try:
trace_id_var.set("outer-trace")
session_id_var.set("outer-session")
with pytest.raises(RuntimeError, match="simulated failure"):
litellm.completion(
model="gpt-3.5-turbo",
messages=[{"role": "user", "content": "hi"}],
mock_response="Hello there!",
litellm_session_id="doomed-call-session",
num_retries=0,
)
setup_failure_records = [r for r in cap.records if "Error in function_setup" in r.get("message", "")]
assert len(setup_failure_records) == 1
record = setup_failure_records[0]
assert record.get("session_id") == "outer-session"
assert record.get("trace_id") == "outer-trace"
finally:
verbose_logger.removeHandler(cap)
trace_id_var.set("")
session_id_var.set("")