mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_fix_oauth_credential_forwarding
# Conflicts: # tests/test_litellm/test_main.py
This commit is contained in:
commit
349e7e6990
100 changed files with 3086 additions and 2172 deletions
31
.github/actions/cache-cargo-build/action.yml
vendored
Normal file
31
.github/actions/cache-cargo-build/action.yml
vendored
Normal file
|
|
@ -0,0 +1,31 @@
|
|||
name: "Cache the Rust build"
|
||||
description: >-
|
||||
Cache the Cargo registry and target directory the root package's build needs,
|
||||
so only the first job on a given Cargo.lock compiles the bridge from scratch.
|
||||
|
||||
litellm builds through maturin, which compiles litellm-rust/crates/python-bridge
|
||||
in release mode before it can produce a wheel. `uv sync` therefore pays a full
|
||||
build in every job that installs the workspace: measured at 2m40s per unit shard
|
||||
on 2026-08-21, more than the whole unit tier spends running tests. Nothing caught
|
||||
it, because the uv cache holds wheels uv downloads rather than wheels it builds,
|
||||
and a path dependency whose source moves every commit could never hit that cache
|
||||
anyway. Cargo rebuilds only what changed when its target directory survives, so a
|
||||
warm job pays for the bridge crate alone.
|
||||
|
||||
The key namespace is separate from test-rust.yml's. Both cache the same directory,
|
||||
but that workflow fills it with debug and clippy artifacts, which a release build
|
||||
cannot reuse, and a shared key would let whichever ran first deny the other a save.
|
||||
|
||||
runs:
|
||||
using: composite
|
||||
steps:
|
||||
- name: Restore the Cargo registry and target directory
|
||||
uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0
|
||||
with:
|
||||
path: |
|
||||
~/.cargo/registry
|
||||
~/.cargo/git
|
||||
litellm-rust/target
|
||||
key: ${{ runner.os }}-cargo-release-${{ hashFiles('litellm-rust/Cargo.lock') }}
|
||||
restore-keys: |
|
||||
${{ runner.os }}-cargo-release-
|
||||
10
.github/ci-coverage-allowlist.yml
vendored
10
.github/ci-coverage-allowlist.yml
vendored
|
|
@ -48,16 +48,6 @@ test_paths:
|
|||
choice it informed is settled
|
||||
paths:
|
||||
- tests/code_coverage_tests/test_aio_http_image_conversion.py
|
||||
- reason: >-
|
||||
The last file of a second mirror that sat beside tests/test_litellm and ran nowhere. Its
|
||||
other 33 files landed in the real mirror during August 2026, 30 as moves and 3 by merging
|
||||
their bodies into the live file of the same name. This one cannot follow either route yet:
|
||||
its live twin was rewritten from 1268 lines to 9434, and of the 19 tests here 5 have no
|
||||
counterpart while 25 assertions fail against today's code, so what survives that rewrite
|
||||
is a judgement about the endpoints, not a merge. Revisit by deciding which of the five
|
||||
behaviours still hold
|
||||
paths:
|
||||
- tests/litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py
|
||||
- reason: >-
|
||||
No job invokes this suite and its files mix pure transformation tests with ones driving live
|
||||
vendor vector stores, so assigning them needs a per-file decision
|
||||
|
|
|
|||
9
.github/workflows/_test-unit-base.yml
vendored
9
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -27,7 +27,7 @@ on:
|
|||
default: 20
|
||||
job-timeout-minutes:
|
||||
description: >-
|
||||
Backstop for the whole job. Keep it >= `timeout-minutes` plus 35: 30 for
|
||||
Backstop for the whole job. Keep it >= `timeout-minutes` plus 40: 35 for
|
||||
the per-step ceilings on the setup steps below, and 5 for the runner
|
||||
overhead the job clock charges but no step owns (job init, step
|
||||
transitions, post-job cleanup). That headroom is what makes the test
|
||||
|
|
@ -36,7 +36,7 @@ on:
|
|||
arithmetic, so the sum is passed in rather than computed.
|
||||
required: false
|
||||
type: number
|
||||
default: 55
|
||||
default: 60
|
||||
max-failures:
|
||||
description: "Stop after this many failures"
|
||||
required: false
|
||||
|
|
@ -103,6 +103,11 @@ jobs:
|
|||
restore-keys: |
|
||||
${{ runner.os }}-uv-
|
||||
|
||||
- name: Cache the Rust build
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
timeout-minutes: 5
|
||||
uses: ./.github/actions/cache-cargo-build
|
||||
|
||||
- name: Install dependencies
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
timeout-minutes: 8
|
||||
|
|
|
|||
4
.github/workflows/check-ui-api-types.yml
vendored
4
.github/workflows/check-ui-api-types.yml
vendored
|
|
@ -67,6 +67,10 @@ jobs:
|
|||
restore-keys: |
|
||||
${{ runner.os }}-uv-
|
||||
|
||||
- name: Cache the Rust build
|
||||
if: steps.changes.outputs.relevant == 'true'
|
||||
uses: ./.github/actions/cache-cargo-build
|
||||
|
||||
- name: Install backend dependencies
|
||||
if: steps.changes.outputs.relevant == 'true'
|
||||
run: .github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
|
||||
|
|
|
|||
3
.github/workflows/mutation-test.yml
vendored
3
.github/workflows/mutation-test.yml
vendored
|
|
@ -53,6 +53,9 @@ jobs:
|
|||
restore-keys: |
|
||||
${{ runner.os }}-uv-
|
||||
|
||||
- name: Cache the Rust build
|
||||
uses: ./.github/actions/cache-cargo-build
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router --extra saml
|
||||
|
|
|
|||
|
|
@ -43,6 +43,9 @@ jobs:
|
|||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
- name: Cache the Rust build
|
||||
uses: ./.github/actions/cache-cargo-build
|
||||
|
||||
- name: Cache Prisma binaries
|
||||
uses: ./.github/actions/cache-prisma-binaries
|
||||
|
||||
|
|
|
|||
3
.github/workflows/test-code-quality.yml
vendored
3
.github/workflows/test-code-quality.yml
vendored
|
|
@ -56,6 +56,9 @@ jobs:
|
|||
restore-keys: |
|
||||
${{ runner.os }}-uv-
|
||||
|
||||
- name: Cache the Rust build
|
||||
uses: ./.github/actions/cache-cargo-build
|
||||
|
||||
- name: Install dependencies
|
||||
run: uv sync --frozen --all-groups --all-extras
|
||||
|
||||
|
|
|
|||
4
.github/workflows/test-linting.yml
vendored
4
.github/workflows/test-linting.yml
vendored
|
|
@ -78,6 +78,10 @@ jobs:
|
|||
run: |
|
||||
uv lock --check || (echo "❌ uv.lock is out of sync with pyproject.toml. Run 'uv lock' locally and commit the result." && exit 1)
|
||||
|
||||
- name: Cache the Rust build
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
uses: ./.github/actions/cache-cargo-build
|
||||
|
||||
- name: Install dependencies
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
run: |
|
||||
|
|
|
|||
4
.github/workflows/test-mcp.yml
vendored
4
.github/workflows/test-mcp.yml
vendored
|
|
@ -47,6 +47,10 @@ jobs:
|
|||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
- name: Cache the Rust build
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
uses: ./.github/actions/cache-cargo-build
|
||||
|
||||
- name: Install dependencies
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
run: |
|
||||
|
|
|
|||
|
|
@ -88,6 +88,9 @@ jobs:
|
|||
restore-keys: |
|
||||
${{ runner.os }}-uv-
|
||||
|
||||
- name: Cache the Rust build
|
||||
uses: ./.github/actions/cache-cargo-build
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra google --extra proxy --extra semantic-router
|
||||
|
|
|
|||
|
|
@ -67,6 +67,10 @@ jobs:
|
|||
restore-keys: |
|
||||
${{ runner.os }}-uv-
|
||||
|
||||
- name: Cache the Rust build
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
uses: ./.github/actions/cache-cargo-build
|
||||
|
||||
- name: Install dependencies
|
||||
if: steps.changes.outputs.decision != 'skip'
|
||||
run: |
|
||||
|
|
|
|||
22
.github/workflows/test-unit.yml
vendored
22
.github/workflows/test-unit.yml
vendored
|
|
@ -55,7 +55,7 @@ jobs:
|
|||
workers: 2
|
||||
reruns: 1
|
||||
timeout-minutes: 20
|
||||
job-timeout-minutes: 55
|
||||
job-timeout-minutes: 60
|
||||
|
||||
- shard: enterprise-routing
|
||||
artifact-name: enterprise-routing
|
||||
|
|
@ -67,7 +67,7 @@ jobs:
|
|||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
job-timeout-minutes: 55
|
||||
job-timeout-minutes: 60
|
||||
|
||||
- shard: integrations
|
||||
artifact-name: integrations
|
||||
|
|
@ -75,7 +75,7 @@ jobs:
|
|||
workers: 2
|
||||
reruns: 3
|
||||
timeout-minutes: 20
|
||||
job-timeout-minutes: 55
|
||||
job-timeout-minutes: 60
|
||||
|
||||
- shard: Vertex AI
|
||||
artifact-name: llm-vertex-ai
|
||||
|
|
@ -83,7 +83,7 @@ jobs:
|
|||
workers: 1
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
job-timeout-minutes: 55
|
||||
job-timeout-minutes: 60
|
||||
|
||||
- shard: All Other Providers
|
||||
artifact-name: llm-other-providers
|
||||
|
|
@ -91,7 +91,7 @@ jobs:
|
|||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
job-timeout-minutes: 55
|
||||
job-timeout-minutes: 60
|
||||
|
||||
- shard: misc
|
||||
artifact-name: misc
|
||||
|
|
@ -122,7 +122,7 @@ jobs:
|
|||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
job-timeout-minutes: 55
|
||||
job-timeout-minutes: 60
|
||||
|
||||
- shard: proxy-auth
|
||||
artifact-name: proxy-auth
|
||||
|
|
@ -134,7 +134,7 @@ jobs:
|
|||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
job-timeout-minutes: 55
|
||||
job-timeout-minutes: 60
|
||||
|
||||
- shard: proxy-endpoints
|
||||
artifact-name: proxy-endpoints
|
||||
|
|
@ -171,7 +171,7 @@ jobs:
|
|||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
job-timeout-minutes: 55
|
||||
job-timeout-minutes: 60
|
||||
|
||||
- shard: proxy-server
|
||||
artifact-name: proxy-server
|
||||
|
|
@ -179,7 +179,7 @@ jobs:
|
|||
workers: 4
|
||||
reruns: 2
|
||||
timeout-minutes: 60
|
||||
job-timeout-minutes: 95
|
||||
job-timeout-minutes: 100
|
||||
|
||||
- shard: proxy-infra
|
||||
artifact-name: proxy-infra
|
||||
|
|
@ -198,7 +198,7 @@ jobs:
|
|||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
job-timeout-minutes: 55
|
||||
job-timeout-minutes: 60
|
||||
|
||||
- shard: responses-caching-types
|
||||
artifact-name: responses-caching-types
|
||||
|
|
@ -209,7 +209,7 @@ jobs:
|
|||
workers: 2
|
||||
reruns: 2
|
||||
timeout-minutes: 20
|
||||
job-timeout-minutes: 55
|
||||
job-timeout-minutes: 60
|
||||
uses: ./.github/workflows/_test-unit-base.yml
|
||||
with:
|
||||
test-path: ${{ matrix.test-path }}
|
||||
|
|
|
|||
3
.github/workflows/weekly_load_anomaly.yml
vendored
3
.github/workflows/weekly_load_anomaly.yml
vendored
|
|
@ -47,6 +47,9 @@ jobs:
|
|||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
- name: Cache the Rust build
|
||||
uses: ./.github/actions/cache-cargo-build
|
||||
|
||||
- name: Install dependencies
|
||||
run: |
|
||||
.github/scripts/uv_sync_with_retries.sh --frozen --group ci --group proxy-dev --extra proxy
|
||||
|
|
|
|||
|
|
@ -21,6 +21,7 @@ from litellm.llms.anthropic.experimental_pass_through.context_management import
|
|||
)
|
||||
from litellm.llms.anthropic.experimental_pass_through.utils import (
|
||||
is_reasoning_auto_summary_enabled,
|
||||
local_model_name,
|
||||
)
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import (
|
||||
AnthropicMessagesResponse,
|
||||
|
|
@ -358,9 +359,9 @@ class LiteLLMMessagesToCompletionTransformationHandler:
|
|||
except Exception:
|
||||
pass
|
||||
|
||||
if isinstance(model, str) and model and not model.startswith("responses/"):
|
||||
# Prefix model with "responses/" to route to OpenAI Responses API
|
||||
completion_kwargs["model"] = f"responses/{model}"
|
||||
if isinstance(model, str) and model and "responses/" not in model:
|
||||
local_model: Final = model.removeprefix(f"{custom_llm_provider}/")
|
||||
completion_kwargs["model"] = f"{custom_llm_provider}/responses/{local_model}"
|
||||
|
||||
auto_summary: Final = is_reasoning_auto_summary_enabled()
|
||||
|
||||
|
|
@ -616,7 +617,7 @@ class LiteLLMMessagesToCompletionTransformationHandler:
|
|||
if stream:
|
||||
transformed_stream: Final = ANTHROPIC_ADAPTER.translate_completion_output_params_streaming(
|
||||
completion_response,
|
||||
model=model,
|
||||
model=local_model_name(model, kwargs.get("custom_llm_provider")),
|
||||
tool_name_mapping=tool_name_mapping,
|
||||
polyfill_result=polyfill_result,
|
||||
is_async=True,
|
||||
|
|
@ -750,7 +751,7 @@ class LiteLLMMessagesToCompletionTransformationHandler:
|
|||
if stream:
|
||||
transformed_stream: Final = ANTHROPIC_ADAPTER.translate_completion_output_params_streaming(
|
||||
completion_response,
|
||||
model=model,
|
||||
model=local_model_name(model, kwargs.get("custom_llm_provider")),
|
||||
tool_name_mapping=tool_name_mapping,
|
||||
polyfill_result=polyfill_result,
|
||||
is_async=False,
|
||||
|
|
|
|||
|
|
@ -42,15 +42,46 @@ from .utils import AnthropicMessagesRequestUtils, mock_response
|
|||
_RESPONSES_API_PROVIDERS: Final = frozenset({"openai"})
|
||||
|
||||
|
||||
def _should_route_to_responses_api(custom_llm_provider: str | None) -> bool:
|
||||
"""Return True when the provider should use the Responses API path.
|
||||
def _bridges_to_responses_api(model: str, custom_llm_provider: str) -> bool:
|
||||
from litellm.main import responses_api_bridge_check
|
||||
|
||||
model_info, _ = responses_api_bridge_check(model=model, custom_llm_provider=custom_llm_provider)
|
||||
return model_info.get("mode") == "responses"
|
||||
|
||||
|
||||
def _responses_mode_is_lost_by_prefix_strip(
|
||||
requested_model: str, resolved_model: str, custom_llm_provider: str
|
||||
) -> bool:
|
||||
"""Whether a Responses-only deployment stops looking like one once its provider prefix is stripped.
|
||||
|
||||
``litellm.completion`` re-derives the Responses bridge from the stripped id alone, so a
|
||||
deployment id such as ``perplexity/perplexity/sonar`` (mode ``responses``) is shadowed by the
|
||||
chat entry ``perplexity/sonar`` and would otherwise be sent to chat/completions.
|
||||
"""
|
||||
if requested_model == resolved_model:
|
||||
return False
|
||||
return _bridges_to_responses_api(requested_model, custom_llm_provider) and not _bridges_to_responses_api(
|
||||
resolved_model, custom_llm_provider
|
||||
)
|
||||
|
||||
|
||||
def _should_route_to_responses_api(
|
||||
custom_llm_provider: str | None,
|
||||
requested_model: str | None = None,
|
||||
resolved_model: str | None = None,
|
||||
) -> bool:
|
||||
"""Return True when the request should use the Responses API path.
|
||||
|
||||
Set ``litellm.use_chat_completions_url_for_anthropic_messages = True`` to
|
||||
opt out and route OpenAI/Azure requests through chat/completions instead.
|
||||
"""
|
||||
if litellm.use_chat_completions_url_for_anthropic_messages:
|
||||
return False
|
||||
return custom_llm_provider in _RESPONSES_API_PROVIDERS
|
||||
if custom_llm_provider in _RESPONSES_API_PROVIDERS:
|
||||
return True
|
||||
if custom_llm_provider is None or requested_model is None or resolved_model is None:
|
||||
return False
|
||||
return _responses_mode_is_lost_by_prefix_strip(requested_model, resolved_model, custom_llm_provider)
|
||||
|
||||
|
||||
def _deployment_passes_through_anthropic_messages(model_info: object) -> bool:
|
||||
|
|
@ -533,7 +564,7 @@ def anthropic_messages_handler(
|
|||
_shared_kwargs: Final = dict(
|
||||
max_tokens=max_tokens,
|
||||
messages=messages,
|
||||
model=model,
|
||||
model=original_model,
|
||||
metadata=metadata,
|
||||
stop_sequences=stop_sequences,
|
||||
stream=stream,
|
||||
|
|
@ -551,7 +582,7 @@ def anthropic_messages_handler(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
**kwargs,
|
||||
)
|
||||
if _should_route_to_responses_api(custom_llm_provider):
|
||||
if _should_route_to_responses_api(custom_llm_provider, original_model, model):
|
||||
return LiteLLMMessagesToResponsesAPIHandler.anthropic_messages_handler(**_shared_kwargs)
|
||||
|
||||
# The in-gateway context_management polyfill runs inside
|
||||
|
|
|
|||
|
|
@ -44,6 +44,7 @@ class AnthropicMessagesRequestUtils:
|
|||
filtered_params: Final = {k: v for k, v in params.items() if k in valid_keys and v is not None}
|
||||
if model is not None:
|
||||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
|
||||
|
||||
AnthropicConfig._maybe_drop_speed_param(
|
||||
model=model,
|
||||
|
|
@ -51,6 +52,16 @@ class AnthropicMessagesRequestUtils:
|
|||
drop_params=drop_params,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
for param in ("temperature", "top_p", "top_k"):
|
||||
if param in filtered_params:
|
||||
AnthropicModelInfo._apply_sampling_param( # pyright: ignore[reportPrivateUsage] # same gating the /chat/completions path applies; forking it would drift
|
||||
optional_params=filtered_params,
|
||||
model=model,
|
||||
param=param,
|
||||
value=filtered_params.pop(param),
|
||||
drop_params=drop_params,
|
||||
output_key=param,
|
||||
)
|
||||
return cast(AnthropicMessagesRequestOptionalParams, filtered_params)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -19,6 +19,7 @@ from litellm.types.llms.anthropic_messages.anthropic_response import (
|
|||
)
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
|
||||
from ..utils import local_model_name
|
||||
from .streaming_iterator import AnthropicResponsesStreamWrapper
|
||||
from .transformation import LiteLLMAnthropicToResponsesAPIAdapter
|
||||
|
||||
|
|
@ -179,7 +180,9 @@ class LiteLLMMessagesToResponsesAPIHandler:
|
|||
result: Final = await litellm.aresponses(**responses_kwargs)
|
||||
|
||||
if stream:
|
||||
wrapper: Final = AnthropicResponsesStreamWrapper(responses_stream=result, model=model)
|
||||
wrapper: Final = AnthropicResponsesStreamWrapper(
|
||||
responses_stream=result, model=local_model_name(model, kwargs.get("custom_llm_provider"))
|
||||
)
|
||||
return wrapper.async_anthropic_sse_wrapper()
|
||||
|
||||
if not isinstance(result, ResponsesAPIResponse):
|
||||
|
|
@ -257,7 +260,9 @@ class LiteLLMMessagesToResponsesAPIHandler:
|
|||
result: Final = litellm.responses(**responses_kwargs)
|
||||
|
||||
if stream:
|
||||
wrapper: Final = AnthropicResponsesStreamWrapper(responses_stream=result, model=model)
|
||||
wrapper: Final = AnthropicResponsesStreamWrapper(
|
||||
responses_stream=result, model=local_model_name(model, kwargs.get("custom_llm_provider"))
|
||||
)
|
||||
return wrapper.async_anthropic_sse_wrapper()
|
||||
|
||||
if not isinstance(result, ResponsesAPIResponse):
|
||||
|
|
|
|||
|
|
@ -13,6 +13,11 @@ def prompt_cache_key_from_user_id(user_id: object) -> str | None:
|
|||
return str(user_id)[:OPENAI_MAX_PROMPT_CACHE_KEY_LENGTH] or None
|
||||
|
||||
|
||||
def local_model_name(model: str, custom_llm_provider: object) -> str:
|
||||
"""The id the provider itself knows, for reporting back to the caller in ``message_start``."""
|
||||
return model.removeprefix(f"{custom_llm_provider}/") if isinstance(custom_llm_provider, str) else model
|
||||
|
||||
|
||||
def is_reasoning_auto_summary_enabled() -> bool:
|
||||
"""Check whether the default 'summary: detailed' injection is enabled (opt-in)."""
|
||||
return litellm.reasoning_auto_summary or os.getenv("LITELLM_REASONING_AUTO_SUMMARY", "false").lower() == "true"
|
||||
|
|
|
|||
|
|
@ -4881,6 +4881,38 @@
|
|||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"azure/gpt-audio-mini": {
|
||||
"deprecation_date": "2027-04-06",
|
||||
"input_cost_per_audio_token": 1e-05,
|
||||
"input_cost_per_token": 6e-07,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 16384,
|
||||
"max_tokens": 16384,
|
||||
"mode": "chat",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"audio"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text",
|
||||
"audio"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_prompt_caching": false,
|
||||
"supports_reasoning": false,
|
||||
"supports_response_schema": false,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"azure/gpt-audio-mini-2025-10-06": {
|
||||
"deprecation_date": "2027-04-06",
|
||||
"input_cost_per_audio_token": 1e-05,
|
||||
|
|
@ -5094,6 +5126,38 @@
|
|||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"azure/gpt-realtime-mini": {
|
||||
"cache_creation_input_audio_token_cost": 3e-07,
|
||||
"cache_read_input_token_cost": 6e-08,
|
||||
"input_cost_per_audio_token": 1e-05,
|
||||
"input_cost_per_image": 8e-07,
|
||||
"input_cost_per_token": 6e-07,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supported_endpoints": [
|
||||
"/v1/realtime"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text",
|
||||
"audio"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_audio_output": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"azure/gpt-realtime-mini-2025-10-06": {
|
||||
"cache_creation_input_audio_token_cost": 3e-07,
|
||||
"cache_read_input_token_cost": 6e-08,
|
||||
|
|
@ -26015,33 +26079,33 @@
|
|||
"supports_minimal_reasoning_effort": true
|
||||
},
|
||||
"gpt-5.6": {
|
||||
"cache_creation_input_token_cost": 6.25e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1.25e-05,
|
||||
"cache_creation_input_token_cost_above_272k_tokens_flex": 6.25e-06,
|
||||
"cache_creation_input_token_cost_flex": 3.125e-06,
|
||||
"cache_creation_input_token_cost_priority": 1.25e-05,
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 1e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens_flex": 5e-07,
|
||||
"cache_read_input_token_cost_flex": 2.5e-07,
|
||||
"cache_read_input_token_cost_priority": 1e-06,
|
||||
"input_cost_per_token": 5e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 1e-05,
|
||||
"input_cost_per_token_above_272k_tokens_flex": 5e-06,
|
||||
"input_cost_per_token_batches": 2.5e-06,
|
||||
"input_cost_per_token_flex": 2.5e-06,
|
||||
"input_cost_per_token_priority": 1e-05,
|
||||
"cache_creation_input_token_cost": 5e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1e-05,
|
||||
"cache_creation_input_token_cost_above_272k_tokens_flex": 5e-06,
|
||||
"cache_creation_input_token_cost_flex": 2.5e-06,
|
||||
"cache_creation_input_token_cost_priority": 1e-05,
|
||||
"cache_read_input_token_cost": 4e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 8e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens_flex": 4e-07,
|
||||
"cache_read_input_token_cost_flex": 2e-07,
|
||||
"cache_read_input_token_cost_priority": 8e-07,
|
||||
"input_cost_per_token": 4e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 8e-06,
|
||||
"input_cost_per_token_above_272k_tokens_flex": 4e-06,
|
||||
"input_cost_per_token_batches": 2e-06,
|
||||
"input_cost_per_token_flex": 2e-06,
|
||||
"input_cost_per_token_priority": 8e-06,
|
||||
"litellm_provider": "openai",
|
||||
"max_input_tokens": 922000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 4.5e-05,
|
||||
"output_cost_per_token_above_272k_tokens_flex": 2.25e-05,
|
||||
"output_cost_per_token_batches": 1.5e-05,
|
||||
"output_cost_per_token_flex": 1.5e-05,
|
||||
"output_cost_per_token_priority": 6e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 3e-05,
|
||||
"output_cost_per_token_above_272k_tokens_flex": 1.5e-05,
|
||||
"output_cost_per_token_batches": 1e-05,
|
||||
"output_cost_per_token_flex": 1e-05,
|
||||
"output_cost_per_token_priority": 4e-05,
|
||||
"regional_processing_uplift_multiplier_eu": 1.1,
|
||||
"regional_processing_uplift_multiplier_us": 1.1,
|
||||
"search_context_cost_per_query": {
|
||||
|
|
@ -26078,33 +26142,33 @@
|
|||
"supports_xhigh_reasoning_effort": true
|
||||
},
|
||||
"gpt-5.6-sol": {
|
||||
"cache_creation_input_token_cost": 6.25e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1.25e-05,
|
||||
"cache_creation_input_token_cost_above_272k_tokens_flex": 6.25e-06,
|
||||
"cache_creation_input_token_cost_flex": 3.125e-06,
|
||||
"cache_creation_input_token_cost_priority": 1.25e-05,
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 1e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens_flex": 5e-07,
|
||||
"cache_read_input_token_cost_flex": 2.5e-07,
|
||||
"cache_read_input_token_cost_priority": 1e-06,
|
||||
"input_cost_per_token": 5e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 1e-05,
|
||||
"input_cost_per_token_above_272k_tokens_flex": 5e-06,
|
||||
"input_cost_per_token_batches": 2.5e-06,
|
||||
"input_cost_per_token_flex": 2.5e-06,
|
||||
"input_cost_per_token_priority": 1e-05,
|
||||
"cache_creation_input_token_cost": 5e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1e-05,
|
||||
"cache_creation_input_token_cost_above_272k_tokens_flex": 5e-06,
|
||||
"cache_creation_input_token_cost_flex": 2.5e-06,
|
||||
"cache_creation_input_token_cost_priority": 1e-05,
|
||||
"cache_read_input_token_cost": 4e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 8e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens_flex": 4e-07,
|
||||
"cache_read_input_token_cost_flex": 2e-07,
|
||||
"cache_read_input_token_cost_priority": 8e-07,
|
||||
"input_cost_per_token": 4e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 8e-06,
|
||||
"input_cost_per_token_above_272k_tokens_flex": 4e-06,
|
||||
"input_cost_per_token_batches": 2e-06,
|
||||
"input_cost_per_token_flex": 2e-06,
|
||||
"input_cost_per_token_priority": 8e-06,
|
||||
"litellm_provider": "openai",
|
||||
"max_input_tokens": 922000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 4.5e-05,
|
||||
"output_cost_per_token_above_272k_tokens_flex": 2.25e-05,
|
||||
"output_cost_per_token_batches": 1.5e-05,
|
||||
"output_cost_per_token_flex": 1.5e-05,
|
||||
"output_cost_per_token_priority": 6e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 3e-05,
|
||||
"output_cost_per_token_above_272k_tokens_flex": 1.5e-05,
|
||||
"output_cost_per_token_batches": 1e-05,
|
||||
"output_cost_per_token_flex": 1e-05,
|
||||
"output_cost_per_token_priority": 4e-05,
|
||||
"regional_processing_uplift_multiplier_eu": 1.1,
|
||||
"regional_processing_uplift_multiplier_us": 1.1,
|
||||
"search_context_cost_per_query": {
|
||||
|
|
@ -26346,19 +26410,19 @@
|
|||
"supports_parallel_function_calling": true
|
||||
},
|
||||
"daybreak-blue-latest": {
|
||||
"cache_creation_input_token_cost": 6.25e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1.25e-05,
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 1e-06,
|
||||
"input_cost_per_token": 5e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 1e-05,
|
||||
"cache_creation_input_token_cost": 5e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1e-05,
|
||||
"cache_read_input_token_cost": 4e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 8e-07,
|
||||
"input_cost_per_token": 4e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 8e-06,
|
||||
"litellm_provider": "openai",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 4.5e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 3e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses"
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ from litellm.proxy._types import (
|
|||
SpecialMCPServerName,
|
||||
SpecialMCPServerNames,
|
||||
UserAPIKeyAuth,
|
||||
user_api_key_has_admin_view,
|
||||
)
|
||||
from litellm.proxy.auth.ip_address_utils import IPAddressUtils
|
||||
from litellm.proxy.auth.user_api_key_auth import (
|
||||
|
|
@ -1785,11 +1786,14 @@ class MCPRequestHandler:
|
|||
global_mcp_server_manager,
|
||||
)
|
||||
|
||||
# An OPEN channel (allow_all_keys, the user's own BYOM) makes the server REACHABLE through the
|
||||
# user, though no grant source names it — without this the union returns [], listable but
|
||||
# uninvokable. Reachability is ALL it confers, NOT a ceiling waiver: the user's own
|
||||
# mcp_tool_permissions and org tool ceiling still bind, exactly as a key's do on an allow_all server.
|
||||
reachable_via_open_channel: Final = server_id in await global_mcp_server_manager.operator_open_server_ids(auth)
|
||||
# An OPEN channel (allow_all_keys, the user's own BYOM, an unscoped admin-view role) makes the
|
||||
# server REACHABLE through the user, though no grant source names it — without this the union
|
||||
# returns [], listable but uninvokable. Reachability is ALL it confers, NOT a ceiling waiver:
|
||||
# the user's own mcp_tool_permissions and org tool ceiling still bind, exactly as a key's do
|
||||
# on an allow_all server or an admin key's do on any server.
|
||||
reachable_via_open_channel: Final = server_id in await global_mcp_server_manager.operator_open_server_ids(
|
||||
auth
|
||||
) or await MCPRequestHandler.admin_view_unscoped(auth)
|
||||
|
||||
allowed: Final[set[str]] = set()
|
||||
for source, granted in await MCPRequestHandler.admitted_source_grants(auth):
|
||||
|
|
@ -2723,6 +2727,32 @@ class MCPRequestHandler:
|
|||
entitled_servers: Final = await MCPRequestHandler._get_allowed_mcp_servers_for_user(user_api_key_auth)
|
||||
return entitled_servers is None or len(entitled_servers) > 0
|
||||
|
||||
@staticmethod
|
||||
async def admin_view_unscoped(user_api_key_auth: UserAPIKeyAuth | None = None) -> bool:
|
||||
"""Whether this principal's admin-view role grants the unscoped MCP resolution, whatever
|
||||
credential carries it (admin key, dashboard session, or OAuth-admitted session subject).
|
||||
|
||||
Two bounds disqualify, one per ownership of the row. A CREDENTIAL's explicit
|
||||
``object_permission.mcp_servers`` scope wins even for admins, including the empty list. An
|
||||
admitted subject's object_permission is the user's own row, whose ``mcp_servers`` column is
|
||||
[] by DB default, so for that shape the row binds through the entitlement ceiling instead
|
||||
(any non-empty entitlement, or an unresolved one, disqualifies), exactly as
|
||||
``operator_open_server_ids`` reads the same row. The one owner of this predicate: the
|
||||
server-axis registry resolution in ``get_allowed_mcp_servers`` and the tools-axis open
|
||||
channel in ``_resolve_admitted_subject_tools`` both consult it, so the two axes cannot
|
||||
disagree."""
|
||||
if user_api_key_auth is None or not user_api_key_has_admin_view(user_api_key_auth):
|
||||
return False
|
||||
object_permission: Final = user_api_key_auth.object_permission
|
||||
credential_scoped: Final = (
|
||||
not _is_mcp_admitted_user_subject(user_api_key_auth)
|
||||
and object_permission is not None
|
||||
and object_permission.mcp_servers is not None
|
||||
)
|
||||
if credential_scoped:
|
||||
return False
|
||||
return not await MCPRequestHandler._user_places_mcp_ceiling(user_api_key_auth)
|
||||
|
||||
@staticmethod
|
||||
async def _apply_user_tool_ceiling(
|
||||
allowed_tools: Sequence[str] | None,
|
||||
|
|
|
|||
|
|
@ -2943,17 +2943,14 @@ class MCPServerManager:
|
|||
2. If admin and no object_permission, return all servers
|
||||
3. Otherwise, use standard permission checks
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view
|
||||
|
||||
allow_all_server_ids: Final = self.get_allow_all_keys_server_ids()
|
||||
|
||||
# A keyless admitted subject is resolved per grant source, and channel decisions that are
|
||||
# absolute for a scoped KEY credential are not absolute for it: its own opt-out silences its
|
||||
# own source (handled per source in the resolver), never its teams' grants, and its admin
|
||||
# role does not swallow the grant model — a session bearer is a third-party client
|
||||
# credential, not the dashboard, so an admin signing in through the connect flow gets their
|
||||
# grants like anyone else rather than handing the client the full registry ahead of every
|
||||
# per-team org ceiling.
|
||||
# own source (handled per source in the resolver), never its teams' grants. Its admin role
|
||||
# rides the HUMAN, not the credential: an admin's session resolves the same registry their
|
||||
# dashboard shows (connect-page parity), bounded like an admin key by explicit
|
||||
# object_permission scope, the entitlement ceiling, and the session resource scope below.
|
||||
is_admitted_subject: Final = _is_mcp_admitted_user_subject(user_api_key_auth)
|
||||
|
||||
# The key explicitly opted out of every MCP server. Return zero before
|
||||
|
|
@ -2982,26 +2979,16 @@ class MCPServerManager:
|
|||
)
|
||||
|
||||
try:
|
||||
# If admin but NO explicit object permission, get all servers (never for an admitted
|
||||
# subject — see is_admitted_subject above)
|
||||
if (
|
||||
user_api_key_auth
|
||||
and not is_admitted_subject
|
||||
and _user_has_admin_view(user_api_key_auth)
|
||||
and not has_explicit_object_permission
|
||||
# An entitlement attached to the HUMAN binds them whatever their role: it is the
|
||||
# person's scope, not the credential's, so an admin role is not a waiver of it. An
|
||||
# UNRESOLVED entitlement also skips the shortcut, so the resolver denies rather than
|
||||
# handing over the whole registry on a transient fault.
|
||||
and not await MCPRequestHandler._user_places_mcp_ceiling(user_api_key_auth)
|
||||
):
|
||||
verbose_logger.debug("Admin user without explicit object_permission - returning all servers")
|
||||
return list(self.get_registry().keys())
|
||||
|
||||
# Get allowed servers from object permissions (respects object_permission even for admins)
|
||||
allowed_mcp_servers: Final = await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth)
|
||||
verbose_logger.debug("Allowed MCP Servers for user api key auth: %s", allowed_mcp_servers)
|
||||
combined_servers: Final = set(allowed_mcp_servers)
|
||||
# Admin view with no explicit object permission and no entitlement ceiling resolves the
|
||||
# whole registry, for keys AND admitted session subjects alike (one predicate owns the
|
||||
# question). Seeded into the union rather than returned early so the session resource
|
||||
# scope below still bounds a per-server envelope held by an admin.
|
||||
combined_servers: Final = (
|
||||
set(self.get_registry().keys())
|
||||
if await MCPRequestHandler.admin_view_unscoped(user_api_key_auth)
|
||||
else set(await MCPRequestHandler.get_allowed_mcp_servers(user_api_key_auth))
|
||||
)
|
||||
verbose_logger.debug("Allowed MCP Servers for user api key auth: %s", combined_servers)
|
||||
combined_servers.update(
|
||||
await self.operator_open_server_ids(
|
||||
user_api_key_auth,
|
||||
|
|
|
|||
|
|
@ -4881,6 +4881,38 @@
|
|||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"azure/gpt-audio-mini": {
|
||||
"deprecation_date": "2027-04-06",
|
||||
"input_cost_per_audio_token": 1e-05,
|
||||
"input_cost_per_token": 6e-07,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 16384,
|
||||
"max_tokens": 16384,
|
||||
"mode": "chat",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"audio"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text",
|
||||
"audio"
|
||||
],
|
||||
"supports_function_calling": true,
|
||||
"supports_native_streaming": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_prompt_caching": false,
|
||||
"supports_reasoning": false,
|
||||
"supports_response_schema": false,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"azure/gpt-audio-mini-2025-10-06": {
|
||||
"deprecation_date": "2027-04-06",
|
||||
"input_cost_per_audio_token": 1e-05,
|
||||
|
|
@ -5094,6 +5126,38 @@
|
|||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"azure/gpt-realtime-mini": {
|
||||
"cache_creation_input_audio_token_cost": 3e-07,
|
||||
"cache_read_input_token_cost": 6e-08,
|
||||
"input_cost_per_audio_token": 1e-05,
|
||||
"input_cost_per_image": 8e-07,
|
||||
"input_cost_per_token": 6e-07,
|
||||
"litellm_provider": "azure",
|
||||
"max_input_tokens": 32000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "realtime",
|
||||
"output_cost_per_audio_token": 2e-05,
|
||||
"output_cost_per_token": 2.4e-06,
|
||||
"supported_endpoints": [
|
||||
"/v1/realtime"
|
||||
],
|
||||
"supported_modalities": [
|
||||
"text",
|
||||
"image",
|
||||
"audio"
|
||||
],
|
||||
"supported_output_modalities": [
|
||||
"text",
|
||||
"audio"
|
||||
],
|
||||
"supports_audio_input": true,
|
||||
"supports_audio_output": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_system_messages": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"azure/gpt-realtime-mini-2025-10-06": {
|
||||
"cache_creation_input_audio_token_cost": 3e-07,
|
||||
"cache_read_input_token_cost": 6e-08,
|
||||
|
|
@ -26015,33 +26079,33 @@
|
|||
"supports_minimal_reasoning_effort": true
|
||||
},
|
||||
"gpt-5.6": {
|
||||
"cache_creation_input_token_cost": 6.25e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1.25e-05,
|
||||
"cache_creation_input_token_cost_above_272k_tokens_flex": 6.25e-06,
|
||||
"cache_creation_input_token_cost_flex": 3.125e-06,
|
||||
"cache_creation_input_token_cost_priority": 1.25e-05,
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 1e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens_flex": 5e-07,
|
||||
"cache_read_input_token_cost_flex": 2.5e-07,
|
||||
"cache_read_input_token_cost_priority": 1e-06,
|
||||
"input_cost_per_token": 5e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 1e-05,
|
||||
"input_cost_per_token_above_272k_tokens_flex": 5e-06,
|
||||
"input_cost_per_token_batches": 2.5e-06,
|
||||
"input_cost_per_token_flex": 2.5e-06,
|
||||
"input_cost_per_token_priority": 1e-05,
|
||||
"cache_creation_input_token_cost": 5e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1e-05,
|
||||
"cache_creation_input_token_cost_above_272k_tokens_flex": 5e-06,
|
||||
"cache_creation_input_token_cost_flex": 2.5e-06,
|
||||
"cache_creation_input_token_cost_priority": 1e-05,
|
||||
"cache_read_input_token_cost": 4e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 8e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens_flex": 4e-07,
|
||||
"cache_read_input_token_cost_flex": 2e-07,
|
||||
"cache_read_input_token_cost_priority": 8e-07,
|
||||
"input_cost_per_token": 4e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 8e-06,
|
||||
"input_cost_per_token_above_272k_tokens_flex": 4e-06,
|
||||
"input_cost_per_token_batches": 2e-06,
|
||||
"input_cost_per_token_flex": 2e-06,
|
||||
"input_cost_per_token_priority": 8e-06,
|
||||
"litellm_provider": "openai",
|
||||
"max_input_tokens": 922000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 4.5e-05,
|
||||
"output_cost_per_token_above_272k_tokens_flex": 2.25e-05,
|
||||
"output_cost_per_token_batches": 1.5e-05,
|
||||
"output_cost_per_token_flex": 1.5e-05,
|
||||
"output_cost_per_token_priority": 6e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 3e-05,
|
||||
"output_cost_per_token_above_272k_tokens_flex": 1.5e-05,
|
||||
"output_cost_per_token_batches": 1e-05,
|
||||
"output_cost_per_token_flex": 1e-05,
|
||||
"output_cost_per_token_priority": 4e-05,
|
||||
"regional_processing_uplift_multiplier_eu": 1.1,
|
||||
"regional_processing_uplift_multiplier_us": 1.1,
|
||||
"search_context_cost_per_query": {
|
||||
|
|
@ -26078,33 +26142,33 @@
|
|||
"supports_xhigh_reasoning_effort": true
|
||||
},
|
||||
"gpt-5.6-sol": {
|
||||
"cache_creation_input_token_cost": 6.25e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1.25e-05,
|
||||
"cache_creation_input_token_cost_above_272k_tokens_flex": 6.25e-06,
|
||||
"cache_creation_input_token_cost_flex": 3.125e-06,
|
||||
"cache_creation_input_token_cost_priority": 1.25e-05,
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 1e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens_flex": 5e-07,
|
||||
"cache_read_input_token_cost_flex": 2.5e-07,
|
||||
"cache_read_input_token_cost_priority": 1e-06,
|
||||
"input_cost_per_token": 5e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 1e-05,
|
||||
"input_cost_per_token_above_272k_tokens_flex": 5e-06,
|
||||
"input_cost_per_token_batches": 2.5e-06,
|
||||
"input_cost_per_token_flex": 2.5e-06,
|
||||
"input_cost_per_token_priority": 1e-05,
|
||||
"cache_creation_input_token_cost": 5e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1e-05,
|
||||
"cache_creation_input_token_cost_above_272k_tokens_flex": 5e-06,
|
||||
"cache_creation_input_token_cost_flex": 2.5e-06,
|
||||
"cache_creation_input_token_cost_priority": 1e-05,
|
||||
"cache_read_input_token_cost": 4e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 8e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens_flex": 4e-07,
|
||||
"cache_read_input_token_cost_flex": 2e-07,
|
||||
"cache_read_input_token_cost_priority": 8e-07,
|
||||
"input_cost_per_token": 4e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 8e-06,
|
||||
"input_cost_per_token_above_272k_tokens_flex": 4e-06,
|
||||
"input_cost_per_token_batches": 2e-06,
|
||||
"input_cost_per_token_flex": 2e-06,
|
||||
"input_cost_per_token_priority": 8e-06,
|
||||
"litellm_provider": "openai",
|
||||
"max_input_tokens": 922000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 4.5e-05,
|
||||
"output_cost_per_token_above_272k_tokens_flex": 2.25e-05,
|
||||
"output_cost_per_token_batches": 1.5e-05,
|
||||
"output_cost_per_token_flex": 1.5e-05,
|
||||
"output_cost_per_token_priority": 6e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 3e-05,
|
||||
"output_cost_per_token_above_272k_tokens_flex": 1.5e-05,
|
||||
"output_cost_per_token_batches": 1e-05,
|
||||
"output_cost_per_token_flex": 1e-05,
|
||||
"output_cost_per_token_priority": 4e-05,
|
||||
"regional_processing_uplift_multiplier_eu": 1.1,
|
||||
"regional_processing_uplift_multiplier_us": 1.1,
|
||||
"search_context_cost_per_query": {
|
||||
|
|
@ -26346,19 +26410,19 @@
|
|||
"supports_parallel_function_calling": true
|
||||
},
|
||||
"daybreak-blue-latest": {
|
||||
"cache_creation_input_token_cost": 6.25e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1.25e-05,
|
||||
"cache_read_input_token_cost": 5e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 1e-06,
|
||||
"input_cost_per_token": 5e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 1e-05,
|
||||
"cache_creation_input_token_cost": 5e-06,
|
||||
"cache_creation_input_token_cost_above_272k_tokens": 1e-05,
|
||||
"cache_read_input_token_cost": 4e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 8e-07,
|
||||
"input_cost_per_token": 4e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 8e-06,
|
||||
"litellm_provider": "openai",
|
||||
"max_input_tokens": 1050000,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 4.5e-05,
|
||||
"output_cost_per_token": 2e-05,
|
||||
"output_cost_per_token_above_272k_tokens": 3e-05,
|
||||
"supported_endpoints": [
|
||||
"/v1/chat/completions",
|
||||
"/v1/responses"
|
||||
|
|
|
|||
|
|
@ -341,9 +341,13 @@ filterwarnings = [
|
|||
paths_to_mutate = [
|
||||
"litellm/proxy/management_endpoints/",
|
||||
]
|
||||
# Only the unit tier that maps to paths_to_mutate. mutmut times and
|
||||
# coverage-maps this whole set once before mutating, so a tier that needs a
|
||||
# seeded database (tests/proxy_behavior/) kills the run before it starts, and
|
||||
# a mutation score is only meaningful against the tests that claim to cover
|
||||
# the mutated code anyway.
|
||||
tests_dir = [
|
||||
"tests/test_litellm/proxy/management_endpoints/",
|
||||
"tests/proxy_behavior/management/",
|
||||
]
|
||||
also_copy = [
|
||||
"litellm/",
|
||||
|
|
@ -360,10 +364,16 @@ mutate_only_covered_lines = true
|
|||
# - rerunning a "failed" test on a mutant would mask which mutants are killed
|
||||
# vs. survive, so reruns are wrong for mutation testing regardless.
|
||||
# - xdist is unnecessary inside mutmut (mutmut handles its own parallelism).
|
||||
# test_saml_sso.py cannot run inside mutmut's mutants/ sandbox: the copied tree
|
||||
# re-imports cryptography's hash classes under a second identity, so x509 .sign()
|
||||
# rejects the SHA256 instance the fixture builds with "Algorithm must be a
|
||||
# registered hash algorithm". Nothing to do with mutation coverage, and one
|
||||
# erroring test is enough to end the stats phase before any mutant runs.
|
||||
pytest_add_cli_args = [
|
||||
"-p", "no:retry",
|
||||
"-p", "no:rerunfailures",
|
||||
"-p", "no:xdist",
|
||||
"--ignore=tests/test_litellm/proxy/management_endpoints/test_saml_sso.py",
|
||||
]
|
||||
|
||||
[tool.coverage.run]
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ import tomllib
|
|||
from collections import defaultdict
|
||||
from difflib import SequenceMatcher
|
||||
from pathlib import Path
|
||||
from typing import Final, NamedTuple
|
||||
from textwrap import dedent
|
||||
|
||||
ROOT = Path(__file__).resolve().parent.parent
|
||||
|
|
@ -33,16 +34,24 @@ def load_mutmut_config() -> dict:
|
|||
return tomllib.load(f)["tool"]["mutmut"]
|
||||
|
||||
|
||||
def get_survivors() -> list[str]:
|
||||
class MutmutResults(NamedTuple):
|
||||
survivors: tuple[str, ...]
|
||||
reported: int
|
||||
|
||||
|
||||
def get_survivors() -> MutmutResults:
|
||||
proc = subprocess.run(
|
||||
[*MUTMUT_INVOCATION, "results"], capture_output=True, text=True, check=False
|
||||
)
|
||||
survivors = []
|
||||
for line in proc.stdout.splitlines():
|
||||
m = re.match(r"\s*(\S+):\s*survived\s*$", line)
|
||||
if m:
|
||||
survivors.append(m.group(1))
|
||||
return survivors
|
||||
verdicts = tuple(
|
||||
m.groups()
|
||||
for line in proc.stdout.splitlines()
|
||||
if (m := re.match(r"\s*(\S+):\s*(\S.*?)\s*$", line))
|
||||
)
|
||||
return MutmutResults(
|
||||
survivors=tuple(name for name, verdict in verdicts if verdict == "survived"),
|
||||
reported=len(verdicts),
|
||||
)
|
||||
|
||||
|
||||
def get_mutmut_show(mutant_name: str) -> str:
|
||||
|
|
@ -222,7 +231,52 @@ def render_meta_style_mutant(
|
|||
return "\n".join(out)
|
||||
|
||||
|
||||
def render(config: dict, survivors: list[str], stats: dict | None) -> str:
|
||||
RESOLVED_KEYS: Final = frozenset({"killed", "survived", "total"})
|
||||
|
||||
|
||||
def unresolved_counts(stats: dict) -> dict[str, int]:
|
||||
"""Every non-zero count that is neither a kill nor a survivor means a mutant did not
|
||||
reach the tests. Reading it as "anything else" rather than as a list of known statuses
|
||||
keeps a status this reporter has never met from passing as a clean sweep."""
|
||||
return {k: v for k, v in sorted(stats.items()) if k not in RESOLVED_KEYS and isinstance(v, int) and v > 0}
|
||||
|
||||
|
||||
def clean_sweep_is_provable(stats: dict | None) -> bool:
|
||||
"""`mutmut results` omits killed mutants, so its silence is equally consistent with a
|
||||
perfect run and with a run that never started. Only the stats file can tell them apart,
|
||||
and only when it agrees that nothing survived and every mutant reached the tests."""
|
||||
if not stats or stats.get("killed", 0) <= 0 or stats.get("survived", 0) != 0:
|
||||
return False
|
||||
return not unresolved_counts(stats)
|
||||
|
||||
|
||||
def no_survivors_verdict(results: MutmutResults, stats: dict | None) -> str:
|
||||
if clean_sweep_is_provable(stats):
|
||||
return "**No surviving mutants, and the run killed some, so the test suite caught every mutation.**"
|
||||
if stats and stats.get("survived", 0) > 0:
|
||||
return (
|
||||
f"**mutmut-cicd-stats.json counts {stats['survived']} surviving mutant(s) that "
|
||||
"`mutmut results` did not list, so the two disagree and neither can be trusted. "
|
||||
"This is not a passing score.**"
|
||||
)
|
||||
if stats and unresolved_counts(stats):
|
||||
unresolved = ", ".join(f"{v} {k.replace('_', ' ')}" for k, v in unresolved_counts(stats).items())
|
||||
return (
|
||||
f"**No survivors, but {unresolved}, so those mutants never reached the tests "
|
||||
"and the suite was not shown to catch them. This is not a passing score.**"
|
||||
)
|
||||
if stats:
|
||||
return "**Not one mutant was killed. This is not a passing score.**"
|
||||
return (
|
||||
f"**mutmut-cicd-stats.json is missing and `mutmut results` printed {results.reported} "
|
||||
"verdict(s), none of them a survivor. Since that command never lists killed mutants, a "
|
||||
"clean sweep and a run that mutated nothing look identical from here. This is not a "
|
||||
"passing score.**"
|
||||
)
|
||||
|
||||
|
||||
def render(config: dict, results: MutmutResults, stats: dict | None) -> str:
|
||||
survivors = list(results.survivors)
|
||||
by_function: dict[tuple[str, str], list[tuple[str, str]]] = defaultdict(list)
|
||||
for survivor in survivors:
|
||||
module_path, function_name, mutant_num = parse_mutant_name(survivor)
|
||||
|
|
@ -235,17 +289,8 @@ def render(config: dict, survivors: list[str], stats: dict | None) -> str:
|
|||
out.append("## Summary")
|
||||
out.append("")
|
||||
if stats:
|
||||
total = stats.get("total", 0) or sum(
|
||||
stats.get(k, 0)
|
||||
for k in (
|
||||
"killed",
|
||||
"survived",
|
||||
"no_tests",
|
||||
"skipped",
|
||||
"suspicious",
|
||||
"timeout",
|
||||
"segfault",
|
||||
)
|
||||
total = stats.get("total", 0) or (
|
||||
stats.get("killed", 0) + stats.get("survived", 0) + sum(unresolved_counts(stats).values())
|
||||
)
|
||||
killed = stats.get("killed", 0)
|
||||
survived = stats.get("survived", 0)
|
||||
|
|
@ -254,17 +299,15 @@ def render(config: dict, survivors: list[str], stats: dict | None) -> str:
|
|||
out.append(f"- Killed: **{killed}**")
|
||||
out.append(f"- Survived: **{survived}**")
|
||||
out.append(f"- Mutation score: **{score:.1f}%**")
|
||||
for k in ("no_tests", "skipped", "suspicious", "timeout", "segfault"):
|
||||
v = stats.get(k, 0)
|
||||
if v:
|
||||
out.append(f"- {k.replace('_', ' ').title()}: {v}")
|
||||
for k, v in unresolved_counts(stats).items():
|
||||
out.append(f"- {k.replace('_', ' ').title()}: {v}")
|
||||
else:
|
||||
out.append(f"- Survivors found: **{len(survivors)}**")
|
||||
out.append("- (mutmut-cicd-stats.json not available — full counts unavailable)")
|
||||
out.append("")
|
||||
|
||||
if not survivors:
|
||||
out.append("**No surviving mutants — the test suite caught every mutation.**")
|
||||
out.append(no_survivors_verdict(results, stats))
|
||||
out.append("")
|
||||
return "\n".join(out)
|
||||
|
||||
|
|
@ -407,15 +450,22 @@ def main() -> int:
|
|||
except json.JSONDecodeError as exc:
|
||||
print(f"warning: could not parse {stats_file}: {exc}", file=sys.stderr)
|
||||
|
||||
survivors = get_survivors()
|
||||
report = render(config, survivors, stats)
|
||||
results = get_survivors()
|
||||
report = render(config, results, stats)
|
||||
|
||||
out_path = ROOT / "mutation-report.md"
|
||||
out_path.write_text(report)
|
||||
print(
|
||||
f"Wrote {out_path} ({len(survivors)} survivor"
|
||||
f"{'s' if len(survivors) != 1 else ''}, {len(report)} chars)"
|
||||
f"Wrote {out_path} ({len(results.survivors)} survivor"
|
||||
f"{'s' if len(results.survivors) != 1 else ''}, {len(report)} chars)"
|
||||
)
|
||||
if not results.survivors and not clean_sweep_is_provable(stats):
|
||||
print(
|
||||
"error: nothing was shown to have been killed, so the report cannot say "
|
||||
"anything about the suite",
|
||||
file=sys.stderr,
|
||||
)
|
||||
return 1
|
||||
return 0
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
{
|
||||
"TQ001": {
|
||||
"limit": 750
|
||||
"limit": 746
|
||||
},
|
||||
"TQ002": {
|
||||
"limit": 742
|
||||
|
|
@ -9,7 +9,7 @@
|
|||
"limit": 1078
|
||||
},
|
||||
"TQ004": {
|
||||
"limit": 757
|
||||
"limit": 557
|
||||
},
|
||||
"TQ005": {
|
||||
"limit": 2810
|
||||
|
|
|
|||
|
|
@ -77,13 +77,26 @@ Mark live tests with `@pytest.mark.e2e` (on the class or the module). Pure cover
|
|||
|
||||
The seam is `provider_edge.py`: `start_provider_edge` boots an in-process HTTP server (one shared instance per pytest process, `e2e_config.provider_edge_base` is the accessor) that mounts each supported provider under a path prefix (`EDGE_MOUNTS`: `/openai` -> `https://api.openai.com`, `/anthropic` -> `https://api.anthropic.com`). A test participates by registering its deployment with `api_base=provider_edge_base("openai")` plus the provider's path suffix; `quota_management/spend_tracking/test_provider_edge_spend_e2e.py` is the reference. In live mode the accessor returns None and the deployment defaults to the real provider, so an edge-wired test runs in all three modes unchanged. Non-wired tests hit their providers live in every mode. The edge binds `E2E_PROVIDER_EDGE_BIND_HOST` (default 127.0.0.1) and advertises `E2E_PROVIDER_EDGE_ADVERTISE_HOST` in the api_base it hands out, for proxies running in containers
|
||||
|
||||
A bundle (default `tests/e2e/.fixtures`, override with `E2E_FIXTURE_DIR`) is a directory: `manifest.json` carries the record timestamp, harness git version, and format version, and each test gets a subdirectory holding one JSON file per provider call in call order (`0000-post-openai-v1-chat-completions.json`). Request headers are never stored (provider credentials never touch disk), non-JSON request bodies store a canonicalized sha256 digest instead of the bytes, and responses store status, filtered headers, and the verbatim body base64-encoded, which is part of why bundles are gitignored. `fixture_bundle.py` owns the format. Record serves the proxy the same filtered stored response replay will serve later, so the two modes are byte-identical from the proxy's side of the socket
|
||||
A bundle (default `tests/e2e/.fixtures`, override with `E2E_FIXTURE_DIR`) is a directory: `manifest.json` carries the record timestamp, harness git version, and format version, and each test gets a subdirectory holding one JSON file per provider call in call order (`0000-post-openai-v1-chat-completions.json`). Request headers are never stored (provider credentials never touch disk), non-JSON request bodies store a canonicalized sha256 digest instead of the bytes, `multipart/form-data` bodies store their ordinary fields plus a JSON list of the uploaded parts' `[field, filename, content-type]` triples and a digest of their content, so the per-request random boundary and the envelope never reach the key, and responses store status, filtered headers, and the verbatim body base64-encoded, which is part of why bundles are gitignored. `fixture_bundle.py` owns the format. Record serves the proxy the same filtered stored response replay will serve later, so the two modes are byte-identical from the proxy's side of the socket
|
||||
|
||||
Multipart identity is the fiddly corner, and the rules exist because each one had a collision behind it. A part counts as an upload when it carries a filename or declares its own content type, and everything else is an ordinary field. Field names get a `name[n]` suffix on repeats, with a literal `[` doubled first, so a form that repeats `purpose` never keys the same as one that literally sends `purpose[1]`. A field whose name reads as a credential is stored as `<secret>`, which stays key-preserving because the key is recomputed from the stored request rather than saved alongside it, so the live request carrying the real value still matches its redacted fixture. A field value that is not UTF-8 is stored as a base64 sha256 digest, base64 and not hex because the canonicalizer rewrites any 64-character hex run to `<sha256>` and would fold every binary value onto one key. The uploaded parts contribute a JSON list rather than a `field:filename` string, so a separator inside a filename cannot impersonate a field boundary, and their byte length is stored for a reader's benefit but deliberately left out of the key, since the canonicalizer absorbs timestamp and id drift inside a file that changes its length
|
||||
|
||||
Replay matches calls per test by canonical key: `fixture_canonical.py` canonicalizes the recorded request (volatile headers and credential fields out, unique markers, generated ids, uuids, and timestamps replaced with fixed placeholders, object keys sorted) and the key is the method, edge path, and a content hash, so identity survives re-records and machine changes while any real content drift comes back as an HTTP 599 naming the computed key, the closest recorded key with its file, and a content diff, and never falls through to a live call. Matching is order-independent across distinct keys (concurrent calls may interleave) and FIFO within one key (a retry loop replays its responses in recorded order); a passed test must also consume its whole recording, or teardown fails it naming a leftover key. Either way the fix is always to re-record with `E2E_FIXTURE_MODE=record`. Every rewrite rule lives in `fixture_canonical.py`, so a new volatile header, credential field name, or generated-id shape is one edit there. Record starts fresh every time: it wipes the previous bundle (refusing to wipe a directory that is not a bundle) and never reads it. A replay bundle whose manifest is older than seven days hard-fails at collection time naming the bundle's age, so replay can never certify against fixtures that have drifted more than a week from the live providers
|
||||
|
||||
A replayed response carries the recorded provider response id, and `LiteLLM_SpendLogs.request_id` (the table's primary key) is that id, so a replay against a database that still holds the record run's rows silently dedupes its spend inserts and any spend assertion goes red with zero matching rows and nothing in the proxy log. Run both modes with `E2E_RESET_SPEND_LOGS=1` (plus `DATABASE_URL` in the runner env) so each session truncates the table after itself, or replay against a fresh database, which is the CI shape
|
||||
|
||||
Current limits: streaming chunk fidelity is LIT-5742 (a streamed response records as one buffered body), CI wiring is LIT-5748, Bedrock cannot be mounted (SigV4 signs the Host header, so a rewritten api_base fails signature verification), multipart uploads have per-run random boundaries (the digest changes every run, so they always miss), and deployments baked into the proxy's config file cannot be edge-wired (only `/model/new` registrations can carry the edge api_base)
|
||||
The same id reuse reaches the managed-object tables. A replayed `/v1/files` or `/v1/batches` response carries the recorded provider object id, and `LiteLLM_ManagedObjectTable.model_object_id` is unique, so a unified batch create replayed against a database that still holds the record run's row fails on a Prisma unique-constraint violation, which surfaces as a 500, makes the router retry, and exhausts the recording. Replay the batches suite against a fresh database, or truncate `LiteLLM_ManagedObjectTable` and `LiteLLM_ManagedFileTable` before the run
|
||||
|
||||
Edge-wired today: `quota_management/spend_tracking/test_provider_edge_spend_e2e.py` (the reference), `llm_translation/test_chat_completions_contract_e2e.py`, the OpenAI registrations in `llm_translation/test_embeddings_endpoint_e2e.py`, the Anthropic deployments in `llm_translation/test_messages_e2e.py` except the streaming test, and the OpenAI batch deployment behind `batches/` (`capabilities.openai_batch_params`). The mount base is not the same for both providers: OpenAI deployments register `f"{base}/v1"`, Anthropic deployments register `base` on its own, because litellm's Anthropic handler appends `/v1/messages` to `api_base` itself where the OpenAI handler appends only `/chat/completions`. Recording one suite locally is two runs against a proxy you already have up:
|
||||
|
||||
```bash
|
||||
E2E_FIXTURE_MODE=record E2E_FIXTURE_DIR=/tmp/e2e-fixtures E2E_RESET_SPEND_LOGS=1 uv run pytest tests/e2e/llm_translation/test_chat_completions_contract_e2e.py
|
||||
E2E_FIXTURE_MODE=replay E2E_FIXTURE_DIR=/tmp/e2e-fixtures E2E_RESET_SPEND_LOGS=1 uv run pytest tests/e2e/llm_translation/test_chat_completions_contract_e2e.py
|
||||
```
|
||||
|
||||
Point the proxy at bogus provider credentials for the replay run and it still has to pass: that is the whole proof that nothing left the process. Bundles are never committed. `tests/e2e/.fixtures` is gitignored because a bundle holds verbatim provider response bodies and hard-fails after seven days, and publishing one for CI is LIT-5748
|
||||
|
||||
Current limits: streaming chunk fidelity is LIT-5742 (a streamed response records as one buffered body), CI wiring is LIT-5748, Bedrock cannot be mounted (SigV4 signs the Host header, so a rewritten api_base fails signature verification), deployments baked into the proxy's config file cannot be edge-wired (only `/model/new` registrations can carry the edge api_base), and a file upload routed by `custom_llm_provider` through the proxy's `files_settings` block never passes a deployment at all, so the batches `model_param` and `provider_fallback` scenarios keep uploading live in every mode
|
||||
|
||||
## Typing
|
||||
|
||||
|
|
|
|||
|
|
@ -57,13 +57,15 @@ Some suites need extra services the bare proxy does not start. The `logging/` OT
|
|||
Record/replay scopes to the proxy's provider-bound traffic only. In `E2E_FIXTURE_MODE=record` the harness boots a local provider-edge server, edge-wired tests register their deployments with an `api_base` pointing at it, and every provider call the proxy makes is forwarded verbatim and written to a fixture bundle (default `tests/e2e/.fixtures`, override with `E2E_FIXTURE_DIR`). `E2E_FIXTURE_MODE=replay` runs the same tests against the same live proxy and database, but the edge answers the proxy's provider calls from the bundle instead of the provider, so the run makes zero provider calls and spends nothing while key auth, routing, cost calculation, and spend-log writes all still execute for real. Unset (or `live`) behaves exactly as before the knob existed. Both record and replay need the proxy up; only the provider is taken out of the loop
|
||||
|
||||
```bash
|
||||
E2E_FIXTURE_MODE=record uv run pytest tests/e2e/quota_management/spend_tracking/test_provider_edge_spend_e2e.py -v
|
||||
E2E_FIXTURE_MODE=replay uv run pytest tests/e2e/quota_management/spend_tracking/test_provider_edge_spend_e2e.py -v
|
||||
E2E_FIXTURE_MODE=record E2E_FIXTURE_DIR=/tmp/e2e-fixtures uv run pytest tests/e2e/quota_management/spend_tracking/test_provider_edge_spend_e2e.py -v
|
||||
E2E_FIXTURE_MODE=replay E2E_FIXTURE_DIR=/tmp/e2e-fixtures uv run pytest tests/e2e/quota_management/spend_tracking/test_provider_edge_spend_e2e.py -v
|
||||
```
|
||||
|
||||
Bundles stay local. `tests/e2e/.fixtures` is gitignored because a bundle holds verbatim provider response bodies and expires seven days after it was recorded, so record the suite you want before you replay it and never commit the result; publishing bundles for CI is LIT-5748
|
||||
|
||||
One sharp edge: a replayed response reuses the recorded provider response id, and that id is the primary key of `LiteLLM_SpendLogs`, so replaying against a database that still holds the record run's rows silently dedupes the spend writes and a spend assertion fails with zero rows. Run both commands above with `E2E_RESET_SPEND_LOGS=1` (and `DATABASE_URL` set in the pytest env) so each session truncates the spend log table after itself, or point replay at a fresh database
|
||||
|
||||
Replay answers any provider call that drifted from the recording with an HTTP 599 whose body names the computed and closest recorded keys, so the test fails loudly instead of silently going live, and a bundle older than seven days fails at collection time naming its age; either way the fix is to re-record. Only tests that register edge-wired deployments participate: everything else hits its provider live in every mode, so record exactly the suite you replay. If the proxy runs in a container, set `E2E_PROVIDER_EDGE_ADVERTISE_HOST` (e.g. `host.docker.internal`) so the api_base the proxy stores can reach the edge on the pytest host, and `E2E_PROVIDER_EDGE_BIND_HOST=0.0.0.0` so the edge accepts it. See `CLAUDE.md` in this directory for the bundle format, the edge design, and the current limits (streaming, Bedrock, multipart)
|
||||
Replay answers any provider call that drifted from the recording with an HTTP 599 whose body names the computed and closest recorded keys, so the test fails loudly instead of silently going live, and a bundle older than seven days fails at collection time naming its age; either way the fix is to re-record. Only tests that register edge-wired deployments participate: everything else hits its provider live in every mode, so record exactly the suite you replay. If the proxy runs in a container, set `E2E_PROVIDER_EDGE_ADVERTISE_HOST` (e.g. `host.docker.internal`) so the api_base the proxy stores can reach the edge on the pytest host, and `E2E_PROVIDER_EDGE_BIND_HOST=0.0.0.0` so the edge accepts it. The suites wired to the edge today are `quota_management/spend_tracking/test_provider_edge_spend_e2e.py`, `llm_translation/test_chat_completions_contract_e2e.py`, the OpenAI registrations in `llm_translation/test_embeddings_endpoint_e2e.py`, the non-streaming Anthropic tests in `llm_translation/test_messages_e2e.py`, and the OpenAI batch deployment behind `batches/`. See `CLAUDE.md` in this directory for the bundle format, the edge design, and the current limits (streaming, Bedrock)
|
||||
|
||||
Tests marked `@pytest.mark.e2e` hard-fail when no proxy answers `/health/liveliness`, so a run that goes red with `No live proxy` at setup means the proxy isn't up; they never skip for a missing proxy, so an absent proxy can't be mistaken for a pass
|
||||
|
||||
|
|
|
|||
|
|
@ -5,9 +5,9 @@ from __future__ import annotations
|
|||
import base64
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from typing import Literal
|
||||
from typing import Final, Literal
|
||||
|
||||
from e2e_config import unique_marker
|
||||
from e2e_config import provider_edge_base, unique_marker
|
||||
from models import LiteLLMParamsBody
|
||||
|
||||
_BATCH_RUN = unique_marker()
|
||||
|
|
@ -17,6 +17,21 @@ def batch_model_name(base: str) -> str:
|
|||
return f"{base}-{_BATCH_RUN}"
|
||||
|
||||
|
||||
OPENAI_BATCH_BACKEND: Final = "gpt-4o-mini"
|
||||
|
||||
|
||||
def openai_batch_params() -> LiteLLMParamsBody:
|
||||
"""The OpenAI batch deployment, wired through the record/replay edge when a fixture
|
||||
mode is active and straight at OpenAI otherwise (LIT-5974). Azure, Vertex, and
|
||||
Bedrock stay live: none of them has an edge mount."""
|
||||
base = provider_edge_base("openai")
|
||||
return LiteLLMParamsBody(
|
||||
model=f"openai/{OPENAI_BATCH_BACKEND}",
|
||||
api_key="os.environ/OPENAI_API_KEY",
|
||||
api_base=None if base is None else f"{base}/v1",
|
||||
)
|
||||
|
||||
|
||||
def _env_ref(*names: str) -> str:
|
||||
for name in names:
|
||||
value = os.environ.get(name)
|
||||
|
|
@ -47,10 +62,7 @@ class Provider:
|
|||
def litellm_params(self) -> LiteLLMParamsBody:
|
||||
match self.name:
|
||||
case "openai":
|
||||
return LiteLLMParamsBody(
|
||||
model="openai/gpt-4o-mini",
|
||||
api_key="os.environ/OPENAI_API_KEY",
|
||||
)
|
||||
return openai_batch_params()
|
||||
case "azure":
|
||||
return LiteLLMParamsBody(
|
||||
model="azure/gpt-5.4-mini-batch",
|
||||
|
|
@ -107,7 +119,11 @@ class Capability:
|
|||
|
||||
PROVIDERS: tuple[Provider, ...] = (
|
||||
Provider(
|
||||
"openai", batch_model_name("openai-batch"), "gpt-4o-mini", can_cancel=True, can_list=True
|
||||
"openai",
|
||||
batch_model_name("openai-batch"),
|
||||
OPENAI_BATCH_BACKEND,
|
||||
can_cancel=True,
|
||||
can_list=True,
|
||||
),
|
||||
Provider(
|
||||
"azure",
|
||||
|
|
|
|||
|
|
@ -42,6 +42,7 @@ from capabilities import (
|
|||
BATCH_ID_SHAPE,
|
||||
CAPABILITIES,
|
||||
FILE_ID_SHAPE,
|
||||
OPENAI_BATCH_BACKEND,
|
||||
OPENAI_BATCH_MODEL,
|
||||
PROVIDERS,
|
||||
Capability,
|
||||
|
|
@ -51,6 +52,7 @@ from capabilities import (
|
|||
decoded_model_from_id,
|
||||
is_managed_id,
|
||||
matches_id_shape,
|
||||
openai_batch_params,
|
||||
raw_id_matches_provider,
|
||||
)
|
||||
from e2e_http import (
|
||||
|
|
@ -479,8 +481,6 @@ def test_rate_limited_batch_create_leaves_no_unattributed_spend_row(
|
|||
)
|
||||
|
||||
|
||||
OPENAI_FILE_CONTENT_BACKEND = "gpt-4o-mini"
|
||||
|
||||
FILE_CONTENT_CELLS = {
|
||||
"azure": "llm.files.azure_openai.content.nonstream.works",
|
||||
"vertex_ai": "llm.files.vertex.content.nonstream.works",
|
||||
|
|
@ -506,17 +506,11 @@ class TestBatchFileContent:
|
|||
self, client: BatchClient, resources: ResourceManager
|
||||
) -> None:
|
||||
proxy_name = f"e2e-file-content-{unique_marker()}"
|
||||
model_id = client.create_model(
|
||||
proxy_name,
|
||||
LiteLLMParamsBody(
|
||||
model=f"openai/{OPENAI_FILE_CONTENT_BACKEND}",
|
||||
api_key="os.environ/OPENAI_API_KEY",
|
||||
),
|
||||
)
|
||||
model_id = client.create_model(proxy_name, openai_batch_params())
|
||||
resources.defer(lambda: client.delete_model(model_id))
|
||||
key = resources.key()
|
||||
|
||||
payload = render_jsonl(OPENAI_FILE_CONTENT_BACKEND)
|
||||
payload = render_jsonl(OPENAI_BATCH_BACKEND)
|
||||
file = unwrap(
|
||||
client.upload_file(
|
||||
content=payload,
|
||||
|
|
|
|||
|
|
@ -64,6 +64,7 @@
|
|||
- {id: llm.responses.openai.basic.stream.works, module: llm, tier: P0, subject_endpoint: responses, route: openai, capability: basic, streaming: stream, assertions: [works], source: "response_api_endpoints/endpoints.py:26", rationale: "Streaming via /v1/responses"}
|
||||
- {id: llm.responses.openai.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: responses, route: openai, capability: basic, streaming: nonstream, assertions: [works, cost_logged], source: "response_api_endpoints/endpoints.py:26", rationale: "Cost logged on responses"}
|
||||
- {id: llm.responses.openai.passthrough.stream.cost_logged, module: llm, tier: P0, subject_endpoint: responses, route: openai, capability: basic, streaming: stream, assertions: [cost_logged], source: "test_passthrough_e2e.py", rationale: "A streamed POST /openai_passthrough/v1/responses is costed and keyed by the provider response id; it used to log a zero-cost row under a random id (GitHub issue #36523)"}
|
||||
- {id: llm.responses.openai.passthrough_websocket.stream.works, module: llm, tier: P1, subject_endpoint: responses, route: openai, capability: basic, streaming: stream, assertions: [works], fail_before_fix: proven, source: "test_passthrough_e2e.py", rationale: "A websocket upgrade on /openai/v1/responses is accepted, so a responses.connect client reaches OpenAI through the same prefix its HTTP traffic uses; the prefix carried no websocket route and refused the upgrade with a 403 (GitHub issue #36088)"}
|
||||
- {id: llm.responses.openai.tool_use.nonstream.works, module: llm, tier: P0, subject_endpoint: responses, route: openai, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Tool calls via Responses API"}
|
||||
- {id: llm.responses.openai.vision.nonstream.works, module: llm, tier: P0, subject_endpoint: responses, route: openai, capability: vision, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Vision via Responses API"}
|
||||
- {id: llm.responses.anthropic.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: anthropic, capability: basic, streaming: nonstream, assertions: [works], source: "response_api_endpoints/endpoints.py:26", rationale: "Responses w/ Anthropic translation (smoke)"}
|
||||
|
|
|
|||
|
|
@ -54,6 +54,7 @@
|
|||
- {id: llm.google_native.gemini.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: google_native, route: gemini, capability: basic, streaming: nonstream, assertions: [cost_logged], source: "LIT-4076 / proxy/google_endpoints/endpoints.py", fail_before_fix: proven, rationale: "google-native generateContent must stamp x-litellm-response-cost so SDK traffic reconciles against spend"}
|
||||
- {id: llm.google_native.gemini.basic.stream.works, module: llm, tier: P0, subject_endpoint: google_native, route: gemini, capability: basic, streaming: stream, assertions: [works], source: "PR #28213 / proxy/proxy_server.py async_data_generator", fail_before_fix: proven, rationale: "streamGenerateContent must relay single-prefixed SSE frames with no [DONE] sentinel; doubled data: prefixes and the OpenAI terminator both break the Vertex Java SDK"}
|
||||
- {id: llm.realtime.openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: realtime, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "vendor strategy §9.19 / LIT-4778", rationale: "HTTP /v1/realtime/client_secrets returns an ephemeral credential"}
|
||||
- {id: llm.realtime.openai.passthrough.stream.works, module: llm, tier: P0, subject_endpoint: realtime, route: openai, capability: basic, streaming: stream, assertions: [works], fail_before_fix: proven, source: "test_passthrough_e2e.py", rationale: "A websocket upgrade on /openai_passthrough/v1/realtime is accepted and relayed to OpenAI; only HTTP routes were registered under the prefix, so realtime clients were refused with a 403 before a socket existed (GitHub issue #36088)"}
|
||||
- {id: llm.vector_stores.openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: vector_stores, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "vendor strategy §9.17 / LIT-4778", rationale: "Vector store create/list/retrieve/delete lifecycle"}
|
||||
- {id: llm.vector_stores.openai.input_validation.nonstream.works, module: llm, tier: P1, subject_endpoint: vector_stores, route: openai, capability: input_validation, streaming: nonstream, assertions: [works], source: "vendor strategy §9.17 / LIT-4778", rationale: "Vector store search and invalid id errors"}
|
||||
- {id: llm.bedrock_native.bedrock_converse.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: bedrock_native, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "vendor strategy §9.12 / LIT-4778", rationale: "Bedrock native converse happy path"}
|
||||
|
|
|
|||
|
|
@ -150,6 +150,15 @@ ANOMALY_SPEND_SETTLE_SECONDS = float(
|
|||
)
|
||||
|
||||
|
||||
def ws_base_url() -> str:
|
||||
"""PROXY_BASE_URL with its scheme swapped for the websocket one, so a suite
|
||||
opening a socket points at the same proxy every HTTP suite uses."""
|
||||
for scheme, ws_scheme in (("https://", "wss://"), ("http://", "ws://")):
|
||||
if PROXY_BASE_URL.startswith(scheme):
|
||||
return ws_scheme + PROXY_BASE_URL[len(scheme) :]
|
||||
return PROXY_BASE_URL
|
||||
|
||||
|
||||
def datadog_mcp_url(*, toolsets: str = "core") -> str:
|
||||
"""Regional Datadog remote MCP endpoint for this process's DD_SITE.
|
||||
|
||||
|
|
|
|||
|
|
@ -5,7 +5,9 @@ version + format version) plus one subdirectory per test, holding one JSON file
|
|||
per provider-bound interaction in call order. Bundles older than
|
||||
``MAX_BUNDLE_AGE`` hard-fail replay at collection time (see conftest), so a
|
||||
green replay run can never certify against fixtures that have drifted more than
|
||||
a week from the live providers.
|
||||
a week from the live providers. Bump ``BUNDLE_FORMAT_VERSION`` whenever a change
|
||||
moves recorded keys: a bundle recorded under the old rules then fails naming
|
||||
both versions instead of quietly missing on every call.
|
||||
|
||||
This module owns the format only. The provider-edge server that produces and
|
||||
consumes it lives in provider_edge.py (LIT-5745) and the canonical match keys
|
||||
|
|
@ -28,7 +30,7 @@ from typing import Final
|
|||
|
||||
from pydantic import BaseModel, JsonValue
|
||||
|
||||
BUNDLE_FORMAT_VERSION: Final = 2
|
||||
BUNDLE_FORMAT_VERSION: Final = 3
|
||||
MAX_BUNDLE_AGE: Final = timedelta(days=7)
|
||||
MANIFEST_FILENAME: Final = "manifest.json"
|
||||
|
||||
|
|
@ -47,7 +49,14 @@ class RecordedRequest(BaseModel):
|
|||
over ``method``, ``path`` (the edge path including the provider mount,
|
||||
query string excluded), and the canonicalized headers, params, body, form,
|
||||
and file identity. Non-JSON bodies store a canonicalized content digest
|
||||
instead of the bytes."""
|
||||
instead of the bytes.
|
||||
|
||||
``file_name`` is a JSON list of the uploaded parts' ``[field, filename,
|
||||
content-type]`` triples rather than a flat label, so a separator inside a
|
||||
filename cannot impersonate a field boundary. ``file_bytes`` is recorded for
|
||||
a reader's benefit and stays out of the key: the canonicalizer absorbs
|
||||
timestamp and id drift inside an uploaded file, and that drift moves the
|
||||
byte count."""
|
||||
|
||||
method: str
|
||||
path: str
|
||||
|
|
|
|||
|
|
@ -129,7 +129,6 @@ def canonicalize(request: RecordedRequest) -> CanonicalRequest:
|
|||
else {
|
||||
"name": None if request.file_name is None else canonical_string(request.file_name),
|
||||
"sha256": request.file_sha256,
|
||||
"bytes": request.file_bytes,
|
||||
}
|
||||
)
|
||||
content: Final[dict[str, JsonValue]] = {
|
||||
|
|
|
|||
|
|
@ -11,9 +11,13 @@ native request models are co-located here because only this suite uses them.
|
|||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from urllib.parse import urlencode
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
from websockets.exceptions import InvalidStatus
|
||||
from websockets.sync.client import connect
|
||||
|
||||
from e2e_config import ws_base_url
|
||||
from proxy_client import ProxyClient
|
||||
from e2e_http import FileUploadForm, Headers, NoBody, Result, StreamingResponse
|
||||
from models import ChatMessage
|
||||
|
|
@ -175,6 +179,26 @@ class OpenAIEmbeddingBody(BaseModel):
|
|||
input: str
|
||||
|
||||
|
||||
class WebsocketEnvelope(BaseModel):
|
||||
"""The one field every provider event carries, so the first frame off a
|
||||
passthrough socket identifies itself without the suite parsing raw dicts."""
|
||||
|
||||
type: str
|
||||
|
||||
|
||||
class WebsocketHandshake(BaseModel):
|
||||
"""What the proxy did with a websocket upgrade on a passthrough prefix.
|
||||
|
||||
`rejected_status` is the HTTP status of a refused upgrade: a prefix carrying no
|
||||
websocket route answers 403, before any socket exists. `first_event_type` is the
|
||||
type of the first frame an accepted socket delivered, which is None when the
|
||||
provider waits for the client to speak first.
|
||||
"""
|
||||
|
||||
rejected_status: int | None = None
|
||||
first_event_type: str | None = None
|
||||
|
||||
|
||||
class PassthroughBatchList(BaseModel):
|
||||
"""OpenAI's own batch page, relayed verbatim. `object` is required so a body
|
||||
that is not an OpenAI list fails validation instead of passing vacuously."""
|
||||
|
|
@ -339,5 +363,39 @@ class PassthroughClient:
|
|||
),
|
||||
)
|
||||
|
||||
# ---- OpenAI websocket passthrough ----------------------------------
|
||||
#
|
||||
# The same prefixes over an upgrade instead of a POST, for the provider APIs
|
||||
# that only speak websocket (realtime, responses.connect).
|
||||
|
||||
def openai_passthrough_websocket(
|
||||
self,
|
||||
key: str,
|
||||
path: str,
|
||||
*,
|
||||
model: str | None = None,
|
||||
open_timeout: float = 30.0,
|
||||
first_event_timeout: float = 30.0,
|
||||
) -> WebsocketHandshake:
|
||||
query = f"?{urlencode({'model': model})}" if model is not None else ""
|
||||
try:
|
||||
connection = connect(
|
||||
f"{ws_base_url()}{path}{query}",
|
||||
additional_headers={"Authorization": f"Bearer {key}"},
|
||||
open_timeout=open_timeout,
|
||||
)
|
||||
except InvalidStatus as rejected:
|
||||
return WebsocketHandshake(rejected_status=rejected.response.status_code)
|
||||
with connection:
|
||||
try:
|
||||
frame = connection.recv(timeout=first_event_timeout)
|
||||
except TimeoutError:
|
||||
return WebsocketHandshake()
|
||||
text = frame.decode("utf-8") if isinstance(frame, bytes) else frame
|
||||
return WebsocketHandshake(
|
||||
first_event_type=WebsocketEnvelope.model_validate_json(text).type
|
||||
)
|
||||
|
||||
|
||||
def build_client(proxy: ProxyClient) -> PassthroughClient:
|
||||
return PassthroughClient(proxy=proxy)
|
||||
|
|
|
|||
|
|
@ -21,20 +21,13 @@ from pydantic import BaseModel, ConfigDict
|
|||
from websockets.sync.client import connect
|
||||
from websockets.sync.connection import Connection
|
||||
|
||||
from e2e_config import PROXY_BASE_URL, unique_marker
|
||||
from e2e_config import unique_marker, ws_base_url
|
||||
from proxy_client import ProxyClient
|
||||
from models import LiteLLMParamsBody
|
||||
|
||||
_M = TypeVar("_M", bound=BaseModel)
|
||||
|
||||
|
||||
def ws_base_url() -> str:
|
||||
for scheme, ws_scheme in (("https://", "wss://"), ("http://", "ws://")):
|
||||
if PROXY_BASE_URL.startswith(scheme):
|
||||
return ws_scheme + PROXY_BASE_URL[len(scheme) :]
|
||||
return PROXY_BASE_URL
|
||||
|
||||
|
||||
def realtime_ws_url(model: str) -> str:
|
||||
return f"{ws_base_url()}/v1/realtime?{urlencode({'model': model})}"
|
||||
|
||||
|
|
|
|||
|
|
@ -27,10 +27,10 @@ from pathlib import Path
|
|||
|
||||
import pytest
|
||||
|
||||
from e2e_config import ws_base_url
|
||||
from realtime_client import (
|
||||
PROVIDERS,
|
||||
RealtimeProvider,
|
||||
ws_base_url,
|
||||
realtime_model,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -25,10 +25,10 @@ import asyncio
|
|||
|
||||
import pytest
|
||||
|
||||
from e2e_config import ws_base_url
|
||||
from realtime_client import (
|
||||
PROVIDERS,
|
||||
RealtimeProvider,
|
||||
ws_base_url,
|
||||
realtime_model,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ Exercises the gateway against a live OpenAI deployment using customer request sh
|
|||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from e2e_config import unique_marker
|
||||
from e2e_config import provider_edge_base, unique_marker
|
||||
from e2e_http import StreamingResponse, assert_client_error, require_successful_call, unwrap
|
||||
from lifecycle import ResourceManager
|
||||
from models import ChatBody, ChatMessage, ChatResponse, LiteLLMParamsBody
|
||||
|
|
@ -38,10 +38,15 @@ class ChatErrorEnvelope(BaseModel):
|
|||
|
||||
|
||||
def _register_chat_model(proxy: ProxyClient, resources: ResourceManager) -> tuple[str, str]:
|
||||
base = provider_edge_base("openai")
|
||||
model = f"e2e-chat-sec-{unique_marker()}"
|
||||
model_id = proxy.create_model(
|
||||
model,
|
||||
LiteLLMParamsBody(model=OPENAI_BACKEND, api_key="os.environ/OPENAI_API_KEY"),
|
||||
LiteLLMParamsBody(
|
||||
model=OPENAI_BACKEND,
|
||||
api_key="os.environ/OPENAI_API_KEY",
|
||||
api_base=None if base is None else f"{base}/v1",
|
||||
),
|
||||
)
|
||||
resources.defer(lambda: proxy.delete_model(model_id))
|
||||
return model, resources.key()
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ covered by tests/e2e/quota_management/spend_tracking/.
|
|||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from e2e_config import unique_marker
|
||||
from e2e_config import provider_edge_base, unique_marker
|
||||
from e2e_http import (
|
||||
assert_client_error,
|
||||
require_successful_call,
|
||||
|
|
@ -27,6 +27,18 @@ class _OptionalEmbeddingsBody(BaseModel):
|
|||
input: str | list[str] | None = None
|
||||
|
||||
|
||||
def _openai_embeddings_params() -> LiteLLMParamsBody:
|
||||
"""The OpenAI embeddings deployment, wired through the record/replay edge when a
|
||||
fixture mode is active and straight at OpenAI otherwise (LIT-5974). Bedrock and
|
||||
Vertex stay live: SigV4 signs the Host header, and neither has an edge mount."""
|
||||
base = provider_edge_base("openai")
|
||||
return LiteLLMParamsBody(
|
||||
model="openai/text-embedding-3-small",
|
||||
api_key="os.environ/OPENAI_API_KEY",
|
||||
api_base=None if base is None else f"{base}/v1",
|
||||
)
|
||||
|
||||
|
||||
class TestEmbeddingsEndpoint:
|
||||
@pytest.mark.covers("llm.embeddings.openai.basic.nonstream.works")
|
||||
def test_embeddings_returns_vector(
|
||||
|
|
@ -35,9 +47,7 @@ class TestEmbeddingsEndpoint:
|
|||
model = f"e2e-embeddings-{unique_marker()}"
|
||||
model_id = endpoints_client.create_model(
|
||||
model,
|
||||
LiteLLMParamsBody(
|
||||
model="openai/text-embedding-3-small", api_key="os.environ/OPENAI_API_KEY"
|
||||
),
|
||||
_openai_embeddings_params(),
|
||||
)
|
||||
resources.defer(lambda: endpoints_client.delete_model(model_id))
|
||||
key = resources.key()
|
||||
|
|
@ -106,9 +116,7 @@ class TestEmbeddingsEndpoint:
|
|||
model = f"e2e-embeddings-array-{unique_marker()}"
|
||||
model_id = endpoints_client.create_model(
|
||||
model,
|
||||
LiteLLMParamsBody(
|
||||
model="openai/text-embedding-3-small", api_key="os.environ/OPENAI_API_KEY"
|
||||
),
|
||||
_openai_embeddings_params(),
|
||||
)
|
||||
resources.defer(lambda: endpoints_client.delete_model(model_id))
|
||||
key = resources.key()
|
||||
|
|
@ -140,9 +148,7 @@ class TestEmbeddingsEndpoint:
|
|||
model = f"e2e-embeddings-missin-{unique_marker()}"
|
||||
model_id = endpoints_client.create_model(
|
||||
model,
|
||||
LiteLLMParamsBody(
|
||||
model="openai/text-embedding-3-small", api_key="os.environ/OPENAI_API_KEY"
|
||||
),
|
||||
_openai_embeddings_params(),
|
||||
)
|
||||
resources.defer(lambda: endpoints_client.delete_model(model_id))
|
||||
key = resources.key()
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ litellm-regression-tests/tests/test_inference_endpoints.py.
|
|||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from e2e_config import unique_marker
|
||||
from e2e_config import provider_edge_base, unique_marker
|
||||
from e2e_http import assert_client_error, require_successful_call, unwrap
|
||||
from endpoints_client import EndpointsClient, MessagesResult
|
||||
from lifecycle import ResourceManager
|
||||
|
|
@ -50,16 +50,27 @@ def _approx_equal(actual: float, expected: float) -> bool:
|
|||
return abs(actual - expected) <= max(1e-9, abs(expected) * 1e-2)
|
||||
|
||||
|
||||
def _anthropic_params() -> LiteLLMParamsBody:
|
||||
"""The Anthropic deployment, wired through the record/replay edge when a fixture
|
||||
mode is active (LIT-5974). The mount base carries no ``/v1``: litellm's Anthropic
|
||||
handler appends ``/v1/messages`` to ``api_base`` itself, where the OpenAI handler
|
||||
appends only ``/chat/completions``."""
|
||||
base = provider_edge_base("anthropic")
|
||||
return LiteLLMParamsBody(
|
||||
model=ANTHROPIC_BACKEND, api_key="os.environ/ANTHROPIC_API_KEY", api_base=base
|
||||
)
|
||||
|
||||
|
||||
class TestAnthropicMessages:
|
||||
def _register(
|
||||
self, endpoints_client: EndpointsClient, resources: ResourceManager
|
||||
self,
|
||||
endpoints_client: EndpointsClient,
|
||||
resources: ResourceManager,
|
||||
params: LiteLLMParamsBody | None = None,
|
||||
) -> tuple[str, str]:
|
||||
model = f"e2e-messages-{unique_marker()}"
|
||||
model_id = endpoints_client.create_model(
|
||||
model,
|
||||
LiteLLMParamsBody(
|
||||
model=ANTHROPIC_BACKEND, api_key="os.environ/ANTHROPIC_API_KEY"
|
||||
),
|
||||
model, _anthropic_params() if params is None else params
|
||||
)
|
||||
resources.defer(lambda: endpoints_client.delete_model(model_id))
|
||||
return model, resources.key()
|
||||
|
|
@ -81,12 +92,7 @@ class TestAnthropicMessages:
|
|||
self, endpoints_client: EndpointsClient, resources: ResourceManager
|
||||
) -> None:
|
||||
model = f"e2e-messages-cost-{unique_marker()}"
|
||||
model_id = endpoints_client.create_model(
|
||||
model,
|
||||
LiteLLMParamsBody(
|
||||
model=ANTHROPIC_BACKEND, api_key="os.environ/ANTHROPIC_API_KEY"
|
||||
),
|
||||
)
|
||||
model_id = endpoints_client.create_model(model, _anthropic_params())
|
||||
resources.defer(lambda: endpoints_client.delete_model(model_id))
|
||||
key = resources.key()
|
||||
|
||||
|
|
@ -131,7 +137,13 @@ class TestAnthropicMessages:
|
|||
def test_messages_streams_completion(
|
||||
self, endpoints_client: EndpointsClient, resources: ResourceManager
|
||||
) -> None:
|
||||
model, key = self._register(endpoints_client, resources)
|
||||
"""Stays on a live Anthropic deployment in every mode: the edge buffers a
|
||||
streamed response into one body, so chunk fidelity waits on LIT-5742."""
|
||||
model, key = self._register(
|
||||
endpoints_client,
|
||||
resources,
|
||||
LiteLLMParamsBody(model=ANTHROPIC_BACKEND, api_key="os.environ/ANTHROPIC_API_KEY"),
|
||||
)
|
||||
|
||||
result = endpoints_client.proxy.messages_stream(
|
||||
key,
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ from passthrough_client import (
|
|||
)
|
||||
|
||||
EMBEDDING_MODEL = "text-embedding-3-small"
|
||||
REALTIME_MODEL = "gpt-realtime-2"
|
||||
|
||||
pytestmark = pytest.mark.e2e
|
||||
|
||||
|
|
@ -339,3 +340,52 @@ class TestOpenAIPassthroughSpend:
|
|||
f"the embeddings row logged no prompt tokens, so whatever cost it carries "
|
||||
f"was not computed from the real usage: {row}"
|
||||
)
|
||||
|
||||
|
||||
class TestOpenAIPassthroughWebsocket:
|
||||
"""The OpenAI passthrough prefixes must answer a websocket upgrade, not only a POST.
|
||||
|
||||
The customer points realtime and responses.connect clients at the same prefixes
|
||||
their HTTP traffic already uses. Only HTTP routes were registered under those
|
||||
prefixes, so every upgrade was refused before a socket existed and those clients
|
||||
could not reach the gateway at all. A refused upgrade is an HTTP response, not a
|
||||
close frame, which is why these assert on the handshake rather than a close code.
|
||||
"""
|
||||
|
||||
@pytest.mark.covers("llm.realtime.openai.passthrough.stream.works")
|
||||
def test_realtime_upgrade_reaches_openai_through_the_passthrough_prefix(
|
||||
self, client: PassthroughClient, scoped_key: str
|
||||
) -> None:
|
||||
"""Pins GitHub issue #36088: /openai_passthrough/v1/realtime accepts the
|
||||
upgrade and relays OpenAI's own session, instead of rejecting it with a 403."""
|
||||
handshake = client.openai_passthrough_websocket(
|
||||
scoped_key, "/openai_passthrough/v1/realtime", model=REALTIME_MODEL
|
||||
)
|
||||
|
||||
assert handshake.rejected_status is None, (
|
||||
f"/openai_passthrough/v1/realtime refused the websocket upgrade with HTTP "
|
||||
f"{handshake.rejected_status}, so a realtime client cannot connect through "
|
||||
"the gateway at all"
|
||||
)
|
||||
assert handshake.first_event_type == "session.created", (
|
||||
"the accepted socket never carried OpenAI's opening session event, so the "
|
||||
f"upgrade was not relayed upstream; the first frame was "
|
||||
f"{handshake.first_event_type}"
|
||||
)
|
||||
|
||||
@pytest.mark.covers("llm.responses.openai.passthrough_websocket.stream.works")
|
||||
def test_responses_upgrade_is_accepted_on_the_openai_prefix(
|
||||
self, client: PassthroughClient, scoped_key: str
|
||||
) -> None:
|
||||
"""Pins GitHub issue #36088 on the second prefix: /openai/v1/responses upgrades
|
||||
as well. A responses.connect socket waits for the client to speak first, so the
|
||||
accepted handshake is the whole signal here."""
|
||||
handshake = client.openai_passthrough_websocket(
|
||||
scoped_key, "/openai/v1/responses", first_event_timeout=2.0
|
||||
)
|
||||
|
||||
assert handshake.rejected_status is None, (
|
||||
f"/openai/v1/responses refused the websocket upgrade with HTTP "
|
||||
f"{handshake.rejected_status}; the prefix relays this route over HTTP but "
|
||||
"drops a responses.connect client before the socket opens"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -20,10 +20,9 @@ headers must never touch disk. An unmatched replay call returns HTTP
|
|||
proxy relays as a provider error the failing test surfaces.
|
||||
|
||||
v1 limits: only the mounts in ``EDGE_MOUNTS`` (SigV4 providers like Bedrock
|
||||
sign the Host header, so a forwarding edge breaks their signatures), JSON and
|
||||
opaque single-part bodies (multipart boundaries are random per request),
|
||||
streaming fidelity is LIT-5742, and CI wiring is LIT-5748. Suites that do not
|
||||
wire the edge keep hitting providers live in every mode.
|
||||
sign the Host header, so a forwarding edge breaks their signatures), streaming
|
||||
fidelity is LIT-5742, and CI wiring is LIT-5748. Suites that do not wire the
|
||||
edge keep hitting providers live in every mode.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
|
@ -32,6 +31,7 @@ import base64
|
|||
import difflib
|
||||
import functools
|
||||
import hashlib
|
||||
import re
|
||||
import threading
|
||||
from collections import deque
|
||||
from collections.abc import Mapping
|
||||
|
|
@ -59,7 +59,13 @@ from fixture_bundle import (
|
|||
prepare_bundle,
|
||||
slug_for_test,
|
||||
)
|
||||
from fixture_canonical import CanonicalRequest, canonical_string, canonicalize
|
||||
from fixture_canonical import (
|
||||
SECRET_PLACEHOLDER,
|
||||
CanonicalRequest,
|
||||
canonical_string,
|
||||
canonicalize,
|
||||
is_secret_field,
|
||||
)
|
||||
from fixture_mode import (
|
||||
FIXTURE_MODES,
|
||||
InvalidFixtureMode,
|
||||
|
|
@ -103,26 +109,244 @@ _RESPONSE_DROPPED_HEADERS: Final[frozenset[str]] = _HOP_BY_HOP_HEADERS | {
|
|||
_JSON: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
|
||||
|
||||
|
||||
def _edge_request(method: str, path: str, query: str, body: bytes | None) -> RecordedRequest:
|
||||
"""The identity replay matches on: the edge path (mount included), the query
|
||||
as params, and the body as parsed JSON, or as a canonicalized content digest
|
||||
when it is not JSON so opaque uploads still match across runs."""
|
||||
params: Final = dict(parse_qsl(query, keep_blank_values=True))
|
||||
if not body:
|
||||
return RecordedRequest(method=method.lower(), path=path, headers={}, params=params)
|
||||
decoded: Final = body.decode("utf-8", errors="replace")
|
||||
_BOUNDARY_PATTERN: Final = re.compile(
|
||||
r'(?:^|;)\s*boundary\s*=\s*(?:"([^"]*)"|([^;,\s]+))', re.IGNORECASE
|
||||
)
|
||||
_DISPOSITION_NAME_PATTERN: Final = re.compile(r'(?:^|;)\s*name="([^"]*)"', re.IGNORECASE)
|
||||
_DISPOSITION_FILENAME_PATTERN: Final = re.compile(
|
||||
r'(?:^|;)\s*filename="([^"]*)"', re.IGNORECASE
|
||||
)
|
||||
_UNPARSED_MULTIPART: Final = "<unparsed-multipart>"
|
||||
_BOUNDARY_PLACEHOLDER: Final = b"--<boundary>"
|
||||
_BINARY_FIELD_PREFIX: Final = "<binary:sha256:"
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _MultipartPart:
|
||||
field_name: str
|
||||
filename: str | None
|
||||
content: bytes
|
||||
content_type: str = ""
|
||||
|
||||
|
||||
def _header_value(headers: Mapping[str, str], name: str) -> str:
|
||||
wanted: Final = name.lower()
|
||||
return next((value for key, value in headers.items() if key.lower() == wanted), "")
|
||||
|
||||
|
||||
def _multipart_boundary(content_type: str) -> str | None:
|
||||
"""The declared boundary, or None when the envelope is not multipart or names no
|
||||
usable boundary. ``boundary`` is matched only as a parameter in its own right, so a
|
||||
longer name ending in it (``myboundary=``) is not mistaken for one, and an empty
|
||||
boundary is refused rather than splitting the body on a bare ``--``."""
|
||||
if "multipart/form-data" not in content_type.lower():
|
||||
return None
|
||||
match: Final = _BOUNDARY_PATTERN.search(content_type)
|
||||
if match is None:
|
||||
return None
|
||||
quoted, bare = match.group(1), match.group(2)
|
||||
return (quoted if quoted is not None else bare) or None
|
||||
|
||||
|
||||
def _part_headers(head: bytes) -> dict[str, str]:
|
||||
return {
|
||||
name.strip().lower(): value.strip()
|
||||
for line in head.decode("utf-8", errors="replace").split("\r\n")
|
||||
for name, separator, value in [line.partition(":")]
|
||||
if separator
|
||||
}
|
||||
|
||||
|
||||
def _parse_multipart_part(segment: bytes) -> _MultipartPart | None:
|
||||
head, separator, content = segment.partition(b"\r\n\r\n")
|
||||
if not separator:
|
||||
return None
|
||||
headers: Final = _part_headers(head)
|
||||
disposition: Final = headers.get("content-disposition", "")
|
||||
name_match: Final = _DISPOSITION_NAME_PATTERN.search(disposition)
|
||||
if name_match is None:
|
||||
return None
|
||||
filename_match: Final = _DISPOSITION_FILENAME_PATTERN.search(disposition)
|
||||
return _MultipartPart(
|
||||
field_name=name_match.group(1),
|
||||
filename=None if filename_match is None else filename_match.group(1),
|
||||
content=content,
|
||||
content_type=headers.get("content-type", ""),
|
||||
)
|
||||
|
||||
|
||||
def _multipart_parts(body: bytes, boundary: str) -> tuple[_MultipartPart, ...] | None:
|
||||
"""The wire body split back into its parts, or None when it does not parse as the
|
||||
declared envelope so the caller can fall back to the opaque content digest."""
|
||||
segments: Final = body.split(b"--" + boundary.encode())
|
||||
if len(segments) < 3 or not segments[-1].startswith(b"--"):
|
||||
return None
|
||||
parsed: Final = tuple(
|
||||
_parse_multipart_part(segment.removeprefix(b"\r\n").removesuffix(b"\r\n"))
|
||||
for segment in segments[1:-1]
|
||||
)
|
||||
if any(part is None for part in parsed):
|
||||
return None
|
||||
return tuple(part for part in parsed if part is not None)
|
||||
|
||||
|
||||
def _content_digest(content: bytes) -> str:
|
||||
"""Text is canonicalized before hashing so a per-run marker inside an uploaded JSONL
|
||||
does not move the key; anything that is not UTF-8 is hashed byte for byte, since a
|
||||
lossy decode collapses every binary payload of one length onto one digest."""
|
||||
try:
|
||||
parsed: Final[JsonValue] = _JSON.validate_json(decoded)
|
||||
except ValueError:
|
||||
return RecordedRequest(
|
||||
method=method.lower(),
|
||||
path=path,
|
||||
headers={},
|
||||
params=params,
|
||||
file_sha256=hashlib.sha256(canonical_string(decoded).encode()).hexdigest(),
|
||||
file_bytes=len(body),
|
||||
text: Final = content.decode("utf-8")
|
||||
except UnicodeDecodeError:
|
||||
return hashlib.sha256(content).hexdigest()
|
||||
return hashlib.sha256(canonical_string(text).encode()).hexdigest()
|
||||
|
||||
|
||||
def _is_file_part(part: _MultipartPart) -> bool:
|
||||
"""Whether a part is an upload rather than an ordinary field. A filename says so
|
||||
outright, and so does a declared content type: clients attach one per part only for
|
||||
a file, and a client that omits the filename (httpx drops the parameter when it is
|
||||
empty) would otherwise have the file's bytes stored inline as a field value and key
|
||||
identically to a plain field of the same name."""
|
||||
return part.filename is not None or bool(part.content_type)
|
||||
|
||||
|
||||
def _field_value(part: _MultipartPart) -> str:
|
||||
"""What a field part contributes to the stored form. A secret-named field never has
|
||||
its value written out, since the bundle is a file on disk and the key redacts that
|
||||
field to the same placeholder either way, so replay still matches. A value that is
|
||||
not UTF-8 is carried as a digest rather than decoded lossily, because a replacing
|
||||
decode collapses every binary value of one length onto one string. That digest is
|
||||
base64 rather than hex, since the canonicalizer rewrites any long hex run to a
|
||||
``<sha256>`` placeholder and would collapse the values right back together."""
|
||||
if is_secret_field(part.field_name):
|
||||
return SECRET_PLACEHOLDER
|
||||
try:
|
||||
return part.content.decode("utf-8")
|
||||
except UnicodeDecodeError:
|
||||
digest: Final = base64.b64encode(hashlib.sha256(part.content).digest()).decode()
|
||||
return f"{_BINARY_FIELD_PREFIX}{digest}>"
|
||||
|
||||
|
||||
def _form_fields(fields: tuple[_MultipartPart, ...]) -> dict[str, str]:
|
||||
"""The ordinary field parts, flattened into the mapping the bundle format stores. A
|
||||
name sent more than once takes an occurrence suffix instead of overwriting the
|
||||
earlier value, so nothing an upload said is dropped from its key. The suffix is
|
||||
escaped so a field literally named ``x[1]`` cannot collide with a second ``x``."""
|
||||
form: dict[str, str] = {}
|
||||
for part in fields:
|
||||
name = part.field_name.replace("[", "[[")
|
||||
occurrence = 1
|
||||
while name in form:
|
||||
name = f"{part.field_name.replace('[', '[[')}[{occurrence}]"
|
||||
occurrence += 1
|
||||
form[name] = _field_value(part)
|
||||
return form
|
||||
|
||||
|
||||
def _file_identity(files: tuple[_MultipartPart, ...]) -> tuple[str | None, str | None, int | None]:
|
||||
"""Name, content digest, and total length for the uploaded file parts.
|
||||
|
||||
The name is a structured list of every part's field name, filename, and declared
|
||||
content type rather than a joined string, so a filename containing the separator
|
||||
cannot be confused for a different split, and two parts that differ only in the type
|
||||
they declare stay apart. It goes through the canonicalizer as one string, which is
|
||||
why per-run markers inside a filename do not move the key in the multi-file case any
|
||||
more than they do in the single-file one.
|
||||
|
||||
The digest covers content only. A lone file keeps its own canonicalized digest;
|
||||
several fold into one ordered digest, so parts arriving in a different order key
|
||||
differently. Total length is recorded for a reader but deliberately kept out of the
|
||||
key: it is the raw byte count, and keying on it would undo exactly the drift the
|
||||
canonicalized digest exists to absorb."""
|
||||
if not files:
|
||||
return None, None, None
|
||||
names: Final = _JSON.dump_json(
|
||||
[[part.field_name, part.filename, part.content_type] for part in files]
|
||||
).decode()
|
||||
total: Final = sum(len(part.content) for part in files)
|
||||
if len(files) == 1:
|
||||
return names, _content_digest(files[0].content), total
|
||||
folded: Final = _JSON.dump_json([_content_digest(part.content) for part in files])
|
||||
return names, hashlib.sha256(folded).hexdigest(), total
|
||||
|
||||
|
||||
def _multipart_request(
|
||||
method: str, path: str, params: dict[str, str], parts: tuple[_MultipartPart, ...]
|
||||
) -> RecordedRequest:
|
||||
"""A multipart upload keyed by what it says rather than by its wire bytes: every
|
||||
ordinary field, plus the identity of the uploaded file. The random per-request
|
||||
boundary is envelope, never content, so it never reaches the digest."""
|
||||
form: Final = _form_fields(tuple(part for part in parts if not _is_file_part(part)))
|
||||
file_name, file_sha256, file_bytes = _file_identity(
|
||||
tuple(part for part in parts if _is_file_part(part))
|
||||
)
|
||||
return RecordedRequest(
|
||||
method=method,
|
||||
path=path,
|
||||
headers={},
|
||||
params=params,
|
||||
form=form,
|
||||
file_name=file_name,
|
||||
file_sha256=file_sha256,
|
||||
file_bytes=file_bytes,
|
||||
)
|
||||
|
||||
|
||||
def _opaque_request(
|
||||
method: str,
|
||||
path: str,
|
||||
params: dict[str, str],
|
||||
body: bytes,
|
||||
digested: bytes,
|
||||
file_name: str | None = None,
|
||||
) -> RecordedRequest:
|
||||
"""A body kept out of the bundle and matched on its digest alone. ``digested`` is
|
||||
what the digest runs over, which is the body itself unless something in it has to be
|
||||
normalized away first."""
|
||||
return RecordedRequest(
|
||||
method=method,
|
||||
path=path,
|
||||
headers={},
|
||||
params=params,
|
||||
file_name=file_name,
|
||||
file_sha256=_content_digest(digested),
|
||||
file_bytes=len(body),
|
||||
)
|
||||
|
||||
|
||||
def edge_request(
|
||||
method: str, path: str, query: str, body: bytes | None, content_type: str = ""
|
||||
) -> RecordedRequest:
|
||||
"""The identity replay matches on: the edge path (mount included), the query as
|
||||
params, and the body as parsed JSON, as parsed multipart fields and file identity
|
||||
when the content type declares an envelope, or as a content digest otherwise so
|
||||
opaque uploads still match across runs. A multipart body that does not parse still
|
||||
has its boundary normalized away, because that boundary is fresh every request and
|
||||
would otherwise guarantee a miss."""
|
||||
params: Final = dict(parse_qsl(query, keep_blank_values=True))
|
||||
lowered_method: Final = method.lower()
|
||||
if not body:
|
||||
return RecordedRequest(method=lowered_method, path=path, headers={}, params=params)
|
||||
boundary: Final = _multipart_boundary(content_type)
|
||||
if boundary is not None:
|
||||
parts = _multipart_parts(body, boundary)
|
||||
if parts is not None:
|
||||
return _multipart_request(lowered_method, path, params, parts)
|
||||
return _opaque_request(
|
||||
lowered_method,
|
||||
path,
|
||||
params,
|
||||
body,
|
||||
body.replace(b"--" + boundary.encode(), _BOUNDARY_PLACEHOLDER),
|
||||
_UNPARSED_MULTIPART,
|
||||
)
|
||||
return RecordedRequest(method=method.lower(), path=path, headers={}, params=params, body=parsed)
|
||||
try:
|
||||
parsed: Final[JsonValue] = _JSON.validate_json(body)
|
||||
except ValueError:
|
||||
return _opaque_request(lowered_method, path, params, body, body)
|
||||
return RecordedRequest(
|
||||
method=lowered_method, path=path, headers={}, params=params, body=parsed
|
||||
)
|
||||
|
||||
|
||||
def _build_pool(recorded: tuple[Interaction, ...]) -> dict[str, deque[Interaction]]:
|
||||
|
|
@ -351,7 +575,9 @@ def handle_edge_request(
|
|||
return _text_reply(
|
||||
404, f"unknown provider mount {mount!r}; known mounts: {', '.join(sorted(mounts))}"
|
||||
)
|
||||
request: Final = _edge_request(method, split.path, split.query, body)
|
||||
request: Final = edge_request(
|
||||
method, split.path, split.query, body, _header_value(headers, "content-type")
|
||||
)
|
||||
match backend:
|
||||
case RecordEdge():
|
||||
return _handle_record(
|
||||
|
|
|
|||
|
|
@ -23,11 +23,13 @@ from concurrent.futures import ThreadPoolExecutor
|
|||
from contextlib import contextmanager
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
from e2e_http import RawResponse, forward
|
||||
from fixture_canonical import canonicalize
|
||||
from fixture_bundle import (
|
||||
BundleRecorder,
|
||||
Interaction,
|
||||
|
|
@ -46,6 +48,7 @@ from provider_edge import (
|
|||
RecordEdge,
|
||||
ReplayEdge,
|
||||
ReplaySource,
|
||||
edge_request,
|
||||
handle_edge_request,
|
||||
provider_edge_api_base,
|
||||
replay_leftover_error,
|
||||
|
|
@ -53,8 +56,10 @@ from provider_edge import (
|
|||
)
|
||||
|
||||
CHAT_PATH = "/openai/v1/chat/completions"
|
||||
UPLOAD_PATH = "/openai/v1/files"
|
||||
REPLAY_MOUNTS = {"openai": "https://replay.invalid"}
|
||||
JSON_OBJECT = TypeAdapter(dict[str, object])
|
||||
BATCH_JSONL = b'{"custom_id":"one"}\n{"custom_id":"two"}\n'
|
||||
|
||||
|
||||
def json_object(body: bytes) -> dict[str, object]:
|
||||
|
|
@ -164,6 +169,46 @@ def chat_body(prompt: str) -> bytes:
|
|||
return json.dumps({"model": "gpt", "messages": [{"role": "user", "content": prompt}]}).encode()
|
||||
|
||||
|
||||
def multipart_body(
|
||||
boundary: str,
|
||||
fields: tuple[tuple[str, str], ...] = (),
|
||||
files: tuple[tuple[str, str, bytes], ...] = (),
|
||||
) -> bytes:
|
||||
"""One multipart/form-data body on the wire, exactly as ``requests`` writes it, with
|
||||
the boundary under the caller's control instead of randomly generated."""
|
||||
parts = [
|
||||
f'--{boundary}\r\nContent-Disposition: form-data; name="{name}"\r\n\r\n'.encode()
|
||||
+ value.encode()
|
||||
for name, value in fields
|
||||
] + [
|
||||
(
|
||||
f'--{boundary}\r\nContent-Disposition: form-data; name="{name}"; '
|
||||
f'filename="{filename}"\r\nContent-Type: application/octet-stream\r\n\r\n'
|
||||
).encode()
|
||||
+ content
|
||||
for name, filename, content in files
|
||||
]
|
||||
return b"\r\n".join(parts) + f"\r\n--{boundary}--\r\n".encode()
|
||||
|
||||
|
||||
def upload_headers(boundary: str) -> dict[str, str]:
|
||||
return {
|
||||
"content-type": f"multipart/form-data; boundary={boundary}",
|
||||
"authorization": "Bearer sk-upload-secret",
|
||||
}
|
||||
|
||||
|
||||
def record_upload(root: Path, body: bytes, boundary: str) -> None:
|
||||
with fake_provider() as provider:
|
||||
with running_edge(record_backend(root), {"openai": provider_url(provider)}) as edge:
|
||||
call_edge(edge, "POST", UPLOAD_PATH, body=body, headers=upload_headers(boundary))
|
||||
|
||||
|
||||
def replay_upload(root: Path, body: bytes, boundary: str) -> RawResponse:
|
||||
with running_edge(ReplayEdge(source=replay_source(root)), REPLAY_MOUNTS) as edge:
|
||||
return call_edge(edge, "POST", UPLOAD_PATH, body=body, headers=upload_headers(boundary))
|
||||
|
||||
|
||||
class TestRecordMode:
|
||||
def test_forwards_to_the_provider_and_writes_one_interaction_file(self, tmp_path: Path) -> None:
|
||||
root = tmp_path / "bundle"
|
||||
|
|
@ -328,6 +373,347 @@ class TestReplayMode:
|
|||
assert replayed.status_code == 200
|
||||
|
||||
|
||||
class TestMultipartIdentity:
|
||||
"""LIT-5974: a multipart upload is keyed by its parsed fields and file identity.
|
||||
``requests`` picks a fresh random boundary per request, so hashing the wire body
|
||||
made every upload miss on replay; parsing the envelope keys the upload on what it
|
||||
actually says, which is stable across runs and still separates real drift."""
|
||||
|
||||
def test_a_fresh_boundary_replays_the_same_upload(self, tmp_path: Path) -> None:
|
||||
root = tmp_path / "bundle"
|
||||
recorded = multipart_body(
|
||||
"d0a1b2c3d4e5f60718293a4b5c6d7e8f",
|
||||
fields=(("purpose", "batch"),),
|
||||
files=(("file", "batch.jsonl", BATCH_JSONL),),
|
||||
)
|
||||
record_upload(root, recorded, "d0a1b2c3d4e5f60718293a4b5c6d7e8f")
|
||||
|
||||
rerun = multipart_body(
|
||||
"ffffeeeeddddccccbbbbaaaa99998888",
|
||||
fields=(("purpose", "batch"),),
|
||||
files=(("file", "batch.jsonl", BATCH_JSONL),),
|
||||
)
|
||||
assert rerun != recorded
|
||||
replayed = replay_upload(root, rerun, "ffffeeeeddddccccbbbbaaaa99998888")
|
||||
assert replayed.status_code == 200, replayed.body[:400]
|
||||
|
||||
def test_the_stored_request_carries_fields_and_file_identity_but_no_secrets(
|
||||
self, tmp_path: Path
|
||||
) -> None:
|
||||
root = tmp_path / "bundle"
|
||||
boundary = "0123456789abcdef0123456789abcdef"
|
||||
record_upload(
|
||||
root,
|
||||
multipart_body(
|
||||
boundary,
|
||||
fields=(("purpose", "batch"),),
|
||||
files=(("file", "batch.jsonl", BATCH_JSONL),),
|
||||
),
|
||||
boundary,
|
||||
)
|
||||
|
||||
raw = this_tests_files(root)[0].read_text(encoding="utf-8")
|
||||
interaction = Interaction.model_validate_json(raw)
|
||||
assert interaction.request.form == {"purpose": "batch"}
|
||||
assert interaction.request.file_name == json.dumps(
|
||||
[["file", "batch.jsonl", "application/octet-stream"]], separators=(",", ":")
|
||||
)
|
||||
assert interaction.request.file_bytes == len(BATCH_JSONL)
|
||||
stored = interaction.request.model_dump_json()
|
||||
assert boundary not in stored
|
||||
assert "sk-upload-secret" not in stored
|
||||
assert "custom_id" not in stored
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("fields", "files"),
|
||||
[
|
||||
pytest.param(
|
||||
(("purpose", "batch"),),
|
||||
(("file", "batch.jsonl", b'{"custom_id":"three"}\n'),),
|
||||
id="file-content",
|
||||
),
|
||||
pytest.param(
|
||||
(("purpose", "batch"),),
|
||||
(("file", "other.jsonl", BATCH_JSONL),),
|
||||
id="file-name",
|
||||
),
|
||||
pytest.param(
|
||||
(("purpose", "fine-tune"),),
|
||||
(("file", "batch.jsonl", BATCH_JSONL),),
|
||||
id="form-field",
|
||||
),
|
||||
pytest.param(
|
||||
(("purpose", "batch"), ("purpose", "batch")),
|
||||
(("file", "batch.jsonl", BATCH_JSONL),),
|
||||
id="repeated-form-field",
|
||||
),
|
||||
pytest.param(
|
||||
(("purpose", "batch"),),
|
||||
(
|
||||
("file", "batch.jsonl", BATCH_JSONL),
|
||||
("mask", "mask.jsonl", BATCH_JSONL),
|
||||
),
|
||||
id="extra-file-part",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_a_structurally_different_upload_misses(
|
||||
self,
|
||||
tmp_path: Path,
|
||||
fields: tuple[tuple[str, str], ...],
|
||||
files: tuple[tuple[str, str, bytes], ...],
|
||||
) -> None:
|
||||
root = tmp_path / "bundle"
|
||||
record_upload(
|
||||
root,
|
||||
multipart_body(
|
||||
"aaaaaaaabbbbbbbbccccccccdddddddd",
|
||||
fields=(("purpose", "batch"),),
|
||||
files=(("file", "batch.jsonl", BATCH_JSONL),),
|
||||
),
|
||||
"aaaaaaaabbbbbbbbccccccccdddddddd",
|
||||
)
|
||||
|
||||
drifted = replay_upload(
|
||||
root,
|
||||
multipart_body("11112222333344445555666677778888", fields=fields, files=files),
|
||||
"11112222333344445555666677778888",
|
||||
)
|
||||
assert drifted.status_code == REPLAY_MISS_STATUS
|
||||
|
||||
def test_several_file_parts_separate_when_their_contents_swap(self, tmp_path: Path) -> None:
|
||||
root = tmp_path / "bundle"
|
||||
image, mask = b"image-bytes", b"mask-bytes"
|
||||
record_upload(
|
||||
root,
|
||||
multipart_body(
|
||||
"1a1a1a1a2b2b2b2b3c3c3c3c4d4d4d4d",
|
||||
fields=(("prompt", "a cat"),),
|
||||
files=(("image", "a.png", image), ("mask", "b.png", mask)),
|
||||
),
|
||||
"1a1a1a1a2b2b2b2b3c3c3c3c4d4d4d4d",
|
||||
)
|
||||
|
||||
swapped = replay_upload(
|
||||
root,
|
||||
multipart_body(
|
||||
"5e5e5e5e6f6f6f6f7070707081818181",
|
||||
fields=(("prompt", "a cat"),),
|
||||
files=(("image", "a.png", mask), ("mask", "b.png", image)),
|
||||
),
|
||||
"5e5e5e5e6f6f6f6f7070707081818181",
|
||||
)
|
||||
assert swapped.status_code == REPLAY_MISS_STATUS
|
||||
|
||||
same = replay_upload(
|
||||
root,
|
||||
multipart_body(
|
||||
"9292929203030303a4a4a4a4b5b5b5b5",
|
||||
fields=(("prompt", "a cat"),),
|
||||
files=(("image", "a.png", image), ("mask", "b.png", mask)),
|
||||
),
|
||||
"9292929203030303a4a4a4a4b5b5b5b5",
|
||||
)
|
||||
assert same.status_code == 200, same.body[:400]
|
||||
|
||||
def test_a_body_that_does_not_match_its_declared_boundary_stays_opaque(
|
||||
self, tmp_path: Path
|
||||
) -> None:
|
||||
root = tmp_path / "bundle"
|
||||
opaque = b"custom_id one\ncustom_id two\n"
|
||||
absent = "boundary-that-is-absent-from-the-body"
|
||||
record_upload(root, opaque, absent)
|
||||
|
||||
raw = this_tests_files(root)[0].read_text(encoding="utf-8")
|
||||
interaction = Interaction.model_validate_json(raw)
|
||||
assert interaction.request.form is None
|
||||
assert interaction.request.file_name == "<unparsed-multipart>"
|
||||
assert interaction.request.file_bytes == len(opaque)
|
||||
assert "custom_id" not in interaction.request.model_dump_json()
|
||||
assert replay_upload(root, opaque, absent).status_code == 200
|
||||
|
||||
|
||||
def raw_multipart(boundary: str, *parts: tuple[str, bytes]) -> bytes:
|
||||
"""A body assembled from literal part headers, so a test can send the shapes a
|
||||
well-formed helper cannot: a file part with no filename, a declared per-part content
|
||||
type, a repeated or bracketed field name, or a non-UTF-8 value."""
|
||||
return (
|
||||
b"".join(
|
||||
f"--{boundary}\r\n{head}\r\n\r\n".encode() + content + b"\r\n"
|
||||
for head, content in parts
|
||||
)
|
||||
+ f"--{boundary}--\r\n".encode()
|
||||
)
|
||||
|
||||
|
||||
def upload_key(body: bytes, boundary: str) -> str:
|
||||
content_type: Final = f"multipart/form-data; boundary={boundary}"
|
||||
return canonicalize(edge_request("POST", UPLOAD_PATH, "", body, content_type)).key
|
||||
|
||||
|
||||
DISPOSITION = 'Content-Disposition: form-data; name="{name}"'
|
||||
FILE_DISPOSITION = DISPOSITION + '; filename="{filename}"'
|
||||
|
||||
|
||||
class TestMultipartIdentityEdges:
|
||||
"""The identity a multipart upload keys on, pinned against the ways two materially
|
||||
different uploads could otherwise collapse onto one key. A collision here is the
|
||||
dangerous failure: replay would answer one request with another's response."""
|
||||
|
||||
def test_a_declared_part_content_type_separates_otherwise_identical_uploads(self) -> None:
|
||||
boundary = "0123456789abcdef0123456789abcdef"
|
||||
as_json = raw_multipart(
|
||||
boundary,
|
||||
(FILE_DISPOSITION.format(name="file", filename="a") + "\r\nContent-Type: application/json", b"xy"),
|
||||
)
|
||||
as_csv = raw_multipart(
|
||||
boundary,
|
||||
(FILE_DISPOSITION.format(name="file", filename="a") + "\r\nContent-Type: text/csv", b"xy"),
|
||||
)
|
||||
|
||||
assert upload_key(as_json, boundary) != upload_key(as_csv, boundary)
|
||||
|
||||
def test_a_file_part_without_a_filename_is_not_mistaken_for_a_plain_field(self) -> None:
|
||||
boundary = "0123456789abcdef0123456789abcdef"
|
||||
upload = raw_multipart(
|
||||
boundary,
|
||||
(DISPOSITION.format(name="file") + "\r\nContent-Type: application/octet-stream", b"CONTENT"),
|
||||
)
|
||||
plain_field = raw_multipart(boundary, (DISPOSITION.format(name="file"), b"CONTENT"))
|
||||
|
||||
request = edge_request(
|
||||
"POST", UPLOAD_PATH, "", upload, f"multipart/form-data; boundary={boundary}"
|
||||
)
|
||||
|
||||
assert upload_key(upload, boundary) != upload_key(plain_field, boundary)
|
||||
assert request.form == {}
|
||||
assert b"CONTENT".decode() not in request.model_dump_json()
|
||||
|
||||
def test_a_filename_carrying_a_per_run_marker_keys_the_same_next_run(self) -> None:
|
||||
boundary = "0123456789abcdef0123456789abcdef"
|
||||
|
||||
def upload(marker: str) -> str:
|
||||
body = raw_multipart(
|
||||
boundary,
|
||||
(FILE_DISPOSITION.format(name="one", filename=f"{marker}.jsonl"), b"first"),
|
||||
(FILE_DISPOSITION.format(name="two", filename="steady.jsonl"), b"second"),
|
||||
)
|
||||
return upload_key(body, boundary)
|
||||
|
||||
assert upload("a1b2c3d4e5f6") == upload("0f9e8d7c6b5a")
|
||||
|
||||
def test_a_separator_inside_a_filename_cannot_forge_a_different_split(self) -> None:
|
||||
boundary = "0123456789abcdef0123456789abcdef"
|
||||
colon_in_filename = raw_multipart(
|
||||
boundary, (FILE_DISPOSITION.format(name="file", filename="a:b.jsonl"), b"same")
|
||||
)
|
||||
colon_in_field = raw_multipart(
|
||||
boundary, (FILE_DISPOSITION.format(name="file:a", filename="b.jsonl"), b"same")
|
||||
)
|
||||
|
||||
assert upload_key(colon_in_filename, boundary) != upload_key(colon_in_field, boundary)
|
||||
|
||||
def test_a_repeated_field_cannot_collide_with_a_literal_indexed_name(self) -> None:
|
||||
boundary = "0123456789abcdef0123456789abcdef"
|
||||
repeated = raw_multipart(
|
||||
boundary,
|
||||
(DISPOSITION.format(name="purpose"), b"x"),
|
||||
(DISPOSITION.format(name="purpose"), b"y"),
|
||||
)
|
||||
literal_index = raw_multipart(
|
||||
boundary,
|
||||
(DISPOSITION.format(name="purpose"), b"x"),
|
||||
(DISPOSITION.format(name="purpose[1]"), b"y"),
|
||||
)
|
||||
|
||||
assert upload_key(repeated, boundary) != upload_key(literal_index, boundary)
|
||||
|
||||
def test_two_binary_field_values_of_one_length_stay_apart(self) -> None:
|
||||
boundary = "0123456789abcdef0123456789abcdef"
|
||||
first = raw_multipart(boundary, (DISPOSITION.format(name="blob"), b"\xff\xfe\xfd"))
|
||||
second = raw_multipart(boundary, (DISPOSITION.format(name="blob"), b"\xf0\xf1\xf2"))
|
||||
|
||||
assert upload_key(first, boundary) != upload_key(second, boundary)
|
||||
|
||||
def test_a_secret_named_field_never_reaches_the_stored_request(self) -> None:
|
||||
boundary = "0123456789abcdef0123456789abcdef"
|
||||
body = raw_multipart(
|
||||
boundary,
|
||||
(DISPOSITION.format(name="openai_api_key"), b"sk-live-DEADBEEF-0123456789abcd"),
|
||||
(DISPOSITION.format(name="purpose"), b"batch"),
|
||||
)
|
||||
|
||||
request = edge_request(
|
||||
"POST", UPLOAD_PATH, "", body, f"multipart/form-data; boundary={boundary}"
|
||||
)
|
||||
|
||||
assert "sk-live-DEADBEEF-0123456789abcd" not in request.model_dump_json()
|
||||
assert request.form == {"openai_api_key": "<secret>", "purpose": "batch"}
|
||||
|
||||
def test_a_redacted_field_still_matches_the_live_request_that_carried_the_secret(
|
||||
self,
|
||||
) -> None:
|
||||
boundary = "0123456789abcdef0123456789abcdef"
|
||||
|
||||
def upload(secret: str) -> str:
|
||||
body = raw_multipart(
|
||||
boundary,
|
||||
(DISPOSITION.format(name="openai_api_key"), secret.encode()),
|
||||
(DISPOSITION.format(name="purpose"), b"batch"),
|
||||
)
|
||||
return upload_key(body, boundary)
|
||||
|
||||
assert upload("sk-live-DEADBEEF-0123456789abcd") == upload("<secret>")
|
||||
|
||||
def test_a_length_change_the_canonicalizer_absorbs_does_not_move_the_key(self) -> None:
|
||||
boundary = "0123456789abcdef0123456789abcdef"
|
||||
|
||||
def upload(created: str) -> str:
|
||||
body = raw_multipart(
|
||||
boundary,
|
||||
(
|
||||
FILE_DISPOSITION.format(name="file", filename="batch.jsonl"),
|
||||
b'{"created_at":"' + created.encode() + b'"}',
|
||||
),
|
||||
)
|
||||
return upload_key(body, boundary)
|
||||
|
||||
assert upload("2026-08-21T02:08:19Z") == upload("2026-08-21T02:08:19.123456Z")
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"content_type",
|
||||
[
|
||||
pytest.param("multipart/form-data; myboundary=zzz; boundary={boundary}", id="lookalike-parameter"),
|
||||
pytest.param("multipart/form-data; BOUNDARY={boundary}", id="uppercase-parameter"),
|
||||
],
|
||||
)
|
||||
def test_the_boundary_parameter_is_read_the_way_the_client_meant_it(
|
||||
self, content_type: str
|
||||
) -> None:
|
||||
boundary = "0123456789abcdef0123456789abcdef"
|
||||
body = raw_multipart(
|
||||
boundary, (FILE_DISPOSITION.format(name="file", filename="batch.jsonl"), BATCH_JSONL)
|
||||
)
|
||||
|
||||
request = edge_request(
|
||||
"POST", UPLOAD_PATH, "", body, content_type.format(boundary=boundary)
|
||||
)
|
||||
|
||||
assert request.form == {}
|
||||
assert request.file_name is not None
|
||||
assert "batch.jsonl" in request.file_name
|
||||
|
||||
def test_an_empty_declared_boundary_falls_back_instead_of_splitting_on_dashes(self) -> None:
|
||||
body = b'--\r\nContent-Disposition: form-data; name="a"\r\n\r\nvalue\r\n----\r\n'
|
||||
|
||||
request = edge_request(
|
||||
"POST", UPLOAD_PATH, "", body, 'multipart/form-data; boundary=""'
|
||||
)
|
||||
|
||||
assert request.form is None
|
||||
assert request.file_sha256 is not None
|
||||
|
||||
|
||||
class TestReplayLeftover:
|
||||
def test_partially_consumed_recording_names_the_leftover(self, tmp_path: Path) -> None:
|
||||
root = tmp_path / "bundle"
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -1508,7 +1508,7 @@ def test_multiple_tool_calls_in_single_choice():
|
|||
print("✓ Multiple tool calls are correctly grouped in a single choice")
|
||||
|
||||
|
||||
def test_map_reasoning_effort_adds_summary_detailed():
|
||||
def test_map_reasoning_effort_adds_summary_detailed(monkeypatch):
|
||||
"""
|
||||
Test that _map_reasoning_effort behavior with reasoning_auto_summary flag.
|
||||
|
||||
|
|
@ -1571,7 +1571,7 @@ def test_map_reasoning_effort_adds_summary_detailed():
|
|||
|
||||
# Test 3: With env var enabled (flag disabled) - summary IS added
|
||||
litellm.reasoning_auto_summary = False
|
||||
os.environ["LITELLM_REASONING_AUTO_SUMMARY"] = "true"
|
||||
monkeypatch.setenv("LITELLM_REASONING_AUTO_SUMMARY", "true")
|
||||
|
||||
result = handler._map_reasoning_effort("high")
|
||||
assert (
|
||||
|
|
@ -1603,7 +1603,7 @@ def test_map_reasoning_effort_adds_summary_detailed():
|
|||
# Restore original values
|
||||
litellm.reasoning_auto_summary = original_flag
|
||||
if original_env is not None:
|
||||
os.environ["LITELLM_REASONING_AUTO_SUMMARY"] = original_env
|
||||
monkeypatch.setenv("LITELLM_REASONING_AUTO_SUMMARY", original_env)
|
||||
elif "LITELLM_REASONING_AUTO_SUMMARY" in os.environ:
|
||||
del os.environ["LITELLM_REASONING_AUTO_SUMMARY"]
|
||||
|
||||
|
|
|
|||
|
|
@ -341,10 +341,10 @@ class TestOpenAIContainerTransformation:
|
|||
assert data["expires_after"] is None
|
||||
assert data["file_ids"] is None
|
||||
|
||||
def test_container_create_response_includes_cost(self):
|
||||
def test_container_create_response_includes_cost(self, monkeypatch):
|
||||
"""Test that container create response includes code interpreter cost calculation."""
|
||||
# Force use of local model cost map for CI/CD consistency
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
from litellm.litellm_core_utils.llm_cost_calc.tool_call_cost_tracking import (
|
||||
|
|
|
|||
|
|
@ -88,49 +88,44 @@ async def test_send_email_success(mock_env_vars):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_email_missing_api_key():
|
||||
async def test_send_email_missing_api_key(monkeypatch):
|
||||
# Remove the API key from environment before initializing logger
|
||||
original_key = os.environ.pop("RESEND_API_KEY", None)
|
||||
monkeypatch.delenv("RESEND_API_KEY", raising=False)
|
||||
|
||||
try:
|
||||
# Initialize the logger after removing the API key
|
||||
logger = ResendEmailLogger()
|
||||
# Initialize the logger after removing the API key
|
||||
logger = ResendEmailLogger()
|
||||
|
||||
# Test data
|
||||
from_email = "test@example.com"
|
||||
to_email = ["recipient@example.com"]
|
||||
subject = "Test Subject"
|
||||
html_body = "<p>Test email body</p>"
|
||||
# Test data
|
||||
from_email = "test@example.com"
|
||||
to_email = ["recipient@example.com"]
|
||||
subject = "Test Subject"
|
||||
html_body = "<p>Test email body</p>"
|
||||
|
||||
# Create mock HTTP client and inject it directly into the logger
|
||||
# This ensures the mock is used regardless of any caching issues
|
||||
mock_response = mock.Mock(spec=Response)
|
||||
mock_response.raise_for_status.return_value = None
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {"id": "test_email_id"}
|
||||
# Create mock HTTP client and inject it directly into the logger
|
||||
# This ensures the mock is used regardless of any caching issues
|
||||
mock_response = mock.Mock(spec=Response)
|
||||
mock_response.raise_for_status.return_value = None
|
||||
mock_response.status_code = 200
|
||||
mock_response.json.return_value = {"id": "test_email_id"}
|
||||
|
||||
mock_async_client = mock.AsyncMock()
|
||||
mock_async_client.post.return_value = mock_response
|
||||
mock_async_client = mock.AsyncMock()
|
||||
mock_async_client.post.return_value = mock_response
|
||||
|
||||
# Directly inject the mock client to bypass any caching
|
||||
logger.async_httpx_client = mock_async_client
|
||||
# Directly inject the mock client to bypass any caching
|
||||
logger.async_httpx_client = mock_async_client
|
||||
|
||||
# Send email
|
||||
await logger.send_email(
|
||||
from_email=from_email,
|
||||
to_email=to_email,
|
||||
subject=subject,
|
||||
html_body=html_body,
|
||||
)
|
||||
# Send email
|
||||
await logger.send_email(
|
||||
from_email=from_email,
|
||||
to_email=to_email,
|
||||
subject=subject,
|
||||
html_body=html_body,
|
||||
)
|
||||
|
||||
# Verify the HTTP client was called with None as the API key
|
||||
mock_async_client.post.assert_called_once()
|
||||
call_args = mock_async_client.post.call_args
|
||||
assert call_args[1]["headers"] == {"Authorization": "Bearer None"}
|
||||
finally:
|
||||
# Restore the original key if it existed
|
||||
if original_key is not None:
|
||||
os.environ["RESEND_API_KEY"] = original_key
|
||||
# Verify the HTTP client was called with None as the API key
|
||||
mock_async_client.post.assert_called_once()
|
||||
call_args = mock_async_client.post.call_args
|
||||
assert call_args[1]["headers"] == {"Authorization": "Bearer None"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -98,22 +98,18 @@ async def test_send_email_success(mock_env_vars, mock_async_client):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_send_email_missing_api_key():
|
||||
original_key = os.environ.pop("SENDGRID_API_KEY", None)
|
||||
async def test_send_email_missing_api_key(monkeypatch):
|
||||
monkeypatch.delenv("SENDGRID_API_KEY", raising=False)
|
||||
|
||||
try:
|
||||
logger = SendGridEmailLogger()
|
||||
logger = SendGridEmailLogger()
|
||||
|
||||
with pytest.raises(ValueError, match='SENDGRID_API_KEY is not set'):
|
||||
await logger.send_email(
|
||||
from_email="test@example.com",
|
||||
to_email=["recipient@example.com"],
|
||||
subject="Test Subject",
|
||||
html_body="<p>Test email body</p>",
|
||||
)
|
||||
finally:
|
||||
if original_key is not None:
|
||||
os.environ["SENDGRID_API_KEY"] = original_key
|
||||
with pytest.raises(ValueError, match='SENDGRID_API_KEY is not set'):
|
||||
await logger.send_email(
|
||||
from_email="test@example.com",
|
||||
to_email=["recipient@example.com"],
|
||||
subject="Test Subject",
|
||||
html_body="<p>Test email body</p>",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -8,10 +8,10 @@ from litellm.integrations.gcs_bucket.gcs_bucket_base import GCSBucketBase
|
|||
|
||||
|
||||
class TestGCSBucketBase:
|
||||
def test_construct_request_headers_with_project_id(self):
|
||||
def test_construct_request_headers_with_project_id(self, monkeypatch):
|
||||
"""Test that construct_request_headers correctly uses project_id if passed from env"""
|
||||
test_project_id = "test-project"
|
||||
os.environ["GOOGLE_SECRET_MANAGER_PROJECT_ID"] = test_project_id
|
||||
monkeypatch.setenv("GOOGLE_SECRET_MANAGER_PROJECT_ID", test_project_id)
|
||||
|
||||
try:
|
||||
# Create handler
|
||||
|
|
|
|||
|
|
@ -236,9 +236,9 @@ class TestOpenMeterIntegration:
|
|||
assert result["data"]["completion_tokens"] == 8
|
||||
assert result["data"]["total_tokens"] == 23
|
||||
|
||||
def test_custom_event_type(self):
|
||||
def test_custom_event_type(self, monkeypatch):
|
||||
"""Test that custom event type is used when set"""
|
||||
os.environ["OPENMETER_EVENT_TYPE"] = "custom_event_type"
|
||||
monkeypatch.setenv("OPENMETER_EVENT_TYPE", "custom_event_type")
|
||||
|
||||
logger = OpenMeterLogger()
|
||||
|
||||
|
|
@ -374,10 +374,10 @@ class TestOpenMeterIntegration:
|
|||
assert isinstance(result["subject"], str)
|
||||
assert result["subject"] == "12345"
|
||||
|
||||
def test_common_logic_trust_request_user_false_ignores_request_user(self):
|
||||
def test_common_logic_trust_request_user_false_ignores_request_user(self, monkeypatch):
|
||||
"""OPENMETER_TRUST_REQUEST_USER=false makes the key-bound user_id win
|
||||
over a request-supplied `user` (forge-attribution mitigation)."""
|
||||
os.environ["OPENMETER_TRUST_REQUEST_USER"] = "false"
|
||||
monkeypatch.setenv("OPENMETER_TRUST_REQUEST_USER", "false")
|
||||
logger = OpenMeterLogger()
|
||||
|
||||
kwargs = {
|
||||
|
|
@ -400,11 +400,11 @@ class TestOpenMeterIntegration:
|
|||
assert result["subject"] == "real-tenant-id"
|
||||
assert result["subject"] != "forged-by-client"
|
||||
|
||||
def test_common_logic_trust_request_user_false_still_raises_without_key_user(self):
|
||||
def test_common_logic_trust_request_user_false_still_raises_without_key_user(self, monkeypatch):
|
||||
"""OPENMETER_TRUST_REQUEST_USER=false still raises when no
|
||||
user_api_key_user_id is available — the request `user` is not a
|
||||
fallback in this mode."""
|
||||
os.environ["OPENMETER_TRUST_REQUEST_USER"] = "false"
|
||||
monkeypatch.setenv("OPENMETER_TRUST_REQUEST_USER", "false")
|
||||
logger = OpenMeterLogger()
|
||||
|
||||
kwargs = {
|
||||
|
|
|
|||
|
|
@ -56,8 +56,8 @@ def test_bedrock_guardrail_cost_no_pricing_entry(monkeypatch):
|
|||
assert bedrock_guardrail_cost(usage_units={"contentPolicyUnits": 1}, aws_region_name="us-east-1") == 0.0
|
||||
|
||||
|
||||
def test_shipped_bedrock_guardrail_prices_match_aws_pricing_page():
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
def test_shipped_bedrock_guardrail_prices_match_aws_pricing_page(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
assert litellm.model_cost["bedrock/guardrails"]["guardrail_cost_per_unit"] == {
|
||||
"automatedReasoningPolicyUnits": 0.00017,
|
||||
|
|
|
|||
|
|
@ -913,8 +913,8 @@ def test_generic_cost_per_token_gpt55_pro():
|
|||
@pytest.mark.parametrize(
|
||||
"model,input_cost,output_cost,cache_read_cost,cache_write_cost",
|
||||
[
|
||||
("gpt-5.6", 5e-6, 3e-5, 5e-7, 6.25e-6),
|
||||
("gpt-5.6-sol", 5e-6, 3e-5, 5e-7, 6.25e-6),
|
||||
("gpt-5.6", 4e-6, 2e-5, 4e-7, 5e-6),
|
||||
("gpt-5.6-sol", 4e-6, 2e-5, 4e-7, 5e-6),
|
||||
("gpt-5.6-terra", 2e-6, 1.2e-5, 2e-7, 2.5e-6),
|
||||
("gpt-5.6-luna", 2e-7, 1.2e-6, 2e-8, 2.5e-7),
|
||||
],
|
||||
|
|
@ -965,11 +965,26 @@ def test_generic_cost_per_token_gpt56(
|
|||
assert round(completion_cost, 10) == round(output_cost * completion_tokens, 10)
|
||||
|
||||
|
||||
def test_gpt_5_6_alias_prices_match_sol(local_model_cost_map):
|
||||
"""Regression: the bare gpt-5.6 alias routes to GPT-5.6 Sol, so every cost field on
|
||||
the two entries has to hold the same value. They drifted once before, when Sol took
|
||||
its promotional cut and gpt-5.6 was left on the pre-cut rates, overbilling callers
|
||||
who used the alias."""
|
||||
alias = litellm.model_cost["gpt-5.6"]
|
||||
sol = litellm.model_cost["gpt-5.6-sol"]
|
||||
|
||||
cost_fields = sorted(field for field in sol if "cost" in field)
|
||||
assert len(cost_fields) == 23
|
||||
|
||||
for field in cost_fields:
|
||||
assert alias.get(field) == sol.get(field), field
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model,flex_long_input_cost,flex_long_output_cost",
|
||||
[
|
||||
("gpt-5.6", 5e-6, 2.25e-5),
|
||||
("gpt-5.6-sol", 5e-6, 2.25e-5),
|
||||
("gpt-5.6", 4e-6, 1.5e-5),
|
||||
("gpt-5.6-sol", 4e-6, 1.5e-5),
|
||||
("gpt-5.6-terra", 2e-6, 9e-6),
|
||||
("gpt-5.6-luna", 2e-7, 9e-7),
|
||||
],
|
||||
|
|
@ -1118,8 +1133,10 @@ def test_generic_cost_per_token_gpt56_cyber(
|
|||
def test_generic_cost_per_token_azure_gpt56(
|
||||
model, input_cost, output_cost, cache_read_cost
|
||||
):
|
||||
"""Azure gpt-5.6 (global + us/eu regional): pricing mirrors the openai
|
||||
family for global deployments and carries the standard 10% regional uplift.
|
||||
"""Azure gpt-5.6 (global + us/eu regional): Azure prices this family on its own
|
||||
schedule and carries the standard 10% regional uplift on top. It did not take the
|
||||
promotional cut OpenAI applied to gpt-5.6-sol, so these rates deliberately sit
|
||||
above the openai ones and must not be lowered to match them.
|
||||
"""
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
|
@ -3301,8 +3318,8 @@ def test_generic_cost_per_token_gemini_35_flash_lite():
|
|||
@pytest.mark.parametrize(
|
||||
"service_tier,input_rate,cache_read_rate,cache_write_rate,output_rate",
|
||||
[
|
||||
("flex", 2.5e-6, 2.5e-7, 3.125e-6, 1.5e-5),
|
||||
("priority", 1e-5, 1e-6, 1.25e-5, 6e-5),
|
||||
("flex", 2e-6, 2e-7, 2.5e-6, 1e-5),
|
||||
("priority", 8e-6, 8e-7, 1e-5, 4e-5),
|
||||
],
|
||||
)
|
||||
def test_service_tier_cache_creation_rates_for_gpt_5_6(
|
||||
|
|
@ -3315,7 +3332,7 @@ def test_service_tier_cache_creation_rates_for_gpt_5_6(
|
|||
):
|
||||
"""Regression: gpt-5.6 publishes cache_creation_input_token_cost_flex/_priority, so a
|
||||
flex or priority request must bill cache writes at that tier's rate instead of falling
|
||||
back to the standard 6.25e-6 rate."""
|
||||
back to the standard cache-write rate."""
|
||||
usage = Usage(
|
||||
prompt_tokens=10_000,
|
||||
completion_tokens=500,
|
||||
|
|
@ -3362,8 +3379,8 @@ def test_fast_service_tier_bills_at_the_priority_rate(_local_model_cost_map):
|
|||
model="gpt-5.6-sol", usage=usage, custom_llm_provider="openai", service_tier="fast"
|
||||
)
|
||||
|
||||
expected_prompt = 800 * 1e-05 + 200 * 1e-06
|
||||
expected_completion = 500 * 6e-05
|
||||
expected_prompt = 800 * 8e-06 + 200 * 8e-07
|
||||
expected_completion = 500 * 4e-05
|
||||
|
||||
assert fast == priority
|
||||
assert fast[0] == pytest.approx(expected_prompt, rel=1e-9)
|
||||
|
|
@ -3398,8 +3415,8 @@ def test_fast_service_tier_matches_priority_above_the_context_threshold(_local_m
|
|||
)
|
||||
|
||||
assert fast == priority
|
||||
assert fast[0] == pytest.approx(300_000 * 1e-05, rel=1e-9)
|
||||
assert fast[1] == pytest.approx(1_000 * 4.5e-05, rel=1e-9)
|
||||
assert fast[0] == pytest.approx(300_000 * 8e-06, rel=1e-9)
|
||||
assert fast[1] == pytest.approx(1_000 * 3e-05, rel=1e-9)
|
||||
|
||||
|
||||
def test_priority_reasoning_tokens_bill_at_the_priority_output_rate(_local_model_cost_map):
|
||||
|
|
|
|||
|
|
@ -377,12 +377,12 @@ def test_get_cost_for_vertex_ai_gemini_web_search(model, custom_llm_provider):
|
|||
assert cost == 0.035, f"Expected $0.035 grounding cost, got ${cost}"
|
||||
|
||||
|
||||
def test_azure_assistant_features_integrated_cost_tracking():
|
||||
def test_azure_assistant_features_integrated_cost_tracking(monkeypatch):
|
||||
"""
|
||||
Test integrated cost tracking for Azure assistant features.
|
||||
"""
|
||||
# Force use of local model cost map for CI/CD consistency
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
model = "azure/gpt-4o"
|
||||
|
|
|
|||
|
|
@ -2844,7 +2844,7 @@ def test_anthropic_messages_pt_file_block_preserves_cache_control():
|
|||
assert text_block["cache_control"]["type"] == "ephemeral"
|
||||
|
||||
|
||||
def test_add_cache_point_tool_block_passes_ttl_for_claude_4_5():
|
||||
def test_add_cache_point_tool_block_passes_ttl_for_claude_4_5(monkeypatch):
|
||||
"""
|
||||
Tools with cache_control ttl should preserve the ttl in the cachePoint
|
||||
block for Claude 4.5+ models on Bedrock, matching the behavior of system
|
||||
|
|
@ -2867,7 +2867,7 @@ def test_add_cache_point_tool_block_passes_ttl_for_claude_4_5():
|
|||
|
||||
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
|
||||
old_cost = litellm.model_cost
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
try:
|
||||
tool_with_1h = {
|
||||
|
|
@ -2927,10 +2927,10 @@ def test_add_cache_point_tool_block_passes_ttl_for_claude_4_5():
|
|||
if old_env is None:
|
||||
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
|
||||
else:
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", old_env)
|
||||
|
||||
|
||||
def test_bedrock_tools_pt_passes_ttl_for_claude_4_5():
|
||||
def test_bedrock_tools_pt_passes_ttl_for_claude_4_5(monkeypatch):
|
||||
"""
|
||||
End-to-end: _bedrock_tools_pt should produce cachePoint blocks with ttl
|
||||
for Claude 4.5+ models when tools have cache_control with ttl.
|
||||
|
|
@ -2944,7 +2944,7 @@ def test_bedrock_tools_pt_passes_ttl_for_claude_4_5():
|
|||
|
||||
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
|
||||
old_cost = litellm.model_cost
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
try:
|
||||
tools = [
|
||||
|
|
@ -2980,7 +2980,7 @@ def test_bedrock_tools_pt_passes_ttl_for_claude_4_5():
|
|||
if old_env is None:
|
||||
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
|
||||
else:
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", old_env)
|
||||
|
||||
|
||||
def test_convert_to_anthropic_tool_result_openai_file_pdf_becomes_document():
|
||||
|
|
|
|||
|
|
@ -64,7 +64,7 @@ def test_post_call_serializes_dict_with_datetime(logging_obj):
|
|||
assert "2026-05-11" in serialized
|
||||
|
||||
|
||||
def test_sentry_sample_rate():
|
||||
def test_sentry_sample_rate(monkeypatch):
|
||||
existing_sample_rate = os.getenv("SENTRY_API_SAMPLE_RATE")
|
||||
try:
|
||||
# test with default value by removing the environment variable
|
||||
|
|
@ -76,7 +76,7 @@ def test_sentry_sample_rate():
|
|||
assert os.environ.get("SENTRY_API_SAMPLE_RATE") == "1.0"
|
||||
|
||||
# test with custom value
|
||||
os.environ["SENTRY_API_SAMPLE_RATE"] = "0.5"
|
||||
monkeypatch.setenv("SENTRY_API_SAMPLE_RATE", "0.5")
|
||||
|
||||
set_callbacks(["sentry"])
|
||||
# Check if the custom sample rate is set correctly
|
||||
|
|
@ -86,13 +86,13 @@ def test_sentry_sample_rate():
|
|||
finally:
|
||||
# Restore the original environment variable
|
||||
if existing_sample_rate:
|
||||
os.environ["SENTRY_API_SAMPLE_RATE"] = existing_sample_rate
|
||||
monkeypatch.setenv("SENTRY_API_SAMPLE_RATE", existing_sample_rate)
|
||||
else:
|
||||
if "SENTRY_API_SAMPLE_RATE" in os.environ:
|
||||
del os.environ["SENTRY_API_SAMPLE_RATE"]
|
||||
|
||||
|
||||
def test_sentry_environment():
|
||||
def test_sentry_environment(monkeypatch):
|
||||
"""Test that SENTRY_ENVIRONMENT is properly handled during Sentry initialization"""
|
||||
existing_environment = os.getenv("SENTRY_ENVIRONMENT")
|
||||
existing_dsn = os.getenv("SENTRY_DSN")
|
||||
|
|
@ -115,7 +115,7 @@ def test_sentry_environment():
|
|||
|
||||
try:
|
||||
# Set a mock DSN to allow Sentry initialization
|
||||
os.environ["SENTRY_DSN"] = "https://test@sentry.io/123456"
|
||||
monkeypatch.setenv("SENTRY_DSN", "https://test@sentry.io/123456")
|
||||
|
||||
# Test with default value (no environment set)
|
||||
if existing_environment:
|
||||
|
|
@ -129,7 +129,7 @@ def test_sentry_environment():
|
|||
assert call_kwargs["environment"] == "production"
|
||||
|
||||
# Test with custom environment value
|
||||
os.environ["SENTRY_ENVIRONMENT"] = "development"
|
||||
monkeypatch.setenv("SENTRY_ENVIRONMENT", "development")
|
||||
|
||||
mock_init.reset_mock()
|
||||
set_callbacks(["sentry"])
|
||||
|
|
@ -139,7 +139,7 @@ def test_sentry_environment():
|
|||
assert call_kwargs["environment"] == "development"
|
||||
|
||||
# Test with staging environment
|
||||
os.environ["SENTRY_ENVIRONMENT"] = "staging"
|
||||
monkeypatch.setenv("SENTRY_ENVIRONMENT", "staging")
|
||||
|
||||
mock_init.reset_mock()
|
||||
set_callbacks(["sentry"])
|
||||
|
|
@ -154,13 +154,13 @@ def test_sentry_environment():
|
|||
finally:
|
||||
# Restore the original environment variables
|
||||
if existing_environment:
|
||||
os.environ["SENTRY_ENVIRONMENT"] = existing_environment
|
||||
monkeypatch.setenv("SENTRY_ENVIRONMENT", existing_environment)
|
||||
else:
|
||||
if "SENTRY_ENVIRONMENT" in os.environ:
|
||||
del os.environ["SENTRY_ENVIRONMENT"]
|
||||
|
||||
if existing_dsn:
|
||||
os.environ["SENTRY_DSN"] = existing_dsn
|
||||
monkeypatch.setenv("SENTRY_DSN", existing_dsn)
|
||||
else:
|
||||
if "SENTRY_DSN" in os.environ:
|
||||
del os.environ["SENTRY_DSN"]
|
||||
|
|
|
|||
|
|
@ -66,5 +66,5 @@ def test_prepare_completion_kwargs_keeps_prompt_cache_key_through_responses_rero
|
|||
{"custom_llm_provider": "openai"},
|
||||
thinking={"type": "enabled", "budget_tokens": 1024},
|
||||
)
|
||||
assert completion_kwargs["model"] == "responses/openai/gpt-5.6-luna"
|
||||
assert completion_kwargs["model"] == "openai/responses/gpt-5.6-luna"
|
||||
assert completion_kwargs["prompt_cache_key"] == "session-abc"
|
||||
|
|
|
|||
|
|
@ -217,7 +217,10 @@ async def _async_return(value):
|
|||
|
||||
def test_anthropic_experimental_pass_through_messages_handler_custom_llm_provider():
|
||||
"""
|
||||
Test that litellm.completion is called when a custom LLM provider is given
|
||||
Test that litellm.completion is called when a custom LLM provider is given.
|
||||
|
||||
Provider resolution now happens exactly once, inside litellm.completion itself
|
||||
(BerriAI/litellm#37716), so the handler passes the original unresolved model through.
|
||||
"""
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.handler import (
|
||||
anthropic_messages_handler,
|
||||
|
|
@ -241,7 +244,7 @@ def test_anthropic_experimental_pass_through_messages_handler_custom_llm_provide
|
|||
# Verify that the custom provider was passed through
|
||||
call_kwargs = mock_completion.call_args.kwargs
|
||||
assert call_kwargs["custom_llm_provider"] == "my-custom-llm"
|
||||
assert call_kwargs["model"] == "my-custom-llm/my-custom-model"
|
||||
assert call_kwargs["model"] == "my-custom-model"
|
||||
assert call_kwargs["api_key"] == "test-api-key"
|
||||
|
||||
|
||||
|
|
@ -525,7 +528,7 @@ class TestThinkingSummaryPreservation:
|
|||
finally:
|
||||
litellm.reasoning_auto_summary = original
|
||||
|
||||
def test_summary_added_when_env_var_set(self):
|
||||
def test_summary_added_when_env_var_set(self, monkeypatch):
|
||||
"""When LITELLM_REASONING_AUTO_SUMMARY env var is true, summary is added."""
|
||||
import litellm
|
||||
from litellm.llms.anthropic.experimental_pass_through.adapters.handler import (
|
||||
|
|
@ -535,7 +538,7 @@ class TestThinkingSummaryPreservation:
|
|||
original = litellm.reasoning_auto_summary
|
||||
try:
|
||||
litellm.reasoning_auto_summary = False
|
||||
os.environ["LITELLM_REASONING_AUTO_SUMMARY"] = "true"
|
||||
monkeypatch.setenv("LITELLM_REASONING_AUTO_SUMMARY", "true")
|
||||
completion_kwargs = {
|
||||
"model": "responses/gpt-5.2",
|
||||
"custom_llm_provider": "openai",
|
||||
|
|
@ -997,3 +1000,108 @@ def test_first_party_claude_4_8_plus_cost_map_entries_carry_mid_conversation_sys
|
|||
and info.get("supports_mid_conversation_system") is not True
|
||||
]
|
||||
assert missing == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"requested_model, expected_wire_model, expected_url",
|
||||
[
|
||||
(
|
||||
"perplexity/perplexity/kimi-k3",
|
||||
"perplexity/kimi-k3",
|
||||
"https://api.perplexity.ai/v1/responses",
|
||||
),
|
||||
(
|
||||
"perplexity/perplexity/sonar",
|
||||
"perplexity/sonar",
|
||||
"https://api.perplexity.ai/v1/responses",
|
||||
),
|
||||
("perplexity/sonar", "sonar", "https://api.perplexity.ai/chat/completions"),
|
||||
],
|
||||
)
|
||||
async def test_messages_strips_provider_prefix_exactly_once(
|
||||
requested_model, expected_wire_model, expected_url
|
||||
):
|
||||
"""
|
||||
BerriAI/litellm#37716: only the leading provider segment may be stripped on the way upstream.
|
||||
|
||||
A multi-segment id such as perplexity/perplexity/kimi-k3 must reach the provider as
|
||||
perplexity/kimi-k3, matching what /v1/chat/completions and /v1/responses already send.
|
||||
|
||||
The endpoint is asserted alongside the body because perplexity/perplexity/sonar is a
|
||||
Responses-only deployment whose bare id perplexity/sonar is an ordinary chat model, so
|
||||
stripping the prefix must not also move the request onto chat/completions.
|
||||
|
||||
The subject is the outbound request, so the transport is cut at the wire rather than
|
||||
stubbed with a response body: these ids take different bridges (chat completions
|
||||
versus the Responses API) and would otherwise need different response shapes.
|
||||
"""
|
||||
captured = {}
|
||||
|
||||
async def fake_send(self, request, **kwargs):
|
||||
captured["body"] = json.loads(request.content)
|
||||
captured["url"] = str(request.url)
|
||||
raise httpx.ConnectError("cut at the wire", request=request)
|
||||
|
||||
with (
|
||||
patch.object(httpx.AsyncClient, "send", fake_send),
|
||||
pytest.raises(litellm.exceptions.InternalServerError),
|
||||
):
|
||||
await litellm.anthropic.messages.acreate(
|
||||
max_tokens=100,
|
||||
messages=[{"role": "user", "content": "ping"}],
|
||||
model=requested_model,
|
||||
api_key="test-api-key",
|
||||
)
|
||||
|
||||
assert captured["body"]["model"] == expected_wire_model
|
||||
assert captured["url"] == expected_url
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"requested_model, expected_reported_model",
|
||||
[
|
||||
("perplexity/perplexity/kimi-k3", "perplexity/kimi-k3"),
|
||||
("perplexity/sonar", "sonar"),
|
||||
],
|
||||
)
|
||||
async def test_messages_streaming_reports_provider_local_model(requested_model, expected_reported_model):
|
||||
"""
|
||||
BerriAI/litellm#37716: the wire keeps every segment, so ``message_start`` must still
|
||||
report the id the provider itself knows rather than the caller's prefixed deployment id.
|
||||
"""
|
||||
|
||||
class _EmptyStream:
|
||||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self):
|
||||
raise StopAsyncIteration
|
||||
|
||||
with patch("litellm.acompletion", new=AsyncMock(return_value=_EmptyStream())):
|
||||
stream = await litellm.anthropic.messages.acreate(
|
||||
max_tokens=100,
|
||||
messages=[{"role": "user", "content": "ping"}],
|
||||
model=requested_model,
|
||||
api_key="test-api-key",
|
||||
stream=True,
|
||||
)
|
||||
first_event = await stream.__anext__()
|
||||
|
||||
assert json.loads(first_event.decode().split("data: ", 1)[1])["message"]["model"] == expected_reported_model
|
||||
|
||||
|
||||
def test_messages_sync_streaming_reports_provider_local_model():
|
||||
"""Same guarantee as the async bridge, at the sync call site."""
|
||||
with patch("litellm.completion", new=MagicMock(return_value=iter(()))):
|
||||
stream = litellm.anthropic.messages.create(
|
||||
max_tokens=100,
|
||||
messages=[{"role": "user", "content": "ping"}],
|
||||
model="perplexity/perplexity/kimi-k3",
|
||||
api_key="test-api-key",
|
||||
stream=True,
|
||||
)
|
||||
first_event = next(iter(stream))
|
||||
|
||||
assert json.loads(first_event.decode().split("data: ", 1)[1])["message"]["model"] == "perplexity/kimi-k3"
|
||||
|
|
|
|||
|
|
@ -6,6 +6,8 @@ Regression tests for the /v1/messages request-parse fast paths:
|
|||
while resolving the (static) type hints only once per process.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.utils import (
|
||||
AnthropicMessagesRequestUtils,
|
||||
|
|
@ -88,3 +90,61 @@ def test_drop_params_keeps_speed_for_supporting_model():
|
|||
litellm.drop_params = original
|
||||
|
||||
assert result == {"speed": "fast"}
|
||||
|
||||
|
||||
def test_drop_params_strips_sampling_params_for_unsupported_model(monkeypatch):
|
||||
# claude-opus-4-7 has supports_sampling_params: false in the model map; the
|
||||
# API 400s on these rather than ignoring them.
|
||||
monkeypatch.setattr(litellm, "drop_params", False)
|
||||
result = AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param(
|
||||
params={"temperature": 0.3, "top_p": 0.9, "top_k": 40, "stream": True},
|
||||
model="claude-opus-4-7",
|
||||
drop_params=True,
|
||||
)
|
||||
|
||||
assert result == {"stream": True}
|
||||
|
||||
|
||||
def test_drop_params_strips_sampling_params_for_provider_prefixed_model(monkeypatch):
|
||||
# Vertex-routed ids must resolve the same capability flag.
|
||||
monkeypatch.setattr(litellm, "drop_params", False)
|
||||
result = AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param(
|
||||
params={"temperature": 0.3, "top_p": 0.9, "top_k": 40},
|
||||
model="vertex_ai/claude-opus-4-7",
|
||||
drop_params=True,
|
||||
)
|
||||
|
||||
assert result == {}
|
||||
|
||||
|
||||
def test_sampling_params_kept_for_supporting_model(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "drop_params", False)
|
||||
result = AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param(
|
||||
params={"temperature": 0.3, "top_p": 0.9, "top_k": 40},
|
||||
model="claude-sonnet-4-6",
|
||||
drop_params=True,
|
||||
)
|
||||
|
||||
assert result == {"temperature": 0.3, "top_p": 0.9, "top_k": 40}
|
||||
|
||||
|
||||
def test_temperature_1_kept_for_unsupported_model(monkeypatch):
|
||||
# temperature=1 is the one value these models still accept.
|
||||
monkeypatch.setattr(litellm, "drop_params", False)
|
||||
result = AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param(
|
||||
params={"temperature": 1},
|
||||
model="claude-opus-4-7",
|
||||
drop_params=True,
|
||||
)
|
||||
|
||||
assert result == {"temperature": 1}
|
||||
|
||||
|
||||
def test_sampling_param_raises_clean_400_without_drop_params(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "drop_params", False)
|
||||
with pytest.raises(litellm.utils.UnsupportedParamsError, match="does not support temperature"):
|
||||
AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param(
|
||||
params={"temperature": 0.3},
|
||||
model="claude-opus-4-7",
|
||||
drop_params=False,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -1,9 +1,15 @@
|
|||
import json
|
||||
import os
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../../..")))
|
||||
|
||||
import litellm
|
||||
from litellm.llms.anthropic.experimental_pass_through.responses_adapters.handler import (
|
||||
LiteLLMMessagesToResponsesAPIHandler,
|
||||
_build_responses_kwargs,
|
||||
)
|
||||
|
||||
|
|
@ -43,3 +49,36 @@ def test_build_responses_kwargs_without_metadata_sets_no_prompt_cache_key():
|
|||
)
|
||||
assert "user" not in responses_kwargs
|
||||
assert "prompt_cache_key" not in responses_kwargs
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"requested_model, expected_reported_model",
|
||||
[
|
||||
("openai/gpt-5.6-luna", "gpt-5.6-luna"),
|
||||
("perplexity/perplexity/kimi-k3", "perplexity/kimi-k3"),
|
||||
],
|
||||
)
|
||||
async def test_streaming_message_start_reports_the_provider_local_model(requested_model, expected_reported_model):
|
||||
"""
|
||||
BerriAI/litellm#37716 sends the caller's unresolved id down this bridge so the provider
|
||||
resolves it once. ``message_start`` is a reporting field rather than a wire value, so it
|
||||
keeps naming the model as the provider knows it, with only the leading provider segment gone.
|
||||
"""
|
||||
|
||||
async def empty_stream():
|
||||
return
|
||||
yield
|
||||
|
||||
with patch.object(litellm, "aresponses", AsyncMock(return_value=empty_stream())):
|
||||
sse = await LiteLLMMessagesToResponsesAPIHandler.async_anthropic_messages_handler(
|
||||
max_tokens=1024,
|
||||
messages=MESSAGES,
|
||||
model=requested_model,
|
||||
stream=True,
|
||||
custom_llm_provider=requested_model.split("/")[0],
|
||||
)
|
||||
events = [json.loads(chunk.decode().split("data: ", 1)[1]) async for chunk in sse]
|
||||
|
||||
message_start = next(e for e in events if e["type"] == "message_start")
|
||||
assert message_start["message"]["model"] == expected_reported_model
|
||||
|
|
|
|||
|
|
@ -845,14 +845,14 @@ class TestTranslateThinkingToReasoning:
|
|||
finally:
|
||||
litellm.reasoning_auto_summary = original
|
||||
|
||||
def test_summary_added_when_env_var_set(self):
|
||||
def test_summary_added_when_env_var_set(self, monkeypatch):
|
||||
"""When LITELLM_REASONING_AUTO_SUMMARY env var is true, summary is included."""
|
||||
import litellm
|
||||
|
||||
original = litellm.reasoning_auto_summary
|
||||
try:
|
||||
litellm.reasoning_auto_summary = False
|
||||
os.environ["LITELLM_REASONING_AUTO_SUMMARY"] = "true"
|
||||
monkeypatch.setenv("LITELLM_REASONING_AUTO_SUMMARY", "true")
|
||||
result = _ADAPTER.translate_thinking_to_reasoning(
|
||||
{
|
||||
"type": "enabled",
|
||||
|
|
|
|||
|
|
@ -239,8 +239,8 @@ class TestAPISerpentSearchIntegration:
|
|||
return mock_response
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_asearch_quick_default(self):
|
||||
os.environ["APISERPENT_API_KEY"] = "test-api-key"
|
||||
async def test_asearch_quick_default(self, monkeypatch):
|
||||
monkeypatch.setenv("APISERPENT_API_KEY", "test-api-key")
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get",
|
||||
new_callable=AsyncMock,
|
||||
|
|
@ -269,8 +269,8 @@ class TestAPISerpentSearchIntegration:
|
|||
assert response.results[0].title == "Test Result"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_asearch_deep(self):
|
||||
os.environ["APISERPENT_API_KEY"] = "test-api-key"
|
||||
async def test_asearch_deep(self, monkeypatch):
|
||||
monkeypatch.setenv("APISERPENT_API_KEY", "test-api-key")
|
||||
with patch(
|
||||
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get",
|
||||
new_callable=AsyncMock,
|
||||
|
|
|
|||
|
|
@ -40,8 +40,8 @@ class TestAzureMAIImageGeneration:
|
|||
assert not AzureFoundryMAIImageGenerationConfig.is_mai_model("flux.2-pro")
|
||||
assert not AzureFoundryMAIImageGenerationConfig.is_mai_model("MAI-DS-R1")
|
||||
|
||||
def test_mai_flash_and_2e_model_pricing_in_cost_map(self):
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
def test_mai_flash_and_2e_model_pricing_in_cost_map(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
flash_info = litellm.get_model_info(
|
||||
|
|
@ -328,8 +328,8 @@ class TestAzureMAIImageGeneration:
|
|||
assert image_response.usage.total_tokens == 1046
|
||||
assert image_response.size == "1792x1024"
|
||||
|
||||
def test_mai_image_cost_calculator_token_based(self):
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
def test_mai_image_cost_calculator_token_based(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
model = "azure_ai/MAI-Image-2.5"
|
||||
model_info = litellm.get_model_info(model=model, custom_llm_provider="azure_ai")
|
||||
|
|
@ -360,8 +360,8 @@ class TestAzureMAIImageGeneration:
|
|||
)
|
||||
assert round(cost, 10) == round(expected_cost, 10)
|
||||
|
||||
def test_mai_image_cost_calculator_falls_back_to_flat_image_pricing(self):
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
def test_mai_image_cost_calculator_falls_back_to_flat_image_pricing(self, monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
model = "azure_ai/MAI-Image-2.5"
|
||||
model_info = litellm.get_model_info(model=model, custom_llm_provider="azure_ai")
|
||||
|
|
|
|||
|
|
@ -678,10 +678,10 @@ def test_transform_request_helper_includes_anthropic_beta_and_tools():
|
|||
assert fields["tools"][0]["type"] == "computer_20250124"
|
||||
|
||||
|
||||
def test_parallel_tool_calls_config_kept_for_sonnet_5():
|
||||
def test_parallel_tool_calls_config_kept_for_sonnet_5(monkeypatch):
|
||||
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
|
||||
old_cost = litellm.model_cost
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
try:
|
||||
config = AmazonConverseConfig()
|
||||
|
|
@ -708,7 +708,7 @@ def test_parallel_tool_calls_config_kept_for_sonnet_5():
|
|||
if old_env is None:
|
||||
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
|
||||
else:
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", old_env)
|
||||
|
||||
|
||||
def test_parallel_tool_calls_config_dropped_for_ttl_only_model(
|
||||
|
|
@ -3575,7 +3575,7 @@ def test_drop_thinking_param_when_thinking_blocks_missing():
|
|||
litellm.modify_params = original_modify_params
|
||||
|
||||
|
||||
def test_supports_native_structured_outputs():
|
||||
def test_supports_native_structured_outputs(monkeypatch):
|
||||
"""Test model detection for native structured outputs support.
|
||||
|
||||
Support is driven by the ``supports_native_structured_output`` flag in the
|
||||
|
|
@ -3583,7 +3583,7 @@ def test_supports_native_structured_outputs():
|
|||
"""
|
||||
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
|
||||
old_cost = litellm.model_cost
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
try:
|
||||
config = AmazonConverseConfig()
|
||||
|
|
@ -3645,7 +3645,7 @@ def test_supports_native_structured_outputs():
|
|||
if old_env is None:
|
||||
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
|
||||
else:
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", old_env)
|
||||
|
||||
|
||||
def test_create_output_config_for_response_format():
|
||||
|
|
@ -3683,11 +3683,11 @@ def test_create_output_config_for_response_format():
|
|||
assert parsed_schema == expected
|
||||
|
||||
|
||||
def test_translate_response_format_native_output_config():
|
||||
def test_translate_response_format_native_output_config(monkeypatch):
|
||||
"""For supported models, _translate_response_format_param should produce outputConfig."""
|
||||
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
|
||||
old_cost = litellm.model_cost
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
try:
|
||||
config = AmazonConverseConfig()
|
||||
|
|
@ -3743,7 +3743,7 @@ def test_translate_response_format_native_output_config():
|
|||
if old_env is None:
|
||||
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
|
||||
else:
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", old_env)
|
||||
|
||||
|
||||
def test_translate_response_format_fallback_tool_call():
|
||||
|
|
@ -3778,11 +3778,11 @@ def test_translate_response_format_fallback_tool_call():
|
|||
assert result["json_mode"] is True
|
||||
|
||||
|
||||
def test_native_structured_output_no_fake_stream():
|
||||
def test_native_structured_output_no_fake_stream(monkeypatch):
|
||||
"""When using native structured outputs with streaming, fake_stream should NOT be set."""
|
||||
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
|
||||
old_cost = litellm.model_cost
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
try:
|
||||
config = AmazonConverseConfig()
|
||||
|
|
@ -3828,7 +3828,7 @@ def test_native_structured_output_no_fake_stream():
|
|||
if old_env is None:
|
||||
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
|
||||
else:
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", old_env)
|
||||
|
||||
|
||||
def test_transform_request_with_output_config():
|
||||
|
|
@ -4116,7 +4116,7 @@ def test_add_additional_properties_definitions():
|
|||
)
|
||||
|
||||
|
||||
def test_json_object_no_schema_skips_tool_injection():
|
||||
def test_json_object_no_schema_skips_tool_injection(monkeypatch):
|
||||
"""response_format: {type: json_object} with no schema should NOT inject
|
||||
the synthetic json_tool_call tool.
|
||||
|
||||
|
|
@ -4126,7 +4126,7 @@ def test_json_object_no_schema_skips_tool_injection():
|
|||
the model respond naturally with the JSON the caller asked for."""
|
||||
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
|
||||
old_cost = litellm.model_cost
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
try:
|
||||
config = AmazonConverseConfig()
|
||||
|
|
@ -4152,7 +4152,7 @@ def test_json_object_no_schema_skips_tool_injection():
|
|||
if old_env is None:
|
||||
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
|
||||
else:
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", old_env)
|
||||
|
||||
|
||||
def test_output_config_applies_additional_properties():
|
||||
|
|
@ -4805,7 +4805,7 @@ def test_cache_control_injection_tool_config_not_added_without_injection_point()
|
|||
assert all("cachePoint" not in tool for tool in tools)
|
||||
|
||||
|
||||
def test_cache_control_injection_tool_config_honors_ttl_for_supported_model():
|
||||
def test_cache_control_injection_tool_config_honors_ttl_for_supported_model(monkeypatch):
|
||||
"""
|
||||
Regression test: cache_control_injection_points with location=tool_config
|
||||
must honor the requested `control.ttl`, mirroring the message/system
|
||||
|
|
@ -4819,7 +4819,7 @@ def test_cache_control_injection_tool_config_honors_ttl_for_supported_model():
|
|||
"""
|
||||
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
|
||||
old_cost = litellm.model_cost
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
try:
|
||||
config = AmazonConverseConfig()
|
||||
|
|
@ -4858,10 +4858,10 @@ def test_cache_control_injection_tool_config_honors_ttl_for_supported_model():
|
|||
if old_env is None:
|
||||
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
|
||||
else:
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", old_env)
|
||||
|
||||
|
||||
def test_cache_control_injection_tool_config_honors_ttl_for_regional_model_lacking_own_pricing():
|
||||
def test_cache_control_injection_tool_config_honors_ttl_for_regional_model_lacking_own_pricing(monkeypatch):
|
||||
"""
|
||||
Regression test: a regional pricing entry that omits
|
||||
`cache_creation_input_token_cost_above_1hr` (e.g. `jp.anthropic.claude-opus-4-7`)
|
||||
|
|
@ -4870,7 +4870,7 @@ def test_cache_control_injection_tool_config_honors_ttl_for_regional_model_lacki
|
|||
"""
|
||||
old_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
|
||||
old_cost = litellm.model_cost
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
try:
|
||||
assert "cache_creation_input_token_cost_above_1hr" not in litellm.model_cost["jp.anthropic.claude-opus-4-7"]
|
||||
|
|
@ -4911,7 +4911,7 @@ def test_cache_control_injection_tool_config_honors_ttl_for_regional_model_lacki
|
|||
if old_env is None:
|
||||
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
|
||||
else:
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = old_env
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", old_env)
|
||||
|
||||
|
||||
def test_cache_control_injection_tool_config_drops_ttl_for_unsupported_model():
|
||||
|
|
|
|||
|
|
@ -60,9 +60,9 @@ class TestAgentCoreSearch:
|
|||
"""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_agentcore_search_request_payload(self):
|
||||
async def test_agentcore_search_request_payload(self, monkeypatch):
|
||||
"""Validates the MCP tools/call payload and SigV4 signing without real AWS calls."""
|
||||
os.environ["AGENTCORE_GATEWAY_URL"] = GATEWAY_URL
|
||||
monkeypatch.setenv("AGENTCORE_GATEWAY_URL", GATEWAY_URL)
|
||||
|
||||
mock_response = _make_mock_response(_mcp_response_body())
|
||||
|
||||
|
|
@ -321,11 +321,11 @@ class TestAgentCoreSearch:
|
|||
assert headers["Authorization"] == "Bearer test-jwt-token"
|
||||
assert signed_body == json.dumps(request_data).encode()
|
||||
|
||||
def test_sign_request_uses_bearer_token_from_env(self):
|
||||
def test_sign_request_uses_bearer_token_from_env(self, monkeypatch):
|
||||
"""Server token is attached when the request targets the configured gateway host."""
|
||||
config = AgentCoreSearchConfig()
|
||||
os.environ["AGENTCORE_GATEWAY_TOKEN"] = "env-jwt-token"
|
||||
os.environ["AGENTCORE_GATEWAY_URL"] = GATEWAY_URL
|
||||
monkeypatch.setenv("AGENTCORE_GATEWAY_TOKEN", "env-jwt-token")
|
||||
monkeypatch.setenv("AGENTCORE_GATEWAY_URL", GATEWAY_URL)
|
||||
try:
|
||||
headers, _ = config.sign_request(
|
||||
headers={},
|
||||
|
|
@ -338,11 +338,11 @@ class TestAgentCoreSearch:
|
|||
os.environ.pop("AGENTCORE_GATEWAY_TOKEN", None)
|
||||
os.environ.pop("AGENTCORE_GATEWAY_URL", None)
|
||||
|
||||
def test_sign_request_refuses_server_token_to_untrusted_host(self):
|
||||
def test_sign_request_refuses_server_token_to_untrusted_host(self, monkeypatch):
|
||||
"""Server-managed token must not be sent to a caller-chosen api_base."""
|
||||
config = AgentCoreSearchConfig()
|
||||
os.environ["AGENTCORE_GATEWAY_TOKEN"] = "env-jwt-token"
|
||||
os.environ["AGENTCORE_GATEWAY_URL"] = GATEWAY_URL
|
||||
monkeypatch.setenv("AGENTCORE_GATEWAY_TOKEN", "env-jwt-token")
|
||||
monkeypatch.setenv("AGENTCORE_GATEWAY_URL", GATEWAY_URL)
|
||||
try:
|
||||
with pytest.raises(ValueError, match="Refusing to send"):
|
||||
config.sign_request(
|
||||
|
|
@ -355,11 +355,11 @@ class TestAgentCoreSearch:
|
|||
os.environ.pop("AGENTCORE_GATEWAY_TOKEN", None)
|
||||
os.environ.pop("AGENTCORE_GATEWAY_URL", None)
|
||||
|
||||
def test_sign_request_uses_env_token_for_gateway_api_base_without_gateway_url(self):
|
||||
def test_sign_request_uses_env_token_for_gateway_api_base_without_gateway_url(self, monkeypatch):
|
||||
"""api_base pointing at a real gateway is a trusted destination for the env token,
|
||||
so operators configuring api_base in yaml don't also need AGENTCORE_GATEWAY_URL."""
|
||||
config = AgentCoreSearchConfig()
|
||||
os.environ["AGENTCORE_GATEWAY_TOKEN"] = "env-jwt-token"
|
||||
monkeypatch.setenv("AGENTCORE_GATEWAY_TOKEN", "env-jwt-token")
|
||||
os.environ.pop("AGENTCORE_GATEWAY_URL", None)
|
||||
try:
|
||||
headers, _ = config.sign_request(
|
||||
|
|
@ -380,12 +380,12 @@ class TestAgentCoreSearch:
|
|||
"https://attacker.example.com/gw.gateway.bedrock-agentcore.us-east-1.amazonaws.com/mcp",
|
||||
],
|
||||
)
|
||||
def test_sign_request_refuses_sigv4_to_untrusted_host(self, untrusted_api_base):
|
||||
def test_sign_request_refuses_sigv4_to_untrusted_host(self, untrusted_api_base, monkeypatch):
|
||||
"""A SigV4 signature carries the proxy's credential scope and session token, so it
|
||||
must never be sent to a host that is not the operator's gateway."""
|
||||
config = AgentCoreSearchConfig()
|
||||
os.environ.pop("AGENTCORE_GATEWAY_TOKEN", None)
|
||||
os.environ["AGENTCORE_GATEWAY_URL"] = GATEWAY_URL
|
||||
monkeypatch.setenv("AGENTCORE_GATEWAY_URL", GATEWAY_URL)
|
||||
try:
|
||||
with patch.object(
|
||||
AgentCoreSearchConfig.__mro__[2], # BaseAWSLLM
|
||||
|
|
@ -410,12 +410,12 @@ class TestAgentCoreSearch:
|
|||
"http://internal-gateway.corp/mcp",
|
||||
],
|
||||
)
|
||||
def test_sign_request_refuses_server_token_over_plaintext_http(self, plaintext_api_base):
|
||||
def test_sign_request_refuses_server_token_over_plaintext_http(self, plaintext_api_base, monkeypatch):
|
||||
"""A trusted hostname over plain http would expose the bearer token to
|
||||
network observers, so credentials only ride https (or localhost)."""
|
||||
config = AgentCoreSearchConfig()
|
||||
os.environ["AGENTCORE_GATEWAY_TOKEN"] = "env-jwt-token"
|
||||
os.environ["AGENTCORE_GATEWAY_URL"] = plaintext_api_base
|
||||
monkeypatch.setenv("AGENTCORE_GATEWAY_TOKEN", "env-jwt-token")
|
||||
monkeypatch.setenv("AGENTCORE_GATEWAY_URL", plaintext_api_base)
|
||||
try:
|
||||
with pytest.raises(ValueError, match="plaintext"):
|
||||
config.sign_request(
|
||||
|
|
@ -446,11 +446,11 @@ class TestAgentCoreSearch:
|
|||
)
|
||||
mock_base_sign.assert_not_called()
|
||||
|
||||
def test_sign_request_allows_plain_http_for_localhost(self):
|
||||
def test_sign_request_allows_plain_http_for_localhost(self, monkeypatch):
|
||||
"""Local development against an MCP stub on 127.0.0.1 keeps working."""
|
||||
config = AgentCoreSearchConfig()
|
||||
os.environ["AGENTCORE_GATEWAY_TOKEN"] = "env-jwt-token"
|
||||
os.environ["AGENTCORE_GATEWAY_URL"] = "http://127.0.0.1:8931/mcp"
|
||||
monkeypatch.setenv("AGENTCORE_GATEWAY_TOKEN", "env-jwt-token")
|
||||
monkeypatch.setenv("AGENTCORE_GATEWAY_URL", "http://127.0.0.1:8931/mcp")
|
||||
try:
|
||||
headers, _ = config.sign_request(
|
||||
headers={},
|
||||
|
|
@ -483,11 +483,11 @@ class TestAgentCoreSearch:
|
|||
# AWS_BEARER_TOKEN_BEDROCK env fallback.
|
||||
assert mock_base_sign.call_args.kwargs["api_key"] == ""
|
||||
|
||||
def test_sign_request_custom_hostname_requires_region(self):
|
||||
def test_sign_request_custom_hostname_requires_region(self, monkeypatch):
|
||||
"""Custom hostname + empty AWS config chain → clear error, no guessed region."""
|
||||
config = AgentCoreSearchConfig()
|
||||
custom_url = "https://gateway.internal.example.com/mcp"
|
||||
os.environ["AGENTCORE_GATEWAY_URL"] = custom_url
|
||||
monkeypatch.setenv("AGENTCORE_GATEWAY_URL", custom_url)
|
||||
|
||||
mock_session = MagicMock()
|
||||
mock_session.region_name = None # nothing configured anywhere
|
||||
|
|
@ -503,11 +503,11 @@ class TestAgentCoreSearch:
|
|||
finally:
|
||||
os.environ.pop("AGENTCORE_GATEWAY_URL", None)
|
||||
|
||||
def test_sign_request_custom_hostname_uses_shared_config_region(self):
|
||||
def test_sign_request_custom_hostname_uses_shared_config_region(self, monkeypatch):
|
||||
"""Custom hostname + region from AWS shared config (profile) must be honored."""
|
||||
config = AgentCoreSearchConfig()
|
||||
custom_url = "https://gateway.internal.example.com/mcp"
|
||||
os.environ["AGENTCORE_GATEWAY_URL"] = custom_url
|
||||
monkeypatch.setenv("AGENTCORE_GATEWAY_URL", custom_url)
|
||||
|
||||
mock_session = MagicMock()
|
||||
mock_session.region_name = "eu-west-1" # e.g. from ~/.aws/config profile
|
||||
|
|
|
|||
|
|
@ -40,12 +40,12 @@ class TestBedrockSSLVerify:
|
|||
ssl_verify = base_aws._get_ssl_verify()
|
||||
assert ssl_verify is True
|
||||
|
||||
def test_base_aws_llm_get_ssl_verify_false(self):
|
||||
def test_base_aws_llm_get_ssl_verify_false(self, monkeypatch):
|
||||
"""Test that _get_ssl_verify returns False when SSL verification is disabled."""
|
||||
base_aws = BaseAWSLLM()
|
||||
|
||||
# Set SSL_VERIFY to False via environment
|
||||
os.environ["SSL_VERIFY"] = "False"
|
||||
monkeypatch.setenv("SSL_VERIFY", "False")
|
||||
|
||||
ssl_verify = base_aws._get_ssl_verify()
|
||||
assert ssl_verify is False
|
||||
|
|
@ -53,7 +53,7 @@ class TestBedrockSSLVerify:
|
|||
# Clean up
|
||||
os.environ.pop("SSL_VERIFY", None)
|
||||
|
||||
def test_base_aws_llm_get_ssl_verify_custom_ca_bundle(self):
|
||||
def test_base_aws_llm_get_ssl_verify_custom_ca_bundle(self, monkeypatch):
|
||||
"""Test that _get_ssl_verify returns custom CA bundle path when SSL_CERT_FILE is set."""
|
||||
base_aws = BaseAWSLLM()
|
||||
|
||||
|
|
@ -66,7 +66,7 @@ class TestBedrockSSLVerify:
|
|||
|
||||
try:
|
||||
# Set SSL_CERT_FILE environment variable
|
||||
os.environ["SSL_CERT_FILE"] = ca_bundle_path
|
||||
monkeypatch.setenv("SSL_CERT_FILE", ca_bundle_path)
|
||||
os.environ.pop("SSL_VERIFY", None)
|
||||
litellm.ssl_verify = True
|
||||
|
||||
|
|
@ -327,7 +327,7 @@ class TestBedrockSSLVerify:
|
|||
os.environ.pop("SSL_CERT_FILE", None)
|
||||
os.unlink(ca_bundle_path)
|
||||
|
||||
def test_ssl_verify_priority_env_over_litellm_config(self):
|
||||
def test_ssl_verify_priority_env_over_litellm_config(self, monkeypatch):
|
||||
"""Test that SSL_VERIFY environment variable takes priority over litellm.ssl_verify."""
|
||||
base_aws = BaseAWSLLM()
|
||||
|
||||
|
|
@ -335,7 +335,7 @@ class TestBedrockSSLVerify:
|
|||
litellm.ssl_verify = True
|
||||
|
||||
# Set SSL_VERIFY environment variable to False
|
||||
os.environ["SSL_VERIFY"] = "False"
|
||||
monkeypatch.setenv("SSL_VERIFY", "False")
|
||||
|
||||
try:
|
||||
ssl_verify = base_aws._get_ssl_verify()
|
||||
|
|
@ -345,7 +345,7 @@ class TestBedrockSSLVerify:
|
|||
os.environ.pop("SSL_VERIFY", None)
|
||||
litellm.ssl_verify = True
|
||||
|
||||
def test_ssl_cert_file_priority_over_default(self):
|
||||
def test_ssl_cert_file_priority_over_default(self, monkeypatch):
|
||||
"""Test that SSL_CERT_FILE takes priority when ssl_verify is True."""
|
||||
base_aws = BaseAWSLLM()
|
||||
|
||||
|
|
@ -358,7 +358,7 @@ class TestBedrockSSLVerify:
|
|||
|
||||
try:
|
||||
# Set SSL_CERT_FILE environment variable
|
||||
os.environ["SSL_CERT_FILE"] = ca_bundle_path
|
||||
monkeypatch.setenv("SSL_CERT_FILE", ca_bundle_path)
|
||||
os.environ.pop("SSL_VERIFY", None)
|
||||
litellm.ssl_verify = True
|
||||
|
||||
|
|
|
|||
|
|
@ -36,13 +36,6 @@ ALL_FIELDS = [
|
|||
IDENTITY = {"user_api_key_alias": "prod-key", "user_api_key_team_alias": "platform"}
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_setting():
|
||||
previous = litellm.bedrock_request_metadata_fields
|
||||
yield
|
||||
litellm.bedrock_request_metadata_fields = previous
|
||||
|
||||
|
||||
def litellm_params(metadata_key, **metadata):
|
||||
return {metadata_key: dict(metadata)}
|
||||
|
||||
|
|
@ -73,8 +66,8 @@ CONVERSE_DRIVERS = [converse_body, converse_body_async]
|
|||
|
||||
|
||||
@pytest.mark.parametrize("setting", [None, []])
|
||||
def test_feature_off_by_default_leaves_body_and_headers_untouched(setting):
|
||||
litellm.bedrock_request_metadata_fields = setting
|
||||
def test_feature_off_by_default_leaves_body_and_headers_untouched(setting, monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", setting)
|
||||
params = litellm_params("metadata", spend_logs_metadata={"team": "x"}, **IDENTITY)
|
||||
|
||||
assert "requestMetadata" not in converse_body(params)
|
||||
|
|
@ -88,18 +81,18 @@ def test_feature_off_by_default_leaves_body_and_headers_untouched(setting):
|
|||
|
||||
|
||||
@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"])
|
||||
def test_resolver_reads_both_metadata_variable_names(metadata_key):
|
||||
def test_resolver_reads_both_metadata_variable_names(metadata_key, monkeypatch: pytest.MonkeyPatch):
|
||||
"""`/v1/chat/completions` populates `metadata`; the LITELLM_METADATA_ROUTES populate
|
||||
`litellm_metadata`. Reading only one silently forwards nothing on the other route."""
|
||||
litellm.bedrock_request_metadata_fields = ALL_FIELDS
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ALL_FIELDS)
|
||||
params = litellm_params(metadata_key, spend_logs_metadata={"cost_center": "cc-1"}, **IDENTITY)
|
||||
|
||||
assert converse_body(params)["requestMetadata"] == {**IDENTITY, "cost_center": "cc-1"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("metadata_key", ["metadata", "litellm_metadata"])
|
||||
def test_invoke_messages_header_reads_both_metadata_variable_names(metadata_key):
|
||||
litellm.bedrock_request_metadata_fields = ALL_FIELDS
|
||||
def test_invoke_messages_header_reads_both_metadata_variable_names(metadata_key, monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ALL_FIELDS)
|
||||
params = litellm_params(metadata_key, **IDENTITY)
|
||||
|
||||
headers, _ = AmazonAnthropicClaudeMessagesConfig().validate_anthropic_messages_environment(
|
||||
|
|
@ -112,10 +105,15 @@ def test_invoke_messages_header_reads_both_metadata_variable_names(metadata_key)
|
|||
@pytest.mark.parametrize("reverse_client_keys", [False, True])
|
||||
@pytest.mark.parametrize("field_order", [ALL_FIELDS, list(reversed(ALL_FIELDS))])
|
||||
@pytest.mark.parametrize("client_source", ["spend_logs_metadata", "requestMetadata"])
|
||||
def test_identity_survives_a_caller_filling_every_slot(reverse_client_keys, field_order, client_source):
|
||||
def test_identity_survives_a_caller_filling_every_slot(
|
||||
reverse_client_keys,
|
||||
field_order,
|
||||
client_source,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
"""A caller sending 16 keys of its own must not evict the identity the feature exists to
|
||||
produce. Driven over every input ordering so the invariant is not an accident of one."""
|
||||
litellm.bedrock_request_metadata_fields = field_order
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", field_order)
|
||||
client_keys = [f"client_{index:02d}" for index in range(BEDROCK_REQUEST_METADATA_MAX_PAIRS)]
|
||||
client_pairs = {key: "v" for key in (reversed(client_keys) if reverse_client_keys else client_keys)}
|
||||
if client_source == "spend_logs_metadata":
|
||||
|
|
@ -141,11 +139,14 @@ def test_identity_survives_a_caller_filling_every_slot(reverse_client_keys, fiel
|
|||
["user_api_key_alias", "user_api_key_team_alias", "spend_logs_metadata", "user_api_key_team_alias"],
|
||||
],
|
||||
)
|
||||
def test_a_field_repeated_in_the_allow_list_does_not_consume_a_client_slot(field_order):
|
||||
def test_a_field_repeated_in_the_allow_list_does_not_consume_a_client_slot(
|
||||
field_order,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
"""An operator repeating a field in YAML must not inflate the reserved count and shrink the
|
||||
client budget. Asserts the client keys that should have fitted actually reach the wire, since
|
||||
asserting only that identity survives passes with or without the deduplication."""
|
||||
litellm.bedrock_request_metadata_fields = field_order
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", field_order)
|
||||
client_keys = [f"client_{index:02d}" for index in range(BEDROCK_REQUEST_METADATA_MAX_PAIRS - 1)]
|
||||
params = litellm_params("metadata", spend_logs_metadata={key: "v" for key in client_keys}, **IDENTITY)
|
||||
|
||||
|
|
@ -162,11 +163,15 @@ def test_a_field_repeated_in_the_allow_list_does_not_consume_a_client_slot(field
|
|||
"forged_key",
|
||||
["user_api_key_team_alias", "user_api_key_org_alias", "user_api_key_hash"],
|
||||
)
|
||||
def test_caller_cannot_forge_or_shadow_a_reserved_identity_key(forged_key, client_source):
|
||||
def test_caller_cannot_forge_or_shadow_a_reserved_identity_key(
|
||||
forged_key,
|
||||
client_source,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
"""`user_api_key_org_alias` and `user_api_key_hash` are names the proxy does not set here,
|
||||
so an exact-key reservation would let the forged value through under a name that reads as
|
||||
proxy-authoritative in the AWS billing record."""
|
||||
litellm.bedrock_request_metadata_fields = ALL_FIELDS
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ALL_FIELDS)
|
||||
forged = {forged_key: "attacker-controlled"}
|
||||
if client_source == "spend_logs_metadata":
|
||||
params, optional_params = litellm_params("metadata", spend_logs_metadata=forged, **IDENTITY), {}
|
||||
|
|
@ -179,10 +184,10 @@ def test_caller_cannot_forge_or_shadow_a_reserved_identity_key(forged_key, clien
|
|||
assert "attacker-controlled" not in resolved.values()
|
||||
|
||||
|
||||
def test_identity_violating_the_character_class_is_dropped_and_the_request_succeeds():
|
||||
def test_identity_violating_the_character_class_is_dropped_and_the_request_succeeds(monkeypatch: pytest.MonkeyPatch):
|
||||
"""A team alias with an apostrophe must not turn a working request into a 400 the moment
|
||||
an operator flips the setting on."""
|
||||
litellm.bedrock_request_metadata_fields = ALL_FIELDS
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ALL_FIELDS)
|
||||
params = litellm_params(
|
||||
"metadata",
|
||||
user_api_key_alias="prod-key",
|
||||
|
|
@ -196,8 +201,8 @@ def test_identity_violating_the_character_class_is_dropped_and_the_request_succe
|
|||
assert body["messages"]
|
||||
|
||||
|
||||
def test_caller_supplied_violation_still_raises_bad_request():
|
||||
litellm.bedrock_request_metadata_fields = ALL_FIELDS
|
||||
def test_caller_supplied_violation_still_raises_bad_request(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ALL_FIELDS)
|
||||
|
||||
with pytest.raises(litellm.exceptions.BadRequestError):
|
||||
converse_body(
|
||||
|
|
@ -206,34 +211,34 @@ def test_caller_supplied_violation_still_raises_bad_request():
|
|||
)
|
||||
|
||||
|
||||
def test_non_string_and_absent_identity_values_are_dropped():
|
||||
litellm.bedrock_request_metadata_fields = ALL_FIELDS + ["user_api_key_spend"]
|
||||
def test_non_string_and_absent_identity_values_are_dropped(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ALL_FIELDS + ["user_api_key_spend"])
|
||||
params = litellm_params("metadata", user_api_key_alias="prod-key", user_api_key_spend=1.25)
|
||||
|
||||
assert converse_body(params)["requestMetadata"] == {"user_api_key_alias": "prod-key"}
|
||||
|
||||
|
||||
def test_email_is_separately_opt_in():
|
||||
def test_email_is_separately_opt_in(monkeypatch: pytest.MonkeyPatch):
|
||||
"""PII crossing into CloudTrail only when the operator names the field."""
|
||||
identity_with_email = {**IDENTITY, "user_api_key_user_email": "owner@example.com"}
|
||||
litellm.bedrock_request_metadata_fields = ["user_api_key_alias", "user_api_key_team_alias"]
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ["user_api_key_alias", "user_api_key_team_alias"])
|
||||
assert (
|
||||
"user_api_key_user_email"
|
||||
not in converse_body(litellm_params("metadata", **identity_with_email))["requestMetadata"]
|
||||
)
|
||||
|
||||
litellm.bedrock_request_metadata_fields = ALL_FIELDS
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ALL_FIELDS)
|
||||
assert converse_body(litellm_params("metadata", **identity_with_email))["requestMetadata"] == identity_with_email
|
||||
|
||||
|
||||
def test_resolver_returns_none_when_nothing_survives():
|
||||
litellm.bedrock_request_metadata_fields = ALL_FIELDS
|
||||
def test_resolver_returns_none_when_nothing_survives(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ALL_FIELDS)
|
||||
assert resolve_bedrock_request_metadata(litellm_params=None) is None
|
||||
assert resolve_bedrock_request_metadata(litellm_params={"metadata": {"unrelated": "x"}}) is None
|
||||
|
||||
|
||||
def test_invoke_header_is_json_encoded_and_signed():
|
||||
litellm.bedrock_request_metadata_fields = ALL_FIELDS
|
||||
def test_invoke_header_is_json_encoded_and_signed(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ALL_FIELDS)
|
||||
params = litellm_params("metadata", spend_logs_metadata={"cost_center": "cc-1"}, **IDENTITY)
|
||||
|
||||
headers = AmazonInvokeConfig().validate_environment(
|
||||
|
|
@ -250,10 +255,10 @@ def test_invoke_header_is_json_encoded_and_signed():
|
|||
assert "anthropic-version" not in signed
|
||||
|
||||
|
||||
def test_a_caller_supplied_guardrail_header_still_wins():
|
||||
def test_a_caller_supplied_guardrail_header_still_wins(monkeypatch: pytest.MonkeyPatch):
|
||||
"""The no-displace rule is deliberate for the guardrail headers and must survive the
|
||||
request-metadata header becoming proxy-owned."""
|
||||
litellm.bedrock_request_metadata_fields = ALL_FIELDS
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ALL_FIELDS)
|
||||
|
||||
headers = AmazonInvokeConfig().validate_environment(
|
||||
headers={"X-Amzn-Bedrock-GuardrailIdentifier": "caller-set"},
|
||||
|
|
@ -318,10 +323,10 @@ def metadata_header_values(headers):
|
|||
return [value for name, value in headers.items() if name.lower() == BEDROCK_REQUEST_METADATA_HEADER.lower()]
|
||||
|
||||
|
||||
def test_converse_still_sets_the_bearer_authorization_header():
|
||||
def test_converse_still_sets_the_bearer_authorization_header(monkeypatch: pytest.MonkeyPatch):
|
||||
"""Converse owns the metadata header now, and that must not disturb the api_key path its
|
||||
validate_environment existed for. Closing the forgery hole cannot break authentication."""
|
||||
litellm.bedrock_request_metadata_fields = ALL_FIELDS
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ALL_FIELDS)
|
||||
|
||||
headers = AmazonConverseConfig().validate_environment(
|
||||
headers={},
|
||||
|
|
@ -341,11 +346,11 @@ def test_converse_still_sets_the_bearer_authorization_header():
|
|||
"caller_header_name",
|
||||
[BEDROCK_REQUEST_METADATA_HEADER, BEDROCK_REQUEST_METADATA_HEADER.lower(), "x-AMZN-bedrock-Request-METADATA"],
|
||||
)
|
||||
def test_a_caller_cannot_forge_the_request_metadata_header(driver, caller_header_name):
|
||||
def test_a_caller_cannot_forge_the_request_metadata_header(driver, caller_header_name, monkeypatch: pytest.MonkeyPatch):
|
||||
"""`extra_headers` puts caller-supplied names into the same dict the proxy merges into, so a
|
||||
deferring merge would sign the caller's forged identity into the AWS billing record. Every
|
||||
spelling must lose, or a second variant is left for the transport to choose between."""
|
||||
litellm.bedrock_request_metadata_fields = ALL_FIELDS
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ALL_FIELDS)
|
||||
|
||||
headers = driver({caller_header_name: FORGED}, litellm_params("metadata", **IDENTITY))
|
||||
|
||||
|
|
@ -355,11 +360,11 @@ def test_a_caller_cannot_forge_the_request_metadata_header(driver, caller_header
|
|||
|
||||
|
||||
@pytest.mark.parametrize("driver", HEADER_DRIVERS)
|
||||
def test_a_caller_cannot_forge_the_header_when_the_resolver_yields_nothing(driver):
|
||||
def test_a_caller_cannot_forge_the_header_when_the_resolver_yields_nothing(driver, monkeypatch: pytest.MonkeyPatch):
|
||||
"""Forwarding enabled but nothing resolvable, which a caller can arrange by supplying values
|
||||
that all fail Bedrock's rules. Owned-but-empty must mean no header on the wire, never a
|
||||
fallback to the caller's."""
|
||||
litellm.bedrock_request_metadata_fields = ALL_FIELDS
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ALL_FIELDS)
|
||||
unresolvable = litellm_params("metadata", user_api_key_alias="O'Brien's key", user_api_key_team_alias="x" * 300)
|
||||
|
||||
headers = driver({BEDROCK_REQUEST_METADATA_HEADER: FORGED}, unresolvable)
|
||||
|
|
@ -373,11 +378,15 @@ def test_a_caller_cannot_forge_the_header_when_the_resolver_yields_nothing(drive
|
|||
"forged_key",
|
||||
["user_api_key_team_alias", "user_api_key_org_alias", "user_api_key_hash"],
|
||||
)
|
||||
def test_a_caller_cannot_keep_reserved_body_keys_when_the_resolver_yields_nothing(forged_key, driver):
|
||||
def test_a_caller_cannot_keep_reserved_body_keys_when_the_resolver_yields_nothing(
|
||||
forged_key,
|
||||
driver,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
):
|
||||
"""The Converse body has the same fail-open shape as the header: with forwarding on and
|
||||
nothing resolvable, leaving the caller's `requestMetadata` in place would keep their
|
||||
reserved-prefix keys on the wire. Owned-but-empty must remove the field outright."""
|
||||
litellm.bedrock_request_metadata_fields = ALL_FIELDS
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ALL_FIELDS)
|
||||
|
||||
body = driver(litellm_params("metadata"), {"requestMetadata": {forged_key: "FORGED"}})
|
||||
|
||||
|
|
@ -386,10 +395,10 @@ def test_a_caller_cannot_keep_reserved_body_keys_when_the_resolver_yields_nothin
|
|||
|
||||
|
||||
@pytest.mark.parametrize("driver", CONVERSE_DRIVERS)
|
||||
def test_benign_caller_body_metadata_still_survives_when_no_identity_resolves(driver):
|
||||
def test_benign_caller_body_metadata_still_survives_when_no_identity_resolves(driver, monkeypatch: pytest.MonkeyPatch):
|
||||
"""Removing the field must be scoped to the reserved keys being the only thing left, not a
|
||||
blanket drop of the caller's own attribution pairs."""
|
||||
litellm.bedrock_request_metadata_fields = ALL_FIELDS
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ALL_FIELDS)
|
||||
|
||||
body = driver(
|
||||
litellm_params("metadata"),
|
||||
|
|
@ -400,10 +409,10 @@ def test_benign_caller_body_metadata_still_survives_when_no_identity_resolves(dr
|
|||
|
||||
|
||||
@pytest.mark.parametrize("driver", CONVERSE_DRIVERS)
|
||||
def test_caller_body_metadata_is_left_alone_when_forwarding_is_off(driver):
|
||||
def test_caller_body_metadata_is_left_alone_when_forwarding_is_off(driver, monkeypatch: pytest.MonkeyPatch):
|
||||
"""With the feature off the proxy does not own the field, so the pre-existing pass-through
|
||||
behaviour for a caller-supplied `requestMetadata` must be unchanged."""
|
||||
litellm.bedrock_request_metadata_fields = None
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", None)
|
||||
caller_supplied = {"user_api_key_team_alias": "caller-set", "cost_center": "cc-9"}
|
||||
|
||||
body = driver(litellm_params("metadata", **IDENTITY), {"requestMetadata": caller_supplied})
|
||||
|
|
@ -412,10 +421,10 @@ def test_caller_body_metadata_is_left_alone_when_forwarding_is_off(driver):
|
|||
|
||||
|
||||
@pytest.mark.parametrize("driver", HEADER_DRIVERS)
|
||||
def test_a_caller_header_is_left_alone_when_forwarding_is_off(driver):
|
||||
def test_a_caller_header_is_left_alone_when_forwarding_is_off(driver, monkeypatch: pytest.MonkeyPatch):
|
||||
"""The proxy only claims the name when the operator turned forwarding on; with the feature
|
||||
off this is an ordinary passthrough header and stripping it would be a regression."""
|
||||
litellm.bedrock_request_metadata_fields = None
|
||||
monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", None)
|
||||
|
||||
headers = driver({BEDROCK_REQUEST_METADATA_HEADER: FORGED}, litellm_params("metadata", **IDENTITY))
|
||||
|
||||
|
|
|
|||
|
|
@ -105,14 +105,14 @@ def test_crusoe_provider_detection_by_prefix():
|
|||
assert model == "meta-llama/Llama-3.3-70B-Instruct"
|
||||
|
||||
|
||||
def test_crusoe_model_list_populated():
|
||||
def test_crusoe_model_list_populated(monkeypatch):
|
||||
"""Test Crusoe models are present in model_prices_and_context_window.json"""
|
||||
import litellm
|
||||
|
||||
original_model_cost = litellm.model_cost
|
||||
original_env = os.environ.get("LITELLM_LOCAL_MODEL_COST_MAP")
|
||||
try:
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
expected = [
|
||||
|
|
@ -132,4 +132,4 @@ def test_crusoe_model_list_populated():
|
|||
if original_env is None:
|
||||
os.environ.pop("LITELLM_LOCAL_MODEL_COST_MAP", None)
|
||||
else:
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = original_env
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", original_env)
|
||||
|
|
|
|||
|
|
@ -24,7 +24,10 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
|||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import (
|
||||
BaseLLMHTTPHandler,
|
||||
_collect_ws_project_quota_callbacks,
|
||||
_google_genai_streaming_hidden_params,
|
||||
_has_pre_call_deployment_hook,
|
||||
_rust_responses_websocket_enabled,
|
||||
)
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
|
@ -2445,3 +2448,82 @@ async def test_generic_http_handler_async_streaming_forwards_provider_response_h
|
|||
|
||||
collected = [chunk async for chunk in response]
|
||||
assert "".join([chunk.choices[0].delta.content or "" for chunk in collected]) == "hi"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"custom_llm_provider, litellm_params, expected",
|
||||
[
|
||||
("openai", GenericLiteLLMParams(rust=True), True),
|
||||
("openai", GenericLiteLLMParams(), False),
|
||||
("openai", GenericLiteLLMParams(rust=False), False),
|
||||
("azure", GenericLiteLLMParams(rust=True), False),
|
||||
("hosted_vllm", GenericLiteLLMParams(rust=True), False),
|
||||
(None, GenericLiteLLMParams(rust=True), False),
|
||||
],
|
||||
)
|
||||
def test_the_rust_responses_websocket_needs_both_openai_and_the_rust_flag(
|
||||
custom_llm_provider, litellm_params, expected
|
||||
):
|
||||
assert _rust_responses_websocket_enabled(custom_llm_provider, litellm_params) is expected
|
||||
|
||||
|
||||
def test_a_plain_callback_does_not_advertise_a_pre_call_deployment_hook(monkeypatch):
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
class _PlainLogger(CustomLogger):
|
||||
pass
|
||||
|
||||
logging_obj = Mock()
|
||||
logging_obj.dynamic_success_callbacks = []
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [])
|
||||
assert _has_pre_call_deployment_hook(logging_obj) is False
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [_PlainLogger()])
|
||||
assert _has_pre_call_deployment_hook(logging_obj) is False
|
||||
|
||||
|
||||
def test_a_callback_that_overrides_the_deployment_hook_is_detected(monkeypatch):
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
class _DeploymentHookLogger(CustomLogger):
|
||||
async def async_pre_call_deployment_hook(self, kwargs, call_type):
|
||||
return None
|
||||
|
||||
class _InheritsTheHook(_DeploymentHookLogger):
|
||||
pass
|
||||
|
||||
logging_obj = Mock()
|
||||
logging_obj.dynamic_success_callbacks = []
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [_DeploymentHookLogger()])
|
||||
assert _has_pre_call_deployment_hook(logging_obj) is True
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [_InheritsTheHook()])
|
||||
assert _has_pre_call_deployment_hook(logging_obj) is True
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [])
|
||||
logging_obj.dynamic_success_callbacks = [_DeploymentHookLogger()]
|
||||
assert _has_pre_call_deployment_hook(logging_obj) is True
|
||||
|
||||
|
||||
def test_only_callbacks_that_can_charge_a_frame_are_collected_for_ws_quota(monkeypatch):
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
class _PlainLogger(CustomLogger):
|
||||
pass
|
||||
|
||||
class _QuotaLogger(CustomLogger):
|
||||
async def enforce_project_io_token_quota_for_frame(self, *args, **kwargs):
|
||||
return None
|
||||
|
||||
class _NotCallableAttribute:
|
||||
enforce_project_io_token_quota_for_frame = "not a method"
|
||||
|
||||
plain, quota, decoy = _PlainLogger(), _QuotaLogger(), _NotCallableAttribute()
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [plain, decoy])
|
||||
assert _collect_ws_project_quota_callbacks() == ()
|
||||
|
||||
monkeypatch.setattr(litellm, "callbacks", [plain, quota, decoy])
|
||||
assert _collect_ws_project_quota_callbacks() == (quota,)
|
||||
|
|
|
|||
|
|
@ -83,8 +83,8 @@ class TestDataRobotConfig:
|
|||
== api_base
|
||||
)
|
||||
|
||||
def test_resolve_api_base_with_environment_variable(self, handler):
|
||||
os.environ["DATAROBOT_ENDPOINT"] = "https://env.datarobot.com"
|
||||
def test_resolve_api_base_with_environment_variable(self, handler, monkeypatch):
|
||||
monkeypatch.setenv("DATAROBOT_ENDPOINT", "https://env.datarobot.com")
|
||||
assert (
|
||||
handler._resolve_api_base(None)
|
||||
== "https://env.datarobot.com/api/v2/genai/llmgw/chat/completions/"
|
||||
|
|
@ -101,7 +101,7 @@ class TestDataRobotConfig:
|
|||
def test_resolve_api_key(self, api_key, expected_api_key, handler):
|
||||
assert handler._resolve_api_key(api_key) == expected_api_key
|
||||
|
||||
def test_resolve_api_key_with_environment_variable(self, handler):
|
||||
os.environ["DATAROBOT_API_TOKEN"] = "env_key"
|
||||
def test_resolve_api_key_with_environment_variable(self, handler, monkeypatch):
|
||||
monkeypatch.setenv("DATAROBOT_API_TOKEN", "env_key")
|
||||
assert handler._resolve_api_key(None) == "env_key"
|
||||
del os.environ["DATAROBOT_API_TOKEN"]
|
||||
|
|
|
|||
|
|
@ -11,14 +11,14 @@ sys.path.insert(0, os.path.abspath("../../../.."))
|
|||
import litellm
|
||||
|
||||
|
||||
def test_deepseek_supported_openai_params():
|
||||
def test_deepseek_supported_openai_params(monkeypatch):
|
||||
"""
|
||||
Test "reasoning_effort" is an openai param supported for the DeepSeek model on deepinfra
|
||||
"""
|
||||
from litellm.llms.deepinfra.chat.transformation import DeepInfraConfig
|
||||
|
||||
# Ensure we're using the local model cost map
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
supported_openai_params = DeepInfraConfig().get_supported_openai_params(
|
||||
|
|
|
|||
|
|
@ -81,8 +81,8 @@ def test_no_usage_details():
|
|||
assert cost == 0.0
|
||||
|
||||
|
||||
def test_gemini_image_edit_cost_prefers_token_usage_metadata():
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
def test_gemini_image_edit_cost_prefers_token_usage_metadata(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
model = "gemini/gemini-3-pro-image-preview"
|
||||
model_info = litellm.get_model_info(model=model, custom_llm_provider="gemini")
|
||||
|
|
@ -120,8 +120,8 @@ def test_gemini_image_edit_cost_prefers_token_usage_metadata():
|
|||
assert cost != flat_image_cost
|
||||
|
||||
|
||||
def test_gemini_image_edit_cost_uses_output_token_details():
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
def test_gemini_image_edit_cost_uses_output_token_details(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
model = "gemini/gemini-3-pro-image-preview"
|
||||
model_info = litellm.get_model_info(model=model, custom_llm_provider="gemini")
|
||||
|
|
@ -176,8 +176,8 @@ def test_gemini_image_edit_cost_uses_output_token_details():
|
|||
assert cost != all_output_as_image_cost
|
||||
|
||||
|
||||
def test_gemini_image_generation_cost_uses_output_token_details():
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
def test_gemini_image_generation_cost_uses_output_token_details(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
model = "gemini/gemini-3-pro-image-preview"
|
||||
model_info = litellm.get_model_info(model=model, custom_llm_provider="gemini")
|
||||
|
|
@ -232,8 +232,8 @@ def test_gemini_image_generation_cost_uses_output_token_details():
|
|||
assert cost != all_output_as_image_cost
|
||||
|
||||
|
||||
def test_gemini_image_edit_cost_falls_back_to_flat_image_pricing():
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
def test_gemini_image_edit_cost_falls_back_to_flat_image_pricing(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
model = "gemini/gemini-3-pro-image-preview"
|
||||
model_info = litellm.get_model_info(model=model, custom_llm_provider="gemini")
|
||||
|
|
@ -264,8 +264,8 @@ def _image_response_with_web_search(web_search_requests):
|
|||
return ImageResponse(data=[ImageObject(b64_json="img1")], usage=usage)
|
||||
|
||||
|
||||
def test_gemini_image_generation_cost_adds_web_search_grounding():
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
def test_gemini_image_generation_cost_adds_web_search_grounding(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
model = "gemini/gemini-3-pro-image-preview"
|
||||
model_info = litellm.get_model_info(model=model, custom_llm_provider="gemini")
|
||||
|
|
@ -286,8 +286,8 @@ def test_gemini_image_generation_cost_adds_web_search_grounding():
|
|||
assert round(grounded - ungrounded, 10) == round(expected_web_search_cost, 10)
|
||||
|
||||
|
||||
def test_gemini_image_generation_cost_no_web_search_when_absent():
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
def test_gemini_image_generation_cost_no_web_search_when_absent(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
model = "gemini/gemini-3-pro-image-preview"
|
||||
|
||||
|
|
|
|||
|
|
@ -231,10 +231,10 @@ def test_inception_in_provider_lists():
|
|||
assert "https://api.inceptionlabs.ai/v1" in litellm.openai_compatible_endpoints
|
||||
|
||||
|
||||
def test_inception_model_configuration():
|
||||
def test_inception_model_configuration(monkeypatch):
|
||||
from litellm import get_model_info
|
||||
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
litellm.inception_models = set()
|
||||
litellm.add_known_models()
|
||||
|
|
@ -251,8 +251,8 @@ def test_inception_model_configuration():
|
|||
assert info.get("supports_response_schema") is True
|
||||
|
||||
|
||||
def test_inception_model_list_populated():
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
def test_inception_model_list_populated(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
litellm.inception_models = set()
|
||||
litellm.add_known_models()
|
||||
|
|
|
|||
|
|
@ -143,10 +143,10 @@ async def test_inception_fim_async():
|
|||
assert r.choices[0].text == "a + b"
|
||||
|
||||
|
||||
def test_inception_fim_model_configuration():
|
||||
def test_inception_fim_model_configuration(monkeypatch):
|
||||
from litellm import get_model_info
|
||||
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
litellm.text_completion_inception_models = set()
|
||||
litellm.add_known_models()
|
||||
|
|
|
|||
|
|
@ -30,8 +30,8 @@ def _image_response_with_web_search(web_search_requests):
|
|||
return ImageResponse(data=[ImageObject(b64_json="img1")], usage=usage)
|
||||
|
||||
|
||||
def test_vertex_image_generation_cost_adds_web_search_grounding():
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
def test_vertex_image_generation_cost_adds_web_search_grounding(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
model = "gemini-3-pro-image-preview"
|
||||
model_info = litellm.get_model_info(model=model, custom_llm_provider="vertex_ai")
|
||||
|
|
@ -55,8 +55,8 @@ def test_vertex_image_generation_cost_adds_web_search_grounding():
|
|||
assert round(grounded - ungrounded, 10) == round(expected_web_search_cost, 10)
|
||||
|
||||
|
||||
def test_vertex_image_generation_cost_no_web_search_when_absent():
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
def test_vertex_image_generation_cost_no_web_search_when_absent(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
model = "gemini-3-pro-image-preview"
|
||||
|
||||
|
|
|
|||
|
|
@ -51,11 +51,11 @@ def test_zai_in_provider_lists():
|
|||
assert "zai" in litellm.provider_list
|
||||
|
||||
|
||||
def test_zai_models_in_model_cost():
|
||||
def test_zai_models_in_model_cost(monkeypatch):
|
||||
"""Test that ZAI models are in the model cost map"""
|
||||
import os
|
||||
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
zai_models = [
|
||||
|
|
@ -75,11 +75,11 @@ def test_zai_models_in_model_cost():
|
|||
assert litellm.model_cost[model]["litellm_provider"] == "zai"
|
||||
|
||||
|
||||
def test_zai_glm46_cost_calculation():
|
||||
def test_zai_glm46_cost_calculation(monkeypatch):
|
||||
"""Test the cost calculation for glm-4.6"""
|
||||
import os
|
||||
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
key = "zai/glm-4.6"
|
||||
|
|
@ -96,11 +96,11 @@ def test_zai_glm46_cost_calculation():
|
|||
assert math.isclose(completion_cost, 2.2, rel_tol=1e-6)
|
||||
|
||||
|
||||
def test_zai_flash_model_is_free():
|
||||
def test_zai_flash_model_is_free(monkeypatch):
|
||||
"""Test that glm-4.5-flash has zero cost"""
|
||||
import os
|
||||
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
key = "zai/glm-4.5-flash"
|
||||
|
|
@ -110,11 +110,11 @@ def test_zai_flash_model_is_free():
|
|||
assert info["output_cost_per_token"] == 0
|
||||
|
||||
|
||||
def test_glm47_supports_reasoning():
|
||||
def test_glm47_supports_reasoning(monkeypatch):
|
||||
"""Test that GLM-4.7 supports reasoning"""
|
||||
import os
|
||||
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
key = "zai/glm-4.7"
|
||||
|
|
@ -124,11 +124,11 @@ def test_glm47_supports_reasoning():
|
|||
assert info["supports_reasoning"] is True
|
||||
|
||||
|
||||
def test_glm47_cost_calculation():
|
||||
def test_glm47_cost_calculation(monkeypatch):
|
||||
"""Test cost calculation for GLM-4.7"""
|
||||
import os
|
||||
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
prompt_cost, completion_cost = cost_per_token(
|
||||
|
|
|
|||
|
|
@ -6430,9 +6430,7 @@ class TestAggregateGatewayDcrChallenge:
|
|||
assert _gateway_dcr_challenge_target("/mcp/srv", None, None) == expected, resolved
|
||||
assert _gateway_dcr_challenge_target("/mcp/a,b", None, None) is None
|
||||
assert _gateway_dcr_challenge_target("/mcp", None, None) is None
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager"
|
||||
) as mock_mgr:
|
||||
with patch("litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager") as mock_mgr:
|
||||
mock_mgr.get_mcp_server_by_name.return_value = _server(MCPAuth.oauth2)
|
||||
assert _gateway_dcr_challenge_target("/mcp/srv", ["other"], None) is None
|
||||
|
||||
|
|
@ -7073,25 +7071,109 @@ class TestUserSubjectTeamUnion:
|
|||
assert await manager.operator_open_server_ids(admitted) == {"srv-byom"}
|
||||
assert await manager.operator_open_server_ids(scoped_key) == set(), "explicit key scope still suppresses BYOM"
|
||||
|
||||
async def test_admitted_admin_is_scoped_to_grants_not_full_registry(self):
|
||||
"""The wrapper's admin short-circuit hands the FULL registry to any admin-role auth before
|
||||
the grant union or the per-team org ceilings run. A session bearer is a third-party client
|
||||
credential, not the dashboard: an admin signing in through the connect flow gets their
|
||||
grants like anyone else. A real admin key keeps the dashboard behavior unchanged."""
|
||||
@pytest.mark.parametrize(
|
||||
"role", ["PROXY_ADMIN", "PROXY_ADMIN_VIEW_ONLY"], ids=["proxy_admin", "proxy_admin_view_only"]
|
||||
)
|
||||
async def test_admitted_admin_gets_registry_like_an_admin_key(self, role):
|
||||
"""Connect-page parity: admin view rides the HUMAN, not the credential. An admitted session
|
||||
subject with an admin-view role resolves the same full registry an admin KEY does, so the
|
||||
servers the dashboard shows an admin are the servers their OAuth session serves. Regression
|
||||
pin for the customer report where an admin's Claude Code session showed zero tools."""
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
manager = self._manager_with(["srv-granted", "srv-secret"])
|
||||
admitted = _make_admitted_subject("admin-user")
|
||||
admitted.user_role = LitellmUserRoles[role]
|
||||
key_admin = UserAPIKeyAuth(user_id="admin-user", api_key="sk-hash", user_role=LitellmUserRoles[role])
|
||||
with patch.object(MCPRequestHandler, "get_allowed_mcp_servers", AsyncMock(return_value=["srv-granted"])):
|
||||
admitted_view = set(await manager.get_allowed_mcp_servers(admitted))
|
||||
key_admin_view = set(await manager.get_allowed_mcp_servers(key_admin))
|
||||
assert admitted_view == {"srv-granted", "srv-secret"}, "an admitted admin resolves the registry"
|
||||
assert key_admin_view == admitted_view, "session and key admin views must be identical"
|
||||
|
||||
async def test_admitted_admin_explicit_scope_still_wins(self):
|
||||
"""An admin whose own user row names servers is entitlement-bound whatever their role: the
|
||||
row binds through the ceiling for an admitted subject (a user row's mcp_servers is the
|
||||
human's grant list, not a credential scope), so the registry seed must not fire. A KEY
|
||||
carrying an explicit scope disqualifies directly, empty list included."""
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LitellmUserRoles
|
||||
|
||||
manager = self._manager_with(["srv-granted", "srv-secret"])
|
||||
admitted = _make_admitted_subject("admin-user", own_servers=["srv-granted"])
|
||||
admitted.user_role = LitellmUserRoles.PROXY_ADMIN
|
||||
with (
|
||||
patch.object(
|
||||
MCPRequestHandler, "_get_allowed_mcp_servers_for_user", AsyncMock(return_value=["srv-granted"])
|
||||
),
|
||||
patch.object(MCPRequestHandler, "get_allowed_mcp_servers", AsyncMock(return_value=["srv-granted"])),
|
||||
):
|
||||
assert set(await manager.get_allowed_mcp_servers(admitted)) == {"srv-granted"}
|
||||
|
||||
scoped_key = UserAPIKeyAuth(
|
||||
user_id="admin-user",
|
||||
api_key="sk-hash",
|
||||
user_role=LitellmUserRoles.PROXY_ADMIN,
|
||||
object_permission=LiteLLM_ObjectPermissionTable(object_permission_id="op-k", mcp_servers=[]),
|
||||
)
|
||||
with patch.object(MCPRequestHandler, "get_allowed_mcp_servers", AsyncMock(return_value=[])):
|
||||
assert await manager.get_allowed_mcp_servers(scoped_key) == []
|
||||
|
||||
async def test_admitted_admin_db_default_empty_scope_still_gets_registry(self):
|
||||
"""The admitted subject's object_permission is the user's own row, whose mcp_servers column
|
||||
is [] by DB default whenever the row exists for any other field: default noise, never an
|
||||
explicit scope. The registry seed must fire through it, or every admin with a shared
|
||||
permission row keeps resolving zero servers while their dashboard shows all of them."""
|
||||
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LitellmUserRoles
|
||||
|
||||
manager = self._manager_with(["srv-granted", "srv-secret"])
|
||||
admitted = _make_admitted_subject("admin-user")
|
||||
admitted.user_role = LitellmUserRoles.PROXY_ADMIN
|
||||
admitted.object_permission = LiteLLM_ObjectPermissionTable(object_permission_id="op-u", mcp_servers=[])
|
||||
with patch.object(MCPRequestHandler, "get_allowed_mcp_servers", AsyncMock(return_value=[])):
|
||||
assert set(await manager.get_allowed_mcp_servers(admitted)) == {"srv-granted", "srv-secret"}
|
||||
|
||||
async def test_non_admin_admitted_subject_never_gets_registry(self):
|
||||
"""The negative control for the registry seed: a plain admitted subject with no admin-view
|
||||
role resolves only their grant union, however many servers the registry holds."""
|
||||
manager = self._manager_with(["srv-granted", "srv-secret"])
|
||||
plain = _make_admitted_subject("plain-user")
|
||||
with patch.object(MCPRequestHandler, "get_allowed_mcp_servers", AsyncMock(return_value=["srv-granted"])):
|
||||
assert set(await manager.get_allowed_mcp_servers(plain)) == {"srv-granted"}
|
||||
|
||||
async def test_admitted_admin_entitlement_ceiling_disables_registry(self):
|
||||
"""An entitlement ceiling, including an UNRESOLVED one, binds the human whatever their role:
|
||||
the registry seed must not fire on a transient fault, and the grant union answers instead."""
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
manager = self._manager_with(["srv-granted", "srv-secret"])
|
||||
admitted = _make_admitted_subject("admin-user")
|
||||
admitted.user_role = LitellmUserRoles.PROXY_ADMIN
|
||||
with patch.object(MCPRequestHandler, "get_allowed_mcp_servers", AsyncMock(return_value=["srv-granted"])):
|
||||
admitted_view = set(await manager.get_allowed_mcp_servers(admitted))
|
||||
key_admin_view = set(
|
||||
await manager.get_allowed_mcp_servers(
|
||||
UserAPIKeyAuth(user_id="admin-user", api_key="sk-hash", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
)
|
||||
)
|
||||
assert admitted_view == {"srv-granted"}, "an admitted admin gets their grants, not the registry"
|
||||
assert key_admin_view == {"srv-granted", "srv-secret"}, "admin KEY behavior must be unchanged"
|
||||
with (
|
||||
patch.object(MCPRequestHandler, "_get_allowed_mcp_servers_for_user", AsyncMock(return_value=None)),
|
||||
patch.object(MCPRequestHandler, "get_allowed_mcp_servers", AsyncMock(return_value=["srv-granted"])),
|
||||
):
|
||||
assert set(await manager.get_allowed_mcp_servers(admitted)) == {"srv-granted"}
|
||||
|
||||
async def test_admitted_admin_tools_ride_own_source_on_ungranted_server(self):
|
||||
"""Admin view is an open channel on the tools axis too: the user's OWN source resolves the
|
||||
tools for a server no grant names, so an admin session's registry-wide servers are invokable
|
||||
rather than listable-but-uninvokable. A non-admin subject on the same server stays denied.
|
||||
An admin whose row carries any entitlement never reaches this channel: the ceiling clause
|
||||
disqualifies the predicate first, so their own tool permissions keep binding on the grants path."""
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
|
||||
admin = _make_admitted_subject("admin-user")
|
||||
admin.user_role = LitellmUserRoles.PROXY_ADMIN
|
||||
plain = _make_admitted_subject("plain-user")
|
||||
with self._patch(teams_by_id={}, user_teams=[]):
|
||||
with patch(
|
||||
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager.operator_open_server_ids",
|
||||
AsyncMock(return_value=set()),
|
||||
):
|
||||
admin_tools = await MCPRequestHandler.get_allowed_tools_for_server("srv-any", admin)
|
||||
plain_tools = await MCPRequestHandler.get_allowed_tools_for_server("srv-any", plain)
|
||||
assert admin_tools is None, "admin channel resolves allow-all through the user's own source"
|
||||
assert plain_tools == [], "a non-admin subject with no granting source stays denied"
|
||||
|
||||
async def test_admitted_opt_out_via_wrapper_keeps_team_servers(self):
|
||||
"""The wrapper's no_mcp_servers early-return is a KEY rule (a scoped credential's opt-out is
|
||||
|
|
|
|||
|
|
@ -2957,6 +2957,55 @@ def test_validate_trusted_redirect_uri_logs_diagnostic_on_rejection(caplog, monk
|
|||
assert "X-Forwarded-Host" in msg
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"direct_ip,expect_accepted",
|
||||
[
|
||||
("10.0.0.7", True),
|
||||
("203.0.113.5", False),
|
||||
],
|
||||
)
|
||||
def test_validate_trusted_redirect_uri_follows_the_xff_trust_gate(direct_ip, expect_accepted, monkeypatch):
|
||||
try:
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.oauth_utils import (
|
||||
validate_trusted_redirect_uri,
|
||||
)
|
||||
except ImportError:
|
||||
pytest.skip("MCP oauth_utils not available")
|
||||
|
||||
monkeypatch.delenv("PROXY_BASE_URL", raising=False)
|
||||
monkeypatch.delenv("MCP_TRUSTED_REDIRECT_ORIGINS", raising=False)
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.base_url = "http://localhost:4000/"
|
||||
mock_request.client = MagicMock()
|
||||
mock_request.client.host = direct_ip
|
||||
|
||||
headers = {
|
||||
"X-Forwarded-Proto": "https",
|
||||
"X-Forwarded-Host": "proxy.example.com",
|
||||
}
|
||||
mock_request.headers.get = lambda name, default=None: headers.get(name, default)
|
||||
mock_request.headers.__contains__ = lambda self_, name: name in headers
|
||||
|
||||
redirect_uri = "https://proxy.example.com/callback"
|
||||
general_settings = {
|
||||
"use_x_forwarded_for": True,
|
||||
"mcp_trusted_proxy_ranges": ["10.0.0.0/8"],
|
||||
}
|
||||
|
||||
with patch("litellm.proxy.proxy_server.general_settings", general_settings, create=True):
|
||||
if expect_accepted:
|
||||
validate_trusted_redirect_uri(mock_request, redirect_uri)
|
||||
return
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
validate_trusted_redirect_uri(mock_request, redirect_uri)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
assert "proxy.example.com" in str(exc_info.value.detail)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bad_value",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -4846,7 +4846,9 @@ class TestMCPServerManager:
|
|||
@staticmethod
|
||||
def _manager_with_deepwiki_and_huggingface() -> MCPServerManager:
|
||||
manager = MCPServerManager()
|
||||
deepwiki = MCPServer(server_id="deepwiki-id", name="deepwiki", server_name="deepwiki", transport=MCPTransport.http)
|
||||
deepwiki = MCPServer(
|
||||
server_id="deepwiki-id", name="deepwiki", server_name="deepwiki", transport=MCPTransport.http
|
||||
)
|
||||
huggingface = MCPServer(
|
||||
server_id="huggingface-id", name="huggingface", server_name="huggingface", transport=MCPTransport.http
|
||||
)
|
||||
|
|
@ -4867,8 +4869,14 @@ class TestMCPServerManager:
|
|||
with pytest.raises(ValueError, match="Tool hub_repo_search not found"):
|
||||
manager._resolve_mcp_server_for_tool_call("deepwiki", "hub_repo_search")
|
||||
|
||||
assert manager._resolve_mcp_server_for_tool_call("deepwiki", "read_wiki_structure") is manager.registry["deepwiki-id"]
|
||||
assert manager._resolve_mcp_server_for_tool_call("huggingface", "hub_repo_search") is manager.registry["huggingface-id"]
|
||||
assert (
|
||||
manager._resolve_mcp_server_for_tool_call("deepwiki", "read_wiki_structure")
|
||||
is manager.registry["deepwiki-id"]
|
||||
)
|
||||
assert (
|
||||
manager._resolve_mcp_server_for_tool_call("huggingface", "hub_repo_search")
|
||||
is manager.registry["huggingface-id"]
|
||||
)
|
||||
|
||||
def test_get_mcp_server_from_tool_name_rejects_other_servers_prefix(self):
|
||||
manager = self._manager_with_deepwiki_and_huggingface()
|
||||
|
|
@ -4876,7 +4884,9 @@ class TestMCPServerManager:
|
|||
assert manager._get_mcp_server_from_tool_name("huggingface-read_wiki_structure") is None
|
||||
assert manager._get_mcp_server_from_tool_name("deepwiki-hub_repo_search") is None
|
||||
assert manager._get_mcp_server_from_tool_name("deepwiki-read_wiki_structure") is manager.registry["deepwiki-id"]
|
||||
assert manager._get_mcp_server_from_tool_name("huggingface-hub_repo_search") is manager.registry["huggingface-id"]
|
||||
assert (
|
||||
manager._get_mcp_server_from_tool_name("huggingface-hub_repo_search") is manager.registry["huggingface-id"]
|
||||
)
|
||||
|
||||
def test_resolve_mcp_server_for_tool_call_shared_bare_name_resolves_via_own_prefixed_spelling(self):
|
||||
manager = MCPServerManager()
|
||||
|
|
@ -10500,6 +10510,39 @@ class TestSessionResourceScopeIntersect:
|
|||
|
||||
assert MCPServerManager._admitted_session_resource_scope(self._admitted_auth("b")) == "b"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_admin_registry_seed_still_bounded_by_session_resource_scope(self):
|
||||
"""The admin-view registry seed flows through the same scoped exit as every union: a
|
||||
session envelope sealed to one server never widens past it, even held by an admin whose
|
||||
role resolves the whole registry. Pin for the connect-page-parity change; without the
|
||||
single-exit shape, the old early return would hand a per-server bearer the registry."""
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
from litellm.types.mcp import MCPTransport
|
||||
from litellm.types.mcp_server.mcp_server_manager import MCPServer
|
||||
|
||||
manager = MCPServerManager()
|
||||
for sid in ("granted-id", "other-id"):
|
||||
manager.registry[sid] = MCPServer(
|
||||
server_id=sid, name=sid, server_name=sid, url="https://example.com/mcp", transport=MCPTransport.http
|
||||
)
|
||||
auth = self._admitted_auth("granted-id")
|
||||
auth.user_role = LitellmUserRoles.PROXY_ADMIN
|
||||
with (
|
||||
patch.object(MCPServerManager, "get_allow_all_keys_server_ids", return_value=[]),
|
||||
patch.object(
|
||||
MCPServerManager,
|
||||
"_get_active_submitted_mcp_server_ids_for_user",
|
||||
new_callable=AsyncMock,
|
||||
return_value=[],
|
||||
),
|
||||
):
|
||||
assert await manager.get_allowed_mcp_servers(auth) == ["granted-id"]
|
||||
auth.mcp_session_resource_server_id = None
|
||||
assert set(await manager.get_allowed_mcp_servers(auth)) == {"granted-id", "other-id"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_allowed_mcp_servers_scopes_past_operator_open_union(self):
|
||||
"""The intersect applies AFTER the operator-open (allow_all_keys) union, so a scoped
|
||||
|
|
@ -10519,7 +10562,12 @@ class TestSessionResourceScopeIntersect:
|
|||
new_callable=AsyncMock,
|
||||
return_value=["granted-id", "other-id"],
|
||||
),
|
||||
patch.object(MCPServerManager, "_get_active_submitted_mcp_server_ids_for_user", new_callable=AsyncMock, return_value=[]),
|
||||
patch.object(
|
||||
MCPServerManager,
|
||||
"_get_active_submitted_mcp_server_ids_for_user",
|
||||
new_callable=AsyncMock,
|
||||
return_value=[],
|
||||
),
|
||||
):
|
||||
allowed = await manager.get_allowed_mcp_servers(auth)
|
||||
assert allowed == ["granted-id"]
|
||||
|
|
@ -10531,7 +10579,12 @@ class TestSessionResourceScopeIntersect:
|
|||
new_callable=AsyncMock,
|
||||
side_effect=RuntimeError("resolver down"),
|
||||
),
|
||||
patch.object(MCPServerManager, "_get_active_submitted_mcp_server_ids_for_user", new_callable=AsyncMock, return_value=[]),
|
||||
patch.object(
|
||||
MCPServerManager,
|
||||
"_get_active_submitted_mcp_server_ids_for_user",
|
||||
new_callable=AsyncMock,
|
||||
return_value=[],
|
||||
),
|
||||
):
|
||||
fallback = await manager.get_allowed_mcp_servers(auth)
|
||||
assert fallback == ["granted-id"]
|
||||
|
|
@ -10645,9 +10698,7 @@ class TestClientForwardedDiscoveryFailureIsNotFatal:
|
|||
|
||||
@pytest.mark.parametrize("auth_type", [MCPAuth.true_passthrough, MCPAuth.oauth_delegate])
|
||||
@pytest.mark.asyncio
|
||||
async def test_client_forwarded_servers_keep_discovering_their_front_door_endpoints(
|
||||
self, auth_type: MCPAuthType
|
||||
):
|
||||
async def test_client_forwarded_servers_keep_discovering_their_front_door_endpoints(self, auth_type: MCPAuthType):
|
||||
"""Exempting these modes from the FAILURE must not exempt them from discovery itself.
|
||||
|
||||
``/authorize``, ``/token`` and ``/register`` read the discovered endpoints for these servers
|
||||
|
|
|
|||
|
|
@ -109,7 +109,7 @@ async def test_authenticate_user_admin_login_with_ui_credentials():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_user_admin_login_with_master_key_as_password():
|
||||
async def test_authenticate_user_admin_login_with_master_key_as_password(monkeypatch):
|
||||
"""Test admin login when UI_PASSWORD is not set, should use master_key"""
|
||||
master_key = "sk-1234"
|
||||
ui_username = "admin"
|
||||
|
|
@ -131,39 +131,35 @@ async def test_authenticate_user_admin_login_with_master_key_as_password():
|
|||
|
||||
with patch.dict(os.environ, env_vars, clear=False):
|
||||
# Explicitly remove UI_PASSWORD if it exists
|
||||
original_ui_password = os.environ.pop("UI_PASSWORD", None)
|
||||
try:
|
||||
monkeypatch.delenv("UI_PASSWORD", raising=False)
|
||||
with patch(
|
||||
"litellm.proxy.auth.login_utils.generate_key_helper_fn",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_generate_key:
|
||||
mock_generate_key.return_value = {
|
||||
"token": "test-token-123",
|
||||
"user_id": LITELLM_PROXY_ADMIN_NAME,
|
||||
}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.login_utils.generate_key_helper_fn",
|
||||
"litellm.proxy.auth.login_utils.user_update",
|
||||
new_callable=AsyncMock,
|
||||
) as mock_generate_key:
|
||||
mock_generate_key.return_value = {
|
||||
"token": "test-token-123",
|
||||
"user_id": LITELLM_PROXY_ADMIN_NAME,
|
||||
}
|
||||
|
||||
return_value=None,
|
||||
) as mock_user_update:
|
||||
with patch(
|
||||
"litellm.proxy.auth.login_utils.user_update",
|
||||
new_callable=AsyncMock,
|
||||
return_value=None,
|
||||
) as mock_user_update:
|
||||
with patch(
|
||||
"litellm.proxy.auth.login_utils.get_secret_bool",
|
||||
return_value=False,
|
||||
):
|
||||
result = await authenticate_user(
|
||||
username=ui_username,
|
||||
password=master_key,
|
||||
master_key=master_key,
|
||||
prisma_client=mock_prisma_client,
|
||||
)
|
||||
"litellm.proxy.auth.login_utils.get_secret_bool",
|
||||
return_value=False,
|
||||
):
|
||||
result = await authenticate_user(
|
||||
username=ui_username,
|
||||
password=master_key,
|
||||
master_key=master_key,
|
||||
prisma_client=mock_prisma_client,
|
||||
)
|
||||
|
||||
assert isinstance(result, LoginResult)
|
||||
assert result.user_id == LITELLM_PROXY_ADMIN_NAME
|
||||
assert result.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
finally:
|
||||
if original_ui_password:
|
||||
os.environ["UI_PASSWORD"] = original_ui_password
|
||||
assert isinstance(result, LoginResult)
|
||||
assert result.user_id == LITELLM_PROXY_ADMIN_NAME
|
||||
assert result.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
@ -319,7 +315,7 @@ async def test_authenticate_user_email_case_insensitive_login():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_authenticate_user_database_required_for_admin():
|
||||
async def test_authenticate_user_database_required_for_admin(monkeypatch):
|
||||
"""Test that database is required for admin login"""
|
||||
master_key = "sk-1234"
|
||||
ui_username = "admin"
|
||||
|
|
@ -353,7 +349,7 @@ async def test_authenticate_user_database_required_for_admin():
|
|||
assert "No Database connected" in exc_info.value.message
|
||||
finally:
|
||||
if original_db_url:
|
||||
os.environ["DATABASE_URL"] = original_db_url
|
||||
monkeypatch.setenv("DATABASE_URL", original_db_url)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
|
|
|
|||
|
|
@ -17,14 +17,14 @@ from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2
|
|||
from litellm.exceptions import GuardrailRaisedException
|
||||
|
||||
|
||||
def test_deepkeep_guard_config():
|
||||
def test_deepkeep_guard_config(monkeypatch):
|
||||
"""Test DeepKeep guard configuration with init_guardrails_v2."""
|
||||
litellm.set_verbose = True
|
||||
litellm.guardrail_name_config_map = {}
|
||||
|
||||
os.environ["DEEPKEEP_API_KEY"] = "test-key"
|
||||
os.environ["DEEPKEEP_API_BASE"] = "https://test.deepkeep.ai"
|
||||
os.environ["DEEPKEEP_FIREWALL_ID"] = "fw-123"
|
||||
monkeypatch.setenv("DEEPKEEP_API_KEY", "test-key")
|
||||
monkeypatch.setenv("DEEPKEEP_API_BASE", "https://test.deepkeep.ai")
|
||||
monkeypatch.setenv("DEEPKEEP_FIREWALL_ID", "fw-123")
|
||||
|
||||
init_guardrails_v2(
|
||||
all_guardrails=[
|
||||
|
|
@ -108,11 +108,11 @@ class TestDeepKeepGuardrail:
|
|||
== "https://test.deepkeep.ai/v3/openai/beta/litellm_basic_guardrail_api"
|
||||
)
|
||||
|
||||
def test_initialization_with_env_vars(self):
|
||||
def test_initialization_with_env_vars(self, monkeypatch):
|
||||
"""should initialize successfully using environment variables."""
|
||||
os.environ["DEEPKEEP_API_KEY"] = "env-key"
|
||||
os.environ["DEEPKEEP_API_BASE"] = "https://env.deepkeep.ai"
|
||||
os.environ["DEEPKEEP_FIREWALL_ID"] = "fw-env-456"
|
||||
monkeypatch.setenv("DEEPKEEP_API_KEY", "env-key")
|
||||
monkeypatch.setenv("DEEPKEEP_API_BASE", "https://env.deepkeep.ai")
|
||||
monkeypatch.setenv("DEEPKEEP_FIREWALL_ID", "fw-env-456")
|
||||
|
||||
guardrail = DeepKeepGuardrail(
|
||||
guardrail_name="deepkeep-env-test",
|
||||
|
|
|
|||
|
|
@ -26,13 +26,13 @@ from litellm.types.utils import (
|
|||
)
|
||||
|
||||
|
||||
def test_hiddenlayer_config_saas():
|
||||
def test_hiddenlayer_config_saas(monkeypatch):
|
||||
"""Test Hiddenlayer SaaS configuration with init_guardrails_v2."""
|
||||
litellm.set_verbose = True
|
||||
litellm.guardrail_name_config_map = {}
|
||||
|
||||
# Set environment variables for testing
|
||||
os.environ["HIDDENLAYER_API_BASE"] = "https://my.hiddenlayer"
|
||||
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
|
||||
|
||||
init_guardrails_v2(
|
||||
all_guardrails=[
|
||||
|
|
@ -71,9 +71,9 @@ class TestHiddenlayerGuardrail:
|
|||
if key in os.environ:
|
||||
del os.environ[key]
|
||||
|
||||
def test_initialization(self):
|
||||
def test_initialization(self, monkeypatch):
|
||||
"""Test successful initialization with default values."""
|
||||
os.environ["HIDDENLAYER_API_BASE"] = "https://my.hiddenlayer"
|
||||
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
|
||||
|
||||
guardrail = HiddenlayerGuardrail(
|
||||
guardrail_name="hiddenlayer", event_hook="pre_call", default_on=True
|
||||
|
|
@ -94,9 +94,9 @@ class TestHiddenlayerGuardrail:
|
|||
HiddenlayerGuardrail(guardrail_name="hiddenlayer", event_hook="pre_call")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_request_no_violations(self):
|
||||
async def test_apply_guardrail_request_no_violations(self, monkeypatch):
|
||||
"""Test apply_guardrail for request with no violations detected."""
|
||||
os.environ["HIDDENLAYER_API_BASE"] = "https://my.hiddenlayer"
|
||||
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
|
||||
|
||||
# Setup guardrail
|
||||
guardrail = HiddenlayerGuardrail(
|
||||
|
|
@ -151,9 +151,9 @@ class TestHiddenlayerGuardrail:
|
|||
assert call_args.args[0] == f"{guardrail.api_base}/detection/v1/interactions"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_request_with_violations(self):
|
||||
async def test_apply_guardrail_request_with_violations(self, monkeypatch):
|
||||
"""Test apply_guardrail for request with violations detected."""
|
||||
os.environ["HIDDENLAYER_API_BASE"] = "https://my.hiddenlayer"
|
||||
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
|
||||
|
||||
# Setup guardrail
|
||||
guardrail = HiddenlayerGuardrail(
|
||||
|
|
@ -209,9 +209,9 @@ class TestHiddenlayerGuardrail:
|
|||
assert "Blocked by Hiddenlayer" in str(exc_info.value.detail)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_response_no_violations(self):
|
||||
async def test_apply_guardrail_response_no_violations(self, monkeypatch):
|
||||
"""Test apply_guardrail for response with no violations detected."""
|
||||
os.environ["HIDDENLAYER_API_BASE"] = "https://my.hiddenlayer"
|
||||
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
|
||||
|
||||
# Setup guardrail
|
||||
guardrail = HiddenlayerGuardrail(
|
||||
|
|
@ -279,10 +279,10 @@ class TestHiddenlayerGuardrail:
|
|||
mock_post.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_response_with_violations(self):
|
||||
async def test_apply_guardrail_response_with_violations(self, monkeypatch):
|
||||
"""Test apply_guardrail for response with violations detected."""
|
||||
|
||||
os.environ["HIDDENLAYER_API_BASE"] = "https://my.hiddenlayer"
|
||||
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
|
||||
|
||||
# Setup guardrail
|
||||
guardrail = HiddenlayerGuardrail(
|
||||
|
|
@ -348,10 +348,10 @@ class TestHiddenlayerGuardrail:
|
|||
assert exc_info.value.status_code == 400
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_api_error_handling(self):
|
||||
async def test_apply_guardrail_api_error_handling(self, monkeypatch):
|
||||
"""Test handling of API errors in apply_guardrail."""
|
||||
# Set required API key
|
||||
os.environ["HIDDENLAYER_API_BASE"] = "https://my.hiddenlayer"
|
||||
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
|
||||
|
||||
guardrail = HiddenlayerGuardrail(
|
||||
guardrail_name="hiddenlayer", event_hook="pre_call", default_on=True
|
||||
|
|
@ -391,10 +391,10 @@ class TestHiddenlayerGuardrail:
|
|||
assert result == inputs
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_validate_with_call_hiddenlayer_method(self):
|
||||
async def test_validate_with_call_hiddenlayer_method(self, monkeypatch):
|
||||
"""Test the _validate_with_guard_server internal method."""
|
||||
# Set required API key
|
||||
os.environ["HIDDENLAYER_API_BASE"] = "https://my.hiddenlayer"
|
||||
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
|
||||
|
||||
guardrail = HiddenlayerGuardrail(
|
||||
guardrail_name="hiddenlayer", event_hook="pre_call", default_on=True
|
||||
|
|
@ -433,9 +433,9 @@ class TestHiddenlayerGuardrail:
|
|||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_request_with_image(self):
|
||||
async def test_apply_guardrail_request_with_image(self, monkeypatch):
|
||||
"""Test apply_guardrail sends multimodal content (image) to HiddenLayer v1."""
|
||||
os.environ["HIDDENLAYER_API_BASE"] = "https://my.hiddenlayer"
|
||||
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
|
||||
|
||||
guardrail = HiddenlayerGuardrail(
|
||||
guardrail_name="hiddenlayer", event_hook="pre_call", default_on=True
|
||||
|
|
@ -498,9 +498,9 @@ class TestHiddenlayerGuardrail:
|
|||
assert result is not None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_redact_with_image_content(self):
|
||||
async def test_apply_guardrail_redact_with_image_content(self, monkeypatch):
|
||||
"""Test that REDACT action with multimodal content extracts text properly into inputs['texts']."""
|
||||
os.environ["HIDDENLAYER_API_BASE"] = "https://my.hiddenlayer"
|
||||
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
|
||||
|
||||
guardrail = HiddenlayerGuardrail(
|
||||
guardrail_name="hiddenlayer", event_hook="pre_call", default_on=True
|
||||
|
|
@ -570,12 +570,12 @@ class TestHiddenlayerGuardrail:
|
|||
assert config_model.__name__ == "HiddenlayerGuardrailConfigModel"
|
||||
|
||||
|
||||
def test_hiddenlayer_config_v2():
|
||||
def test_hiddenlayer_config_v2(monkeypatch):
|
||||
"""Test HiddenLayer V2 configuration with init_guardrails_v2."""
|
||||
litellm.set_verbose = True
|
||||
litellm.guardrail_name_config_map = {}
|
||||
|
||||
os.environ["HIDDENLAYER_API_BASE"] = "https://my.hiddenlayer"
|
||||
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
|
||||
|
||||
init_guardrails_v2(
|
||||
all_guardrails=[
|
||||
|
|
@ -612,9 +612,9 @@ class TestHiddenlayerGuardrailV2:
|
|||
if key in os.environ:
|
||||
del os.environ[key]
|
||||
|
||||
def test_initialization(self):
|
||||
def test_initialization(self, monkeypatch):
|
||||
"""Test successful initialization with default values."""
|
||||
os.environ["HIDDENLAYER_API_BASE"] = "https://my.hiddenlayer"
|
||||
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
|
||||
|
||||
guardrail = HiddenlayerGuardrailV2(
|
||||
guardrail_name="hiddenlayer", event_hook="pre_call", default_on=True
|
||||
|
|
@ -633,9 +633,9 @@ class TestHiddenlayerGuardrailV2:
|
|||
HiddenlayerGuardrailV2(guardrail_name="hiddenlayer", event_hook="pre_call")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_request_no_violations(self):
|
||||
async def test_apply_guardrail_request_no_violations(self, monkeypatch):
|
||||
"""Test apply_guardrail for request with no violations detected."""
|
||||
os.environ["HIDDENLAYER_API_BASE"] = "https://my.hiddenlayer"
|
||||
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
|
||||
|
||||
guardrail = HiddenlayerGuardrailV2(
|
||||
guardrail_name="hiddenlayer", event_hook="pre_call", default_on=True
|
||||
|
|
@ -691,9 +691,9 @@ class TestHiddenlayerGuardrailV2:
|
|||
assert "detection/v2/request-evaluations" in call_args.args[0]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_request_with_violations(self):
|
||||
async def test_apply_guardrail_request_with_violations(self, monkeypatch):
|
||||
"""Test apply_guardrail for request with violations detected (block via header)."""
|
||||
os.environ["HIDDENLAYER_API_BASE"] = "https://my.hiddenlayer"
|
||||
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
|
||||
|
||||
guardrail = HiddenlayerGuardrailV2(
|
||||
guardrail_name="hiddenlayer", event_hook="pre_call", default_on=True
|
||||
|
|
@ -751,9 +751,9 @@ class TestHiddenlayerGuardrailV2:
|
|||
assert "Blocked by Hiddenlayer" in str(exc_info.value.detail)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_response_no_violations(self):
|
||||
async def test_apply_guardrail_response_no_violations(self, monkeypatch):
|
||||
"""Test apply_guardrail for response with no violations detected."""
|
||||
os.environ["HIDDENLAYER_API_BASE"] = "https://my.hiddenlayer"
|
||||
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
|
||||
|
||||
guardrail = HiddenlayerGuardrailV2(
|
||||
guardrail_name="hiddenlayer", event_hook="post_call", default_on=True
|
||||
|
|
@ -816,9 +816,9 @@ class TestHiddenlayerGuardrailV2:
|
|||
assert "detection/v2/response-evaluations" in call_args.args[0]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_response_with_violations(self):
|
||||
async def test_apply_guardrail_response_with_violations(self, monkeypatch):
|
||||
"""Test apply_guardrail for response with violations detected (block via header)."""
|
||||
os.environ["HIDDENLAYER_API_BASE"] = "https://my.hiddenlayer"
|
||||
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
|
||||
|
||||
guardrail = HiddenlayerGuardrailV2(
|
||||
guardrail_name="hiddenlayer", event_hook="post_call", default_on=True
|
||||
|
|
@ -863,9 +863,9 @@ class TestHiddenlayerGuardrailV2:
|
|||
assert "Blocked by Hiddenlayer" in str(exc_info.value.detail)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_response_with_tool_calls(self):
|
||||
async def test_apply_guardrail_response_with_tool_calls(self, monkeypatch):
|
||||
"""Test apply_guardrail for response containing tool calls."""
|
||||
os.environ["HIDDENLAYER_API_BASE"] = "https://my.hiddenlayer"
|
||||
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
|
||||
|
||||
guardrail = HiddenlayerGuardrailV2(
|
||||
guardrail_name="hiddenlayer", event_hook="post_call", default_on=True
|
||||
|
|
@ -924,9 +924,9 @@ class TestHiddenlayerGuardrailV2:
|
|||
assert "detection/v2/response-evaluations" in call_args.args[0]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_call_hiddenlayer_uses_correct_endpoints(self):
|
||||
async def test_call_hiddenlayer_uses_correct_endpoints(self, monkeypatch):
|
||||
"""Test that _call_hiddenlayer uses the v2 request/response evaluation endpoints."""
|
||||
os.environ["HIDDENLAYER_API_BASE"] = "https://my.hiddenlayer"
|
||||
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
|
||||
|
||||
guardrail = HiddenlayerGuardrailV2(
|
||||
guardrail_name="hiddenlayer", event_hook="pre_call", default_on=True
|
||||
|
|
@ -959,9 +959,9 @@ class TestHiddenlayerGuardrailV2:
|
|||
assert "detection/v2/response-evaluations" in mock_post.call_args.args[0]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_request_with_image(self):
|
||||
async def test_apply_guardrail_request_with_image(self, monkeypatch):
|
||||
"""Test apply_guardrail sends multimodal content (image) to HiddenLayer v2."""
|
||||
os.environ["HIDDENLAYER_API_BASE"] = "https://my.hiddenlayer"
|
||||
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
|
||||
|
||||
guardrail = HiddenlayerGuardrailV2(
|
||||
guardrail_name="hiddenlayer", event_hook="pre_call", default_on=True
|
||||
|
|
@ -1030,9 +1030,9 @@ class TestHiddenlayerGuardrailV2:
|
|||
assert texts == ["how much is on this receipt?"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_request_with_image_multimodal_response(self):
|
||||
async def test_apply_guardrail_request_with_image_multimodal_response(self, monkeypatch):
|
||||
"""Test that new_texts extraction handles multimodal content (list) returned by HiddenLayer v2."""
|
||||
os.environ["HIDDENLAYER_API_BASE"] = "https://my.hiddenlayer"
|
||||
monkeypatch.setenv("HIDDENLAYER_API_BASE", "https://my.hiddenlayer")
|
||||
|
||||
guardrail = HiddenlayerGuardrailV2(
|
||||
guardrail_name="hiddenlayer", event_hook="pre_call", default_on=True
|
||||
|
|
|
|||
|
|
@ -19,13 +19,13 @@ from litellm.proxy.guardrails.guardrail_hooks.lasso.lasso import (
|
|||
from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2
|
||||
|
||||
|
||||
def test_lasso_guard_config():
|
||||
def test_lasso_guard_config(monkeypatch):
|
||||
"""Test Lasso guard configuration with init_guardrails_v2."""
|
||||
litellm.set_verbose = True
|
||||
litellm.guardrail_name_config_map = {}
|
||||
|
||||
# Set environment variable for testing
|
||||
os.environ["LASSO_API_KEY"] = "test-key"
|
||||
monkeypatch.setenv("LASSO_API_KEY", "test-key")
|
||||
|
||||
init_guardrails_v2(
|
||||
all_guardrails=[
|
||||
|
|
|
|||
|
|
@ -18,14 +18,14 @@ from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2
|
|||
from litellm.types.utils import Choices, GenericGuardrailAPIInputs, Message
|
||||
|
||||
|
||||
def test_onyx_guard_config():
|
||||
def test_onyx_guard_config(monkeypatch):
|
||||
"""Test Onyx guard configuration with init_guardrails_v2."""
|
||||
litellm.set_verbose = True
|
||||
litellm.guardrail_name_config_map = {}
|
||||
|
||||
# Set environment variables for testing
|
||||
os.environ["ONYX_API_BASE"] = "https://test.onyx.security"
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
monkeypatch.setenv("ONYX_API_BASE", "https://test.onyx.security")
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
|
||||
init_guardrails_v2(
|
||||
all_guardrails=[
|
||||
|
|
@ -48,11 +48,11 @@ def test_onyx_guard_config():
|
|||
del os.environ["ONYX_API_KEY"]
|
||||
|
||||
|
||||
def test_onyx_guard_with_custom_timeout_from_kwargs():
|
||||
def test_onyx_guard_with_custom_timeout_from_kwargs(monkeypatch):
|
||||
"""Test Onyx guard instantiation with custom timeout passed via kwargs."""
|
||||
# Set environment variables for testing
|
||||
os.environ["ONYX_API_BASE"] = "https://test.onyx.security"
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
monkeypatch.setenv("ONYX_API_BASE", "https://test.onyx.security")
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client"
|
||||
|
|
@ -81,16 +81,16 @@ def test_onyx_guard_with_custom_timeout_from_kwargs():
|
|||
del os.environ["ONYX_API_KEY"]
|
||||
|
||||
|
||||
def test_onyx_guard_with_timeout_none_uses_env_var():
|
||||
def test_onyx_guard_with_timeout_none_uses_env_var(monkeypatch):
|
||||
"""Test Onyx guard with timeout=None uses ONYX_TIMEOUT env var.
|
||||
|
||||
When timeout=None is passed (as it would be from config model with default None),
|
||||
the ONYX_TIMEOUT environment variable should be used.
|
||||
"""
|
||||
# Set environment variables for testing
|
||||
os.environ["ONYX_API_BASE"] = "https://test.onyx.security"
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
os.environ["ONYX_TIMEOUT"] = "60"
|
||||
monkeypatch.setenv("ONYX_API_BASE", "https://test.onyx.security")
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
monkeypatch.setenv("ONYX_TIMEOUT", "60")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client"
|
||||
|
|
@ -121,11 +121,11 @@ def test_onyx_guard_with_timeout_none_uses_env_var():
|
|||
del os.environ["ONYX_TIMEOUT"]
|
||||
|
||||
|
||||
def test_onyx_guard_with_timeout_none_defaults_to_10():
|
||||
def test_onyx_guard_with_timeout_none_defaults_to_10(monkeypatch):
|
||||
"""Test Onyx guard with timeout=None and no env var defaults to 10 seconds."""
|
||||
# Set environment variables for testing
|
||||
os.environ["ONYX_API_BASE"] = "https://test.onyx.security"
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
monkeypatch.setenv("ONYX_API_BASE", "https://test.onyx.security")
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
# Ensure ONYX_TIMEOUT is not set
|
||||
if "ONYX_TIMEOUT" in os.environ:
|
||||
del os.environ["ONYX_TIMEOUT"]
|
||||
|
|
@ -174,10 +174,10 @@ class TestOnyxGuardrail:
|
|||
if key in os.environ:
|
||||
del os.environ[key]
|
||||
|
||||
def test_initialization_with_defaults(self):
|
||||
def test_initialization_with_defaults(self, monkeypatch):
|
||||
"""Test successful initialization with default values."""
|
||||
# Set required API key
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
|
||||
guardrail = OnyxGuardrail(
|
||||
guardrail_name="test-guard", event_hook="pre_call", default_on=True
|
||||
|
|
@ -189,10 +189,10 @@ class TestOnyxGuardrail:
|
|||
assert guardrail.guardrail_name == "test-guard"
|
||||
assert guardrail.event_hook == "pre_call"
|
||||
|
||||
def test_initialization_with_env_vars(self):
|
||||
def test_initialization_with_env_vars(self, monkeypatch):
|
||||
"""Test initialization with environment variables."""
|
||||
os.environ["ONYX_API_BASE"] = "https://custom.onyx.security"
|
||||
os.environ["ONYX_API_KEY"] = "custom-api-key"
|
||||
monkeypatch.setenv("ONYX_API_BASE", "https://custom.onyx.security")
|
||||
monkeypatch.setenv("ONYX_API_KEY", "custom-api-key")
|
||||
|
||||
guardrail = OnyxGuardrail(
|
||||
guardrail_name="test-guard", event_hook="post_call", default_on=True
|
||||
|
|
@ -213,9 +213,9 @@ class TestOnyxGuardrail:
|
|||
):
|
||||
OnyxGuardrail(guardrail_name="test-guard", event_hook="pre_call")
|
||||
|
||||
def test_initialization_with_default_timeout(self):
|
||||
def test_initialization_with_default_timeout(self, monkeypatch):
|
||||
"""Test that default timeout is 10.0 seconds."""
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client"
|
||||
|
|
@ -232,9 +232,9 @@ class TestOnyxGuardrail:
|
|||
assert timeout_param.read == 10.0
|
||||
assert timeout_param.connect == 5.0
|
||||
|
||||
def test_initialization_with_custom_timeout_parameter(self):
|
||||
def test_initialization_with_custom_timeout_parameter(self, monkeypatch):
|
||||
"""Test initialization with custom timeout parameter."""
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client"
|
||||
|
|
@ -254,14 +254,14 @@ class TestOnyxGuardrail:
|
|||
assert timeout_param.read == 30.0
|
||||
assert timeout_param.connect == 5.0
|
||||
|
||||
def test_initialization_with_timeout_from_env_var(self):
|
||||
def test_initialization_with_timeout_from_env_var(self, monkeypatch):
|
||||
"""Test initialization with timeout from ONYX_TIMEOUT environment variable.
|
||||
|
||||
Note: The env var is only used when timeout=None is explicitly passed,
|
||||
since the default parameter value is 10.0 (not None).
|
||||
"""
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
os.environ["ONYX_TIMEOUT"] = "25"
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
monkeypatch.setenv("ONYX_TIMEOUT", "25")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client"
|
||||
|
|
@ -282,10 +282,10 @@ class TestOnyxGuardrail:
|
|||
assert timeout_param.read == 25.0
|
||||
assert timeout_param.connect == 5.0
|
||||
|
||||
def test_initialization_timeout_parameter_overrides_env_var(self):
|
||||
def test_initialization_timeout_parameter_overrides_env_var(self, monkeypatch):
|
||||
"""Test that timeout parameter overrides ONYX_TIMEOUT environment variable."""
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
os.environ["ONYX_TIMEOUT"] = "25"
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
monkeypatch.setenv("ONYX_TIMEOUT", "25")
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.guardrails.guardrail_hooks.onyx.onyx.get_async_httpx_client"
|
||||
|
|
@ -306,10 +306,10 @@ class TestOnyxGuardrail:
|
|||
assert timeout_param.connect == 5.0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_request_no_violations(self):
|
||||
async def test_apply_guardrail_request_no_violations(self, monkeypatch):
|
||||
"""Test apply_guardrail for request with no violations detected."""
|
||||
# Set required API key
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
|
||||
# Setup guardrail
|
||||
guardrail = OnyxGuardrail(
|
||||
|
|
@ -372,10 +372,10 @@ class TestOnyxGuardrail:
|
|||
assert call_args.kwargs["json"]["conversation_id"] == "test-call-id"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_request_with_violations(self):
|
||||
async def test_apply_guardrail_request_with_violations(self, monkeypatch):
|
||||
"""Test apply_guardrail for request with violations detected."""
|
||||
# Set required API key
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
|
||||
# Setup guardrail
|
||||
guardrail = OnyxGuardrail(
|
||||
|
|
@ -423,10 +423,10 @@ class TestOnyxGuardrail:
|
|||
assert "prompt_injection" in str(exc_info.value.detail)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_response_no_violations(self):
|
||||
async def test_apply_guardrail_response_no_violations(self, monkeypatch):
|
||||
"""Test apply_guardrail for response with no violations detected."""
|
||||
# Set required API key
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
|
||||
# Setup guardrail
|
||||
guardrail = OnyxGuardrail(
|
||||
|
|
@ -497,10 +497,10 @@ class TestOnyxGuardrail:
|
|||
assert call_args.kwargs["json"]["conversation_id"] == "test-call-id-2"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_response_with_violations(self):
|
||||
async def test_apply_guardrail_response_with_violations(self, monkeypatch):
|
||||
"""Test apply_guardrail for response with violations detected."""
|
||||
# Set required API key
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
|
||||
# Setup guardrail
|
||||
guardrail = OnyxGuardrail(
|
||||
|
|
@ -558,10 +558,10 @@ class TestOnyxGuardrail:
|
|||
assert "illegal_instructions" in str(exc_info.value.detail)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_api_error_handling(self):
|
||||
async def test_apply_guardrail_api_error_handling(self, monkeypatch):
|
||||
"""Test handling of API errors in apply_guardrail."""
|
||||
# Set required API key
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
|
||||
guardrail = OnyxGuardrail(
|
||||
guardrail_name="test-guard", event_hook="pre_call", default_on=True
|
||||
|
|
@ -591,10 +591,10 @@ class TestOnyxGuardrail:
|
|||
assert result == inputs
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_timeout_error_handling(self):
|
||||
async def test_apply_guardrail_timeout_error_handling(self, monkeypatch):
|
||||
"""Test handling of timeout errors in apply_guardrail (graceful degradation)."""
|
||||
# Set required API key
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
|
||||
guardrail = OnyxGuardrail(
|
||||
guardrail_name="test-guard",
|
||||
|
|
@ -629,10 +629,10 @@ class TestOnyxGuardrail:
|
|||
assert result == inputs
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_read_timeout_error_handling(self):
|
||||
async def test_apply_guardrail_read_timeout_error_handling(self, monkeypatch):
|
||||
"""Test handling of read timeout errors in apply_guardrail."""
|
||||
# Set required API key
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
|
||||
guardrail = OnyxGuardrail(
|
||||
guardrail_name="test-guard",
|
||||
|
|
@ -667,10 +667,10 @@ class TestOnyxGuardrail:
|
|||
assert result == inputs
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_connect_timeout_error_handling(self):
|
||||
async def test_apply_guardrail_connect_timeout_error_handling(self, monkeypatch):
|
||||
"""Test handling of connect timeout errors in apply_guardrail."""
|
||||
# Set required API key
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
|
||||
guardrail = OnyxGuardrail(
|
||||
guardrail_name="test-guard",
|
||||
|
|
@ -705,10 +705,10 @@ class TestOnyxGuardrail:
|
|||
assert result == inputs
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_no_logging_obj(self):
|
||||
async def test_apply_guardrail_no_logging_obj(self, monkeypatch):
|
||||
"""Test apply_guardrail without logging object (uses UUID)."""
|
||||
# Set required API key
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
|
||||
guardrail = OnyxGuardrail(
|
||||
guardrail_name="test-guard", event_hook="pre_call", default_on=True
|
||||
|
|
@ -747,10 +747,10 @@ class TestOnyxGuardrail:
|
|||
assert call_args.kwargs["json"]["conversation_id"] == "test-uuid"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_validate_with_guard_server_method(self):
|
||||
async def test_validate_with_guard_server_method(self, monkeypatch):
|
||||
"""Test the _validate_with_guard_server internal method."""
|
||||
# Set required API key
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
|
||||
guardrail = OnyxGuardrail(
|
||||
guardrail_name="test-guard", event_hook="pre_call", default_on=True
|
||||
|
|
@ -788,10 +788,10 @@ class TestOnyxGuardrail:
|
|||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_validate_with_guard_server_blocked(self):
|
||||
async def test_validate_with_guard_server_blocked(self, monkeypatch):
|
||||
"""Test _validate_with_guard_server when request is blocked."""
|
||||
# Set required API key
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
|
||||
guardrail = OnyxGuardrail(
|
||||
guardrail_name="test-guard", event_hook="pre_call", default_on=True
|
||||
|
|
@ -825,10 +825,10 @@ class TestOnyxGuardrail:
|
|||
assert config_model.__name__ == "OnyxGuardrailConfigModel"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_with_modelresponse(self):
|
||||
async def test_apply_guardrail_with_modelresponse(self, monkeypatch):
|
||||
"""Test apply_guardrail with ModelResponse object for response type."""
|
||||
# Set required API key
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
|
||||
guardrail = OnyxGuardrail(
|
||||
guardrail_name="test-guard", event_hook="post_call", default_on=True
|
||||
|
|
@ -880,10 +880,10 @@ class TestOnyxGuardrail:
|
|||
assert "payload" in call_args.kwargs["json"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_response_error_handling(self):
|
||||
async def test_apply_guardrail_response_error_handling(self, monkeypatch):
|
||||
"""Test error handling when processing response data."""
|
||||
# Set required API key
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
|
||||
guardrail = OnyxGuardrail(
|
||||
guardrail_name="test-guard", event_hook="post_call", default_on=True
|
||||
|
|
@ -925,11 +925,11 @@ class TestOnyxIntegration:
|
|||
"""Test integration scenarios."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_full_guardrail_flow(self):
|
||||
async def test_full_guardrail_flow(self, monkeypatch):
|
||||
"""Test full guardrail flow with multiple hooks."""
|
||||
# Set environment variables
|
||||
os.environ["ONYX_API_BASE"] = "https://test.onyx.security"
|
||||
os.environ["ONYX_API_KEY"] = "test-key"
|
||||
monkeypatch.setenv("ONYX_API_BASE", "https://test.onyx.security")
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-key")
|
||||
|
||||
init_guardrails_v2(
|
||||
all_guardrails=[
|
||||
|
|
@ -973,10 +973,10 @@ class TestOnyxIntegration:
|
|||
del os.environ["ONYX_API_KEY"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_empty_request_data(self):
|
||||
async def test_apply_guardrail_empty_request_data(self, monkeypatch):
|
||||
"""Test apply_guardrail with empty request data."""
|
||||
# Set required API key
|
||||
os.environ["ONYX_API_KEY"] = "test-api-key"
|
||||
monkeypatch.setenv("ONYX_API_KEY", "test-api-key")
|
||||
|
||||
guardrail = OnyxGuardrail(
|
||||
guardrail_name="test-guard", event_hook="pre_call", default_on=True
|
||||
|
|
|
|||
|
|
@ -93,24 +93,24 @@ class TestRepelloAIInitialization:
|
|||
with pytest.raises(ValueError, match="asset_id"):
|
||||
RepelloAIGuardrail(api_key="test-api-key", guardrail_name="t")
|
||||
|
||||
def test_api_key_from_env(self):
|
||||
os.environ["REPELLOAI_API_KEY"] = "env-key"
|
||||
def test_api_key_from_env(self, monkeypatch):
|
||||
monkeypatch.setenv("REPELLOAI_API_KEY", "env-key")
|
||||
guardrail = RepelloAIGuardrail(asset_id="asset-123", guardrail_name="t")
|
||||
assert guardrail.repelloai_api_key == "env-key"
|
||||
|
||||
def test_api_key_from_argus_env(self):
|
||||
os.environ["ARGUS_API_KEY"] = "argus-key"
|
||||
def test_api_key_from_argus_env(self, monkeypatch):
|
||||
monkeypatch.setenv("ARGUS_API_KEY", "argus-key")
|
||||
guardrail = RepelloAIGuardrail(asset_id="asset-123", guardrail_name="t")
|
||||
assert guardrail.repelloai_api_key == "argus-key"
|
||||
|
||||
def test_argus_env_preferred_over_legacy(self):
|
||||
os.environ["ARGUS_API_KEY"] = "argus-key"
|
||||
os.environ["REPELLOAI_API_KEY"] = "legacy-key"
|
||||
def test_argus_env_preferred_over_legacy(self, monkeypatch):
|
||||
monkeypatch.setenv("ARGUS_API_KEY", "argus-key")
|
||||
monkeypatch.setenv("REPELLOAI_API_KEY", "legacy-key")
|
||||
guardrail = RepelloAIGuardrail(asset_id="asset-123", guardrail_name="t")
|
||||
assert guardrail.repelloai_api_key == "argus-key"
|
||||
|
||||
def test_explicit_api_key_preferred_over_env(self):
|
||||
os.environ["ARGUS_API_KEY"] = "argus-key"
|
||||
def test_explicit_api_key_preferred_over_env(self, monkeypatch):
|
||||
monkeypatch.setenv("ARGUS_API_KEY", "argus-key")
|
||||
guardrail = RepelloAIGuardrail(
|
||||
api_key="explicit-key", asset_id="asset-123", guardrail_name="t"
|
||||
)
|
||||
|
|
@ -145,10 +145,10 @@ class TestRepelloAIInitialization:
|
|||
assert guardrail.api_base == DEFAULT_REPELLOAI_API_BASE
|
||||
assert guardrail.unreachable_fallback == "fail_closed"
|
||||
|
||||
def test_init_guardrails_v2_wiring(self):
|
||||
def test_init_guardrails_v2_wiring(self, monkeypatch):
|
||||
"""The guardrail registers and constructs via the config.yaml path."""
|
||||
litellm.guardrail_name_config_map = {}
|
||||
os.environ["REPELLOAI_API_KEY"] = "test-key"
|
||||
monkeypatch.setenv("REPELLOAI_API_KEY", "test-key")
|
||||
init_guardrails_v2(
|
||||
all_guardrails=[
|
||||
{
|
||||
|
|
|
|||
|
|
@ -19,14 +19,14 @@ import litellm
|
|||
from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2
|
||||
|
||||
|
||||
def test_prompt_security_guard_config():
|
||||
def test_prompt_security_guard_config(monkeypatch):
|
||||
"""Test guardrail initialization with proper configuration"""
|
||||
litellm.set_verbose = True
|
||||
litellm.guardrail_name_config_map = {}
|
||||
|
||||
# Set environment variables for testing
|
||||
os.environ["PROMPT_SECURITY_API_KEY"] = "test-key"
|
||||
os.environ["PROMPT_SECURITY_API_BASE"] = "https://test.prompt.security"
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
|
||||
|
||||
init_guardrails_v2(
|
||||
all_guardrails=[
|
||||
|
|
@ -78,10 +78,10 @@ def test_prompt_security_guard_config_no_api_key():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_block_request():
|
||||
async def test_apply_guardrail_block_request(monkeypatch):
|
||||
"""Test that apply_guardrail blocks malicious prompts"""
|
||||
os.environ["PROMPT_SECURITY_API_KEY"] = "test-key"
|
||||
os.environ["PROMPT_SECURITY_API_BASE"] = "https://test.prompt.security"
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
|
||||
|
||||
guardrail = PromptSecurityGuardrail(
|
||||
guardrail_name="test-guard", event_hook="pre_call", default_on=True
|
||||
|
|
@ -132,10 +132,10 @@ async def test_apply_guardrail_block_request():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_modify_request():
|
||||
async def test_apply_guardrail_modify_request(monkeypatch):
|
||||
"""Test that apply_guardrail modifies prompts when needed"""
|
||||
os.environ["PROMPT_SECURITY_API_KEY"] = "test-key"
|
||||
os.environ["PROMPT_SECURITY_API_BASE"] = "https://test.prompt.security"
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
|
||||
|
||||
guardrail = PromptSecurityGuardrail(
|
||||
guardrail_name="test-guard", event_hook="pre_call", default_on=True
|
||||
|
|
@ -183,10 +183,10 @@ async def test_apply_guardrail_modify_request():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_allow_request():
|
||||
async def test_apply_guardrail_allow_request(monkeypatch):
|
||||
"""Test that apply_guardrail allows safe prompts"""
|
||||
os.environ["PROMPT_SECURITY_API_KEY"] = "test-key"
|
||||
os.environ["PROMPT_SECURITY_API_BASE"] = "https://test.prompt.security"
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
|
||||
|
||||
guardrail = PromptSecurityGuardrail(
|
||||
guardrail_name="test-guard", event_hook="pre_call", default_on=True
|
||||
|
|
@ -226,10 +226,10 @@ async def test_apply_guardrail_allow_request():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_block_response():
|
||||
async def test_apply_guardrail_block_response(monkeypatch):
|
||||
"""Test that apply_guardrail blocks malicious responses"""
|
||||
os.environ["PROMPT_SECURITY_API_KEY"] = "test-key"
|
||||
os.environ["PROMPT_SECURITY_API_BASE"] = "https://test.prompt.security"
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
|
||||
|
||||
guardrail = PromptSecurityGuardrail(
|
||||
guardrail_name="test-guard", event_hook="post_call", default_on=True
|
||||
|
|
@ -273,10 +273,10 @@ async def test_apply_guardrail_block_response():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_apply_guardrail_modify_response():
|
||||
async def test_apply_guardrail_modify_response(monkeypatch):
|
||||
"""Test that apply_guardrail modifies responses when needed"""
|
||||
os.environ["PROMPT_SECURITY_API_KEY"] = "test-key"
|
||||
os.environ["PROMPT_SECURITY_API_BASE"] = "https://test.prompt.security"
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
|
||||
|
||||
guardrail = PromptSecurityGuardrail(
|
||||
guardrail_name="test-guard", event_hook="post_call", default_on=True
|
||||
|
|
@ -317,10 +317,10 @@ async def test_apply_guardrail_modify_response():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_file_sanitization():
|
||||
async def test_file_sanitization(monkeypatch):
|
||||
"""Test file sanitization for images"""
|
||||
os.environ["PROMPT_SECURITY_API_KEY"] = "test-key"
|
||||
os.environ["PROMPT_SECURITY_API_BASE"] = "https://test.prompt.security"
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
|
||||
|
||||
guardrail = PromptSecurityGuardrail(
|
||||
guardrail_name="test-guard", event_hook="pre_call", default_on=True
|
||||
|
|
@ -407,10 +407,10 @@ async def test_file_sanitization():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_file_sanitization_block():
|
||||
async def test_file_sanitization_block(monkeypatch):
|
||||
"""Test that file sanitization blocks malicious files"""
|
||||
os.environ["PROMPT_SECURITY_API_KEY"] = "test-key"
|
||||
os.environ["PROMPT_SECURITY_API_BASE"] = "https://test.prompt.security"
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
|
||||
|
||||
guardrail = PromptSecurityGuardrail(
|
||||
guardrail_name="test-guard", event_hook="pre_call", default_on=True
|
||||
|
|
@ -491,10 +491,10 @@ async def test_file_sanitization_block():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_user_api_key_alias_forwarding():
|
||||
async def test_user_api_key_alias_forwarding(monkeypatch):
|
||||
"""Test that user API key alias is properly sent via headers and payload"""
|
||||
os.environ["PROMPT_SECURITY_API_KEY"] = "test-key"
|
||||
os.environ["PROMPT_SECURITY_API_BASE"] = "https://test.prompt.security"
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
|
||||
|
||||
guardrail = PromptSecurityGuardrail(
|
||||
guardrail_name="test-guard", event_hook="pre_call", default_on=True
|
||||
|
|
@ -535,10 +535,10 @@ async def test_user_api_key_alias_forwarding():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_role_filtering():
|
||||
async def test_role_filtering(monkeypatch):
|
||||
"""Test that tool/function messages are filtered out by default"""
|
||||
os.environ["PROMPT_SECURITY_API_KEY"] = "test-key"
|
||||
os.environ["PROMPT_SECURITY_API_BASE"] = "https://test.prompt.security"
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
|
||||
|
||||
guardrail = PromptSecurityGuardrail(
|
||||
guardrail_name="test-guard", event_hook="pre_call", default_on=True
|
||||
|
|
@ -600,11 +600,11 @@ async def test_role_filtering():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_check_tool_results_enabled():
|
||||
async def test_check_tool_results_enabled(monkeypatch):
|
||||
"""Test with check_tool_results=True: transforms tool/function to 'other' role"""
|
||||
os.environ["PROMPT_SECURITY_API_KEY"] = "test-key"
|
||||
os.environ["PROMPT_SECURITY_API_BASE"] = "https://test.prompt.security"
|
||||
os.environ["PROMPT_SECURITY_CHECK_TOOL_RESULTS"] = "true"
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_KEY", "test-key")
|
||||
monkeypatch.setenv("PROMPT_SECURITY_API_BASE", "https://test.prompt.security")
|
||||
monkeypatch.setenv("PROMPT_SECURITY_CHECK_TOOL_RESULTS", "true")
|
||||
|
||||
guardrail = PromptSecurityGuardrail(
|
||||
guardrail_name="test-guard", event_hook="pre_call", default_on=True
|
||||
|
|
|
|||
|
|
@ -42,7 +42,7 @@ def time_controller(monkeypatch):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_priority_weight_allocation():
|
||||
async def test_priority_weight_allocation(monkeypatch):
|
||||
"""
|
||||
Test that priority weights are correctly applied instead of equal splitting.
|
||||
|
||||
|
|
@ -53,7 +53,7 @@ async def test_priority_weight_allocation():
|
|||
This validates the core fix where before it would split 50/50.
|
||||
"""
|
||||
# Set up environment for premium feature
|
||||
os.environ["LITELLM_LICENSE"] = "test-license-key"
|
||||
monkeypatch.setenv("LITELLM_LICENSE", "test-license-key")
|
||||
|
||||
# Set up priority reservations
|
||||
litellm.priority_reservation = {"high": 0.9, "low": 0.1}
|
||||
|
|
@ -128,7 +128,7 @@ async def test_priority_weight_allocation():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_priority_requests():
|
||||
async def test_concurrent_priority_requests(monkeypatch):
|
||||
"""
|
||||
Test the core issue: 5 concurrent requests with different priorities should get
|
||||
proper allocation based on priority weights, not equal splitting.
|
||||
|
|
@ -136,7 +136,7 @@ async def test_concurrent_priority_requests():
|
|||
This tests the exact scenario mentioned: priorities 0.9 and 0.1 should be 0.9/0.1, not 0.5/0.5.
|
||||
"""
|
||||
# Set up environment for premium feature
|
||||
os.environ["LITELLM_LICENSE"] = "test-license-key"
|
||||
monkeypatch.setenv("LITELLM_LICENSE", "test-license-key")
|
||||
|
||||
# Set up the exact scenario from the issue
|
||||
litellm.priority_reservation = {"high": 0.9, "low": 0.1}
|
||||
|
|
@ -214,7 +214,7 @@ async def test_concurrent_priority_requests():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_100_concurrent_priority_requests(time_controller):
|
||||
async def test_100_concurrent_priority_requests(time_controller, monkeypatch):
|
||||
"""
|
||||
Stress test: 100 concurrent requests with mixed priorities over 10 seconds.
|
||||
|
||||
|
|
@ -224,7 +224,7 @@ async def test_100_concurrent_priority_requests(time_controller):
|
|||
- Spread across 10 seconds to simulate real-world load
|
||||
"""
|
||||
# Set up environment for premium feature
|
||||
os.environ["LITELLM_LICENSE"] = "test-license-key"
|
||||
monkeypatch.setenv("LITELLM_LICENSE", "test-license-key")
|
||||
|
||||
# Set up priority reservations
|
||||
litellm.priority_reservation = {"high": 0.9, "low": 0.1}
|
||||
|
|
@ -384,7 +384,7 @@ async def test_100_concurrent_priority_requests(time_controller):
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_pre_call_hooks_stress():
|
||||
async def test_concurrent_pre_call_hooks_stress(monkeypatch):
|
||||
"""
|
||||
Stress test: 50 concurrent pre-call hooks with saturation-aware priority enforcement.
|
||||
|
||||
|
|
@ -394,7 +394,7 @@ async def test_concurrent_pre_call_hooks_stress():
|
|||
Standard users (20% allocation) should have ~70% success rate with 30% random limiting.
|
||||
"""
|
||||
# Set up environment for premium feature
|
||||
os.environ["LITELLM_LICENSE"] = "test-license-key"
|
||||
monkeypatch.setenv("LITELLM_LICENSE", "test-license-key")
|
||||
|
||||
litellm.priority_reservation = {"premium": 0.8, "standard": 0.2}
|
||||
|
||||
|
|
@ -634,7 +634,7 @@ async def test_concurrent_pre_call_hooks_stress():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fake_calls_case_1_no_rate_limiting_at_capacity():
|
||||
async def test_fake_calls_case_1_no_rate_limiting_at_capacity(monkeypatch):
|
||||
"""
|
||||
Test Case 1: Saturation-Aware Rate Limiting at 50% Threshold
|
||||
|
||||
|
|
@ -650,7 +650,7 @@ async def test_fake_calls_case_1_no_rate_limiting_at_capacity():
|
|||
|
||||
Once saturation hits 50%, strict mode enforces priority-based limits.
|
||||
"""
|
||||
os.environ["LITELLM_LICENSE"] = "test-license-key"
|
||||
monkeypatch.setenv("LITELLM_LICENSE", "test-license-key")
|
||||
|
||||
# Set up priority reservations
|
||||
litellm.priority_reservation = {"key_a": 0.75, "key_b": 0.25}
|
||||
|
|
@ -759,7 +759,7 @@ async def test_fake_calls_case_1_no_rate_limiting_at_capacity():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fake_calls_case_2_priority_queue_during_saturation():
|
||||
async def test_fake_calls_case_2_priority_queue_during_saturation(monkeypatch):
|
||||
"""
|
||||
Test Case 2: Priority Queue Behavior During Saturation
|
||||
|
||||
|
|
@ -773,7 +773,7 @@ async def test_fake_calls_case_2_priority_queue_during_saturation():
|
|||
|
||||
When total traffic exceeds capacity, rate limiting enforces priority reservations.
|
||||
"""
|
||||
os.environ["LITELLM_LICENSE"] = "test-license-key"
|
||||
monkeypatch.setenv("LITELLM_LICENSE", "test-license-key")
|
||||
|
||||
litellm.priority_reservation = {"key_a": 0.75, "key_b": 0.25}
|
||||
|
||||
|
|
@ -886,7 +886,7 @@ async def test_fake_calls_case_2_priority_queue_during_saturation():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fake_calls_case_3_spillover_capacity_default_keys():
|
||||
async def test_fake_calls_case_3_spillover_capacity_default_keys(monkeypatch):
|
||||
"""
|
||||
Test Case 3: Spillover Capacity for Default Keys
|
||||
|
||||
|
|
@ -906,7 +906,7 @@ async def test_fake_calls_case_3_spillover_capacity_default_keys():
|
|||
|
||||
Tests spillover behavior where default keys share remaining capacity.
|
||||
"""
|
||||
os.environ["LITELLM_LICENSE"] = "test-license-key"
|
||||
monkeypatch.setenv("LITELLM_LICENSE", "test-license-key")
|
||||
|
||||
litellm.priority_reservation = {"key_a": 0.75}
|
||||
litellm.priority_reservation_settings.default_priority = 0.25
|
||||
|
|
@ -1025,7 +1025,7 @@ async def test_fake_calls_case_3_spillover_capacity_default_keys():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fake_calls_case_4_over_allocated_with_normalization():
|
||||
async def test_fake_calls_case_4_over_allocated_with_normalization(monkeypatch):
|
||||
"""
|
||||
Test Case 4: Over-Allocated Priority reservations with Normalization
|
||||
|
||||
|
|
@ -1042,7 +1042,7 @@ async def test_fake_calls_case_4_over_allocated_with_normalization():
|
|||
- Due to concurrent burst, total successful may exceed 100 RPM in the test window
|
||||
- This test verifies normalization works and total capacity is reasonably bounded
|
||||
"""
|
||||
os.environ["LITELLM_LICENSE"] = "test-license-key"
|
||||
monkeypatch.setenv("LITELLM_LICENSE", "test-license-key")
|
||||
|
||||
litellm.priority_reservation = {"key_a": 0.60, "key_b": 0.80}
|
||||
|
||||
|
|
@ -1156,7 +1156,7 @@ async def test_fake_calls_case_4_over_allocated_with_normalization():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fake_calls_case_5_default_value_priority_reservation():
|
||||
async def test_fake_calls_case_5_default_value_priority_reservation(monkeypatch):
|
||||
"""
|
||||
Test Case 5: Default value for priority reservation
|
||||
|
||||
|
|
@ -1176,7 +1176,7 @@ async def test_fake_calls_case_5_default_value_priority_reservation():
|
|||
|
||||
Tests complex scenario with explicit priorities and default priority.
|
||||
"""
|
||||
os.environ["LITELLM_LICENSE"] = "test-license-key"
|
||||
monkeypatch.setenv("LITELLM_LICENSE", "test-license-key")
|
||||
|
||||
litellm.priority_reservation = {"key_a": 0.50, "key_b": 0.20, "key_c": 0.05}
|
||||
litellm.priority_reservation_settings.default_priority = 0.05
|
||||
|
|
@ -1296,7 +1296,7 @@ async def test_fake_calls_case_5_default_value_priority_reservation():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_default_priority_shared_pool():
|
||||
async def test_default_priority_shared_pool(monkeypatch):
|
||||
"""
|
||||
Test that keys without explicit priority share ONE default pool, not get individual allocations.
|
||||
|
||||
|
|
@ -1304,7 +1304,7 @@ async def test_default_priority_shared_pool():
|
|||
- Key A, B, C (no priority) should share ONE 25 RPM pool
|
||||
- NOT get 25 RPM each (which would be 75 RPM total)
|
||||
"""
|
||||
os.environ["LITELLM_LICENSE"] = "test-license-key"
|
||||
monkeypatch.setenv("LITELLM_LICENSE", "test-license-key")
|
||||
|
||||
litellm.priority_reservation = {"prod": 0.75}
|
||||
litellm.priority_reservation_settings.default_priority = 0.25
|
||||
|
|
@ -1382,7 +1382,7 @@ async def test_default_priority_shared_pool():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_log_success_event_increments_by_actual_tokens():
|
||||
async def test_async_log_success_event_increments_by_actual_tokens(monkeypatch):
|
||||
"""
|
||||
Test that async_log_success_event increments token counters by actual token usage.
|
||||
|
||||
|
|
@ -1394,7 +1394,7 @@ async def test_async_log_success_event_increments_by_actual_tokens():
|
|||
|
||||
from litellm.types.utils import ModelResponse, Usage
|
||||
|
||||
os.environ["LITELLM_LICENSE"] = "test-license-key"
|
||||
monkeypatch.setenv("LITELLM_LICENSE", "test-license-key")
|
||||
litellm.priority_reservation = {"dev": 0.1, "prod": 0.9}
|
||||
|
||||
dual_cache = DualCache()
|
||||
|
|
@ -1483,7 +1483,7 @@ async def test_async_log_success_event_increments_by_actual_tokens():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_saturation_check_cache_ttl_configuration():
|
||||
async def test_saturation_check_cache_ttl_configuration(monkeypatch):
|
||||
"""
|
||||
Test that saturation_check_cache_ttl controls how long saturation values are cached locally.
|
||||
|
||||
|
|
@ -1492,7 +1492,7 @@ async def test_saturation_check_cache_ttl_configuration():
|
|||
- After expiration, fresh values should be fetched from Redis
|
||||
- This prevents nodes from having stale saturation data in multi-node deployments
|
||||
"""
|
||||
os.environ["LITELLM_LICENSE"] = "test-license-key"
|
||||
monkeypatch.setenv("LITELLM_LICENSE", "test-license-key")
|
||||
|
||||
# Set a short TTL for testing (5 seconds)
|
||||
original_ttl = litellm.priority_reservation_settings.saturation_check_cache_ttl
|
||||
|
|
@ -1587,7 +1587,7 @@ async def test_saturation_check_cache_ttl_configuration():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_log_success_event_uses_team_priority_from_auth_metadata():
|
||||
async def test_async_log_success_event_uses_team_priority_from_auth_metadata(monkeypatch):
|
||||
"""
|
||||
Test that async_log_success_event correctly retrieves priority from user_api_key_auth_metadata.
|
||||
|
||||
|
|
@ -1598,7 +1598,7 @@ async def test_async_log_success_event_uses_team_priority_from_auth_metadata():
|
|||
|
||||
from litellm.types.utils import ModelResponse, Usage
|
||||
|
||||
os.environ["LITELLM_LICENSE"] = "test-license-key"
|
||||
monkeypatch.setenv("LITELLM_LICENSE", "test-license-key")
|
||||
litellm.priority_reservation = {"team_priority": 0.8, "default": 0.2}
|
||||
|
||||
dual_cache = DualCache()
|
||||
|
|
@ -1680,7 +1680,7 @@ async def test_async_log_success_event_uses_team_priority_from_auth_metadata():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_priority_429_includes_model_name_and_configured_limits():
|
||||
async def test_priority_429_includes_model_name_and_configured_limits(monkeypatch):
|
||||
"""
|
||||
The priority-based 429 should tell operators which model was hit and what
|
||||
the model's configured TPM/RPM are, so they can decide whether to tune the
|
||||
|
|
@ -1694,7 +1694,7 @@ async def test_priority_429_includes_model_name_and_configured_limits():
|
|||
"""
|
||||
from fastapi import HTTPException
|
||||
|
||||
os.environ["LITELLM_LICENSE"] = "test-license-key"
|
||||
monkeypatch.setenv("LITELLM_LICENSE", "test-license-key")
|
||||
litellm.priority_reservation = {"prod": 0.5}
|
||||
|
||||
dual_cache = DualCache()
|
||||
|
|
@ -1774,7 +1774,7 @@ async def test_priority_429_includes_model_name_and_configured_limits():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tpm_only_model_enforces_priority_and_model_capacity():
|
||||
async def test_tpm_only_model_enforces_priority_and_model_capacity(monkeypatch):
|
||||
"""Regression: a model configured with ONLY tpm (no rpm) must still be
|
||||
rate limited.
|
||||
|
||||
|
|
@ -1789,7 +1789,7 @@ async def test_tpm_only_model_enforces_priority_and_model_capacity():
|
|||
|
||||
from litellm.types.utils import ModelResponse, Usage
|
||||
|
||||
os.environ["LITELLM_LICENSE"] = "test-license-key"
|
||||
monkeypatch.setenv("LITELLM_LICENSE", "test-license-key")
|
||||
litellm.priority_reservation = {"dev": 0.25, "prod": 0.5}
|
||||
|
||||
dual_cache = DualCache()
|
||||
|
|
|
|||
|
|
@ -189,7 +189,7 @@ async def test_batch_limiter_uses_atomic_check_and_increment():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dynamic_rate_limiter_v3_concurrent_bypasses_model_capacity():
|
||||
async def test_dynamic_rate_limiter_v3_concurrent_bypasses_model_capacity(monkeypatch):
|
||||
"""
|
||||
DynamicRateLimitHandler PHASE 1 (read_only check) → PHASE 3 (increment)
|
||||
is non-atomic: dynamic_rate_limiter_v3.py:463-548.
|
||||
|
|
@ -209,7 +209,7 @@ async def test_dynamic_rate_limiter_v3_concurrent_bypasses_model_capacity():
|
|||
# RPM + 1 successes before the next sees counter > RPM.
|
||||
MAX_SEQUENTIAL_SUCCESSES = MODEL_RPM + 1
|
||||
|
||||
os.environ["LITELLM_LICENSE"] = "test-license-key"
|
||||
monkeypatch.setenv("LITELLM_LICENSE", "test-license-key")
|
||||
litellm.priority_reservation = {"high": 0.9, "low": 0.1}
|
||||
|
||||
dual_cache = DualCache()
|
||||
|
|
@ -273,7 +273,7 @@ async def test_dynamic_rate_limiter_v3_concurrent_bypasses_model_capacity():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dynamic_rate_limiter_v3_uses_atomic_check_and_increment():
|
||||
async def test_dynamic_rate_limiter_v3_uses_atomic_check_and_increment(monkeypatch):
|
||||
"""
|
||||
Regression test: dynamic limiter's enforced descriptors flow through
|
||||
`atomic_check_and_increment_by_n`, not the legacy
|
||||
|
|
@ -283,7 +283,7 @@ async def test_dynamic_rate_limiter_v3_uses_atomic_check_and_increment():
|
|||
bundled into the atomic call alongside model_saturation_check. When not
|
||||
enforced, priority counter is incremented for tracking only.
|
||||
"""
|
||||
os.environ["LITELLM_LICENSE"] = "test-license-key"
|
||||
monkeypatch.setenv("LITELLM_LICENSE", "test-license-key")
|
||||
litellm.priority_reservation = {"high": 0.9, "low": 0.1}
|
||||
|
||||
dual_cache = DualCache()
|
||||
|
|
@ -413,7 +413,7 @@ async def test_batch_zero_token_consumes_rpm_only():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dynamic_rate_limiter_v3_fails_closed_on_unknown_descriptor():
|
||||
async def test_dynamic_rate_limiter_v3_fails_closed_on_unknown_descriptor(monkeypatch):
|
||||
"""
|
||||
Fail-closed guard: when atomic_check_and_increment_by_n returns
|
||||
overall_code=OVER_LIMIT but with a descriptor_key the dispatcher does
|
||||
|
|
@ -425,7 +425,7 @@ async def test_dynamic_rate_limiter_v3_fails_closed_on_unknown_descriptor():
|
|||
"""
|
||||
from fastapi import HTTPException
|
||||
|
||||
os.environ["LITELLM_LICENSE"] = "test-license-key"
|
||||
monkeypatch.setenv("LITELLM_LICENSE", "test-license-key")
|
||||
litellm.priority_reservation = {"high": 0.9, "low": 0.1}
|
||||
|
||||
dual_cache = DualCache()
|
||||
|
|
|
|||
|
|
@ -25,12 +25,9 @@ from litellm.types.utils import StandardAuditLogPayload
|
|||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_audit_log_callbacks():
|
||||
"""Reset audit_log_callbacks before and after each test."""
|
||||
original = litellm.audit_log_callbacks
|
||||
litellm.audit_log_callbacks = []
|
||||
yield
|
||||
litellm.audit_log_callbacks = original
|
||||
def reset_audit_log_callbacks(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Every test starts with no audit log callbacks registered."""
|
||||
monkeypatch.setattr(litellm, "audit_log_callbacks", [])
|
||||
|
||||
|
||||
def _make_audit_log(
|
||||
|
|
@ -115,10 +112,10 @@ class TestBuildAuditLogPayload:
|
|||
|
||||
class TestDispatchAuditLogToCallbacks:
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatches_to_custom_logger_instance(self):
|
||||
async def test_dispatches_to_custom_logger_instance(self, monkeypatch: pytest.MonkeyPatch):
|
||||
mock_logger = MagicMock(spec=CustomLogger)
|
||||
mock_logger.async_log_audit_log_event = AsyncMock()
|
||||
litellm.audit_log_callbacks = [mock_logger]
|
||||
monkeypatch.setattr(litellm, "audit_log_callbacks", [mock_logger])
|
||||
|
||||
audit_log = _make_audit_log()
|
||||
await _dispatch_audit_log_to_callbacks(audit_log)
|
||||
|
|
@ -132,18 +129,18 @@ class TestDispatchAuditLogToCallbacks:
|
|||
assert payload["action"] == "created"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_dispatch_when_callbacks_empty(self):
|
||||
litellm.audit_log_callbacks = []
|
||||
async def test_no_dispatch_when_callbacks_empty(self, monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(litellm, "audit_log_callbacks", [])
|
||||
audit_log = _make_audit_log()
|
||||
# Should return immediately without error
|
||||
await _dispatch_audit_log_to_callbacks(audit_log)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_resolves_string_callback(self):
|
||||
async def test_resolves_string_callback(self, monkeypatch: pytest.MonkeyPatch):
|
||||
mock_logger = MagicMock(spec=CustomLogger)
|
||||
mock_logger.async_log_audit_log_event = AsyncMock()
|
||||
|
||||
litellm.audit_log_callbacks = ["s3_v2"]
|
||||
monkeypatch.setattr(litellm, "audit_log_callbacks", ["s3_v2"])
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_helpers.audit_logs._resolve_audit_log_callback",
|
||||
|
|
@ -156,13 +153,13 @@ class TestDispatchAuditLogToCallbacks:
|
|||
mock_logger.async_log_audit_log_event.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_nonblocking_on_callback_failure(self):
|
||||
async def test_nonblocking_on_callback_failure(self, monkeypatch: pytest.MonkeyPatch):
|
||||
"""Callback errors should not propagate."""
|
||||
mock_logger = MagicMock(spec=CustomLogger)
|
||||
mock_logger.async_log_audit_log_event = AsyncMock(
|
||||
side_effect=RuntimeError("boom")
|
||||
)
|
||||
litellm.audit_log_callbacks = [mock_logger]
|
||||
monkeypatch.setattr(litellm, "audit_log_callbacks", [mock_logger])
|
||||
|
||||
audit_log = _make_audit_log()
|
||||
# Should not raise
|
||||
|
|
@ -170,8 +167,8 @@ class TestDispatchAuditLogToCallbacks:
|
|||
await asyncio.sleep(0.1)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_skips_unresolvable_string_callback(self):
|
||||
litellm.audit_log_callbacks = ["nonexistent_callback"]
|
||||
async def test_skips_unresolvable_string_callback(self, monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setattr(litellm, "audit_log_callbacks", ["nonexistent_callback"])
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.management_helpers.audit_logs._resolve_audit_log_callback",
|
||||
|
|
@ -184,10 +181,10 @@ class TestDispatchAuditLogToCallbacks:
|
|||
|
||||
class TestCreateAuditLogForUpdateWithCallbacks:
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatches_to_callbacks_after_db_write(self):
|
||||
async def test_dispatches_to_callbacks_after_db_write(self, monkeypatch: pytest.MonkeyPatch):
|
||||
mock_logger = MagicMock(spec=CustomLogger)
|
||||
mock_logger.async_log_audit_log_event = AsyncMock()
|
||||
litellm.audit_log_callbacks = [mock_logger]
|
||||
monkeypatch.setattr(litellm, "audit_log_callbacks", [mock_logger])
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.premium_user", True),
|
||||
|
|
@ -206,10 +203,10 @@ class TestCreateAuditLogForUpdateWithCallbacks:
|
|||
mock_logger.async_log_audit_log_event.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_dispatch_when_not_premium(self):
|
||||
async def test_no_dispatch_when_not_premium(self, monkeypatch: pytest.MonkeyPatch):
|
||||
mock_logger = MagicMock(spec=CustomLogger)
|
||||
mock_logger.async_log_audit_log_event = AsyncMock()
|
||||
litellm.audit_log_callbacks = [mock_logger]
|
||||
monkeypatch.setattr(litellm, "audit_log_callbacks", [mock_logger])
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.premium_user", False),
|
||||
|
|
@ -224,10 +221,10 @@ class TestCreateAuditLogForUpdateWithCallbacks:
|
|||
mock_prisma.db.litellm_auditlog.create.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_dispatch_when_store_audit_logs_false(self):
|
||||
async def test_no_dispatch_when_store_audit_logs_false(self, monkeypatch: pytest.MonkeyPatch):
|
||||
mock_logger = MagicMock(spec=CustomLogger)
|
||||
mock_logger.async_log_audit_log_event = AsyncMock()
|
||||
litellm.audit_log_callbacks = [mock_logger]
|
||||
monkeypatch.setattr(litellm, "audit_log_callbacks", [mock_logger])
|
||||
|
||||
with patch("litellm.store_audit_logs", False):
|
||||
audit_log = _make_audit_log()
|
||||
|
|
@ -237,11 +234,11 @@ class TestCreateAuditLogForUpdateWithCallbacks:
|
|||
mock_logger.async_log_audit_log_event.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatches_even_when_prisma_client_is_none(self):
|
||||
async def test_dispatches_even_when_prisma_client_is_none(self, monkeypatch: pytest.MonkeyPatch):
|
||||
"""Callbacks should fire even if DB is unavailable."""
|
||||
mock_logger = MagicMock(spec=CustomLogger)
|
||||
mock_logger.async_log_audit_log_event = AsyncMock()
|
||||
litellm.audit_log_callbacks = [mock_logger]
|
||||
monkeypatch.setattr(litellm, "audit_log_callbacks", [mock_logger])
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.premium_user", True),
|
||||
|
|
@ -256,11 +253,11 @@ class TestCreateAuditLogForUpdateWithCallbacks:
|
|||
mock_logger.async_log_audit_log_event.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dispatches_even_when_db_write_fails(self):
|
||||
async def test_dispatches_even_when_db_write_fails(self, monkeypatch: pytest.MonkeyPatch):
|
||||
"""Callbacks should fire even if the DB write raises."""
|
||||
mock_logger = MagicMock(spec=CustomLogger)
|
||||
mock_logger.async_log_audit_log_event = AsyncMock()
|
||||
litellm.audit_log_callbacks = [mock_logger]
|
||||
monkeypatch.setattr(litellm, "audit_log_callbacks", [mock_logger])
|
||||
|
||||
with (
|
||||
patch("litellm.proxy.proxy_server.premium_user", True),
|
||||
|
|
@ -384,21 +381,21 @@ class TestS3AuditCallbackParamsDecoupling:
|
|||
S3Logger instance, distinct from the singleton serving normal logs."""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _isolate_caches_and_globals(self):
|
||||
def _isolate_caches_and_globals(self, monkeypatch: pytest.MonkeyPatch):
|
||||
from litellm.litellm_core_utils import litellm_logging as ll_logging
|
||||
from litellm.proxy.management_helpers import audit_logs as ll_audit_logs
|
||||
|
||||
original_s3 = litellm.s3_callback_params
|
||||
original_audit = getattr(litellm, "s3_audit_callback_params", None)
|
||||
monkeypatch.setattr(litellm, "s3_callback_params", litellm.s3_callback_params)
|
||||
monkeypatch.setattr(
|
||||
litellm, "s3_audit_callback_params", getattr(litellm, "s3_audit_callback_params", None)
|
||||
)
|
||||
ll_audit_logs._audit_log_callback_cache.clear()
|
||||
ll_logging._in_memory_loggers.clear()
|
||||
yield
|
||||
litellm.s3_callback_params = original_s3
|
||||
litellm.s3_audit_callback_params = original_audit
|
||||
ll_audit_logs._audit_log_callback_cache.clear()
|
||||
ll_logging._in_memory_loggers.clear()
|
||||
|
||||
def test_opt_in_constructs_separate_instance_with_audit_config(self):
|
||||
def test_opt_in_constructs_separate_instance_with_audit_config(self, monkeypatch: pytest.MonkeyPatch):
|
||||
"""Audit config set → audit resolver returns a fresh S3Logger pointing
|
||||
at the audit bucket, distinct from the normal-log singleton."""
|
||||
from litellm.integrations.s3_v2 import S3Logger
|
||||
|
|
@ -409,8 +406,8 @@ class TestS3AuditCallbackParamsDecoupling:
|
|||
_resolve_audit_log_callback,
|
||||
)
|
||||
|
||||
litellm.s3_callback_params = {"s3_bucket_name": "normal-bucket"}
|
||||
litellm.s3_audit_callback_params = {"s3_bucket_name": "audit-bucket"}
|
||||
monkeypatch.setattr(litellm, "s3_callback_params", {"s3_bucket_name": "normal-bucket"})
|
||||
monkeypatch.setattr(litellm, "s3_audit_callback_params", {"s3_bucket_name": "audit-bucket"})
|
||||
|
||||
with patch("asyncio.create_task"):
|
||||
audit_instance = _resolve_audit_log_callback("s3_v2")
|
||||
|
|
@ -426,7 +423,7 @@ class TestS3AuditCallbackParamsDecoupling:
|
|||
assert audit_instance.s3_bucket_name == "audit-bucket"
|
||||
assert normal_instance.s3_bucket_name == "normal-bucket"
|
||||
|
||||
def test_opt_out_preserves_singleton_behavior(self):
|
||||
def test_opt_out_preserves_singleton_behavior(self, monkeypatch: pytest.MonkeyPatch):
|
||||
"""No `s3_audit_callback_params` → audit and normal share the singleton
|
||||
(existing behavior, regression guard)."""
|
||||
from litellm.integrations.s3_v2 import S3Logger
|
||||
|
|
@ -437,8 +434,8 @@ class TestS3AuditCallbackParamsDecoupling:
|
|||
_resolve_audit_log_callback,
|
||||
)
|
||||
|
||||
litellm.s3_callback_params = {"s3_bucket_name": "shared-bucket"}
|
||||
litellm.s3_audit_callback_params = None
|
||||
monkeypatch.setattr(litellm, "s3_callback_params", {"s3_bucket_name": "shared-bucket"})
|
||||
monkeypatch.setattr(litellm, "s3_audit_callback_params", None)
|
||||
|
||||
with patch("asyncio.create_task"):
|
||||
normal_instance = _init_custom_logger_compatible_class(
|
||||
|
|
@ -452,7 +449,7 @@ class TestS3AuditCallbackParamsDecoupling:
|
|||
assert id(audit_instance) == id(normal_instance)
|
||||
assert audit_instance.s3_bucket_name == "shared-bucket"
|
||||
|
||||
def test_empty_dict_opts_in(self):
|
||||
def test_empty_dict_opts_in(self, monkeypatch: pytest.MonkeyPatch):
|
||||
"""`s3_audit_callback_params = {}` is opt-in (truthy-by-presence) and
|
||||
produces a separate instance with no bucket configured (env/IAM-only)."""
|
||||
from litellm.integrations.s3_v2 import S3Logger
|
||||
|
|
@ -463,8 +460,8 @@ class TestS3AuditCallbackParamsDecoupling:
|
|||
_resolve_audit_log_callback,
|
||||
)
|
||||
|
||||
litellm.s3_callback_params = {"s3_bucket_name": "normal-bucket"}
|
||||
litellm.s3_audit_callback_params = {}
|
||||
monkeypatch.setattr(litellm, "s3_callback_params", {"s3_bucket_name": "normal-bucket"})
|
||||
monkeypatch.setattr(litellm, "s3_audit_callback_params", {})
|
||||
|
||||
with patch("asyncio.create_task"):
|
||||
audit_instance = _resolve_audit_log_callback("s3_v2")
|
||||
|
|
@ -478,7 +475,7 @@ class TestS3AuditCallbackParamsDecoupling:
|
|||
assert audit_instance.s3_bucket_name is None
|
||||
assert normal_instance.s3_bucket_name == "normal-bucket"
|
||||
|
||||
def test_reset_audit_log_callback_cache_clears_audit_instance(self):
|
||||
def test_reset_audit_log_callback_cache_clears_audit_instance(self, monkeypatch: pytest.MonkeyPatch):
|
||||
"""`reset_audit_log_callback_cache()` must drop the cached audit
|
||||
instance so a config reload picks up the new params."""
|
||||
from litellm.proxy.management_helpers.audit_logs import (
|
||||
|
|
@ -487,7 +484,7 @@ class TestS3AuditCallbackParamsDecoupling:
|
|||
reset_audit_log_callback_cache,
|
||||
)
|
||||
|
||||
litellm.s3_audit_callback_params = {"s3_bucket_name": "first"}
|
||||
monkeypatch.setattr(litellm, "s3_audit_callback_params", {"s3_bucket_name": "first"})
|
||||
with patch("asyncio.create_task"):
|
||||
first = _resolve_audit_log_callback("s3_v2")
|
||||
assert first is not None and "s3_v2" in _audit_log_callback_cache
|
||||
|
|
@ -495,7 +492,7 @@ class TestS3AuditCallbackParamsDecoupling:
|
|||
reset_audit_log_callback_cache()
|
||||
assert "s3_v2" not in _audit_log_callback_cache
|
||||
|
||||
litellm.s3_audit_callback_params = {"s3_bucket_name": "second"}
|
||||
monkeypatch.setattr(litellm, "s3_audit_callback_params", {"s3_bucket_name": "second"})
|
||||
second = _resolve_audit_log_callback("s3_v2")
|
||||
assert second is not None
|
||||
assert id(second) != id(first)
|
||||
|
|
|
|||
|
|
@ -177,3 +177,107 @@ def test_project_io_token_limits_are_stored_in_metadata(request_type):
|
|||
|
||||
assert request.metadata == limits
|
||||
assert request.model_dump(exclude_none=True)["metadata"] == limits
|
||||
|
||||
|
||||
def test_a_jwt_issuer_must_pick_audience_validation_or_opt_out():
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.proxy._types import JWTIssuerConfig
|
||||
|
||||
with pytest.raises(ValidationError, match="must configure audience or set disable_audience_validation"):
|
||||
JWTIssuerConfig(issuer="https://issuer.example.com")
|
||||
|
||||
with pytest.raises(ValidationError, match="cannot set audience and disable_audience_validation"):
|
||||
JWTIssuerConfig(
|
||||
issuer="https://issuer.example.com",
|
||||
audience="litellm-proxy",
|
||||
disable_audience_validation=True,
|
||||
)
|
||||
|
||||
assert JWTIssuerConfig(issuer="https://issuer.example.com", audience="litellm-proxy").audience == "litellm-proxy"
|
||||
assert (
|
||||
JWTIssuerConfig(issuer="https://issuer.example.com", disable_audience_validation=True).audience
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
def test_a_jwt_issuer_rejects_a_field_it_does_not_define():
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.proxy._types import JWTIssuerConfig
|
||||
|
||||
with pytest.raises(ValidationError, match="Extra inputs are not permitted"):
|
||||
JWTIssuerConfig(issuer="https://issuer.example.com", audience="a", jwks_uri="https://issuer/jwks")
|
||||
|
||||
|
||||
def test_a_temp_budget_needs_both_halves_or_neither():
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.proxy._types import UpdateKeyRequest
|
||||
|
||||
with pytest.raises(ValidationError, match="temp_budget_increase and temp_budget_expiry must be set together"):
|
||||
UpdateKeyRequest(key="sk-1234", temp_budget_increase=10)
|
||||
|
||||
with pytest.raises(ValidationError, match="temp_budget_increase and temp_budget_expiry must be set together"):
|
||||
UpdateKeyRequest(key="sk-1234", temp_budget_expiry="2026-01-01")
|
||||
|
||||
both = UpdateKeyRequest(key="sk-1234", temp_budget_increase=10, temp_budget_expiry="2026-01-01")
|
||||
assert both.temp_budget_increase == 10
|
||||
|
||||
|
||||
def test_an_empty_max_budget_is_read_as_no_limit():
|
||||
from litellm.proxy._types import GenerateKeyRequest
|
||||
|
||||
assert GenerateKeyRequest(max_budget="").max_budget is None
|
||||
assert GenerateKeyRequest(max_budget=25).max_budget == 25
|
||||
|
||||
|
||||
def test_an_organization_member_can_only_take_a_role_the_organization_has():
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.proxy._types import LitellmUserRoles, OrganizationMemberUpdateRequest
|
||||
|
||||
with pytest.raises(ValidationError, match="Invalid role"):
|
||||
OrganizationMemberUpdateRequest(
|
||||
organization_id="org-1", user_id="user-1", role=LitellmUserRoles.PROXY_ADMIN
|
||||
)
|
||||
|
||||
allowed = OrganizationMemberUpdateRequest(
|
||||
organization_id="org-1", user_id="user-1", role=LitellmUserRoles.ORG_ADMIN
|
||||
)
|
||||
assert allowed.role == LitellmUserRoles.ORG_ADMIN
|
||||
|
||||
|
||||
def test_an_llm_backed_injection_check_needs_the_call_it_would_make():
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.proxy._types import LiteLLMPromptInjectionParams
|
||||
|
||||
for missing in ("llm_api_name", "llm_api_system_prompt", "llm_api_fail_call_string"):
|
||||
complete = {
|
||||
"llm_api_name": "gpt-4o",
|
||||
"llm_api_system_prompt": "is this an injection",
|
||||
"llm_api_fail_call_string": "yes",
|
||||
}
|
||||
del complete[missing]
|
||||
with pytest.raises(ValidationError, match=f"{missing} must be provided"):
|
||||
LiteLLMPromptInjectionParams(llm_api_check=True, **complete)
|
||||
|
||||
assert LiteLLMPromptInjectionParams(llm_api_check=False).llm_api_name is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"field, forged, default",
|
||||
[
|
||||
("mcp_admitted_user_subject", "someone-else", False),
|
||||
("mcp_source_team_rpm_limits", {"team-1": 10_000}, None),
|
||||
("mcp_session_resource_server_id", "server-1", None),
|
||||
("via_virtual_key", "sk-someone-elses-key", False),
|
||||
],
|
||||
)
|
||||
def test_a_server_only_marker_is_not_taken_from_the_caller(field, forged, default):
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
auth = UserAPIKeyAuth(api_key="sk-1234", **{field: forged})
|
||||
|
||||
assert getattr(auth, field) == default
|
||||
|
|
|
|||
|
|
@ -1644,3 +1644,193 @@ async def test_post_mcp_call_hook_skips_opted_out_guardrail(restore_callbacks):
|
|||
|
||||
assert guardrail.call_count == 0
|
||||
assert [item.text for item in returned.content] == ["jane@example.com"]
|
||||
|
||||
|
||||
FAILURE_USAGE_MODEL = "gpt-4o"
|
||||
ONE_USER_MESSAGE = [{"role": "user", "content": "hi"}]
|
||||
|
||||
|
||||
class _LoggingObj:
|
||||
def __init__(self, model_call_details):
|
||||
self.model_call_details = model_call_details
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"system_input, expected",
|
||||
[
|
||||
("be brief", "be brief"),
|
||||
([{"type": "text", "text": "a"}, {"type": "text", "text": "b"}], "ab"),
|
||||
(["a", {"text": "b"}], "ab"),
|
||||
([{"type": "image"}], ""),
|
||||
(None, ""),
|
||||
(17, ""),
|
||||
],
|
||||
)
|
||||
def test_a_system_prompt_reads_the_same_whatever_shape_it_arrived_in(system_input, expected):
|
||||
from litellm.proxy.utils import _system_prompt_text
|
||||
|
||||
assert _system_prompt_text(system_input) == expected
|
||||
|
||||
|
||||
def test_a_system_prompt_is_counted_on_top_of_the_request():
|
||||
from litellm.proxy.utils import _count_request_input_tokens
|
||||
|
||||
without = _count_request_input_tokens(FAILURE_USAGE_MODEL, "hello world", None)
|
||||
with_system = _count_request_input_tokens(FAILURE_USAGE_MODEL, "hello world", "be brief")
|
||||
|
||||
assert without > 0
|
||||
assert with_system > without
|
||||
|
||||
|
||||
def test_a_request_with_nothing_in_it_counts_zero():
|
||||
from litellm.proxy.utils import _count_request_input_tokens
|
||||
|
||||
assert _count_request_input_tokens(FAILURE_USAGE_MODEL, [], None) == 0
|
||||
assert _count_request_input_tokens(FAILURE_USAGE_MODEL, None, None) == 0
|
||||
|
||||
|
||||
def test_a_failed_dispatch_is_estimated_as_input_only():
|
||||
from litellm.proxy.utils import _count_request_input_tokens, _estimate_dispatched_failure_usage
|
||||
|
||||
usage = _estimate_dispatched_failure_usage(FAILURE_USAGE_MODEL, ONE_USER_MESSAGE, None)
|
||||
|
||||
assert usage is not None
|
||||
assert usage.prompt_tokens == _count_request_input_tokens(
|
||||
FAILURE_USAGE_MODEL, ONE_USER_MESSAGE, None
|
||||
)
|
||||
assert usage.completion_tokens == 0
|
||||
assert usage.total_tokens == usage.prompt_tokens
|
||||
|
||||
|
||||
@pytest.mark.parametrize("request_input", [[], object()])
|
||||
def test_nothing_is_estimated_when_there_is_nothing_to_count(request_input):
|
||||
from litellm.proxy.utils import _estimate_dispatched_failure_usage
|
||||
|
||||
assert _estimate_dispatched_failure_usage(FAILURE_USAGE_MODEL, request_input, None) is None
|
||||
|
||||
|
||||
def test_usage_the_stream_already_recovered_beats_an_estimate():
|
||||
from litellm.proxy.utils import _failure_usage_to_lift
|
||||
from litellm.types.utils import Usage
|
||||
|
||||
recovered = Usage(prompt_tokens=5, completion_tokens=7, total_tokens=12)
|
||||
|
||||
lifted = _failure_usage_to_lift(
|
||||
model_call_details={"combined_usage_object": recovered, "response_cost": 0.25},
|
||||
request_body={},
|
||||
dispatched=True,
|
||||
)
|
||||
|
||||
assert lifted == (recovered, 0.25)
|
||||
|
||||
|
||||
def test_a_request_that_reached_a_provider_bills_its_input_at_no_cost():
|
||||
from litellm.proxy.utils import _failure_usage_to_lift
|
||||
|
||||
lifted = _failure_usage_to_lift(
|
||||
model_call_details={
|
||||
"call_type": "acompletion",
|
||||
"model": FAILURE_USAGE_MODEL,
|
||||
"messages": ONE_USER_MESSAGE,
|
||||
},
|
||||
request_body={},
|
||||
dispatched=True,
|
||||
)
|
||||
|
||||
assert lifted is not None
|
||||
usage, response_cost = lifted
|
||||
assert usage.prompt_tokens > 0
|
||||
assert usage.completion_tokens == 0
|
||||
assert response_cost == 0.0
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"model_call_details, dispatched",
|
||||
[
|
||||
({"call_type": "acompletion", "model": FAILURE_USAGE_MODEL, "messages": ONE_USER_MESSAGE}, False),
|
||||
(
|
||||
{
|
||||
"litellm_no_upstream_llm_call": True,
|
||||
"call_type": "acompletion",
|
||||
"model": FAILURE_USAGE_MODEL,
|
||||
"messages": ONE_USER_MESSAGE,
|
||||
},
|
||||
True,
|
||||
),
|
||||
({"call_type": "afile_content", "model": FAILURE_USAGE_MODEL, "messages": ONE_USER_MESSAGE}, True),
|
||||
],
|
||||
ids=["never dispatched", "no upstream call", "call type has no input to price"],
|
||||
)
|
||||
def test_a_failure_that_cost_the_provider_nothing_lifts_nothing(model_call_details, dispatched):
|
||||
from litellm.proxy.utils import _failure_usage_to_lift
|
||||
|
||||
assert _failure_usage_to_lift(
|
||||
model_call_details=model_call_details, request_body={}, dispatched=dispatched
|
||||
) is None
|
||||
|
||||
|
||||
def test_the_no_upstream_call_key_the_module_uses_is_the_one_asserted_above():
|
||||
from litellm.constants import LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL
|
||||
|
||||
assert LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL == "litellm_no_upstream_llm_call"
|
||||
|
||||
|
||||
def test_the_dispatched_system_prompt_wins_over_the_one_in_the_request_body():
|
||||
from litellm.proxy.utils import _failure_usage_to_lift
|
||||
|
||||
def lift(model_call_details, request_body):
|
||||
lifted = _failure_usage_to_lift(
|
||||
model_call_details=model_call_details, request_body=request_body, dispatched=True
|
||||
)
|
||||
assert lifted is not None
|
||||
return lifted[0].prompt_tokens
|
||||
|
||||
base = {
|
||||
"call_type": "aanthropic_messages",
|
||||
"model": FAILURE_USAGE_MODEL,
|
||||
"messages": ONE_USER_MESSAGE,
|
||||
}
|
||||
long_system = "answer as briefly as you possibly can, in one short sentence"
|
||||
|
||||
from_body = lift(base, {"system": long_system})
|
||||
from_params = lift({**base, "optional_params": {"system": "x"}}, {"system": long_system})
|
||||
body_only_short = lift(base, {"system": "x"})
|
||||
|
||||
assert from_body > body_only_short
|
||||
assert from_params == body_only_short
|
||||
|
||||
|
||||
def test_a_failure_with_no_logging_object_lifts_nothing():
|
||||
from litellm.proxy.utils import _failure_fields_to_lift
|
||||
|
||||
assert dict(_failure_fields_to_lift({})) == {}
|
||||
assert dict(_failure_fields_to_lift({"litellm_logging_obj": _LoggingObj({})})) == {}
|
||||
|
||||
|
||||
def test_a_dispatched_failure_lifts_the_four_fields_the_spend_log_needs():
|
||||
from litellm.proxy.utils import _failure_fields_to_lift
|
||||
|
||||
lifted = _failure_fields_to_lift(
|
||||
{
|
||||
"litellm_logging_obj": _LoggingObj(
|
||||
{
|
||||
"first_api_call_start_time": 1700000000.0,
|
||||
"call_type": "acompletion",
|
||||
"model": FAILURE_USAGE_MODEL,
|
||||
"messages": ONE_USER_MESSAGE,
|
||||
"standard_logging_object": {"id": "log-1"},
|
||||
}
|
||||
)
|
||||
}
|
||||
)
|
||||
|
||||
assert set(lifted) == {
|
||||
"first_api_call_start_time",
|
||||
"combined_usage_object",
|
||||
"response_cost",
|
||||
"standard_logging_object",
|
||||
}
|
||||
assert lifted["first_api_call_start_time"] == 1700000000.0
|
||||
assert lifted["response_cost"] == 0.0
|
||||
assert lifted["combined_usage_object"].prompt_tokens > 0
|
||||
assert lifted["standard_logging_object"] == {"id": "log-1"}
|
||||
|
|
|
|||
|
|
@ -62,7 +62,7 @@ async def test_add_deployment_without_master_key():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_add_deployment_without_salt_key_or_master_key():
|
||||
async def test_add_deployment_without_salt_key_or_master_key(monkeypatch):
|
||||
"""
|
||||
Test that add_deployment() works when both master_key and LITELLM_SALT_KEY are None.
|
||||
|
||||
|
|
@ -70,55 +70,50 @@ async def test_add_deployment_without_salt_key_or_master_key():
|
|||
such as in a local/dev environment or when just saving spend logs.
|
||||
"""
|
||||
# Remove LITELLM_SALT_KEY from environment
|
||||
old_salt_key = os.environ.pop("LITELLM_SALT_KEY", None)
|
||||
monkeypatch.delenv("LITELLM_SALT_KEY", raising=False)
|
||||
|
||||
try:
|
||||
# Set master_key to None
|
||||
with patch("litellm.proxy.proxy_server.master_key", None):
|
||||
# Mock the required dependencies
|
||||
mock_prisma_client = MagicMock(spec=PrismaClient)
|
||||
mock_prisma_client.db = MagicMock()
|
||||
mock_prisma_client.db.litellm_config = MagicMock()
|
||||
mock_prisma_client.db.litellm_config.find_first = AsyncMock(
|
||||
return_value=None
|
||||
# Set master_key to None
|
||||
with patch("litellm.proxy.proxy_server.master_key", None):
|
||||
# Mock the required dependencies
|
||||
mock_prisma_client = MagicMock(spec=PrismaClient)
|
||||
mock_prisma_client.db = MagicMock()
|
||||
mock_prisma_client.db.litellm_config = MagicMock()
|
||||
mock_prisma_client.db.litellm_config.find_first = AsyncMock(
|
||||
return_value=None
|
||||
)
|
||||
|
||||
mock_proxy_logging = MagicMock(spec=ProxyLogging)
|
||||
|
||||
# Create ProxyConfig instance
|
||||
proxy_config = ProxyConfig()
|
||||
|
||||
# Mock the internal methods
|
||||
proxy_config._should_load_db_object = MagicMock(return_value=False)
|
||||
proxy_config._init_non_llm_objects_in_db = AsyncMock()
|
||||
|
||||
# This should NOT raise an exception
|
||||
try:
|
||||
await proxy_config.add_deployment(
|
||||
prisma_client=mock_prisma_client,
|
||||
proxy_logging_obj=mock_proxy_logging,
|
||||
)
|
||||
|
||||
mock_proxy_logging = MagicMock(spec=ProxyLogging)
|
||||
|
||||
# Create ProxyConfig instance
|
||||
proxy_config = ProxyConfig()
|
||||
|
||||
# Mock the internal methods
|
||||
proxy_config._should_load_db_object = MagicMock(return_value=False)
|
||||
proxy_config._init_non_llm_objects_in_db = AsyncMock()
|
||||
|
||||
# This should NOT raise an exception
|
||||
try:
|
||||
await proxy_config.add_deployment(
|
||||
prisma_client=mock_prisma_client,
|
||||
proxy_logging_obj=mock_proxy_logging,
|
||||
assert True
|
||||
except ValueError as e:
|
||||
if "Master key is not initialized" in str(
|
||||
e
|
||||
) or "Encryption key is not initialized" in str(e):
|
||||
pytest.fail(
|
||||
f"add_deployment raised ValueError about encryption key: {e}"
|
||||
)
|
||||
assert True
|
||||
except ValueError as e:
|
||||
if "Master key is not initialized" in str(
|
||||
e
|
||||
) or "Encryption key is not initialized" in str(e):
|
||||
pytest.fail(
|
||||
f"add_deployment raised ValueError about encryption key: {e}"
|
||||
)
|
||||
raise
|
||||
except Exception as e:
|
||||
if "Master key is not initialized" in str(
|
||||
e
|
||||
) or "Encryption key is not initialized" in str(e):
|
||||
pytest.fail(
|
||||
f"add_deployment raised exception about encryption key: {e}"
|
||||
)
|
||||
raise
|
||||
finally:
|
||||
# Restore LITELLM_SALT_KEY if it was set
|
||||
if old_salt_key:
|
||||
os.environ["LITELLM_SALT_KEY"] = old_salt_key
|
||||
raise
|
||||
except Exception as e:
|
||||
if "Master key is not initialized" in str(
|
||||
e
|
||||
) or "Encryption key is not initialized" in str(e):
|
||||
pytest.fail(
|
||||
f"add_deployment raised exception about encryption key: {e}"
|
||||
)
|
||||
raise
|
||||
|
||||
|
||||
def test_add_deployment_sync_without_master_key():
|
||||
|
|
|
|||
75
tests/test_litellm/test_azure_audio_price_aliases.py
Normal file
75
tests/test_litellm/test_azure_audio_price_aliases.py
Normal file
|
|
@ -0,0 +1,75 @@
|
|||
"""Undated azure aliases for the audio models must exist and match their dated
|
||||
variants. Azure deployments are commonly created under an admin-chosen name, so
|
||||
the served model name means nothing to the cost lookup and `base_model:
|
||||
azure/gpt-audio-mini` is what prices the call. That key resolved to nothing, the
|
||||
lookup raised "This model isn't mapped yet", and the proxy logged the request at
|
||||
$0. Issue #33170."""
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
|
||||
pytestmark = pytest.mark.usefixtures("local_model_cost_map")
|
||||
|
||||
|
||||
COST_FIELDS = (
|
||||
"input_cost_per_token",
|
||||
"output_cost_per_token",
|
||||
"input_cost_per_audio_token",
|
||||
"output_cost_per_audio_token",
|
||||
)
|
||||
|
||||
ALIAS_PAIRS = (
|
||||
("azure/gpt-audio-mini", "azure/gpt-audio-mini-2025-10-06"),
|
||||
("azure/gpt-realtime-mini", "azure/gpt-realtime-mini-2025-10-06"),
|
||||
)
|
||||
|
||||
|
||||
def _load_root_cost_map() -> dict:
|
||||
root_map_path = Path(__file__).parents[2] / "model_prices_and_context_window.json"
|
||||
with open(root_map_path) as f:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("undated, dated", ALIAS_PAIRS)
|
||||
def test_undated_azure_audio_alias_matches_dated_entry(undated, dated):
|
||||
undated_info = litellm.get_model_info(undated)
|
||||
dated_info = litellm.get_model_info(dated)
|
||||
|
||||
for field in COST_FIELDS:
|
||||
assert undated_info.get(field) == dated_info.get(field), field
|
||||
assert (undated_info.get(field) or 0) > 0, f"{undated}.{field} must be non-zero"
|
||||
|
||||
assert undated_info.get("litellm_provider") == "azure"
|
||||
assert undated_info.get("mode") == dated_info.get("mode")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("undated, dated", ALIAS_PAIRS)
|
||||
def test_undated_azure_audio_alias_is_exact_mirror(undated, dated):
|
||||
"""The undated alias must be a byte-for-byte mirror of its dated entry, covering
|
||||
every field (incl. realtime-specific cache/audio cost keys) so any future drift
|
||||
between the pair is caught, not just the core COST_FIELDS."""
|
||||
model_map = litellm.model_cost
|
||||
assert undated in model_map, f"{undated} missing from model cost map"
|
||||
assert model_map[undated] == model_map[dated], (
|
||||
f"{undated} must exactly mirror {dated}; "
|
||||
f"diff keys: {[k for k in set(model_map[undated]) | set(model_map[dated]) if model_map[undated].get(k) != model_map[dated].get(k)]}"
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("undated, dated", ALIAS_PAIRS)
|
||||
def test_undated_azure_audio_alias_is_in_the_root_cost_map(undated, dated):
|
||||
"""`local_model_cost_map` pins `litellm.model_cost` to the packaged backup, but a
|
||||
proxy left on its defaults fetches the root map instead, and that is the copy
|
||||
that ships to the CDN. An alias added to only one of the two files still bills
|
||||
$0 for every proxy reading the other, which is the very bug this file guards, so
|
||||
assert the root map directly and assert the two files agree."""
|
||||
root_map = _load_root_cost_map()
|
||||
assert undated in root_map, f"{undated} missing from the root cost map"
|
||||
assert root_map[undated] == root_map[dated], f"{undated} must exactly mirror {dated} in the root cost map"
|
||||
assert root_map[undated] == litellm.model_cost[undated], (
|
||||
f"{undated} differs between the root cost map and the packaged backup"
|
||||
)
|
||||
|
|
@ -3774,4 +3774,4 @@ def test_completion_cost_prices_anthropic_shaped_cache_read_tokens():
|
|||
custom_llm_provider="openai",
|
||||
)
|
||||
|
||||
assert cost == pytest.approx(3 * 5e-6 + 4014 * 5e-7 + 5 * 3e-5, rel=1e-9)
|
||||
assert cost == pytest.approx(3 * 4e-6 + 4014 * 4e-7 + 5 * 2e-5, rel=1e-9)
|
||||
|
|
|
|||
|
|
@ -144,20 +144,16 @@ def test_acount_tokens_api_error_falls_back():
|
|||
assert result.total_tokens > 0
|
||||
|
||||
|
||||
def test_acount_tokens_no_api_key_falls_back():
|
||||
def test_acount_tokens_no_api_key_falls_back(monkeypatch):
|
||||
"""Test that missing API key falls back to local counting."""
|
||||
env_backup = os.environ.pop("OPENAI_API_KEY", None)
|
||||
try:
|
||||
result = asyncio.run(
|
||||
litellm.acount_tokens(
|
||||
model="openai/gpt-4o",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
)
|
||||
monkeypatch.delenv("OPENAI_API_KEY", raising=False)
|
||||
result = asyncio.run(
|
||||
litellm.acount_tokens(
|
||||
model="openai/gpt-4o",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
)
|
||||
)
|
||||
|
||||
# Should fall back to local tokenizer since no API key
|
||||
assert result.total_tokens > 0
|
||||
assert result.tokenizer_type == "local_tokenizer"
|
||||
finally:
|
||||
if env_backup:
|
||||
os.environ["OPENAI_API_KEY"] = env_backup
|
||||
# Should fall back to local tokenizer since no API key
|
||||
assert result.total_tokens > 0
|
||||
assert result.tokenizer_type == "local_tokenizer"
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ from litellm.types.utils import ModelInfoBase
|
|||
REALTIME_ONLY_GPT_MODELS = (
|
||||
"azure/gpt-realtime-2025-08-28",
|
||||
"azure/gpt-realtime-1.5-2026-02-23",
|
||||
"azure/gpt-realtime-mini",
|
||||
"azure/gpt-realtime-mini-2025-10-06",
|
||||
"gpt-realtime",
|
||||
"gpt-realtime-1.5",
|
||||
|
|
|
|||
|
|
@ -2801,3 +2801,109 @@ def test_anthropic_oauth_credential_does_not_persist_into_next_provider_hop():
|
|||
leaked = [name for name, value in shared_headers.items() if value == _SUBSCRIPTION_OAUTH_CREDENTIAL]
|
||||
assert leaked == []
|
||||
assert "anthropic-version" not in shared_headers
|
||||
|
||||
|
||||
STREAM_COST_MODEL = "gpt-4o"
|
||||
STREAMED_USAGE = {"prompt_tokens": 137, "completion_tokens": 42, "total_tokens": 179}
|
||||
|
||||
|
||||
def _text_chunk(content, finish_reason=None, usage=None):
|
||||
chunk = {
|
||||
"id": "chatcmpl-stream-cost",
|
||||
"object": "chat.completion.chunk",
|
||||
"created": 1700000000,
|
||||
"model": STREAM_COST_MODEL,
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"delta": {"role": "assistant", "content": content},
|
||||
"finish_reason": finish_reason,
|
||||
}
|
||||
],
|
||||
}
|
||||
if usage is not None:
|
||||
chunk["usage"] = usage
|
||||
return chunk
|
||||
|
||||
|
||||
def _priced_at(prompt_tokens, completion_tokens):
|
||||
prices = litellm.model_cost[STREAM_COST_MODEL]
|
||||
return (
|
||||
prompt_tokens * prices["input_cost_per_token"]
|
||||
+ completion_tokens * prices["output_cost_per_token"]
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def local_cost_map(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
|
||||
|
||||
def test_a_streamed_response_bills_the_usage_the_provider_reported(local_cost_map):
|
||||
rebuilt = litellm.stream_chunk_builder(
|
||||
chunks=[
|
||||
_text_chunk("Hello"),
|
||||
_text_chunk(" there"),
|
||||
_text_chunk(None, finish_reason="stop", usage=STREAMED_USAGE),
|
||||
],
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
)
|
||||
|
||||
assert rebuilt.choices[0].message.content == "Hello there"
|
||||
assert rebuilt.usage.prompt_tokens == STREAMED_USAGE["prompt_tokens"]
|
||||
assert rebuilt.usage.completion_tokens == STREAMED_USAGE["completion_tokens"]
|
||||
|
||||
cost = litellm.completion_cost(completion_response=rebuilt, model=STREAM_COST_MODEL)
|
||||
|
||||
assert cost == pytest.approx(_priced_at(137, 42))
|
||||
assert cost == pytest.approx(0.0007625)
|
||||
|
||||
|
||||
def test_streaming_and_not_streaming_bill_the_same_usage_the_same(local_cost_map):
|
||||
rebuilt = litellm.stream_chunk_builder(
|
||||
chunks=[
|
||||
_text_chunk("Hello"),
|
||||
_text_chunk(" there"),
|
||||
_text_chunk(None, finish_reason="stop", usage=STREAMED_USAGE),
|
||||
],
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
)
|
||||
whole = litellm.ModelResponse(
|
||||
id="chatcmpl-stream-cost",
|
||||
model=STREAM_COST_MODEL,
|
||||
object="chat.completion",
|
||||
created=1700000000,
|
||||
choices=[
|
||||
{
|
||||
"index": 0,
|
||||
"message": {"role": "assistant", "content": "Hello there"},
|
||||
"finish_reason": "stop",
|
||||
}
|
||||
],
|
||||
usage=STREAMED_USAGE,
|
||||
)
|
||||
|
||||
assert litellm.completion_cost(
|
||||
completion_response=rebuilt, model=STREAM_COST_MODEL
|
||||
) == pytest.approx(litellm.completion_cost(completion_response=whole, model=STREAM_COST_MODEL))
|
||||
|
||||
|
||||
def test_a_stream_that_reported_no_usage_is_still_billed(local_cost_map):
|
||||
rebuilt = litellm.stream_chunk_builder(
|
||||
chunks=[
|
||||
_text_chunk("Hello"),
|
||||
_text_chunk(" there"),
|
||||
_text_chunk(None, finish_reason="stop"),
|
||||
],
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
)
|
||||
|
||||
assert rebuilt.usage.prompt_tokens > 0
|
||||
assert rebuilt.usage.completion_tokens > 0
|
||||
|
||||
cost = litellm.completion_cost(completion_response=rebuilt, model=STREAM_COST_MODEL)
|
||||
|
||||
assert cost > 0
|
||||
assert cost == pytest.approx(
|
||||
_priced_at(rebuilt.usage.prompt_tokens, rebuilt.usage.completion_tokens)
|
||||
)
|
||||
|
|
|
|||
139
tests/test_litellm/test_mutation_report.py
Normal file
139
tests/test_litellm/test_mutation_report.py
Normal file
|
|
@ -0,0 +1,139 @@
|
|||
"""Tests for scripts/mutation_report.py.
|
||||
|
||||
The report is the only thing anyone reads after a mutation run, so the one thing it
|
||||
must never do is describe a run that produced nothing as a run that killed everything.
|
||||
`render` decides that wording and `get_survivors` supplies the evidence for it, so both
|
||||
are tested directly.
|
||||
"""
|
||||
|
||||
import importlib.util
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
_REPO_ROOT = Path(__file__).resolve().parents[2]
|
||||
_MODULE_PATH = _REPO_ROOT / "scripts" / "mutation_report.py"
|
||||
_spec = importlib.util.spec_from_file_location("mutation_report", _MODULE_PATH)
|
||||
report = importlib.util.module_from_spec(_spec)
|
||||
sys.modules[_spec.name] = report
|
||||
_spec.loader.exec_module(report)
|
||||
|
||||
_CONFIG = {"paths_to_mutate": ["litellm/proxy/management_endpoints/"], "tests_dir": ["tests/"]}
|
||||
|
||||
|
||||
def test_a_run_that_reported_nothing_is_not_a_clean_sweep():
|
||||
rendered = report.render(_CONFIG, report.MutmutResults(survivors=(), reported=0), None)
|
||||
|
||||
assert "not a passing score" in rendered
|
||||
assert "caught every mutation" not in rendered
|
||||
|
||||
|
||||
def test_a_run_that_killed_every_mutant_says_so():
|
||||
rendered = report.render(
|
||||
_CONFIG, report.MutmutResults(survivors=(), reported=0), {"killed": 48, "survived": 0}
|
||||
)
|
||||
|
||||
assert "caught every mutation" in rendered
|
||||
assert "not a passing score" not in rendered
|
||||
|
||||
|
||||
def test_stats_counting_survivors_results_never_listed_is_not_a_clean_sweep():
|
||||
rendered = report.render(
|
||||
_CONFIG, report.MutmutResults(survivors=(), reported=0), {"killed": 48, "survived": 3}
|
||||
)
|
||||
|
||||
assert "not a passing score" in rendered
|
||||
assert "caught every mutation" not in rendered
|
||||
assert "3 surviving mutant(s)" in rendered
|
||||
|
||||
|
||||
def test_mutants_that_never_reached_the_tests_are_not_a_clean_sweep():
|
||||
rendered = report.render(
|
||||
_CONFIG,
|
||||
report.MutmutResults(survivors=(), reported=0),
|
||||
{"killed": 48, "survived": 0, "no_tests": 4, "timeout": 1},
|
||||
)
|
||||
|
||||
assert "not a passing score" in rendered
|
||||
assert "caught every mutation" not in rendered
|
||||
assert "4 no tests" in rendered
|
||||
assert "1 timeout" in rendered
|
||||
|
||||
|
||||
def test_a_status_the_reporter_has_never_met_still_blocks_a_clean_sweep():
|
||||
rendered = report.render(
|
||||
_CONFIG,
|
||||
report.MutmutResults(survivors=(), reported=0),
|
||||
{"killed": 48, "survived": 0, "check_was_interrupted_by_user": 2},
|
||||
)
|
||||
|
||||
assert "not a passing score" in rendered
|
||||
assert "caught every mutation" not in rendered
|
||||
assert "2 check was interrupted by user" in rendered
|
||||
|
||||
|
||||
def test_no_survivors_without_a_kill_is_not_a_clean_sweep():
|
||||
rendered = report.render(
|
||||
_CONFIG, report.MutmutResults(survivors=(), reported=48), {"killed": 0, "survived": 0}
|
||||
)
|
||||
|
||||
assert "not a passing score" in rendered
|
||||
assert "caught every mutation" not in rendered
|
||||
|
||||
|
||||
def test_no_survivors_and_no_stats_cannot_claim_a_sweep():
|
||||
"""`mutmut results` never lists killed mutants, so with the stats file missing an
|
||||
empty survivor list is equally consistent with a perfect run and a dead one."""
|
||||
rendered = report.render(_CONFIG, report.MutmutResults(survivors=(), reported=48), None)
|
||||
|
||||
assert "not a passing score" in rendered
|
||||
assert "caught every mutation" not in rendered
|
||||
|
||||
|
||||
def test_survivors_are_read_out_of_the_verdicts_they_came_with(monkeypatch):
|
||||
class _Proc:
|
||||
stdout = (
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.x_1: killed\n"
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.x_2: survived\n"
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.x_3: no tests\n"
|
||||
"not a verdict line at all\n"
|
||||
)
|
||||
|
||||
monkeypatch.setattr(report.subprocess, "run", lambda *a, **k: _Proc())
|
||||
|
||||
results = report.get_survivors()
|
||||
|
||||
assert results.survivors == (
|
||||
"litellm.proxy.management_endpoints.key_management_endpoints.x_2",
|
||||
)
|
||||
assert results.reported == 3
|
||||
|
||||
|
||||
def test_every_multi_word_verdict_mutmut_can_emit_still_counts(monkeypatch):
|
||||
class _Proc:
|
||||
stdout = "".join(
|
||||
f"litellm.proxy.management_endpoints.key_management_endpoints.x_{i}: {verdict}\n"
|
||||
for i, verdict in enumerate(
|
||||
(
|
||||
"no tests",
|
||||
"not checked",
|
||||
"caught by type check",
|
||||
"check was interrupted by user",
|
||||
)
|
||||
)
|
||||
)
|
||||
|
||||
monkeypatch.setattr(report.subprocess, "run", lambda *a, **k: _Proc())
|
||||
|
||||
results = report.get_survivors()
|
||||
|
||||
assert results.survivors == ()
|
||||
assert results.reported == 4
|
||||
|
||||
|
||||
def test_an_empty_mutmut_results_reports_nothing_rather_than_zero_survivors(monkeypatch):
|
||||
class _Proc:
|
||||
stdout = ""
|
||||
|
||||
monkeypatch.setattr(report.subprocess, "run", lambda *a, **k: _Proc())
|
||||
|
||||
assert report.get_survivors() == report.MutmutResults(survivors=(), reported=0)
|
||||
|
|
@ -318,7 +318,7 @@ def test_register_model_strips_none_litellm_provider_from_get_model_info(monkeyp
|
|||
litellm.model_cost.pop(model_key, None)
|
||||
|
||||
|
||||
def test_register_model_inherits_builtin_cache_pricing_for_unmapped_key():
|
||||
def test_register_model_inherits_builtin_cache_pricing_for_unmapped_key(monkeypatch):
|
||||
"""Registering a custom override under a key shape that
|
||||
``get_model_info`` cannot resolve (e.g. a triple provider prefix like
|
||||
``bedrock/bedrock/bedrock/us.anthropic.claude-sonnet-4-6``; a double
|
||||
|
|
@ -338,7 +338,7 @@ def test_register_model_inherits_builtin_cache_pricing_for_unmapped_key():
|
|||
from litellm.types.utils import PromptTokensDetailsWrapper, Usage
|
||||
|
||||
original_model_cost = litellm.model_cost
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
builtin_key = "us.anthropic.claude-sonnet-4-6"
|
||||
|
|
|
|||
|
|
@ -672,8 +672,8 @@ def test_all_model_configs():
|
|||
) == {"max_output_tokens": 10}
|
||||
|
||||
|
||||
def test_anthropic_web_search_in_model_info():
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
def test_anthropic_web_search_in_model_info(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
supported_models = [
|
||||
|
|
@ -1193,11 +1193,11 @@ def test_max_tokens_consistency():
|
|||
raise AssertionError(error_msg)
|
||||
|
||||
|
||||
def test_get_model_info_gemini():
|
||||
def test_get_model_info_gemini(monkeypatch):
|
||||
"""
|
||||
Tests if ALL gemini models have 'tpm' and 'rpm' in the model info
|
||||
"""
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
model_map = litellm.model_cost
|
||||
|
|
@ -1252,8 +1252,8 @@ def test_get_model_info_bedrock_double_provider_prefix_resolves(local_model_cost
|
|||
assert info["key"] == "us.anthropic.claude-sonnet-4-6"
|
||||
|
||||
|
||||
def test_openai_models_in_model_info():
|
||||
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
|
||||
def test_openai_models_in_model_info(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="")
|
||||
|
||||
model_map = litellm.model_cost
|
||||
|
|
@ -1408,7 +1408,7 @@ for commitment in BEDROCK_COMMITMENTS:
|
|||
print("block_list", block_list)
|
||||
|
||||
|
||||
def test_supports_computer_use_utility():
|
||||
def test_supports_computer_use_utility(monkeypatch):
|
||||
"""
|
||||
Tests the litellm.utils.supports_computer_use utility function.
|
||||
"""
|
||||
|
|
@ -1420,7 +1420,7 @@ def test_supports_computer_use_utility():
|
|||
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"
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
litellm.model_cost = litellm.get_model_cost_map(url="") # Load with local/backup
|
||||
|
||||
try:
|
||||
|
|
@ -1438,7 +1438,7 @@ def test_supports_computer_use_utility():
|
|||
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
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", original_env_var)
|
||||
|
||||
if original_model_cost is not None:
|
||||
litellm.model_cost = original_model_cost
|
||||
|
|
@ -1446,13 +1446,13 @@ def test_supports_computer_use_utility():
|
|||
delattr(litellm, "model_cost")
|
||||
|
||||
|
||||
def test_get_model_info_shows_supports_computer_use():
|
||||
def test_get_model_info_shows_supports_computer_use(monkeypatch):
|
||||
"""
|
||||
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"
|
||||
monkeypatch.setenv("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="")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue