chore: merge main into litellm_responses_precall_block_stream
Some checks failed
LiteLLM Rust / rust-lint (push) Has been cancelled
LiteLLM Rust / rust-test (push) Has been cancelled
LiteLLM Rust / rust-wheel (push) Has been cancelled

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-09-24 11:20:13 +00:00
commit 1fd0658a88
101 changed files with 10853 additions and 280 deletions

View file

@ -12,6 +12,9 @@ parameters:
migration_source_sha:
type: string
default: ""
routing_parity_base:
type: string
default: ""
orbs:
codecov: codecov/codecov@4.0.1
node: circleci/node@5.1.0 # Add this line to declare the node orb
@ -176,6 +179,9 @@ commands:
image:
type: string
default: postgres:14@sha256:6a70deda415ec296f977890e11aba04a0db9f632a362e3fce45e845e3db74f26
server_args:
type: string
default: ""
steps:
- run:
name: Start PostgreSQL
@ -186,7 +192,7 @@ commands:
-e POSTGRES_PASSWORD=postgres \
-e POSTGRES_DB=<< parameters.db_name >> \
-p 5432:5432 \
<< parameters.image >>
<< parameters.image >> << parameters.server_args >>
- wait_for_service:
url: tcp://localhost:5432
timeout: "60"
@ -3108,6 +3114,10 @@ jobs:
parameters:
suite:
type: string
mode:
type: enum
enum: [standard, replica]
default: standard
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
@ -3142,18 +3152,19 @@ jobs:
command: cd ui/litellm-dashboard && NEXT_TELEMETRY_DISABLED=1 npm run build
- start_postgres:
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
server_args: "-c shared_preload_libraries=pg_stat_statements -c pg_stat_statements.track=all -c pg_stat_statements.max=20000"
- start_redis
- run:
name: Run owned integration contracts
command: bash .circleci/scripts/run_integration.sh << parameters.suite >>
command: bash .circleci/scripts/run_integration.sh << parameters.suite >> << parameters.mode >>
no_output_timeout: 15m
- run:
name: Stop owned database and Redis
when: always
command: |
mkdir -p test-results/integration-<< parameters.suite >>
docker logs postgres-db > test-results/integration-<< parameters.suite >>/postgres.log 2>&1 || true
docker logs redis-cache > test-results/integration-<< parameters.suite >>/redis.log 2>&1 || true
mkdir -p test-results/services-<< parameters.suite >>-<< parameters.mode >>
docker logs postgres-db > test-results/services-<< parameters.suite >>-<< parameters.mode >>/postgres.log 2>&1 || true
docker logs redis-cache > test-results/services-<< parameters.suite >>-<< parameters.mode >>/redis.log 2>&1 || true
docker rm -f postgres-db redis-cache
test -z "$(docker ps -aq --filter name=postgres-db --filter name=redis-cache)"
- store_test_results:
@ -3161,6 +3172,76 @@ jobs:
- store_artifacts:
path: test-results
routing_parity:
parameters:
suite:
type: string
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
steps:
- setup_litellm_test_deps
- run:
name: Check out base product code
environment:
ROUTING_PARITY_BASE: << pipeline.parameters.routing_parity_base >>
command: |
[[ "$ROUTING_PARITY_BASE" =~ ^[0-9a-f]{40}$ ]] || exit 1
git fetch --depth 1 origin "$ROUTING_PARITY_BASE"
git rm -r -f --quiet litellm enterprise litellm-proxy-extras
git checkout "$ROUTING_PARITY_BASE" -- litellm enterprise litellm-proxy-extras
git reset --quiet
test -f litellm/rust_bridge/_native.abi3.so
- start_postgres:
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
server_args: "-c shared_preload_libraries=pg_stat_statements -c pg_stat_statements.track=all -c pg_stat_statements.max=20000"
- start_redis
- run:
name: Run base side
command: bash .circleci/scripts/run_integration.sh << parameters.suite >> parity base
no_output_timeout: 15m
- run:
name: Stop base database and Redis
when: always
command: |
mkdir -p test-results/services-<< parameters.suite >>-parity-base
docker logs postgres-db > test-results/services-<< parameters.suite >>-parity-base/postgres.log 2>&1 || true
docker logs redis-cache > test-results/services-<< parameters.suite >>-parity-base/redis.log 2>&1 || true
docker rm -f postgres-db redis-cache
test -z "$(docker ps -aq --filter name=postgres-db --filter name=redis-cache)"
- run:
name: Check out head product code
command: |
git rm -r -f --quiet litellm enterprise litellm-proxy-extras
git checkout "$CIRCLE_SHA1" -- litellm enterprise litellm-proxy-extras
git reset --quiet
test -f litellm/rust_bridge/_native.abi3.so
- start_postgres:
image: postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
server_args: "-c shared_preload_libraries=pg_stat_statements -c pg_stat_statements.track=all -c pg_stat_statements.max=20000"
- start_redis
- run:
name: Run head side
command: bash .circleci/scripts/run_integration.sh << parameters.suite >> parity head
no_output_timeout: 15m
- run:
name: Stop head database and Redis
when: always
command: |
mkdir -p test-results/services-<< parameters.suite >>-parity-head
docker logs postgres-db > test-results/services-<< parameters.suite >>-parity-head/postgres.log 2>&1 || true
docker logs redis-cache > test-results/services-<< parameters.suite >>-parity-head/redis.log 2>&1 || true
docker rm -f postgres-db redis-cache
test -z "$(docker ps -aq --filter name=postgres-db --filter name=redis-cache)"
- run:
name: Compare routing parity
command: PYTHONPATH="$PWD/tests" .venv/bin/python -m integration._support.routing check test-results/parity-<< parameters.suite >>/base test-results/parity-<< parameters.suite >>/head
- store_test_results:
path: test-results
- store_artifacts:
path: test-results
unit:
machine:
image: ubuntu-2204:2024.04.1
@ -3224,8 +3305,22 @@ workflows:
branches:
only: main
jobs: *migration_jobs
routing_parity:
when:
not:
equal: ["", << pipeline.parameters.routing_parity_base >>]
jobs:
- routing_parity:
name: routing-parity-<< matrix.suite >>
matrix:
parameters:
suite: [management, accounting, database, providers, extensions, cost, mcp]
integration:
unless: << pipeline.parameters.run_migration_tests >>
unless:
or:
- << pipeline.parameters.run_migration_tests >>
- not:
equal: ["", << pipeline.parameters.routing_parity_base >>]
jobs:
- integration_contracts:
name: integration-<< matrix.suite >>
@ -3237,8 +3332,23 @@ workflows:
only:
- main
- /litellm_.*/
- integration_contracts:
name: integration-<< matrix.suite >>-replica
matrix:
parameters:
suite: [management, database]
mode: [replica]
filters:
branches:
only:
- main
- /litellm_.*/
build_and_test:
unless: << pipeline.parameters.run_migration_tests >>
unless:
or:
- << pipeline.parameters.run_migration_tests >>
- not:
equal: ["", << pipeline.parameters.routing_parity_base >>]
jobs:
- using_litellm_on_windows:
filters: &main_branches

View file

@ -0,0 +1,52 @@
from __future__ import annotations
import os
from typing import Final
from urllib.parse import urlsplit, urlunsplit
import psycopg
DATABASE_URL: Final = os.environ["DATABASE_URL"]
def postgres_url() -> str:
parsed: Final = urlsplit(DATABASE_URL)
return urlunsplit(parsed._replace(path="/postgres"))
def main() -> None:
with psycopg.connect(postgres_url(), autocommit=True) as admin:
admin.execute("CREATE EXTENSION IF NOT EXISTS pg_stat_statements")
admin.execute("CREATE ROLE litellm_writer LOGIN PASSWORD 'litellm-writer' NOSUPERUSER")
admin.execute("CREATE ROLE litellm_reader LOGIN PASSWORD 'litellm-reader' NOSUPERUSER NOINHERIT")
admin.execute("ALTER ROLE litellm_reader SET default_transaction_read_only = on")
admin.execute("ALTER DATABASE circle_test OWNER TO litellm_writer")
admin.execute("GRANT CONNECT ON DATABASE circle_test TO litellm_reader")
with psycopg.connect(DATABASE_URL, autocommit=True) as admin:
admin.execute("GRANT USAGE ON SCHEMA public TO litellm_reader")
admin.execute(
"ALTER DEFAULT PRIVILEGES FOR ROLE litellm_writer IN SCHEMA public GRANT SELECT ON TABLES TO litellm_reader"
)
admin.execute("GRANT SELECT ON ALL TABLES IN SCHEMA public TO litellm_reader")
parsed: Final = urlsplit(DATABASE_URL)
reader_url: Final = urlunsplit(
parsed._replace(netloc=f"litellm_reader:litellm-reader@{parsed.hostname}:{parsed.port}")
)
writer_url: Final = urlunsplit(
parsed._replace(netloc=f"litellm_writer:litellm-writer@{parsed.hostname}:{parsed.port}")
)
with psycopg.connect(reader_url, autocommit=True) as reader:
assert reader.execute("SHOW transaction_read_only").fetchone() == ("on",)
try:
reader.execute("CREATE TABLE integration_readonly_probe (id int)")
except psycopg.errors.ReadOnlySqlTransaction:
pass
else:
raise AssertionError("litellm_reader executed a write statement")
with psycopg.connect(writer_url, autocommit=True) as writer:
assert writer.execute("SELECT current_user").fetchone() == ("litellm_writer",)
if __name__ == "__main__":
main()

View file

@ -7,7 +7,15 @@ if [ "${GITHUB_ACTIONS:-}" = true ]; then
fi
suite="${1:?integration suite required}"
results="test-results/integration-${suite}"
mode="${2:-standard}"
side="${3:-}"
if [ "$mode" = replica ]; then
results="test-results/integration-${suite}-replica"
elif [ "$mode" = parity ]; then
results="test-results/parity-${suite}/${side:?parity side required}"
else
results="test-results/integration-${suite}"
fi
mkdir -p "$results"
integration_identity="$(.venv/bin/python -c 'import uuid; print(uuid.uuid4().hex)')"
upstream_pid=""
@ -80,6 +88,18 @@ export INTEGRATION_ORDER_SEED="$INTEGRATION_SEED"
uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma > "$results/prisma-generate.log" 2>&1
export INTEGRATION_PROXY_DATABASE_URL=""
export INTEGRATION_PROXY_READ_REPLICA_URL=""
export INTEGRATION_ROUTING=""
if [ "$mode" = replica ] || [ "$mode" = parity ]; then
.venv/bin/python .circleci/scripts/prepare_replica_roles.py > "$results/prepare-replica-roles.log" 2>&1
export INTEGRATION_PROXY_DATABASE_URL="postgresql://litellm_writer:litellm-writer@127.0.0.1:5432/circle_test"
export INTEGRATION_PROXY_READ_REPLICA_URL="postgresql://litellm_reader:litellm-reader@127.0.0.1:5432/circle_test"
fi
if [ "$mode" = parity ]; then
export INTEGRATION_ROUTING=capture
fi
sudo iptables -N integration_only
guard_created=true
sudo iptables -A integration_only -o lo -j ACCEPT
@ -137,8 +157,12 @@ start_proxy() {
else
cost_map_env=("LITELLM_LOCAL_MODEL_COST_MAP=True")
fi
local -a database_env=("DATABASE_URL=${INTEGRATION_PROXY_DATABASE_URL:-$DATABASE_URL}")
if [ -n "$INTEGRATION_PROXY_READ_REPLICA_URL" ]; then
database_env+=("DATABASE_URL_READ_REPLICA=$INTEGRATION_PROXY_READ_REPLICA_URL")
fi
setsid env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" INTEGRATION_RUN_ID="$integration_identity" \
DATABASE_URL="$DATABASE_URL" REDIS_HOST="$REDIS_HOST" REDIS_PORT="$REDIS_PORT" \
"${database_env[@]}" REDIS_HOST="$REDIS_HOST" REDIS_PORT="$REDIS_PORT" \
INTEGRATION_UPSTREAM_URL="$INTEGRATION_UPSTREAM_URL" \
LITELLM_MASTER_KEY="$LITELLM_MASTER_KEY" LITELLM_SALT_KEY="$LITELLM_SALT_KEY" LITELLM_UI_PATH="$LITELLM_UI_PATH" PROXY_BASE_URL="http://127.0.0.1:$port" \
LITELLM_MODE=PRODUCTION STORE_MODEL_IN_DB=True "${cost_map_env[@]}" \
@ -195,6 +219,9 @@ env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" \
INTEGRATION_SEED="$INTEGRATION_SEED" \
INTEGRATION_ORDER_SEED="$INTEGRATION_ORDER_SEED" \
LITELLM_LOCAL_MODEL_COST_MAP=True AWS_EC2_METADATA_DISABLED=true DO_NOT_TRACK=1 \
INTEGRATION_PROXY_DATABASE_URL="$INTEGRATION_PROXY_DATABASE_URL" \
INTEGRATION_PROXY_READ_REPLICA_URL="$INTEGRATION_PROXY_READ_REPLICA_URL" \
INTEGRATION_ROUTING="$INTEGRATION_ROUTING" \
.venv/bin/python tests/integration/run.py "$suite" --results "$results"
if [ "${INTEGRATION_COVERAGE:-0}" = 1 ]; then

View file

@ -245,6 +245,8 @@ pub struct ModelInfo {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cache_creation_input_token_cost_above_272k_tokens_priority: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cache_creation_input_token_cost_above_32k_tokens: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cache_creation_input_token_cost_batches: Option<f64>,
/// Flex service-tier rate for the same-named base field.
#[serde(default, skip_serializing_if = "Option::is_none")]
@ -283,6 +285,8 @@ pub struct ModelInfo {
/// Priority service-tier rate for the same-named base field.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cache_read_input_token_cost_above_272k_tokens_priority: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cache_read_input_token_cost_above_32k_tokens: Option<f64>,
/// Rate applied once the prompt exceeds the token threshold in the field name.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub cache_read_input_token_cost_above_512k_tokens: Option<f64>,
@ -377,6 +381,8 @@ pub struct ModelInfo {
/// Priority service-tier rate for the same-named base field.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub input_cost_per_token_above_272k_tokens_priority: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub input_cost_per_token_above_32k_tokens: Option<f64>,
/// Rate applied once the prompt exceeds the token threshold in the field name.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub input_cost_per_token_above_512k_tokens: Option<f64>,
@ -498,6 +504,8 @@ pub struct ModelInfo {
/// Priority service-tier rate for the same-named base field.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_cost_per_token_above_272k_tokens_priority: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_cost_per_token_above_32k_tokens: Option<f64>,
/// Rate applied once the prompt exceeds the token threshold in the field name.
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_cost_per_token_above_512k_tokens: Option<f64>,

View file

@ -429,6 +429,7 @@ class LLMCachingHandler:
kwargs=kwargs,
cached_result=cached_result,
is_async=False,
custom_llm_provider=custom_llm_provider,
)
if not _should_defer_streaming_cache_hit_callbacks(cached_result=cached_result):

View file

@ -48,6 +48,7 @@ DEFAULT_BATCH_SIZE: Final = int(os.getenv("DEFAULT_BATCH_SIZE", 512))
DEFAULT_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_FLUSH_INTERVAL_SECONDS", 5))
DEFAULT_S3_FLUSH_INTERVAL_SECONDS: Final = int(os.getenv("DEFAULT_S3_FLUSH_INTERVAL_SECONDS", 10))
DEFAULT_S3_BATCH_SIZE: Final = int(os.getenv("DEFAULT_S3_BATCH_SIZE", 512))
DEFAULT_S3_MAX_CONCURRENT_UPLOADS: Final = int(os.getenv("DEFAULT_S3_MAX_CONCURRENT_UPLOADS", "16"))
# https://docs.aws.amazon.com/AmazonS3/latest/userguide/object-keys.html
MAX_S3_OBJECT_KEY_BYTES: Final = 1024
S3_BOUNDED_OBJECT_KEY_HEAD_BYTES: Final = 64
@ -1518,6 +1519,7 @@ PROMETHEUS_BUDGET_METRICS_REFRESH_INTERVAL_MINUTES: Final = int(
CLOUDZERO_EXPORT_INTERVAL_MINUTES: Final = int(os.getenv("CLOUDZERO_EXPORT_INTERVAL_MINUTES", 60))
MCP_TOOL_NAME_PREFIX: Final = "mcp_tool"
MAXIMUM_TRACEBACK_LINES_TO_LOG: Final = int(os.getenv("MAXIMUM_TRACEBACK_LINES_TO_LOG", 100))
PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS: Final = 4096
# Headers to control callbacks
X_LITELLM_DISABLE_CALLBACKS: Final = "x-litellm-disable-callbacks"

View file

@ -2017,8 +2017,8 @@ def _deployment_model_info(
return cast(ModelInfo, registered_deployment_info) # cast-ok: router registers deployment prices under its id
if litellm_logging_obj is None:
return None
litellm_params: Final = getattr(litellm_logging_obj, "litellm_params", None)
if litellm_params is None:
litellm_params: Final = litellm_logging_obj.litellm_params
if not litellm_params:
return None
return next(
(
@ -2036,7 +2036,9 @@ def _ocr_model_info(
router_model_id: str | None,
) -> OCRPricing | None:
deployment_info: Final = _deployment_model_info(litellm_logging_obj, custom_pricing, router_model_id)
litellm_params: Final = getattr(litellm_logging_obj, "litellm_params", None) if custom_pricing else None
litellm_params: Final = (
litellm_logging_obj.litellm_params if custom_pricing and litellm_logging_obj is not None else None
)
if litellm_params is None:
return deployment_info
return _layered_ocr_pricing(litellm_params, deployment_info)

View file

@ -129,7 +129,7 @@ async def list_tools_with_pagination(
)
tools.extend(result.tools)
next_cursor = getattr(result, "next_cursor", None)
next_cursor = result.next_cursor
if not isinstance(next_cursor, str) or not next_cursor:
return tools
if next_cursor in seen_cursors:

View file

@ -112,7 +112,7 @@ class ArizeLogger(OpenTelemetry):
if value is None or value in ("", "None"):
return None
try:
rate = float(value)
rate: Final = float(value)
except (TypeError, ValueError):
verbose_logger.warning(
"ArizeLogger: %s value %r is not a number; exporting the request",

View file

@ -21,7 +21,7 @@ from __future__ import annotations
import os
from datetime import datetime, timedelta, timezone
from typing import TYPE_CHECKING, Any, Final, Protocol
from typing import TYPE_CHECKING, Any, Final
import litellm
from litellm._logging import verbose_proxy_logger
@ -35,17 +35,6 @@ else:
AsyncIOScheduler = Any
class _PodLockManager(Protocol):
"""The subset of PodLockManager this logger drives to serialize the export across pods."""
@property
def redis_cache(self) -> object: ...
async def acquire_lock(self, cronjob_id: str) -> bool | None: ...
async def release_lock(self, cronjob_id: str) -> None: ...
def _parse_metrics_marker(
marker: object | None,
) -> datetime | None:
@ -237,13 +226,10 @@ class MavvrikFocusLogger(FocusLogger):
"""Scheduler entry point — uses Mavvrik-specific pod-lock key."""
from litellm.proxy.proxy_server import proxy_logging_obj # noqa: PLC0415
pod_lock_manager: _PodLockManager | None = None
if proxy_logging_obj is not None:
writer: Final[object] = getattr(proxy_logging_obj, "db_spend_update_writer", None)
if writer is not None:
pod_lock_manager = getattr(writer, "pod_lock_manager", None)
if pod_lock_manager and pod_lock_manager.redis_cache:
pod_lock_manager: Final = (
proxy_logging_obj.db_spend_update_writer.pod_lock_manager if proxy_logging_obj is not None else None
)
if pod_lock_manager is not None and pod_lock_manager.redis_cache:
acquired: Final = await pod_lock_manager.acquire_lock(cronjob_id=MAVVRIK_FOCUS_EXPORT_JOB_NAME)
if not acquired:
verbose_proxy_logger.debug("Mavvrik FOCUS export: unable to acquire pod lock")

View file

@ -20,7 +20,8 @@ from litellm.constants import (
)
from litellm.types.utils import StandardLoggingPayload
_S3_LOG_PROMPTS_ONLY: Final = TypeAdapter(bool)
_S3_BOOL: Final = TypeAdapter(bool)
_UPLOAD_BOUND: Final = TypeAdapter(int)
def resolve_s3_log_prompts_only(configured: object, environ: Mapping[str, str] | None = None) -> bool:
@ -29,12 +30,42 @@ def resolve_s3_log_prompts_only(configured: object, environ: Mapping[str, str] |
if raw is None or raw == "":
return False
try:
return _S3_LOG_PROMPTS_ONLY.validate_python(raw.strip() if isinstance(raw, str) else raw)
return _S3_BOOL.validate_python(raw.strip() if isinstance(raw, str) else raw)
except ValidationError:
verbose_logger.warning("s3 logging: s3_log_prompts_only=%r is not a boolean, logging prompts only", raw)
return True
def resolve_s3_max_concurrent_uploads(configured: object, fallback: int) -> int:
if configured is None or configured == "":
return fallback
try:
bound: Final = _UPLOAD_BOUND.validate_python(configured.strip() if isinstance(configured, str) else configured)
except ValidationError:
verbose_logger.warning(
"s3 logging: s3_max_concurrent_uploads=%r is not an integer, using %s", configured, fallback
)
return fallback
if bound < 1:
verbose_logger.warning(
"s3 logging: s3_max_concurrent_uploads=%r must be at least 1, using %s", configured, fallback
)
return fallback
return bound
def resolve_s3_batch_file_upload(configured: object) -> bool:
if configured is None or configured == "":
return False
try:
return _S3_BOOL.validate_python(configured.strip() if isinstance(configured, str) else configured)
except ValidationError:
verbose_logger.warning(
"s3 logging: s3_batch_file_upload=%r is not a boolean, keeping per-request objects", configured
)
return False
def prompts_only_payload(payload: StandardLoggingPayload) -> StandardLoggingPayload:
return {**payload, "response": None}

View file

@ -3,26 +3,33 @@ s3 Bucket Logging Integration
async_log_success_event: Processes the event, stores it in memory for DEFAULT_S3_FLUSH_INTERVAL_SECONDS seconds or until DEFAULT_S3_BATCH_SIZE and then flushes to s3
async_log_failure_event: Processes the event, stores it in memory for DEFAULT_S3_FLUSH_INTERVAL_SECONDS seconds or until DEFAULT_S3_BATCH_SIZE and then flushes to s3
NOTE 1: S3 does not provide a BATCH PUT API endpoint, so we create tasks to upload each element individually
NOTE 1: S3 does not provide a BATCH PUT API endpoint; by default each element is uploaded concurrently (bounded by s3_max_concurrent_uploads), or with s3_batch_file_upload the whole flush is written as one .jsonl file
"""
import asyncio
import time
from collections.abc import Mapping
from datetime import datetime
from datetime import datetime, timezone
from typing import TYPE_CHECKING, Final, cast
from urllib.parse import quote
from uuid import uuid4
import httpx
import litellm
from litellm._logging import print_verbose, verbose_logger
from litellm.constants import DEFAULT_S3_BATCH_SIZE, DEFAULT_S3_FLUSH_INTERVAL_SECONDS
from litellm.constants import (
DEFAULT_S3_BATCH_SIZE,
DEFAULT_S3_FLUSH_INTERVAL_SECONDS,
DEFAULT_S3_MAX_CONCURRENT_UPLOADS,
)
from litellm.integrations.s3 import (
get_s3_object_download_filename,
get_s3_object_key,
prompts_only_payload,
resolve_s3_batch_file_upload,
resolve_s3_log_prompts_only,
resolve_s3_max_concurrent_uploads,
resolve_sse_params,
)
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
@ -43,7 +50,20 @@ if TYPE_CHECKING:
from botocore.credentials import Credentials
def _s3_key_parent(s3_object_key: str) -> str:
return s3_object_key.rsplit("/", 1)[0] if "/" in s3_object_key else ""
class S3BatchUploadError(Exception):
def __init__(self, failed: int, total: int) -> None:
self.failed = failed
self.total = total
super().__init__(f"{failed} of {total} S3 uploads failed; events kept in queue for the next flush")
class S3Logger(CustomBatchLogger, BaseAWSLLM):
preserve_events_added_during_flush = True
def __init__(
self,
s3_bucket_name: str | None = None,
@ -71,6 +91,8 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_server_side_encryption: str | None = None,
s3_sse_kms_key_id: str | None = None,
s3_log_prompts_only: bool | None = None,
s3_max_concurrent_uploads: int = DEFAULT_S3_MAX_CONCURRENT_UPLOADS,
s3_batch_file_upload: bool = False,
s3_callback_params_override: dict | None = None,
**kwargs,
):
@ -112,7 +134,10 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_server_side_encryption=s3_server_side_encryption,
s3_sse_kms_key_id=s3_sse_kms_key_id,
s3_log_prompts_only=s3_log_prompts_only,
s3_max_concurrent_uploads=s3_max_concurrent_uploads,
s3_batch_file_upload=s3_batch_file_upload,
)
self._upload_semaphore = asyncio.Semaphore(self.s3_max_concurrent_uploads)
verbose_logger.debug("s3 logger using endpoint url %s", s3_endpoint_url)
# IMPORTANT
@ -168,6 +193,8 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
s3_server_side_encryption: str | None = None,
s3_sse_kms_key_id: str | None = None,
s3_log_prompts_only: bool | None = None,
s3_max_concurrent_uploads: int = DEFAULT_S3_MAX_CONCURRENT_UPLOADS,
s3_batch_file_upload: bool = False,
params_source: dict | None = None,
):
"""
@ -226,6 +253,16 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
params.get("s3_sse_kms_key_id") or s3_sse_kms_key_id,
)
configured_bound: Final = params.get("s3_max_concurrent_uploads")
self.s3_max_concurrent_uploads = resolve_s3_max_concurrent_uploads(
s3_max_concurrent_uploads if configured_bound is None or configured_bound == "" else configured_bound,
DEFAULT_S3_MAX_CONCURRENT_UPLOADS,
)
self.s3_batch_file_upload = s3_batch_file_upload or resolve_s3_batch_file_upload(
params.get("s3_batch_file_upload")
)
def _build_object_url(self, s3_object_key: str) -> str:
"""
Build the exact URL that is both signed and sent, with the key percent-encoded once.
@ -347,7 +384,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
verbose_logger.exception("s3 Layer Error - %s", e)
self.handle_callback_failure(callback_name="S3Logger")
async def async_upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement):
async def async_upload_data_to_s3(self, batch_logging_element: s3BatchLoggingElement) -> bool:
try:
import base64
import hashlib
@ -364,7 +401,11 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
url: Final = self._build_object_url(batch_logging_element.s3_object_key)
# Convert JSON to string
json_string: Final = safe_dumps(batch_logging_element.payload)
json_string: Final = (
batch_logging_element.body
if batch_logging_element.body is not None
else safe_dumps(batch_logging_element.payload)
)
# Calculate SHA256 hash of the content
content_hash: Final = hashlib.sha256(json_string.encode("utf-8")).hexdigest()
@ -374,7 +415,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
# Prepare the request
headers: Final = {
"Content-Type": "application/json",
"Content-Type": batch_logging_element.content_type,
"Content-MD5": content_md5,
"x-amz-content-sha256": content_hash,
"Content-Language": "en",
@ -421,27 +462,72 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
except Exception as e:
verbose_logger.exception("Error uploading to s3: %s", e)
self.handle_callback_failure(callback_name="S3Logger")
return False
return True
async def async_send_batch(self):
async def async_send_batch(self) -> None:
"""
Sends runs from self.log_queue.
Sends runs from self.log_queue
Returns: None
Raises: Does not raise an exception, will only verbose_logger.exception()
Raises S3BatchUploadError when any upload failed; CustomBatchLogger.flush_queue
keeps the surviving queue entries for the next flush.
"""
verbose_logger.debug("s3_v2 logger - sending batch of %s", len(self.log_queue))
if not self.log_queue:
batch: Final = tuple(self.log_queue)
if not batch:
return
verbose_logger.debug("s3_v2 logger - sending batch of %s", len(batch))
#########################################################
# Flush the log queue to s3
# the log queue can be bounded by DEFAULT_S3_BATCH_SIZE
# see custom_batch_logger.py which triggers the flush
#########################################################
for payload in self.log_queue:
asyncio.create_task(self.async_upload_data_to_s3(payload))
uploads: Final = self._batch_file_elements(batch) if self._batch_file_mode_active() else batch
results: Final = await asyncio.gather(*(self._upload_bounded(element) for element in uploads))
failed: Final = tuple(element for element, ok in zip(uploads, results, strict=True) if not ok)
if not failed:
return
self.log_queue = [*failed, *self.log_queue[len(batch) :]]
raise S3BatchUploadError(failed=len(failed), total=len(uploads))
def _batch_file_mode_active(self) -> bool:
if not self.s3_batch_file_upload:
return False
if litellm.cold_storage_custom_logger == "s3_v2":
verbose_logger.warning(
"s3 logging: s3_batch_file_upload is ignored because s3_v2 is the cold storage logger; "
"per-request objects are required for spend log lookups"
)
return False
return True
async def _upload_bounded(self, element: s3BatchLoggingElement) -> bool:
async with self._upload_semaphore:
return await self.async_upload_data_to_s3(element)
def _batch_file_elements(self, batch: tuple[s3BatchLoggingElement, ...]) -> tuple[s3BatchLoggingElement, ...]:
now: Final = datetime.now(timezone.utc)
groups: Final = {
parent: tuple(
element for element in batch if element.body is None and _s3_key_parent(element.s3_object_key) == parent
)
for parent in sorted({_s3_key_parent(element.s3_object_key) for element in batch if element.body is None})
}
return tuple(element for element in batch if element.body is not None) + tuple(
self._build_batch_file_element(elements, parent, now) for parent, elements in groups.items()
)
def _build_batch_file_element(
self, elements: tuple[s3BatchLoggingElement, ...], parent: str, now: datetime
) -> s3BatchLoggingElement:
batch_name: Final = f"batch_{now.strftime('%H-%M-%S')}_{uuid4().hex}"
return s3BatchLoggingElement(
payload={},
body="\n".join(safe_dumps(element.payload) for element in elements),
content_type="application/x-ndjson",
s3_object_key=f"{parent}/{batch_name}.jsonl" if parent else f"{batch_name}.jsonl",
s3_object_download_filename=f"{batch_name}.jsonl",
)
def create_s3_batch_logging_element(
self,
@ -521,7 +607,11 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
url: Final = self._build_object_url(batch_logging_element.s3_object_key)
# Convert JSON to string
json_string: Final = safe_dumps(batch_logging_element.payload)
json_string: Final = (
batch_logging_element.body
if batch_logging_element.body is not None
else safe_dumps(batch_logging_element.payload)
)
# Calculate SHA256 hash of the content
content_hash: Final = hashlib.sha256(json_string.encode("utf-8")).hexdigest()
@ -531,7 +621,7 @@ class S3Logger(CustomBatchLogger, BaseAWSLLM):
# Prepare the request
headers: Final = {
"Content-Type": "application/json",
"Content-Type": batch_logging_element.content_type,
"Content-MD5": content_md5,
"x-amz-content-sha256": content_hash,
"Content-Language": "en",

View file

@ -1849,7 +1849,7 @@ class WebSearchInterceptionLogger(CustomLogger):
for tool_call in tool_calls:
# Handle both Anthropic-style input and OpenAI-style function.arguments
query = None
tool_args: dict | None = None # mutable-ok: the tool call's own arguments dict
tool_args: dict[str, object] | None = None # mutable-ok: the tool call's own arguments dict
if "input" in tool_call and isinstance(tool_call["input"], dict):
tool_args = tool_call["input"]
query = tool_args.get("query")

View file

@ -365,7 +365,7 @@ def _budget_reservation_on_auth_object(user_api_key_auth: object) -> object:
return getattr(user_api_key_auth, "budget_reservation", None)
def budget_reservation_from_metadata(metadata: Mapping[str, object]) -> dict | None:
def budget_reservation_from_metadata(metadata: Mapping[str, object]) -> dict[str, object] | None:
stamped: Final = metadata.get("user_api_key_budget_reservation")
if isinstance(stamped, dict):
return stamped

View file

@ -60,13 +60,13 @@ class _HasProxyErrorType(Protocol):
_MESSAGE_PATTERNS: Final[tuple[tuple[re.Pattern[str], str], ...]] = (
(re.compile(r"upstream passthrough request failed", re.IGNORECASE), UPSTREAM_PASSTHROUGH),
(
re.compile(r"budget has been exceeded|max budget|exceeded.*budget|crossed budget", re.IGNORECASE),
re.compile(r"budget has been exceeded|max budget|crossed budget", re.IGNORECASE),
BUDGET_EXCEEDED,
),
(re.compile(r"no healthy deployments?|no deployments available", re.IGNORECASE), NO_HEALTHY_DEPLOYMENTS),
(re.compile(r"not allowed to access model due to tags configuration", re.IGNORECASE), MODEL_ACCESS_DENIED),
(re.compile(r"upstream passthrough request failed", re.IGNORECASE), UPSTREAM_PASSTHROUGH),
(re.compile(r"is not supported for provider|not implemented", re.IGNORECASE), UNSUPPORTED_OPERATION),
(
re.compile(r"context window|context length|(prompt|input) is too long|tokens? ?> ?\d+ ?maximum", re.IGNORECASE),
@ -155,6 +155,14 @@ _CLASS_CODE_TABLE: Final[tuple[tuple[tuple[type[BaseException], ...], str], ...]
)
def _exceeded_before_budget(message: str) -> bool:
"""Linear-time equivalent of ``re.search(r"exceeded.*budget", message, re.IGNORECASE)``."""
return any(
(start := line.find("exceeded")) != -1 and line.find("budget", start + len("exceeded")) != -1
for line in message.lower().split("\n")
)
def _classify_by_message(message: str, patterns: tuple[tuple[re.Pattern[str], str], ...]) -> str | None:
return next((code for pattern, code in patterns if pattern.search(message)), None)
@ -183,7 +191,9 @@ def normalize_error(exc: Exception | None, status_code: str, message: str) -> st
by_proxy_type: Final = _PROXY_ERROR_TYPE_MAP.get(proxy_type) if isinstance(proxy_type, str) else None
if by_proxy_type is not None:
return by_proxy_type
by_message: Final = _classify_by_message(message, _MESSAGE_PATTERNS)
by_message: Final = (
BUDGET_EXCEEDED if _exceeded_before_budget(message) else _classify_by_message(message, _MESSAGE_PATTERNS)
)
if by_message is not None:
return by_message
by_class: Final = _classify_by_class(exc)

View file

@ -5191,7 +5191,7 @@ def _maybe_construct_otel_v2(callback_name: str, _in_memory_loggers: list[Custom
for callback in _in_memory_loggers:
if (
isinstance(callback, OpenTelemetryV2)
and getattr(callback, "callback_name", None) == callback_name
and callback.callback_name == callback_name
and (serves_a_destination or not _exports_nowhere(callback.config))
):
return callback
@ -6663,7 +6663,7 @@ def get_standard_logging_object_payload(
cost_breakdown=request_cost_breakdown,
autorouter_savings=autorouter_savings,
autorouter_savings_estimate=(
{
{ # mutable-ok: spend-log JSON serialization requires plain mappings
"version": 3,
"status": "unknown",
"reason": "pending_projection",

View file

@ -5,7 +5,12 @@ from typing import Final
from pydantic import TypeAdapter, ValidationError
from typing_extensions import assert_never
from litellm.types.utils import StandardLoggingZeroCostDiagnostic, Usage
from litellm.types.utils import (
CompletionTokensDetailsWrapper,
PromptTokensDetailsWrapper,
StandardLoggingZeroCostDiagnostic,
Usage,
)
ZERO_COST_COUNTER_NAME: Final = "litellm_zero_cost_requests_total"
@ -18,8 +23,8 @@ _NESTED_PRICING: Final = TypeAdapter(Mapping[str, object] | tuple[object, ...])
_MAX_PRICING_DEPTH: Final = 4
def _audio_tokens(details: object) -> int:
audio_tokens: Final = getattr(details, "audio_tokens", None)
def _audio_tokens(details: PromptTokensDetailsWrapper | CompletionTokensDetailsWrapper | None) -> int:
audio_tokens: Final = details.audio_tokens if details is not None else None
return audio_tokens if isinstance(audio_tokens, int) and audio_tokens > 0 else 0

View file

@ -2003,11 +2003,11 @@ def strip_encrypted_reasoning_from_messages(messages: object) -> None:
"""
if not isinstance(messages, list):
return
for content in _anthropic_content_lists(cast(list[object], messages)): # cast-ok: untyped client json
for content in anthropic_content_lists(cast(list[object], messages)): # cast-ok: untyped client json
_strip_encrypted_reasoning_from_blocks(content)
def _anthropic_content_lists(messages: Sequence[object]) -> Iterator[object]:
def anthropic_content_lists(messages: Sequence[object]) -> Iterator[object]:
return (
cast(list[object], content) # cast-ok: narrowed by isinstance
for message in messages

View file

@ -1329,7 +1329,7 @@ class CustomStreamWrapper:
"is_finished": chunk_finish_reason is not None,
"finish_reason": chunk_finish_reason,
"original_chunk": cached_chunk,
"tool_calls": (getattr(cached_choice.delta, "tool_calls", None) if cached_choice is not None else None),
"tool_calls": cached_choice.delta.tool_calls if cached_choice is not None else None,
}
completion_obj["content"] = response_obj["text"]

View file

@ -48,7 +48,7 @@ def _registry_api_key(agent_litellm_params: Mapping[str, object]) -> str | None:
return configured_api_key if isinstance(configured_api_key, str) else None
def _registry_headers(agent_litellm_params: Mapping[str, object]) -> dict[str, Any] | None:
def _registry_headers(agent_litellm_params: Mapping[str, object]) -> dict[str, object] | None:
stored_headers: Final = agent_litellm_params.get("headers")
if not isinstance(stored_headers, Mapping):
return None

View file

@ -685,7 +685,7 @@ class AnthropicMessagesHandler(BaseTranslation):
return data
def _hoisted_top_level_system_message(self, data: dict) -> AllMessageValues | None:
def _hoisted_top_level_system_message(self, data: Mapping[str, object]) -> AllMessageValues | None:
"""Return the system message produced by translating the top-level prompt."""
system: Final = data.get("system")
if not system:

View file

@ -11,7 +11,6 @@ from typing import (
Final,
Literal,
Protocol,
cast, # noqa: TID251 # rebuilt message_delta dict spans the ContentBlockDelta/MessageBlockDelta union
get_args,
)
@ -27,6 +26,7 @@ from litellm.types.llms.anthropic import (
ContentBlockDelta,
ContextManagementResponse,
MessageBlockDelta,
MessageDelta,
StreamingContentBlockDeltaType,
UsageDelta,
UsageIteration,
@ -1028,26 +1028,22 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
self,
processed_chunk: ContentBlockDelta | MessageBlockDelta,
) -> ContentBlockDelta | MessageBlockDelta:
if processed_chunk.get("type") != "message_delta" or not self._refusal_text:
if processed_chunk["type"] != "message_delta" or not self._refusal_text:
return processed_chunk
delta: Final = cast(Mapping[str, object], processed_chunk["delta"]) # cast-ok: keys checked before use
delta: Final = processed_chunk["delta"]
if delta.get("stop_reason") == "max_tokens":
return processed_chunk
from litellm.llms.anthropic.experimental_pass_through.messages.utils import (
refusal_stop_details,
)
return cast( # cast-ok: rebuilt dict matches the message_delta TypedDict shape for this branch
ContentBlockDelta | MessageBlockDelta,
{ # mutable-ok: fresh translation payload; never mutated after construction
**processed_chunk,
"delta": { # mutable-ok: fresh message_delta payload; never mutated after construction
**delta,
"stop_reason": "refusal",
"stop_details": refusal_stop_details(self._refusal_text),
},
},
)
refusal_delta: Final[MessageDelta] = {
**delta,
"stop_reason": "refusal",
"stop_details": refusal_stop_details(self._refusal_text),
}
refusal_chunk: Final[MessageBlockDelta] = {**processed_chunk, "delta": refusal_delta}
return refusal_chunk
@staticmethod
def _delta_has_content(processed_chunk: Mapping[str, object]) -> bool:

View file

@ -37,7 +37,7 @@ def _mapping_field(container: object, key: str) -> object | None:
"""One key of a raw provider payload, or None when the payload is not a mapping."""
if not isinstance(container, Mapping):
return None
return cast(Mapping[str, object], container).get(key) # cast-ok: raw payload, callers re-check every value
return container.get(key)
def _mapping_str_field(container: object, key: str) -> str | None:

View file

@ -169,7 +169,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
cls,
summary: Iterable[object],
encrypted_content: object,
) -> dict[str, Any] | None: # mutable-ok: API message payload
) -> dict[str, object] | None: # mutable-ok: API message payload
"""The one Anthropic block for a Responses reasoning item.
The item's encrypted reasoning rides the block's opaque field (`signature`, or
@ -198,7 +198,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter:
@classmethod
def _assistant_group_to_input_items(
cls, group: tuple[Mapping[str, object], ...]
) -> tuple[dict[str, Any], ...]: # mutable-ok: API message payload
) -> tuple[dict[str, object], ...]: # mutable-ok: API message payload
first: Final = group[0]
btype: Final = first.get("type")
if btype in ("thinking", "redacted_thinking"):

View file

@ -59,6 +59,7 @@ def resolve_fireworks_api_key(api_key: str | None) -> str | None:
AZURE_FOUNDRY_FIREWORKS_MODEL_ID_PREFIX: Final = "FW-"
FIREROUTER: Final = "firerouter"
def resolve_fireworks_resource_name(model: str) -> str:
@ -67,7 +68,7 @@ def resolve_fireworks_resource_name(model: str) -> str:
return stripped
if stripped.startswith(("routers/", "models/")):
return f"accounts/fireworks/{stripped}"
if stripped.endswith("-fast"):
if stripped.endswith("-fast") or stripped == FIREROUTER or stripped.startswith(f"{FIREROUTER}/"):
return f"accounts/fireworks/routers/{stripped}"
return f"accounts/fireworks/models/{stripped}"

View file

@ -59,6 +59,13 @@ def get_base_model_for_pricing(model_name: str) -> str:
def _resolve_model_info(model: str) -> ModelInfo:
try:
return get_model_info(model=model, custom_llm_provider="fireworks_ai")
except Exception:
return _resolve_routed_model_info(model)
def _resolve_routed_model_info(model: str) -> ModelInfo:
try:
return get_model_info(model=model.removeprefix("fireworks_ai/"))
except Exception:
base_model: Final = get_base_model_for_pricing(model_name=model)
return get_model_info(model=base_model, custom_llm_provider="fireworks_ai")
@ -81,7 +88,7 @@ def cost_per_token(model: str, usage: Usage, current_time: datetime | None = Non
return generic_cost_per_token(
model=model,
usage=usage,
custom_llm_provider="fireworks_ai",
custom_llm_provider=model_info["litellm_provider"],
model_info=model_info,
current_time=current_time,
)

View file

@ -994,7 +994,7 @@ class OpenAIResponsesHandler(BaseTranslation):
def _spread_text_rewrite_over_stream_events(
self,
stream_events: Sequence[Any],
stream_events: Sequence[object],
rewritten_text: str,
guardrail_name: str,
) -> None:

View file

@ -7,6 +7,7 @@ Vercel AI Gateway is OpenAI-compatible and supports embeddings via the /v1/embed
Docs: https://vercel.com/docs/ai-gateway/openai-compat/embeddings
"""
from collections.abc import Mapping
from typing import TYPE_CHECKING, Any, Final
import httpx
@ -161,12 +162,14 @@ class VercelAIGatewayEmbeddingConfig(BaseEmbeddingConfig):
optional_params[param] = value
return optional_params
def get_error_class(self, error_message: str, status_code: int, headers: Any) -> BaseLLMException:
def get_error_class(
self, error_message: str, status_code: int, headers: Mapping[str, str] | httpx.Headers
) -> BaseLLMException:
"""
Get the error class for Vercel AI Gateway errors.
"""
return VercelAIGatewayException(
message=error_message,
status_code=status_code,
headers=headers,
headers=headers if isinstance(headers, httpx.Headers) else httpx.Headers(headers),
)

View file

@ -286,7 +286,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
"""
Check if the model is Gemini 3 or newer.
"""
model_name = model.split("/")[-1].lower()
model_name: Final = model.split("/")[-1].lower()
is_vertex_fine_tuned_model: Final = model_name.isdigit() or (
model.startswith("gemini/") and not model_name.startswith("gemini-")
)

View file

@ -1182,7 +1182,7 @@ def _is_claude_tool_target(custom_llm_provider: str | None, model: str) -> bool:
return False
def _without_anthropic_only_tool_keys(tool: dict) -> dict:
def _without_anthropic_only_tool_keys(tool: dict[str, object]) -> dict[str, object]:
kept: Final = {key: value for key, value in tool.items() if key not in _ANTHROPIC_ONLY_TOOL_KEYS}
function: Final = tool.get("function")
if not isinstance(function, dict):
@ -1193,7 +1193,7 @@ def _without_anthropic_only_tool_keys(tool: dict) -> dict:
}
def _drop_anthropic_only_tool_keys(tools: list[dict] | None) -> list[dict] | None:
def _drop_anthropic_only_tool_keys(tools: list[dict[str, object]] | None) -> list[dict[str, object]] | None:
if tools is None:
return None
return [_without_anthropic_only_tool_keys(tool) if isinstance(tool, dict) else tool for tool in tools]

File diff suppressed because it is too large Load diff

View file

@ -540,7 +540,9 @@ def llm_passthrough_route(
)
## IS STREAMING REQUEST
_streaming_request_data: dict = data if isinstance(data, dict) else (json if isinstance(json, dict) else {})
_streaming_request_data: Final[dict[str, object]] = (
data if isinstance(data, dict) else (json if isinstance(json, dict) else {})
)
is_streaming_request: Final = provider_config.is_streaming_request(
endpoint=endpoint,
request_data=_streaming_request_data,

View file

@ -86,8 +86,8 @@ async def oauth_authorization_uses_gateway_credential(request: Request) -> bool:
async def _opaque_bearer_is_gateway_credential(token: str) -> bool:
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import (
is_envelope, # noqa: PLC0415 # envelope imports bridge types
from litellm.proxy._experimental.mcp_server.outbound_credentials.envelope import ( # noqa: PLC0415 # envelope imports bridge types
is_envelope,
is_refresh_envelope,
)
from litellm.proxy._types import hash_token # noqa: PLC0415 # proxy import cycle

View file

@ -2634,7 +2634,9 @@ def _build_aggregate_protected_resource_response(request: Request) -> dict:
}
def _build_aggregate_authorization_server_response(request: Request, token_exchange_available: bool) -> dict:
def _build_aggregate_authorization_server_response(
request: Request, token_exchange_available: bool
) -> dict[str, object]:
"""RFC 8414 metadata for the gateway as the aggregate authorization server.
The issuer is ``{base}/mcp`` and must stay equal to the value the

View file

@ -3490,7 +3490,7 @@ class MCPServerManager:
passthrough_server_ids: Final = [
server.server_id
for server in self.get_registry().values()
if getattr(server, "auth_type", None) == MCPAuth.true_passthrough
if server.auth_type == MCPAuth.true_passthrough
]
combined_servers.update(passthrough_server_ids)

View file

@ -128,14 +128,15 @@ def with_mcp_proxy_identity(tool: Tool, server_id: str) -> Tool:
def _mcp_proxy_identity(tool: Tool) -> MCPProxyToolIdentity:
identity: Final = (tool.meta or {}).get(_MCP_PROXY_IDENTITY_META_KEY) # mutable-ok: absent metadata default
identity: Final = None if tool.meta is None else tool.meta.get(_MCP_PROXY_IDENTITY_META_KEY)
if not isinstance(identity, Mapping):
raise TypeError("MCP proxy tool identity is missing")
server_id: Final = identity.get("server_id")
tool_name: Final = identity.get("tool_name")
if not isinstance(server_id, str) or not isinstance(tool_name, str):
raise TypeError("MCP proxy tool identity is invalid")
return {"server_id": server_id, "tool_name": tool_name} # mutable-ok: TypedDict identity payload
resolved: Final[MCPProxyToolIdentity] = {"server_id": server_id, "tool_name": tool_name}
return resolved
def mcp_proxy_tool_id(tool: Tool) -> str:

View file

@ -4073,7 +4073,7 @@ async def get_org_object_for_request(
)
except OrganizationNotFoundError:
return None
except Exception as e: # noqa: BLE001 # only a DB outage may fail auth here, anything else degrades to no org limits
except Exception as e:
if not PrismaDBExceptionHandler.is_database_service_unavailable_error_in_chain(e):
verbose_proxy_logger.debug("org lookup failed, continuing without org limits", exc_info=True)
return None

View file

@ -2323,7 +2323,7 @@ class ProxyBaseLLMRequestProcessing:
return fallbacks if isinstance(fallbacks, list) and fallbacks else None
@staticmethod
def _resolve_fallback_models(model: str, fallbacks: list) -> list | None:
def _resolve_fallback_models(model: str, fallbacks: list) -> list[str] | None:
from litellm.router_utils.fallback_event_handlers import get_fallback_model_group
fallback_model_group, generic_fallback_idx = get_fallback_model_group(

View file

@ -486,7 +486,7 @@ class BaselineAccountingStore:
async def _pages(
self, db: SupportsRawQueries, scope: str, after_revision: int, withdraw_from: float | None = None
) -> AsyncIterator[tuple[_StoredRecord, ...]]:
cursor: float | None = None
cursor: float | None = None # rebind-ok: keyset pagination advances after each complete timestamp group
while page := _RECORDS.validate_python(
tuple(await db.query_raw(_READ_PAGE, scope, after_revision, cursor, _PAGE_TIMESTAMPS, withdraw_from))
):
@ -627,7 +627,7 @@ async def flush_baseline_accounting(client: PrismaClient) -> None:
more_queued: Final = bool(client.baseline_accounting_transactions)
try:
remaining: Final = await asyncio.wait_for(_flush_records(store, batch), timeout=5)
except (Exception, asyncio.CancelledError) as error: # noqa: BLE001 # unknown acknowledgements can be replayed safely
except (Exception, asyncio.CancelledError) as error:
async with client.baseline_accounting_lock:
client.baseline_accounting_transactions.extend(batch)
if isinstance(error, asyncio.CancelledError):

View file

@ -12,7 +12,8 @@ import os
import random
import time
import traceback
from collections.abc import Callable, Mapping, Sequence
from collections.abc import Callable, Coroutine, Mapping, Sequence
from contextvars import ContextVar
from datetime import datetime, timedelta, timezone
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal, Protocol, TypeAlias, TypeVar, cast, overload
@ -269,6 +270,80 @@ def _spend_update_tx(prisma_client: PrismaClient) -> _SpendTransactionManager:
return tx
_daily_spend_commit_started: Final[ContextVar[asyncio.Event | None]] = ContextVar(
"_daily_spend_commit_started", default=None
)
def _mark_daily_spend_commit_started() -> None:
started: Final = _daily_spend_commit_started.get()
if started is not None:
started.set()
def _mark_daily_spend_commit_finished() -> None:
started: Final = _daily_spend_commit_started.get()
if started is not None:
started.clear()
def _start_daily_spend_commit(
commit_started: asyncio.Event, commit: Callable[[], Coroutine[object, object, None]]
) -> "asyncio.Task[None]":
token: Final = _daily_spend_commit_started.set(commit_started)
try:
return asyncio.ensure_future(commit())
finally:
_daily_spend_commit_started.reset(token)
def _track_interrupted_commit(commits: set[asyncio.Task[None]], settle: Coroutine[object, object, None]) -> None:
task: Final = asyncio.ensure_future(settle)
commits.add(task)
task.add_done_callback(commits.discard)
async def _settle_interrupted_commits(commits: set[asyncio.Task[None]]) -> None:
while commits:
await asyncio.wait(tuple(commits))
async def _restore_tag_spend_the_commit_left_behind(
commit_task: "asyncio.Task[None]",
redis_update_buffer: RedisUpdateBuffer,
transactions: dict[str, DailyTagSpendTransaction],
) -> None:
await asyncio.wait({commit_task})
if commit_task.cancelled() or commit_task.exception() is None:
return
await redis_update_buffer.restore_transactions_to_redis(
daily_tag_spend_update_transactions=transactions,
)
async def _requeue_daily_spend_the_commit_left_behind(
commit_task: "asyncio.Task[None]",
queue: DailySpendUpdateQueue,
entity_type: str,
transactions: dict[str, BaseDailySpendTransaction],
) -> None:
await asyncio.wait({commit_task})
if commit_task.cancelled() or not transactions:
return
failure: Final = commit_task.exception()
if failure is None:
return
spend_log_error(
"Spend tracking - daily %s spend commit interrupted by shutdown failed. Re-queued %d rows for the "
"shutdown flush. Error: %s",
entity_type,
len(transactions),
str(failure),
exc=failure,
)
await queue.add_update(transactions)
# The per-team advisory lock the team endpoints hold while changing a roster (TEAM_ADVISORY_LOCK_SQL),
# so the roster check below cannot interleave with their writes. A row lock would deadlock with the
# access-group endpoints, which lock a team row after an access-group lock.
@ -391,6 +466,9 @@ class DBSpendUpdateWriter:
self.daily_org_spend_update_queue = DailySpendUpdateQueue()
self.daily_tag_spend_update_queue = DailySpendUpdateQueue()
self.window_spend_update_queue = WindowSpendUpdateQueue()
self.interrupted_tag_commits: set[asyncio.Task[None]] = (
set()
) # mutable-ok: same registry as DailySpendUpdateQueue.interrupted_commits
async def update_database(
# LiteLLM management object fields
@ -1606,17 +1684,24 @@ class DBSpendUpdateWriter:
proxy_logging_obj: ProxyLogging,
) -> None:
transactions: Final = await queue.flush_and_get_aggregated_daily_spend_update_transactions()
commit_task: Final = asyncio.ensure_future(
commit(
commit_started: Final = asyncio.Event()
commit_task: Final = _start_daily_spend_commit(
commit_started,
lambda: commit(
n_retry_times=n_retry_times,
prisma_client=prisma_client,
proxy_logging_obj=proxy_logging_obj,
daily_spend_transactions=cast(dict[str, _DailySpendTransactionT], transactions),
)
),
)
try:
await asyncio.shield(commit_task)
except asyncio.CancelledError:
if commit_started.is_set():
queue.track_interrupted_commit(
_requeue_daily_spend_the_commit_left_behind(commit_task, queue, entity_type, transactions)
)
raise
commit_task.cancel()
if transactions:
await queue.add_update(transactions)
@ -1841,23 +1926,36 @@ class DBSpendUpdateWriter:
The drain is destructive, so a failed commit must push the transactions back for the next tick
or their spend is lost permanently.
"""
await _settle_interrupted_commits(self.interrupted_tag_commits)
daily_tag_spend_update_transactions: Final = (
await self.redis_update_buffer.get_all_daily_tag_spend_update_transactions_from_redis_buffer()
)
if not daily_tag_spend_update_transactions:
return
commit_task: Final = asyncio.ensure_future(
DBSpendUpdateWriter.update_daily_tag_spend(
commit_started: Final = asyncio.Event()
commit_task: Final = _start_daily_spend_commit(
commit_started,
lambda: DBSpendUpdateWriter.update_daily_tag_spend(
n_retry_times=n_retry_times,
prisma_client=prisma_client,
proxy_logging_obj=proxy_logging_obj,
daily_spend_transactions=daily_tag_spend_update_transactions,
)
),
)
try:
await asyncio.shield(commit_task)
except BaseException: # noqa: BLE001 # a cancel must restore the drained rows before its rollback returns
if commit_started.is_set():
_track_interrupted_commit(
self.interrupted_tag_commits,
_restore_tag_spend_the_commit_left_behind(
commit_task,
self.redis_update_buffer,
daily_tag_spend_update_transactions,
),
)
raise
commit_task.cancel()
await self.redis_update_buffer.restore_transactions_to_redis(
daily_tag_spend_update_transactions=daily_tag_spend_update_transactions,
@ -2382,6 +2480,8 @@ class DBSpendUpdateWriter:
sql, params = build_bulk_upsert(table=table, batch=merged_batch)
async with _spend_update_tx(prisma_client) as transaction:
await transaction.execute_raw(sql, *params)
_mark_daily_spend_commit_started()
_mark_daily_spend_commit_finished()
except Exception as batch_error:
if _spend_commit_failure_is_requeue_safe(batch_error):
spend_log_error(

View file

@ -1,4 +1,5 @@
import asyncio
from collections.abc import Coroutine
from copy import deepcopy
from typing import Final
@ -57,6 +58,18 @@ class DailySpendUpdateQueue(BaseUpdateQueue):
self.update_queue: asyncio.Queue[dict[str, BaseDailySpendTransaction]] = asyncio.Queue(
maxsize=LITELLM_ASYNCIO_QUEUE_MAXSIZE
)
self.interrupted_commits: set[asyncio.Task[None]] = (
set()
) # mutable-ok: registry of in-flight commit outcomes, entries leave via their done callback
def track_interrupted_commit(self, settle: Coroutine[object, object, None]) -> None:
task: Final = asyncio.ensure_future(settle)
self.interrupted_commits.add(task)
task.add_done_callback(self.interrupted_commits.discard)
async def settle_interrupted_commits(self) -> None:
while self.interrupted_commits:
await asyncio.wait(tuple(self.interrupted_commits))
async def add_update(self, update: dict[str, BaseDailySpendTransaction]):
"""Enqueue an update."""
@ -81,6 +94,7 @@ class DailySpendUpdateQueue(BaseUpdateQueue):
self,
) -> dict[str, BaseDailySpendTransaction]:
"""Get all updates from the queue and return all updates aggregated by daily_transaction_key. Works for both user and team spend updates."""
await self.settle_interrupted_commits()
updates: Final = await self.flush_all_updates_from_in_memory_queue()
if len(updates) > 0:
verbose_proxy_logger.info(

View file

@ -42,6 +42,7 @@ from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.anthropic_sse import (
anthropic_sse_chunks_from_response,
assemble_anthropic_sse_stream,
is_anthropic_sse_stream,
model_response_text,
)
from litellm.types.guardrails import (
@ -93,6 +94,42 @@ def _json_escaped_len(text: str) -> int:
return len(json.dumps(text).encode("utf-8")) - 2 # strip the surrounding quotes
_MAX_FIRST_SSE_FRAME_BYTES: Final = 64 * 1024
def _holds_complete_sse_frame(raw: bytes) -> bool:
"""Whether ``raw`` holds one blank-line terminated SSE event, or is too large to keep joining."""
return b"\n\n" in raw or b"\r\n\r\n" in raw or len(raw) >= _MAX_FIRST_SSE_FRAME_BYTES
async def _coalesce_first_sse_frame(stream: AsyncIterator[object]) -> AsyncGenerator[object, None]:
"""
Join leading raw ``bytes`` chunks until they hold one complete SSE event, so
the stream shape is decided on a whole frame rather than a transport fragment.
Everything after that first frame is forwarded untouched.
"""
pending = b""
try:
async for chunk in stream:
if not isinstance(chunk, bytes):
yield chunk
continue
pending += chunk
if _holds_complete_sse_frame(pending):
break
else:
if pending:
yield pending
return
except Exception:
if pending:
yield pending
raise
yield pending
async for chunk in stream:
yield chunk
class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
user_api_key_cache = None
ad_hoc_recognizers: list[str] | None = None
@ -1356,7 +1393,7 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
all_chunks: list[ModelResponseStream] = []
passthrough_due_to_unknown_stream_shape = False
try:
stream: Final = response.__aiter__()
stream: Final = _coalesce_first_sse_frame(response.__aiter__())
async for chunk in stream:
if isinstance(chunk, ModelResponseStream):
if passthrough_due_to_unknown_stream_shape:
@ -1364,7 +1401,15 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
else:
all_chunks.append(chunk)
elif isinstance(chunk, bytes):
if passthrough_due_to_unknown_stream_shape or all_chunks:
first_frame_is_anthropic = (
not passthrough_due_to_unknown_stream_shape
and not all_chunks
and is_anthropic_sse_stream((chunk,))
)
if not first_frame_is_anthropic:
passthrough_due_to_unknown_stream_shape = (
passthrough_due_to_unknown_stream_shape or not all_chunks
)
yield chunk
continue
for masked_chunk in await self._mask_anthropic_sse_stream(chunk, stream, request_data):
@ -1387,8 +1432,9 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail):
yield chunk
if passthrough_due_to_unknown_stream_shape:
verbose_proxy_logger.warning(
"Presidio apply_to_output: streaming response contained unknown event objects "
"(e.g. /v1/responses events). Output PII masking was skipped for this response."
"Presidio apply_to_output: streaming response was not a parsed chat completion stream "
"(raw non-Anthropic SSE passthrough or /v1/responses events). "
"Output PII masking was skipped for this response."
)
return
if not all_chunks:

View file

@ -120,7 +120,7 @@ def initialize_presidio(litellm_params: LitellmParams, guardrail: Guardrail) ->
_OPTIONAL_PresidioPIIMasking,
)
explicit_filter_scope: Final = getattr(litellm_params, "presidio_filter_scope", None)
explicit_filter_scope: Final = litellm_params.presidio_filter_scope
filter_scope: Final = explicit_filter_scope or ("input" if _is_mcp_only_mode(litellm_params.mode) else "both")
run_input: Final = filter_scope in ("input", "both")
run_output: Final = filter_scope in ("output", "both")

View file

@ -770,7 +770,7 @@ async def _reconcile_budget_reservation_before_db_update(
"Failed to invalidate budget reservation counters after pre-persist reconcile failed"
)
finally:
budget_reservation["finalized"] = True # rebind-ok: the counter update reads the stamp off the shared dict
budget_reservation["finalized"] = True # rebind-ok: stamps the caller's shared dict for the counter update
async def _release_budget_reservation(budget_reservation: dict | None) -> None:

View file

@ -528,7 +528,7 @@ def _prisma_value(value: object) -> object:
return list(value) if isinstance(value, tuple) else value
def member_budget_patch(source: BaseModel) -> dict[str, Any]:
def member_budget_patch(source: BaseModel) -> Mapping[str, object]:
"""Map the per-member limit fields a request actually set to their budget-table
columns (merge-patch: a sent value updates, an explicit null clears, an absent
field is left untouched)."""
@ -561,7 +561,7 @@ async def _upsert_budget_and_membership(
user_id: str,
existing_budget_id: str | None,
user_api_key_dict: UserAPIKeyAuth,
budget_patch: dict[str, Any],
budget_patch: Mapping[str, object],
team_default_budget_id: str | None = None,
shared_budget_ids: frozenset[str] | None = None,
):
@ -624,9 +624,9 @@ async def _upsert_budget_and_membership(
if is_shared_default and not temp_only
else None
)
source: Final[Mapping[str, Any]] = source_row.model_dump() if source_row is not None else MappingProxyType({})
source: Final[Mapping[str, object]] = source_row.model_dump() if source_row is not None else MappingProxyType({})
create_data: Final[dict[str, Any]] = { # mutable-ok: Prisma create payloads are dict-shaped
create_data: Final[dict[str, object]] = { # mutable-ok: Prisma create payloads are dict-shaped
"created_by": user_api_key_dict.user_id or "",
"updated_by": user_api_key_dict.user_id or "",
**MappingProxyType(

View file

@ -348,7 +348,7 @@ async def _prepare_user(user: _PendingUser, prisma_client: PrismaClient) -> _Pre
data: Final = {**dumped, "user_id": user.user_id} # mutable-ok: /user/new defaults helper mutates in place
data_json: Final = _JSON_OBJECT.validate_python(_update_internal_new_user_params(data, user.request))
with_permission: Final = _JSON_OBJECT.validate_python(
await _set_object_permission(data_json=data_json, prisma_client=prisma_client) # pyright: ignore[reportUnknownArgumentType] # validated by the adapter
await _set_object_permission(data_json=data_json, prisma_client=prisma_client)
)
return _PreparedUser(user, _USER_ROW.validate_python(with_permission))
except Exception as exc: # noqa: BLE001 # any preparation failure is reported on this row only
@ -509,7 +509,7 @@ class _TeamsData(TypedDict):
def _default_member_budget_id(team: LiteLLM_TeamTable) -> str | None:
metadata: Final = (
_JSON_OBJECT.validate_python(
team.metadata # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # LiteLLM_TeamTable.metadata is a bare dict; validated by the adapter
team.metadata # pyright: ignore[reportUnknownMemberType] # LiteLLM_TeamTable.metadata is a bare dict; validated by the adapter
)
if team.metadata # pyright: ignore[reportUnknownMemberType] # same bare dict
else None

View file

@ -204,7 +204,7 @@ def _error_message(exc: BaseException) -> str:
if isinstance(exc, HTTPException) and isinstance(exc.detail, dict):
return str(exc.detail.get("error", exc.detail)) # pyright: ignore[reportUnknownMemberType, reportUnknownArgumentType] # HTTPException.detail is untyped
if isinstance(exc, HTTPException):
return str(exc.detail) # pyright: ignore[reportUnknownArgumentType] # HTTPException.detail is untyped
return str(exc.detail)
return str(exc) or type(exc).__name__

View file

@ -3902,7 +3902,7 @@ async def handle_gigachat_passthrough_router_model(
is_streaming: Final = request_body.get("stream", False) # pyright: ignore[reportUnknownVariableType] # request_body is dict[Unknown, Unknown]
data: dict[str, Any] = await _read_request_body(request=request) # Any needed for proxy pipeline
data: Final[dict[str, object]] = await _read_request_body(request=request)
if user_api_key_dict is not None:
auth_metadata: Final = {
metadata_key: value

View file

@ -458,7 +458,7 @@ class VertexPassthroughLoggingHandler:
@staticmethod
def _is_audio_predict_response(
model: str,
json_response: dict, # mutable-ok: predicate inspects the decoded provider response dictionary without mutation
json_response: Mapping[str, object],
) -> bool:
return (
VertexPassthroughLoggingHandler._get_audio_prediction_count(json_response=json_response) > 0
@ -467,7 +467,7 @@ class VertexPassthroughLoggingHandler:
@staticmethod
def _get_audio_prediction_count(
json_response: dict, # mutable-ok: counter inspects the decoded provider response dictionary without mutation
json_response: Mapping[str, object],
) -> int:
predictions: Final = json_response.get("predictions")
if not isinstance(predictions, list):

View file

@ -5,7 +5,7 @@ import json
import posixpath
import traceback
from base64 import b64encode
from collections.abc import AsyncGenerator, Callable, Iterable, Mapping, Sequence
from collections.abc import AsyncGenerator, AsyncIterator, Callable, Iterable, Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime
from itertools import count, groupby
@ -41,6 +41,8 @@ from litellm._logging import verbose_proxy_logger
from litellm._uuid import uuid
from litellm.constants import (
MAXIMUM_TRACEBACK_LINES_TO_LOG,
PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS,
REDACTED_BY_LITELLM,
SESSION_ID_OMITTED_METADATA_KEY,
WEBSOCKET_CLOSE_REASON_MAX_BYTES,
)
@ -56,6 +58,7 @@ from litellm.litellm_core_utils.internal_call_metadata import MODEL_ACCESS_GROUP
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.litellm_logging import _get_masked_values
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
from litellm.litellm_core_utils.redact_messages import should_redact_message_logging
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.llms.base_llm.managed_resources.utils import (
resolve_passthrough_managed_id_provider,
@ -849,23 +852,106 @@ def _resolve_team_callback_wiring(
)
def _truncate_upstream_error_body(body: str) -> str:
if len(body) <= PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS:
return body
return (
f"{body[:PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS]}... "
f"(truncated at {PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS} chars)"
)
def _sanitize_upstream_error_body(body: str) -> str:
return " ".join("".join(char if char.isprintable() else " " for char in body).split())
class _PrefixReplayStream(httpx.AsyncByteStream):
def __init__(self, prefix: bytes, rest: AsyncIterator[bytes], upstream: httpx.Response) -> None:
self._prefix: Final = prefix
self._rest: Final = rest
self._upstream: Final = upstream
async def __aiter__(self) -> AsyncIterator[bytes]:
if self._prefix:
yield self._prefix
async for chunk in self._rest:
yield chunk
async def aclose(self) -> None:
await self._upstream.aclose()
async def _no_more_chunks() -> AsyncIterator[bytes]:
return
yield b""
async def _read_error_body_preview(
stream: AsyncIterator[bytes],
) -> tuple[bytes, AsyncIterator[bytes]]:
collected: Final[list[bytes]] = [] # mutable-ok: accumulated until the preview byte budget, then joined once
total = 0 # rebind-ok: running byte count against the preview budget
try:
async for chunk in stream:
collected.append(chunk)
total += len(chunk)
if total > PASSTHROUGH_UPSTREAM_ERROR_BODY_MAX_LOG_CHARS:
break
except httpx.HTTPError as err:
partial: Final = b"".join(collected)
verbose_proxy_logger.warning(
"pass_through_endpoint: upstream error body read failed after %d bytes: %s",
len(partial),
type(err).__name__,
)
return partial, _no_more_chunks()
return b"".join(collected), stream
def _headers_without_body_framing(headers: httpx.Headers) -> httpx.Headers:
return httpx.Headers(
[(name, value) for name, value in headers.raw if name.lower() not in (b"content-encoding", b"content-length")]
)
async def _error_body_preview_and_relay(response: httpx.Response) -> tuple[str, httpx.Response]:
if response.is_stream_consumed:
return response.text, response
body_iter: Final = response.aiter_bytes()
prefix, rest = await _read_error_body_preview(body_iter)
preview_text: Final = prefix.decode(response.encoding or "utf-8", errors="replace")
return preview_text, httpx.Response(
status_code=response.status_code,
headers=_headers_without_body_framing(response.headers),
stream=_PrefixReplayStream(prefix=prefix, rest=rest, upstream=response),
request=response.request,
extensions=response.extensions,
)
async def _log_passthrough_upstream_failure(
response: httpx.Response,
user_api_key_dict: UserAPIKeyAuth,
request_payload: dict,
) -> None:
"""Fire LiteLLM-side failure hooks (spend tracking, alerting callbacks) for
an upstream 4xx/5xx passthrough response.
Passthrough must return the upstream status/body/headers to the client
unchanged, so this never raises or transforms the response - it only
mirrors the monitoring side effect that ``post_call_failure_hook`` would
have received had the error originated inside LiteLLM.
"""
logging_obj: LiteLLMLoggingObj,
) -> httpx.Response:
if response.status_code < 400:
return
return response
from litellm.proxy.proxy_server import proxy_logging_obj
preview_text, relay_response = await _error_body_preview_and_relay(response)
upstream_error_body: Final = (
REDACTED_BY_LITELLM
if should_redact_message_logging(logging_obj.model_call_details)
else _truncate_upstream_error_body(_sanitize_upstream_error_body(preview_text))
)
verbose_proxy_logger.warning(
"pass_through_endpoint: upstream %s %s returned %s: %s",
response.request.method,
response.url.copy_with(query=None, fragment=None),
response.status_code,
upstream_error_body,
)
try:
response.raise_for_status()
except httpx.HTTPStatusError:
@ -878,7 +964,7 @@ async def _log_passthrough_upstream_failure(
# rate-limit errors already are.
synthetic_exception: Final = HTTPException(
status_code=response.status_code,
detail=f"Upstream passthrough request failed with status {response.status_code}",
detail=f"Upstream passthrough request failed with status {response.status_code}: {upstream_error_body}",
)
try:
await proxy_logging_obj.post_call_failure_hook(
@ -892,6 +978,7 @@ async def _log_passthrough_upstream_failure(
"pass_through_endpoint: post_call_failure_hook raised for upstream error",
exc_info=True,
)
return relay_response
async def _relay_reporting_failures(
@ -1321,7 +1408,7 @@ async def pass_through_request(
headers=response.headers,
)
await _log_passthrough_upstream_failure(
relay_response: Final = await _log_passthrough_upstream_failure(
response=response,
user_api_key_dict=user_api_key_dict,
request_payload=_build_passthrough_failure_request_payload(
@ -1331,17 +1418,18 @@ async def pass_through_request(
custom_llm_provider=custom_llm_provider,
upstream_usage=upstream_usage,
),
logging_obj=logging_obj,
)
# Call response headers hook for streaming pass-through
_response_headers = HttpPassThroughEndpointHelpers.get_response_headers(
headers=response.headers,
headers=relay_response.headers,
litellm_call_id=litellm_call_id,
)
callback_headers = await proxy_logging_obj.post_call_response_headers_hook(
data=_parsed_body or {},
user_api_key_dict=user_api_key_dict,
response=response,
response=relay_response,
request_headers=dict(request.headers),
)
if callback_headers:
@ -1352,7 +1440,7 @@ async def pass_through_request(
stream=_own_streamed_managed_ids(
stream=_relay_reporting_failures(
stream=PassThroughStreamingHandler.chunk_processor(
response=response,
response=relay_response,
request_body=_parsed_body,
litellm_logging_obj=logging_obj,
endpoint_type=endpoint_type,
@ -1360,7 +1448,7 @@ async def pass_through_request(
passthrough_success_handler_obj=pass_through_endpoint_logging,
url_route=str(url),
),
upstream_status=response.status_code,
upstream_status=relay_response.status_code,
user_api_key_dict=user_api_key_dict,
request_payload=_build_passthrough_failure_request_payload(
parsed_body=_parsed_body,
@ -1374,10 +1462,10 @@ async def pass_through_request(
user_api_key_dict=user_api_key_dict,
),
ping_interval_seconds=litellm.sse_keepalive_ping_interval_seconds,
upstream_headers=response.headers,
upstream_headers=relay_response.headers,
),
headers=_response_headers,
status_code=response.status_code,
status_code=relay_response.status_code,
)
if state_raw_body is not None:
@ -1412,7 +1500,7 @@ async def pass_through_request(
logging_obj.stream = True
logging_obj.model_call_details["stream"] = True
await _log_passthrough_upstream_failure(
detected_relay_response: Final = await _log_passthrough_upstream_failure(
response=response,
user_api_key_dict=user_api_key_dict,
request_payload=_build_passthrough_failure_request_payload(
@ -1422,17 +1510,18 @@ async def pass_through_request(
custom_llm_provider=custom_llm_provider,
upstream_usage=upstream_usage,
),
logging_obj=logging_obj,
)
# Call response headers hook for detected streaming pass-through
_response_headers = HttpPassThroughEndpointHelpers.get_response_headers(
headers=response.headers,
headers=detected_relay_response.headers,
litellm_call_id=litellm_call_id,
)
callback_headers = await proxy_logging_obj.post_call_response_headers_hook(
data=_parsed_body or {},
user_api_key_dict=user_api_key_dict,
response=response,
response=detected_relay_response,
request_headers=dict(request.headers),
)
if callback_headers:
@ -1443,7 +1532,7 @@ async def pass_through_request(
stream=_own_streamed_managed_ids(
stream=_relay_reporting_failures(
stream=PassThroughStreamingHandler.chunk_processor(
response=response,
response=detected_relay_response,
request_body=_parsed_body,
litellm_logging_obj=logging_obj,
endpoint_type=endpoint_type,
@ -1451,7 +1540,7 @@ async def pass_through_request(
passthrough_success_handler_obj=pass_through_endpoint_logging,
url_route=str(url),
),
upstream_status=response.status_code,
upstream_status=detected_relay_response.status_code,
user_api_key_dict=user_api_key_dict,
request_payload=_build_passthrough_failure_request_payload(
parsed_body=_parsed_body,
@ -1465,10 +1554,10 @@ async def pass_through_request(
user_api_key_dict=user_api_key_dict,
),
ping_interval_seconds=litellm.sse_keepalive_ping_interval_seconds,
upstream_headers=response.headers,
upstream_headers=detected_relay_response.headers,
),
headers=_response_headers,
status_code=response.status_code,
status_code=detected_relay_response.status_code,
)
if not _should_buffer_passthrough_response(response):
@ -1526,6 +1615,7 @@ async def pass_through_request(
response=response,
user_api_key_dict=user_api_key_dict,
request_payload=failure_request_payload,
logging_obj=logging_obj,
)
if response.status_code < 400 and response_body is not None and guardrails_to_run:
@ -3435,7 +3525,7 @@ async def _filter_endpoints_by_team_allowed_routes(
for endpoint in pass_through_endpoints
if endpoint.path
in cast( # cast-ok: guarded above; team metadata stores this key as a list of route paths
"Sequence[str]", team_metadata.get("allowed_passthrough_routes")
Sequence[str], team_metadata.get("allowed_passthrough_routes")
)
]

View file

@ -5,7 +5,7 @@ from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Mapping, S
from enum import Enum
from functools import partial
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, NamedTuple, Protocol, cast, get_args
from typing import TYPE_CHECKING, Any, Final, NamedTuple, Protocol, TypeAlias, cast, get_args
import fastapi
from fastapi import APIRouter, Depends, HTTPException, Request, Response
@ -48,7 +48,7 @@ if TYPE_CHECKING:
router: Final = APIRouter()
_ResponseDocSchemas = dict[int | str, dict[str, Any]] # pyright: ignore[reportExplicitAny] # fastapi's responses kwarg
_ResponseDocSchemas: TypeAlias = dict[int | str, dict[str, object]] # fastapi's responses kwarg
RESPONSES_API_RESPONSE_SCHEMAS: Final[_ResponseDocSchemas] = {200: {"model": ResponsesAPIResponse}}
RESPONSES_API_CREATE_RESPONSE_SCHEMAS: Final[_ResponseDocSchemas] = {

View file

@ -198,10 +198,6 @@ async def _scan_pending(prisma_client: "PrismaClient") -> _PendingScan:
return _PendingScan(marker, db_now.now, tuple(_DateRow.model_validate(row).date for row in rows))
async def pending_days(prisma_client: "PrismaClient") -> tuple[str, ...]:
return (await _scan_pending(prisma_client)).days
async def reconcile_day(prisma_client: "PrismaClient", day: str) -> None:
"""Rewrite one day of the global table from the per-key sums. Idempotent: a rerun
overwrites every group with the same totals."""

View file

@ -974,6 +974,14 @@ def _failure_usage_to_lift(
_EMPTY_LIFT: Final = MappingProxyType({})
def _reached_deployment(litellm_logging_obj: Logging) -> bool:
"""A provider handoff or a cached response both mean the router selected a deployment."""
caching_details: Final = litellm_logging_obj.caching_details
return litellm_logging_obj.model_call_details.get("first_api_call_start_time") is not None or (
caching_details is not None and caching_details.get("cache_hit") is True
)
def _stamp_deployment_attribution(
litellm_params: dict[str, object], model_group: str | None, team_id: str | None, dispatched: bool
) -> Mapping[str, object]:
@ -2381,7 +2389,6 @@ class ProxyLogging:
)
try:
# Execute guardrail pipelines before the normal callback loop
if not skip_guardrails:
data, _ = await self._maybe_execute_pipelines( # rebind-ok: pipeline edits feed the callback loop below
data=data,
@ -3325,7 +3332,7 @@ class ProxyLogging:
_litellm_params,
request_data.get("model"),
user_api_key_dict.team_id,
dispatched=litellm_logging_obj.model_call_details.get("first_api_call_start_time") is not None,
dispatched=_reached_deployment(litellm_logging_obj),
)
litellm_logging_obj.update_environment_variables(

View file

@ -2254,7 +2254,7 @@ class LiteLLMCompletionResponsesConfig:
) -> Mapping[str, ResponseFunctionWebSearch]:
calls: Final[dict[str, ResponseFunctionWebSearch]] = {} # mutable-ok: indexes provider-built calls
for choice in chat_completion_response.choices:
provider_fields = getattr(choice.message, "provider_specific_fields", None)
provider_fields = choice.message.provider_specific_fields
if not isinstance(provider_fields, Mapping):
continue
web_search_calls = provider_fields.get("web_search_calls")

View file

@ -1360,7 +1360,7 @@ def _billed_terminal_response(
return None
usage: Final[object] = response_obj.get("usage") # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # a model_constructed terminal event leaves response as an untyped dict
return ResponsesAPIResponse.model_construct(
**{**response_obj, "usage": usage if usage is not None or estimate is None else estimate()} # pyright: ignore[reportUnknownArgumentType, reportArgumentType] # same untyped dict spread
**{**response_obj, "usage": usage if usage is not None or estimate is None else estimate()} # pyright: ignore[reportArgumentType] # same untyped dict spread
)

View file

@ -1,7 +1,7 @@
from collections.abc import Mapping
from datetime import datetime, timezone
from types import MappingProxyType
from typing import Annotated, Final, Literal, NamedTuple, Protocol
from typing import Annotated, Final, Literal, NamedTuple, Protocol, TypeAlias
from uuid import uuid4
import httpx
@ -24,7 +24,7 @@ from litellm.proxy.pass_through_endpoints.llm_provider_handlers.typesafe_passthr
from litellm.router_strategy.complexity_router.config import DEFAULT_JEV_INSTRUCTIONS as _DEFAULT_JEV_INSTRUCTIONS
from litellm.types.utils import AUTOROUTER_CLASSIFIER_CALL_ORIGIN
JevProbability = Annotated[float, Field(ge=0.0, le=1.0)]
JevProbability: TypeAlias = Annotated[float, Field(ge=0.0, le=1.0)]
DEFAULT_JEV_INSTRUCTIONS: Final = _DEFAULT_JEV_INSTRUCTIONS

View file

@ -50,6 +50,7 @@ from litellm.exceptions import (
from litellm.integrations.custom_logger import CustomLogger, Span
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.litellm_core_utils.prompt_templates.common_utils import (
anthropic_content_lists,
encrypted_content_of_block,
strip_encrypted_reasoning_from_messages,
)
@ -155,10 +156,7 @@ class EncryptedContentAffinityCheck(CustomLogger):
return iter(())
return (
cast(Mapping[str, object], block) # cast-ok: narrowed by isinstance
for message in cast(list[object], messages) # cast-ok: narrowed by isinstance
if isinstance(message, Mapping)
for content in (cast(Mapping[str, object], message).get("content"),) # cast-ok: narrowed by isinstance
if isinstance(content, list)
for content in anthropic_content_lists(cast(list[object], messages)) # cast-ok: narrowed by isinstance
for block in cast(list[object], content) # cast-ok: narrowed by isinstance
if isinstance(block, Mapping)
)

View file

@ -2,7 +2,7 @@ from __future__ import annotations
from collections.abc import Awaitable, Callable, Mapping
from dataclasses import dataclass
from typing import Final, Generic, TypeVar
from typing import Final, Generic, TypeAlias, TypeVar
from litellm.rust_bridge import catalog, runtime
from litellm.rust_bridge.bindings import NativeBinding
@ -10,11 +10,11 @@ from litellm.rust_bridge.catalog import Route, RouteContext, RouteRule, Rules
from litellm.rust_bridge.configuration import Decision
from litellm.rust_bridge.configuration import decision as rollout_decision
RequestT = TypeVar("RequestT")
NativeT = TypeVar("NativeT")
ResultT = TypeVar("ResultT")
RequestT: Final = TypeVar("RequestT")
NativeT: Final = TypeVar("NativeT")
ResultT: Final = TypeVar("ResultT")
NativeHook = Callable[[RequestT, tuple[object, ...], Mapping[str, object]], ResultT]
NativeHook: TypeAlias = Callable[[RequestT, tuple[object, ...], Mapping[str, object]], ResultT]
def call_hook(

View file

@ -1,10 +1,10 @@
from typing import TypeVar
from typing import Final, TypeVar
from litellm.router_utils.add_retry_fallback_headers import (
_add_headers_to_response, # pyright: ignore[reportPrivateUsage] # reuse the proxy's identity-preserving response metadata writer
)
ResultT = TypeVar("ResultT")
ResultT: Final = TypeVar("ResultT")
def mark_rust_response(response: ResultT) -> ResultT:

View file

@ -9,3 +9,5 @@ class s3BatchLoggingElement(BaseModel):
payload: dict
s3_object_key: str
s3_object_download_filename: str
body: str | None = None
content_type: str = "application/json"

View file

@ -558,7 +558,6 @@ class AmazonTitanMultimodalEmbeddingResponse(TypedDict):
message: str # Specifies any errors that occur during generation.
# TwelveLabs Marengo Embed types
TWELVELABS_EMBEDDING_INPUT_TYPES = Literal["text", "image", "video", "audio"]
TWELVELABS_EMBEDDING_OPTIONS = Literal["visual-text", "visual-image", "audio"]

View file

@ -3860,7 +3860,7 @@ def without_server_derived_pricing(model_info: Mapping[str, Any]) -> Mapping[str
)
def echoed_cost_map_pricing_fields(model_info: Mapping[str, Any]) -> tuple[str, ...]:
def echoed_cost_map_pricing_fields(model_info: Mapping[str, object]) -> tuple[str, ...]:
"""Pricing fields a stored ``model_info`` blob copied from a ``/model/info`` response.
Only ``litellm.get_model_info`` emits ``key`` (the resolved cost-map entry), so a stored
@ -3891,7 +3891,7 @@ def echoed_cost_map_fields(
)
def pricing_override_fields(*sources: Mapping[str, Any]) -> tuple[str, ...]:
def pricing_override_fields(*sources: Mapping[str, object]) -> tuple[str, ...]:
return tuple(
sorted(
frozenset(

File diff suppressed because it is too large Load diff

View file

@ -35,3 +35,5 @@ The extensions shard uses the built-in generic callback and guardrail transports
The mcp shard runs the MCP gateway against SDK peers owned by each test (`_support/mcp.py`): streamable HTTP, SSE and stdio peers, an OpenAPI-spec app, and an OAuth 2.1 authorization-server double. Every peer records the requests it receives so a test can assert what reached the peer, not only what the proxy answered. The shard runs with `INTEGRATION_WORKERS` set and with `INTEGRATION_COVERAGE=1`, which starts the proxy under `coverage run --parallel-mode` limited to the MCP modules and stores `coverage.txt` plus an HTML report with the job artifacts. A test that fails because the product is wrong is skipped with `pytest.skip("BUG: <symptom>")` so the skip list in `execution.json` is the open MCP bug list
Browser contracts live in `tests/e2e/ui/tests/integrationCritical` and run only through `tests/e2e/ui/integration.config.ts`. The expected browser results are listed in `expected.json` in that directory and checked by `.circleci/scripts/verify_integration_browser.py`. The CircleCI browser shard builds the checked-out dashboard, starts the owned proxy with that build, and verifies one exact browser result without retries or skips. The default Playwright selection excludes this directory. The focused project flow asserts the submitted create and clear values, fresh SQL state and actual blocked/restored serving while preserving model restrictions
Two always-on `-replica` CircleCI jobs (management, database) run their groups in replica mode, where every proxy connects through a real `litellm_writer` role and a real read-only `litellm_reader` role against the same PostgreSQL. Nothing is captured there: the job passes when the tests pass, and a write routed to the read-only reader fails the test that issued it. A deeper check runs on demand as the `routing_parity` workflow, triggered through the CircleCI API v2 pipeline endpoint on the PR branch with `{"parameters": {"routing_parity_base": "<40-hex merge-base sha>"}}`. The workflow fans out over the seven groups, and each `routing-parity-<group>` job runs its own group twice against the same test harness, once with `litellm/`, `enterprise/`, and `litellm-proxy-extras/` checked out from the base revision and once from the head, with a pytest plugin snapshotting `pg_stat_statements` into `routing-observed.json` per side. The `check` step then compares the two observations and writes `routing-diff.txt`: a statement seen on both sides fails when its role set changed, globally or for the same test (per-test capture is skipped under xdist), unless it is listed in `tests/integration/routing/either_role.json`, where each entry names the statement and a one-line reason it legitimately runs on whichever role asks for it, printed under `== either role ==`. Queries seen on only one side are listed, never failed, `pg_stat_statements` evictions and a role that never ran a statement are failures

View file

@ -9,6 +9,7 @@ from collections.abc import Iterator, Mapping
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from types import MappingProxyType
from typing import Final
import httpx
@ -16,6 +17,17 @@ import psutil
from integration._support.client import Gateway
def proxy_database_environment() -> Mapping[str, str]:
writer: Final = os.environ.get("INTEGRATION_PROXY_DATABASE_URL", "")
reader: Final = os.environ.get("INTEGRATION_PROXY_READ_REPLICA_URL", "")
return MappingProxyType(
{
**({"DATABASE_URL": writer} if writer else {}),
**({"DATABASE_URL_READ_REPLICA": reader} if reader else {}),
}
)
def in_group(process: psutil.Process, group: int) -> bool:
try:
return os.getpgid(process.pid) == group
@ -83,7 +95,11 @@ def owned_proxy_process(
port: Final = reserve.getsockname()[1]
root: Final = Path(os.environ.get("INTEGRATION_PROXY_ROOT") or Path(__file__).resolve().parents[3])
environment: Final = {
**{name: value for name, value in os.environ.items() if name not in remove_environment},
**{
name: value
for name, value in {**os.environ, **proxy_database_environment()}.items()
if name not in remove_environment
},
"LITELLM_MASTER_KEY": gateway.key,
"LITELLM_SALT_KEY": os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt"),
"STORE_MODEL_IN_DB": "True",

View file

@ -0,0 +1,333 @@
from __future__ import annotations
import argparse
import itertools
import json
import os
import re
import sys
from collections.abc import Iterator, Mapping
from dataclasses import dataclass
from pathlib import Path
from types import MappingProxyType
from typing import Final
from urllib.parse import urlsplit, urlunsplit
import psycopg
import pytest
from pydantic import TypeAdapter
WRITER_ROLE: Final = "litellm_writer"
READER_ROLE: Final = "litellm_reader"
ROLES: Final = (READER_ROLE, WRITER_ROLE)
DATABASE_NAME: Final = "circle_test"
OBSERVED_FILE: Final = "routing-observed.json"
DIFF_FILE: Final = "routing-diff.txt"
EITHER_ROLE_FILE: Final = Path(__file__).resolve().parents[1] / "routing" / "either_role.json"
RoleSet = frozenset[str]
RoutingMap = Mapping[str, frozenset[str]]
Snapshot = Mapping[tuple[str, str], int]
_PLACEHOLDERS: Final = re.compile(r"\$\d+(?:\s*,\s*\$\d+)*")
_QUERIES: Final = TypeAdapter(dict[str, tuple[str, ...]])
_OBSERVED: Final = TypeAdapter(dict[str, object])
def normalize(query: str) -> str:
return _PLACEHOLDERS.sub("$n", " ".join(query.split()))
@dataclass(frozen=True, slots=True)
class Observation:
queries: RoutingMap
tests: Mapping[str, RoutingMap]
calls: Mapping[str, int]
dealloc: int
@dataclass(frozen=True, slots=True)
class Mismatch:
test: str | None
query: str
base: tuple[str, ...]
head: tuple[str, ...]
@dataclass(frozen=True, slots=True)
class Report:
mismatches: tuple[Mismatch, ...]
only_base: tuple[str, ...]
only_head: tuple[str, ...]
calls: Mapping[str, Mapping[str, int]]
dealloc: Mapping[str, int]
either_role: tuple[str, ...] = ()
def failures(self) -> tuple[str, ...]:
mismatch_failures: Final = tuple(
f"{mismatch.test if mismatch.test is not None else 'global'}: {mismatch.query}: "
f"base [{', '.join(mismatch.base)}] head [{', '.join(mismatch.head)}]"
for mismatch in self.mismatches
)
side_failures: Final = tuple(
failure
for side in ("base", "head")
for failure in (
*(
(f"{side}: pg_stat_statements evicted {self.dealloc[side]} entries (dealloc > 0)",)
if self.dealloc[side] > 0
else ()
),
*(f"{side}: no {role} calls observed" for role in ROLES if self.calls[side].get(role, 0) == 0),
)
)
return (*mismatch_failures, *side_failures)
def _sorted_map(value: RoutingMap) -> RoutingMap:
return MappingProxyType(dict(sorted(value.items())))
def compare(base: Observation, head: Observation, either_role: frozenset[str] = frozenset()) -> Report:
mismatches: Final = (
*(
Mismatch(
None,
query,
tuple(sorted(base_roles)),
tuple(sorted(head.queries[query])),
)
for query, base_roles in base.queries.items()
if query in head.queries and head.queries[query] != base_roles and query not in either_role
),
*(
Mismatch(
test,
query,
tuple(sorted(base_roles)),
tuple(sorted(head.tests[test][query])),
)
for test, queries in base.tests.items()
if test in head.tests
for query, base_roles in queries.items()
if query in head.tests[test] and head.tests[test][query] != base_roles and query not in either_role
),
)
varying: Final = frozenset(
query
for query in either_role
if (query in base.queries and query in head.queries and head.queries[query] != base.queries[query])
or any(
query in base.tests[test]
and query in head.tests[test]
and head.tests[test][query] != base.tests[test][query]
for test in frozenset(base.tests) & frozenset(head.tests)
)
)
return Report(
mismatches,
tuple(sorted(query for query in base.queries if query not in head.queries)),
tuple(sorted(query for query in head.queries if query not in base.queries)),
MappingProxyType({"base": base.calls, "head": head.calls}),
MappingProxyType({"base": base.dealloc, "head": head.dealloc}),
tuple(sorted(varying)),
)
def render(report: Report) -> str:
failures: Final = report.failures()
lines: Final = (
"== failures ==",
*(failures or ("none",)),
"",
"== either role ==",
*(report.either_role or ("none",)),
"",
"== queries only in base ==",
*(report.only_base or ("none",)),
"",
"== queries only in head ==",
*(report.only_head or ("none",)),
"",
"== calls ==",
*(
line
for side in ("base", "head")
for line in (
*(f"{side} {role}: {report.calls[side].get(role, 0)}" for role in ROLES),
f"{side} dealloc: {report.dealloc[side]}",
)
),
)
return "\n".join(lines) + "\n"
def _roles(document: Mapping[str, tuple[str, ...]]) -> RoutingMap:
return _sorted_map({query: frozenset(roles) for query, roles in document.items()})
def _tests(document: Mapping[str, Mapping[str, tuple[str, ...]]]) -> Mapping[str, RoutingMap]:
return MappingProxyType({node: _roles(queries) for node, queries in document.items()})
def load_observation(path: Path) -> Observation:
document: Final = _OBSERVED.validate_python(json.loads(path.read_text()))
queries: Final = _QUERIES.validate_python(document.get("queries", {}))
tests: Final = TypeAdapter(dict[str, dict[str, tuple[str, ...]]]).validate_python(document.get("tests", {}))
calls: Final = TypeAdapter(dict[str, int]).validate_python(document.get("calls", {}))
dealloc: Final = TypeAdapter(int).validate_python(document.get("dealloc", 0))
return Observation(_roles(queries), _tests(tests), MappingProxyType(calls), dealloc)
def load_either_role(path: Path) -> frozenset[str]:
if not path.exists():
return frozenset()
document: Final = TypeAdapter(dict[str, str]).validate_python(json.loads(path.read_text()))
return frozenset(document)
def _serializable(queries: RoutingMap, tests: Mapping[str, RoutingMap]) -> dict[str, object]:
return {
"queries": {query: sorted(roles) for query, roles in queries.items()},
"tests": {node: {query: sorted(roles) for query, roles in mapping.items()} for node, mapping in tests.items()},
}
def dump_observation(observation: Observation) -> str:
document: Final = _serializable(observation.queries, observation.tests)
return (
json.dumps(
{**document, "calls": dict(observation.calls), "dealloc": observation.dealloc},
indent=2,
sort_keys=True,
)
+ "\n"
)
def _maintenance_url() -> str:
parsed: Final = urlsplit(os.environ["DATABASE_URL"])
return urlunsplit(parsed._replace(path="/postgres"))
def snapshot(connection: psycopg.Connection[object]) -> Mapping[tuple[str, str], int]:
rows: Final = connection.execute(
"""
SELECT r.rolname, s.query, s.calls
FROM pg_stat_statements s
JOIN pg_roles r ON r.oid = s.userid
WHERE s.dbid = (SELECT oid FROM pg_database WHERE datname = %s)
AND r.rolname = ANY(%s)
""",
(DATABASE_NAME, list(ROLES)),
).fetchall()
return MappingProxyType(
{
key: sum(calls for _, _, calls in grouped)
for key, grouped in itertools.groupby(
sorted((str(role), normalize(str(query)), int(calls)) for role, query, calls in rows),
key=lambda row: (row[0], row[1]),
)
}
)
def delta(before: Snapshot, after: Snapshot) -> RoutingMap:
pairs: Final = {key: after.get(key, 0) - before.get(key, 0) for key in frozenset(before) | frozenset(after)}
queries: Final = frozenset(query for (_, query), change in pairs.items() if change > 0)
return MappingProxyType(
{query: frozenset(role for role in ROLES if pairs.get((role, query), 0) > 0) for query in sorted(queries)}
)
def role_calls(before: Snapshot, after: Snapshot) -> Mapping[str, int]:
return MappingProxyType(
{
role: sum(
max(after.get((role, query), 0) - before.get((role, query), 0), 0)
for query in frozenset(q for _, q in before) | frozenset(q for _, q in after)
)
for role in ROLES
}
)
def dealloc(connection: psycopg.Connection[object]) -> int:
return int(connection.execute("SELECT dealloc FROM pg_stat_statements_info").fetchone()[0])
class RoutingPlugin:
def __init__(self, config: pytest.Config) -> None:
self.config = config
self._session_start: Snapshot | None = None
self._tests: tuple[tuple[str, RoutingMap], ...] = ()
def _snapshot(self) -> Snapshot:
with psycopg.connect(_maintenance_url(), autocommit=True) as connection:
return snapshot(connection)
def pytest_sessionstart(self, session: pytest.Session) -> None:
if hasattr(self.config, "workerinput"):
return
self._session_start = self._snapshot()
@pytest.hookimpl(hookwrapper=True)
def pytest_runtest_protocol(self, item: pytest.Item, nextitem: pytest.Item | None) -> Iterator[None]:
if self.config.getoption("numprocesses", default=None) or hasattr(self.config, "workerinput"):
yield
return
before: Final = self._snapshot()
yield
after: Final = self._snapshot()
self._tests = (*self._tests, (item.nodeid, delta(before, after)))
def pytest_sessionfinish(self, session: pytest.Session, exitstatus: int) -> None:
if hasattr(self.config, "workerinput"):
return
end: Final = self._snapshot()
with psycopg.connect(_maintenance_url(), autocommit=True) as connection:
evictions: Final = dealloc(connection)
start: Final = self._session_start or {}
tests: Final = MappingProxyType({node: mapping for node, mapping in self._tests})
destination: Final = Path(os.environ["INTEGRATION_RESULTS_DIR"])
destination.mkdir(parents=True, exist_ok=True)
(destination / OBSERVED_FILE).write_text(
dump_observation(
Observation(
delta(start, end),
tests,
role_calls(start, end),
evictions,
)
)
)
def main(argv: tuple[str, ...] | list[str]) -> int:
parser: Final = argparse.ArgumentParser()
parser.add_argument("command", choices=("check",))
parser.add_argument("base_dir", type=Path)
parser.add_argument("head_dir", type=Path)
parser.add_argument("--either-role", type=Path, default=EITHER_ROLE_FILE)
parser.add_argument("--diff", type=Path, default=None)
options: Final = parser.parse_args(argv)
base_path: Final = options.base_dir / OBSERVED_FILE
head_path: Final = options.head_dir / OBSERVED_FILE
for path in (base_path, head_path):
if not path.exists():
sys.stderr.write(f"observed routing file missing: {path}\n")
if not base_path.exists() or not head_path.exists():
return 1
report: Final = compare(
load_observation(base_path),
load_observation(head_path),
load_either_role(options.either_role),
)
diff: Final = render(report)
(options.diff or options.head_dir.parent / DIFF_FILE).write_text(diff)
sys.stdout.write(diff)
return 1 if report.failures() else 0
if __name__ == "__main__":
raise SystemExit(main(sys.argv[1:]))

View file

@ -43,7 +43,9 @@ class Wire:
@contextmanager
def wire_server(respond: Callable[[Request], Reply], tls: ssl.SSLContext | None = None) -> Generator[Wire, None, None]:
def wire_server(
respond: Callable[[Request], Reply], tls: ssl.SSLContext | None = None, port: int = 0
) -> Generator[Wire, None, None]:
"""Owned TCP peer; requests traverse the real HTTP client and serialization."""
received: Final[SimpleQueue[Request]] = SimpleQueue()
errors: Final[SimpleQueue[Exception]] = SimpleQueue()
@ -114,7 +116,7 @@ def wire_server(respond: Callable[[Request], Reply], tls: ssl.SSLContext | None
if tls is not None:
self.socket = tls.wrap_socket(self.socket, server_side=True)
with OwnedHTTPServer(("127.0.0.1", 0), Handler) as server:
with OwnedHTTPServer(("127.0.0.1", port), Handler) as server:
thread: Final = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.05})
thread.start()
try:

View file

@ -0,0 +1,154 @@
import json
import uuid
from collections.abc import Mapping
from pathlib import Path
from typing import Final
import pytest
from pydantic import JsonValue
from tests.integration._support.client import Gateway
from tests.integration._support.process import owned_proxy
SERVING_CONSUMERS: Final = {
"compression_interception": "CompressionInterceptionLogger",
"code_interpreter_interception": "CodeInterpreterInterceptionLogger",
"websearch_interception": "WebSearchInterceptionLogger",
}
OTEL_CONSUMER: Final = {"otel": "OpenTelemetry"}
GUARDRAIL_CONSUMERS: Final = {
"presidio": "_OPTIONAL_PresidioPIIMasking",
"lakera_prompt_injection": "lakeraAI_Moderation",
}
TOP_LEVEL_SHAPES: Final = (
pytest.param({}, id="empty-object"),
pytest.param(None, id="null"),
pytest.param("otel", id="string"),
pytest.param(["otel"], id="list"),
pytest.param(True, id="bool"),
pytest.param(0, id="zero"),
)
CONSUMER_SHAPES: Final = (
pytest.param({}, id="empty-object"),
pytest.param(None, id="null"),
pytest.param("on", id="string"),
pytest.param(True, id="bool"),
pytest.param([], id="empty-list"),
pytest.param(["on"], id="list"),
pytest.param(7, id="int"),
)
TOP_LEVEL_BOOT_CRASH: Final = (
"BUG: a non-object callback_settings is stored verbatim and proxy startup crashes calling .get on it"
)
TOP_LEVEL_BOOT_CRASH_IDS: Final = frozenset({"string", "list", "bool"})
OTEL_DROPPED: Final = (
"BUG: a non-object callback_settings.otel fails dict() and the otel callback is silently not registered"
)
OTEL_DROPPED_IDS: Final = frozenset({"null", "string", "bool", "int"})
def _write_config(
directory: Path, upstream_url: str, model: str, callbacks: tuple[str, ...], callback_settings: JsonValue
) -> Path:
config: Final = directory / f"callback_settings_{uuid.uuid4().hex}.yaml"
config.write_text(
json.dumps(
{
"model_list": [
{
"model_name": model,
"litellm_params": {
"model": f"openai/{model}",
"api_base": f"{upstream_url}/v1",
"api_key": "integration-provider-key",
},
}
],
"litellm_settings": {"callbacks": list(callbacks)},
"callback_settings": callback_settings,
"general_settings": {
"master_key": "os.environ/LITELLM_MASTER_KEY",
"database_url": "os.environ/DATABASE_URL",
},
}
)
)
return config
def _assert_registered(candidate: Gateway, consumers: Mapping[str, str]) -> None:
response: Final = candidate.request("GET", "/active/callbacks")
assert response.status_code == 200, response.text
missing: Final = sorted(name for name, class_name in consumers.items() if class_name not in response.text)
assert missing == [], response.text
@pytest.mark.parametrize("callback_settings", TOP_LEVEL_SHAPES)
def test_top_level_callback_settings_shape_boots_registers_and_serves_chat(
gateway: Gateway, tmp_path: Path, callback_settings: JsonValue, request: pytest.FixtureRequest
) -> None:
if request.node.callspec.id in TOP_LEVEL_BOOT_CRASH_IDS:
pytest.skip(TOP_LEVEL_BOOT_CRASH)
consumers: Final = {**SERVING_CONSUMERS, **OTEL_CONSUMER}
model: Final = f"integration-callback-settings-{uuid.uuid4().hex}"
config: Final = _write_config(tmp_path, gateway.upstream_url, model, tuple(consumers), callback_settings)
with owned_proxy(gateway, tmp_path, {"STORE_MODEL_IN_DB": "False"}, config=config) as candidate:
_assert_registered(candidate, consumers)
reply: Final = candidate.chat(model, text=f"callback settings {uuid.uuid4().hex}")
assert reply["model"] == model, reply
@pytest.mark.parametrize("value", CONSUMER_SHAPES)
def test_serving_consumer_settings_shape_boots_registers_and_serves_chat(
gateway: Gateway, tmp_path: Path, value: JsonValue
) -> None:
model: Final = f"integration-callback-settings-{uuid.uuid4().hex}"
config: Final = _write_config(
tmp_path,
gateway.upstream_url,
model,
tuple(SERVING_CONSUMERS),
{consumer: value for consumer in SERVING_CONSUMERS},
)
with owned_proxy(gateway, tmp_path, {"STORE_MODEL_IN_DB": "False"}, config=config) as candidate:
_assert_registered(candidate, SERVING_CONSUMERS)
reply: Final = candidate.chat(model, text=f"callback settings {uuid.uuid4().hex}")
assert reply["model"] == model, reply
@pytest.mark.parametrize("value", CONSUMER_SHAPES)
def test_otel_settings_shape_boots_registers_and_serves_chat(
gateway: Gateway, tmp_path: Path, value: JsonValue, request: pytest.FixtureRequest
) -> None:
if request.node.callspec.id in OTEL_DROPPED_IDS:
pytest.skip(OTEL_DROPPED)
model: Final = f"integration-callback-settings-{uuid.uuid4().hex}"
config: Final = _write_config(tmp_path, gateway.upstream_url, model, tuple(OTEL_CONSUMER), {"otel": value})
with owned_proxy(gateway, tmp_path, {"STORE_MODEL_IN_DB": "False"}, config=config) as candidate:
_assert_registered(candidate, OTEL_CONSUMER)
reply: Final = candidate.chat(model, text=f"callback settings {uuid.uuid4().hex}")
assert reply["model"] == model, reply
@pytest.mark.parametrize("value", CONSUMER_SHAPES)
def test_guardrail_consumer_settings_shape_boots_and_registers(
gateway: Gateway, tmp_path: Path, value: JsonValue
) -> None:
model: Final = f"integration-callback-settings-{uuid.uuid4().hex}"
config: Final = _write_config(
tmp_path,
gateway.upstream_url,
model,
tuple(GUARDRAIL_CONSUMERS),
{consumer: value for consumer in GUARDRAIL_CONSUMERS},
)
environment: Final = {
"STORE_MODEL_IN_DB": "False",
"PRESIDIO_ANALYZER_API_BASE": gateway.upstream_url,
"PRESIDIO_ANONYMIZER_API_BASE": gateway.upstream_url,
}
with owned_proxy(gateway, tmp_path, environment, config=config) as candidate:
_assert_registered(candidate, GUARDRAIL_CONSUMERS)

View file

@ -15,6 +15,7 @@ from redis import Redis
from tests.integration._support.client import Gateway, eventually, gateway_from_environment
from tests.integration._support.generation import LIFECYCLE_SETTINGS
from tests.integration._support.manifest import OWNED_DIRECTORIES
from tests.integration._support.routing import RoutingPlugin
COLLECTED: Final = pytest.StashKey[tuple[str, ...]]()
REPORTS: Final = pytest.StashKey[list[pytest.TestReport]]()
@ -29,6 +30,8 @@ def pytest_configure(config: pytest.Config) -> None:
config.addinivalue_line("markers", "covers(*ids): legacy contract IDs kept for existing tests, not enforced")
config.stash[REPORTS] = []
config.pluginmanager.register(IntegrationReportPlugin(config))
if os.environ.get("INTEGRATION_ROUTING"):
config.pluginmanager.register(RoutingPlugin(config))
class IntegrationReportPlugin:

View file

@ -43,6 +43,7 @@ def test_access_group_second_key_constraint_failure_rolls_back_all_writes(gatewa
)
with psycopg.connect(os.environ["DATABASE_URL"], autocommit=True) as connection, ExitStack() as cleanup:
connection.execute(sql.SQL("CREATE SEQUENCE {}").format(sql.Identifier(witness)))
connection.execute(sql.SQL("GRANT USAGE ON SEQUENCE {} TO PUBLIC").format(sql.Identifier(witness)))
cleanup.callback(connection.execute, sql.SQL("DROP SEQUENCE {}").format(sql.Identifier(witness)))
connection.execute(
sql.SQL(

View file

@ -121,6 +121,7 @@ def test_user_budget_updates_return_promptly_while_coordination_redis_is_wedged(
"REDIS_PORT": str(coordination.port),
},
config=Path("tests/integration/coordination_redis_proxy_config.yaml"),
remove_environment=("DATABASE_URL_READ_REPLICA",),
workers=2,
) as candidate,
Redis(host=coordination.host, port=coordination.port, socket_timeout=1) as subscriber_client,

View file

@ -318,7 +318,14 @@ def test_redis_outage_keeps_config_store_served_and_recovers(
"REDIS_PORT": str(cache.port),
"REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT": "1",
}
with owned_proxy(gateway, tmp_path, overrides, config=PROXY_CONFIG, workers=2) as candidate:
with owned_proxy(
gateway,
tmp_path,
overrides,
config=PROXY_CONFIG,
workers=2,
remove_environment=("DATABASE_URL_READ_REPLICA",),
) as candidate:
db_store_id: Final = f"vs_db_{uuid.uuid4().hex}"
for phase in ("before", "during", "after"):
if phase == "during":

View file

@ -0,0 +1,344 @@
import asyncio
import json
import threading
import time
from collections.abc import Mapping
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass, field
from pathlib import Path
from types import MappingProxyType
from typing import Final
import anthropic
import openai
import yaml
from integration._support.client import Gateway, JsonValue, eventually, object_value
from integration._support.wire import Reply, Request
BUCKET: Final = "integration-bucket"
PREFIX: Final = "integration-logs"
@dataclass(slots=True)
class RecordingS3Sink:
"""Records every accepted PUT body by target, tracks peak concurrency, and can reject a leading
run of PUT attempts with a chosen status before accepting. Serves stored bodies back on GET."""
fail_attempts: int = 0
fail_until: float = 0.0
fail_status: int = 503
delay_seconds: float = 0.5
lock: threading.Lock = field(default_factory=threading.Lock)
in_flight: int = 0
peak: int = 0
attempts: int = 0
store: dict[str, bytes] = field(default_factory=dict) # mutable-ok: GET reads must see writes from earlier PUTs
def respond(self, request: Request) -> Reply:
if request.method == "GET":
body: Final = self.store.get(request.target)
if body is None:
return Reply(status=404)
return Reply(body=body)
assert request.method == "PUT", request.method
assert request.target.startswith(f"/{BUCKET}/{PREFIX}/"), request.target
with self.lock:
self.attempts += 1
if self.attempts <= self.fail_attempts or time.time() < self.fail_until:
return Reply(
status=self.fail_status,
body=b"<Error><Code>SinkFailure</Code></Error>",
content_type="application/xml",
)
self.in_flight += 1
self.peak = max(self.peak, self.in_flight)
self.store[request.target] = request.body
time.sleep(self.delay_seconds)
with self.lock:
self.in_flight -= 1
return Reply()
def objects(self) -> Mapping[str, bytes]:
with self.lock:
return MappingProxyType(dict(self.store))
def payloads(self) -> tuple[dict[str, JsonValue], ...]:
return tuple(object_value(json.loads(line)) for body in self.objects().values() for line in body.splitlines())
def s3_config(
path: Path, sink_url: str, extra: Mapping[str, JsonValue], settings: Mapping[str, JsonValue] | None = None
) -> Path:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["litellm_settings"].update(
{
"callbacks": ["s3_v2"],
"s3_callback_params": {
"s3_bucket_name": BUCKET,
"s3_region_name": "us-east-1",
"s3_endpoint_url": sink_url,
"s3_path": PREFIX,
"s3_aws_access_key_id": "AKIAIOSFODNN7EXAMPLE",
"s3_aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY",
**extra,
},
**(settings or {}),
}
)
target: Final = path / "s3_v2.yaml"
target.write_text(yaml.safe_dump(config))
return target
def _chat_completion(identity: str) -> dict[str, JsonValue]:
return {
"id": identity,
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15},
}
def _chat_stream_frames(identity: str) -> tuple[bytes, ...]:
chunks: Final = (
{
"id": identity,
"object": "chat.completion.chunk",
"created": 1,
"model": "gpt-4o-mini",
"choices": [{"index": 0, "delta": {"role": "assistant", "content": "ok"}, "finish_reason": None}],
},
{
"id": identity,
"object": "chat.completion.chunk",
"created": 1,
"model": "gpt-4o-mini",
"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15},
},
)
return tuple(f"data: {json.dumps(chunk)}\n\n".encode() for chunk in chunks) + (b"data: [DONE]\n\n",)
def _messages_completion(identity: str) -> dict[str, JsonValue]:
return {
"id": identity,
"type": "message",
"role": "assistant",
"model": "claude-sonnet-4-5-20250929",
"content": [{"type": "text", "text": "ok"}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 11, "output_tokens": 4},
}
def _messages_stream_frames(identity: str) -> tuple[bytes, ...]:
events: Final = (
(
"message_start",
{
"type": "message_start",
"message": {
"id": identity,
"type": "message",
"role": "assistant",
"model": "claude-sonnet-4-5-20250929",
"content": [],
"stop_reason": None,
"usage": {"input_tokens": 11, "output_tokens": 1},
},
},
),
(
"content_block_start",
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
),
(
"content_block_delta",
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "ok"}},
),
("content_block_stop", {"type": "content_block_stop", "index": 0}),
(
"message_delta",
{"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 4}},
),
("message_stop", {"type": "message_stop"}),
)
return tuple(f"event: {name}\ndata: {json.dumps(payload)}\n\n".encode() for name, payload in events)
def _responses_completion(identity: str) -> dict[str, JsonValue]:
return {
"id": identity,
"object": "response",
"created_at": 1,
"status": "completed",
"model": "gpt-4o-mini",
"output": [
{
"type": "message",
"id": f"msg_{identity}",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "ok", "annotations": []}],
}
],
"usage": {"input_tokens": 11, "output_tokens": 4, "total_tokens": 15},
}
def _responses_stream_frames(identity: str) -> tuple[bytes, ...]:
events: Final = (
(
"response.created",
{
"type": "response.created",
"response": {**_responses_completion(identity), "status": "in_progress", "output": []},
},
),
(
"response.output_text.delta",
{
"type": "response.output_text.delta",
"item_id": f"msg_{identity}",
"output_index": 0,
"content_index": 0,
"delta": "ok",
},
),
("response.completed", {"type": "response.completed", "response": _responses_completion(identity)}),
)
return tuple(f"event: {name}\ndata: {json.dumps(payload)}\n\n".encode() for name, payload in events)
def surface_reply(request: Request) -> Reply:
"""Scripted upstream that echoes the caller's marker string back as the response id."""
if request.method != "POST" or not request.body:
return Reply(status=404)
body: Final = json.loads(request.body)
if request.target.endswith("/chat/completions"):
identity: Final = body["messages"][0]["content"]
if body.get("stream"):
return Reply(content_type="text/event-stream", chunks=_chat_stream_frames(identity))
return Reply(body=json.dumps(_chat_completion(identity)).encode())
if request.target.endswith("/messages"):
identity_messages: Final = body["messages"][0]["content"]
if body.get("stream"):
return Reply(content_type="text/event-stream", chunks=_messages_stream_frames(identity_messages))
return Reply(body=json.dumps(_messages_completion(identity_messages)).encode())
assert request.target.endswith("/responses"), request.target
identity_responses: Final = body["input"]
if body.get("stream"):
return Reply(content_type="text/event-stream", chunks=_responses_stream_frames(identity_responses))
return Reply(body=json.dumps(_responses_completion(identity_responses)).encode())
SURFACES: Final = ("chat", "chat_stream", "messages", "messages_stream", "responses", "responses_stream")
def call_surface(
candidate: Gateway, surface: str, openai_model: str, anthropic_model: str, key: str, marker: str
) -> tuple[str, str | None]:
"""Drive one request through the given surface; return (client-visible response id, x-litellm-call-id)."""
base: Final = str(candidate.client.base_url).rstrip("/")
headers: Final = {"Authorization": f"Bearer {key}"}
if surface == "chat":
reply: Final = openai.OpenAI(base_url=f"{base}/v1", api_key=key).chat.completions.create(
model=openai_model,
messages=[{"role": "user", "content": marker}],
extra_body={"cache": {"no-cache": True}},
)
return reply.id, None
async def chat_stream() -> str:
stream = await openai.AsyncOpenAI(base_url=f"{base}/v1", api_key=key).chat.completions.create(
model=openai_model,
messages=[{"role": "user", "content": marker}],
stream=True,
extra_body={"cache": {"no-cache": True}},
)
seen = ""
async for chunk in stream:
seen = chunk.id # rebind-ok: the stream yields one chunk at a time
return seen
if surface == "chat_stream":
return asyncio.run(chat_stream()), None
if surface in ("messages", "messages_stream"):
client: Final = anthropic.Anthropic(base_url=base, api_key="anthropic-placeholder", default_headers=headers)
if surface == "messages":
reply_messages: Final = client.messages.create(
model=anthropic_model, max_tokens=16, messages=[{"role": "user", "content": marker}]
)
return reply_messages.id, None
with client.messages.stream(
model=anthropic_model, max_tokens=16, messages=[{"role": "user", "content": marker}]
) as stream:
final: Final = stream.get_final_message()
return final.id, None
if surface == "responses":
response: Final = candidate.request(
"POST",
"/v1/responses",
{"model": openai_model, "input": marker, "cache": {"no-cache": True}},
key=key,
)
assert response.status_code == 200, response.text
return str(response.json()["id"]), response.headers.get("x-litellm-call-id")
assert surface == "responses_stream", surface
with candidate.client.stream(
"POST",
"/v1/responses",
json={"model": openai_model, "input": marker, "stream": True},
headers=headers,
) as response:
text: Final = response.read().decode()
assert response.status_code == 200, text
call_id: Final = response.headers.get("x-litellm-call-id")
assert marker in text, text
return marker, call_id
def collect_payloads(sink: RecordingS3Sink, count: int, seconds: float = 60) -> tuple[dict[str, JsonValue], ...]:
"""Wait until `count` stored payload lines exist, then return every stored payload object."""
def delivered() -> int:
return sum(len(body.splitlines()) for body in sink.objects().values())
eventually(delivered, lambda total: total >= count, seconds=seconds)
return sink.payloads()
def mixed_burst(
candidate: Gateway, openai_model: str, anthropic_model: str, key: str, marker: str, per_surface: int = 8
) -> tuple[tuple[str, str | None], ...]:
"""Fire `per_surface` requests on every surface; returns (response id, x-litellm-call-id) per request."""
jobs: Final = tuple(
(surface, f"{marker}-{surface}-{index}") for surface in SURFACES for index in range(per_surface)
)
def call(job: tuple[str, str]) -> tuple[str, str | None]:
surface, identity = job
return call_surface(candidate, surface, openai_model, anthropic_model, key, identity)
with ThreadPoolExecutor(max_workers=48) as pool:
return tuple(pool.map(call, jobs))
def matched_ids(
payloads: tuple[dict[str, JsonValue], ...], answered: tuple[tuple[str, str | None], ...]
) -> frozenset[str]:
"""Every payload must be accountable to an answered request by response id or litellm_call_id."""
response_ids: Final = frozenset(observed for observed, _ in answered)
call_ids: Final = frozenset(call_id for _, call_id in answered if call_id is not None)
landed: Final = []
for payload in payloads:
if payload["id"] in response_ids:
landed.append(payload["id"])
continue
assert payload["litellm_call_id"] in call_ids, f"unmatched payload {payload['id']!r}"
landed.append(str(payload["id"]))
return frozenset(landed)

View file

@ -0,0 +1,652 @@
import asyncio
import json
import subprocess
import uuid
from collections.abc import Callable, Generator, Mapping
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from typing import Final
import anthropic
import httpx
import openai
import pytest
import yaml
from integration._support.client import Gateway, Scenario, eventually, object_value, string_value
from integration._support.database import read_rows
from integration._support.process import owned_proxy_process
from integration._support.wire import Reply, Request, Wire, wire_server
from prometheus_client.parser import text_string_to_metric_families
GUARDRAIL_PATH: Final = "/beta/litellm_basic_guardrail_api"
DEPLOYMENT_FAILURE: Final = "litellm_deployment_failure_responses_total"
DEPLOYMENT_REQUESTS: Final = "litellm_deployment_total_requests_total"
DEPLOYMENT_STATE: Final = "litellm_deployment_state"
PROXY_FAILED: Final = "litellm_proxy_failed_requests_metric_total"
def _chat_sse(marker: str) -> tuple[bytes, ...]:
chunk: Final = {
"id": "chatcmpl_" + marker,
"object": "chat.completion.chunk",
"created": 1,
"model": "gpt-4o-mini",
}
frames: Final = (
{**chunk, "choices": [{"index": 0, "delta": {"role": "assistant", "content": ""}, "finish_reason": None}]},
{**chunk, "choices": [{"index": 0, "delta": {"content": "provider control"}, "finish_reason": None}]},
{**chunk, "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]},
)
return tuple(f"data: {json.dumps(frame)}".encode() for frame in frames) + (b"data: [DONE]",)
def _provider_body(target: str, marker: str, streamed: bool) -> Reply:
match target:
case "/v1/chat/completions":
if streamed:
return Reply(content_type="text/event-stream", chunks=_chat_sse(marker))
body: dict = {
"id": "chatcmpl_" + marker,
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": "provider control " + marker},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15},
}
case "/v1/messages":
body = {
"id": "msg_" + marker,
"type": "message",
"role": "assistant",
"model": "claude-sonnet-4-5-20250929",
"content": [{"type": "text", "text": "provider control " + marker}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {"input_tokens": 11, "output_tokens": 4},
}
case "/v1/responses":
body = {
"id": "resp_" + marker,
"object": "response",
"created_at": 1,
"status": "completed",
"model": "gpt-4o-mini",
"output": [
{
"type": "message",
"id": "msg_" + marker,
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "provider control " + marker, "annotations": []}],
}
],
"usage": {"input_tokens": 11, "output_tokens": 4, "total_tokens": 15},
}
case "/v1/embeddings":
body = {
"object": "list",
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}],
"model": "text-embedding-3-small",
"usage": {"prompt_tokens": 3, "total_tokens": 3},
}
case _:
return Reply(status=404, body=json.dumps({"error": "unexpected provider target " + target}).encode())
return Reply(body=json.dumps(body).encode())
def _provider(marker: str) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
streamed: Final = b'"stream":true' in request.body.replace(b" ", b"")
return _provider_body(request.target.split("?", 1)[0], marker, streamed)
return respond
def _blocking_sink(request: Request) -> Reply:
assert request.target == GUARDRAIL_PATH, request.target
return Reply(body=json.dumps({"action": "BLOCKED", "blocked_reason": "synthetic block"}).encode())
def _failing_sink(status: int) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert request.target == GUARDRAIL_PATH, request.target
return Reply(status=status, body=json.dumps({"error": "synthetic guardrail outage"}).encode())
return respond
def _first_call_pass_sink() -> Callable[[Request], Reply]:
calls: list[int] = [] # mutable-ok: the wire handler must remember call order across requests
def respond(request: Request) -> Reply:
assert request.target == GUARDRAIL_PATH, request.target
calls.append(1)
action: dict = (
{"action": "NONE"} if len(calls) == 1 else {"action": "BLOCKED", "blocked_reason": "synthetic block"}
)
return Reply(body=json.dumps(action).encode())
return respond
def _guardrail_config(
tmp_path: Path,
name: str,
sink_url: str,
*,
mode: str = "post_call",
default_on: bool = False,
local_cache: bool = False,
ttl: int | None = None,
) -> Path:
config: dict = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["litellm_settings"]["callbacks"] = ["prometheus"]
if local_cache:
config["litellm_settings"]["cache_params"] = {"type": "local"}
if ttl is not None:
config["litellm_settings"]["cache_params"]["ttl"] = ttl
config["guardrails"] = [
{
"guardrail_name": name,
"litellm_params": {
"guardrail": "generic_guardrail_api",
"mode": mode,
"default_on": default_on,
"api_base": sink_url,
"api_key": "synthetic-guardrail-key",
},
}
]
path: Final = tmp_path / "guardrail.yaml"
path.write_text(yaml.safe_dump(config))
return path
@dataclass(frozen=True, slots=True)
class Rig:
candidate: Gateway
scenario: Scenario
model_name: str
deployment_id: str
guardrail_name: str
policy: Wire
provider: Wire
process: subprocess.Popen[bytes]
@contextmanager
def _rig(
gateway: Gateway,
tmp_path: Path,
marker: str,
*,
sink: Callable[[Request], Reply] = _blocking_sink,
mode: str = "post_call",
default_on: bool = False,
local_cache: bool = False,
ttl: int | None = None,
workers: int = 1,
upstream_model: str = "openai/gpt-4o-mini",
api_base_suffix: str = "/v1",
env: Mapping[str, str] | None = None,
) -> Generator[Rig, None, None]:
identity: Final = "guardrail-" + marker
with wire_server(sink) as policy, wire_server(_provider(marker)) as provider:
config: Final = _guardrail_config(
tmp_path, identity, policy.url, mode=mode, default_on=default_on, local_cache=local_cache, ttl=ttl
)
prom_dir: Final = tmp_path / "prom"
prom_dir.mkdir()
with (
owned_proxy_process(
gateway,
tmp_path,
{"PROMETHEUS_MULTIPROC_DIR": str(prom_dir), **(env or {})},
config=config,
workers=workers,
) as owned,
owned.gateway.scenario() as scenario,
):
model: Final = scenario.model(
model=upstream_model, api_base=provider.url + api_base_suffix, api_key="synthetic-provider-key"
)
entries: Final = owned.gateway.get("/model/info")["data"]
assert isinstance(entries, list)
entry: Final = next(item for item in entries if object_value(item)["model_name"] == model)
yield Rig(
owned.gateway,
scenario,
model,
string_value(object_value(object_value(entry)["model_info"])["id"]),
identity,
policy,
provider,
owned.process,
)
def _metric_samples(candidate: Gateway, model_name: str) -> tuple:
response: Final = candidate.client.request(
"GET", "/metrics", headers={"Authorization": f"Bearer {candidate.key}"}, follow_redirects=True
)
assert response.status_code == 200, f"GET /metrics: {response.status_code} {response.text[:300]}"
return tuple(
sample
for family in text_string_to_metric_families(response.text)
for sample in family.samples
if sample.labels.get("requested_model") == model_name
or (sample.name == DEPLOYMENT_STATE and sample.labels.get("model_id") != "")
)
def _count(samples: tuple, name: str, model_id: str) -> float:
return float(
sum(sample.value for sample in samples if sample.name == name and sample.labels.get("model_id") == model_id)
)
def _populated_failures(samples: tuple, rig: Rig, api_provider: str) -> float:
return float(
sum(
sample.value
for sample in samples
if sample.name == DEPLOYMENT_FAILURE
and sample.labels.get("model_id") == rig.deployment_id
and sample.labels.get("api_provider") == api_provider
and sample.labels.get("litellm_model_name") != ""
)
)
def _expect_metrics(
rig: Rig,
populated: float,
blank: float,
*,
api_provider: str = "openai",
pf_id: str | None = None,
pf_populated: float | None = None,
pf_blank: float | None = None,
) -> tuple:
expected_id: Final = rig.deployment_id if pf_id is None else pf_id
expected_pf_populated: Final = populated if pf_populated is None else pf_populated
expected_pf_blank: Final = blank if pf_blank is None else pf_blank
def read() -> tuple:
samples: Final = _metric_samples(rig.candidate, rig.model_name)
satisfied: Final = (
_populated_failures(samples, rig, api_provider) == populated
and _count(samples, DEPLOYMENT_FAILURE, "") == blank
and _count(samples, PROXY_FAILED, expected_id) == expected_pf_populated
and _count(samples, PROXY_FAILED, "") == expected_pf_blank
)
return samples if satisfied else ()
return eventually(read, bool, seconds=70)
def _spend_rows(call_id: str) -> tuple[dict, ...]:
rows: Final = eventually(
lambda: read_rows(
'SELECT request_id, custom_llm_provider, model_id, status FROM "LiteLLM_SpendLogs" '
"WHERE request_id = %s OR request_id LIKE %s",
(call_id, call_id + "\\_%"),
),
lambda values: len(values) >= 1,
seconds=70,
)
return tuple(dict(row) for row in rows)
def _assert_spend(call_id: str, rig: Rig, api_provider: str = "openai") -> None:
rows: Final = _spend_rows(call_id)
failures: Final = tuple(row for row in rows if row["status"] == "failure")
assert len(failures) == 1, rows
assert (failures[0]["custom_llm_provider"], failures[0]["model_id"]) == (api_provider, rig.deployment_id), rows
def _call_id(reject: httpx.Response) -> str:
return reject.headers["x-litellm-call-id"]
def _chat_body(model: str, text: str, guardrail: str | None, stream: bool = False) -> dict:
body: dict = {"model": model, "messages": [{"role": "user", "content": text}]}
if stream:
body["stream"] = True
if guardrail is not None:
body["guardrails"] = [guardrail]
return body
def test_cache_hit_post_call_reject_keeps_deployment_labels(gateway: Gateway, tmp_path: Path) -> None:
"""H1: warm then identical post_call-rejected cache hit keeps populated deployment labels."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker) as rig:
text: Final = "cache hit control h1 " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name)
)
assert reject.status_code == 400, reject.text
assert rig.provider.received.qsize() == 1, rig.provider.drain()
_expect_metrics(rig, 1, 0)
_assert_spend(_call_id(reject), rig)
def test_cache_hit_post_call_reject_keeps_deployment_labels_openai_sdk(gateway: Gateway, tmp_path: Path) -> None:
"""H2: same as H1 through the openai AsyncOpenAI client."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker) as rig:
text: Final = "cache hit control h2 " + marker
sdk: Final = openai.AsyncOpenAI(
base_url=str(rig.candidate.client.base_url) + "/v1",
api_key=rig.candidate.key,
http_client=httpx.AsyncClient(trust_env=False, timeout=15),
)
async def run() -> int:
await sdk.chat.completions.create(model=rig.model_name, messages=[{"role": "user", "content": text}])
try:
await sdk.chat.completions.create(
model=rig.model_name,
messages=[{"role": "user", "content": text}],
extra_body={"guardrails": [rig.guardrail_name]},
)
return 200
except openai.BadRequestError:
return 400
assert asyncio.run(run()) == 400
assert rig.provider.received.qsize() == 1, rig.provider.drain()
_expect_metrics(rig, 1, 0)
def test_cache_hit_post_call_reject_streaming(gateway: Gateway, tmp_path: Path) -> None:
"""H3: streamed responses are not cached; the reject call hits upstream again and no failure hook fires."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker) as rig:
text: Final = "cache hit control h3 " + marker
warm: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None, stream=True)
)
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name, stream=True)
)
assert reject.status_code == 200, reject.text
assert rig.provider.received.qsize() == 2, rig.provider.drain()
samples: Final = _metric_samples(rig.candidate, rig.model_name)
assert _populated_failures(samples, rig, "openai") == 0, samples
assert _count(samples, DEPLOYMENT_FAILURE, "") == 0, samples
def test_cache_hit_post_call_reject_keeps_deployment_labels_anthropic(gateway: Gateway, tmp_path: Path) -> None:
"""H4: /v1/messages cache hit reject through the anthropic SDK."""
marker: Final = uuid.uuid4().hex
with _rig(
gateway, tmp_path, marker, upstream_model="anthropic/claude-sonnet-4-5-20250929", api_base_suffix=""
) as rig:
text: Final = "cache hit control h4 " + marker
sdk: Final = anthropic.Anthropic(
base_url=str(rig.candidate.client.base_url),
api_key=rig.candidate.key,
http_client=httpx.Client(trust_env=False, timeout=15),
)
sdk.messages.create(model=rig.model_name, max_tokens=16, messages=[{"role": "user", "content": text}])
raised: bool = False # mutable-ok: a flag set inside the except block cannot be Final
try:
sdk.messages.create(
model=rig.model_name,
max_tokens=16,
messages=[{"role": "user", "content": text}],
extra_body={"guardrails": [rig.guardrail_name]},
)
except anthropic.BadRequestError:
raised = True
assert raised, "cache-hit post_call guardrail did not reject /v1/messages"
assert rig.provider.received.qsize() == 1, rig.provider.drain()
_expect_metrics(rig, 1, 0, api_provider="anthropic", pf_id="None")
def test_cache_hit_post_call_reject_keeps_deployment_labels_responses(gateway: Gateway, tmp_path: Path) -> None:
"""H5: /v1/responses cache hit reject."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker) as rig:
text: Final = "cache hit control h5 " + marker
warm: Final = rig.candidate.request("POST", "/v1/responses", {"model": rig.model_name, "input": text})
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request(
"POST",
"/v1/responses",
{"model": rig.model_name, "input": text, "guardrails": [rig.guardrail_name]},
)
assert reject.status_code == 400, reject.text
assert rig.provider.received.qsize() == 1, rig.provider.drain()
_expect_metrics(rig, 1, 0, pf_id="None")
_assert_spend(_call_id(reject), rig)
def test_cache_hit_post_call_reject_embeddings(gateway: Gateway, tmp_path: Path) -> None:
"""H6: post_call guardrails do not run on embeddings; the cached response returns 200 unguarded."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker) as rig:
text: Final = "cache hit control h6 " + marker
warm: Final = rig.candidate.request("POST", "/v1/embeddings", {"model": rig.model_name, "input": text})
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request(
"POST",
"/v1/embeddings",
{"model": rig.model_name, "input": text, "guardrails": [rig.guardrail_name]},
)
assert reject.status_code == 200, reject.text
assert rig.provider.received.qsize() == 1, rig.provider.drain()
samples: Final = _metric_samples(rig.candidate, rig.model_name)
assert _populated_failures(samples, rig, "openai") == 0, samples
assert _count(samples, DEPLOYMENT_FAILURE, "") == 0, samples
def test_cache_hit_during_call_reject_keeps_deployment_labels(gateway: Gateway, tmp_path: Path) -> None:
"""H7: during_call guardrail reject on a cache hit."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker, mode="during_call") as rig:
text: Final = "cache hit control h7 " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name)
)
assert reject.status_code == 400, reject.text
_expect_metrics(rig, 1, 0)
def test_pre_call_reject_on_cache_hit_stays_blank(gateway: Gateway, tmp_path: Path) -> None:
"""C1: pre_call reject never reaches the deployment; labels stay blank on both legs."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker, mode="pre_call") as rig:
text: Final = "cache hit control c1 " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name)
)
assert reject.status_code == 400, reject.text
_expect_metrics(rig, 0, 1, pf_populated=1, pf_blank=0)
def test_post_call_reject_without_cache_keeps_deployment_labels(gateway: Gateway, tmp_path: Path) -> None:
"""C2: a real provider call rejected post_call keeps populated labels on both legs."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker) as rig:
text: Final = "non cache control c2 " + marker
reject: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name)
)
assert reject.status_code == 400, reject.text
assert rig.provider.received.qsize() == 1, rig.provider.drain()
_expect_metrics(rig, 1, 0)
_assert_spend(_call_id(reject), rig)
def test_cache_hit_post_call_reject_default_on(gateway: Gateway, tmp_path: Path) -> None:
"""C3: default_on post_call guardrail rejects the cached response (sink passes the warm call)."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker, sink=_first_call_pass_sink(), default_on=True) as rig:
text: Final = "cache hit control c3 " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert reject.status_code == 400, reject.text
assert rig.provider.received.qsize() == 1, rig.provider.drain()
_expect_metrics(rig, 1, 0)
def test_cache_hit_post_call_reject_key_metadata_guardrails(gateway: Gateway, tmp_path: Path) -> None:
"""C4: guardrail attached via key metadata guardrails on a cache hit."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker, sink=_first_call_pass_sink()) as rig:
key: Final = rig.candidate.post("/key/generate", {"metadata": {"guardrails": [rig.guardrail_name]}})["key"]
text: Final = "cache hit control c4 " + marker
warm: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None), key=key
)
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None), key=key
)
assert reject.status_code == 400, reject.text
assert rig.provider.received.qsize() == 1, rig.provider.drain()
assert rig.policy.received.qsize() == 2
_expect_metrics(rig, 1, 0)
def test_cache_hit_post_call_reject_local_cache(gateway: Gateway, tmp_path: Path) -> None:
"""C5: same cache-hit reject with cache_params type local."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker, local_cache=True) as rig:
text: Final = "cache hit control c5 " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name)
)
assert reject.status_code == 400, reject.text
assert rig.provider.received.qsize() == 1, rig.provider.drain()
_expect_metrics(rig, 1, 0)
@pytest.mark.parametrize("status", (500, 403))
def test_cache_hit_post_call_guardrail_outage_keeps_deployment_labels(
gateway: Gateway, tmp_path: Path, status: int
) -> None:
"""S1/S2: guardrail sink answers 500/403 on the cache-hit call; failure hook still counts as dispatched."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker, sink=_failing_sink(status)) as rig:
text: Final = "cache hit control s " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name)
)
assert reject.status_code >= 400, reject.text
assert rig.provider.received.qsize() == 1, rig.provider.drain()
_expect_metrics(rig, 0, 0, pf_populated=1, pf_blank=0)
def test_two_identical_cache_hit_rejects_increment_populated_series(gateway: Gateway, tmp_path: Path) -> None:
"""E1: two identical cache-hit rejects count +2 on the populated series, two spend rows."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker) as rig:
text: Final = "cache hit control e1 " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
rejects: Final = tuple(
rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name))
for _ in range(2)
)
assert all(response.status_code == 400 for response in rejects), [r.text for r in rejects]
assert rig.provider.received.qsize() == 1, rig.provider.drain()
_expect_metrics(rig, 2, 0)
def test_two_identical_cache_hit_rejects_write_matching_spend_rows(gateway: Gateway, tmp_path: Path) -> None:
"""E1b: both cache-hit rejects land a failure spend row."""
pytest.skip("BUG: roughly one in four back-to-back cache-hit rejects never lands its LiteLLM_SpendLogs row")
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker) as rig:
text: Final = "cache hit control e1b " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
rejects: Final = tuple(
rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name))
for _ in range(2)
)
assert all(response.status_code == 400 for response in rejects), [r.text for r in rejects]
for response in rejects:
_assert_spend(_call_id(response), rig)
def test_cache_hit_reject_after_ttl_expiry_is_a_miss(gateway: Gateway, tmp_path: Path) -> None:
"""E2: cache_params ttl=1; post-expiry the same body misses, hits upstream again, labels populated."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker, ttl=1) as rig:
text: Final = "cache hit control e2 " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
assert rig.provider.received.qsize() == 1
rejects: list[int] = [] # mutable-ok: the poll helper must remember how many rejects it issued
def expired_miss() -> int:
rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name))
rejects.append(1)
return rig.provider.received.qsize()
eventually(lambda: expired_miss() == 2, bool, seconds=70)
_expect_metrics(rig, len(rejects), 0)
def test_cache_hit_reject_metrics_aggregate_across_workers(gateway: Gateway, tmp_path: Path) -> None:
"""E3: workers=2, 8 cache-hit rejects, aggregated /metrics shows +8 on the populated series."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker, workers=2) as rig:
text: Final = "cache hit control e3 " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
rejects: Final = tuple(
rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name))
for _ in range(8)
)
assert all(response.status_code == 400 for response in rejects), [r.text for r in rejects]
_expect_metrics(rig, 8, 0)
def test_cache_hit_reject_deployment_metric_set_diff(gateway: Gateway, tmp_path: Path) -> None:
"""E4: exact expected label sets on litellm_deployment_* and litellm_proxy_failed_requests_metric."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker) as rig:
text: Final = "cache hit control e4 " + marker
warm: Final = rig.candidate.request("POST", "/v1/chat/completions", _chat_body(rig.model_name, text, None))
assert warm.status_code == 200, warm.text
reject: Final = rig.candidate.request(
"POST", "/v1/chat/completions", _chat_body(rig.model_name, text, rig.guardrail_name)
)
assert reject.status_code == 400, reject.text
samples: Final = _expect_metrics(rig, 1, 0)
blank: Final = tuple(sample for sample in samples if sample.labels.get("model_id") == "")
assert blank == (), blank
states: Final = tuple(
sample.value
for sample in samples
if sample.name == DEPLOYMENT_STATE
and sample.labels.get("model_id") == rig.deployment_id
and sample.labels.get("api_base") == ""
)
assert states == (1.0,), states

View file

@ -0,0 +1,348 @@
import json
import signal
import socket
import subprocess
import threading
import uuid
from collections.abc import Callable, Generator
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from pathlib import Path
from typing import Final
import httpx
import psutil
from integration._support.client import Gateway, eventually, object_value, string_value
from integration._support.process import owned_proxy_process
from integration._support.wire import Reply, Request, wire_server
from prometheus_client.parser import text_string_to_metric_families
from test_cache_hit_guardrail_metrics import (
DEPLOYMENT_FAILURE,
GUARDRAIL_PATH,
PROXY_FAILED,
Rig,
_blocking_sink,
_chat_body,
_guardrail_config,
_provider,
_rig,
)
BURST: Final = 10
@contextmanager
def _redis(port: int) -> Generator[subprocess.Popen[bytes], None, None]:
process: Final = subprocess.Popen(["redis-server", "--port", str(port), "--save", ""], stdout=subprocess.DEVNULL)
try:
yield process
finally:
process.kill()
process.wait(timeout=10)
def _free_port() -> int:
with socket.socket() as reserve:
reserve.bind(("127.0.0.1", 0))
return reserve.getsockname()[1]
def _deployment_id(candidate: Gateway, model_name: str) -> str:
entries: Final = candidate.get("/model/info")["data"]
assert isinstance(entries, list)
entry: Final = next(item for item in entries if object_value(item)["model_name"] == model_name)
return string_value(object_value(object_value(entry)["model_info"])["id"])
def _stall_sink(stall: threading.Event, release: threading.Event) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert request.target == GUARDRAIL_PATH, request.target
if stall.is_set():
release.wait(timeout=60)
return Reply(body=json.dumps({"action": "BLOCKED", "blocked_reason": "synthetic block"}).encode())
return respond
def _samples(candidate: Gateway, model_names: tuple[str, ...]) -> tuple:
response: Final = candidate.client.request(
"GET", "/metrics", headers={"Authorization": f"Bearer {candidate.key}"}, follow_redirects=True
)
assert response.status_code == 200, f"GET /metrics: {response.status_code}"
return tuple(
sample
for family in text_string_to_metric_families(response.text)
for sample in family.samples
if sample.labels.get("requested_model") in model_names
)
def _populated(samples: tuple, deployment_id: str) -> float:
return float(
sum(
sample.value
for sample in samples
if sample.name == DEPLOYMENT_FAILURE and sample.labels.get("model_id") == deployment_id
)
)
def _blank(samples: tuple) -> float:
return float(
sum(
sample.value
for sample in samples
if sample.name == DEPLOYMENT_FAILURE and sample.labels.get("model_id") == ""
)
)
def _proxy_failed(samples: tuple) -> float:
return float(sum(sample.value for sample in samples if sample.name == PROXY_FAILED))
def _burst_bodies(rig: Rig, marker: str, anthropic_name: str | None) -> tuple[tuple[str, dict], ...]:
chat: Final = tuple(
("/v1/chat/completions", _chat_body(rig.model_name, f"burst {marker} {index}", rig.guardrail_name))
for index in range(BURST)
)
responses: Final = tuple(
(
"/v1/responses",
{"model": rig.model_name, "input": f"burst {marker} r{index}", "guardrails": [rig.guardrail_name]},
)
for index in range(BURST)
)
messages: Final = (
tuple(
(
"/v1/messages",
{
"model": anthropic_name,
"max_tokens": 16,
"messages": [{"role": "user", "content": f"burst {marker} m{index}"}],
"guardrails": [rig.guardrail_name],
},
)
for index in range(BURST)
)
if anthropic_name is not None
else ()
)
return chat + responses + messages
def _warm(rig: Rig, bodies: tuple[tuple[str, dict], ...]) -> None:
for path, body in bodies:
warmed: Final = dict(body)
warmed.pop("guardrails", None)
response: Final = rig.candidate.request("POST", path, warmed)
assert response.status_code == 200, f"warm {path}: {response.status_code} {response.text}"
def _fire(rig: Rig, bodies: tuple[tuple[str, dict], ...]) -> tuple[tuple[int, str | None], ...]:
def call(item: tuple[str, dict]) -> tuple[int, str | None]:
path, body = item
try:
response: Final = rig.candidate.request("POST", path, body)
return response.status_code, response.headers.get("x-litellm-call-id")
except httpx.HTTPError:
return -1, None
with ThreadPoolExecutor(max_workers=8) as pool:
return tuple(pool.map(call, bodies))
def _expect_counted_within(
rig: Rig, model_names: tuple[str, ...], deployment_ids: tuple[str, ...], low: int, high: int
) -> None:
def converged() -> tuple:
samples: Final = _samples(rig.candidate, model_names)
populated: Final = sum(_populated(samples, deployment) for deployment in deployment_ids)
if low <= populated <= high and _blank(samples) == 0:
return samples
return ()
eventually(converged, bool, seconds=70)
def _expect_exactly_once(rig: Rig, model_names: tuple[str, ...], deployment_ids: tuple[str, ...], four_xx: int) -> None:
_expect_counted_within(rig, model_names, deployment_ids, four_xx, four_xx)
def test_burst_cache_hit_rejects_count_exactly_once(gateway: Gateway, tmp_path: Path) -> None:
"""X0: 30 mixed-endpoint cache-hit rejects across two deployments, each counted once."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker) as rig:
anthropic_name: Final = rig.scenario.model(
model="anthropic/claude-sonnet-4-5-20250929", api_base=rig.provider.url, api_key="synthetic-provider-key"
)
anthropic_id: Final = _deployment_id(rig.candidate, anthropic_name)
bodies: Final = _burst_bodies(rig, marker, anthropic_name)
_warm(rig, bodies)
outcomes: Final = _fire(rig, bodies)
rejected: Final = sum(1 for status, _ in outcomes if status >= 400)
assert all(status == 400 for status, _ in outcomes), outcomes
_expect_exactly_once(rig, (rig.model_name, anthropic_name), (rig.deployment_id, anthropic_id), rejected)
def test_stalled_guardrail_sink_recovers_and_counts(gateway: Gateway, tmp_path: Path) -> None:
"""X1: guardrail sink stalls mid-burst; requests fail exactly once, then recovery counts again."""
marker: Final = uuid.uuid4().hex
stall: Final = threading.Event()
release: Final = threading.Event()
with _rig(gateway, tmp_path, marker, sink=_stall_sink(stall, release)) as rig:
bodies: Final = _burst_bodies(rig, marker, None)
_warm(rig, bodies)
stall.set()
with ThreadPoolExecutor(max_workers=8) as pool:
futures: Final = tuple(
pool.submit(lambda b: rig.candidate.request("POST", b[0], b[1]), body) for body in bodies
)
eventually(lambda: rig.policy.received.qsize() >= 5, bool, seconds=30)
release.set()
outcomes: Final = tuple(
(future.result().status_code, future.result().headers.get("x-litellm-call-id")) for future in futures
)
assert all(status >= 400 for status, _ in outcomes), outcomes
blocked: Final = sum(1 for status, _ in outcomes if status == 400)
outages: Final = sum(1 for status, _ in outcomes if status >= 500)
assert blocked + outages == len(bodies), outcomes
samples: Final = eventually(
lambda: _samples(rig.candidate, (rig.model_name,)),
lambda observed: _proxy_failed(observed) == blocked + outages,
seconds=70,
)
assert _proxy_failed(samples) == blocked + outages, (samples, outcomes)
follow_up: Final = rig.candidate.request(
"POST",
"/v1/chat/completions",
_chat_body(rig.model_name, "post stall unrelated " + marker, None),
)
assert follow_up.status_code == 200, follow_up.text
_expect_exactly_once(rig, (rig.model_name,), (rig.deployment_id,), blocked)
def test_redis_outage_keeps_serving_in_memory_hits(gateway: Gateway, tmp_path: Path) -> None:
"""X2: the redis cache keeps an in-memory shadow, so a redis kill does not stop cache-hit rejects."""
marker: Final = uuid.uuid4().hex
port: Final = _free_port()
with _redis(port) as redis_one:
with _rig(gateway, tmp_path, marker, env={"REDIS_HOST": "127.0.0.1", "REDIS_PORT": str(port)}) as rig:
bodies: Final = _burst_bodies(rig, marker, None)[:BURST]
_warm(rig, bodies)
reject: Final = rig.candidate.request("POST", *bodies[0])
assert reject.status_code == 400, reject.text
warmed_hits: Final = rig.provider.received.qsize()
redis_one.kill()
redis_one.wait(timeout=10)
outcomes: Final = _fire(rig, bodies[1:])
assert all(status == 400 for status, _ in outcomes), outcomes
assert rig.provider.received.qsize() == warmed_hits, (
"redis outage reached the provider",
warmed_hits,
rig.provider.received.qsize(),
)
with _redis(port):
recovered: Final = rig.candidate.request(
"POST",
"/v1/chat/completions",
_chat_body(rig.model_name, "x2 rehit " + marker, rig.guardrail_name),
)
assert recovered.status_code == 400, recovered.text
_expect_exactly_once(rig, (rig.model_name,), (rig.deployment_id,), 1 + len(bodies))
def test_worker_kill_mid_burst_keeps_counting(gateway: Gateway, tmp_path: Path) -> None:
"""X3: workers=2, SIGKILL one uvicorn child mid-burst; survivors keep rejecting; the count is answered plus at most the in-flight requests the killed worker had already counted."""
marker: Final = uuid.uuid4().hex
with _rig(gateway, tmp_path, marker, workers=2) as rig:
bodies: Final = _burst_bodies(rig, marker, None)
_warm(rig, bodies)
children: Final = psutil.Process(rig.process.pid).children(recursive=True)
assert children, "no uvicorn worker children found"
with ThreadPoolExecutor(max_workers=8) as pool:
futures: Final = tuple(
pool.submit(lambda b: rig.candidate.request("POST", b[0], b[1]), body) for body in bodies
)
eventually(lambda: rig.policy.received.qsize() >= 3, bool, seconds=30)
children[0].send_signal(signal.SIGKILL)
statuses: list[int] = [] # mutable-ok: collect per-request outcomes from concurrent futures
for future in futures:
try:
statuses.append(future.result().status_code)
except httpx.HTTPError:
statuses.append(-1)
answered: Final = sum(1 for status in statuses if status >= 0)
transport_lost: Final = sum(1 for status in statuses if status == -1)
assert all(status == 400 for status in statuses if status >= 0), (
statuses,
transport_lost,
)
_expect_counted_within(rig, (rig.model_name,), (rig.deployment_id,), answered, answered + transport_lost)
def test_proxy_restart_mid_burst_keeps_counting(gateway: Gateway, tmp_path: Path) -> None:
"""X4: restart the owned proxy between the two halves; pre-restart count asserted, then recounted."""
marker: Final = uuid.uuid4().hex
prom_dir: Final = tmp_path / "prom"
prom_dir.mkdir()
with wire_server(_blocking_sink) as policy, wire_server(_provider(marker)) as provider:
config: Final = _guardrail_config(tmp_path, "guardrail-" + marker, policy.url)
bodies: Final = tuple(
(
"/v1/chat/completions",
_chat_body("pending-model", f"burst {marker} {index}", "guardrail-" + marker),
)
for index in range(BURST)
)
with owned_proxy_process(
gateway, tmp_path, {"PROMETHEUS_MULTIPROC_DIR": str(prom_dir)}, config=config
) as owned_one:
model: Final = "restart-" + marker
owned_one.gateway.post(
"/model/new",
{
"model_name": model,
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_base": provider.url + "/v1",
"api_key": "synthetic-provider-key",
},
},
)
deployment: Final = _deployment_id(owned_one.gateway, model)
named: Final = tuple((path, {**body, "model": model}) for path, body in bodies)
first_half, second_half = named[: BURST // 2], named[BURST // 2 :]
for path, body in named:
warmed: Final = dict(body)
warmed.pop("guardrails", None)
assert owned_one.gateway.request("POST", path, warmed).status_code == 200
outcomes_one: Final = tuple(owned_one.gateway.request("POST", path, body) for path, body in first_half)
assert all(response.status_code == 400 for response in outcomes_one), [r.text for r in outcomes_one]
pre: Final = eventually(
lambda: (
_populated(_samples(owned_one.gateway, (model,)), deployment),
_blank(_samples(owned_one.gateway, (model,))),
),
lambda observed: observed[0] == len(first_half) and observed[1] == 0,
seconds=70,
)
with owned_proxy_process(
gateway, tmp_path, {"PROMETHEUS_MULTIPROC_DIR": str(prom_dir)}, config=config
) as owned_two:
outcomes_two: Final = tuple(owned_two.gateway.request("POST", path, body) for path, body in second_half)
assert all(response.status_code == 400 for response in outcomes_two), (
pre,
[(r.status_code, r.text[:200]) for r in outcomes_two],
)
post: Final = eventually(
lambda: (
_populated(_samples(owned_two.gateway, (model,)), deployment),
_blank(_samples(owned_two.gateway, (model,))),
),
lambda observed: observed[0] == len(named) and observed[1] == 0,
seconds=70,
)
assert post[0] == len(named), (pre, post, outcomes_two)
owned_two.gateway.post("/model/delete", {"id": deployment})

View file

@ -0,0 +1,150 @@
import asyncio
import json
import re
import signal
from pathlib import Path
from typing import Final
import httpx
import psutil
import pytest
import yaml
from integration._support.client import Gateway, eventually, object_value
from integration._support.database import read_rows
from integration._support.process import owned_proxy_process
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue
_GENERATE_CONTENT: Final[dict[str, JsonValue]] = {"contents": [{"role": "user", "parts": [{"text": "hi"}]}]}
_NOT_FOUND_BODY: Final = json.dumps(
{
"error": {
"code": 404,
"message": "models/nope-9 is not found for this scripted upstream",
"status": "NOT_FOUND",
}
}
).encode()
_INTERNAL_BODY: Final = (
'{"error":{"code":500,"message":"' + "chunked upstream failure body " * 200 + '","status":"INTERNAL"}}'
).encode()
_OK_CHUNKS: Final = tuple(f"data: ok-{index}\n\n".encode() for index in range(3))
_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
def _chaos_reply(request: Request) -> Reply:
if "streamGenerateContent" in request.target:
return Reply(status=500, chunks=tuple(_INTERNAL_BODY[i : i + 512] for i in range(0, len(_INTERNAL_BODY), 512)))
if "healthy-model" in request.target:
return Reply(status=200, chunks=_OK_CHUNKS, content_type="text/event-stream")
return Reply(status=404, body=_NOT_FOUND_BODY)
def _error_information(call_id: str) -> dict[str, JsonValue]:
rows: Final = eventually(
lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (call_id,)),
lambda values: len(values) == 1,
seconds=70,
)
metadata: Final = rows[0]["metadata"]
parsed: Final = json.loads(metadata) if isinstance(metadata, str) else object_value(metadata)
return object_value(parsed["error_information"])
def _single_spend_row(call_id: str) -> None:
rows: Final = eventually(
lambda: read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (call_id,)),
lambda values: len(values) == 1,
seconds=70,
)
assert len(rows) == 1, call_id
async def _fire_burst(
base_url: str, key: str, count: int, *, tolerate_transport_errors: bool = False
) -> tuple[httpx.Response, ...]:
async def one(client: httpx.AsyncClient, index: int) -> httpx.Response:
if index % 3 == 0:
path: Final = "/gemini/v1beta/models/nope-9:generateContent"
elif index % 3 == 1:
path = "/gemini/v1beta/models/nope-9:streamGenerateContent?alt=sse"
else:
path = "/gemini/v1beta/models/healthy-model:streamGenerateContent?alt=sse"
return await client.post(
path,
json=_GENERATE_CONTENT,
headers={"Authorization": f"Bearer {key}", "x-goog-api-key": key},
)
async with httpx.AsyncClient(base_url=base_url, timeout=30, trust_env=False) as client:
results: Final = await asyncio.gather(
*(one(client, index) for index in range(count)), return_exceptions=tolerate_transport_errors
)
for result in results:
assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result)
return tuple(result for result in results if isinstance(result, httpx.Response))
async def test_passthrough_upstream_outage_mid_burst_still_logs_errors_once(gateway: Gateway, tmp_path: Path) -> None:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
path: Final = tmp_path / "chaos-outage.yaml"
with wire_server(_chaos_reply) as wire:
port: Final = int(wire.url.rsplit(":", 1)[1])
config["environment_variables"] = {"GEMINI_API_BASE": wire.url, "GEMINI_API_KEY": "scripted"}
path.write_text(yaml.safe_dump(config))
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
candidate: Final = owned.gateway
burst: Final = asyncio.create_task(_fire_burst(str(candidate.client.base_url), candidate.key, 30))
await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= 10, 30)
with wire_server(_chaos_reply, port=port):
responses: Final = await burst
assert len(responses) == 30
for response in responses:
assert response.status_code in (200, 404, 500, 502), response.status_code
assert "x-litellm-call-id" in response.headers, response.status_code
assert len(_STARTED_WORKER.findall(owned.log.read_text())) >= 2
for response in responses:
_single_spend_row(response.headers["x-litellm-call-id"])
if response.status_code == 404:
error_information: Final = _error_information(response.headers["x-litellm-call-id"])
assert "not found for this scripted upstream" in str(error_information["error_message"]), response.text
elif response.status_code == 500:
assert "chunked upstream failure body" in str(
_error_information(response.headers["x-litellm-call-id"])["error_message"]
), response.text
async def test_passthrough_worker_sigkill_leaves_sibling_serving_and_logging(gateway: Gateway, tmp_path: Path) -> None:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
path: Final = tmp_path / "chaos-kill.yaml"
with wire_server(_chaos_reply) as wire:
config["environment_variables"] = {"GEMINI_API_BASE": wire.url, "GEMINI_API_KEY": "scripted"}
path.write_text(yaml.safe_dump(config))
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
candidate: Final = owned.gateway
workers: Final = eventually(
lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())),
lambda pids: len(pids) == 2,
seconds=30,
)
burst: Final = asyncio.create_task(
_fire_burst(str(candidate.client.base_url), candidate.key, 20, tolerate_transport_errors=True)
)
await asyncio.to_thread(eventually, lambda: wire.received.qsize(), lambda size: size >= 5, 30)
psutil.Process(workers[0]).send_signal(signal.SIGKILL)
responses: Final = await burst
for response in responses:
assert response.status_code in (200, 404, 500, 502), response.status_code
follow_up: Final = candidate.request(
"POST",
"/gemini/v1beta/models/nope-9:generateContent",
_GENERATE_CONTENT,
headers={"x-goog-api-key": candidate.key},
)
assert follow_up.status_code == 404, follow_up.text
assert follow_up.json() == json.loads(_NOT_FOUND_BODY), follow_up.text
for response in responses:
if "x-litellm-call-id" in response.headers:
_single_spend_row(response.headers["x-litellm-call-id"])
error_information: Final = _error_information(follow_up.headers["x-litellm-call-id"])
assert "not found for this scripted upstream" in str(error_information["error_message"]), follow_up.text

View file

@ -0,0 +1,600 @@
import gzip
import json
from hashlib import sha256
from pathlib import Path
from typing import Final
import httpx
import pytest
import yaml
from integration._support.client import Gateway, eventually, object_value
from integration._support.database import read_rows
from integration._support.process import owned_proxy_process
from integration._support.wire import Reply, Request, wire_server
from openai import AsyncOpenAI, NotFoundError, OpenAI
from pydantic import JsonValue
_UPSTREAM_ERROR: Final[dict[str, JsonValue]] = {
"error": {
"code": 404,
"message": "Publisher Model `publishers/anthropic/models/claude-nope-9` was not found or your project does not have access to it. Please ensure you are using a valid model version.",
"status": "NOT_FOUND",
}
}
def test_gemini_passthrough_upstream_error_body_reaches_proxy_log_and_spend_row(
gateway: Gateway, tmp_path: Path
) -> None:
def respond(request: Request) -> Reply:
return Reply(status=404, body=json.dumps(_UPSTREAM_ERROR).encode())
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
path: Final = tmp_path / "gemini-passthrough.yaml"
with wire_server(respond) as wire:
config["environment_variables"] = {"GEMINI_API_BASE": wire.url, "GEMINI_API_KEY": "scripted"}
path.write_text(yaml.safe_dump(config))
with owned_proxy_process(gateway, tmp_path, {}, config=path) as owned:
candidate: Final = owned.gateway
response: Final = candidate.request(
"POST",
"/gemini/v1beta/models/claude-nope-9:generateContent",
{"contents": [{"role": "user", "parts": [{"text": "hi"}]}]},
headers={"x-goog-api-key": candidate.key},
)
assert response.status_code == 404, response.text
assert response.json() == _UPSTREAM_ERROR, response.text
try:
eventually(
lambda: owned.log.read_text(),
lambda text: "was not found or your project" in text,
seconds=30,
)
except AssertionError:
pytest.fail(
f"upstream 404 body never reached the proxy log after {response.status_code} passthrough; "
f"log tail: {owned.log.read_text()[-2000:]}"
)
rows: Final = eventually(
lambda: read_rows(
'SELECT metadata FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
(response.headers["x-litellm-call-id"],),
),
lambda values: len(values) == 1,
seconds=70,
)
metadata: Final = rows[0]["metadata"]
parsed: Final = json.loads(metadata) if isinstance(metadata, str) else object_value(metadata)
error_information: Final = object_value(parsed["error_information"])
assert "was not found or your project" in str(error_information["error_message"]), response.text
assert error_information["error_code"] == "404", response.text
_GEMINI_MODEL_PATH: Final = "/gemini/v1beta/models/claude-nope-9:generateContent"
_GEMINI_STREAM_PATH: Final = "/gemini/v1beta/models/claude-nope-9:streamGenerateContent"
_GENERATE_CONTENT: Final[dict[str, JsonValue]] = {"contents": [{"role": "user", "parts": [{"text": "hi"}]}]}
_UPSTREAM_500_BODY: Final = (
'{"error":{"code":500,"message":"' + "chunked upstream failure body " * 200 + '","status":"INTERNAL"}}'
).encode()
def _gemini_config(path: Path, wire_url: str) -> None:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["environment_variables"] = {"GEMINI_API_BASE": wire_url, "GEMINI_API_KEY": "scripted"}
path.write_text(yaml.safe_dump(config))
def _gemini_headers(candidate: Gateway) -> dict[str, str]:
return {"Authorization": f"Bearer {candidate.key}", "x-goog-api-key": candidate.key}
def _spend_error_information(call_id: str) -> dict[str, JsonValue]:
rows: Final = eventually(
lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (call_id,)),
lambda values: len(values) == 1,
seconds=70,
)
metadata: Final = rows[0]["metadata"]
parsed: Final = json.loads(metadata) if isinstance(metadata, str) else object_value(metadata)
return object_value(parsed["error_information"])
def _spend_status(call_id: str) -> str:
rows: Final = eventually(
lambda: read_rows('SELECT status FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (call_id,)),
lambda values: len(values) == 1,
seconds=70,
)
return str(rows[0]["status"])
def _upstream_warning(log: Path, needle: str = "pass_through_endpoint: upstream") -> str:
text: Final = eventually(lambda: log.read_text(), lambda content: needle in content, seconds=30)
return next(line for line in text.splitlines() if needle in line)
def _upstream_warnings(log: Path, needle: str = "pass_through_endpoint: upstream") -> tuple[str, ...]:
return tuple(line for line in log.read_text().splitlines() if needle in line)
async def test_gemini_passthrough_async_client_404_body_reaches_proxy_log_and_spend_row(
gateway: Gateway, tmp_path: Path
) -> None:
def respond(request: Request) -> Reply:
return Reply(status=404, body=json.dumps(_UPSTREAM_ERROR).encode())
path: Final = tmp_path / "gemini-async.yaml"
with wire_server(respond) as wire:
_gemini_config(path, wire.url)
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
candidate: Final = owned.gateway
async with httpx.AsyncClient(
base_url=str(candidate.client.base_url), timeout=15, trust_env=False
) as async_client:
response: Final = await async_client.post(
_GEMINI_MODEL_PATH, json=_GENERATE_CONTENT, headers=_gemini_headers(candidate)
)
assert response.status_code == 404, response.text
assert response.json() == _UPSTREAM_ERROR, response.text
warning: Final = _upstream_warning(owned.log)
assert "was not found or your project" in warning, warning
error_information: Final = _spend_error_information(response.headers["x-litellm-call-id"])
assert "was not found or your project" in str(error_information["error_message"]), response.text
assert error_information["error_code"] == "404", response.text
def test_gemini_passthrough_streaming_500_relays_full_body_and_logs_bounded_preview(
gateway: Gateway, tmp_path: Path
) -> None:
body: Final = _UPSTREAM_500_BODY
assert len(body) == 6055
chunks: Final = tuple(body[index * 512 : (index + 1) * 512] for index in range(11)) + (body[5632:],)
def respond(request: Request) -> Reply:
return Reply(status=500, chunks=chunks)
path: Final = tmp_path / "gemini-stream-500.yaml"
with wire_server(respond) as wire:
_gemini_config(path, wire.url)
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
candidate: Final = owned.gateway
with candidate.client.stream(
"POST",
_GEMINI_STREAM_PATH,
params={"alt": "sse"},
json=_GENERATE_CONTENT,
headers=_gemini_headers(candidate),
) as response:
assert response.status_code == 500, response.text
streamed: Final = response.read()
assert streamed == body
warning: Final = _upstream_warning(owned.log)
assert warning.endswith("... (truncated at 4096 chars)"), warning
error_information: Final = _spend_error_information(response.headers["x-litellm-call-id"])
error_message: Final = str(error_information["error_message"])
assert error_message.endswith("... (truncated at 4096 chars)"), error_message
assert error_information["error_code"] == "500", error_message
def test_gemini_passthrough_success_logs_nothing_and_spend_row_is_success(gateway: Gateway, tmp_path: Path) -> None:
upstream_ok: Final = {
"candidates": [{"content": {"parts": [{"text": "hello"}], "role": "model"}}],
"usageMetadata": {"promptTokenCount": 3, "candidatesTokenCount": 2, "totalTokenCount": 5},
}
def respond(request: Request) -> Reply:
return Reply(status=200, body=json.dumps(upstream_ok).encode())
path: Final = tmp_path / "gemini-200.yaml"
with wire_server(respond) as wire:
_gemini_config(path, wire.url)
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
candidate: Final = owned.gateway
response: Final = candidate.request(
"POST", _GEMINI_MODEL_PATH, _GENERATE_CONTENT, headers={"x-goog-api-key": candidate.key}
)
assert response.status_code == 200, response.text
assert response.json() == upstream_ok, response.text
assert _spend_status(response.headers["x-litellm-call-id"]) == "success"
assert not _upstream_warnings(owned.log), owned.log.read_text()[-2000:]
def test_gemini_passthrough_streaming_200_relays_every_chunk(gateway: Gateway, tmp_path: Path) -> None:
chunks: Final = tuple(f"data: chunk-{index}\n\n".encode() for index in range(5))
def respond(request: Request) -> Reply:
return Reply(status=200, chunks=chunks, content_type="text/event-stream")
path: Final = tmp_path / "gemini-stream-200.yaml"
with wire_server(respond) as wire:
_gemini_config(path, wire.url)
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
candidate: Final = owned.gateway
with candidate.client.stream(
"POST",
_GEMINI_STREAM_PATH,
params={"alt": "sse"},
json=_GENERATE_CONTENT,
headers=_gemini_headers(candidate),
) as response:
assert response.status_code == 200
streamed: Final = response.read()
assert streamed == b"".join(chunks)
assert not _upstream_warnings(owned.log), owned.log.read_text()[-2000:]
def test_config_pass_through_route_logs_body_and_strips_query(gateway: Gateway, tmp_path: Path) -> None:
upstream_error: Final = {"error": {"message": "max budget reached for this deployment"}}
def respond(request: Request) -> Reply:
return Reply(status=403, body=json.dumps(upstream_error).encode())
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
path: Final = tmp_path / "config-route.yaml"
with wire_server(respond) as wire:
config["general_settings"]["pass_through_endpoints"] = [
{
"path": "/audit-pt",
"target": f"{wire.url}/upstream?trace=secret-q",
"include_subpath": True,
"headers": {"Authorization": "Bearer scripted"},
}
]
path.write_text(yaml.safe_dump(config))
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
candidate: Final = owned.gateway
response: Final = candidate.request("POST", "/audit-pt", _GENERATE_CONTENT)
assert response.status_code == 403, response.text
assert response.json() == upstream_error, response.text
warning: Final = _upstream_warning(owned.log)
assert "max budget reached for this deployment" in warning, warning
assert "?" not in warning and "secret-q" not in warning, warning
error_information: Final = _spend_error_information(response.headers["x-litellm-call-id"])
assert error_information["normalized_error"] == "500_UPSTREAM_PASSTHROUGH", response.text
assert "max budget reached for this deployment" in str(error_information["error_message"]), response.text
_OPENAI_UPSTREAM_404: Final[dict[str, JsonValue]] = {
"error": {"message": "The model `nope-9` does not exist", "type": "invalid_request_error"}
}
def _openai_config(path: Path, wire_url: str) -> None:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["environment_variables"] = {"OPENAI_API_BASE": wire_url, "OPENAI_API_KEY": "scripted"}
path.write_text(yaml.safe_dump(config))
def test_openai_passthrough_sdk_error_body_reaches_proxy_log_and_spend_row(gateway: Gateway, tmp_path: Path) -> None:
def respond(request: Request) -> Reply:
return Reply(status=404, body=json.dumps(_OPENAI_UPSTREAM_404).encode())
path: Final = tmp_path / "openai-404.yaml"
with wire_server(respond) as wire:
_openai_config(path, wire.url)
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
candidate: Final = owned.gateway
with OpenAI(
api_key=candidate.key,
base_url=f"{str(candidate.client.base_url).rstrip('/')}/openai",
max_retries=0,
http_client=httpx.Client(timeout=15, trust_env=False),
) as sdk:
with pytest.raises(NotFoundError) as raised:
sdk.chat.completions.create(model="nope-9", messages=[{"role": "user", "content": "hi"}])
assert "does not exist" in str(raised.value), raised.value
warning: Final = _upstream_warning(owned.log)
assert "does not exist" in warning, warning
error_information: Final = _spend_error_information(raised.value.response.headers["x-litellm-call-id"])
assert "does not exist" in str(error_information["error_message"])
async def test_openai_passthrough_async_sdk_error_body_reaches_proxy_log_and_spend_row(
gateway: Gateway, tmp_path: Path
) -> None:
def respond(request: Request) -> Reply:
return Reply(status=404, body=json.dumps(_OPENAI_UPSTREAM_404).encode())
path: Final = tmp_path / "openai-async-404.yaml"
with wire_server(respond) as wire:
_openai_config(path, wire.url)
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
candidate: Final = owned.gateway
async with AsyncOpenAI(
api_key=candidate.key,
base_url=f"{str(candidate.client.base_url).rstrip('/')}/openai",
max_retries=0,
http_client=httpx.AsyncClient(timeout=15, trust_env=False),
) as sdk:
with pytest.raises(NotFoundError) as raised:
await sdk.chat.completions.create(model="nope-9", messages=[{"role": "user", "content": "hi"}])
assert "does not exist" in str(raised.value), raised.value
warning: Final = _upstream_warning(owned.log)
assert "does not exist" in warning, warning
error_information: Final = _spend_error_information(raised.value.response.headers["x-litellm-call-id"])
assert "does not exist" in str(error_information["error_message"])
def test_gemini_passthrough_control_characters_cannot_forge_log_lines(gateway: Gateway, tmp_path: Path) -> None:
forged: Final = b'{"error": "line one"}\n2026-01-01 FAKE LOG LINE\x1b[31m\r' + b"x" * 4943 + b"\x00tail"
assert len(forged) == 5000
def respond(request: Request) -> Reply:
return Reply(status=502, body=forged, content_type="text/html")
path: Final = tmp_path / "gemini-forged.yaml"
with wire_server(respond) as wire:
_gemini_config(path, wire.url)
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
candidate: Final = owned.gateway
response: Final = candidate.request(
"POST", _GEMINI_MODEL_PATH, _GENERATE_CONTENT, headers={"x-goog-api-key": candidate.key}
)
assert response.status_code == 502, response.text
assert response.content == forged, response.text
warning: Final = _upstream_warning(owned.log)
assert "\n" not in warning and "\x1b" not in warning, warning
assert "line one" in warning and "FAKE LOG LINE" in warning, warning
assert warning.endswith("... (truncated at 4096 chars)"), warning
error_information: Final = _spend_error_information(response.headers["x-litellm-call-id"])
assert error_information["error_code"] == "502", response.text
def test_gemini_passthrough_empty_error_body_still_logged_and_proxy_serves(gateway: Gateway, tmp_path: Path) -> None:
def respond(request: Request) -> Reply:
if "claude-nope-9" in request.target:
return Reply(status=404, body=b"")
return Reply(status=200, body=b'{"candidates": [{"content": {"parts": [{"text": "ok"}]}}]}')
path: Final = tmp_path / "gemini-empty.yaml"
with wire_server(respond) as wire:
_gemini_config(path, wire.url)
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
candidate: Final = owned.gateway
response: Final = candidate.request(
"POST", _GEMINI_MODEL_PATH, _GENERATE_CONTENT, headers={"x-goog-api-key": candidate.key}
)
assert response.status_code == 404, response.text
assert response.content == b"", response.text
warning: Final = _upstream_warning(owned.log)
assert "returned 404" in warning, warning
error_information: Final = _spend_error_information(response.headers["x-litellm-call-id"])
assert error_information["error_code"] == "404", response.text
follow_up: Final = candidate.request(
"POST",
"/gemini/v1beta/models/healthy-model:generateContent",
_GENERATE_CONTENT,
headers={"x-goog-api-key": candidate.key},
)
assert follow_up.status_code == 200, follow_up.text
def test_gemini_passthrough_gzip_error_body_decoded_for_log_and_client(gateway: Gateway, tmp_path: Path) -> None:
upstream_error: Final = {"error": {"message": "gzipped upstream says the model is gone"}}
def respond(request: Request) -> Reply:
return Reply(
status=400,
body=gzip.compress(json.dumps(upstream_error).encode()),
headers={"content-encoding": "gzip"},
)
path: Final = tmp_path / "gemini-gzip.yaml"
with wire_server(respond) as wire:
_gemini_config(path, wire.url)
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
candidate: Final = owned.gateway
response: Final = candidate.request(
"POST", _GEMINI_MODEL_PATH, _GENERATE_CONTENT, headers={"x-goog-api-key": candidate.key}
)
assert response.status_code == 400, response.text
assert response.json() == upstream_error, response.text
warning: Final = _upstream_warning(owned.log)
assert "gzipped upstream says the model is gone" in warning, warning
def test_gemini_passthrough_streaming_gzip_error_body_decoded_for_log_and_client(
gateway: Gateway, tmp_path: Path
) -> None:
upstream_error: Final = {"error": {"message": "streamed gzip upstream denies the deployment"}}
compressed: Final = gzip.compress(json.dumps(upstream_error).encode())
third: Final = len(compressed) // 3
def respond(request: Request) -> Reply:
return Reply(
status=403,
chunks=(compressed[:third], compressed[third : 2 * third], compressed[2 * third :]),
headers={"content-encoding": "gzip"},
)
path: Final = tmp_path / "gemini-stream-gzip.yaml"
with wire_server(respond) as wire:
_gemini_config(path, wire.url)
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
candidate: Final = owned.gateway
with candidate.client.stream(
"POST",
_GEMINI_STREAM_PATH,
params={"alt": "sse"},
json=_GENERATE_CONTENT,
headers=_gemini_headers(candidate),
) as response:
assert response.status_code == 403
streamed: Final = response.read()
assert json.loads(streamed) == upstream_error, streamed
warning: Final = _upstream_warning(owned.log)
assert "streamed gzip upstream denies the deployment" in warning, warning
def test_gemini_passthrough_error_body_redacted_when_message_logging_off(gateway: Gateway, tmp_path: Path) -> None:
upstream_error: Final = {"error": {"message": "sensitive upstream explanation"}}
def respond(request: Request) -> Reply:
return Reply(status=404, body=json.dumps(upstream_error).encode())
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
path: Final = tmp_path / "gemini-redacted.yaml"
with wire_server(respond) as wire:
config["environment_variables"] = {"GEMINI_API_BASE": wire.url, "GEMINI_API_KEY": "scripted"}
config["litellm_settings"]["turn_off_message_logging"] = True
path.write_text(yaml.safe_dump(config))
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
candidate: Final = owned.gateway
response: Final = candidate.request(
"POST", _GEMINI_MODEL_PATH, _GENERATE_CONTENT, headers={"x-goog-api-key": candidate.key}
)
assert response.status_code == 404, response.text
assert response.json() == upstream_error, response.text
warning: Final = _upstream_warning(owned.log)
assert "redacted-by-litellm" in warning, warning
assert "sensitive upstream explanation" not in warning, warning
error_information: Final = _spend_error_information(response.headers["x-litellm-call-id"])
error_message: Final = str(error_information["error_message"])
assert "redacted-by-litellm" in error_message, error_message
assert "sensitive upstream explanation" not in error_message, error_message
def test_gemini_passthrough_exact_4096_byte_body_logged_without_marker(gateway: Gateway, tmp_path: Path) -> None:
body: Final = b'{"error": "' + b"y" * 4083 + b'"}'
assert len(body) == 4096
def respond(request: Request) -> Reply:
return Reply(status=404, body=body)
path: Final = tmp_path / "gemini-exact.yaml"
with wire_server(respond) as wire:
_gemini_config(path, wire.url)
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
candidate: Final = owned.gateway
response: Final = candidate.request(
"POST", _GEMINI_MODEL_PATH, _GENERATE_CONTENT, headers={"x-goog-api-key": candidate.key}
)
assert response.status_code == 404, response.text
warning: Final = _upstream_warning(owned.log)
assert body[:512].decode() in warning, warning
assert "(truncated at 4096 chars)" not in warning, warning
def test_gemini_passthrough_4097_byte_body_truncated_with_marker(gateway: Gateway, tmp_path: Path) -> None:
body: Final = b'{"error": "' + b"y" * 4084 + b'"}'
assert len(body) == 4097
def respond(request: Request) -> Reply:
return Reply(status=404, body=body)
path: Final = tmp_path / "gemini-over.yaml"
with wire_server(respond) as wire:
_gemini_config(path, wire.url)
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
candidate: Final = owned.gateway
response: Final = candidate.request(
"POST", _GEMINI_MODEL_PATH, _GENERATE_CONTENT, headers={"x-goog-api-key": candidate.key}
)
assert response.status_code == 404, response.text
warning: Final = _upstream_warning(owned.log)
assert body[:512].decode() in warning, warning
assert warning.endswith("... (truncated at 4096 chars)"), warning
def test_gemini_passthrough_one_byte_stream_chunks_reassembled_and_logged(gateway: Gateway, tmp_path: Path) -> None:
body: Final = json.dumps(_UPSTREAM_ERROR).encode()
def respond(request: Request) -> Reply:
return Reply(status=404, chunks=tuple(bytes([byte]) for byte in body))
path: Final = tmp_path / "gemini-one-byte.yaml"
with wire_server(respond) as wire:
_gemini_config(path, wire.url)
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
candidate: Final = owned.gateway
with candidate.client.stream(
"POST",
_GEMINI_STREAM_PATH,
params={"alt": "sse"},
json=_GENERATE_CONTENT,
headers=_gemini_headers(candidate),
) as response:
assert response.status_code == 404
streamed: Final = response.read()
assert streamed == body
warning: Final = _upstream_warning(owned.log)
assert "was not found or your project" in warning, warning
def test_gemini_passthrough_repeated_errors_each_get_row_and_log_line(gateway: Gateway, tmp_path: Path) -> None:
def respond(request: Request) -> Reply:
return Reply(status=404, body=json.dumps(_UPSTREAM_ERROR).encode())
path: Final = tmp_path / "gemini-twice.yaml"
with wire_server(respond) as wire:
_gemini_config(path, wire.url)
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
candidate: Final = owned.gateway
responses: Final = tuple(
candidate.request(
"POST", _GEMINI_MODEL_PATH, _GENERATE_CONTENT, headers={"x-goog-api-key": candidate.key}
)
for _ in range(2)
)
call_ids: Final = tuple(response.headers["x-litellm-call-id"] for response in responses)
assert len(set(call_ids)) == 2
for response in responses:
assert response.status_code == 404, response.text
error_information: Final = _spend_error_information(response.headers["x-litellm-call-id"])
assert "was not found or your project" in str(error_information["error_message"]), response.text
eventually(
lambda: _upstream_warnings(owned.log, "returned 404"),
lambda lines: len(lines) == 2,
seconds=30,
)
def test_budget_rejected_call_keeps_budget_normalized_error(gateway: Gateway, tmp_path: Path) -> None:
path: Final = tmp_path / "budget.yaml"
path.write_text(Path("tests/integration/proxy_config.yaml").read_text())
with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned:
candidate: Final = owned.gateway
with candidate.scenario() as scenario:
model: Final = scenario.model()
key: Final = scenario.key(max_budget=0.000001)
first: Final = candidate.chat(model, key=key)
assert "id" in first, first
rejected: Final = candidate.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": "over budget"}]},
key=key,
)
assert rejected.status_code == 422 and "budget_exceeded" in rejected.text, rejected.text
digest: Final = sha256(key.encode()).hexdigest()
rows: Final = eventually(
lambda: read_rows(
'SELECT metadata FROM "LiteLLM_SpendLogs" WHERE api_key=%s',
(digest,),
),
lambda values: any(
"BUDGET_EXCEEDED"
in str(
object_value(
json.loads(row["metadata"])
if isinstance(row["metadata"], str)
else object_value(row["metadata"])
)["error_information"]
)
for row in values
),
seconds=70,
)
budget_rows: Final = tuple(
row
for row in rows
if "BUDGET_EXCEEDED"
in str(
object_value(
json.loads(row["metadata"])
if isinstance(row["metadata"], str)
else object_value(row["metadata"])
)["error_information"]
)
)
assert len(budget_rows) == 1, budget_rows

View file

@ -0,0 +1,538 @@
import json
import re
import signal
import threading
import uuid
from collections.abc import Callable, Iterator, Mapping
from concurrent.futures import ThreadPoolExecutor
from contextlib import ExitStack, contextmanager
from dataclasses import dataclass
from pathlib import Path
from typing import Final
import psutil
import yaml
from integration._support.client import Gateway, eventually
from integration._support.process import OwnedProxy, group_members, owned_proxy_process
from integration._support.wire import Reply, Request, Wire, wire_server
from openai import OpenAI
from pydantic import BaseModel
PERSON: Final = "John Smith"
MASK: Final = "<PERSON>"
GEMINI_MODEL: Final = "gemini-2.5-flash"
def gemini_frame(text: str) -> bytes:
payload: Final = {
"candidates": [{"content": {"parts": [{"text": text}], "role": "model"}, "index": 0}],
"usageMetadata": {"promptTokenCount": 10, "candidatesTokenCount": 5, "totalTokenCount": 15},
"modelVersion": GEMINI_MODEL,
}
return b"data: " + json.dumps(payload).encode() + b"\r\n\r\n"
class GeminiPart(BaseModel):
text: str
class GeminiContent(BaseModel):
parts: list[GeminiPart]
class GeminiCandidate(BaseModel):
content: GeminiContent
class GeminiFrame(BaseModel):
candidates: list[GeminiCandidate]
def data_payloads(raw: bytes) -> tuple[dict[str, object], ...]:
"""JSON payload of each ``data:`` frame, whatever line ending the sender used."""
return tuple(json.loads(line[len("data: ") :]) for line in raw.decode().splitlines() if line.startswith("data: "))
def gemini_text(payload: Mapping[str, object]) -> str:
return GeminiFrame.model_validate(payload).candidates[0].content.parts[0].text
def gemini_texts(raw: bytes) -> tuple[str, ...]:
return tuple(gemini_text(payload) for payload in data_payloads(raw))
def anthropic_frame(event_type: str, payload: dict[str, object]) -> bytes:
return f"event: {event_type}\ndata: {json.dumps(payload)}\n\n".encode()
def anthropic_stream(identity: str, text: str) -> tuple[bytes, ...]:
return (
anthropic_frame(
"message_start",
{
"type": "message_start",
"message": {
"id": identity,
"type": "message",
"role": "assistant",
"model": "claude-sonnet-4-5-20250929",
"content": [],
"stop_reason": None,
"usage": {"input_tokens": 11, "output_tokens": 0},
},
},
),
anthropic_frame(
"content_block_start",
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
),
anthropic_frame(
"content_block_delta",
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": text}},
),
anthropic_frame("content_block_stop", {"type": "content_block_stop", "index": 0}),
anthropic_frame(
"message_delta",
{
"type": "message_delta",
"delta": {"stop_reason": "end_turn", "stop_sequence": None},
"usage": {"output_tokens": 4},
},
),
anthropic_frame("message_stop", {"type": "message_stop"}),
)
def openai_frame(identity: str, delta: dict[str, str], finish: str | None = None) -> bytes:
payload: Final = {
"id": identity,
"object": "chat.completion.chunk",
"created": 1,
"model": "gpt-4o-mini",
"choices": [{"index": 0, "delta": delta, "finish_reason": finish}],
}
return b"data: " + json.dumps(payload).encode() + b"\n\n"
def analyzer(request: Request) -> Reply:
assert request.target == "/analyze", request.target
text: Final = json.loads(request.body)["text"]
findings: Final = [
{"entity_type": "PERSON", "start": match.start(), "end": match.end(), "score": 0.85}
for match in re.finditer(re.escape(PERSON), text)
]
return Reply(body=json.dumps(findings).encode())
def anonymizer(request: Request) -> Reply:
assert request.target == "/anonymize", request.target
body: Final = json.loads(request.body)
text: Final = body["text"]
items: Final = [
{"entity_type": "PERSON", "start": item["start"], "end": item["end"], "operator": "replace", "text": MASK}
for item in body["analyzer_results"]
]
return Reply(body=json.dumps({"text": text.replace(PERSON, MASK), "items": items}).encode())
def broken(request: Request) -> Reply:
return Reply(status=500, body=b'{"error": "scripted outage"}')
@dataclass(frozen=True, slots=True)
class Received:
status: int
frames: tuple[bytes, ...]
@property
def text(self) -> str:
return b"".join(self.frames).decode()
@dataclass(frozen=True, slots=True)
class Rig:
proxy: OwnedProxy
upstream: Wire
analyzer: Wire
anonymizer: Wire
guardrail: str
gemini: str
anthropic: str
openai: str
@property
def gateway(self) -> Gateway:
return self.proxy.gateway
def stream(self, path: str, body: dict[str, object] | None = None, *, key: str | None = None) -> Received:
with self.gateway.client.stream(
"POST", path, json=body, headers={"Authorization": f"Bearer {key or self.gateway.key}"}
) as response:
return Received(response.status_code, tuple(response.iter_raw()))
def gemini_path(self) -> str:
return f"/v1beta/models/{self.gemini}:streamGenerateContent?alt=sse"
def gemini_body(self) -> dict[str, object]:
return {"contents": [{"role": "user", "parts": [{"text": "who designed it"}]}]}
def messages_body(self, *, guardrails: tuple[str, ...] | None = None) -> dict[str, object]:
return {
"model": self.anthropic,
"max_tokens": 64,
"stream": True,
"messages": [{"role": "user", "content": "who designed it"}],
**({"guardrails": list(guardrails)} if guardrails is not None else {}),
}
def anthropic_text(received: Received) -> str:
events: Final = tuple(
json.loads(line.removeprefix("data: ")) for line in received.text.split("\n") if line.startswith("data: ")
)
return "".join(event["delta"]["text"] for event in events if event.get("type") == "content_block_delta")
@contextmanager
def presidio_rig(
gateway: Gateway,
tmp_path: Path,
provider: Callable[[Request], Reply],
*,
analyze: Callable[[Request], Reply] = analyzer,
anonymize: Callable[[Request], Reply] = anonymizer,
default_on: bool = True,
) -> Iterator[Rig]:
guardrail: Final = "presidio" + uuid.uuid4().hex
with ExitStack() as stack:
upstream: Final = stack.enter_context(wire_server(provider))
analyze_sink: Final = stack.enter_context(wire_server(analyze))
anonymize_sink: Final = stack.enter_context(wire_server(anonymize))
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["guardrails"] = [
{
"guardrail_name": guardrail,
"litellm_params": {
"guardrail": "presidio",
"mode": "post_call",
"default_on": default_on,
"presidio_analyzer_api_base": analyze_sink.url,
"presidio_anonymizer_api_base": anonymize_sink.url,
"presidio_filter_scope": "output",
},
}
]
path: Final = tmp_path / f"{guardrail}.yaml"
path.write_text(yaml.safe_dump(config))
proxy: Final = stack.enter_context(owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2))
scenario: Final = stack.enter_context(proxy.gateway.scenario())
yield Rig(
proxy=proxy,
upstream=upstream,
analyzer=analyze_sink,
anonymizer=anonymize_sink,
guardrail=guardrail,
gemini=scenario.model(
model=f"gemini/{GEMINI_MODEL}", api_base=upstream.url, api_key="synthetic-gemini-key"
),
anthropic=scenario.model(
model="anthropic/claude-sonnet-4-5-20250929", api_base=upstream.url, api_key="synthetic-anthropic-key"
),
openai=scenario.model(model="openai/gpt-4o-mini", api_base=upstream.url + "/v1", api_key="synthetic-key"),
)
def gemini_provider(reply: Reply) -> Callable[[Request], Reply]:
def provider(request: Request) -> Reply:
assert "streamGenerateContent" in request.target, request.target
return reply
return provider
def test_native_gemini_first_frame_reaches_caller_before_upstream_sends_the_second(
gateway: Gateway, tmp_path: Path
) -> None:
gate: Final = threading.Event()
first: Final = gemini_frame("first ")
second: Final = gemini_frame("second ")
provider: Final = gemini_provider(
Reply(content_type="text/event-stream", chunks=(first, second), gate_after_first=gate)
)
with presidio_rig(gateway, tmp_path, provider) as rig:
with rig.gateway.client.stream(
"POST", rig.gemini_path(), json=rig.gemini_body(), headers={"Authorization": f"Bearer {rig.gateway.key}"}
) as response:
assert response.status_code == 200, response.read().decode()
chunks: Final = response.iter_raw()
arrived: Final = next(chunks)
assert gemini_texts(arrived) == ("first ",), f"first chunk while upstream is gated: {arrived!r}"
gate.set()
rest: Final = b"".join(chunks)
assert gemini_texts(rest) == ("second ",), rest
assert len(rig.upstream.drain()) == 1
assert rig.analyzer.drain() == () and rig.anonymizer.drain() == ()
def test_native_gemini_frames_received_before_upstream_abort_reach_caller(gateway: Gateway, tmp_path: Path) -> None:
frames: Final = (gemini_frame(f"chunk {index} from {PERSON}. ") for index in range(3))
provider: Final = gemini_provider(
Reply(content_type="text/event-stream", chunks=tuple(frames), abort_after=2, pause_between_chunks=0.2)
)
with presidio_rig(gateway, tmp_path, provider) as rig:
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body())
assert received.status == 200, received.text
*frames_before_abort, trailer = data_payloads(b"".join(received.frames))
assert [gemini_text(frame) for frame in frames_before_abort] == [
f"chunk 0 from {PERSON}. ",
f"chunk 1 from {PERSON}. ",
], received.text
assert "candidates" not in trailer and json.dumps(trailer).count('"code": "500"') == 1, received.text
assert len(rig.upstream.drain()) == 1
def test_native_gemini_first_frame_split_into_transport_fragments_streams_every_byte(
gateway: Gateway, tmp_path: Path
) -> None:
first: Final = gemini_frame(f"fragmented {PERSON}")
second: Final = gemini_frame("whole")
chunks: Final = (first[:7], first[7:19], first[19:], second)
provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=chunks))
with presidio_rig(gateway, tmp_path, provider) as rig:
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body())
assert received.status == 200, received.text
assert gemini_texts(b"".join(received.frames)) == (f"fragmented {PERSON}", "whole")
def test_native_gemini_non_json_frame_passes_through_unchanged(gateway: Gateway, tmp_path: Path) -> None:
frames: Final = (b"data: not json at all\r\n\r\n", gemini_frame("after"))
provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=frames))
with presidio_rig(gateway, tmp_path, provider) as rig:
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body())
assert received.status == 200, received.text
assert received.text.replace("\r\n", "\n") == b"".join(frames).decode().replace("\r\n", "\n")
def test_native_gemini_empty_stream_returns_200_with_no_body(gateway: Gateway, tmp_path: Path) -> None:
provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=()))
with presidio_rig(gateway, tmp_path, provider) as rig:
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body())
assert received.status == 200, received.text
assert received.text == ""
def test_native_gemini_streams_while_presidio_analyzer_is_down(gateway: Gateway, tmp_path: Path) -> None:
frames: Final = (gemini_frame(f"{PERSON} one. "), gemini_frame("two."))
provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=frames))
with presidio_rig(gateway, tmp_path, provider, analyze=broken) as rig:
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body())
assert received.status == 200, received.text
assert gemini_texts(b"".join(received.frames)) == (f"{PERSON} one. ", "two.")
assert rig.analyzer.drain() == ()
def test_native_gemini_unauthenticated_request_is_rejected_before_upstream(gateway: Gateway, tmp_path: Path) -> None:
provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=(gemini_frame("never"),)))
with presidio_rig(gateway, tmp_path, provider) as rig:
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body(), key="sk-not-a-key")
assert received.status == 401, received.text
assert rig.upstream.drain() == ()
def anthropic_provider(chunks: tuple[bytes, ...]) -> Callable[[Request], Reply]:
def provider(request: Request) -> Reply:
assert request.target == "/v1/messages", request.target
return Reply(content_type="text/event-stream", chunks=chunks)
return provider
def test_anthropic_messages_stream_masks_person_in_text_delta(gateway: Gateway, tmp_path: Path) -> None:
identity: Final = "msg_" + uuid.uuid4().hex
provider: Final = anthropic_provider(anthropic_stream(identity, f"{PERSON} designed it."))
with presidio_rig(gateway, tmp_path, provider) as rig:
received: Final = rig.stream("/v1/messages", rig.messages_body())
assert received.status == 200, received.text
assert anthropic_text(received) == f"{MASK} designed it."
assert PERSON not in received.text
assert identity in received.text
analyzed: Final = rig.analyzer.drain()
anonymized: Final = rig.anonymizer.drain()
assert len(analyzed) == len(anonymized) == 1
assert json.loads(analyzed[0].body)["text"] == f"{PERSON} designed it."
def test_anthropic_messages_first_frame_split_across_transport_chunks_is_still_masked(
gateway: Gateway, tmp_path: Path
) -> None:
identity: Final = "msg_" + uuid.uuid4().hex
whole: Final = anthropic_stream(identity, f"{PERSON} designed it.")
split_at: Final = whole[0].index(b'"message_') + len(b'"message_')
chunks: Final = (whole[0][:split_at], whole[0][split_at:], *whole[1:])
with presidio_rig(gateway, tmp_path, anthropic_provider(chunks)) as rig:
received: Final = rig.stream("/v1/messages", rig.messages_body())
assert received.status == 200, received.text
assert anthropic_text(received) == f"{MASK} designed it."
assert received.text.count("event: message_start") == 1
def test_anthropic_messages_stream_fails_closed_when_analyzer_is_down(gateway: Gateway, tmp_path: Path) -> None:
identity: Final = "msg_" + uuid.uuid4().hex
provider: Final = anthropic_provider(anthropic_stream(identity, f"{PERSON} designed it."))
with presidio_rig(gateway, tmp_path, provider, analyze=broken) as rig:
received: Final = rig.stream("/v1/messages", rig.messages_body())
assert PERSON not in received.text, received.text
assert "Presidio analyzer" in received.text, received.text
assert rig.anonymizer.drain() == ()
def test_anthropic_messages_per_request_guardrails_selects_masking(gateway: Gateway, tmp_path: Path) -> None:
identity: Final = "msg_" + uuid.uuid4().hex
provider: Final = anthropic_provider(anthropic_stream(identity, f"{PERSON} designed it."))
with presidio_rig(gateway, tmp_path, provider, default_on=False) as rig:
unguarded: Final = rig.stream("/v1/messages", rig.messages_body())
assert unguarded.status == 200, unguarded.text
assert anthropic_text(unguarded) == f"{PERSON} designed it."
assert rig.analyzer.drain() == ()
guarded: Final = rig.stream("/v1/messages", rig.messages_body(guardrails=(rig.guardrail,)))
assert guarded.status == 200, guarded.text
assert anthropic_text(guarded) == f"{MASK} designed it."
assert len(rig.analyzer.drain()) == 1
def openai_provider(identity: str) -> Callable[[Request], Reply]:
def provider(request: Request) -> Reply:
assert request.target == "/v1/chat/completions", request.target
if json.loads(request.body).get("stream"):
return Reply(
content_type="text/event-stream",
chunks=(
openai_frame(identity, {"role": "assistant", "content": ""}),
openai_frame(identity, {"content": f"{PERSON} designed"}),
openai_frame(identity, {"content": " it."}, "stop"),
b"data: [DONE]\n\n",
),
)
return Reply(
body=json.dumps(
{
"id": identity,
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": f"{PERSON} designed it."},
"finish_reason": "stop",
}
],
"usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15},
}
).encode()
)
return provider
def test_chat_completions_openai_sdk_stream_and_non_stream_are_masked(gateway: Gateway, tmp_path: Path) -> None:
identity: Final = "chatcmpl-" + uuid.uuid4().hex
with presidio_rig(gateway, tmp_path, openai_provider(identity)) as rig:
client: Final = OpenAI(api_key=rig.gateway.key, base_url=f"{rig.gateway.client.base_url}/v1", max_retries=0)
streamed: Final = client.chat.completions.create(
model=rig.openai, messages=[{"role": "user", "content": "who designed it"}], stream=True
)
pieces: Final = tuple(
chunk.choices[0].delta.content for chunk in streamed if chunk.choices and chunk.choices[0].delta.content
)
assert "".join(pieces) == f"{MASK} designed it.", pieces
whole: Final = client.chat.completions.create(
model=rig.openai, messages=[{"role": "user", "content": "who designed it"}]
)
assert whole.id == identity
assert whole.choices[0].message.content == f"{MASK} designed it."
assert len(rig.upstream.drain()) == 2
assert len(rig.analyzer.drain()) == len(rig.anonymizer.drain()) == 2
def test_mixed_burst_survives_anonymizer_outage_and_recovers(gateway: Gateway, tmp_path: Path) -> None:
outage: Final = threading.Event()
def flaky_anonymizer(request: Request) -> Reply:
return Reply(status=503, body=b'{"error": "scripted outage"}') if outage.is_set() else anonymizer(request)
def provider(request: Request) -> Reply:
if request.target == "/v1/messages":
identity: Final = "msg_" + json.loads(request.body)["messages"][0]["content"]
return Reply(content_type="text/event-stream", chunks=anthropic_stream(identity, f"{PERSON} designed it."))
return Reply(
content_type="text/event-stream",
chunks=(gemini_frame(f"{PERSON} "), gemini_frame("designed it.")),
pause_between_chunks=0.05,
)
with presidio_rig(gateway, tmp_path, provider, anonymize=flaky_anonymizer) as rig:
def gemini_call(index: int) -> tuple[str, str, int]:
received: Final = rig.stream(rig.gemini_path(), rig.gemini_body())
return (
"gemini",
f"g{index}",
received.status if gemini_texts(b"".join(received.frames)) == (f"{PERSON} ", "designed it.") else -1,
)
def anthropic_call(index: int) -> tuple[str, str, int]:
body: Final = {**rig.messages_body(), "messages": [{"role": "user", "content": f"a{index}"}]}
received: Final = rig.stream("/v1/messages", body)
leaked: Final = PERSON in received.text
return ("anthropic", f"a{index}", -1 if leaked else (1 if MASK in received.text else 0))
def phase(offset: int) -> tuple[tuple[str, str, int], ...]:
with ThreadPoolExecutor(max_workers=12) as pool:
futures: Final = tuple(
pool.submit(gemini_call if index % 2 == 0 else anthropic_call, offset + index)
for index in range(12)
)
return tuple(future.result() for future in futures)
healthy_before: Final = phase(0)
outage.set()
during: Final = phase(100)
outage.clear()
healthy_after: Final = phase(200)
for name, results in (("before", healthy_before), ("during", during), ("after", healthy_after)):
assert all(status == 200 for kind, _, status in results if kind == "gemini"), (name, results)
assert all(status == 1 for kind, _, status in healthy_before + healthy_after if kind == "anthropic"), (
healthy_before,
healthy_after,
)
assert all(status == 0 for kind, _, status in during if kind == "anthropic"), during
identities: Final = tuple(identity for _, identity, _ in healthy_before + during + healthy_after)
assert len(identities) == len(set(identities)) == 36
def test_native_gemini_keeps_streaming_after_one_worker_is_killed(gateway: Gateway, tmp_path: Path) -> None:
frames: Final = (gemini_frame("alive "), gemini_frame("still."))
provider: Final = gemini_provider(Reply(content_type="text/event-stream", chunks=frames, pause_between_chunks=0.05))
with presidio_rig(gateway, tmp_path, provider) as rig:
workers: Final = eventually(
lambda: tuple(
member for member in group_members(rig.proxy.process.pid) if member.pid != rig.proxy.process.pid
),
lambda members: len(members) >= 2,
seconds=30,
)
victim: Final = workers[0]
with ThreadPoolExecutor(max_workers=8) as pool:
futures: Final = tuple(pool.submit(rig.stream, rig.gemini_path(), rig.gemini_body()) for _ in range(8))
victim.send_signal(signal.SIGKILL)
psutil.wait_procs((victim,), timeout=10)
first_wave: Final = tuple(future.result() for future in futures)
survivors: Final = tuple(received for received in first_wave if received.status == 200)
assert survivors, [received.text[:200] for received in first_wave]
assert all(gemini_texts(b"".join(received.frames)) == ("alive ", "still.") for received in survivors)
second_wave: Final = tuple(rig.stream(rig.gemini_path(), rig.gemini_body()) for _ in range(6))
assert all(received.status == 200 for received in second_wave), [r.text[:200] for r in second_wave]
assert all(gemini_texts(b"".join(received.frames)) == ("alive ", "still.") for received in second_wave)
assert rig.proxy.process.poll() is None

View file

@ -0,0 +1,97 @@
import re
import uuid
from pathlib import Path
from typing import Final
import pytest
from _s3_v2_support import (
BUCKET,
PREFIX,
RecordingS3Sink,
collect_payloads,
matched_ids,
mixed_burst,
s3_config,
surface_reply,
)
from integration._support.client import Gateway
from integration._support.process import owned_proxy
from integration._support.wire import wire_server
PER_REQUEST_KEY: Final = re.compile(rf"^/{BUCKET}/{PREFIX}/\d{{4}}-\d{{2}}-\d{{2}}/.+\.json$")
BATCH_KEY: Final = re.compile(
rf"^/{BUCKET}/{PREFIX}/\d{{4}}-\d{{2}}-\d{{2}}/batch_\d{{2}}-\d{{2}}-\d{{2}}_[0-9a-f]{{32}}\.jsonl$"
)
@pytest.mark.covers("other.observability.s3_v2.mixed_surface_burst_bounds_puts_one_object_per_response_id")
def test_s3_v2_mixed_surface_burst_bounds_puts_one_object_per_response_id(gateway: Gateway, tmp_path: Path) -> None:
marker: Final = "s3mix" + uuid.uuid4().hex[:8]
sink: Final = RecordingS3Sink()
with wire_server(surface_reply) as provider, wire_server(sink.respond) as bucket:
config: Final = s3_config(tmp_path, bucket.url, {})
with (
owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "3"}, config=config) as candidate,
candidate.scenario() as scenario,
):
openai_model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key")
anthropic_model: Final = scenario.model(
model="anthropic/claude-sonnet-4-5-20250929", api_base=provider.url, api_key="synthetic-provider-key"
)
key: Final = scenario.key(models=[openai_model, anthropic_model])
answered: Final = mixed_burst(candidate, openai_model, anthropic_model, key, marker)
payloads: Final = collect_payloads(sink, len(answered))
targets: Final = tuple(sink.objects())
assert sum(1 for r in provider.drain() if r.method == "POST") == 48
assert sink.peak <= 16, f"peak concurrent PUTs {sink.peak} exceeded the default bound"
assert all(PER_REQUEST_KEY.match(target) for target in targets), list(targets)
assert len(targets) == 48
assert matched_ids(payloads, answered)
@pytest.mark.covers("other.observability.s3_v2.mixed_surface_batch_writes_ndjson_lines_per_response_id")
def test_s3_v2_mixed_surface_batch_writes_ndjson_lines_per_response_id(gateway: Gateway, tmp_path: Path) -> None:
marker: Final = "s3mixb" + uuid.uuid4().hex[:8]
sink: Final = RecordingS3Sink()
with wire_server(surface_reply) as provider, wire_server(sink.respond) as bucket:
config: Final = s3_config(tmp_path, bucket.url, {"s3_batch_file_upload": True})
with (
owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "3"}, config=config) as candidate,
candidate.scenario() as scenario,
):
openai_model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key")
anthropic_model: Final = scenario.model(
model="anthropic/claude-sonnet-4-5-20250929", api_base=provider.url, api_key="synthetic-provider-key"
)
key: Final = scenario.key(models=[openai_model, anthropic_model])
answered: Final = mixed_burst(candidate, openai_model, anthropic_model, key, marker)
payloads: Final = collect_payloads(sink, len(answered))
targets: Final = tuple(sink.objects())
puts: Final = bucket.drain()
assert sum(1 for r in provider.drain() if r.method == "POST") == 48
assert all(BATCH_KEY.match(target) for target in targets), list(targets)
assert all(put.headers["content-type"] == "application/x-ndjson" for put in puts), [put.headers for put in puts]
assert matched_ids(payloads, answered)
assert len(payloads) == 48
@pytest.mark.covers("other.observability.s3_v2.sink_outage_mid_mixed_burst_recovers_every_response_id")
def test_s3_v2_sink_outage_mid_mixed_burst_recovers_every_response_id(gateway: Gateway, tmp_path: Path) -> None:
marker: Final = "s3mixo" + uuid.uuid4().hex[:8]
sink: Final = RecordingS3Sink(fail_attempts=30, fail_status=503, delay_seconds=0.2)
with wire_server(surface_reply) as provider, wire_server(sink.respond) as bucket:
config: Final = s3_config(tmp_path, bucket.url, {})
with (
owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "3"}, config=config) as candidate,
candidate.scenario() as scenario,
):
openai_model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key")
anthropic_model: Final = scenario.model(
model="anthropic/claude-sonnet-4-5-20250929", api_base=provider.url, api_key="synthetic-provider-key"
)
key: Final = scenario.key(models=[openai_model, anthropic_model])
answered: Final = mixed_burst(candidate, openai_model, anthropic_model, key, marker)
payloads: Final = collect_payloads(sink, len(answered), seconds=90)
assert sum(1 for r in provider.drain() if r.method == "POST") == 48
assert matched_ids(payloads, answered)
assert len(payloads) == 48, "a stored id was overwritten or duplicated"

View file

@ -0,0 +1,630 @@
import json
import re
import threading
import time
import uuid
from collections.abc import Mapping
from concurrent.futures import ThreadPoolExecutor
from dataclasses import dataclass, field
from pathlib import Path
from typing import Final
import httpx
import pytest
import yaml
from _s3_v2_support import RecordingS3Sink, collect_payloads
from _s3_v2_support import s3_config as _recording_s3_config
from integration._support.client import Gateway, JsonValue, eventually
from integration._support.process import group_members, owned_proxy, owned_proxy_process
from integration._support.wire import Reply, Request, Wire, wire_server
BUCKET: Final = "integration-bucket"
PREFIX: Final = "integration-logs"
REQUESTS: Final = 64
PUT_DELAY_SECONDS: Final = 0.5
@dataclass(slots=True)
class S3Sink:
"""Accepts every PUT after a fixed delay and records the peak number of PUTs in flight."""
lock: threading.Lock = field(default_factory=threading.Lock)
in_flight: int = 0
peak: int = 0
def respond(self, request: Request) -> Reply:
assert request.method == "PUT", request.method
assert request.target.startswith(f"/{BUCKET}/{PREFIX}/"), request.target
with self.lock:
self.in_flight += 1
self.peak = max(self.peak, self.in_flight)
time.sleep(PUT_DELAY_SECONDS)
with self.lock:
self.in_flight -= 1
return Reply()
def _chat_reply(request: Request) -> Reply:
if request.method != "POST" or not request.body:
return Reply(status=404)
text: Final = json.loads(request.body)["messages"][0]["content"]
return Reply(
body=json.dumps(
{
"id": text,
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [{"index": 0, "message": {"role": "assistant", "content": text}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15},
}
).encode()
)
def _s3_config(path: Path, sink_url: str, extra: Mapping[str, JsonValue]) -> Path:
config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())
config["litellm_settings"].update(
{
"callbacks": ["s3_v2"],
"s3_callback_params": {
"s3_bucket_name": BUCKET,
"s3_region_name": "us-east-1",
"s3_endpoint_url": sink_url,
"s3_path": PREFIX,
"s3_aws_access_key_id": "AKIAIOSFODNN7EXAMPLE",
"s3_aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY",
**extra,
},
}
)
target: Final = path / "s3_v2.yaml"
target.write_text(yaml.safe_dump(config))
return target
def _burst(candidate: Gateway, model: str, key: str, marker: str) -> frozenset[str]:
ids: Final = tuple(f"{marker}-{index}" for index in range(REQUESTS))
def request(identity: str) -> str:
response: Final = candidate.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": identity}], "cache": {"no-cache": True}},
key=key,
)
assert response.status_code == 200, response.text
return response.json()["id"]
with ThreadPoolExecutor(max_workers=32) as pool:
returned: Final = frozenset(pool.map(request, ids))
assert returned == frozenset(ids)
return returned
def _collect(bucket: Wire, count_lines: bool, expected: int) -> tuple[Request, ...]:
puts: Final[list[Request]] = [] # mutable-ok: drain() consumes the queue, later polls must keep earlier PUTs
def delivered() -> int:
puts.extend(bucket.drain())
return sum(len(put.body.splitlines()) if count_lines else 1 for put in puts)
eventually(delivered, lambda total: total >= expected, seconds=30)
return tuple(puts)
PER_REQUEST_KEY: Final = re.compile(rf"^/{BUCKET}/{PREFIX}/\d{{4}}-\d{{2}}-\d{{2}}/.+\.json$")
BATCH_KEY: Final = re.compile(
rf"^/{BUCKET}/{PREFIX}/\d{{4}}-\d{{2}}-\d{{2}}/batch_\d{{2}}-\d{{2}}-\d{{2}}_[0-9a-f]{{32}}\.jsonl$"
)
@pytest.mark.covers("other.observability.s3_v2.flush_bounds_concurrent_puts_to_default_and_keeps_every_log")
def test_s3_v2_flush_bounds_concurrent_puts_to_the_default_of_sixteen(gateway: Gateway, tmp_path: Path) -> None:
marker: Final = "s3fan" + uuid.uuid4().hex[:8]
sink: Final = S3Sink()
with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket:
config: Final = _s3_config(tmp_path, bucket.url, {})
with (
owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "3"}, config=config) as candidate,
candidate.scenario() as scenario,
):
model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key")
key: Final = scenario.key(models=[model])
ids: Final = _burst(candidate, model, key, marker)
puts: Final = _collect(bucket, count_lines=False, expected=REQUESTS)
assert sum(1 for r in provider.drain() if r.method == "POST") == REQUESTS
assert sink.peak <= 16, f"peak concurrent PUTs {sink.peak} exceeded the default bound for {REQUESTS} queued logs"
assert all(PER_REQUEST_KEY.match(put.target) for put in puts), [put.target for put in puts]
assert frozenset(json.loads(put.body)["id"] for put in puts) == ids
assert len({put.target for put in puts}) == REQUESTS
@pytest.mark.covers("other.observability.s3_v2.configured_bound_and_env_backed_false_keeps_per_request_objects")
def test_s3_v2_honors_configured_bound_and_env_backed_false_batch_flag(gateway: Gateway, tmp_path: Path) -> None:
marker: Final = "s3cap" + uuid.uuid4().hex[:8]
sink: Final = S3Sink()
with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket:
config: Final = _s3_config(
tmp_path,
bucket.url,
{"s3_max_concurrent_uploads": 4, "s3_batch_file_upload": "os.environ/INTEGRATION_S3_BATCH_FILE_UPLOAD"},
)
with (
owned_proxy(
gateway,
tmp_path,
{"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "3", "INTEGRATION_S3_BATCH_FILE_UPLOAD": "false"},
config=config,
) as candidate,
candidate.scenario() as scenario,
):
model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key")
key: Final = scenario.key(models=[model])
ids: Final = _burst(candidate, model, key, marker)
puts: Final = _collect(bucket, count_lines=False, expected=REQUESTS)
assert sum(1 for r in provider.drain() if r.method == "POST") == REQUESTS
assert sink.peak <= 4, f"peak concurrent PUTs {sink.peak} exceeded s3_max_concurrent_uploads=4"
assert all(PER_REQUEST_KEY.match(put.target) for put in puts), [put.target for put in puts]
assert frozenset(json.loads(put.body)["id"] for put in puts) == ids
@pytest.mark.covers("other.observability.s3_v2.batch_file_upload_writes_one_ndjson_object_per_flush")
def test_s3_v2_batch_file_upload_writes_one_jsonl_object_per_flush(gateway: Gateway, tmp_path: Path) -> None:
marker: Final = "s3jsonl" + uuid.uuid4().hex[:8]
sink: Final = S3Sink()
with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket:
config: Final = _s3_config(tmp_path, bucket.url, {"s3_batch_file_upload": True})
with (
owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "3"}, config=config) as candidate,
candidate.scenario() as scenario,
):
model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key")
key: Final = scenario.key(models=[model])
ids: Final = _burst(candidate, model, key, marker)
puts: Final = _collect(bucket, count_lines=True, expected=REQUESTS)
assert sum(1 for r in provider.drain() if r.method == "POST") == REQUESTS
assert len(puts) <= 2, f"{len(puts)} PUTs for {REQUESTS} logs; batch mode must write one object per flush"
assert all(BATCH_KEY.match(put.target) for put in puts), [put.target for put in puts]
assert all(put.headers["content-type"] == "application/x-ndjson" for put in puts), [put.headers for put in puts]
lines: Final = tuple(line for put in puts for line in put.body.decode().splitlines())
assert frozenset(json.loads(line)["id"] for line in lines) == ids
assert len(lines) == REQUESTS
@pytest.mark.covers("other.observability.s3_v2.batch_file_upload_keeps_team_prefix_in_object_key")
def test_s3_v2_batch_file_upload_keeps_team_alias_prefix(gateway: Gateway, tmp_path: Path) -> None:
marker: Final = "s3team" + uuid.uuid4().hex[:8]
team_alias: Final = f"alpha-{uuid.uuid4().hex[:8]}"
team_batch_key: Final = re.compile(
rf"^/{BUCKET}/{PREFIX}/{team_alias}/\d{{4}}-\d{{2}}-\d{{2}}/batch_\d{{2}}-\d{{2}}-\d{{2}}_[0-9a-f]{{32}}\.jsonl$"
)
sink: Final = S3Sink()
with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket:
config: Final = _s3_config(tmp_path, bucket.url, {"s3_batch_file_upload": True, "s3_use_team_prefix": True})
with (
owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "3"}, config=config) as candidate,
candidate.scenario() as scenario,
):
model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key")
team: Final = scenario.team(team_alias=team_alias, models=[model])
key: Final = scenario.key(team_id=team, models=[model])
ids: Final = _burst(candidate, model, key, marker)
puts: Final = _collect(bucket, count_lines=True, expected=REQUESTS)
assert sum(1 for r in provider.drain() if r.method == "POST") == REQUESTS
assert len(puts) >= 1
assert all(team_batch_key.match(put.target) for put in puts), [put.target for put in puts]
lines: Final = tuple(line for put in puts for line in put.body.decode().splitlines())
assert frozenset(json.loads(line)["id"] for line in lines) == ids
assert len(lines) == REQUESTS
@pytest.mark.covers("other.observability.s3_v2.upstream_failure_events_land_alongside_successes")
def test_s3_v2_upstream_failure_events_land_alongside_successes(gateway: Gateway, tmp_path: Path) -> None:
marker: Final = "s3fail" + uuid.uuid4().hex[:8]
sink: Final = RecordingS3Sink()
def provider(request: Request) -> Reply:
text: Final = json.loads(request.body)["messages"][0]["content"]
if text.endswith("-fail"):
return Reply(
status=401,
body=b'{"error": {"message": "synthetic upstream rejection", "code": "synthetic_401"}}',
)
return _chat_reply(request)
with wire_server(provider) as upstream, wire_server(sink.respond) as bucket:
config: Final = _s3_config(tmp_path, bucket.url, {})
with (
owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "3"}, config=config) as candidate,
candidate.scenario() as scenario,
):
model: Final = scenario.model(api_base=upstream.url + "/v1", api_key="synthetic-provider-key")
key: Final = scenario.key(models=[model])
success_ids: Final = tuple(f"{marker}-{index}" for index in range(8))
failure_ids: Final = tuple(f"{marker}-{index}-fail" for index in range(4))
def send(identity: str) -> httpx.Response:
return candidate.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": identity}], "cache": {"no-cache": True}},
key=key,
)
with ThreadPoolExecutor(max_workers=12) as pool:
responses: Final = tuple(pool.map(send, (*success_ids, *failure_ids)))
ok: Final = responses[:8]
rejected: Final = responses[8:]
assert all(response.status_code == 200 for response in ok), [r.text for r in ok]
assert tuple(response.json()["id"] for response in ok) == success_ids
for response in rejected:
assert response.status_code in (400, 401), response.status_code
assert "synthetic upstream rejection" in response.text, response.text
failure_call_ids: Final = frozenset(response.headers["x-litellm-call-id"] for response in rejected)
payloads: Final = collect_payloads(sink, len(success_ids) + len(failure_ids))
assert len(upstream.drain()) == len(success_ids) + len(failure_ids)
delivered: Final = frozenset(payload["id"] for payload in payloads if payload["status"] == "success")
assert delivered == frozenset(success_ids)
failures: Final = tuple(payload for payload in payloads if payload["status"] == "failure")
assert len(failures) == len(failure_ids)
assert frozenset(payload["litellm_call_id"] for payload in failures) == failure_call_ids
assert all("synthetic upstream rejection" in json.dumps(payload["error_information"]) for payload in failures)
@pytest.mark.covers("other.observability.s3_v2.invalid_or_empty_bound_falls_back_to_sixteen")
@pytest.mark.parametrize(
("bad", "warns"),
[
pytest.param("abc", True, id="non_integer"),
pytest.param(0, True, id="below_one"),
pytest.param("", False, id="empty"),
],
)
def test_s3_v2_invalid_or_empty_bound_falls_back_to_sixteen(
gateway: Gateway, tmp_path: Path, bad: JsonValue, warns: bool
) -> None:
marker: Final = "s3bound" + uuid.uuid4().hex[:8]
sink: Final = RecordingS3Sink()
with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket:
config: Final = _s3_config(tmp_path, bucket.url, {"s3_max_concurrent_uploads": bad})
with (
owned_proxy_process(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "3"}, config=config) as owned,
owned.gateway.scenario() as scenario,
):
model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key")
key: Final = scenario.key(models=[model])
ids: Final = _burst(owned.gateway, model, key, marker)
payloads: Final = collect_payloads(sink, REQUESTS)
if warns:
eventually(
lambda: owned.log.read_text(),
lambda text: "s3_max_concurrent_uploads" in text,
seconds=15,
)
else:
assert "s3_max_concurrent_uploads" not in owned.log.read_text()
assert sum(1 for r in provider.drain() if r.method == "POST") == REQUESTS
assert sink.peak <= 16, f"peak concurrent PUTs {sink.peak} exceeded the fallback bound"
assert frozenset(payload["id"] for payload in payloads) == ids
@pytest.mark.covers("other.observability.s3_v2.sink_rejection_requeues_and_delivers_every_id_once")
def test_s3_v2_sink_rejection_requeues_and_delivers_every_id_once(gateway: Gateway, tmp_path: Path) -> None:
marker: Final = "s3deny" + uuid.uuid4().hex[:8]
sink: Final = RecordingS3Sink(fail_status=403, delay_seconds=0.2)
with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket:
config: Final = _s3_config(tmp_path, bucket.url, {})
with (
owned_proxy_process(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "3"}, config=config) as owned,
owned.gateway.scenario() as scenario,
):
model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key")
key: Final = scenario.key(models=[model])
sink.fail_until = time.time() + 10
ids: Final = _burst(owned.gateway, model, key, marker)
payloads: Final = collect_payloads(sink, REQUESTS, seconds=90)
eventually(
lambda: owned.log.read_text(),
lambda text: "S3BatchUploadError" in text,
seconds=15,
)
readiness: Final = owned.gateway.client.get("/health/readiness")
assert readiness.status_code == 200, readiness.text
assert sum(1 for r in provider.drain() if r.method == "POST") == REQUESTS
assert len(sink.objects()) == REQUESTS
assert frozenset(payload["id"] for payload in payloads) == ids
@pytest.mark.covers("other.observability.s3_v2.batch_retry_resends_identical_key_and_body")
def test_s3_v2_batch_retry_resends_identical_key_and_body(gateway: Gateway, tmp_path: Path) -> None:
marker: Final = "s3retry" + uuid.uuid4().hex[:8]
sink: Final = RecordingS3Sink(fail_status=500, delay_seconds=0.2)
with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket:
config: Final = _s3_config(tmp_path, bucket.url, {"s3_batch_file_upload": True})
with (
owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "3"}, config=config) as candidate,
candidate.scenario() as scenario,
):
model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key")
key: Final = scenario.key(models=[model])
sink.fail_until = time.time() + 8
ids: Final = _burst(candidate, model, key, marker)
payloads: Final = collect_payloads(sink, REQUESTS, seconds=90)
puts: Final = bucket.drain()
assert sum(1 for r in provider.drain() if r.method == "POST") == REQUESTS
by_target: Final = {}
for put in puts:
by_target.setdefault(put.target, set()).add(put.body) # mutable-ok: grouping attempts seen so far per target
assert all(len(bodies) == 1 for bodies in by_target.values()), "a retried batch PUT changed key or body"
assert max(sum(1 for put in puts if put.target == target) for target in by_target) >= 2, "no retried PUT observed"
assert frozenset(payload["id"] for payload in payloads) == ids
assert len(payloads) == REQUESTS
@pytest.mark.covers("other.observability.s3_v2.unknown_model_rejection_keeps_other_requests_logging")
def test_s3_v2_unknown_model_rejection_keeps_other_requests_logging(gateway: Gateway, tmp_path: Path) -> None:
marker: Final = "s3ghost" + uuid.uuid4().hex[:8]
sink: Final = RecordingS3Sink()
with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket:
config: Final = _s3_config(tmp_path, bucket.url, {})
with (
owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "3"}, config=config) as candidate,
candidate.scenario() as scenario,
):
model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key")
key: Final = scenario.key(models=[model])
ghost: Final = candidate.request(
"POST",
"/v1/chat/completions",
{"model": f"ghost-{uuid.uuid4().hex}", "messages": [{"role": "user", "content": "hi"}]},
key=key,
)
assert ghost.status_code in (400, 403, 404), ghost.text
ids: Final = _burst(candidate, model, key, marker)
eventually(
lambda: frozenset(payload["id"] for payload in sink.payloads()),
lambda landed: ids <= landed,
seconds=90,
)
payloads: Final = sink.payloads()
assert sum(1 for r in provider.drain() if r.method == "POST") == REQUESTS
assert ids <= frozenset(payload["id"] for payload in payloads)
extras: Final = tuple(payload for payload in payloads if payload["id"] not in ids)
assert all(payload["status"] == "failure" for payload in extras), extras
@pytest.mark.covers("other.observability.s3_v2.batch_flag_ignored_when_s3_v2_is_cold_storage_logger")
def test_s3_v2_batch_flag_ignored_when_s3_v2_is_cold_storage_logger(gateway: Gateway, tmp_path: Path) -> None:
marker: Final = "s3cold" + uuid.uuid4().hex[:8]
sink: Final = RecordingS3Sink()
with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket:
config: Final = _recording_s3_config(
tmp_path,
bucket.url,
{"s3_batch_file_upload": True},
{"cold_storage_custom_logger": "s3_v2"},
)
with (
owned_proxy_process(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "3"}, config=config) as owned,
owned.gateway.scenario() as scenario,
):
model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key")
key: Final = scenario.key(models=[model])
response: Final = owned.gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": marker}], "cache": {"no-cache": True}},
key=key,
)
assert response.status_code == 200, response.text
request_id: Final = str(response.json()["id"])
payloads: Final = collect_payloads(sink, 1)
assert all(PER_REQUEST_KEY.match(target) for target in sink.objects()), list(sink.objects())
eventually(
lambda: owned.log.read_text(),
lambda text: "s3_batch_file_upload is ignored because s3_v2 is the cold storage logger" in text,
seconds=15,
)
spend: Final = eventually(
lambda: owned.gateway.request("GET", f"/spend/logs/ui/{request_id}"),
lambda reply: reply.status_code == 200 and bool((reply.json() or {}).get("messages")),
seconds=60,
)
assert spend.status_code == 200, spend.text
body: Final = spend.json()
assert body["messages"], spend.text
assert body["response"], spend.text
assert payloads[0]["id"] == request_id
@pytest.mark.covers("other.observability.s3_v2.identical_requests_land_distinct_objects")
def test_s3_v2_identical_requests_land_distinct_objects(gateway: Gateway, tmp_path: Path) -> None:
marker: Final = "s3same" + uuid.uuid4().hex[:8]
sink: Final = RecordingS3Sink()
with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket:
config: Final = _s3_config(tmp_path, bucket.url, {})
with (
owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "3"}, config=config) as candidate,
candidate.scenario() as scenario,
):
model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key")
key: Final = scenario.key(models=[model])
def send(_: int) -> str:
response: Final = candidate.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": marker}], "cache": {"no-cache": True}},
key=key,
)
assert response.status_code == 200, response.text
return str(response.json()["id"])
with ThreadPoolExecutor(max_workers=16) as pool:
returned: Final = frozenset(pool.map(send, range(16)))
payloads: Final = collect_payloads(sink, 16)
assert sum(1 for r in provider.drain() if r.method == "POST") == 16
assert returned == {marker}, "the upstream echo keeps the same id for identical requests"
assert len(sink.objects()) == 16, "identical requests must still land as distinct objects"
assert all(payload["id"] == marker for payload in payloads)
@pytest.mark.covers("other.observability.s3_v2.two_workers_bound_and_deliver_every_id")
def test_s3_v2_two_workers_bound_and_deliver_every_id(gateway: Gateway, tmp_path: Path) -> None:
marker: Final = "s3work" + uuid.uuid4().hex[:8]
sink: Final = RecordingS3Sink()
with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket:
config: Final = _s3_config(tmp_path, bucket.url, {})
with (
owned_proxy(
gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "3"}, config=config, workers=2
) as candidate,
candidate.scenario() as scenario,
):
model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key")
key: Final = scenario.key(models=[model])
ids: Final = _burst(candidate, model, key, marker)
payloads: Final = collect_payloads(sink, REQUESTS)
assert sum(1 for r in provider.drain() if r.method == "POST") == REQUESTS
assert sink.peak <= 32, f"peak concurrent PUTs {sink.peak} exceeded two workers at the default bound"
assert len(sink.objects()) == REQUESTS
assert frozenset(payload["id"] for payload in payloads) == ids
@pytest.mark.covers("other.observability.s3_v2.slow_sink_never_duplicates_or_stalls_readiness")
def test_s3_v2_slow_sink_never_duplicates_or_stalls_readiness(gateway: Gateway, tmp_path: Path) -> None:
marker: Final = "s3slow" + uuid.uuid4().hex[:8]
sink: Final = RecordingS3Sink(delay_seconds=1.5)
with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket:
config: Final = _s3_config(tmp_path, bucket.url, {"s3_batch_file_upload": True})
with (
owned_proxy(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "1"}, config=config) as candidate,
candidate.scenario() as scenario,
):
model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key")
key: Final = scenario.key(models=[model])
ids: Final = _burst(candidate, model, key, marker)
def delivered() -> int:
readiness: Final = candidate.client.get("/health/readiness")
assert readiness.status_code == 200, readiness.text
return sum(len(body.splitlines()) for body in sink.objects().values())
eventually(delivered, lambda total: total >= REQUESTS, seconds=90)
payloads: Final = sink.payloads()
puts: Final = bucket.drain()
targets: Final = tuple(put.target for put in puts)
assert sum(1 for r in provider.drain() if r.method == "POST") == REQUESTS
assert len(set(targets)) == len(targets), "the same object was PUT more than once"
assert frozenset(payload["id"] for payload in payloads) == ids
assert len(payloads) == REQUESTS
@pytest.mark.covers("other.observability.s3_v2.worker_kill_mid_burst_keeps_surviving_deliveries")
def test_s3_v2_worker_kill_mid_burst_keeps_surviving_deliveries(gateway: Gateway, tmp_path: Path) -> None:
marker: Final = "s3kill" + uuid.uuid4().hex[:8]
sink: Final = RecordingS3Sink()
with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket:
config: Final = _s3_config(tmp_path, bucket.url, {})
with (
owned_proxy_process(
gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "3"}, config=config, workers=2
) as owned,
owned.gateway.scenario() as scenario,
):
model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key")
key: Final = scenario.key(models=[model])
sent: Final = tuple(f"{marker}-{index}" for index in range(REQUESTS))
def send(identity: str) -> tuple[str, bool]:
try:
response: Final = owned.gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": identity}],
"cache": {"no-cache": True},
},
key=key,
)
except Exception:
return identity, False
return identity, response.status_code == 200
with ThreadPoolExecutor(max_workers=32) as pool:
futures: Final = tuple(pool.submit(send, identity) for identity in sent)
time.sleep(0.5)
children: Final = tuple(
process for process in group_members(owned.process.pid) if process.pid != owned.process.pid
)
assert children, "no worker children found to kill"
children[0].kill()
results: Final = tuple(future.result() for future in futures)
survivors: Final = frozenset(identity for identity, ok in results if ok)
assert survivors, "no request survived the worker kill"
readiness: Final = owned.gateway.client.get("/health/readiness")
assert readiness.status_code == 200, readiness.text
payloads: Final = collect_payloads(sink, len(survivors), seconds=90)
landed: Final = frozenset(payload["id"] for payload in payloads)
assert survivors <= landed, "an id whose response succeeded never landed"
assert landed <= frozenset(sent), "an id that was never sent landed"
@pytest.mark.covers("other.observability.s3_v2.sigterm_mid_burst_loses_only_inflight_without_duplicates")
def test_s3_v2_sigterm_mid_burst_loses_only_inflight_without_duplicates(gateway: Gateway, tmp_path: Path) -> None:
marker: Final = "s3term" + uuid.uuid4().hex[:8]
sink: Final = RecordingS3Sink()
with wire_server(_chat_reply) as provider, wire_server(sink.respond) as bucket:
config: Final = _s3_config(tmp_path, bucket.url, {})
owned: Final = owned_proxy_process(gateway, tmp_path, {"DEFAULT_S3_FLUSH_INTERVAL_SECONDS": "3"}, config=config)
candidate_owned: Final = owned.__enter__()
try:
created: Final = candidate_owned.gateway.post(
"/model/new",
{
"model_name": f"integration-{marker}",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"api_key": "synthetic-provider-key",
"api_base": provider.url + "/v1",
},
"model_info": {},
},
)
model: Final = str(created["model_name"])
key: Final = str(candidate_owned.gateway.post("/key/generate", {"models": [model]})["key"])
sent: Final = tuple(f"{marker}-{index}" for index in range(REQUESTS))
def send(identity: str) -> tuple[str, bool]:
try:
response: Final = candidate_owned.gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": identity}],
"cache": {"no-cache": True},
},
key=key,
)
except Exception:
return identity, False
return identity, response.status_code == 200
with ThreadPoolExecutor(max_workers=32) as pool:
futures: Final = tuple(pool.submit(send, identity) for identity in sent)
time.sleep(0.5)
candidate_owned.process.terminate()
results: Final = tuple(future.result() for future in futures)
candidate_owned.process.wait(timeout=30)
finally:
owned.__exit__(None, None, None)
answered: Final = frozenset(identity for identity, ok in results if ok)
landed: Final = frozenset(payload["id"] for payload in sink.payloads())
assert landed <= answered, (
"a delivered object has no matching answered request; lost in-flight ids are expected, extras are not"
)
targets: Final = tuple(sink.objects())
assert len(set(targets)) == len(targets), "the same object was PUT more than once"

View file

@ -0,0 +1,244 @@
import json
import uuid
from collections.abc import Mapping
from pathlib import Path
from typing import Final
import pytest
from integration._support.client import Gateway, object_value
from integration._support.process import owned_proxy
from pydantic import JsonValue
SIBLING_LIMITS: Final = {"max_input_tokens": 4321, "max_output_tokens": 987}
NON_NUMERIC_LIMITS: Final = (
pytest.param("", id="empty-string"),
pytest.param(" ", id="blank-string"),
pytest.param("128,000", id="thousands-separator"),
pytest.param("unlimited", id="word"),
pytest.param("NaN", id="nan-string"),
pytest.param("inf", id="inf-string"),
pytest.param([], id="empty-list"),
pytest.param([4096], id="list"),
pytest.param({}, id="empty-object"),
pytest.param({"tokens": 4096}, id="object"),
pytest.param(True, id="bool"),
pytest.param(None, id="null"),
)
NUMERIC_EDGE_LIMITS: Final = (
pytest.param(0, id="zero"),
pytest.param(-1, id="negative"),
pytest.param(1.5, id="float"),
pytest.param("1.5", id="float-string"),
pytest.param("1e9", id="exponent-string"),
pytest.param(10**12, id="huge"),
)
NUMERIC_EDGE_EXPECTED: Final = {
"zero": 0,
"negative": -1,
"float": 1,
"float-string": 1,
"exponent-string": 1_000_000_000,
"huge": 10**12,
}
MODEL_GROUP_INFO_500: Final = (
"BUG: /model_group/info returns 500 for every caller when one deployment's token limit is non-numeric"
)
CHAT_500: Final = (
"BUG: chat completions return 500 from ModelGroupInfo validation when the deployment's token limit is non-numeric"
)
MODEL_GROUP_INFO_500_IDS: Final = frozenset(
{"empty-string", "blank-string", "thousands-separator", "word", "nan-string", "inf-string"}
| {"empty-list", "list", "empty-object", "object"}
)
CHAT_500_IDS: Final = frozenset(
{"empty-string", "blank-string", "thousands-separator", "word", "empty-list", "list", "empty-object", "object"}
)
def _listed(gateway: Gateway, path: str) -> dict[str, dict[str, JsonValue]]:
entries: Final = gateway.get(path)["data"]
assert isinstance(entries, list)
return {str(object_value(entry)["id"]): object_value(entry) for entry in entries}
def _limits(entry: Mapping[str, JsonValue]) -> tuple[JsonValue, JsonValue]:
return entry.get("max_input_tokens"), entry.get("max_output_tokens")
def _assert_listing_spares_the_sibling(
gateway: Gateway, broken: str, sibling: str, broken_limits: tuple[JsonValue, JsonValue]
) -> None:
for path in ("/v1/models", "/models"):
listed: Final = _listed(gateway, path)
assert _limits(listed[sibling]) == (4321, 987), (path, listed[sibling])
assert _limits(listed[broken]) == broken_limits, (path, listed[broken])
single: Final = gateway.get(f"/v1/models/{broken}")
assert single["id"] == broken, single
assert _limits(single) == broken_limits, single
registered: Final = gateway.get("/model/info")["data"]
assert isinstance(registered, list)
assert {broken, sibling} <= {str(object_value(entry)["model_name"]) for entry in registered}
def _assert_serves_chat(gateway: Gateway, *models: str) -> None:
for model in models:
reply: Final = gateway.chat(model, text=f"token limit edge {uuid.uuid4().hex}")
assert reply["model"] == model, reply
def _listed_model(gateway: Gateway, model: str) -> dict[str, JsonValue]:
entries: Final = gateway.get("/v1/models")["data"]
assert isinstance(entries, list)
return next(object_value(entry) for entry in entries if object_value(entry)["id"] == model)
def test_v1_models_carries_cost_map_context_window_for_a_known_model(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(model="openai/gpt-4o-mini")
listed: Final = _listed_model(gateway, model)
# OpenAI publishes these for gpt-4o-mini: https://platform.openai.com/docs/models/gpt-4o-mini (checked 2026-09-24)
assert listed["max_input_tokens"] == 128000, listed
assert listed["max_output_tokens"] == 16384, listed
def test_v1_models_carries_deployment_model_info_limits_for_an_unknown_model(gateway: Gateway) -> None:
unknown: Final = f"openai/custom-{uuid.uuid4().hex}"
with gateway.scenario() as scenario:
model: Final = scenario.model(model=unknown, model_info={"max_input_tokens": 4321, "max_output_tokens": 987})
listed: Final = _listed_model(gateway, model)
assert listed["max_input_tokens"] == 4321, listed
assert listed["max_output_tokens"] == 987, listed
def test_numeric_string_token_limit_is_coerced_to_an_int(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model(
model=f"openai/custom-{uuid.uuid4().hex}",
model_info={"max_input_tokens": "4096", "max_output_tokens": "512"},
)
assert _limits(_listed_model(gateway, model)) == (4096, 512)
@pytest.mark.parametrize("value", NON_NUMERIC_LIMITS)
def test_non_numeric_token_limit_is_listed_as_absent_without_breaking_the_listing(
gateway: Gateway, value: JsonValue
) -> None:
with gateway.scenario() as scenario:
sibling: Final = scenario.model(model=f"openai/custom-{uuid.uuid4().hex}", model_info=SIBLING_LIMITS)
broken: Final = scenario.model(
model=f"openai/custom-{uuid.uuid4().hex}",
model_info={"max_input_tokens": value, "max_output_tokens": value},
)
_assert_listing_spares_the_sibling(gateway, broken, sibling, (None, None))
@pytest.mark.parametrize("value", NON_NUMERIC_LIMITS)
def test_non_numeric_token_limit_still_serves_chat(
gateway: Gateway, value: JsonValue, request: pytest.FixtureRequest
) -> None:
if request.node.callspec.id in CHAT_500_IDS:
pytest.skip(CHAT_500)
with gateway.scenario() as scenario:
sibling: Final = scenario.model(model=f"openai/custom-{uuid.uuid4().hex}", model_info=SIBLING_LIMITS)
broken: Final = scenario.model(
model=f"openai/custom-{uuid.uuid4().hex}",
model_info={"max_input_tokens": value, "max_output_tokens": value},
)
_assert_serves_chat(gateway, broken, sibling)
@pytest.mark.parametrize("value", NUMERIC_EDGE_LIMITS)
def test_numeric_edge_token_limit_is_listed_as_its_integer_without_breaking_the_listing(
gateway: Gateway, value: JsonValue, request: pytest.FixtureRequest
) -> None:
expected: Final = NUMERIC_EDGE_EXPECTED[request.node.callspec.id]
with gateway.scenario() as scenario:
sibling: Final = scenario.model(model=f"openai/custom-{uuid.uuid4().hex}", model_info=SIBLING_LIMITS)
broken: Final = scenario.model(
model=f"openai/custom-{uuid.uuid4().hex}",
model_info={"max_input_tokens": value, "max_output_tokens": value},
)
_assert_listing_spares_the_sibling(gateway, broken, sibling, (expected, expected))
_assert_serves_chat(gateway, broken, sibling)
@pytest.mark.parametrize("field", ("max_input_tokens", "max_output_tokens"))
def test_one_malformed_limit_does_not_disturb_the_other(gateway: Gateway, field: str) -> None:
other: Final = "max_output_tokens" if field == "max_input_tokens" else "max_input_tokens"
with gateway.scenario() as scenario:
model: Final = scenario.model(
model=f"openai/custom-{uuid.uuid4().hex}", model_info={field: "128,000", other: 2048}
)
listed: Final = _listed_model(gateway, model)
assert listed.get(field) is None, listed
assert listed[other] == 2048, listed
@pytest.mark.parametrize("value", NON_NUMERIC_LIMITS + NUMERIC_EDGE_LIMITS)
def test_malformed_token_limit_keeps_model_group_info_serving(
gateway: Gateway, value: JsonValue, request: pytest.FixtureRequest
) -> None:
if request.node.callspec.id in MODEL_GROUP_INFO_500_IDS:
pytest.skip(MODEL_GROUP_INFO_500)
with gateway.scenario() as scenario:
sibling: Final = scenario.model(model=f"openai/custom-{uuid.uuid4().hex}", model_info=SIBLING_LIMITS)
broken: Final = scenario.model(
model=f"openai/custom-{uuid.uuid4().hex}",
model_info={"max_input_tokens": value, "max_output_tokens": value},
)
groups: Final = gateway.get("/model_group/info")["data"]
assert isinstance(groups, list)
assert {broken, sibling} <= {str(object_value(group)["model_group"]) for group in groups}
single: Final = gateway.get("/model_group/info", {"model_group": broken})["data"]
assert isinstance(single, list)
assert [object_value(group)["model_group"] for group in single] == [broken]
def _yaml_deployment(name: str, upstream_url: str, model_info: Mapping[str, JsonValue]) -> dict[str, JsonValue]:
return {
"model_name": name,
"litellm_params": {
"model": f"openai/custom-{uuid.uuid4().hex}",
"api_base": f"{upstream_url}/v1",
"api_key": "integration-provider-key",
},
"model_info": dict(model_info),
}
def test_non_numeric_token_limits_in_config_yaml_are_listed_as_absent(gateway: Gateway, tmp_path: Path) -> None:
run: Final = uuid.uuid4().hex
sibling: Final = f"integration-yaml-sibling-{run}"
broken: Final = {f"integration-yaml-{parameter.id}-{run}": parameter.values[0] for parameter in NON_NUMERIC_LIMITS}
serving: Final = tuple(
f"integration-yaml-{parameter.id}-{run}" for parameter in NON_NUMERIC_LIMITS if parameter.id not in CHAT_500_IDS
)
config: Final = tmp_path / "malformed_token_limits.yaml"
config.write_text(
json.dumps(
{
"model_list": [
_yaml_deployment(sibling, gateway.upstream_url, SIBLING_LIMITS),
*(
_yaml_deployment(
name, gateway.upstream_url, {"max_input_tokens": value, "max_output_tokens": value}
)
for name, value in broken.items()
),
],
"general_settings": {
"master_key": "os.environ/LITELLM_MASTER_KEY",
"database_url": "os.environ/DATABASE_URL",
"store_model_in_db": True,
},
"router_settings": {"disable_cooldowns": True},
}
)
)
with owned_proxy(gateway, tmp_path, {}, config=config) as candidate:
for name in broken:
_assert_listing_spares_the_sibling(candidate, name, sibling, (None, None))
_assert_serves_chat(candidate, sibling, *serving)

View file

@ -52,9 +52,7 @@ def test_chat_completions_system_block_list_carries_cache_control_to_anthropic_s
"messages": [
{
"role": "system",
"content": [
{"type": "text", "text": policy, "cache_control": {"type": "ephemeral"}}
],
"content": [{"type": "text", "text": policy, "cache_control": {"type": "ephemeral"}}],
},
{"role": "user", "content": "hi"},
],
@ -112,9 +110,7 @@ def test_responses_system_input_item_carries_cache_control_to_anthropic_system(g
"input": [
{
"role": "system",
"content": [
{"type": "input_text", "text": policy, "cache_control": {"type": "ephemeral"}}
],
"content": [{"type": "input_text", "text": policy, "cache_control": {"type": "ephemeral"}}],
},
{"role": "user", "content": "hi"},
],
@ -126,3 +122,67 @@ def test_responses_system_input_item_carries_cache_control_to_anthropic_system(g
assert any(item.get("type") == "message" for item in payload.get("output", []) if isinstance(item, dict))
assert len(wire.drain()) == 1
def _anthropic_usage_reply(identity: str, cache_creation: int, cache_read: int) -> bytes:
return json.dumps(
{
"id": identity,
"type": "message",
"role": "assistant",
"model": _MODEL,
"content": [{"type": "text", "text": "done"}],
"stop_reason": "end_turn",
"stop_sequence": None,
"usage": {
"input_tokens": 3,
"output_tokens": 1,
"cache_creation_input_tokens": cache_creation,
"cache_read_input_tokens": cache_read,
},
}
).encode()
def test_responses_usage_reports_anthropic_system_cache_write_then_read(gateway: Gateway) -> None:
identity: Final = f"responses-system-cache-usage-{uuid.uuid4().hex}"
policy: Final = f"policy {identity}"
replies: Final = iter(
(
_anthropic_usage_reply(identity, cache_creation=1200, cache_read=0),
_anthropic_usage_reply(identity, cache_creation=0, cache_read=1200),
)
)
def respond(request: Request) -> Reply:
assert request.method == "POST" and request.target == "/v1/messages"
_assert_system_block(_JSON_OBJECT.validate_json(request.body), policy)
return Reply(body=next(replies))
def input_tokens_details(model: str, user_turn: str) -> JsonValue:
response: Final = gateway.request(
"POST",
"/v1/responses",
{
"model": model,
"input": [
{
"role": "system",
"content": [{"type": "input_text", "text": policy, "cache_control": {"type": "ephemeral"}}],
},
{"role": "user", "content": user_turn},
],
},
)
assert response.status_code == 200, response.text
usage: Final = _JSON_OBJECT.validate_json(response.content)["usage"]
assert isinstance(usage, dict), response.text
return usage["input_tokens_details"]
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_API_KEY)
first: Final = input_tokens_details(model, "first turn")
second: Final = input_tokens_details(model, "second turn")
assert len(wire.drain()) == 2
assert isinstance(first, dict) and isinstance(second, dict), (first, second)
assert (first["cache_write_tokens"], first["cached_tokens"]) == (1200, 0), first
assert (second.get("cache_write_tokens", 0), second["cached_tokens"]) == (0, 1200), second

View file

@ -1,16 +1,70 @@
import json
import uuid
from pathlib import Path
from typing import Final
import pytest
from integration._support.client import Gateway
from integration._support.client import Gateway, eventually
from integration._support.database import read_rows
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue, TypeAdapter
_ROUTER_SLUG: Final = "routers/glm-latest"
_ROUTER_RESOURCE: Final = "accounts/fireworks/routers/glm-latest"
_FIREROUTER_SLUGS: Final = ("firerouter", "firerouter/kimi-k3/deepseek-v4")
_API_KEY: Final = "synthetic-fireworks-key"
_PROMPT: Final = "route me through the router"
_COST_MAP_PATH: Final = Path(__file__).resolve().parents[3] / "model_prices_and_context_window.json"
_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue])
_COST_MAP: Final = TypeAdapter(dict[str, dict[str, object]])
def _positive_rate(entry: dict[str, object], field: str) -> bool:
value: Final = entry.get(field)
return isinstance(value, (int, float)) and value > 0
def _pick_routed_model() -> str:
catalog: Final = _COST_MAP.validate_json(_COST_MAP_PATH.read_bytes())
return next(
key
for key, entry in catalog.items()
if "/" not in key
and entry.get("litellm_provider") == "anthropic"
and _positive_rate(entry, "input_cost_per_token")
and _positive_rate(entry, "output_cost_per_token")
and f"fireworks_ai/{key}" not in catalog
)
def _catalog_cost(model: str, field: str) -> float:
cost_value: Final = _COST_MAP.validate_json(_COST_MAP_PATH.read_bytes())[model][field]
assert isinstance(cost_value, (int, float))
return float(cost_value)
_ROUTED_MODEL: Final = _pick_routed_model()
def _approx(value: float) -> object:
return pytest.approx(value, rel=1e-6) # pyright: ignore[reportUnknownMemberType] # pytest lacks typed approx stubs
def _chat_completion(identity: str, model: str, prompt_tokens: int, completion_tokens: int) -> bytes:
return json.dumps(
{
"id": identity,
"object": "chat.completion",
"created": 1,
"model": model,
"choices": [{"index": 0, "message": {"role": "assistant", "content": "routed"}, "finish_reason": "stop"}],
"usage": {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens,
},
}
).encode()
def _provider_body(request: Request, target: str) -> dict[str, JsonValue]:
@ -82,3 +136,64 @@ def test_fireworks_router_slug_text_completion_sends_router_resource_not_models_
payload: Final = _JSON_OBJECT.validate_json(response.content)
assert payload["choices"] == [{"index": 0, "text": "routed", "finish_reason": "stop", "logprobs": None}]
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/completions")]
@pytest.mark.parametrize("slug", _FIREROUTER_SLUGS)
def test_fireworks_firerouter_short_name_sends_router_resource_not_models_path(gateway: Gateway, slug: str) -> None:
resource: Final = f"accounts/fireworks/routers/{slug}"
def respond(request: Request) -> Reply:
body: Final = _provider_body(request, "/chat/completions")
assert body["model"] == resource, body
return Reply(body=_chat_completion(f"fw-{slug}", resource, 5, 1))
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=f"fireworks_ai/{slug}", api_base=wire.url, api_key=_API_KEY)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": _PROMPT}]},
)
assert response.status_code == 200, response.text
payload: Final = _JSON_OBJECT.validate_json(response.content)
assert payload["choices"] == [
{"finish_reason": "stop", "index": 0, "message": {"role": "assistant", "content": "routed"}}
]
assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")]
def test_fireworks_firerouter_claude_leg_is_charged_at_the_routed_models_own_rate(gateway: Gateway) -> None:
identity: Final = f"fw-firerouter-claude-{uuid.uuid4().hex}"
def respond(request: Request) -> Reply:
body: Final = _provider_body(request, "/chat/completions")
assert body["model"] == "accounts/fireworks/routers/firerouter", body
assert request.headers["x-anthropic-api-key"] == "synthetic-anthropic-key"
return Reply(body=_chat_completion(identity, _ROUTED_MODEL, 23, 41))
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model="fireworks_ai/firerouter",
api_base=wire.url,
api_key=_API_KEY,
extra_headers={"x-anthropic-api-key": "synthetic-anthropic-key"},
)
response: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": _PROMPT}]},
)
assert response.status_code == 200, response.text
expected_cost: Final = 23 * _catalog_cost(_ROUTED_MODEL, "input_cost_per_token") + 41 * _catalog_cost(
_ROUTED_MODEL, "output_cost_per_token"
)
assert expected_cost > 0
assert float(response.headers["x-litellm-response-cost"]) == _approx(expected_cost)
rows: Final = eventually(
lambda: read_rows('SELECT spend FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (identity,)),
lambda values: len(values) == 1,
seconds=70,
)
spend: Final = rows[0]["spend"]
assert isinstance(spend, (int, float, str))
assert float(spend) == _approx(expected_cost)

View file

@ -0,0 +1,3 @@
{
"SELECT $n": "SELECT 1 health probe: the database watchdog and probe target query whichever pool reader_unavailable selects and the reconnect smoke test always uses the writer, all timer driven, so it lands under whichever test is in flight"
}

View file

@ -26,7 +26,7 @@ def test_owned_redis_outage_recovers_requests_and_real_response_cache(gateway: G
try:
with owned_redis(tmp_path) as cache, monkeypatch.context() as environment:
environment.setenv("DATABASE_URL", database_url)
with owned_proxy(gateway, tmp_path, {"DATABASE_URL": database_url, "REDIS_HOST": cache.host, "REDIS_PORT": str(cache.port), "REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT": "1"}) as candidate, candidate.scenario() as scenario, httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream:
with owned_proxy(gateway, tmp_path, {"DATABASE_URL": database_url, "REDIS_HOST": cache.host, "REDIS_PORT": str(cache.port), "REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT": "1"}, remove_environment=("DATABASE_URL_READ_REPLICA",)) as candidate, candidate.scenario() as scenario, httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream:
model: Final = scenario.model()
key: Final = scenario.key(models=[model])
for generation in ("before", "after"):

View file

@ -26,6 +26,7 @@ def _install_daily_user_rollup_fault(user_id: str) -> str:
_execute(
(
sql.SQL("CREATE SEQUENCE {}").format(sequence),
sql.SQL("GRANT USAGE ON SEQUENCE {} TO PUBLIC").format(sequence),
sql.SQL(
"CREATE FUNCTION {}() RETURNS trigger LANGUAGE plpgsql AS $fault$ "
"BEGIN PERFORM nextval({}); "

View file

@ -0,0 +1,463 @@
import asyncio
import json
import os
import signal
import threading
import time
import uuid
from collections.abc import Callable, Iterator, Mapping
from concurrent.futures import FIRST_COMPLETED, ThreadPoolExecutor, wait
from contextlib import contextmanager
from dataclasses import dataclass
from hashlib import sha256
from pathlib import Path
from typing import Final
import anthropic
import httpx
import openai
import psutil
import pytest
from integration._support.client import Gateway, eventually, object_value, string_value
from integration._support.database import read_rows
from integration._support.process import OwnedProxy, owned_proxy_process
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue
CRAFTED_MODEL: Final = ("exceeded " * 32_000)[:288_000]
HOSTILE_5KB_MODEL: Final = ("exceeded budget " * 400)[:5_000]
FAST_SECONDS: Final = 10.0
LIVELINESS_MAX_SECONDS: Final = 5.0
ROW_SECONDS: Final = 70
CHAT: Final = "/v1/chat/completions"
MESSAGES: Final = "/v1/messages"
RESPONSES: Final = "/v1/responses"
def _body(path: str, model: str, marker: str, stream: bool = False) -> dict[str, JsonValue]:
content: Final = f"normalized error audit {marker}"
match path:
case "/v1/messages":
return {
"model": model,
"max_tokens": 8,
"messages": [{"role": "user", "content": content}],
"stream": stream,
}
case "/v1/responses":
return {"model": model, "input": content, "stream": stream}
case _:
return {"model": model, "messages": [{"role": "user", "content": content}], "stream": stream}
@dataclass(frozen=True, slots=True)
class _Timed:
response: httpx.Response
seconds: float
def _timed_post(client: httpx.Client, path: str, body: Mapping[str, JsonValue], key: str) -> _Timed:
started: Final = time.perf_counter()
response: Final = client.post(path, json=body, headers={"Authorization": f"Bearer {key}"})
return _Timed(response, time.perf_counter() - started)
@contextmanager
def _patient_client(gateway: Gateway) -> Iterator[httpx.Client]:
with httpx.Client(base_url=str(gateway.client.base_url), timeout=120, trust_env=False) as client:
yield client
def _error_information(call_id: str) -> dict[str, JsonValue]:
rows: Final = eventually(
lambda: read_rows(
"SELECT status, metadata->'error_information' AS info FROM \"LiteLLM_SpendLogs\" WHERE request_id=%s",
(call_id,),
),
lambda values: len(values) == 1,
seconds=ROW_SECONDS,
)
assert rows[0]["status"] == "failure", rows
return object_value(rows[0]["info"])
def _assert_crafted_failure(timed: _Timed, expected_status: int = 400) -> None:
response: Final = timed.response
assert response.status_code == expected_status, response.text[:300]
assert "Invalid model name passed in" in response.text, response.text[:300]
assert timed.seconds < FAST_SECONDS, f"crafted 288 KB model took {timed.seconds:.2f}s"
info: Final = _error_information(response.headers["x-litellm-call-id"])
assert info["normalized_error"] == "400_INVALID_REQUEST" and info["error_code"] == "400", info
@pytest.mark.parametrize("path", [CHAT, MESSAGES, RESPONSES])
def test_crafted_288kb_model_fails_fast_and_logs_invalid_request(gateway: Gateway, path: str) -> None:
with gateway.scenario() as scenario, _patient_client(gateway) as client:
key: Final = scenario.key()
_assert_crafted_failure(_timed_post(client, path, _body(path, CRAFTED_MODEL, uuid.uuid4().hex), key))
@pytest.mark.parametrize("path", [CHAT, MESSAGES, RESPONSES])
def test_crafted_288kb_model_with_stream_true_fails_fast(gateway: Gateway, path: str) -> None:
with gateway.scenario() as scenario, _patient_client(gateway) as client:
key: Final = scenario.key()
body: Final = _body(path, CRAFTED_MODEL, uuid.uuid4().hex, stream=True)
_assert_crafted_failure(_timed_post(client, path, body, key))
def test_crafted_288kb_model_through_async_openai_sdk_fails_fast(gateway: Gateway) -> None:
async def call(key: str) -> tuple[openai.BadRequestError, float]:
client: Final = openai.AsyncOpenAI(base_url=f"{gateway.client.base_url}/v1", api_key=key, timeout=120)
started: Final = time.perf_counter()
try:
with pytest.raises(openai.BadRequestError) as raised:
await client.chat.completions.create(
model=CRAFTED_MODEL, messages=[{"role": "user", "content": f"audit {uuid.uuid4().hex}"}]
)
return raised.value, time.perf_counter() - started
finally:
await client.close()
with gateway.scenario() as scenario:
error, seconds = asyncio.run(call(scenario.key()))
assert seconds < FAST_SECONDS, f"crafted 288 KB model took {seconds:.2f}s"
assert "Invalid model name passed in" in str(error), str(error)[:300]
info: Final = _error_information(error.response.headers["x-litellm-call-id"])
assert info["normalized_error"] == "400_INVALID_REQUEST", info
def test_crafted_288kb_model_through_anthropic_sdk_fails_fast(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
client: Final = anthropic.Anthropic(base_url=str(gateway.client.base_url), api_key=scenario.key(), timeout=120)
started: Final = time.perf_counter()
with pytest.raises(anthropic.BadRequestError) as raised:
client.messages.create(
model=CRAFTED_MODEL, max_tokens=8, messages=[{"role": "user", "content": f"audit {uuid.uuid4().hex}"}]
)
seconds: Final = time.perf_counter() - started
assert seconds < FAST_SECONDS, f"crafted 288 KB model took {seconds:.2f}s"
assert "Invalid model name passed in" in str(raised.value), str(raised.value)[:300]
info: Final = _error_information(raised.value.response.headers["x-litellm-call-id"])
assert info["normalized_error"] == "400_INVALID_REQUEST", info
def _poll_liveliness(client: httpx.Client, stop: threading.Event) -> list[float]:
latencies: Final[list[float]] = [] # mutable-ok: thread-local sample buffer drained once by the caller
while not stop.is_set():
started = time.perf_counter()
assert client.get("/health/liveliness").status_code == 200
latencies.append(time.perf_counter() - started)
stop.wait(0.1)
return latencies
def _cmdline(process: psutil.Process) -> str:
try:
return " ".join(process.cmdline())
except psutil.Error:
return ""
def _worker_pids(owned: OwnedProxy) -> tuple[int, ...]:
return tuple(child.pid for child in psutil.Process(owned.process.pid).children() if "spawn_main" in _cmdline(child))
def test_two_concurrent_crafted_requests_do_not_stall_liveliness_on_a_two_worker_proxy(
gateway: Gateway, tmp_path: Path
) -> None:
with owned_proxy_process(gateway, tmp_path, {}, workers=2) as owned, _patient_client(owned.gateway) as client:
assert len(_worker_pids(owned)) == 2, _worker_pids(owned)
stop: Final = threading.Event()
with ThreadPoolExecutor(max_workers=3) as pool:
liveliness: Final = pool.submit(_poll_liveliness, client, stop)
crafted: Final = tuple(
pool.submit(_timed_post, client, CHAT, _body(CHAT, CRAFTED_MODEL, uuid.uuid4().hex), gateway.key)
for _ in range(2)
)
results: Final = tuple(future.result() for future in crafted)
stop.set()
latencies: Final = liveliness.result()
for timed in results:
_assert_crafted_failure(timed)
assert latencies and max(latencies) < LIVELINESS_MAX_SECONDS, f"liveliness max {max(latencies):.2f}s"
def _completion(request: Request) -> Reply:
body: Final = object_value(json.loads(request.body or b"{}"))
return Reply(
body=json.dumps(
{
"id": "chatcmpl-" + uuid.uuid4().hex,
"object": "chat.completion",
"created": 1,
"model": body.get("model", "unknown"),
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 20, "completion_tokens": 20, "total_tokens": 40},
}
).encode()
)
def _rate_limited(message: str) -> Callable[[Request], Reply]:
def respond(_request: Request) -> Reply:
return Reply(
status=429,
body=json.dumps({"error": {"message": message, "type": "rate_limit_error", "code": "429"}}).encode(),
)
return respond
def _budget_denied_row(key: str) -> dict[str, JsonValue]:
rows: Final = eventually(
lambda: read_rows(
"SELECT request_id, metadata->'error_information' AS info FROM \"LiteLLM_SpendLogs\" "
"WHERE api_key=%s AND status='failure'",
(sha256(key.encode()).hexdigest(),),
),
lambda values: len(values) == 1,
seconds=ROW_SECONDS,
)
return object_value(rows[0]["info"])
def _exhaust(
gateway: Gateway, client: httpx.Client, model: str, key: str, table: str, column: str, identity: str
) -> None:
first: Final = gateway.chat(model, key=key, text=f"spend {uuid.uuid4().hex}")
assert object_value(first["usage"])["total_tokens"] == 40, first
eventually(
lambda: read_rows(f'SELECT spend FROM "{table}" WHERE {column}=%s', (identity,)),
lambda values: len(values) == 1 and float(string_value(str(values[0]["spend"]))) >= 0.06,
seconds=ROW_SECONDS,
)
denied: Final = eventually(
lambda: client.post(
CHAT, json=_body(CHAT, model, uuid.uuid4().hex), headers={"Authorization": f"Bearer {key}"}
),
lambda response: response.status_code in {400, 422},
seconds=ROW_SECONDS,
)
assert denied.json()["error"]["type"] == "budget_exceeded", denied.text
info: Final = _budget_denied_row(key)
assert info["normalized_error"] == "429_BUDGET_EXCEEDED", info
assert "budget" in string_value(info["error_message"]).lower(), info
def test_exhausted_key_budget_denial_clusters_as_budget_exceeded(gateway: Gateway) -> None:
with gateway.scenario() as scenario, _patient_client(gateway) as client:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
key: Final = scenario.key(models=[model], max_budget=0.06)
_exhaust(gateway, client, model, key, "LiteLLM_VerificationToken", "token", sha256(key.encode()).hexdigest())
def test_exhausted_team_budget_denial_clusters_as_budget_exceeded(gateway: Gateway) -> None:
with gateway.scenario() as scenario, _patient_client(gateway) as client:
model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
team: Final = scenario.team(models=[model], max_budget=0.06)
key: Final = scenario.key(team_id=team, models=[model])
_exhaust(gateway, client, model, key, "LiteLLM_TeamTable", "team_id", team)
def _upstream_failure_row(gateway: Gateway, message: str) -> dict[str, JsonValue]:
with wire_server(_rate_limited(message)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(api_base=wire.url + "/v1", num_retries=0)
key: Final = scenario.key(models=[model])
failed: Final = gateway.request("POST", CHAT, _body(CHAT, model, uuid.uuid4().hex), key=key)
assert failed.status_code == 429 and message in failed.json()["error"]["message"], failed.text[:300]
assert len(wire.drain()) == 1
return _error_information(failed.headers["x-litellm-call-id"])
@pytest.mark.parametrize(
"message",
[
"Budget has been exceeded! Current cost: 11.0, Max budget: 10.0",
"ExceededBudget: User=audit over budget. Spend=12.5, Budget=10.0",
"Exceeded budget for provider openai: 105.2 >= 100.0",
"exceeded" + "x" * 64 + "budget",
],
)
def test_upstream_budget_wording_clusters_as_budget_exceeded(gateway: Gateway, message: str) -> None:
info: Final = _upstream_failure_row(gateway, message)
assert info["normalized_error"] == "429_BUDGET_EXCEEDED", info
def test_upstream_exceeded_and_budget_65_chars_apart_still_clusters_as_budget_exceeded(gateway: Gateway) -> None:
info: Final = _upstream_failure_row(gateway, "exceeded" + "x" * 65 + "budget")
assert info["normalized_error"] == "429_BUDGET_EXCEEDED", info
def test_upstream_exceeded_and_budget_on_different_lines_cluster_by_exception_class(gateway: Gateway) -> None:
info: Final = _upstream_failure_row(gateway, "exceeded the limit\nbudget unaffected")
assert info["normalized_error"] == "429_RATE_LIMIT_EXCEEDED", info
def test_hostile_model_values_are_rejected_without_taking_the_proxy_down(gateway: Gateway) -> None:
with gateway.scenario() as scenario, _patient_client(gateway) as client:
key: Final = scenario.key()
hostile_values: Final[tuple[JsonValue, ...]] = (5, ["gpt-4o-mini"])
for hostile in hostile_values:
rejected = _timed_post(client, CHAT, {"model": hostile, "messages": []}, key)
assert rejected.response.status_code == 400 and "must be a string" in rejected.response.text
assert _error_information(rejected.response.headers["x-litellm-call-id"])["normalized_error"] == (
"400_INVALID_REQUEST"
)
empty: Final = _timed_post(client, CHAT, _body(CHAT, "", uuid.uuid4().hex), key)
assert empty.response.status_code == 400, empty.response.text
assert _error_information(empty.response.headers["x-litellm-call-id"])["normalized_error"] == (
"400_INVALID_REQUEST"
)
repeated: Final = tuple(
_timed_post(client, CHAT, _body(CHAT, HOSTILE_5KB_MODEL, uuid.uuid4().hex), key) for _ in range(2)
)
call_ids: Final = tuple(timed.response.headers["x-litellm-call-id"] for timed in repeated)
assert len(set(call_ids)) == 2 and all(timed.response.status_code == 400 for timed in repeated)
assert all(timed.seconds < FAST_SECONDS for timed in repeated), [timed.seconds for timed in repeated]
codes: Final = tuple(_error_information(call_id)["normalized_error"] for call_id in call_ids)
assert len(set(codes)) == 1 and codes[0] in {"400_INVALID_REQUEST", "429_BUDGET_EXCEEDED"}, codes
unauthenticated: Final = client.post(CHAT, json=_body(CHAT, "gpt-4o-mini", "x"))
assert unauthenticated.status_code == 401, unauthenticated.text
assert client.get("/health/liveliness").status_code == 200
@dataclass(frozen=True, slots=True)
class _BurstResult:
label: str
status: int | None
call_id: str | None
response_id: str | None
seconds: float
def _burst_call(client: httpx.Client, label: str, path: str, body: Mapping[str, JsonValue], key: str) -> _BurstResult:
started: Final = time.perf_counter()
try:
response: Final = client.post(path, json=body, headers={"Authorization": f"Bearer {key}"})
except httpx.TransportError:
return _BurstResult(label, None, None, None, time.perf_counter() - started)
identity: Final = object_value(response.json()).get("id") if response.status_code == 200 else None
return _BurstResult(
label,
response.status_code,
response.headers.get("x-litellm-call-id"),
identity if isinstance(identity, str) else None,
time.perf_counter() - started,
)
def _burst(
client: httpx.Client, happy_model: str, happy_key: str, open_key: str, during: Callable[[], None]
) -> tuple[_BurstResult, ...]:
crafted: Final = tuple(
(f"crafted-{path}-{index}", path, _body(path, CRAFTED_MODEL, uuid.uuid4().hex, stream=index % 2 == 1), open_key)
for path in (CHAT, MESSAGES, RESPONSES)
for index in range(4)
)
happy: Final = tuple(
(f"happy-{index}", CHAT, _body(CHAT, happy_model, f"happy-{index}"), happy_key) for index in range(8)
)
late: Final = tuple(
(f"late-{index}", CHAT, _body(CHAT, happy_model, f"late-{index}"), happy_key) for index in range(8)
)
with (
ThreadPoolExecutor(max_workers=28) as pool,
httpx.Client(base_url=client.base_url, timeout=client.timeout, trust_env=False) as fresh,
):
first: Final = tuple(
pool.submit(_burst_call, client, label, path, body, key) for label, path, body, key in crafted + happy
)
wait(first, return_when=FIRST_COMPLETED)
during()
second: Final = tuple(
pool.submit(_burst_call, fresh, label, path, body, key) for label, path, body, key in late
)
return tuple(future.result() for future in first + second)
def _assert_rows_land_exactly_once(results: tuple[_BurstResult, ...], prefix: str) -> None:
landed: Final = tuple(result for result in results if result.label.startswith(prefix) and result.status == 200)
assert landed, results
response_ids: Final = tuple(string_value(result.response_id) for result in landed)
assert len(set(response_ids)) == len(response_ids), response_ids
rows: Final = eventually(
lambda: read_rows(
'SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE request_id = ANY(%s::text[])',
("{" + ",".join(response_ids) + "}",),
),
lambda values: len(values) == len(response_ids),
seconds=ROW_SECONDS,
)
assert sorted(string_value(row["request_id"]) for row in rows) == sorted(response_ids), rows
assert all(row["status"] == "success" for row in rows), rows
def test_killing_one_worker_mid_burst_leaves_the_other_serving_crafted_and_happy_traffic(
gateway: Gateway, tmp_path: Path
) -> None:
with (
wire_server(_completion) as wire,
owned_proxy_process(gateway, tmp_path, {}, workers=2) as owned,
owned.gateway.scenario() as scenario,
_patient_client(owned.gateway) as client,
):
model: Final = scenario.model(api_base=wire.url + "/v1", num_retries=0)
key: Final = scenario.key(models=[model])
workers: Final = _worker_pids(owned)
assert len(workers) == 2, workers
def kill_one_worker() -> None:
os.kill(workers[0], signal.SIGKILL)
results: Final = _burst(client, model, key, scenario.key(), kill_one_worker)
dropped: Final = tuple(result for result in results if result.status is None)
assert len(dropped) < len(results), results
crafted: Final = tuple(result for result in results if result.label.startswith("crafted") and result.status)
assert crafted and all(result.status == 400 and result.seconds < FAST_SECONDS for result in crafted), crafted
late: Final = tuple(result for result in results if result.label.startswith("late"))
assert all(result.status == 200 for result in late), late
_assert_rows_land_exactly_once(results, "late")
survivor: Final = tuple(pid for pid in _worker_pids(owned) if pid != workers[0])
assert survivor, "no worker left serving"
after: Final = owned.gateway.chat(model, key=key, text=f"after kill {uuid.uuid4().hex}")
assert isinstance(after["id"], str) and after["id"].startswith("chatcmpl-"), after
assert client.get("/health/liveliness").status_code == 200
def test_upstream_returning_503_mid_burst_logs_every_failure_with_its_own_cluster_key(
gateway: Gateway, tmp_path: Path
) -> None:
def overloaded(_request: Request) -> Reply:
return Reply(
status=503,
body=b'{"error":{"message":"Controlled provider outage","type":"server_error","code":"503"}}',
)
with (
wire_server(overloaded) as wire,
owned_proxy_process(gateway, tmp_path, {}, workers=2) as owned,
owned.gateway.scenario() as scenario,
_patient_client(owned.gateway) as client,
):
model: Final = scenario.model(api_base=wire.url + "/v1", num_retries=0)
key: Final = scenario.key(models=[model])
results: Final = _burst(client, model, key, scenario.key(), lambda: None)
assert all(result.status is not None for result in results), results
happy: Final = tuple(result for result in results if result.label.startswith("happy"))
assert all(result.status == 503 for result in happy), happy
late: Final = tuple(result for result in results if result.label.startswith("late"))
assert all(result.status == 503 for result in late), late
seen: Final = tuple(request.body.decode() for request in wire.drain())
assert all(any(f"normalized error audit {result.label}" in body for body in seen) for result in happy + late), (
seen
)
crafted: Final = tuple(result for result in results if result.label.startswith("crafted"))
assert all(result.status == 400 and result.seconds < FAST_SECONDS for result in crafted), crafted
codes: Final = {
result.label: _error_information(string_value(result.call_id))["normalized_error"] for result in results
}
assert all(code == "503_PROVIDER_OVERLOADED" for label, code in codes.items() if label.startswith("happy")), (
codes
)
assert all(code == "400_INVALID_REQUEST" for label, code in codes.items() if label.startswith("crafted")), codes
assert client.get("/health/liveliness").status_code == 200

View file

@ -4,6 +4,7 @@ import signal
import threading
import uuid
from collections.abc import Callable, Iterator
from concurrent.futures import ThreadPoolExecutor
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
@ -13,16 +14,18 @@ import httpx
import psycopg
import pytest
import yaml
from integration._support.client import Gateway, delete_key_if_present, eventually, string_value
from integration._support.database import read_rows
from integration._support.process import OwnedProxy, owned_proxy_process
from integration._support.wire import Reply, Request, wire_server
from psycopg import sql
REQUESTS_WHILE_BLOCKED: Final = 6
CANCEL_LOG_LINE: Final = "in-flight scheduled job(s) for shutdown"
BATCH_DRAINED_LOG_LINE: Final = f"flushed {REQUESTS_WHILE_BLOCKED} daily spend update items from in-memory queue"
MODEL_INSERT_ARRIVED_LOG_LINE: Final = "path=/model/new"
COMMIT_DELAY_SECONDS: Final = 15
BURST_REQUESTS: Final = 30
def _api_requests(table: str, column: str, identity: str) -> int:
@ -44,6 +47,50 @@ def _waiting_on(table: str) -> int:
return waiting
def _committing_daily_user_spend() -> int:
rows: Final = read_rows(
"SELECT count(*)::int AS committing FROM pg_stat_activity "
"WHERE query='COMMIT' AND state='active' AND wait_event='PgSleep' AND pid IN "
"(SELECT pid FROM pg_locks WHERE relation = %s::regclass AND mode='RowExclusiveLock')",
('"LiteLLM_DailyUserSpend"',),
)
committing: Final = rows[0]["committing"]
assert isinstance(committing, int)
return committing
def _install_slow_commit(user_id: str, fails_once: bool) -> str:
suffix: Final = f"slow_commit_{uuid.uuid4().hex}"
with psycopg.connect(os.environ["DATABASE_URL"], autocommit=True) as connection:
connection.execute(sql.SQL("CREATE SEQUENCE {}").format(sql.Identifier(suffix)))
connection.execute(
sql.SQL(
"CREATE FUNCTION {}() RETURNS trigger LANGUAGE plpgsql AS $slow$ "
"BEGIN PERFORM pg_sleep({}); "
"IF {} AND nextval({}) = 1 THEN RAISE EXCEPTION 'integration: first COMMIT fails'; END IF; "
"RETURN NULL; END $slow$"
).format(
sql.Identifier(suffix), sql.Literal(COMMIT_DELAY_SECONDS), sql.Literal(fails_once), sql.Literal(suffix)
)
)
connection.execute(
sql.SQL(
'CREATE CONSTRAINT TRIGGER {} AFTER INSERT OR UPDATE ON "LiteLLM_DailyUserSpend" '
"DEFERRABLE INITIALLY DEFERRED FOR EACH ROW WHEN (NEW.user_id = {}) EXECUTE FUNCTION {}()"
).format(sql.Identifier(suffix), sql.Literal(user_id), sql.Identifier(suffix))
)
return suffix
def _drop_slow_commit(suffix: str) -> None:
with psycopg.connect(os.environ["DATABASE_URL"], autocommit=True) as connection:
connection.execute(
sql.SQL('DROP TRIGGER IF EXISTS {} ON "LiteLLM_DailyUserSpend"').format(sql.Identifier(suffix))
)
connection.execute(sql.SQL("DROP FUNCTION IF EXISTS {}()").format(sql.Identifier(suffix)))
connection.execute(sql.SQL("DROP SEQUENCE IF EXISTS {}").format(sql.Identifier(suffix)))
def _provider(request: Request) -> Reply:
if request.method != "POST":
return Reply(status=404, body=b'{"error":"not scripted"}')
@ -77,6 +124,17 @@ class _Shutdown:
def daily_user_requests(self) -> int:
return _api_requests("LiteLLM_DailyUserSpend", "user_id", self.owner)
def spend_logs(self) -> int:
rows: Final = read_rows('SELECT count(*)::int AS total FROM "LiteLLM_SpendLogs" WHERE "user"=%s', (self.owner,))
total: Final = rows[0]["total"]
assert isinstance(total, int)
return total
def burst(self, requests: int) -> None:
with ThreadPoolExecutor(max_workers=8) as pool:
for outcome in pool.map(lambda _: self.chat(), range(requests)):
assert outcome is None
def logged(self, line: str, times: int = 1) -> bool:
return self.owned.log.read_text(errors="replace").count(line) >= times
@ -121,7 +179,15 @@ def _config_with_pool_limit(tmp_path: Path, pool_limit: int) -> Path:
@contextmanager
def _proxy_with_one_seeded_row(gateway: Gateway, tmp_path: Path, pool_limit: int) -> Iterator[_Shutdown]:
def _proxy_with_one_seeded_row(
gateway: Gateway,
tmp_path: Path,
pool_limit: int,
cancel_timeout_seconds: int = 5,
settle_seconds: int = 0,
requests: int = REQUESTS_WHILE_BLOCKED,
workers: int = 1,
) -> Iterator[_Shutdown]:
owner: Final = f"integration-owner-{uuid.uuid4().hex}"
with gateway.scenario() as scenario, wire_server(_provider) as wire:
model: Final = scenario.model(api_base=wire.url + "/v1", num_retries=0)
@ -130,12 +196,15 @@ def _proxy_with_one_seeded_row(gateway: Gateway, tmp_path: Path, pool_limit: int
gateway,
tmp_path,
{
"DATABASE_URL": os.environ["DATABASE_URL"],
"LITELLM_LOG": "DEBUG",
"GRACEFUL_SHUTDOWN_TIMEOUT": "1",
"SCHEDULED_JOB_SHUTDOWN_FINISH_TIMEOUT_SECONDS": "1",
"SCHEDULED_JOB_SHUTDOWN_CANCEL_TIMEOUT_SECONDS": "5",
"SCHEDULED_JOB_SHUTDOWN_CANCEL_TIMEOUT_SECONDS": str(cancel_timeout_seconds),
},
config=_config_with_pool_limit(tmp_path, pool_limit),
remove_environment=("DATABASE_URL_READ_REPLICA",),
workers=workers,
) as owned:
key: Final = string_value(
owned.gateway.post("/key/generate", {"user_id": owner, "team_id": team, "models": [model]})["key"]
@ -145,8 +214,18 @@ def _proxy_with_one_seeded_row(gateway: Gateway, tmp_path: Path, pool_limit: int
shutdown.chat()
eventually(shutdown.daily_user_requests, lambda total: total == 1, seconds=60)
yield shutdown
assert _api_requests("LiteLLM_DailyUserSpend", "user_id", owner) == 1 + REQUESTS_WHILE_BLOCKED
assert _api_requests("LiteLLM_DailyTeamSpend", "team_id", team) == 1 + REQUESTS_WHILE_BLOCKED
written: Final = 1 + requests
if settle_seconds:
eventually(
lambda: (
_api_requests("LiteLLM_DailyUserSpend", "user_id", owner),
_api_requests("LiteLLM_DailyTeamSpend", "team_id", team),
),
lambda totals: totals == (written, written),
seconds=settle_seconds,
)
assert _api_requests("LiteLLM_DailyUserSpend", "user_id", owner) == written
assert _api_requests("LiteLLM_DailyTeamSpend", "team_id", team) == written
@pytest.mark.covers("quota_management.spend_tracking.shutdown_cancel_keeps_in_flight_daily_batch")
@ -187,3 +266,48 @@ def test_daily_spend_batch_cancelled_while_waiting_for_a_row_lock_is_written_exa
lambda: shutdown.logged(BATCH_DRAINED_LOG_LINE) and _waiting_on("LiteLLM_DailyUserSpend") == 1,
holder.rollback,
)
@pytest.mark.parametrize(
("cancel_timeout_seconds", "commit_fails_once"),
[
pytest.param(60, False, id="cancel_budget_outlives_commit"),
pytest.param(5, False, id="commit_outlives_cancel_budget"),
pytest.param(60, True, id="commit_fails_within_cancel_budget"),
pytest.param(5, True, id="commit_fails_after_cancel_budget"),
],
)
def test_daily_spend_batch_cancelled_while_postgres_is_committing_it_is_written_exactly_once(
gateway: Gateway, tmp_path: Path, cancel_timeout_seconds: int, commit_fails_once: bool
) -> None:
with (
_proxy_with_one_seeded_row(
gateway, tmp_path, pool_limit=10, cancel_timeout_seconds=cancel_timeout_seconds, settle_seconds=90
) as shutdown,
psycopg.connect(os.environ["DATABASE_URL"]) as memberships,
):
suffix: Final = _install_slow_commit(shutdown.owner, fails_once=commit_fails_once)
try:
shutdown.chat_while_spend_update_is_blocked(memberships, "LiteLLM_TeamMembership")
memberships.rollback()
shutdown.terminate_once(
lambda: shutdown.logged(BATCH_DRAINED_LOG_LINE) and _committing_daily_user_spend() == 1,
lambda: None,
)
finally:
_drop_slow_commit(suffix)
def test_daily_spend_burst_across_two_workers_survives_shutdown_during_commit_exactly_once(
gateway: Gateway, tmp_path: Path
) -> None:
with _proxy_with_one_seeded_row(
gateway, tmp_path, pool_limit=10, settle_seconds=120, requests=BURST_REQUESTS, workers=2
) as shutdown:
suffix: Final = _install_slow_commit(shutdown.owner, fails_once=False)
try:
shutdown.burst(BURST_REQUESTS)
eventually(shutdown.spend_logs, lambda total: total == 1 + BURST_REQUESTS, seconds=60)
shutdown.terminate_once(lambda: _committing_daily_user_spend() >= 1, lambda: None)
finally:
_drop_slow_commit(suffix)

View file

@ -1,7 +1,9 @@
import uuid
from pathlib import Path
from typing import Final
from integration._support.client import Gateway, eventually
from integration._support.process import owned_proxy
def test_spend_over_a_tag_max_budget_rejects_the_next_request(gateway: Gateway) -> None:
@ -58,3 +60,99 @@ def test_spend_over_a_tag_max_budget_rejects_the_next_request(gateway: Gateway)
},
)
assert control.status_code == 200, control.text
def test_key_tag_rpm_limit_rejects_the_second_request_carrying_that_tag(gateway: Gateway) -> None:
tag: Final = f"tag-rpm-{uuid.uuid4().hex}"
with gateway.scenario() as scenario:
model: Final = scenario.model(input_cost_per_token=0.01, output_cost_per_token=0.01)
key: Final = scenario.key(metadata={"tag_rpm_limit": {tag: 1}})
first: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": f"tag rpm {tag}"}],
"metadata": {"tags": [tag]},
},
key=key,
)
assert first.status_code == 200, first.text
second: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": f"tag rpm {tag}"}],
"metadata": {"tags": [tag]},
},
key=key,
)
assert second.status_code == 429, second.text
assert "rpm" in second.text.lower() or "rate" in second.text.lower(), second.text
control: Final = gateway.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": f"other tag rpm {tag}"}],
"metadata": {"tags": [f"other-{tag}"]},
},
key=key,
)
assert control.status_code == 200, control.text
def test_tag_budget_duration_resets_spend_and_unblocks_the_tag(gateway: Gateway, tmp_path: Path) -> None:
tag: Final = f"tag-reset-{uuid.uuid4().hex}"
with (
owned_proxy(
gateway,
tmp_path,
{"PROXY_BUDGET_RESCHEDULER_MIN_TIME": "2", "PROXY_BUDGET_RESCHEDULER_MAX_TIME": "3"},
) as candidate,
candidate.scenario() as scenario,
):
def delete_tag() -> None:
candidate.post("/tag/delete", {"name": tag})
model: Final = scenario.model(input_cost_per_token=0.01, output_cost_per_token=0.01)
candidate.post("/tag/new", {"name": tag, "max_budget": 0.0001, "budget_duration": "5s"})
scenario.cleanups.callback(delete_tag)
first: Final = candidate.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": f"tag spend {tag}"}],
"metadata": {"tags": [tag]},
},
)
assert first.status_code == 200, first.text
def rejection() -> int:
return candidate.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": f"tag reset probe {tag}"}],
"metadata": {"tags": [tag]},
},
).status_code
status: Final = eventually(rejection, lambda code: code != 200, seconds=70)
assert status in (400, 422, 429), status
blocked: Final = candidate.request(
"POST",
"/v1/chat/completions",
{
"model": model,
"messages": [{"role": "user", "content": f"tag reset probe {tag}"}],
"metadata": {"tags": [tag]},
},
)
assert "budget" in blocked.text.lower(), blocked.text
recovered: Final = eventually(rejection, lambda code: code == 200, seconds=70)
assert recovered == 200, recovered

View file

@ -591,6 +591,48 @@ async def test_embedding_cache_hit_sets_custom_llm_provider_on_logging_obj():
assert logging_obj.model_call_details["custom_llm_provider"] == "openai"
def test_sync_stream_responses_cache_hit_sets_custom_llm_provider_on_logging_obj(monkeypatch):
import litellm
from litellm.caching.caching import Cache
from litellm.types.utils import CallTypes
monkeypatch.setattr(litellm, "cache", Cache(type="local"))
kwargs = {"model": "azure/gpt-5.4-mini", "input": "hello", "stream": True}
cached_response = {
"id": "resp_sync_stream",
"created_at": int(time.time()),
"status": "completed",
"model": "gpt-5.4-mini",
"object": "response",
"output": [
{
"type": "message",
"id": "msg_sync_stream",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "hi", "annotations": []}],
}
],
}
litellm.cache.add_cache(json.dumps(cached_response), **kwargs)
handler = LLMCachingHandler(original_function=litellm.responses, request_kwargs=kwargs, start_time=datetime.now())
logging_obj = _build_logging_obj(CallTypes.responses.value, stream=True)
hit = handler._sync_get_cache(
model="azure/gpt-5.4-mini",
original_function=litellm.responses,
logging_obj=logging_obj,
start_time=datetime.now(),
call_type=CallTypes.responses.value,
kwargs=kwargs,
args=(),
)
assert hit.cached_result is not None
assert logging_obj.model_call_details["custom_llm_provider"] == "azure"
assert logging_obj.model_call_details["litellm_params"]["custom_llm_provider"] == "azure"
def test_request_kwargs_does_not_retain_logging_obj():
"""
The caching handler lives on logging_obj._llm_caching_handler, so keeping

View file

@ -2468,3 +2468,453 @@ def test_prompts_only_toggle_is_exposed_to_admin_ui_for_both_s3_callbacks(callba
from litellm.integrations.custom_logger import CustomLogger
assert "S3_LOG_PROMPTS_ONLY" in CustomLogger.get_callback_env_vars(callback_name)
def _element(payload: dict[str, object], key_suffix: str) -> s3BatchLoggingElement:
return s3BatchLoggingElement(
s3_object_key=f"2025-09-14/test-{key_suffix}.json",
payload=payload,
s3_object_download_filename=f"test-{key_suffix}.json",
)
def _ok_response() -> MagicMock:
response = MagicMock()
response.status_code = 200
response.raise_for_status = MagicMock()
return response
class _CountingPut:
def __init__(self) -> None:
self.in_flight = 0
self.peak = 0
self.calls = 0
async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
self.in_flight += 1
self.peak = max(self.peak, self.in_flight)
self.calls += 1
await asyncio.sleep(0.01)
self.in_flight -= 1
return _ok_response()
class _RecordingPut:
def __init__(self) -> None:
self.calls: tuple[tuple[str, str | None, dict[str, str] | None], ...] = ()
async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
self.calls = (*self.calls, (url, data, headers))
return _ok_response()
class _LateAppendingPut:
def __init__(self, logger: S3Logger, element: s3BatchLoggingElement, fail_first: bool = False) -> None:
self.logger = logger
self.element = element
self.fail_first = fail_first
self.appended = False
async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
if not self.appended:
self.appended = True
self.logger.log_queue.append(self.element)
if self.fail_first:
return _failure_response()
return _ok_response()
class _FailOnSuffixPut:
def __init__(self, suffixes: tuple[str, ...]) -> None:
self.failing = True
self.suffixes = suffixes
async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
if self.failing and url.endswith(self.suffixes):
return _failure_response()
return _ok_response()
class _FailUntilClearedPut:
def __init__(self) -> None:
self.failing = True
self.calls: tuple[tuple[str, str | None], ...] = ()
async def __call__(self, url: str, data: str | None = None, headers: dict[str, str] | None = None) -> MagicMock:
self.calls = (*self.calls, (url, data))
if self.failing:
return _failure_response()
return _ok_response()
@pytest.mark.asyncio
async def test_async_send_batch_bounds_concurrent_uploads() -> None:
logger = S3Logger(
s3_bucket_name="test-bucket",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1",
s3_max_concurrent_uploads=4,
)
put = _CountingPut()
logger.async_httpx_client = AsyncMock()
logger.async_httpx_client.put = put
logger.log_queue = [_element({"i": i}, f"{i}") for i in range(40)]
await logger.async_send_batch()
assert put.peak == 4
assert put.calls == 40
@pytest.mark.asyncio
async def test_async_send_batch_uploads_single_jsonl_file() -> None:
import json
logger = S3Logger(
s3_bucket_name="test-bucket",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1",
s3_batch_file_upload=True,
)
put = _RecordingPut()
logger.async_httpx_client = AsyncMock()
logger.async_httpx_client.put = put
payloads = [{"id": "req-1"}, {"id": "req-2"}, {"id": "req-3"}]
logger.log_queue = [_element(payload, f"{i}") for i, payload in enumerate(payloads)]
await logger.async_send_batch()
assert len(put.calls) == 1
url, data, headers = put.calls[0]
assert url.endswith(".jsonl")
assert data is not None
assert headers is not None
assert [json.loads(line) for line in data.splitlines()] == payloads
assert headers["Content-Type"] == "application/x-ndjson"
@pytest.mark.asyncio
async def test_flush_queue_preserves_events_added_during_upload() -> None:
logger = S3Logger(
s3_bucket_name="test-bucket",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1",
)
late_element = _element({"id": "late"}, "late")
logger.async_httpx_client = AsyncMock()
logger.async_httpx_client.put = _LateAppendingPut(logger, late_element)
logger.log_queue = [_element({"id": "first"}, "first")]
await logger.flush_queue()
assert logger.log_queue == [late_element]
def _override_logger(**overrides: object) -> S3Logger:
return S3Logger(
s3_bucket_name="test-bucket",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1",
s3_callback_params_override=overrides,
)
def test_env_backed_false_string_keeps_per_request_uploads() -> None:
assert _override_logger(s3_batch_file_upload="false").s3_batch_file_upload is False
assert _override_logger(s3_batch_file_upload="true").s3_batch_file_upload is True
logger = S3Logger(
s3_bucket_name="test-bucket",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1",
s3_batch_file_upload=True,
s3_callback_params_override={"s3_batch_file_upload": "false"},
)
assert logger.s3_batch_file_upload is True
@pytest.mark.parametrize("bad", [0, -3, "0", "abc", ""])
def test_invalid_concurrency_falls_back_to_default(bad: object) -> None:
from litellm.constants import DEFAULT_S3_MAX_CONCURRENT_UPLOADS
logger = _override_logger(s3_max_concurrent_uploads=bad)
assert logger.s3_max_concurrent_uploads == DEFAULT_S3_MAX_CONCURRENT_UPLOADS
assert logger._upload_semaphore._value == DEFAULT_S3_MAX_CONCURRENT_UPLOADS
def test_env_backed_concurrency_string_is_parsed() -> None:
logger = _override_logger(s3_max_concurrent_uploads="4")
assert logger.s3_max_concurrent_uploads == 4
assert logger._upload_semaphore._value == 4
@pytest.mark.parametrize("empty", [None, ""])
def test_empty_config_concurrency_falls_back_to_constructor_value(empty: object) -> None:
logger = S3Logger(
s3_bucket_name="test-bucket",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1",
s3_max_concurrent_uploads=4,
s3_callback_params_override={"s3_max_concurrent_uploads": empty},
)
assert logger.s3_max_concurrent_uploads == 4
assert logger._upload_semaphore._value == 4
def _failure_response() -> MagicMock:
response = MagicMock()
response.status_code = 400
response.raise_for_status = MagicMock(side_effect=Exception("s3 rejected the object"))
return response
@pytest.mark.asyncio
async def test_failed_uploads_stay_queued_for_next_flush() -> None:
logger = S3Logger(
s3_bucket_name="test-bucket",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1",
)
elements = [_element({"i": i}, f"{i}") for i in range(5)]
put = _FailOnSuffixPut(("test-2.json", "test-4.json"))
logger.async_httpx_client = AsyncMock()
logger.async_httpx_client.put = put
logger.log_queue = list(elements)
await logger.flush_queue()
assert logger.log_queue == [elements[2], elements[4]]
put.failing = False
await logger.flush_queue()
assert logger.log_queue == []
@pytest.mark.asyncio
async def test_batch_file_upload_failure_keeps_whole_batch() -> None:
logger = S3Logger(
s3_bucket_name="test-bucket",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1",
s3_batch_file_upload=True,
)
put = _FailUntilClearedPut()
logger.async_httpx_client = AsyncMock()
logger.async_httpx_client.put = put
elements = [_element({"i": i}, f"{i}") for i in range(3)]
logger.log_queue = list(elements)
await logger.flush_queue()
assert len(put.calls) == 1
assert len(logger.log_queue) == 1
assert logger.log_queue[0].body == "\n".join(json.dumps(element.payload) for element in elements)
@pytest.mark.asyncio
async def test_events_appended_during_failed_flush_survive() -> None:
logger = S3Logger(
s3_bucket_name="test-bucket",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1",
)
late = _element({"id": "late"}, "late")
logger.async_httpx_client = AsyncMock()
logger.async_httpx_client.put = _LateAppendingPut(logger, late, fail_first=True)
first = _element({"id": "first"}, "first")
logger.log_queue = [first]
await logger.flush_queue()
assert logger.log_queue == [first, late]
@pytest.mark.asyncio
async def test_batch_file_key_shape() -> None:
logger = S3Logger(
s3_bucket_name="test-bucket",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1",
s3_path="logs",
s3_batch_file_upload=True,
)
put = _RecordingPut()
logger.async_httpx_client = AsyncMock()
logger.async_httpx_client.put = put
logger.log_queue = [_element({"id": "req-1"}, "0")]
await logger.async_send_batch()
((url, _data, headers),) = put.calls
assert headers is not None
assert re.search(r".*/2025-09-14/batch_\d{2}-\d{2}-\d{2}_[0-9a-f]{32}\.jsonl$", url)
assert headers["Content-Disposition"].endswith('.jsonl"')
@pytest.mark.asyncio
async def test_batch_file_groups_raw_elements_by_key_parent() -> None:
logger = S3Logger(
s3_bucket_name="test-bucket",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1",
s3_batch_file_upload=True,
)
put = _RecordingPut()
logger.async_httpx_client = AsyncMock()
logger.async_httpx_client.put = put
alpha = s3BatchLoggingElement(
s3_object_key="logs/alpha/2026-01-01/a.json", payload={"id": "a"}, s3_object_download_filename="a.json"
)
beta = s3BatchLoggingElement(
s3_object_key="logs/beta/2026-01-01/b.json", payload={"id": "b"}, s3_object_download_filename="b.json"
)
plain = s3BatchLoggingElement(
s3_object_key="logs/2026-01-01/c.json", payload={"id": "c"}, s3_object_download_filename="c.json"
)
root = s3BatchLoggingElement(
s3_object_key="solo.json", payload={"id": "d"}, s3_object_download_filename="solo.json"
)
logger.log_queue = [alpha, beta, plain, root]
await logger.async_send_batch()
assert len(put.calls) == 4
by_parent = {
re.sub(r"(^|/)batch_\d{2}-\d{2}-\d{2}_[0-9a-f]{32}\.jsonl$", "", url.split(".com/", 1)[-1]): (url, data)
for url, data, _headers in put.calls
}
assert sorted(by_parent) == ["", "logs/2026-01-01", "logs/alpha/2026-01-01", "logs/beta/2026-01-01"]
assert [line for line in by_parent[""][1].splitlines()] == [json.dumps({"id": "d"})]
assert [line for line in by_parent["logs/alpha/2026-01-01"][1].splitlines()] == [json.dumps({"id": "a"})]
assert [line for line in by_parent["logs/beta/2026-01-01"][1].splitlines()] == [json.dumps({"id": "b"})]
assert [line for line in by_parent["logs/2026-01-01"][1].splitlines()] == [json.dumps({"id": "c"})]
@pytest.mark.asyncio
async def test_failed_batch_file_is_requeued_and_resent_unchanged() -> None:
logger = S3Logger(
s3_bucket_name="test-bucket",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1",
s3_batch_file_upload=True,
)
put = _FailUntilClearedPut()
logger.async_httpx_client = AsyncMock()
logger.async_httpx_client.put = put
logger.log_queue = [_element({"i": i}, f"{i}") for i in range(3)]
await logger.flush_queue()
assert len(logger.log_queue) == 1
assert logger.log_queue[0].body is not None
assert logger.log_queue[0].s3_object_key.endswith(".jsonl")
put.failing = False
await logger.flush_queue()
assert logger.log_queue == []
assert len(put.calls) == 2
assert put.calls[0] == put.calls[1]
@pytest.mark.asyncio
async def test_elements_appended_after_failed_batch_file_get_their_own_file() -> None:
logger = S3Logger(
s3_bucket_name="test-bucket",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1",
s3_batch_file_upload=True,
)
put = _FailUntilClearedPut()
logger.async_httpx_client = AsyncMock()
logger.async_httpx_client.put = put
logger.log_queue = [_element({"id": "first"}, "first")]
await logger.flush_queue()
late = _element({"id": "late"}, "late")
logger.log_queue.append(late)
put.failing = False
await logger.flush_queue()
assert logger.log_queue == []
assert len(put.calls) == 3
assert put.calls[0] == put.calls[1]
assert put.calls[2][0] != put.calls[0][0]
assert put.calls[2][1] == json.dumps({"id": "late"})
@pytest.mark.asyncio
async def test_batch_file_mode_disabled_when_s3_v2_is_cold_storage_logger(monkeypatch: pytest.MonkeyPatch) -> None:
logger = S3Logger(
s3_bucket_name="test-bucket",
s3_aws_access_key_id="test-key",
s3_aws_secret_access_key="test-secret",
s3_region_name="us-east-1",
s3_batch_file_upload=True,
)
put = _RecordingPut()
logger.async_httpx_client = AsyncMock()
logger.async_httpx_client.put = put
import litellm
monkeypatch.setattr(litellm, "cold_storage_custom_logger", "s3_v2")
logger.log_queue = [_element({"id": "req-1"}, "0")]
await logger.async_send_batch()
assert len(put.calls) == 1
assert put.calls[0][0].endswith("test-0.json")
monkeypatch.setattr(litellm, "cold_storage_custom_logger", None)
logger.log_queue = [_element({"id": "req-2"}, "1")]
await logger.async_send_batch()
assert len(put.calls) == 2
assert put.calls[1][0].endswith(".jsonl")

View file

@ -1,3 +1,5 @@
import time
import httpx
import pytest
@ -163,6 +165,18 @@ def test_variants_of_one_failure_share_a_normalized_error(messages: tuple[Except
assert normalized == {expected}
def test_normalize_error_passthrough_prefix_wins_over_upstream_body_text() -> None:
from fastapi import HTTPException
for detail in (
'Upstream passthrough request failed with status 400: {"error": {"message": "no deployments available for this model"}}',
'Upstream passthrough request failed with status 400: {"error": {"message": "max budget reached"}}',
):
exc = HTTPException(status_code=400, detail=detail)
message = f"400: {detail}"
assert normalize_error(exc, "400", message) == "500_UPSTREAM_PASSTHROUGH", message
def test_router_no_healthy_deployment_wording_clusters_as_no_healthy_deployments() -> None:
for message in (RouterErrors.no_healthy_deployments.value, "No healthy deployments found."):
exc = litellm.BadRequestError(message, llm_provider="openai", model="gpt-4o")
@ -229,3 +243,36 @@ def test_normalized_error_never_embeds_dynamic_parts() -> None:
info = StandardLoggingPayloadSetup.get_error_information(exc)
assert info["error_message"] == "No team has access to anthropic.claude-sonnet-4-5"
assert "claude" not in (info["normalized_error"] or "")
def test_repeated_exceeded_in_a_288kb_message_classifies_in_linear_time() -> None:
model = ("exceeded " * 32_000)[:288_000]
message = (
f"/chat/completions: Invalid model name passed in model={model}. Call `/v1/models` to view available models"
)
exc = litellm.BadRequestError(message=message, model="unknown-model", llm_provider="openai")
started = time.perf_counter()
code = normalize_error(exc, "400", message)
elapsed = time.perf_counter() - started
assert code == "400_INVALID_REQUEST", code
assert elapsed < 1.0, f"normalize_error took {elapsed:.2f}s on a 288 KB message"
@pytest.mark.parametrize(
"message",
[
"ExceededBudget: User=abc over budget. Spend=12.5, Budget=10.0",
"Exceeded budget for provider openai: 105.2 >= 100.0",
"LiteLLM Team: team-1, exceeded budget for model=gpt-4o-mini",
"ExceededBudget: Key over 1d budget. Spend=3.0, Budget=2.0",
"Budget has been exceeded! Key=sk-... Current cost: 11.0, Max budget: 10.0",
"EXCEEDED " + "x" * 65 + " BuDgEt",
],
)
def test_real_budget_wordings_still_cluster_as_budget_exceeded(message: str) -> None:
assert normalize_error(Exception(message), "400", message) == "429_BUDGET_EXCEEDED"
@pytest.mark.parametrize("message", ["budget then exceeded", "exceeded the limit\nbudget unaffected", "exceededbudge"])
def test_exceeded_without_a_following_budget_on_the_same_line_is_not_budget(message: str) -> None:
assert normalize_error(Exception(message), "400", message) == "400_INVALID_REQUEST"

View file

@ -4501,3 +4501,241 @@ async def test_tag_batch_drained_from_redis_and_cancelled_mid_flight_is_restored
await asyncio.wait_for(db.rolled_back.wait(), timeout=5)
assert db.transaction_outcomes == ["rollback"]
assert _daily_upserts(db, "LiteLLM_DailyTagSpend") == []
class _CommittingDailySpendFakeDB(_DailySpendFakeDB):
"""Runs the daily upsert at once but holds the COMMIT until released. A COMMIT that has left
the client lands on the server whether or not the client keeps waiting for the reply."""
def __init__(self) -> None:
super().__init__(failing_table=None)
self.committing = asyncio.Event()
self.commit_release = asyncio.Event()
self.transaction_outcomes: list[str] = []
@asynccontextmanager
async def _tx(self) -> AsyncIterator["_CommittingDailySpendFakeDB"]:
try:
yield self
except BaseException:
self.transaction_outcomes.append("rollback")
raise
self.committing.set()
try:
await self.commit_release.wait()
finally:
self.transaction_outcomes.append("commit")
@pytest.mark.parametrize(("queue_name", "entity_type", "entity_id_field", "table"), _DAILY_SPEND_ENTITIES)
@pytest.mark.asyncio
async def test_cancel_that_lands_while_the_daily_batch_is_committing_waits_for_the_commit_and_does_not_requeue_it(
queue_name: str, entity_type: str, entity_id_field: str, table: str
):
"""Shutdown cancels the tick after the COMMIT has left for Postgres. The server finishes that
commit whatever the client does, so putting the batch back on the queue makes the final flush
write the same spend a second time. The tick has to wait for the commit's outcome instead."""
db_writer = DBSpendUpdateWriter()
queue = _DAILY_SPEND_QUEUES[queue_name](db_writer)
await queue.add_update({"key-a": _daily_entity_txn(entity_id_field)})
await queue.add_update({"key-a": _daily_entity_txn(entity_id_field)})
db = _CommittingDailySpendFakeDB()
def flush(prisma_db: _DailySpendFakeDB):
return db_writer._flush_daily_spend_queue(
queue=queue,
entity_type=entity_type,
commit=_DAILY_SPEND_COMMITS[entity_type],
n_retry_times=0,
prisma_client=_WindowSpendFakePrisma(prisma_db),
proxy_logging_obj=MagicMock(),
)
tick = asyncio.ensure_future(flush(db))
await asyncio.wait_for(db.committing.wait(), timeout=5)
tick.cancel()
finished, _ = await asyncio.wait({tick}, timeout=0.2)
assert finished == {tick}, "the cancelled tick must hand the in-flight commit's outcome to the next flush"
with pytest.raises(asyncio.CancelledError):
tick.result()
assert len(queue.interrupted_commits) == 1
db.commit_release.set()
await queue.settle_interrupted_commits()
assert db.transaction_outcomes == ["commit"]
(upsert,) = _daily_upserts(db, table)
assert _row_values(upsert, "api_requests") == [2]
assert queue.update_queue.empty(), "a batch whose COMMIT already left for the server must not be requeued"
final_db = _DailySpendFakeDB(failing_table=None)
await flush(final_db)
assert _daily_upserts(final_db, table) == [], "the final flush must not write the committed batch again"
@pytest.mark.asyncio
async def test_tag_batch_drained_from_redis_and_cancelled_while_committing_is_not_restored():
"""Same in-flight COMMIT as the in-memory path, but the drained rows live in Redis. Restoring
them after the server committed writes the tag spend twice on the next tick."""
db_writer = DBSpendUpdateWriter()
drained = {"key-a": cast(DailyTagSpendTransaction, _daily_entity_txn("tag"))}
redis_buffer = _DrainedTagRedisBuffer(drained)
db_writer.redis_update_buffer = cast(RedisUpdateBuffer, redis_buffer)
db = _CommittingDailySpendFakeDB()
tick = asyncio.ensure_future(
db_writer._drain_and_commit_daily_tag_spend_from_redis(
prisma_client=_WindowSpendFakePrisma(db),
n_retry_times=0,
proxy_logging_obj=MagicMock(),
)
)
await asyncio.wait_for(db.committing.wait(), timeout=5)
tick.cancel()
finished, _ = await asyncio.wait({tick}, timeout=0.2)
assert finished == {tick}, "the cancelled drain must hand the in-flight commit's outcome to the next drain"
with pytest.raises(asyncio.CancelledError):
tick.result()
assert len(db_writer.interrupted_tag_commits) == 1
db.commit_release.set()
(settle,) = tuple(db_writer.interrupted_tag_commits)
await settle
assert db.transaction_outcomes == ["commit"]
assert redis_buffer.restored == [], (
"a tag batch whose COMMIT already left for the server must not be restored to Redis"
)
class _CommitFailingDailySpendFakeDB(_CommittingDailySpendFakeDB):
"""COMMIT leaves for the server but the reply comes back as a failure."""
@asynccontextmanager
async def _tx(self) -> AsyncIterator["_CommitFailingDailySpendFakeDB"]:
yield self
self.committing.set()
await self.commit_release.wait()
self.transaction_outcomes.append("commit_failed")
raise Exception("connection reset")
@pytest.mark.asyncio
async def test_cancel_while_committing_requeues_the_batch_when_the_commit_itself_fails():
"""Waiting for the in-flight commit's outcome must not swallow a real commit failure:
the batch still goes back on the queue and the next flush writes it once."""
db_writer = DBSpendUpdateWriter()
queue = db_writer.daily_spend_update_queue
await queue.add_update({"key-a": _daily_txn()})
await queue.add_update({"key-a": _daily_txn()})
db = _CommitFailingDailySpendFakeDB()
def flush(prisma_db: _DailySpendFakeDB):
return db_writer._flush_daily_spend_queue(
queue=queue,
entity_type="user",
commit=DBSpendUpdateWriter.update_daily_user_spend,
n_retry_times=0,
prisma_client=_WindowSpendFakePrisma(prisma_db),
proxy_logging_obj=MagicMock(),
)
tick = asyncio.ensure_future(flush(db))
await asyncio.wait_for(db.committing.wait(), timeout=5)
tick.cancel()
finished, _ = await asyncio.wait({tick}, timeout=0.2)
assert finished == {tick}, "the cancelled tick must not eat the shutdown budget waiting on the commit"
with pytest.raises(asyncio.CancelledError):
tick.result()
db.commit_release.set()
await queue.settle_interrupted_commits()
assert db.transaction_outcomes == ["commit_failed"]
assert not queue.update_queue.empty(), "a batch whose COMMIT came back failed must be requeued"
final_db = _DailySpendFakeDB(failing_table=None)
await flush(final_db)
(upsert,) = _daily_upserts(final_db, "LiteLLM_DailyUserSpend")
assert _row_values(upsert, "api_requests") == [2]
assert queue.update_queue.empty()
@pytest.mark.asyncio
async def test_shutdown_flush_that_lands_before_the_interrupted_commit_resolves_still_writes_a_failed_batch_once():
"""The cancelled tick returns right away, so a COMMIT can still be in flight when the
shutdown flush runs. If that commit later fails, the flush must first settle it, pick the
requeued rows back up, and write them exactly once instead of losing them."""
db_writer = DBSpendUpdateWriter()
queue = db_writer.daily_spend_update_queue
await queue.add_update({"key-a": _daily_txn()})
await queue.add_update({"key-a": _daily_txn()})
db = _CommitFailingDailySpendFakeDB()
def flush(prisma_db: _DailySpendFakeDB):
return db_writer._flush_daily_spend_queue(
queue=queue,
entity_type="user",
commit=DBSpendUpdateWriter.update_daily_user_spend,
n_retry_times=0,
prisma_client=_WindowSpendFakePrisma(prisma_db),
proxy_logging_obj=MagicMock(),
)
tick = asyncio.ensure_future(flush(db))
await asyncio.wait_for(db.committing.wait(), timeout=5)
tick.cancel()
with pytest.raises(asyncio.CancelledError):
await asyncio.wait_for(tick, timeout=5)
assert db.transaction_outcomes == [], "the COMMIT is still on the wire when the shutdown flush starts"
final_db = _DailySpendFakeDB(failing_table=None)
shutdown_flush = asyncio.ensure_future(flush(final_db))
finished, _ = await asyncio.wait({shutdown_flush}, timeout=0.2)
assert finished == set(), "the shutdown flush must wait for the interrupted commit's outcome"
assert _daily_upserts(final_db, "LiteLLM_DailyUserSpend") == []
db.commit_release.set()
await asyncio.wait_for(shutdown_flush, timeout=5)
(upsert,) = _daily_upserts(final_db, "LiteLLM_DailyUserSpend")
assert _row_values(upsert, "api_requests") == [2]
assert queue.update_queue.empty()
@pytest.mark.asyncio
async def test_shutdown_drain_that_lands_before_the_interrupted_tag_commit_resolves_restores_a_failed_batch():
"""Same ordering for the Redis tag path: the shutdown drain must settle the interrupted
commit before the destructive drain, or a commit that fails late is never restored."""
db_writer = DBSpendUpdateWriter()
drained = {"key-a": cast(DailyTagSpendTransaction, _daily_entity_txn("tag"))}
redis_buffer = _DrainedTagRedisBuffer(drained)
db_writer.redis_update_buffer = cast(RedisUpdateBuffer, redis_buffer)
db = _CommitFailingDailySpendFakeDB()
def drain(prisma_db: _DailySpendFakeDB):
return db_writer._drain_and_commit_daily_tag_spend_from_redis(
prisma_client=_WindowSpendFakePrisma(prisma_db),
n_retry_times=0,
proxy_logging_obj=MagicMock(),
)
tick = asyncio.ensure_future(drain(db))
await asyncio.wait_for(db.committing.wait(), timeout=5)
tick.cancel()
with pytest.raises(asyncio.CancelledError):
await asyncio.wait_for(tick, timeout=5)
assert db.transaction_outcomes == []
final_db = _DailySpendFakeDB(failing_table=None)
shutdown_drain = asyncio.ensure_future(drain(final_db))
finished, _ = await asyncio.wait({shutdown_drain}, timeout=0.2)
assert finished == set(), "the shutdown drain must wait for the interrupted commit's outcome"
assert _daily_upserts(final_db, "LiteLLM_DailyTagSpend") == []
db.commit_release.set()
await asyncio.wait_for(shutdown_drain, timeout=5)
assert redis_buffer.restored == [drained], "a tag batch whose COMMIT came back failed must be restored to Redis"
(upsert,) = _daily_upserts(final_db, "LiteLLM_DailyTagSpend")
assert _row_values(upsert, "api_requests") == [1]

View file

@ -2261,7 +2261,7 @@ async def test_apply_to_output_streaming_mixed_chunks_flushes_and_warns():
assert mock_logger.warning.call_count == 2
warning_messages = [call.args[0] for call in mock_logger.warning.call_args_list]
assert any("mixed stream detected" in msg for msg in warning_messages)
assert any("unknown event objects" in msg for msg in warning_messages)
assert any("Output PII masking was skipped" in msg for msg in warning_messages)
# ---------------------------------------------------------------------------
@ -2519,6 +2519,147 @@ async def test_apply_to_output_streaming_anthropic_sse_bytes_without_pii_are_for
assert collected == byte_chunks
def _gemini_sse(text: str) -> bytes:
payload = {"candidates": [{"content": {"parts": [{"text": text}], "role": "model"}, "index": 0}]}
return f"data: {json.dumps(payload)}\n\n".encode()
@pytest.mark.asyncio
async def test_apply_to_output_streaming_gemini_sse_bytes_are_forwarded_incrementally_until_upstream_aborts():
guardrail = _OPTIONAL_PresidioPIIMasking(
mock_testing=True,
apply_to_output=True,
mock_redacted_text={"text": "<PERSON>"},
)
frames = [_gemini_sse("Partial one from John Smith. "), _gemini_sse("Partial two. ")]
collected: list[object] = []
async def mock_stream():
for frame in frames:
yield frame
raise ConnectionError("upstream closed mid-stream")
async def collect() -> None:
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
response=mock_stream(),
request_data={},
):
collected.append(chunk)
with pytest.raises(ConnectionError):
await collect()
assert collected == frames
@pytest.mark.asyncio
async def test_apply_to_output_streaming_anthropic_first_frame_split_across_transport_chunks_is_still_masked():
guardrail = _OPTIONAL_PresidioPIIMasking(
mock_testing=True,
apply_to_output=True,
mock_redacted_text={"text": "<PERSON>"},
)
message_start = _anthropic_sse(
"message_start",
{"type": "message_start", "message": {"id": "msg_1", "model": "claude", "content": [], "usage": {}}},
)
split_at = message_start.index(b'"message_') + len(b'"message_')
byte_chunks = [
message_start[:split_at],
message_start[split_at:],
_anthropic_sse(
"content_block_start",
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
),
_anthropic_sse(
"content_block_delta",
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "John Smith"}},
),
_anthropic_sse("content_block_stop", {"type": "content_block_stop", "index": 0}),
_anthropic_sse("message_delta", {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {}}),
_anthropic_sse("message_stop", {"type": "message_stop"}),
]
async def mock_stream():
for b in byte_chunks:
yield b
collected = []
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
response=mock_stream(),
request_data={},
):
collected.append(chunk)
joined = b"".join(collected).decode()
assert "John Smith" not in joined, joined
assert "".join(text for _, text in _anthropic_text_deltas(collected)) == "<PERSON>"
assert joined.count("event: message_start") == 1
@pytest.mark.asyncio
async def test_apply_to_output_streaming_gemini_first_frame_split_across_transport_chunks_streams_incrementally():
guardrail = _OPTIONAL_PresidioPIIMasking(
mock_testing=True,
apply_to_output=True,
mock_redacted_text={"text": "<PERSON>"},
)
first = _gemini_sse("Partial one from John Smith. ")
second = _gemini_sse("Partial two. ")
collected: list[object] = []
async def mock_stream():
yield first[:20]
yield first[20:]
yield second
raise ConnectionError("upstream closed mid-stream")
async def collect() -> None:
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
response=mock_stream(),
request_data={},
):
collected.append(chunk)
with pytest.raises(ConnectionError):
await collect()
assert collected == [first, second]
@pytest.mark.asyncio
async def test_apply_to_output_streaming_unterminated_first_frame_is_released_once_it_exceeds_the_cap():
guardrail = _OPTIONAL_PresidioPIIMasking(
mock_testing=True,
apply_to_output=True,
mock_redacted_text={"text": "<PERSON>"},
)
piece = b"data: " + b"x" * 1023 + b"\n"
pieces_to_cap = -(-(64 * 1024) // len(piece))
released_at: list[int] = []
async def mock_stream():
for index in range(pieces_to_cap * 4):
if collected:
released_at.append(index)
yield piece
collected: list[object] = []
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=UserAPIKeyAuth(api_key="test-key"),
response=mock_stream(),
request_data={},
):
collected.append(chunk)
assert released_at, "nothing reached the caller before the upstream finished"
assert released_at[0] == pieces_to_cap, released_at[:3]
assert b"".join(collected) == piece * (pieces_to_cap * 4)
@pytest.mark.asyncio
async def test_apply_to_output_streaming_anthropic_sse_bytes_fail_closed_when_presidio_is_unreachable():
"""

View file

@ -1,8 +1,10 @@
import asyncio
import gzip
import json
import logging
import os
import sys
import zlib
from collections.abc import Callable
from contextlib import ExitStack, contextmanager
from io import BytesIO
@ -12,12 +14,14 @@ from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from fastapi import Request, Response, UploadFile
from fastapi import HTTPException, Request, Response, UploadFile
from fastapi.responses import StreamingResponse
from pydantic import ValidationError
from starlette.datastructures import FormData, Headers, QueryParams
from starlette.datastructures import UploadFile as StarletteUploadFile
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
@ -27,6 +31,8 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
HttpPassThroughEndpointHelpers,
InitPassThroughEndpointHelpers,
_registered_pass_through_routes,
_truncate_upstream_error_body,
_with_trace_context,
chat_completion_pass_through_endpoint,
create_pass_through_route,
initialize_pass_through_endpoints,
@ -34,7 +40,6 @@ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
resolve_llm_passthrough_timeout,
resolve_pass_through_request_timeout,
websocket_passthrough_request,
_with_trace_context,
)
from litellm.proxy.pass_through_endpoints.success_handler import (
PassThroughEndpointLogging,
@ -4126,6 +4131,705 @@ async def test_pass_through_request_streaming_upstream_error_returned_unchanged(
assert failure_call_kwargs["original_exception"].status_code == 403
class _UpstreamErrorBodyStream(httpx.AsyncByteStream):
def __init__(self, body: bytes) -> None:
self._body: Final = body
async def __aiter__(self):
yield self._body
def _upstream_error_request() -> MagicMock:
mock_request: Final = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://test-proxy.com/mock-upstream/v1beta/models/claude-nope-9:generateContent"
mock_request.body = AsyncMock(return_value=b'{"contents": []}')
mock_request.headers = Headers({"content-type": "application/json"})
mock_request.query_params = QueryParams({})
return mock_request
@pytest.mark.asyncio
async def test_pass_through_request_non_streaming_upstream_error_body_logged_and_in_failure_detail(
caplog: pytest.LogCaptureFixture,
):
upstream_body: Final = {
"error": {
"code": 404,
"message": "Publisher Model `publishers/anthropic/models/claude-nope-9` was not found or your project does not have access",
"status": "NOT_FOUND",
}
}
upstream_content: Final = json.dumps(upstream_body).encode("utf-8")
upstream_response: Final = httpx.Response(
status_code=404,
headers={"content-type": "application/json"},
content=upstream_content,
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:generateContent"),
)
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
) as mock_get_client:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.ProxyBaseLLMRequestProcessing"
) as mock_processing:
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_processing.get_custom_headers.return_value = {}
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:generateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
)
warning_messages: Final = [record.getMessage() for record in caplog.records if record.levelno == logging.WARNING]
upstream_warnings: Final = [
message for message in warning_messages if "upstream" in message and "returned 404" in message
]
assert len(upstream_warnings) == 1, warning_messages
assert "was not found or your project" in upstream_warnings[0]
assert "/v1beta/models/claude-nope-9:generateContent" in upstream_warnings[0]
assert response.status_code == 404
assert response.body == upstream_content
mock_proxy_logging.post_call_failure_hook.assert_called_once()
failure_call_kwargs: Final = mock_proxy_logging.post_call_failure_hook.call_args.kwargs
original_exception: Final = failure_call_kwargs["original_exception"]
assert isinstance(original_exception, HTTPException)
assert original_exception.status_code == 404
assert "was not found or your project" in original_exception.detail
@pytest.mark.asyncio
async def test_pass_through_request_streaming_upstream_error_body_reaches_client_and_failure_detail():
upstream_content: Final = (
b'data: {"error": {"code": 403, "message": "stream access was not found or your project lacks"}}\n\n'
)
upstream_response: Final = httpx.Response(
status_code=403,
headers={"content-type": "text/event-stream"},
stream=_UpstreamErrorBodyStream(upstream_content),
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
) as mock_get_client:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
) as mock_success_handler:
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler.return_value = None
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=True,
)
assert isinstance(response, StreamingResponse)
assert response.status_code == 403
streamed_chunks: Final = [chunk async for chunk in response.body_iterator]
streamed_bytes: Final = b"".join(
chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") for chunk in streamed_chunks
)
assert streamed_bytes == upstream_content
mock_proxy_logging.post_call_failure_hook.assert_called_once()
original_exception: Final = mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"]
assert "was not found or your project" in original_exception.detail
@pytest.mark.asyncio
async def test_truncate_upstream_error_body_caps_at_log_limit():
short_body: Final = "x" * 4096
assert _truncate_upstream_error_body(short_body) == short_body
long_body: Final = "a" * 5000
truncated: Final = _truncate_upstream_error_body(long_body)
assert truncated == f"{'a' * 4096}... (truncated at 4096 chars)"
upstream_response: Final = httpx.Response(
status_code=500,
headers={"content-type": "text/plain"},
content=long_body.encode("utf-8"),
request=httpx.Request("POST", "http://target-api.com/api/big-error"),
)
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
) as mock_get_client:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.ProxyBaseLLMRequestProcessing"
) as mock_processing:
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_processing.get_custom_headers.return_value = {}
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/api/big-error",
custom_headers={},
user_api_key_dict=MagicMock(),
)
detail: Final = mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
assert detail == f"Upstream passthrough request failed with status 500: {'a' * 4096}... (truncated at 4096 chars)"
@pytest.mark.asyncio
async def test_pass_through_request_upstream_error_log_strips_provider_key_from_url():
upstream_content: Final = b'{"error": "denied"}'
upstream_response: Final = httpx.Response(
status_code=404,
headers={"content-type": "application/json"},
content=upstream_content,
request=httpx.Request(
"POST",
"http://target-api.com/v1beta/models/claude-nope-9:generateContent?key=AIzaSySecretProviderKey123",
),
)
with patch.object(verbose_proxy_logger, "warning") as mock_warning:
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
) as mock_get_client:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.ProxyBaseLLMRequestProcessing"
) as mock_processing:
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_processing.get_custom_headers.return_value = {}
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:generateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
)
upstream_warnings: Final = [
call
for call in mock_warning.call_args_list
if call.args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s"
]
assert len(upstream_warnings) == 1, mock_warning.call_args_list
logged_url: Final = str(upstream_warnings[0].args[2])
assert "/v1beta/models/claude-nope-9:generateContent" in logged_url
assert "AIzaSySecretProviderKey123" not in logged_url
assert "key=" not in logged_url
@pytest.mark.asyncio
@pytest.mark.parametrize("turn_off_message_logging", [True, False])
async def test_passthrough_upstream_error_body_redacted_when_message_logging_off(
turn_off_message_logging: bool,
):
upstream_content: Final = b'{"error": {"message": "upstream body says the project was not found"}}'
upstream_response: Final = httpx.Response(
status_code=404,
headers={"content-type": "application/json"},
content=upstream_content,
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:generateContent"),
)
user_api_key_dict: Final = MagicMock()
user_api_key_dict.metadata = {
"logging": [
{
"callback_name": "prometheus",
"callback_type": "success_and_failure",
"callback_vars": {"turn_off_message_logging": turn_off_message_logging},
}
]
}
user_api_key_dict.team_metadata = None
user_api_key_dict.team_id = None
with patch.object(verbose_proxy_logger, "warning") as mock_warning:
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
) as mock_get_client:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.ProxyBaseLLMRequestProcessing"
) as mock_processing:
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_processing.get_custom_headers.return_value = {}
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:generateContent",
custom_headers={},
user_api_key_dict=user_api_key_dict,
)
assert response.status_code == 404
assert response.body == upstream_content
upstream_warnings: Final = [
call
for call in mock_warning.call_args_list
if call.args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s"
]
assert len(upstream_warnings) == 1, mock_warning.call_args_list
logged_body: Final = str(upstream_warnings[0].args[4])
mock_proxy_logging.post_call_failure_hook.assert_called_once()
detail: Final = mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
if turn_off_message_logging:
assert logged_body == "redacted-by-litellm"
assert "upstream body says the project was not found" not in logged_body
assert detail == "Upstream passthrough request failed with status 404: redacted-by-litellm"
else:
assert "upstream body says the project was not found" in logged_body
assert detail == f"Upstream passthrough request failed with status 404: {upstream_content.decode()}"
class _ChunkedUpstreamErrorBodyStream(httpx.AsyncByteStream):
def __init__(self, chunks: tuple[bytes, ...]) -> None:
self._chunks: Final = chunks
self.served: int = 0
async def __aiter__(self):
for chunk in self._chunks:
self.served += 1
yield chunk
@pytest.mark.asyncio
async def test_pass_through_request_streaming_upstream_error_reads_only_preview_and_relays_full_body():
chunk_size: Final = 1024
chunks: Final = tuple(b"x" * chunk_size for _ in range(10))
upstream_content: Final = b"".join(chunks)
body_stream: Final = _ChunkedUpstreamErrorBodyStream(chunks)
upstream_response: Final = httpx.Response(
status_code=500,
headers={"content-type": "text/plain"},
stream=body_stream,
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
served_at_warning: list[int] = []
real_warning: Final = verbose_proxy_logger.warning
def _recording_warning(*args, **kwargs):
if args and args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s":
served_at_warning.append(body_stream.served)
return real_warning(*args, **kwargs)
with patch.object(verbose_proxy_logger, "warning", side_effect=_recording_warning):
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
) as mock_get_client:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
) as mock_success_handler:
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler.return_value = None
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=True,
)
assert isinstance(response, StreamingResponse)
assert response.status_code == 500
streamed_chunks: Final = [chunk async for chunk in response.body_iterator]
streamed_bytes: Final = b"".join(
chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") for chunk in streamed_chunks
)
assert streamed_bytes == upstream_content
assert served_at_warning == [5], (
"each raw chunk is yielded as-is; five 1024-byte chunks are the first point the preview budget is exceeded"
)
expected_body: Final = f"{'x' * 4096}... (truncated at 4096 chars)"
assert (
mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
== f"Upstream passthrough request failed with status 500: {expected_body}"
)
@pytest.mark.asyncio
async def test_pass_through_request_streaming_upstream_error_single_large_chunk_stays_bounded():
first_chunk: Final = b"x" * 65536
second_chunk: Final = b'{"error": "tail"}'
upstream_content: Final = first_chunk + second_chunk
body_stream: Final = _ChunkedUpstreamErrorBodyStream((first_chunk, second_chunk))
upstream_response: Final = httpx.Response(
status_code=500,
headers={"content-type": "text/plain"},
stream=body_stream,
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
served_at_warning: list[int] = []
real_warning: Final = verbose_proxy_logger.warning
def _recording_warning(*args, **kwargs):
if args and args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s":
served_at_warning.append(body_stream.served)
return real_warning(*args, **kwargs)
with patch.object(verbose_proxy_logger, "warning", side_effect=_recording_warning):
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
) as mock_get_client:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
) as mock_success_handler:
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler.return_value = None
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=True,
)
assert isinstance(response, StreamingResponse)
assert response.status_code == 500
streamed_chunks: Final = [chunk async for chunk in response.body_iterator]
streamed_bytes: Final = b"".join(
chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") for chunk in streamed_chunks
)
assert streamed_bytes == upstream_content
assert served_at_warning == [1], (
"the rechunked preview is served from the first raw chunk; the second must not be pulled before the warning"
)
expected_body: Final = f"{'x' * 4096}... (truncated at 4096 chars)"
assert (
mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
== f"Upstream passthrough request failed with status 500: {expected_body}"
)
class _UpstreamErrorBodyStreamDropping(httpx.AsyncByteStream):
async def __aiter__(self):
yield b'{"error": "half'
raise httpx.ReadError("peer reset")
@pytest.mark.asyncio
async def test_pass_through_request_streaming_upstream_error_body_read_failure_keeps_status_and_partial_body():
"""
Regression: a 502 whose upstream dies while the error preview is being read
must still reach the client with status 502 and the bytes already received;
the read failure must not escape as a ProxyException 500.
"""
upstream_response: Final = httpx.Response(
status_code=502,
headers={"content-type": "application/json"},
stream=_UpstreamErrorBodyStreamDropping(),
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
recorded_warnings: list[tuple] = []
real_warning: Final = verbose_proxy_logger.warning
def _recording_warning(*args, **kwargs):
if args and str(args[0]).startswith("pass_through_endpoint: upstream"):
recorded_warnings.append(args)
return real_warning(*args, **kwargs)
with patch.object(verbose_proxy_logger, "warning", side_effect=_recording_warning):
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
) as mock_get_client:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
) as mock_success_handler:
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler.return_value = None
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=True,
)
assert isinstance(response, StreamingResponse)
assert response.status_code == 502
streamed_chunks: Final = [chunk async for chunk in response.body_iterator]
streamed_bytes: Final = b"".join(
chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") for chunk in streamed_chunks
)
assert streamed_bytes == b'{"error": "half'
await upstream_response.aclose()
rendered: Final = [str(args[0]) for args in recorded_warnings]
formats: Final = [args[0] for args in recorded_warnings]
assert any(
fmt == "pass_through_endpoint: upstream %s %s returned %s: %s" and '{"error": "half' in str(args[4])
for args, fmt in zip(recorded_warnings, formats)
), rendered
assert any(
fmt == "pass_through_endpoint: upstream error body read failed after %d bytes: %s"
and args[1] == 15
and args[2] == "ReadError"
for args, fmt in zip(recorded_warnings, formats)
), rendered
class _UpstreamErrorGzipStreamDropping(httpx.AsyncByteStream):
def __init__(self, flushed_prefix: bytes) -> None:
self._flushed_prefix: Final = flushed_prefix
async def __aiter__(self):
yield self._flushed_prefix
raise httpx.ReadError("peer reset")
@pytest.mark.asyncio
async def test_pass_through_request_streaming_upstream_error_gzip_read_failure_relays_decoded_partial():
"""
Regression: a mid-read failure on a gzip upstream must relay the decoded
plaintext, not the compressed bytes; the relay strips content-encoding so
raw compressed bytes would reach the client as garbage.
"""
plaintext: Final = b'{"error": "half'
compressor: Final = zlib.compressobj(level=6, wbits=31)
flushed_prefix: Final = compressor.compress(plaintext) + compressor.flush(zlib.Z_SYNC_FLUSH)
upstream_response: Final = httpx.Response(
status_code=502,
headers={"content-type": "application/json", "content-encoding": "gzip"},
stream=_UpstreamErrorGzipStreamDropping(flushed_prefix),
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
recorded_warnings: list[tuple] = []
real_warning: Final = verbose_proxy_logger.warning
def _recording_warning(*args, **kwargs):
if args and str(args[0]).startswith("pass_through_endpoint: upstream"):
recorded_warnings.append(args)
return real_warning(*args, **kwargs)
with patch.object(verbose_proxy_logger, "warning", side_effect=_recording_warning):
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
) as mock_get_client:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
) as mock_success_handler:
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler.return_value = None
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=True,
)
assert isinstance(response, StreamingResponse)
assert response.status_code == 502
assert "content-encoding" not in response.headers
streamed_chunks: Final = [chunk async for chunk in response.body_iterator]
streamed_bytes: Final = b"".join(
chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") for chunk in streamed_chunks
)
assert streamed_bytes == plaintext
await upstream_response.aclose()
rendered: Final = [str(args[0]) for args in recorded_warnings]
assert any(
args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s" and plaintext.decode() in str(args[4])
for args in recorded_warnings
), rendered
@pytest.mark.asyncio
async def test_pass_through_request_streaming_upstream_error_gzip_body_decoded_for_log_and_client():
upstream_content: Final = b'{"error": {"message": "gzipped upstream says the project was not found"}}'
compressed: Final = gzip.compress(upstream_content)
upstream_response: Final = httpx.Response(
status_code=502,
headers={"content-type": "text/event-stream", "content-encoding": "gzip"},
stream=_ChunkedUpstreamErrorBodyStream((compressed[:10], compressed[10:])),
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent"),
)
with patch.object(verbose_proxy_logger, "warning") as mock_warning:
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
) as mock_get_client:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler"
) as mock_success_handler:
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_success_handler.return_value = None
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
response: Final = await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:streamGenerateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
stream=True,
)
assert isinstance(response, StreamingResponse)
streamed_chunks: Final = [chunk async for chunk in response.body_iterator]
streamed_bytes: Final = b"".join(
chunk if isinstance(chunk, bytes) else chunk.encode("utf-8") for chunk in streamed_chunks
)
assert streamed_bytes == upstream_content
upstream_warnings: Final = [
call
for call in mock_warning.call_args_list
if call.args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s"
]
assert len(upstream_warnings) == 1, mock_warning.call_args_list
logged_body: Final = str(upstream_warnings[0].args[4])
assert "gzipped upstream says the project was not found" in logged_body
@pytest.mark.asyncio
async def test_pass_through_request_upstream_error_body_sanitized_against_log_forging():
upstream_content: Final = b'{"error": "line one"}\n2026-01-01 FAKE LOG LINE\x1b[31m'
upstream_response: Final = httpx.Response(
status_code=404,
headers={"content-type": "application/json"},
content=upstream_content,
request=httpx.Request("POST", "http://target-api.com/v1beta/models/claude-nope-9:generateContent"),
)
with patch.object(verbose_proxy_logger, "warning") as mock_warning:
with patch("litellm.proxy.proxy_server.proxy_logging_obj") as mock_proxy_logging:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client"
) as mock_get_client:
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.ProxyBaseLLMRequestProcessing"
) as mock_processing:
mock_proxy_logging.pre_call_hook = AsyncMock(return_value={})
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_processing.get_custom_headers.return_value = {}
async_client: Final = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
mock_get_client.return_value = MagicMock(client=async_client)
await pass_through_request(
request=_upstream_error_request(),
target="http://target-api.com/v1beta/models/claude-nope-9:generateContent",
custom_headers={},
user_api_key_dict=MagicMock(),
)
upstream_warnings: Final = [
call
for call in mock_warning.call_args_list
if call.args[0] == "pass_through_endpoint: upstream %s %s returned %s: %s"
]
assert len(upstream_warnings) == 1, mock_warning.call_args_list
logged_body: Final = str(upstream_warnings[0].args[4])
assert logged_body == '{"error": "line one"} 2026-01-01 FAKE LOG LINE [31m'
assert "\n" not in logged_body
assert "\x1b" not in logged_body
detail: Final = mock_proxy_logging.post_call_failure_hook.call_args.kwargs["original_exception"].detail
assert (
detail
== 'Upstream passthrough request failed with status 404: {"error": "line one"} 2026-01-01 FAKE LOG LINE [31m'
)
class _UpstreamDroppingMidStream(httpx.AsyncByteStream):
async def __aiter__(self):
yield b'data: {"id": "chatcmpl-1", "choices": [{"delta": {"content": "hi"}}]}\n\n'
@ -4287,7 +4991,9 @@ async def test_pass_through_request_claims_the_budget_reservation_only_when_its_
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_proxy_logging.post_call_response_headers_hook = AsyncMock(return_value=None)
mock_processing.get_custom_headers.return_value = {}
mock_worker.ensure_initialized_and_enqueue = MagicMock(side_effect=lambda async_coroutine: async_coroutine.close())
mock_worker.ensure_initialized_and_enqueue = MagicMock(
side_effect=lambda async_coroutine: async_coroutine.close()
)
async_client = MagicMock()
async_client.build_request = MagicMock(return_value=MagicMock())
async_client.send = AsyncMock(return_value=upstream_response)
@ -5304,9 +6010,7 @@ async def test_websocket_passthrough_propagates_active_trace_context(
mock_proxy_logging.post_call_success_hook = AsyncMock()
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mock_worker = MagicMock()
mock_worker.ensure_initialized_and_enqueue = MagicMock(
side_effect=lambda async_coroutine: async_coroutine.close()
)
mock_worker.ensure_initialized_and_enqueue = MagicMock(side_effect=lambda async_coroutine: async_coroutine.close())
monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging)
monkeypatch.setattr(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.connect",

View file

@ -18,6 +18,7 @@ from litellm.exceptions import GuardrailRaisedException
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import ProxyErrorTypes, UserAPIKeyAuth
from litellm.proxy.utils import ProxyLogging
from litellm.types.utils import CachingDetails
@pytest.fixture(autouse=True)
@ -249,6 +250,65 @@ async def test_post_call_failure_hook_keeps_router_stamped_metadata_for_post_cal
assert kwargs["standard_logging_object"]["model_id"] == "routed-deployment"
@pytest.mark.asyncio
async def test_post_call_failure_hook_keeps_deployment_attribution_for_cache_hit_post_call_failures(
proxy_logging, make_user_api_key_auth, monkeypatch
):
"""A post-call guardrail blocks a response served from the litellm cache. No provider call was made,
so ``first_api_call_start_time`` is unset, but the router did pick the deployment: the pre-routing
flag must stay off so ``litellm_deployment_failure_responses`` keeps its model_id and provider labels."""
from litellm.proxy import proxy_server
recorded: list[dict] = []
class _RecordingLogger(CustomLogger):
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
recorded.append(kwargs)
monkeypatch.setattr(
proxy_server,
"llm_router",
litellm.Router(
model_list=[
{
"model_name": "internal-model",
"litellm_params": {"model": "openai/gpt-4.1", "api_key": "sk-test"},
"model_info": {"id": "routed-deployment"},
}
]
),
)
monkeypatch.setattr(litellm, "callbacks", [_RecordingLogger()])
proxy_logging.alert_types = []
request_data = {
"litellm_call_id": "cache-hit-post-call-guardrail",
"model": "internal-model",
"messages": [{"role": "user", "content": "hi"}],
"metadata": {"model_info": {"id": "routed-deployment"}},
}
logging_obj, request_data = litellm.utils.function_setup(
original_function="acompletion", rules_obj=litellm.utils.Rules(), start_time=datetime.now(), **request_data
)
logging_obj.caching_details = CachingDetails(cache_hit=True, cache_duration_ms=1.0)
request_data["litellm_logging_obj"] = logging_obj
await proxy_logging.post_call_failure_hook(
request_data=request_data,
original_exception=GuardrailRaisedException(guardrail_name="g", message="response blocked"),
user_api_key_dict=make_user_api_key_auth(request_route="/chat/completions"),
route="/chat/completions",
)
assert len(recorded) == 1
kwargs = recorded[0]
assert PROXY_REJECTED_BEFORE_ROUTING_KEY not in kwargs["litellm_params"], kwargs["litellm_params"]
assert kwargs["standard_logging_object"]["model_id"] == "routed-deployment"
assert kwargs["standard_logging_object"]["custom_llm_provider"] == "openai"
assert kwargs["model"] == "internal-model"
assert kwargs["litellm_params"]["custom_llm_provider"] == "openai"
@pytest.mark.asyncio
async def test_post_call_failure_hook_flags_pre_routing_reject_despite_caller_model_info(
proxy_logging, make_user_api_key_auth, monkeypatch

View file

@ -1192,22 +1192,52 @@ def test_get_model_info_bedrock_regional_inference_profile_pricing(local_model_c
"""Regression LIT-4056: with the bedrock/ routing prefix (plain, converse/, or
invoke/), the exact regional cost-map entry must win over the region-stripped
base entry, matching the unprefixed control form."""
regional = litellm.model_cost["au.anthropic.claude-opus-4-8"]
base = litellm.model_cost["anthropic.claude-opus-4-8"]
regional = litellm.model_cost["eu.amazon.nova-pro-v1:0"]
base = litellm.model_cost["amazon.nova-pro-v1:0"]
assert regional["input_cost_per_token"] > base["input_cost_per_token"]
for model in (
"bedrock/au.anthropic.claude-opus-4-8",
"bedrock/converse/au.anthropic.claude-opus-4-8",
"bedrock/invoke/au.anthropic.claude-opus-4-8",
"bedrock/eu.amazon.nova-pro-v1:0",
"bedrock/converse/eu.amazon.nova-pro-v1:0",
"bedrock/invoke/eu.amazon.nova-pro-v1:0",
):
info = litellm.get_model_info(model=model)
assert info["key"] == "au.anthropic.claude-opus-4-8", model
assert info["key"] == "eu.amazon.nova-pro-v1:0", model
assert info["input_cost_per_token"] == regional["input_cost_per_token"], model
assert info["output_cost_per_token"] == regional["output_cost_per_token"], model
control = litellm.get_model_info(model="au.anthropic.claude-opus-4-8", custom_llm_provider="bedrock")
assert control["key"] == "au.anthropic.claude-opus-4-8"
control = litellm.get_model_info(model="eu.amazon.nova-pro-v1:0", custom_llm_provider="bedrock")
assert control["key"] == "eu.amazon.nova-pro-v1:0"
@pytest.mark.parametrize(
"bare_key",
[
"anthropic.claude-fable-5",
"anthropic.claude-fable-5-1",
"anthropic.claude-haiku-4-5-20251001-v1:0",
"anthropic.claude-opus-4-5-20251101-v1:0",
"anthropic.claude-opus-4-6-v1",
"anthropic.claude-opus-4-7",
"anthropic.claude-opus-4-8",
"anthropic.claude-opus-5",
"anthropic.claude-opus-5-5",
"anthropic.claude-sonnet-4-5-20250929-v1:0",
"anthropic.claude-sonnet-4-6",
"anthropic.claude-sonnet-5",
],
)
def test_bedrock_bare_claude_id_is_priced_global(local_model_cost_map, bare_key):
"""A bare Bedrock Claude id is billed at the Global SKU, so it carries the same
rate as its global. inference profile and sits below the regional us. rate."""
bare = litellm.model_cost[bare_key]
us = litellm.model_cost[f"us.{bare_key}"]
global_ = litellm.model_cost[f"global.{bare_key}"]
cost_fields = [f for f in bare if "cost" in f]
assert cost_fields
for field in cost_fields:
assert bare[field] == global_[field], field
assert bare["input_cost_per_token"] < us["input_cost_per_token"]
def test_get_model_info_bedrock_mantle_region_prefix_falls_back_to_the_mantle_row(local_model_cost_map):

View file

@ -0,0 +1,438 @@
from __future__ import annotations
import json
from collections.abc import Iterator
from pathlib import Path
from types import MappingProxyType
from typing import Final
import pytest
from tests.integration._support.routing import (
DIFF_FILE,
OBSERVED_FILE,
READER_ROLE,
WRITER_ROLE,
Mismatch,
Observation,
compare,
delta,
dump_observation,
load_either_role,
load_observation,
main,
normalize,
render,
role_calls,
)
TOKEN_QUERY: Final = 'UPDATE "LiteLLM_VerificationToken" SET token = $n WHERE token = $n'
NODE_ID: Final = "tests/integration/management/test_keys.py::test_generate"
def _routing(entries: dict[str, tuple[str, ...]]) -> MappingProxyType[str, frozenset[str]]:
return MappingProxyType({query: frozenset(roles) for query, roles in entries.items()})
def _observation(
queries: dict[str, tuple[str, ...]],
tests: dict[str, dict[str, tuple[str, ...]]] | None = None,
calls: dict[str, int] | None = None,
dealloc: int = 0,
) -> Observation:
return Observation(
_routing(queries),
MappingProxyType({node: _routing(mapping) for node, mapping in (tests or {}).items()}),
MappingProxyType(calls if calls is not None else {"litellm_reader": 3, "litellm_writer": 7}),
dealloc,
)
@pytest.mark.parametrize(
("raw", "expected"),
[
("SELECT a\n FROM t", "SELECT a FROM t"),
("SELECT * FROM t WHERE id IN ($1, $2, $3)", "SELECT * FROM t WHERE id IN ($n)"),
("SELECT * FROM t WHERE id IN ($1,$2)", "SELECT * FROM t WHERE id IN ($n)"),
("SELECT * FROM t WHERE id IN ($4)", "SELECT * FROM t WHERE id IN ($n)"),
(
"INSERT INTO t VALUES ($1, $2) ON CONFLICT ($3, $4, $5) DO NOTHING",
"INSERT INTO t VALUES ($n) ON CONFLICT ($n) DO NOTHING",
),
],
)
def test_normalize_collapses_whitespace_and_placeholders(raw: str, expected: str) -> None:
assert normalize(raw) == expected
def test_compare_reports_global_role_mismatch() -> None:
base: Final = _observation({TOKEN_QUERY: ("litellm_reader",), "SELECT 1": ("litellm_writer",)})
head: Final = _observation({TOKEN_QUERY: ("litellm_writer",), "SELECT 1": ("litellm_writer",)})
report: Final = compare(base, head)
assert report.mismatches == (Mismatch(None, TOKEN_QUERY, ("litellm_reader",), ("litellm_writer",)),)
assert report.failures() == (f"global: {TOKEN_QUERY}: base [litellm_reader] head [litellm_writer]",)
def test_compare_reports_global_shrink_mismatch() -> None:
base: Final = _observation({TOKEN_QUERY: ("litellm_reader", "litellm_writer")})
head: Final = _observation({TOKEN_QUERY: ("litellm_writer",)})
report: Final = compare(base, head)
assert report.mismatches == (
Mismatch(None, TOKEN_QUERY, ("litellm_reader", "litellm_writer"), ("litellm_writer",)),
)
assert report.failures() == (f"global: {TOKEN_QUERY}: base [litellm_reader, litellm_writer] head [litellm_writer]",)
def test_compare_reports_per_test_mismatch_with_nodeid() -> None:
base: Final = _observation(
{TOKEN_QUERY: ("litellm_reader",)},
{NODE_ID: {TOKEN_QUERY: ("litellm_reader",)}},
)
head: Final = _observation(
{TOKEN_QUERY: ("litellm_reader",)},
{NODE_ID: {TOKEN_QUERY: ("litellm_writer",)}},
)
report: Final = compare(base, head)
assert report.mismatches == (Mismatch(NODE_ID, TOKEN_QUERY, ("litellm_reader",), ("litellm_writer",)),)
assert report.failures() == (f"{NODE_ID}: {TOKEN_QUERY}: base [litellm_reader] head [litellm_writer]",)
def test_compare_per_test_mismatch_ignores_global_observation() -> None:
base: Final = _observation(
{TOKEN_QUERY: ("litellm_reader", "litellm_writer")},
{NODE_ID: {TOKEN_QUERY: ("litellm_reader",)}},
)
head: Final = _observation(
{TOKEN_QUERY: ("litellm_reader", "litellm_writer")},
{NODE_ID: {TOKEN_QUERY: ("litellm_writer",)}},
)
report: Final = compare(base, head)
assert report.mismatches == (Mismatch(NODE_ID, TOKEN_QUERY, ("litellm_reader",), ("litellm_writer",)),)
assert report.failures() == (f"{NODE_ID}: {TOKEN_QUERY}: base [litellm_reader] head [litellm_writer]",)
def test_compare_reports_per_test_shrink_mismatch() -> None:
base: Final = _observation(
{TOKEN_QUERY: ("litellm_reader", "litellm_writer")},
{NODE_ID: {TOKEN_QUERY: ("litellm_reader", "litellm_writer")}},
)
head: Final = _observation(
{TOKEN_QUERY: ("litellm_reader", "litellm_writer")},
{NODE_ID: {TOKEN_QUERY: ("litellm_reader",)}},
)
report: Final = compare(base, head)
assert report.mismatches == (
Mismatch(NODE_ID, TOKEN_QUERY, ("litellm_reader", "litellm_writer"), ("litellm_reader",)),
)
assert report.failures() == (
f"{NODE_ID}: {TOKEN_QUERY}: base [litellm_reader, litellm_writer] head [litellm_reader]",
)
def test_compare_reports_global_gain_mismatch() -> None:
base: Final = _observation({TOKEN_QUERY: ("litellm_reader",)})
head: Final = _observation({TOKEN_QUERY: ("litellm_reader", "litellm_writer")})
report: Final = compare(base, head)
assert report.mismatches == (
Mismatch(None, TOKEN_QUERY, ("litellm_reader",), ("litellm_reader", "litellm_writer")),
)
assert report.failures() == (f"global: {TOKEN_QUERY}: base [litellm_reader] head [litellm_reader, litellm_writer]",)
def test_compare_reports_per_test_gain_mismatch() -> None:
base: Final = _observation(
{TOKEN_QUERY: ("litellm_reader",)},
{NODE_ID: {TOKEN_QUERY: ("litellm_reader",)}},
)
head: Final = _observation(
{TOKEN_QUERY: ("litellm_reader",)},
{NODE_ID: {TOKEN_QUERY: ("litellm_reader", "litellm_writer")}},
)
report: Final = compare(base, head)
assert report.mismatches == (
Mismatch(NODE_ID, TOKEN_QUERY, ("litellm_reader",), ("litellm_reader", "litellm_writer")),
)
assert report.failures() == (
f"{NODE_ID}: {TOKEN_QUERY}: base [litellm_reader] head [litellm_reader, litellm_writer]",
)
def test_compare_either_role_suppresses_and_reports_variance() -> None:
base: Final = _observation(
{TOKEN_QUERY: ("litellm_reader", "litellm_writer"), "SELECT quiet": ("litellm_reader",)},
{NODE_ID: {TOKEN_QUERY: ("litellm_reader",)}},
)
head: Final = _observation(
{TOKEN_QUERY: ("litellm_writer",), "SELECT quiet": ("litellm_reader",)},
{NODE_ID: {TOKEN_QUERY: ("litellm_writer",)}},
)
report: Final = compare(base, head, either_role=frozenset({TOKEN_QUERY, "SELECT quiet"}))
assert report.mismatches == ()
assert report.failures() == ()
assert report.either_role == (TOKEN_QUERY,)
assert "== either role ==\n" + TOKEN_QUERY + "\n" in render(report)
def test_compare_either_role_matches_exact_keys_only() -> None:
base: Final = _observation(
{
"SELECT $n": ("litellm_reader",),
"SELECT $n FROM x": ("litellm_reader",),
"SELECT $n FROM x WHERE y = $n": ("litellm_reader",),
}
)
head: Final = _observation(
{
"SELECT $n": ("litellm_writer",),
"SELECT $n FROM x": ("litellm_writer",),
"SELECT $n FROM x WHERE y = $n": ("litellm_writer",),
}
)
report: Final = compare(base, head, either_role=frozenset({"SELECT $n FROM x"}))
assert frozenset(mismatch.query for mismatch in report.mismatches) == frozenset(
{"SELECT $n", "SELECT $n FROM x WHERE y = $n"}
)
other: Final = compare(base, head, either_role=frozenset({"SELECT $n"}))
assert frozenset(mismatch.query for mismatch in other.mismatches) == frozenset(
{"SELECT $n FROM x", "SELECT $n FROM x WHERE y = $n"}
)
def test_compare_one_sided_queries_are_listed_not_failed() -> None:
base: Final = _observation({"SELECT a": ("litellm_reader",), "SELECT gone": ("litellm_writer",)})
head: Final = _observation({"SELECT a": ("litellm_reader",), "SELECT new": ("litellm_writer",)})
report: Final = compare(base, head)
assert report.only_base == ("SELECT gone",)
assert report.only_head == ("SELECT new",)
assert report.mismatches == ()
assert report.failures() == ()
def test_failures_flags_dealloc_evictions_on_base() -> None:
report: Final = compare(_observation({}, dealloc=1), _observation({}))
assert report.failures() == ("base: pg_stat_statements evicted 1 entries (dealloc > 0)",)
def test_failures_flags_dealloc_evictions_on_head() -> None:
report: Final = compare(_observation({}), _observation({}, dealloc=1))
assert report.failures() == ("head: pg_stat_statements evicted 1 entries (dealloc > 0)",)
def test_failures_flags_silent_reader_on_base() -> None:
report: Final = compare(
_observation({}, calls={"litellm_reader": 0, "litellm_writer": 5}),
_observation({}),
)
assert report.failures() == ("base: no litellm_reader calls observed",)
def test_failures_flags_silent_reader_on_head() -> None:
report: Final = compare(
_observation({}),
_observation({}, calls={"litellm_reader": 0, "litellm_writer": 5}),
)
assert report.failures() == ("head: no litellm_reader calls observed",)
def test_failures_flags_silent_writer() -> None:
report: Final = compare(
_observation({}, calls={"litellm_reader": 5, "litellm_writer": 0}),
_observation({}),
)
assert report.failures() == ("base: no litellm_writer calls observed",)
def test_failures_counts_missing_role_as_silent() -> None:
report: Final = compare(_observation({}), _observation({}, calls={"litellm_writer": 5}))
assert report.failures() == ("head: no litellm_reader calls observed",)
def test_compare_skips_per_test_mismatches_for_xdist_shape() -> None:
base: Final = _observation(
{TOKEN_QUERY: ("litellm_reader",)},
{NODE_ID: {TOKEN_QUERY: ("litellm_reader",)}},
)
head: Final = _observation({TOKEN_QUERY: ("litellm_reader",)})
assert head.tests == {}
report: Final = compare(base, head)
assert report.mismatches == ()
assert report.failures() == ()
class _WriterFirst(frozenset[str]):
def __iter__(self) -> Iterator[str]:
return iter((WRITER_ROLE, READER_ROLE))
def test_dump_observation_sorts_role_lists_and_round_trips(tmp_path: Path) -> None:
queries: Final = [f"SELECT {index}" for index in range(4)]
observation: Final = Observation(
MappingProxyType({query: _WriterFirst({WRITER_ROLE, READER_ROLE}) for query in queries}),
MappingProxyType(
{NODE_ID: MappingProxyType({query: _WriterFirst({WRITER_ROLE, READER_ROLE}) for query in queries})}
),
MappingProxyType({READER_ROLE: 1, WRITER_ROLE: 2}),
0,
)
expected: Final = (
json.dumps(
{
"queries": {query: ["litellm_reader", "litellm_writer"] for query in queries},
"tests": {NODE_ID: {query: ["litellm_reader", "litellm_writer"] for query in queries}},
"calls": {"litellm_reader": 1, "litellm_writer": 2},
"dealloc": 0,
},
sort_keys=True,
indent=2,
)
+ "\n"
)
dumped: Final = dump_observation(observation)
assert dumped == expected
path: Final = tmp_path / OBSERVED_FILE
path.write_text(dumped)
loaded: Final = load_observation(path)
assert loaded.queries == _routing({query: (WRITER_ROLE, READER_ROLE) for query in queries})
assert loaded.tests == {NODE_ID: loaded.queries}
def _write_observed(results: Path, observation: Observation) -> None:
results.mkdir(parents=True, exist_ok=True)
(results / OBSERVED_FILE).write_text(dump_observation(observation))
def test_main_check_returns_zero_for_matching_routes(tmp_path: Path) -> None:
base_dir: Final = tmp_path / "base"
head_dir: Final = tmp_path / "head"
observation: Final = _observation({TOKEN_QUERY: ("litellm_reader",)})
_write_observed(base_dir, observation)
_write_observed(head_dir, observation)
assert main(["check", str(base_dir), str(head_dir)]) == 0
diff: Final = (tmp_path / DIFF_FILE).read_text()
assert "== failures ==\nnone\n" in diff
def test_main_check_returns_one_and_writes_exact_diff(tmp_path: Path) -> None:
base_dir: Final = tmp_path / "parity" / "base"
head_dir: Final = tmp_path / "parity" / "head"
_write_observed(
base_dir,
_observation({TOKEN_QUERY: ("litellm_reader",), "SELECT absent": ("litellm_writer",)}),
)
_write_observed(
head_dir,
_observation(
{TOKEN_QUERY: ("litellm_writer",)},
calls={"litellm_reader": 0, "litellm_writer": 5},
dealloc=2,
),
)
assert main(["check", str(base_dir), str(head_dir)]) == 1
assert (head_dir.parent / DIFF_FILE).read_text() == (
"== failures ==\n"
f"global: {TOKEN_QUERY}: base [litellm_reader] head [litellm_writer]\n"
"head: pg_stat_statements evicted 2 entries (dealloc > 0)\n"
"head: no litellm_reader calls observed\n"
"\n"
"== either role ==\n"
"none\n"
"\n"
"== queries only in base ==\n"
"SELECT absent\n"
"\n"
"== queries only in head ==\n"
"none\n"
"\n"
"== calls ==\n"
"base litellm_reader: 3\n"
"base litellm_writer: 7\n"
"base dealloc: 0\n"
"head litellm_reader: 0\n"
"head litellm_writer: 5\n"
"head dealloc: 2\n"
)
def test_main_check_missing_observed_returns_one(tmp_path: Path, capsys: pytest.CaptureFixture[str]) -> None:
base_dir: Final = tmp_path / "base"
head_dir: Final = tmp_path / "head"
_write_observed(base_dir, _observation({}))
head_dir.mkdir()
assert main(["check", str(base_dir), str(head_dir)]) == 1
assert "observed routing file missing" in capsys.readouterr().err
def test_main_check_either_role_suppresses_shrink(tmp_path: Path) -> None:
base_dir: Final = tmp_path / "base"
head_dir: Final = tmp_path / "head"
_write_observed(base_dir, _observation({TOKEN_QUERY: ("litellm_reader", "litellm_writer")}))
_write_observed(head_dir, _observation({TOKEN_QUERY: ("litellm_writer",)}))
argv: Final = ["check", str(base_dir), str(head_dir)]
allowlist: Final = tmp_path / "either.json"
allowlist.write_text(json.dumps({TOKEN_QUERY: "timer probe may use either pool"}))
assert main([*argv, "--either-role", str(allowlist)]) == 0
assert "== either role ==\n" + TOKEN_QUERY + "\n" in (tmp_path / DIFF_FILE).read_text()
assert main(argv) == 1
def test_delta_maps_positive_increases_per_role() -> None:
before: Final = MappingProxyType(
{
("litellm_reader", "SELECT both"): 1,
("litellm_writer", "SELECT both"): 2,
("litellm_reader", "SELECT reader"): 3,
("litellm_writer", "SELECT gone"): 4,
("litellm_reader", "SELECT same"): 5,
}
)
after: Final = MappingProxyType(
{
("litellm_reader", "SELECT both"): 2,
("litellm_writer", "SELECT both"): 5,
("litellm_reader", "SELECT reader"): 6,
("litellm_reader", "SELECT same"): 5,
("litellm_writer", "SELECT writer"): 7,
}
)
assert delta(before, after) == {
"SELECT both": frozenset({"litellm_reader", "litellm_writer"}),
"SELECT reader": frozenset({"litellm_reader"}),
"SELECT writer": frozenset({"litellm_writer"}),
}
def test_role_calls_sums_positive_increases_per_role() -> None:
before: Final = MappingProxyType(
{
("litellm_reader", "SELECT a"): 10,
("litellm_reader", "SELECT b"): 4,
("litellm_writer", "SELECT a"): 1,
}
)
after: Final = MappingProxyType(
{
("litellm_reader", "SELECT a"): 11,
("litellm_reader", "SELECT b"): 2,
("litellm_writer", "SELECT a"): 1,
("litellm_writer", "SELECT c"): 6,
}
)
assert role_calls(before, after) == {"litellm_reader": 1, "litellm_writer": 6}
def test_load_either_role_missing_path_returns_empty(tmp_path: Path) -> None:
assert load_either_role(tmp_path / "absent.json") == frozenset()
def test_load_either_role_reads_query_keys(tmp_path: Path) -> None:
path: Final = tmp_path / "either.json"
path.write_text(json.dumps({"SELECT $n": "probe", "SELECT now()": "clock"}))
assert load_either_role(path) == frozenset({"SELECT $n", "SELECT now()"})
def test_load_observation_reads_calls_and_dealloc(tmp_path: Path) -> None:
observation: Final = _observation({TOKEN_QUERY: ("litellm_reader",)}, dealloc=0)
path: Final = tmp_path / OBSERVED_FILE
path.write_text(dump_observation(observation))
loaded: Final = load_observation(path)
assert loaded == observation

View file

@ -1,7 +1,5 @@
import pytest
from litellm.llms.fireworks_ai.common_utils import resolve_fireworks_resource_name
@ -16,6 +14,10 @@ from litellm.llms.fireworks_ai.common_utils import resolve_fireworks_resource_na
("glm-4p6", "accounts/fireworks/models/glm-4p6"),
("fireworks_ai/glm-4p6", "accounts/fireworks/models/glm-4p6"),
("kimi-k2p6-fast", "accounts/fireworks/routers/kimi-k2p6-fast"),
("firerouter", "accounts/fireworks/routers/firerouter"),
("fireworks_ai/firerouter", "accounts/fireworks/routers/firerouter"),
("firerouter/kimi-k3/deepseek-v4", "accounts/fireworks/routers/firerouter/kimi-k3/deepseek-v4"),
("firerouter-v2", "accounts/fireworks/models/firerouter-v2"),
(
"accounts/fireworks/routers/glm-latest",
"accounts/fireworks/routers/glm-latest",

Some files were not shown because too many files have changed in this diff Show more