mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
chore: merge main into litellm_responses_precall_block_stream
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
commit
1fd0658a88
101 changed files with 10853 additions and 280 deletions
|
|
@ -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
|
||||
|
|
|
|||
52
.circleci/scripts/prepare_replica_roles.py
Normal file
52
.circleci/scripts/prepare_replica_roles.py
Normal 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()
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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>,
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"):
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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-")
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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__
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
)
|
||||
]
|
||||
|
||||
|
|
|
|||
|
|
@ -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] = {
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
333
tests/integration/_support/routing.py
Normal file
333
tests/integration/_support/routing.py
Normal 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:]))
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
154
tests/integration/configuration/test_callback_settings_boot.py
Normal file
154
tests/integration/configuration/test_callback_settings_boot.py
Normal 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)
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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":
|
||||
|
|
|
|||
344
tests/integration/observability/_s3_v2_support.py
Normal file
344
tests/integration/observability/_s3_v2_support.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
@ -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})
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
97
tests/integration/observability/test_s3_v2_flush_surfaces.py
Normal file
97
tests/integration/observability/test_s3_v2_flush_surfaces.py
Normal 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"
|
||||
630
tests/integration/observability/test_s3_v2_upload_fanout.py
Normal file
630
tests/integration/observability/test_s3_v2_upload_fanout.py
Normal 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"
|
||||
244
tests/integration/pricing/test_model_listing_token_limits.py
Normal file
244
tests/integration/pricing/test_model_listing_token_limits.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
3
tests/integration/routing/either_role.json
Normal file
3
tests/integration/routing/either_role.json
Normal 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"
|
||||
}
|
||||
|
|
@ -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"):
|
||||
|
|
|
|||
|
|
@ -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({}); "
|
||||
|
|
|
|||
463
tests/integration/spend/test_normalized_error_long_message.py
Normal file
463
tests/integration/spend/test_normalized_error_long_message.py
Normal 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
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
0
tests/unit/integration_support/__init__.py
Normal file
0
tests/unit/integration_support/__init__.py
Normal file
438
tests/unit/integration_support/test_routing.py
Normal file
438
tests/unit/integration_support/test_routing.py
Normal 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
|
||||
|
|
@ -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
Loading…
Add table
Reference in a new issue