mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
Merge remote-tracking branch 'origin/main' into litellm_replica_db_opt_in
This commit is contained in:
commit
9a38c429ce
28 changed files with 3072 additions and 40 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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -6159,6 +6159,7 @@
|
|||
]
|
||||
},
|
||||
"azure/gpt-realtime-whisper": {
|
||||
"deprecation_date": "2027-05-06",
|
||||
"input_cost_per_second": 0.0002833333333333333,
|
||||
"litellm_provider": "azure",
|
||||
"mode": "audio_transcription",
|
||||
|
|
@ -8060,6 +8061,7 @@
|
|||
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-05,
|
||||
"cache_read_input_token_cost": 1e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 2e-06,
|
||||
"deprecation_date": "2028-01-11",
|
||||
"input_cost_per_token": 1e-05,
|
||||
"input_cost_per_token_above_272k_tokens": 2e-05,
|
||||
"litellm_provider": "azure",
|
||||
|
|
@ -8108,6 +8110,7 @@
|
|||
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-05,
|
||||
"cache_read_input_token_cost": 1e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 2e-06,
|
||||
"deprecation_date": "2028-01-11",
|
||||
"input_cost_per_token": 1e-05,
|
||||
"input_cost_per_token_above_272k_tokens": 2e-05,
|
||||
"litellm_provider": "azure",
|
||||
|
|
@ -8156,6 +8159,7 @@
|
|||
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-07,
|
||||
"cache_read_input_token_cost": 1e-08,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 2e-08,
|
||||
"deprecation_date": "2028-03-11",
|
||||
"input_cost_per_token": 1e-07,
|
||||
"input_cost_per_token_above_272k_tokens": 2e-07,
|
||||
"litellm_provider": "azure",
|
||||
|
|
@ -8204,6 +8208,7 @@
|
|||
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-07,
|
||||
"cache_read_input_token_cost": 1e-08,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 2e-08,
|
||||
"deprecation_date": "2028-03-11",
|
||||
"input_cost_per_token": 1e-07,
|
||||
"input_cost_per_token_above_272k_tokens": 2e-07,
|
||||
"litellm_provider": "azure",
|
||||
|
|
@ -8252,6 +8257,7 @@
|
|||
"cache_creation_input_token_cost_above_272k_tokens": 5e-06,
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 4e-07,
|
||||
"deprecation_date": "2028-03-11",
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 4e-06,
|
||||
"litellm_provider": "azure",
|
||||
|
|
@ -8300,6 +8306,7 @@
|
|||
"cache_creation_input_token_cost_above_272k_tokens": 5e-06,
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 4e-07,
|
||||
"deprecation_date": "2028-03-11",
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 4e-06,
|
||||
"litellm_provider": "azure",
|
||||
|
|
@ -56802,7 +56809,8 @@
|
|||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_xhigh_reasoning_effort": true,
|
||||
"supports_web_search": true
|
||||
"supports_web_search": true,
|
||||
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-56-sol.html"
|
||||
},
|
||||
"bedrock_mantle/openai.gpt-5.6-terra": {
|
||||
"input_cost_per_token": 2.2e-06,
|
||||
|
|
@ -56844,7 +56852,8 @@
|
|||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_xhigh_reasoning_effort": true,
|
||||
"supports_web_search": true
|
||||
"supports_web_search": true,
|
||||
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-56-terra.html"
|
||||
},
|
||||
"bedrock_mantle/openai.gpt-5.6-cyber": {
|
||||
"input_cost_per_token": 1.375e-05,
|
||||
|
|
@ -56951,7 +56960,8 @@
|
|||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_xhigh_reasoning_effort": true,
|
||||
"supports_web_search": true
|
||||
"supports_web_search": true,
|
||||
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-56-luna.html"
|
||||
},
|
||||
"us.openai.gpt-5.6-sol": {
|
||||
"input_cost_per_token": 4.4e-06,
|
||||
|
|
@ -57722,7 +57732,8 @@
|
|||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
"supports_vision": true,
|
||||
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-xai-grok-4-6.html"
|
||||
},
|
||||
"bedrock_mantle/anthropic.claude-haiku-4-5": {
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -6159,6 +6159,7 @@
|
|||
]
|
||||
},
|
||||
"azure/gpt-realtime-whisper": {
|
||||
"deprecation_date": "2027-05-06",
|
||||
"input_cost_per_second": 0.0002833333333333333,
|
||||
"litellm_provider": "azure",
|
||||
"mode": "audio_transcription",
|
||||
|
|
@ -8060,6 +8061,7 @@
|
|||
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-05,
|
||||
"cache_read_input_token_cost": 1e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 2e-06,
|
||||
"deprecation_date": "2028-01-11",
|
||||
"input_cost_per_token": 1e-05,
|
||||
"input_cost_per_token_above_272k_tokens": 2e-05,
|
||||
"litellm_provider": "azure",
|
||||
|
|
@ -8108,6 +8110,7 @@
|
|||
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-05,
|
||||
"cache_read_input_token_cost": 1e-06,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 2e-06,
|
||||
"deprecation_date": "2028-01-11",
|
||||
"input_cost_per_token": 1e-05,
|
||||
"input_cost_per_token_above_272k_tokens": 2e-05,
|
||||
"litellm_provider": "azure",
|
||||
|
|
@ -8156,6 +8159,7 @@
|
|||
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-07,
|
||||
"cache_read_input_token_cost": 1e-08,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 2e-08,
|
||||
"deprecation_date": "2028-03-11",
|
||||
"input_cost_per_token": 1e-07,
|
||||
"input_cost_per_token_above_272k_tokens": 2e-07,
|
||||
"litellm_provider": "azure",
|
||||
|
|
@ -8204,6 +8208,7 @@
|
|||
"cache_creation_input_token_cost_above_272k_tokens": 2.5e-07,
|
||||
"cache_read_input_token_cost": 1e-08,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 2e-08,
|
||||
"deprecation_date": "2028-03-11",
|
||||
"input_cost_per_token": 1e-07,
|
||||
"input_cost_per_token_above_272k_tokens": 2e-07,
|
||||
"litellm_provider": "azure",
|
||||
|
|
@ -8252,6 +8257,7 @@
|
|||
"cache_creation_input_token_cost_above_272k_tokens": 5e-06,
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 4e-07,
|
||||
"deprecation_date": "2028-03-11",
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 4e-06,
|
||||
"litellm_provider": "azure",
|
||||
|
|
@ -8300,6 +8306,7 @@
|
|||
"cache_creation_input_token_cost_above_272k_tokens": 5e-06,
|
||||
"cache_read_input_token_cost": 2e-07,
|
||||
"cache_read_input_token_cost_above_272k_tokens": 4e-07,
|
||||
"deprecation_date": "2028-03-11",
|
||||
"input_cost_per_token": 2e-06,
|
||||
"input_cost_per_token_above_272k_tokens": 4e-06,
|
||||
"litellm_provider": "azure",
|
||||
|
|
@ -56802,7 +56809,8 @@
|
|||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_xhigh_reasoning_effort": true,
|
||||
"supports_web_search": true
|
||||
"supports_web_search": true,
|
||||
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-56-sol.html"
|
||||
},
|
||||
"bedrock_mantle/openai.gpt-5.6-terra": {
|
||||
"input_cost_per_token": 2.2e-06,
|
||||
|
|
@ -56844,7 +56852,8 @@
|
|||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_xhigh_reasoning_effort": true,
|
||||
"supports_web_search": true
|
||||
"supports_web_search": true,
|
||||
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-56-terra.html"
|
||||
},
|
||||
"bedrock_mantle/openai.gpt-5.6-cyber": {
|
||||
"input_cost_per_token": 1.375e-05,
|
||||
|
|
@ -56951,7 +56960,8 @@
|
|||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_xhigh_reasoning_effort": true,
|
||||
"supports_web_search": true
|
||||
"supports_web_search": true,
|
||||
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-56-luna.html"
|
||||
},
|
||||
"us.openai.gpt-5.6-sol": {
|
||||
"input_cost_per_token": 4.4e-06,
|
||||
|
|
@ -57722,7 +57732,8 @@
|
|||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
"supports_vision": true,
|
||||
"source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-xai-grok-4-6.html"
|
||||
},
|
||||
"bedrock_mantle/anthropic.claude-haiku-4-5": {
|
||||
"cache_creation_input_token_cost": 1.25e-06,
|
||||
|
|
|
|||
|
|
@ -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:]))
|
||||
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)
|
||||
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"
|
||||
|
|
@ -1,9 +1,93 @@
|
|||
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"]
|
||||
|
|
@ -27,3 +111,134 @@ def test_v1_models_carries_deployment_model_info_limits_for_an_unknown_model(gat
|
|||
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)
|
||||
|
|
|
|||
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({}); "
|
||||
|
|
|
|||
|
|
@ -196,12 +196,14 @@ def _proxy_with_one_seeded_row(
|
|||
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": 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(
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
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
|
||||
Loading…
Add table
Reference in a new issue