merge: origin/main into litellm_fal_gpt_image_25_flux_dev_edits

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
kerry 2026-09-20 10:15:52 +00:00
commit c3193a8e32
165 changed files with 2549 additions and 1390 deletions

View file

@ -1,10 +1,33 @@
version: 2.1
parameters:
run_migration_tests:
type: boolean
default: false
migration_candidate_image:
type: string
default: ""
migration_source_sha:
type: string
default: ""
orbs:
codecov: codecov/codecov@4.0.1
node: circleci/node@5.1.0 # Add this line to declare the node orb
win: circleci/windows@5.0 # Add Windows orb
commands:
checkout_migration_source:
steps:
- run:
name: Select the requested migration test revision
environment:
MIGRATION_SOURCE_SHA: << pipeline.parameters.migration_source_sha >>
command: |
revision="${MIGRATION_SOURCE_SHA:-$CIRCLE_SHA1}"
[[ "$revision" =~ ^[0-9a-f]{40}$ ]] || exit 1
git init
git remote add origin https://github.com/BerriAI/litellm.git
git fetch --depth 1 origin "$revision"
git checkout --detach FETCH_HEAD
skip_if_unrelated_changes:
parameters:
category:
@ -2871,20 +2894,41 @@ jobs:
destination: e2e-server-root-path-playwright-report
build_docker_database_image:
parameters:
migration_qualification:
type: boolean
default: false
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
steps:
- checkout
- skip_if_unrelated_changes
- when:
condition: << parameters.migration_qualification >>
steps:
- checkout_migration_source
- unless:
condition: << parameters.migration_qualification >>
steps:
- checkout
- skip_if_unrelated_changes
- run:
name: Build Docker image
environment:
MIGRATION_CANDIDATE_IMAGE: << pipeline.parameters.migration_candidate_image >>
command: |
docker build \
-t litellm-docker-database:ci \
-f docker/Dockerfile.database .
if [ -n "$MIGRATION_CANDIDATE_IMAGE" ]; then
[[ "$MIGRATION_CANDIDATE_IMAGE" =~ ^ghcr.io/berriai/[a-z0-9._/-]+@sha256:[0-9a-f]{64}$ ]] || exit 1
docker pull "$MIGRATION_CANDIDATE_IMAGE"
docker tag "$MIGRATION_CANDIDATE_IMAGE" litellm-docker-database:ci
else
docker build \
--label org.opencontainers.image.revision="$(git rev-parse HEAD)" \
-t litellm-docker-database:ci \
-f docker/Dockerfile.database .
fi
python3 .circleci/scripts/run_migration_tests.py record-image
- run:
name: Save Docker image to workspace root
@ -2895,6 +2939,78 @@ jobs:
root: .
paths:
- litellm-docker-database.tar.zst
- migration-image.json
migration_startup_tests:
parameters:
suite:
type: enum
enum: [startup, recovery, legacy]
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
environment:
LITELLM_MIGRATION_TESTS: "1"
LITELLM_MIGRATION_TEST_IMAGE: litellm-docker-database:ci
MIGRATION_TEST_ADMIN_URL: postgresql://postgres:postgres@127.0.0.1:5432/postgres
MIGRATION_TEST_CONTAINER_ADMIN_URL: postgresql://postgres:postgres@host.docker.internal:5432/postgres
MIGRATION_TEST_OUTPUT: /tmp/migration-results
PYTHONPATH: tests/e2e
steps:
- checkout_migration_source
- install_uv
- install_rust
- restore_cache:
keys:
- v1-uv-cache-{{ checksum "uv.lock" }}
- run:
name: Install test dependencies
command: uv sync --frozen --all-groups --all-extras --python 3.12
- attach_workspace:
at: ~/project
- run:
name: Load the shared candidate and start PostgreSQL
command: |
zstd -d litellm-docker-database.tar.zst --stdout | docker load
docker run -d --name migration-postgres \
-e POSTGRES_USER=postgres -e POSTGRES_PASSWORD=postgres \
-p 5432:5432 \
postgres:16@sha256:e17e86066e5ef83e0952a9347f5c792b7ece00972e2aa787a6986f471b3dd3d5
- wait_for_service:
url: tcp://localhost:5432
timeout: "60"
- run:
name: Run migration startup regressions
environment:
MIGRATION_TEST_SUITE: << parameters.suite >>
MIGRATION_CANDIDATE_IMAGE: << pipeline.parameters.migration_candidate_image >>
command: |
mkdir -p /tmp/migration-results
uv run --no-sync python .circleci/scripts/run_migration_tests.py
no_output_timeout: 15m
- store_test_results:
path: /tmp/migration-results/junit
- run:
name: Package migration diagnostics
when: always
command: |
mkdir -p /tmp/migration-artifacts
if [ -d /tmp/migration-results ]; then
tar -czf /tmp/migration-artifacts/diagnostics.tar.gz -C /tmp/migration-results .
if [ -f /tmp/migration-results/verdict.json ]; then
cp /tmp/migration-results/verdict.json /tmp/migration-artifacts/verdict.json
fi
fi
- store_artifacts:
path: /tmp/migration-artifacts
destination: migration-results
- run:
name: Remove migration test containers
when: always
command: |
docker ps -aq --filter label=litellm-migration-test=true | xargs -r docker rm -f
docker rm -f migration-postgres || true
test_bad_database_url:
machine:
@ -3026,8 +3142,61 @@ jobs:
- store_artifacts:
path: test-results
unit:
machine:
image: ubuntu-2204:2024.04.1
resource_class: large
working_directory: ~/project
steps:
- setup_litellm_test_deps
- run:
name: Generate Prisma client
command: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
- run:
name: Run unit tests
command: |
mkdir -p test-results/unit
mapfile -t files < <(find tests/unit -name 'test_*.py' | sort)
if [ "${#files[@]}" -eq 0 ]; then echo "tests/unit holds no test_*.py files; nothing to run"; exit 0; fi
set +e
LITELLM_LOCAL_MODEL_COST_MAP=True uv run --no-sync pytest "${files[@]}" -p no:rerunfailures -p no:pytest-retry --timeout=90 -n 4 --dist=loadscope --tb=short --junitxml=test-results/unit/junit.xml
status=$?
set -e
if [ "$status" -eq 5 ]; then echo "pytest collected no tests from tests/unit; passing"; exit 0; fi
exit "$status"
- store_test_results:
path: test-results
- store_artifacts:
path: test-results
workflows:
migration_startup:
when: << pipeline.parameters.run_migration_tests >>
jobs: &migration_jobs
- build_docker_database_image:
migration_qualification: true
- migration_startup_tests:
name: migration-startup
suite: startup
requires: [build_docker_database_image]
- migration_startup_tests:
name: migration-recovery
suite: recovery
requires: [build_docker_database_image]
- migration_startup_tests:
name: migration-legacy-and-pooling
suite: legacy
requires: [build_docker_database_image]
migration_startup_scheduled:
triggers:
- schedule:
cron: "17 0,6,12,18 * * *"
filters:
branches:
only: litellm_internal_staging
jobs: *migration_jobs
integration:
unless: << pipeline.parameters.run_migration_tests >>
jobs:
- integration_contracts:
name: integration-<< matrix.suite >>
@ -3040,6 +3209,7 @@ workflows:
- main
- /litellm_.*/
build_and_test:
unless: << pipeline.parameters.run_migration_tests >>
jobs:
- using_litellm_on_windows:
filters: &main_branches
@ -3047,6 +3217,8 @@ workflows:
only:
- main
- /litellm_.*/
- unit:
filters: *main_branches
- provider_replay_harness
- base_sdk_install:
filters: *main_branches

View file

@ -0,0 +1,113 @@
from __future__ import annotations
import json
import os
import re
import subprocess
import sys
from pathlib import Path
from typing import Final
from xml.etree import ElementTree
SUITES: Final = {
"startup": (("test_startup.py",), 12),
"recovery": (("test_recovery.py",), 15),
"legacy": (("test_legacy.py", "test_pooling.py"), 11),
}
def successful_junit(path: Path, expected: int, exit_code: int) -> bool:
if exit_code != 0 or not path.is_file():
return False
try:
root: Final = ElementTree.parse(path).getroot()
except ElementTree.ParseError:
return False
cases: Final = tuple(root.iter("testcase"))
identities: Final = frozenset((case.get("classname"), case.get("name")) for case in cases)
return len(cases) == len(identities) == expected and all(
not any(case.find(tag) is not None for tag in ("failure", "error", "skipped")) for case in cases
)
def output(*command: str) -> str:
return subprocess.check_output(command, text=True, timeout=90).strip()
def record_image() -> None:
source: Final = output("git", "rev-parse", "HEAD")
image: Final = output("docker", "image", "inspect", "litellm-docker-database:ci", "--format", "{{.Id}}")
revision: Final = output(
"docker",
"image",
"inspect",
"litellm-docker-database:ci",
"--format",
'{{index .Config.Labels "org.opencontainers.image.revision"}}',
)
assert re.fullmatch(r"[0-9a-f]{40}", source), "Invalid source revision"
assert revision == source, "Candidate image revision differs from the tested source"
Path("migration-image.json").write_text(
json.dumps(
{
"source_sha": source,
"image_id": image,
"candidate_image": os.environ.get("MIGRATION_CANDIDATE_IMAGE", ""),
}
)
)
def main() -> int:
suite: Final = os.environ["MIGRATION_TEST_SUITE"]
files, expected = SUITES[suite]
metadata: Final = json.loads(Path("migration-image.json").read_text())
assert metadata["source_sha"] == output("git", "rev-parse", "HEAD"), "Image and test source revisions differ"
assert metadata["image_id"] == output(
"docker", "image", "inspect", os.environ["LITELLM_MIGRATION_TEST_IMAGE"], "--format", "{{.Id}}"
), "Loaded image differs from the build output"
assert metadata["candidate_image"] == os.environ.get("MIGRATION_CANDIDATE_IMAGE", ""), "Wrong release candidate"
destination: Final = Path(os.environ["MIGRATION_TEST_OUTPUT"])
junit: Final = destination / "junit" / "results.xml"
junit.parent.mkdir(parents=True, exist_ok=True)
result: Final = subprocess.run(
(
sys.executable,
"-m",
"pytest",
*(f"tests/e2e/migrations/{name}" for name in files),
"-vv",
"--tb=short",
"--durations=10",
f"--junitxml={junit}",
"-o",
"addopts=",
"--reruns=0",
),
check=False,
timeout=1200,
)
passed: Final = successful_junit(junit, expected, result.returncode)
(destination / "verdict.json").write_text(
json.dumps(
{
**metadata,
"suite": suite,
"expected_cases": expected,
"passed": passed,
"pytest_exit_code": result.returncode,
"test_revision": metadata["source_sha"],
"workflow_id": os.environ.get("CIRCLE_WORKFLOW_ID", ""),
"job_number": os.environ.get("CIRCLE_BUILD_NUM", ""),
},
indent=2,
)
)
return 0 if passed else 1
if __name__ == "__main__":
if sys.argv[1:] == ["record-image"]:
record_image()
else:
raise SystemExit(main())

View file

@ -4,7 +4,7 @@ from typing import Final
SELECTABLE: Final = re.compile(r"^tests/e2e/([A-Za-z0-9_.-]+/)*test_[A-Za-z0-9_.-]+\.py$")
UNSUPPORTED: Final = re.compile(
r"^tests/e2e/(ui|claude_code|load)/"
r"^tests/e2e/(ui|claude_code|load|migrations)/"
r"|^tests/e2e/mcp/test_mcp_oauth_happy_path_e2e\.py$"
r"|^tests/e2e/llm_translation/realtime/test_realtime_pipecat_audio_e2e\.py$"
r"|^tests/e2e/batches/test_managed_files_enforcement_e2e\.py$"

View file

@ -21,7 +21,7 @@ Parameters:
WorkerConfigParameter:
Type: String
Description: Sample environment variable
Default: '{"model": null, "alias": null, "api_base": null, "api_version": "2023-07-01-preview", "debug": false, "temperature": null, "max_tokens": null, "request_timeout": 600, "max_budget": null, "telemetry": true, "drop_params": false, "add_function_to_prompt": false, "headers": null, "save": false, "config": null, "use_queue": false}'
Default: '{"model": null, "alias": null, "api_base": null, "api_version": "2023-07-01-preview", "debug": false, "temperature": null, "max_tokens": null, "request_timeout": 600, "max_budget": null, "drop_params": false, "add_function_to_prompt": false, "headers": null, "save": false, "config": null, "use_queue": false}'
Resources:
MyUrlFunctionPermissions:

View file

@ -165,33 +165,41 @@ jobs:
DIST: ${{ inputs.dist }}
COVERAGE_CORE: sysmon
run: |
if [ "${WORKERS}" = "0" ]; then
uv run --no-sync pytest ${TEST_PATH:?} \
--tb=short -vv \
--maxfail="${MAX_FAILURES}" \
--reruns "${RERUNS}" \
--reruns-delay 1 \
--timeout="${TEST_TIMEOUT_SECONDS}" \
--rerun-except "from pytest-timeout" \
--durations=20 \
--cov=./litellm --cov=./enterprise/litellm_enterprise \
--cov-report=xml:coverage.xml \
--cov-config=pyproject.toml
else
uv run --no-sync pytest ${TEST_PATH:?} \
--tb=short -vv \
--maxfail="${MAX_FAILURES}" \
-n "${WORKERS}" \
--reruns "${RERUNS}" \
--reruns-delay 1 \
--timeout="${TEST_TIMEOUT_SECONDS}" \
--rerun-except "from pytest-timeout" \
--dist="${DIST}" \
--durations=20 \
--cov=./litellm --cov=./enterprise/litellm_enterprise \
--cov-report=xml:coverage.xml \
--cov-config=pyproject.toml
found_path=false
for path in ${TEST_PATH}; do
if [ -e "${path%%::*}" ]; then
found_path=true
break
fi
done
if [ "$found_path" = false ]; then
echo "No path in TEST_PATH exists (${TEST_PATH}); nothing to run"
exit 0
fi
xdist_args=()
if [ "${WORKERS}" != "0" ]; then
xdist_args=(-n "${WORKERS}" --dist="${DIST}")
fi
set +e
uv run --no-sync pytest ${TEST_PATH:?} \
--tb=short -vv \
--maxfail="${MAX_FAILURES}" \
"${xdist_args[@]}" \
--reruns "${RERUNS}" \
--reruns-delay 1 \
--timeout="${TEST_TIMEOUT_SECONDS}" \
--rerun-except "from pytest-timeout" \
--durations=20 \
--cov=./litellm --cov=./enterprise/litellm_enterprise \
--cov-report=xml:coverage.xml \
--cov-config=pyproject.toml
status=$?
set -e
if [ "$status" -eq 5 ]; then
echo "pytest collected no tests from ${TEST_PATH}; passing"
exit 0
fi
exit "$status"
- name: Save coverage report
if: always() && steps.changes.outputs.decision != 'skip'

View file

@ -15,7 +15,6 @@ model_list:
litellm_settings:
drop_params: True
telemetry: False
general_settings:
master_key: sk-1234 # Change this to a secure key

View file

@ -55,7 +55,6 @@ litellm_settings:
# budget_duration: 30d
num_retries: 5
request_timeout: 600
telemetry: False
context_window_fallbacks: [{"gpt-3.5-turbo": ["gpt-3.5-turbo-large"]}]
general_settings:

View file

@ -0,0 +1,89 @@
import random
import time
from collections.abc import Generator, Mapping
from contextlib import contextmanager
from dataclasses import dataclass
from typing import TYPE_CHECKING, Final
from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit
from litellm_proxy_extras._logging import logger
from litellm_proxy_extras.prisma_toolchain import MIGRATION_LOCK_TIMEOUT_ENV_VAR, migration_lock_timeout
MIGRATION_LOCK_KEY: Final = int.from_bytes(b"llm_mig2", "big")
if TYPE_CHECKING:
import psycopg
def migration_environment(environment: Mapping[str, str]) -> Mapping[str, str]:
database_url: Final = environment.get("DATABASE_URL")
direct_url: Final = environment.get("DIRECT_URL")
if not database_url or not direct_url:
return environment
schema: Final = next((value for key, value in parse_qsl(urlsplit(database_url).query) if key == "schema"), "public")
direct: Final = urlsplit(direct_url)
parameters: Final = tuple((key, value) for key, value in parse_qsl(direct.query) if key != "schema")
return {
**environment,
"DATABASE_URL": urlunsplit(direct._replace(query=urlencode((*parameters, ("schema", schema))))),
}
@dataclass(frozen=True, slots=True)
class _LockResult:
acquired: bool
def _try_lock(connection: "psycopg.Connection[tuple[object, ...]]", key: int = MIGRATION_LOCK_KEY) -> bool:
from psycopg.rows import class_row
with connection.cursor(row_factory=class_row(_LockResult)) as cursor:
row: Final = cursor.execute("SELECT pg_try_advisory_xact_lock(%s) AS acquired", (key,)).fetchone()
return row is not None and row.acquired
@dataclass(frozen=True, slots=True)
class MigrationCoordinator:
connection: "psycopg.Connection[tuple[object, ...]]"
def check_connection(self) -> None:
self.connection.execute("SELECT 1")
def acquire_prisma_lock(self) -> None:
deadline: Final = time.monotonic() + migration_lock_timeout()
while time.monotonic() < deadline:
if _try_lock(self.connection, 72707369):
return
time.sleep(min(random.uniform(0.5, 1.5), max(0.0, deadline - time.monotonic())))
raise RuntimeError(
"Timed out waiting for Prisma's lock to recover migration history. LiteLLM startup has stopped. "
"Another migration or a pooled database session may still hold the lock. Check the database lock holder. "
"When using a transaction pooler, configure DIRECT_URL to reach the same database without the pooler."
)
@contextmanager
def migration_lock(database_url: str) -> Generator[MigrationCoordinator, None, None]:
import psycopg
wait_seconds: Final = migration_lock_timeout()
deadline: Final = time.monotonic() + wait_seconds
try:
with psycopg.connect(database_url, connect_timeout=10, autocommit=True) as connection:
coordinator: Final = MigrationCoordinator(connection)
logger.info("Waiting for the v2 migration coordinator lock (up to %ss)", wait_seconds)
while time.monotonic() < deadline:
with connection.transaction():
if _try_lock(connection):
logger.info("Acquired the v2 migration coordinator lock")
yield coordinator
coordinator.check_connection()
return
time.sleep(min(random.uniform(0.5, 1.5), max(0.0, deadline - time.monotonic())))
except psycopg.Error as exc:
raise RuntimeError(f"Lost or could not establish v2 migration coordination with the database: {exc}") from exc
raise RuntimeError(
f"Timed out waiting for another v2 migration resolver after {wait_seconds}s. "
f"Check the running migration or increase {MIGRATION_LOCK_TIMEOUT_ENV_VAR}."
)

View file

@ -0,0 +1,158 @@
import hashlib
import subprocess
from collections.abc import Mapping
from dataclasses import dataclass
from pathlib import Path
from typing import TYPE_CHECKING, Final
from uuid import uuid4
from litellm_proxy_extras import prisma_toolchain
from litellm_proxy_extras._logging import logger
from litellm_proxy_extras.migration_lock import MigrationCoordinator
if TYPE_CHECKING:
import psycopg
@dataclass(frozen=True, slots=True)
class MigrationProgress:
checksum: str
applied_steps_count: int
logs: str
id: str = ""
finished: bool = False
def confirms_completion(self, script: bytes) -> bool:
return (
self.applied_steps_count == 1
and not self.logs.strip()
and self.checksum == hashlib.sha256(script).hexdigest()
)
def _migration_records(
connection: "psycopg.Connection[tuple[object, ...]]", schema: str, migration: Path
) -> tuple[MigrationProgress, ...]:
from psycopg import sql
from psycopg.rows import class_row
with connection.cursor(row_factory=class_row(MigrationProgress)) as cursor:
records: Final = cursor.execute(
sql.SQL(
"SELECT id, checksum, applied_steps_count, coalesce(logs, '') AS logs, "
"finished_at IS NOT NULL AS finished FROM {} "
"WHERE migration_name = %s AND rolled_back_at IS NULL"
).format(sql.Identifier(schema, "_prisma_migrations")),
(migration.parent.name,),
).fetchall()
return tuple(records)
def recover_completed_migration(coordinator: MigrationCoordinator, schema: str, migration: Path) -> bool:
"""Finish a proven successful row without erasing its durable completion evidence.
The caller commits this checkpoint before running another Prisma command.
"""
from psycopg import sql
coordinator.acquire_prisma_lock()
records: Final = _migration_records(coordinator.connection, schema, migration)
unfinished: Final = tuple(record for record in records if not record.finished)
script: Final = migration.read_bytes()
if not unfinished:
return any(record.checksum == hashlib.sha256(script).hexdigest() for record in records)
if len(unfinished) != 1 or not unfinished[0].confirms_completion(script):
return False
progress: Final = unfinished[0]
result: Final = coordinator.connection.execute(
sql.SQL(
"UPDATE {} SET finished_at = current_timestamp "
"WHERE id = %s AND checksum = %s AND applied_steps_count = 1 "
"AND finished_at IS NULL AND rolled_back_at IS NULL AND coalesce(logs, '') = %s"
).format(sql.Identifier(schema, "_prisma_migrations")),
(progress.id, progress.checksum, progress.logs),
)
if result.rowcount != 1:
raise RuntimeError("Could not complete the confirmed migration history row; retry startup.")
logger.info("Completed migration %s using its successful SQL step and matching checksum", migration.parent.name)
return True
def migration_files(directory: Path) -> tuple[tuple[str, str], ...]:
return tuple(
(path.parent.name, hashlib.sha256(path.read_bytes()).hexdigest())
for path in sorted((directory / "migrations").glob("*/migration.sql"))
)
def baseline_current_schema(
coordinator: MigrationCoordinator,
schema: str,
migrations_dir: Path,
prisma_command: str,
prisma_env: Mapping[str, str],
) -> None:
from psycopg import sql
packaged_dir: Final = Path(__file__).parent
migrations: Final = migration_files(migrations_dir)
if (
not migrations
or migrations != migration_files(packaged_dir)
or (migrations_dir / "schema.prisma").read_bytes() != (packaged_dir / "schema.prisma").read_bytes()
):
raise RuntimeError("Cannot automatically baseline an existing database with custom migration history.")
coordinator.acquire_prisma_lock()
existing: Final = coordinator.connection.execute(
"SELECT to_regclass(%s)", (sql.Identifier(schema, "_prisma_migrations").as_string(coordinator.connection),)
).fetchone()
if existing is not None and existing[0] is not None:
return
try:
prisma_toolchain.run_prisma(
(
prisma_command,
"migrate",
"diff",
"--from-schema-datasource",
str(migrations_dir / "schema.prisma"),
"--to-schema-datamodel",
str(migrations_dir / "schema.prisma"),
"--exit-code",
),
timeout=prisma_toolchain.prisma_command_timeout(),
env=prisma_env,
)
except (subprocess.CalledProcessError, subprocess.TimeoutExpired) as exc:
raise RuntimeError(
"Cannot automatically baseline this database: its schema has not been verified to match this build. "
"Establish the existing migration history before retrying. No schema reconciliation was performed. "
"If using a transaction pooler, configure DIRECT_URL to reach the same database without the pooler. "
f"Schema verification detail: {exc.stderr}"
) from exc
coordinator.check_connection()
ledger: Final = sql.Identifier(schema, "_prisma_migrations")
coordinator.connection.execute(
sql.SQL(
"CREATE TABLE {} (id varchar(36) PRIMARY KEY NOT NULL, checksum varchar(64) NOT NULL, "
"finished_at timestamptz, migration_name varchar(255) NOT NULL, logs text, rolled_back_at timestamptz, "
"started_at timestamptz NOT NULL DEFAULT now(), applied_steps_count integer NOT NULL DEFAULT 0)"
).format(ledger)
)
with coordinator.connection.cursor() as cursor:
cursor.executemany(
sql.SQL(
"INSERT INTO {} (id, checksum, migration_name, logs, started_at, finished_at) "
"VALUES (%s, %s, %s, '', current_timestamp, current_timestamp)"
).format(ledger),
tuple((str(uuid4()), checksum, name) for name, checksum in migrations),
)
logger.warning(
"Legacy migration history was missing. The existing Prisma schema matches this build; "
"adopted %s packaged migrations as a baseline. No schema changes were applied, and "
"historical data backfills were not replayed or verified. Continuing startup; "
"review any feature-specific backfill requirements.",
len(migrations),
)

View file

@ -59,6 +59,7 @@ except ImportError:
PRISMA_COMMAND_TIMEOUT_ENV_VAR = "LITELLM_PRISMA_COMMAND_TIMEOUT"
PRISMA_BOOTSTRAP_TIMEOUT_ENV_VAR = "LITELLM_PRISMA_BOOTSTRAP_TIMEOUT"
PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR = "LITELLM_PRISMA_MIGRATE_DEPLOY_TIMEOUT"
MIGRATION_LOCK_TIMEOUT_ENV_VAR = "LITELLM_MIGRATION_LOCK_TIMEOUT"
NODEENV_CACHE_DIR_ENV_VAR = "PRISMA_NODEENV_CACHE_DIR"
DEFAULT_PRISMA_COMMAND_TIMEOUT = 60.0
@ -106,6 +107,10 @@ def prisma_command_timeout() -> float:
)
def migration_lock_timeout() -> float:
return _timeout_from_env(MIGRATION_LOCK_TIMEOUT_ENV_VAR, 600.0)
def prisma_bootstrap_timeout() -> float:
"""Seconds the one-time Node toolchain install may run for."""
return _timeout_from_env(

View file

@ -6,6 +6,7 @@ import shutil
import subprocess
import tempfile
import time
from collections.abc import Callable
from dataclasses import dataclass, replace
from pathlib import Path
from typing import TYPE_CHECKING, Final, Optional
@ -78,15 +79,10 @@ MAX_MIGRATE_DEPLOY_ATTEMPTS = 4
@dataclass(frozen=True)
class _MigrateAttemptBudget:
"""Retries left, and the recoveries already run.
A recovery that lands something new costs nothing, so a database full of
objects `prisma db push` created works through them one per pass. Anything
that made no progress spends an attempt, so a stuck run still gives up.
"""
"""Independent bounds for failed attempts and Prisma lock contention."""
attempts_left: int
recoveries: frozenset[str] = frozenset()
contention_seconds_left: float = 600.0
@property
def exhausted(self) -> bool:
@ -99,10 +95,14 @@ class _MigrateAttemptBudget:
def spend(self) -> "_MigrateAttemptBudget":
return replace(self, attempts_left=self.attempts_left - 1)
def after_recovery(self, recovery: str) -> "_MigrateAttemptBudget":
if recovery in self.recoveries:
return self.spend()
return replace(self, recoveries=self.recoveries | {recovery})
def after_contention(self, elapsed: float) -> "_MigrateAttemptBudget":
remaining: Final = self.contention_seconds_left - elapsed
if remaining <= 0:
raise RuntimeError(
"Timed out waiting for Prisma's migration advisory lock. Check the running migration "
"or increase LITELLM_MIGRATION_LOCK_TIMEOUT."
)
return replace(self, contention_seconds_left=remaining)
_SPEND_LOGS_ALTER_RE = re.compile(r'^ALTER\s+TABLE\s+"LiteLLM_SpendLogs"\s', re.IGNORECASE)
@ -836,12 +836,51 @@ class ProxyExtrasDBManager:
@staticmethod
def _setup_database_v2(use_migrate: bool) -> bool:
if not use_migrate:
return ProxyExtrasDBManager._run_database_v2(False)
from litellm_proxy_extras.migration_lock import migration_environment, migration_lock
from litellm_proxy_extras.migration_recovery import baseline_current_schema, recover_completed_migration
database_url: Final = os.environ.get("DATABASE_URL")
if not database_url:
raise RuntimeError("DATABASE_URL is required for v2 migrations")
lock_url: Final = ProxyExtrasDBManager._strip_prisma_query_params(os.environ.get("DIRECT_URL") or database_url)
schema: Final = ProxyExtrasDBManager._prisma_schema_param(database_url) or "public"
def recover_completed(name: str) -> bool:
if Path(name).name != name or "\\" in name:
return False
migration: Final = Path(os.getcwd()) / "migrations" / name / "migration.sql"
if not migration.is_file():
return False
with migration_lock(lock_url) as coordinator:
return recover_completed_migration(coordinator, schema, migration)
def baseline_existing(migrations_dir: str) -> None:
with migration_lock(lock_url) as coordinator:
baseline_current_schema(
coordinator,
schema,
Path(migrations_dir),
_get_prisma_command(),
migration_environment(_get_prisma_env()),
)
while not ProxyExtrasDBManager._run_database_v2(True, recover_completed, baseline_existing):
continue
return True
@staticmethod
def _run_database_v2(
use_migrate: bool,
recover_completed: Callable[[str], bool] = lambda name: False,
baseline_existing: "Callable[[str], None] | None" = None,
) -> bool:
"""
v2 migration resolver (opt-in via --use_v2_migration_resolver).
Runs `prisma migrate deploy` and handles standard recovery paths
(P3005 baseline, P3009/P3018 idempotent errors, deadlocks against a
concurrent migrate deploy). Critically, it does
Runs `prisma migrate deploy`, baselines verified existing schemas,
and recovers confirmed SQL completion or reported deadlocks. It does
NOT call `_resolve_all_migrations` — the diff-and-force recovery that
caused schema thrashing when two LiteLLM versions contended for the
same DB during rolling deploys.
@ -850,10 +889,9 @@ class ProxyExtrasDBManager:
is logged as a warning, not a fatal error — users whose DBs got into
weird shapes from the old thrashing should still be able to start.
The retry budget only counts attempts that made no progress: see
_MigrateAttemptBudget.
False requests a committed recovery checkpoint and another deploy
pass. True means every pending migration is complete.
"""
schema_path = ProxyExtrasDBManager._get_prisma_dir() + "/schema.prisma"
migrations_dir = ProxyExtrasDBManager._get_prisma_dir()
if not use_migrate:
@ -886,14 +924,22 @@ class ProxyExtrasDBManager:
original_dir = os.getcwd()
os.chdir(migrations_dir)
deploy_timeout = prisma_migrate_deploy_timeout()
budget = _MigrateAttemptBudget(attempts_left=MAX_MIGRATE_DEPLOY_ATTEMPTS)
from litellm_proxy_extras.migration_lock import migration_environment, migration_lock_timeout
migration_env: Final = migration_environment(_get_prisma_env())
budget = _MigrateAttemptBudget(
attempts_left=MAX_MIGRATE_DEPLOY_ATTEMPTS,
contention_seconds_left=migration_lock_timeout(),
)
try:
while not budget.exhausted:
attempt_started = time.monotonic()
try:
result = prisma_toolchain.run_prisma(
[_get_prisma_command(), "migrate", "deploy"],
timeout=deploy_timeout,
env=_get_prisma_env(),
env=migration_env,
)
logger.info(f"prisma migrate deploy stdout: {result.stdout}")
return True
@ -909,8 +955,16 @@ class ProxyExtrasDBManager:
next_budget = budget.spend()
except subprocess.CalledProcessError as e:
if "P3005" in (e.stderr or "") and baseline_existing is not None:
baseline_existing(migrations_dir)
return False
failed_migration = ProxyExtrasDBManager._v2_failed_migration_name(e.stderr or "")
if failed_migration and recover_completed(failed_migration):
return False
next_budget = ProxyExtrasDBManager._budget_after_deploy_failure(
e, budget, schema_path
e,
budget,
time.monotonic() - attempt_started,
)
if next_budget.attempts_left < budget.attempts_left:
@ -919,19 +973,41 @@ class ProxyExtrasDBManager:
raise RuntimeError(
f"Database migration failed after {MAX_MIGRATE_DEPLOY_ATTEMPTS} "
"attempts that made no progress (timeouts, deadlock retries, or a "
"recovery that had already run once). Check database connectivity, "
"attempts that made no progress (timeouts or deadlock retries). Check database connectivity, "
"load, and _prisma_migrations ledger state, and raise "
f"{PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR} if the attempts timed out."
)
finally:
os.chdir(original_dir)
@staticmethod
def _v2_failed_migration_name(stderr: str) -> "str | None":
if "P3009" in stderr:
match = re.search(r"`(\d+_[^`\r\n]+)`", stderr)
return match.group(1) if match else None
if "P3018" in stderr:
match = re.search(r"Migration name: (\d+_[^\r\n]+)", stderr)
return match.group(1) if match else None
return None
@staticmethod
def _v2_roll_back_migration_best_effort(migration_name: str) -> None:
from litellm_proxy_extras.migration_lock import migration_environment
try:
prisma_toolchain.run_prisma(
[_get_prisma_command(), "migrate", "resolve", "--rolled-back", migration_name],
timeout=prisma_command_timeout(),
env=migration_environment(_get_prisma_env()),
)
except (subprocess.CalledProcessError, subprocess.TimeoutExpired):
pass
@staticmethod
def _budget_after_deploy_failure(
error: subprocess.CalledProcessError,
budget: "_MigrateAttemptBudget",
schema_path: str,
attempt_seconds: float = 0.0,
) -> "_MigrateAttemptBudget":
"""Recover from one failed `prisma migrate deploy`, and price the pass.
@ -940,37 +1016,35 @@ class ProxyExtrasDBManager:
"""
stderr = error.stderr or ""
if "P3005" in stderr and "database schema is not empty" in stderr:
logger.info("Schema exists but no migrations ledger — creating baseline")
if ProxyExtrasDBManager._create_baseline_migration(schema_path):
return budget.after_recovery("baseline")
return budget.spend()
if "P3009" in stderr:
migration_match = re.search(r"`(\d+_\S+?)`", stderr)
if migration_match and ProxyExtrasDBManager._is_idempotent_error(stderr):
name = migration_match.group(1)
logger.info(
f"Migration {name} failed idempotently — marking applied and retrying"
)
ProxyExtrasDBManager._mark_migration_applied(name)
return budget.after_recovery(f"resolved:{name}")
if migration_match:
migration_name = migration_match.group(1)
migration_name = ProxyExtrasDBManager._v2_failed_migration_name(stderr)
if migration_name:
ledger_logs = ProxyExtrasDBManager._failed_migration_logs(migration_name)
if ledger_logs is not None and (
ledger_logs == "" or _MIGRATION_DEADLOCK_MARKER in ledger_logs
):
if ledger_logs and _MIGRATION_DEADLOCK_MARKER in ledger_logs:
logger.info(
"Migration %s failed in a concurrent migrate deploy "
"deadlock race, rolling its ledger row back and retrying",
migration_name,
)
ProxyExtrasDBManager._roll_back_migration_best_effort(migration_name)
ProxyExtrasDBManager._v2_roll_back_migration_best_effort(migration_name)
return budget.spend()
raise RuntimeError(
"Database migration failed and cannot be auto-recovered. "
f"Manual intervention required.\n\nPrisma error:\n{stderr}"
"Migration completion could not be verified. LiteLLM startup has stopped.\n\n"
f"Prisma migration history (migration name and start time):\n{stderr}\n\n"
"A migration has a start record but no successful completion record. "
"LiteLLM cannot determine whether its SQL committed from this record alone. "
"Startup stopped to avoid repeating or skipping database changes.\n\n"
"Before resolving, stop other migration runners and inspect _prisma_migrations, "
"the named migration.sql from this build, database logs, and the actual database objects and data. "
"Use the same database and this build's schema and migration files for recovery:\n"
"- Only after verifying every migration change is present, run "
"prisma migrate resolve --applied <migration_name>, then retry startup.\n"
"- Only after verifying no migration changes remain (or fully undoing partial changes), run "
"prisma migrate resolve --rolled-back <migration_name>, then retry startup. "
"This command updates history; it does not undo SQL.\n"
"Replace <migration_name> with the reported name. If the outcome remains uncertain, "
"leave migration history unchanged and contact your database administrator. "
"Repeated restarts alone will not resolve this state."
) from error
if "P3018" in stderr:
@ -981,25 +1055,14 @@ class ProxyExtrasDBManager:
f"and retry.\n\nPrisma error:\n{stderr}"
) from error
migration_match = re.search(r"Migration name: (\d+_\S+)", stderr)
if migration_match and ProxyExtrasDBManager._is_idempotent_error(stderr):
name = migration_match.group(1)
migration_name = ProxyExtrasDBManager._v2_failed_migration_name(stderr)
if migration_name and _MIGRATION_DEADLOCK_MARKER in stderr:
logger.info(
f"Migration {name} SQL hit idempotent error — marking applied and retrying"
)
ProxyExtrasDBManager._mark_migration_applied(name)
return budget.after_recovery(f"resolved:{name}")
if migration_match and _MIGRATION_DEADLOCK_MARKER in stderr:
logger.info(
"Migration %s deadlocked against a concurrent "
"migrate deploy, rolling its ledger row back "
"and retrying",
migration_match.group(1),
)
ProxyExtrasDBManager._roll_back_migration_best_effort(
migration_match.group(1)
"Migration %s deadlocked against a concurrent migrate deploy, "
"rolling its ledger row back and retrying",
migration_name,
)
ProxyExtrasDBManager._v2_roll_back_migration_best_effort(migration_name)
return budget.spend()
raise RuntimeError(
@ -1009,19 +1072,17 @@ class ProxyExtrasDBManager:
if _MIGRATION_DEADLOCK_MARKER in stderr:
logger.info(
"prisma migrate deploy attempt %s deadlocked against "
"a concurrent migrate deploy, retrying",
"prisma migrate deploy attempt %s deadlocked against a concurrent migrate deploy, retrying",
budget.attempt_number,
)
return budget.spend()
if "P1002" in stderr and "advisory lock" in stderr:
logger.info(
"prisma migrate deploy attempt %s timed out waiting for "
"the advisory lock a concurrent migrate deploy holds, retrying",
budget.attempt_number,
"Waiting for the advisory lock held by another Prisma migration; "
"contention does not spend a migration failure attempt"
)
return budget.spend()
return budget.after_contention(attempt_seconds)
raise RuntimeError(
"Database migration failed and cannot be auto-recovered. "

View file

@ -6,7 +6,8 @@ The v2 resolver is opt-in via `--use_v2_migration_resolver` / the
"""
import subprocess
from unittest.mock import patch
from types import SimpleNamespace
from unittest.mock import MagicMock, Mock, patch
import pytest
@ -31,12 +32,7 @@ def _fake_migrate_deploy_failure(returncode: int, stderr: str):
def test_v2_p3018_permission_error_raises_runtime_error(monkeypatch, tmp_path):
"""v2: a permission failure during migrate deploy raises RuntimeError."""
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@localhost:9/x")
monkeypatch.setattr(
ProxyExtrasDBManager, "_warn_if_db_ahead_of_head", lambda _: None
)
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
(tmp_path / "schema.prisma").write_text("// stub")
_stub_v2_env(monkeypatch, tmp_path)
stderr = (
"Error: P3018\nMigration name: 20250326162113_baseline\n"
@ -49,19 +45,14 @@ def test_v2_p3018_permission_error_raises_runtime_error(monkeypatch, tmp_path):
def test_v2_non_idempotent_p3009_raises_runtime_error(monkeypatch, tmp_path):
"""v2: a non-idempotent migration failure raises (no silent recovery)."""
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@localhost:9/x")
monkeypatch.setattr(
ProxyExtrasDBManager, "_warn_if_db_ahead_of_head", lambda _: None
)
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
(tmp_path / "schema.prisma").write_text("// stub")
_stub_v2_env(monkeypatch, tmp_path)
stderr = (
"Error: P3009\nMigration `20260101000000_genuinely_broken` failed\n"
'Reason: syntax error at or near "BRKN" LINE 42'
)
with patch("litellm_proxy_extras.prisma_toolchain.run_prisma", side_effect=_fake_migrate_deploy_failure(1, stderr)):
with pytest.raises(RuntimeError, match="cannot be auto-recovered"):
with pytest.raises(RuntimeError, match="Migration completion could not be verified"):
ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
@ -135,8 +126,7 @@ def test_v1_default_still_calls_resolve_all_migrations(monkeypatch, tmp_path):
def test_v2_db_push_wraps_subprocess_error_as_runtime_error(monkeypatch, tmp_path):
"""v2: a failing `prisma db push` must raise RuntimeError, not leak
CalledProcessError past proxy_cli.py's `except RuntimeError`."""
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
(tmp_path / "schema.prisma").write_text("// stub")
monkeypatch.setenv("LITELLM_MIGRATION_DIR", str(tmp_path))
stderr = "db push error"
with patch("litellm_proxy_extras.prisma_toolchain.run_prisma", side_effect=_fake_migrate_deploy_failure(1, stderr)):
@ -153,8 +143,7 @@ def test_v2_warn_ahead_of_head_swallows_db_errors(monkeypatch, tmp_path):
import psycopg
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@localhost:9/x")
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
(tmp_path / "schema.prisma").write_text("// stub")
monkeypatch.setenv("LITELLM_MIGRATION_DIR", str(tmp_path))
class _FakeConn:
def __enter__(self):
@ -176,70 +165,28 @@ def test_v2_warn_ahead_of_head_swallows_db_errors(monkeypatch, tmp_path):
ProxyExtrasDBManager._warn_if_db_ahead_of_head(str(tmp_path))
def test_v2_resolve_specific_migration_failure_raises_runtime_error(
monkeypatch, tmp_path
):
"""If marking a migration as applied fails inside P3009 idempotent
recovery, the subprocess error must be re-raised as RuntimeError so
proxy_cli.py catches it cleanly (instead of leaking CalledProcessError)."""
monkeypatch.setattr(
ProxyExtrasDBManager, "_warn_if_db_ahead_of_head", lambda _: None
)
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
(tmp_path / "schema.prisma").write_text("// stub")
monkeypatch.setattr(
ProxyExtrasDBManager, "_roll_back_migration", lambda *a, **kw: None
)
# First call: migrate deploy -> P3009 idempotent error.
# Recovery path tries _resolve_specific_migration; that also raises.
def _failing_resolve(*a, **kw):
raise subprocess.CalledProcessError(
returncode=1,
cmd="prisma migrate resolve --applied",
stderr="resolve failed",
output="",
)
monkeypatch.setattr(
ProxyExtrasDBManager, "_resolve_specific_migration", _failing_resolve
)
stderr = (
"Error: P3009\nMigration `20260101000000_some_migration` failed\n"
"relation already exists"
)
with patch("litellm_proxy_extras.prisma_toolchain.run_prisma", side_effect=_fake_migrate_deploy_failure(1, stderr)):
with pytest.raises(
RuntimeError, match="Failed to mark migration .* as applied"
):
def test_v2_duplicate_object_p3009_is_not_marked_applied(monkeypatch, tmp_path):
_stub_v2_env(monkeypatch, tmp_path, ledger_logs="relation already exists")
stderr = "Error: P3009\nMigration `20260101000000_some_migration` failed\nrelation already exists"
with patch(
"litellm_proxy_extras.prisma_toolchain.run_prisma", side_effect=_fake_migrate_deploy_failure(1, stderr)
) as run:
with pytest.raises(RuntimeError, match="Migration completion could not be verified"):
ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
assert tuple(call.args[0][1:] for call in run.call_args_list if "migrate" in call.args[0]) == (
["migrate", "deploy"],
)
def test_v2_does_not_call_resolve_all_migrations(monkeypatch, tmp_path):
"""v2 must never call _resolve_all_migrations — that's the bug it fixes."""
monkeypatch.setattr(
ProxyExtrasDBManager, "_warn_if_db_ahead_of_head", lambda _: None
_stub_v2_env(monkeypatch, tmp_path)
run = Mock(side_effect=_succeed_after(0, ""))
monkeypatch.setattr("litellm_proxy_extras.prisma_toolchain.run_prisma", run)
assert ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True) is True
assert tuple(call.args[0][1:] for call in run.call_args_list if "migrate" in call.args[0]) == (
["migrate", "deploy"],
)
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
(tmp_path / "schema.prisma").write_text("// stub")
class FakeResult:
stdout = "Applied migration.\n"
stderr = ""
monkeypatch.setattr("litellm_proxy_extras.prisma_toolchain.run_prisma", lambda *a, **kw: FakeResult())
resolve_called = {"n": 0}
monkeypatch.setattr(
ProxyExtrasDBManager,
"_resolve_all_migrations",
lambda *a, **kw: resolve_called.__setitem__("n", resolve_called["n"] + 1),
)
ok = ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
assert ok is True
assert resolve_called["n"] == 0, "v2 must not invoke the diff-and-force recovery"
_DEADLOCK_P3018_STDERR = (
@ -250,14 +197,34 @@ _DEADLOCK_P3018_STDERR = (
)
def _stub_v2_env(monkeypatch, tmp_path):
def _stub_v2_env(monkeypatch, tmp_path, ledger_logs=""):
import psycopg
monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@localhost:9/x")
monkeypatch.setattr(
ProxyExtrasDBManager, "_warn_if_db_ahead_of_head", lambda _: None
)
monkeypatch.setattr(ProxyExtrasDBManager, "_get_prisma_dir", lambda: str(tmp_path))
(tmp_path / "schema.prisma").write_text("// stub")
monkeypatch.delenv("DIRECT_URL", raising=False)
monkeypatch.setenv("LITELLM_MIGRATION_DIR", str(tmp_path))
monkeypatch.setattr("time.sleep", lambda _: None)
connection = MagicMock()
connection.__enter__.return_value = connection
cursor = connection.cursor.return_value.__enter__.return_value
cursor.execute.return_value = cursor
cursor.fetchone.return_value = SimpleNamespace(acquired=True)
cursor.fetchall.return_value = []
empty = MagicMock()
empty.fetchall.return_value = []
empty.fetchone.return_value = None
ledger = MagicMock()
ledger.fetchone.return_value = (ledger_logs,)
def execute(query, *args, **kwargs):
if "SELECT logs FROM" in str(query):
if ledger_logs is None:
raise psycopg.OperationalError("ledger is unavailable")
return ledger
return empty
connection.execute.side_effect = execute
monkeypatch.setattr("psycopg.connect", lambda *args, **kwargs: connection)
def _succeed_after(failures: int, stderr: str):
@ -272,9 +239,7 @@ def _succeed_after(failures: int, stderr: str):
return _OkResult()
calls["n"] += 1
if calls["n"] <= failures:
raise subprocess.CalledProcessError(
returncode=1, cmd=args[0], stderr=stderr, output=""
)
raise subprocess.CalledProcessError(returncode=1, cmd=args[0], stderr=stderr, output="")
return _OkResult()
return _run
@ -285,28 +250,21 @@ def test_v2_p3018_deadlock_rolls_back_and_retries(monkeypatch, tmp_path):
instance rolls the ledger row back and retries instead of dying."""
_stub_v2_env(monkeypatch, tmp_path)
rolled_back = []
monkeypatch.setattr(
ProxyExtrasDBManager,
"_roll_back_migration",
lambda name: rolled_back.append(name),
)
monkeypatch.setattr(
ProxyExtrasDBManager,
"_resolve_specific_migration",
lambda name: pytest.fail("a deadlocked migration must never be marked applied"),
)
monkeypatch.setattr("litellm_proxy_extras.prisma_toolchain.run_prisma", _succeed_after(1, _DEADLOCK_P3018_STDERR))
run = Mock(side_effect=_succeed_after(1, _DEADLOCK_P3018_STDERR))
monkeypatch.setattr("litellm_proxy_extras.prisma_toolchain.run_prisma", run)
ok = ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
assert ok is True
assert rolled_back == ["20260415120000_health_check_latest_per_model_index"]
assert tuple(call.args[0][1:] for call in run.call_args_list if "migrate" in call.args[0]) == (
["migrate", "deploy"],
["migrate", "resolve", "--rolled-back", "20260415120000_health_check_latest_per_model_index"],
["migrate", "deploy"],
)
def test_v2_p3018_persistent_deadlock_exhausts_attempts(monkeypatch, tmp_path):
"""v2: a deadlock on every attempt still fails after the retry budget."""
_stub_v2_env(monkeypatch, tmp_path)
monkeypatch.setattr(ProxyExtrasDBManager, "_roll_back_migration", lambda name: None)
with patch(
"litellm_proxy_extras.prisma_toolchain.run_prisma",
@ -319,7 +277,7 @@ def test_v2_p3018_persistent_deadlock_exhausts_attempts(monkeypatch, tmp_path):
def test_v2_p3009_deadlocked_ledger_row_rolls_back_and_retries(monkeypatch, tmp_path):
"""v2: the surviving instance sees the victim's failed ledger row as P3009.
When that row's logs show a deadlock, roll it back and retry."""
_stub_v2_env(monkeypatch, tmp_path)
_stub_v2_env(monkeypatch, tmp_path, ledger_logs="ERROR: deadlock detected\nDETAIL: Process 72 waits for ShareLock")
stderr = (
"Error: P3009\n"
@ -327,61 +285,39 @@ def test_v2_p3009_deadlocked_ledger_row_rolls_back_and_retries(monkeypatch, tmp_
"The `20260415120000_health_check_latest_per_model_index` migration "
"started at 2026-09-01 18:46:13 UTC failed"
)
monkeypatch.setattr(
ProxyExtrasDBManager,
"_failed_migration_logs",
lambda name: "ERROR: deadlock detected\nDETAIL: Process 72 waits for ShareLock",
)
rolled_back = []
monkeypatch.setattr(
ProxyExtrasDBManager,
"_roll_back_migration",
lambda name: rolled_back.append(name),
)
monkeypatch.setattr(
ProxyExtrasDBManager,
"_resolve_specific_migration",
lambda name: pytest.fail("a deadlocked migration must never be marked applied"),
)
monkeypatch.setattr("litellm_proxy_extras.prisma_toolchain.run_prisma", _succeed_after(1, stderr))
run = Mock(side_effect=_succeed_after(1, stderr))
monkeypatch.setattr("litellm_proxy_extras.prisma_toolchain.run_prisma", run)
ok = ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
assert ok is True
assert rolled_back == ["20260415120000_health_check_latest_per_model_index"]
assert tuple(call.args[0][1:] for call in run.call_args_list if "migrate" in call.args[0]) == (
["migrate", "deploy"],
["migrate", "resolve", "--rolled-back", "20260415120000_health_check_latest_per_model_index"],
["migrate", "deploy"],
)
def test_v2_p3009_empty_ledger_logs_rolls_back_and_retries(monkeypatch, tmp_path):
"""v2: empty failed ledger logs mean a concurrent deploy moved it on."""
def test_v2_p3009_empty_ledger_logs_do_not_prove_completion(monkeypatch, tmp_path):
_stub_v2_env(monkeypatch, tmp_path)
stderr = (
"Error: P3009\n"
"migrate found failed migrations in the target database\n"
"The `20260415120000_health_check_latest_per_model_index` migration "
"started at 2026-09-01 18:46:13 UTC failed"
)
monkeypatch.setattr(ProxyExtrasDBManager, "_failed_migration_logs", lambda name: "")
rolled_back = []
monkeypatch.setattr(
ProxyExtrasDBManager,
"_roll_back_migration",
lambda name: rolled_back.append(name),
with patch(
"litellm_proxy_extras.prisma_toolchain.run_prisma", side_effect=_fake_migrate_deploy_failure(1, stderr)
) as run:
with pytest.raises(RuntimeError, match="Migration completion could not be verified"):
ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
assert tuple(call.args[0][1:] for call in run.call_args_list if "migrate" in call.args[0]) == (
["migrate", "deploy"],
)
monkeypatch.setattr(
ProxyExtrasDBManager,
"_resolve_specific_migration",
lambda name: pytest.fail("a deadlocked migration must never be marked applied"),
)
monkeypatch.setattr("litellm_proxy_extras.prisma_toolchain.run_prisma", _succeed_after(1, stderr))
ok = ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
assert ok is True
assert rolled_back == ["20260415120000_health_check_latest_per_model_index"]
def test_v2_p3009_unreadable_ledger_still_raises(monkeypatch, tmp_path):
"""v2: an unreadable ledger cannot establish that P3009 was a deadlock."""
_stub_v2_env(monkeypatch, tmp_path)
_stub_v2_env(monkeypatch, tmp_path, ledger_logs=None)
stderr = (
"Error: P3009\n"
@ -389,21 +325,15 @@ def test_v2_p3009_unreadable_ledger_still_raises(monkeypatch, tmp_path):
"The `20260415120000_health_check_latest_per_model_index` migration "
"started at 2026-09-01 18:46:13 UTC failed"
)
monkeypatch.setattr(ProxyExtrasDBManager, "_failed_migration_logs", lambda name: None)
monkeypatch.setattr(
ProxyExtrasDBManager,
"_roll_back_migration",
lambda name: pytest.fail("an unreadable ledger must not trigger a retry"),
)
monkeypatch.setattr("litellm_proxy_extras.prisma_toolchain.run_prisma", _succeed_after(1, stderr))
with pytest.raises(RuntimeError, match="cannot be auto-recovered"):
with pytest.raises(RuntimeError, match="Migration completion could not be verified"):
ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)
def test_v2_p3009_non_deadlock_ledger_row_still_raises(monkeypatch, tmp_path):
"""v2: a failed ledger row whose logs show a real SQL error stays fatal."""
_stub_v2_env(monkeypatch, tmp_path)
_stub_v2_env(monkeypatch, tmp_path, ledger_logs='ERROR: syntax error at or near "BRKN"')
stderr = (
"Error: P3009\n"
@ -411,14 +341,9 @@ def test_v2_p3009_non_deadlock_ledger_row_still_raises(monkeypatch, tmp_path):
"The `20260101000000_genuinely_broken` migration started at "
"2026-09-01 18:46:13 UTC failed"
)
monkeypatch.setattr(
ProxyExtrasDBManager,
"_failed_migration_logs",
lambda name: 'ERROR: syntax error at or near "BRKN"',
)
with patch("litellm_proxy_extras.prisma_toolchain.run_prisma", side_effect=_fake_migrate_deploy_failure(1, stderr)):
with pytest.raises(RuntimeError, match="cannot be auto-recovered"):
with pytest.raises(RuntimeError, match="Migration completion could not be verified"):
ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=True)

View file

@ -240,7 +240,6 @@ email: Optional[str] = (
token: Optional[str] = (
None # Not used anymore, will be removed in next MAJOR release - https://github.com/BerriAI/litellm/discussions/648
)
telemetry = True
max_tokens: int = DEFAULT_MAX_TOKENS # OpenAI Defaults
drop_params = drop_params_env_flag(os.environ, verbose_logger)
modify_params = bool(os.getenv("LITELLM_MODIFY_PARAMS", False))

View file

@ -68102,9 +68102,9 @@
"supports_web_search": true
},
"openrouter/z-ai/glm-5.2": {
"input_cost_per_token": 5.544e-07,
"output_cost_per_token": 1.7424e-06,
"cache_read_input_token_cost": 1.0296e-07,
"input_cost_per_token": 6.496e-07,
"output_cost_per_token": 2.0416e-06,
"cache_read_input_token_cost": 1.2064e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 131072,

View file

@ -212,7 +212,6 @@ sandbox_tools:
litellm_settings:
drop_params: True
telemetry: False
code_interpreter_interception_params:
enabled: true
sandbox_tool_name: e2b_sandbox

View file

@ -38,7 +38,6 @@ litellm_settings:
# budget_duration: 30d
num_retries: 5
request_timeout: 600
telemetry: False
context_window_fallbacks: [{"gpt-5-mini": ["gpt-5.5"]}]
default_team_settings:
- team_id: team-1

View file

@ -56,8 +56,6 @@ if litellm_mode == "DEV":
load_dotenv()
from enum import Enum
telemetry: Final = None
class LiteLLMDatabaseConnectionPool(Enum):
database_connection_pool_limit = 10
@ -758,9 +756,11 @@ class ProxyInitializationHelpers:
)
@click.option(
"--telemetry",
default=True,
default=None,
type=bool,
help="Helps us know if people are using this feature. Turn this off by doing `--telemetry False`",
hidden=True,
expose_value=False,
help="Deprecated no-op kept so existing start commands still parse",
)
@click.option(
"--log_config",
@ -977,7 +977,6 @@ def run_server(
add_function_to_prompt,
config,
max_budget,
telemetry,
test,
local,
num_workers,
@ -1082,7 +1081,6 @@ def run_server(
max_tokens=max_tokens,
request_timeout=request_timeout,
max_budget=max_budget,
telemetry=telemetry,
drop_params=drop_params,
add_function_to_prompt=add_function_to_prompt,
headers=headers,

View file

@ -1230,12 +1230,12 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
general_settings,
) = await proxy_config.load_config(router=llm_router, config_file_path=worker_config)
elif isinstance(worker_config, dict):
await initialize(**worker_config)
await initialize_from_worker_config(worker_config)
else:
# if not, assume it's a json string
worker_config = json.loads(worker_config)
if isinstance(worker_config, dict):
await initialize(**worker_config)
await initialize_from_worker_config(worker_config)
enforce_master_key_boot_verdict(
await with_stored_secrets_counted(
@ -2426,7 +2426,6 @@ user_debug = False
user_max_tokens = None
user_request_timeout = None
user_temperature = None
user_telemetry = True
user_config: Final = None
user_headers = None
user_config_file_path: str | None = None
@ -8621,6 +8620,14 @@ def save_worker_config(**data):
os.environ["WORKER_CONFIG"] = json.dumps(data)
LEGACY_WORKER_CONFIG_KEYS: Final = frozenset({"telemetry"})
async def initialize_from_worker_config(worker_config: Mapping[str, object]) -> None:
supported: Final = MappingProxyType({k: v for k, v in worker_config.items() if k not in LEGACY_WORKER_CONFIG_KEYS})
await initialize(**supported)
async def initialize(
model=None,
alias=None,
@ -8632,7 +8639,6 @@ async def initialize(
max_tokens=None,
request_timeout=600,
max_budget=None,
telemetry=False,
drop_params=True,
add_function_to_prompt=True,
headers=None,
@ -8648,7 +8654,6 @@ async def initialize(
user_user_max_tokens, \
user_request_timeout, \
user_temperature, \
user_telemetry, \
user_headers, \
experimental, \
llm_model_list, \
@ -8755,7 +8760,6 @@ async def initialize(
dynamic_config["general"]["max_budget"] = litellm.max_budget
if experimental:
pass
user_telemetry = telemetry
# for streaming

View file

@ -49,4 +49,3 @@ general_settings:
litellm_settings:
drop_params: True
telemetry: False

View file

@ -68102,9 +68102,9 @@
"supports_web_search": true
},
"openrouter/z-ai/glm-5.2": {
"input_cost_per_token": 5.544e-07,
"output_cost_per_token": 1.7424e-06,
"cache_read_input_token_cost": 1.0296e-07,
"input_cost_per_token": 6.496e-07,
"output_cost_per_token": 2.0416e-06,
"cache_read_input_token_cost": 1.2064e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 131072,

View file

@ -173,7 +173,6 @@ litellm_settings:
# budget_duration: 30d
num_retries: 5
request_timeout: 600
telemetry: False
context_window_fallbacks: [{"gpt-3.5-turbo": ["gpt-3.5-turbo-large"]}]
default_team_settings:
- team_id: team-1

View file

@ -196,6 +196,7 @@ dev = [
"mypy==1.20.1",
"keyring==25.7.0",
"pytest==9.0.3",
"pytest-socket==0.8.1",
"tomli==2.4.1; python_version < '3.11'",
"pytest-mock==3.15.1",
"pytest-asyncio==1.3.0",

View file

@ -256,7 +256,6 @@ general_settings:
master_key: {api_key}
litellm_settings:
telemetry: false
""",
encoding="utf-8",
)

View file

@ -228,7 +228,6 @@ general_settings:
litellm_settings:
drop_params: true
telemetry: false
""",
encoding="utf-8",
)

View file

@ -7,13 +7,17 @@ driven DOWN over time. This check compares every budget file against its own
content at the merge-base with the target branch and fails (exits 1, red) if:
* a rule's `limit` went up,
* a rule was dropped from a budget (its ceiling effectively became infinite), or
* a rule was dropped from a budget (its ceiling effectively became infinite) while
its checker still emits it, or
* an entire budget file was deleted.
New rules and lowered/equal limits are fine. So is a rule that graduated: once a
paired config (ruff.toml for the ruff-strict budget) selects the rule outright it
hard-fails at the first violation, which is stricter than any ceiling the budget
could hold, so dropping its entry tightens the guard rather than removing it.
Likewise a retired rule: once the paired checker (check_test_quality.py for the
test-quality budget) no longer emits a code, its entry has no ceiling left to
loosen.
This is deliberately NOT a gating check. It should turn the run red so that a
loosening is impossible to miss in review, but it must stay OUT of the
@ -29,11 +33,12 @@ Usage:
from __future__ import annotations
import argparse
import importlib.util
import json
import subprocess
import sys
from pathlib import Path
from types import MappingProxyType
from types import MappingProxyType, ModuleType
from typing import Final, NamedTuple
if sys.version_info >= (3, 11):
@ -49,6 +54,7 @@ DEFAULT_BUDGETS: tuple[str, ...] = (
"test-quality-budget.json",
)
GRADUATION_CONFIGS = MappingProxyType({"ruff-strict-budget.json": "ruff.toml"})
RETIREMENT_SOURCES = MappingProxyType({"test-quality-budget.json": "check_test_quality"})
class Regression(NamedTuple):
@ -139,20 +145,40 @@ def graduated_selectors(rel: str) -> tuple[str, ...]:
)
def _load_script(name: str) -> ModuleType:
if name in sys.modules:
return sys.modules[name]
spec: Final = importlib.util.spec_from_file_location(name, REPO_ROOT / "scripts" / f"{name}.py")
assert spec is not None and spec.loader is not None
module: Final = importlib.util.module_from_spec(spec)
sys.modules[name] = module
spec.loader.exec_module(module)
return module
def retired_rules(rel: str, base: dict[str, object]) -> frozenset[str]:
"""Rules in the base budget that the paired checker can no longer emit, so there is no ceiling to loosen."""
source: Final = RETIREMENT_SOURCES.get(rel)
if source is None:
return frozenset()
return frozenset(_limits(base)) - _load_script(source).RULE_CODES
def _regression_detail(
rule: str,
base_limits: dict[str, int],
head_limits: dict[str, int],
graduated: tuple[str, ...],
retired: frozenset[str] = frozenset(),
) -> str | None:
"""Why `rule` regressed vs base, or None when it held flat, fell, or graduated.
"""Why `rule` regressed vs base, or None when it held flat, fell, or left the budget legitimately.
A dropped rule is terminal unless it graduated; otherwise the only loosening
left is a raised limit.
A dropped rule is terminal unless it graduated or retired; otherwise the only
loosening left is a raised limit.
"""
base_limit = base_limits[rule]
if rule not in head_limits:
if graduated and rule.startswith(graduated):
if rule in retired or (graduated and rule.startswith(graduated)):
return None
return f"rule dropped (limit {base_limit} -> removed)"
if head_limits[rule] > base_limit:
@ -165,6 +191,7 @@ def regressions_for(
base: dict | None,
head: dict | None,
graduated: tuple[str, ...] = (),
retired: frozenset[str] = frozenset(),
) -> list[Regression]:
if base is None:
return [] # new budget file: nothing to ratchet against yet
@ -175,7 +202,7 @@ def regressions_for(
return [
Regression(rel, rule, detail)
for rule in sorted(base_limits)
if (detail := _regression_detail(rule, base_limits, head_limits, graduated)) is not None
if (detail := _regression_detail(rule, base_limits, head_limits, graduated, retired)) is not None
]
@ -209,7 +236,7 @@ def main() -> int:
print(f"skip {rel}: new file (no base at {base_ref} to ratchet against)")
continue
checked.append(rel)
regressions.extend(regressions_for(rel, base, head, graduated_selectors(rel)))
regressions.extend(regressions_for(rel, base, head, graduated_selectors(rel), retired_rules(rel, base)))
if regressions:
print(

View file

@ -48,10 +48,6 @@ TQ006 A `pytest.skip` reached only when a credential-shaped environment variab
deliberate branch. The gate follows one local or module-level binding, which is
the `key = os.getenv(...)` then `if not key: pytest.skip(...)` shape most of
these use.
TQ008 A `patch(...)` whose target is a `litellm.` internal. Patching the SDK's own
functions pins the test to the current wiring instead of the behaviour, and it
is the idiom the suite reaches for instead of faking the HTTP boundary. Mocking
a third-party client, a transport, or anything outside `litellm.` is untouched.
TQ007 A module global that a conftest saves before every test and restores after it.
The save/restore list is a hand-maintained inventory of the leaks the suite
already knows about, so it is allowed to shrink and never to grow: a new entry
@ -150,6 +146,10 @@ SDK_MODULE: Final = "litellm"
SUBPROCESS_SPAWNS: Final = frozenset(("run", "Popen", "check_output", "check_call", "call"))
INTERPRETER_ISOLATION_FLAGS: Final = frozenset(("-I", "-P"))
RULE_CODES: Final = frozenset((
"TQ000", "TQ001", "TQ002", "TQ003", "TQ004", "TQ005", "TQ006", "TQ007", "TQ009",
))
CREDENTIAL_NAME_RE: Final = re.compile(
r"(?:API_KEY|_KEY|TOKEN|SECRET|PASSWORD|CREDENTIAL|DATABASE_URL|ACCESS_KEY_ID)$"
)
@ -483,67 +483,6 @@ def iter_global_mutation_violations(path: Path, tree: ast.Module) -> Iterator[Vi
)
def _is_sdk_internal(dotted: str) -> bool:
return dotted == SDK_MODULE or dotted.startswith(f"{SDK_MODULE}.")
def _sdk_import_bindings(tree: ast.Module) -> Iterator[tuple[str, str]]:
"""(local name, dotted path) for every import that binds something under `litellm`."""
for node in ast.walk(tree):
if isinstance(node, ast.Import):
yield from (
(alias.asname, alias.name) if alias.asname else (root, root)
for alias in node.names
if _is_sdk_internal(alias.name)
for root in (alias.name.partition(".")[0],)
)
elif isinstance(node, ast.ImportFrom) and node.module and _is_sdk_internal(node.module):
yield from ((alias.asname or alias.name, f"{node.module}.{alias.name}") for alias in node.names)
def _sdk_aliases(tree: ast.Module) -> Mapping[str, str]:
"""Local names bound to something under `litellm`, mapped to the path they stand for.
`from litellm.llms.openai.chat import handler` then `patch.object(handler.X, ...)`
reaches the same internal as the dotted string form and has to read the same way.
"""
return MappingProxyType({name: dotted for name, dotted in _sdk_import_bindings(tree)})
def _resolved(dotted: str, aliases: Mapping[str, str]) -> str:
root, _, rest = dotted.partition(".")
base: Final = aliases.get(root, root)
return f"{base}.{rest}" if rest else base
def _patch_targets(call: ast.Call, aliases: Mapping[str, str]) -> Iterator[str]:
"""What a patch installer is replacing: the dotted string it names, or the
attribute chain handed to `patch.object` / `patch.dict`, resolved through the
module's imports so a locally bound SDK object reads as its full path."""
for first in call.args[:1]:
if isinstance(first, ast.Constant) and isinstance(first.value, str):
yield first.value
elif dotted := _dotted_name(first):
yield _resolved(dotted, aliases)
def iter_internal_patch_violations(path: Path, tree: ast.Module) -> Iterator[Violation]:
aliases: Final = _sdk_aliases(tree)
for node in ast.walk(tree):
if not (isinstance(node, ast.Call) and _is_patch_installer(_dotted_name(node.func))):
continue
for target in _patch_targets(node, aliases):
if _is_sdk_internal(target):
yield Violation(
path,
node.lineno,
"TQ008",
f"patches `{target}`, an SDK internal, so the test is pinned to how the code is "
"wired rather than what it does; fake the HTTP boundary (respx / MockTransport) "
f"or inject the collaborator (suppress: `# {SUPPRESSION_TOKEN}: <reason>`)",
)
def _environ_keys(node: ast.AST) -> Iterator[str]:
for inner in ast.walk(node):
if isinstance(inner, ast.Call) and _dotted_name(inner.func) in ENVIRON_READERS:
@ -784,7 +723,6 @@ def check_file(path: Path) -> tuple[Violation, ...]:
*iter_global_mutation_violations(path, tree),
*iter_credential_skip_violations(path, tree),
*iter_conftest_inventory_violations(path, tree),
*iter_internal_patch_violations(path, tree),
*iter_child_interpreter_violations(path, tree),
)
if violation.line not in skip

View file

@ -20,9 +20,6 @@
"TQ007": {
"limit": 117
},
"TQ008": {
"limit": 10993
},
"TQ009": {
"limit": 59
}

38
tests/AGENTS.md Normal file
View file

@ -0,0 +1,38 @@
# Tests
Nothing on the other side of the call: `tests/unit`. A proxy we start with an upstream we script:
`tests/integration`. Someone else's service with real credentials: `tests/e2e`. Two fit, split it
## What good looks like
Red when the claim in the name is broken. Prove it: mutate the behaviour, red; restore, green. Put the
mutation in the PR body
```python
def test_custom_price_is_reported_and_charged(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002)
response = gateway.request("POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": "price control"}]})
assert response.status_code == 200, response.text
assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(20 * 0.001 + 20 * 0.002)
```
Rates in the test, expected computed by hand, one call, `response.text` in the assert
Assert the whole value. Iterating `expected_body.items()` (`test_responses_api_request_body.py`) cannot
see an extra key; that is the shape of `stream_options.include_usage` (#19777, #28553)
The linter catches no-assert, mock-echo and credential skips. It cannot see an assert
behind an `if` (a poll that ends in `pytest.fail` is fine), `except Exception` around the call
(`test_router.py`: `except Exception as e: print(f"FAILED TEST")`), or blanket `--reruns`
## Where it goes
What the assertion depends on goes in the test; everything else in conftest. A rate in a fixture three
directories up makes a failed assertion unreadable. Extend the file that already covers the behaviour
## Writing it so a human can read it
Name says what broke: `test_send_batched_with_valid_data` says nothing. Build, one call, assert, on one
screen. Helpers named for what they return, `_pii_prompt(marker, email)`, not `_setup()`. Context in the
assert message, not a comment

View file

@ -132,6 +132,7 @@ def select_tests(changed: tuple[str, ...]) -> tuple[str, ...]:
(
(("tests/e2e/logging/test_datadog_e2e.py", "litellm/router.py"), ("tests/e2e/logging/test_datadog_e2e.py",)),
(("tests/e2e/ui/test_keys.py", "tests/e2e/claude_code/test_cli.py", "tests/e2e/load/test_burst.py"), ()),
(("tests/e2e/migrations/test_startup.py", "tests/e2e/migrations/test_recovery.py"), ()),
(("tests/e2e/batches/test_managed_files_enforcement_e2e.py",), ()),
(("tests/e2e/guardrails/test_presidio_masking_e2e.py",), ()),
(("tests/e2e/llm_translation/realtime/test_realtime_pipecat_audio_e2e.py",), ()),
@ -181,6 +182,10 @@ def test_a_changed_canary_file_is_selected_once_alongside_a_harness_change() ->
assert select_tests((CANARY[1], "tests/e2e/proxy_client.py")) == CANARY
def test_dedicated_migration_tests_do_not_suppress_shared_harness_canaries() -> None:
assert select_tests(("tests/e2e/migrations/test_startup.py", "tests/e2e/conftest.py")) == CANARY
def test_the_canary_joins_directly_selected_files_in_sorted_order() -> None:
assert select_tests(("tests/e2e/logging/test_datadog_e2e.py", ".github/e2e-stack/up.sh")) == (
*CANARY,

View file

@ -2,10 +2,36 @@
Code-style rules for writing tests under `tests/e2e/`. The harness already encodes the plumbing; your job is the feature-specific behavior, not reinventing it. For what a complete test must do (the lifecycle contract, asserting both recorded state and enforced behavior) and how to run a suite, see `CONTRIBUTING.md` in this directory. Repo-wide conventions live in the root `AGENTS.md`
## What good looks like
Only what a real provider proves. If it holds against our scripted upstream: `tests/integration`
```python
def test_pre_call_masks_pii_on_chat_completions(self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str) -> None:
name = f"e2e-presidio-pre-chat-{unique_marker()}"
_register_presidio(client, resources, name=name)
email = _fake_email()
_assert_eventually_masked(
lambda: client.chat(scoped_key, MODEL, _pii_prompt(unique_marker(), email), guardrails=[name], max_tokens=128),
_first_content,
email=email,
)
```
Marker per run, so a leftover guardrail cannot pass it. `resources.defer(...)` at creation, so a failed
assert still tears down. Assert what the caller receives
## Where it goes
By the surface a customer would name: `guardrails`, `llm_translation`, `management`. Mutation check
deferred; it needs credentials
## Suite folders
Each subdirectory under `tests/e2e/` is one suite, scoped to an endpoint family or behavior area. If you add a new folder, you must add a line here describing what kind of tests belong in it, so the layout stays self-describing. `gateway/` is the exception: it holds proxy configuration only and never tests
- `migrations/` - isolated Docker startup, concurrent migration, crash recovery, and legacy database compatibility. The CircleCI migration workflow enables `LITELLM_MIGRATION_TESTS=1`; these tests own their proxy containers and databases, so they do not use the shared proxy preflight or shared database cleanup
- `llm_translation/` - LLM endpoint and provider-translation behavior: passthrough, custom pricing, OCR, and the non-chat inference endpoints (`/v1/responses`, `/v1/messages`, `/embeddings`, `/v1/rerank`, `/v1/audio/speech`, `/v1/images/generations`), each against a deployment the test creates via `/model/new` and deletes on teardown
- `access_control/` - the gateway's authorization and error-shape contract: per-key model allow-lists, route-group permissions (`allowed_routes`), and unknown-model validation
- `embeddings/` - the `/embeddings` endpoint across providers

View file

@ -91,6 +91,9 @@ def jwt_identity(idp: Keycloak, resources: ResourceManager, proxy: ProxyClient)
def pytest_configure(config: pytest.Config) -> None:
config.addinivalue_line(
"markers", "migration_startup: isolated container startup tests run by the migration CI workflow"
)
config.addinivalue_line(
"markers",
"provider_live: requires actual provider timing, limits, state, or a response that echoes this"
@ -185,6 +188,11 @@ def pytest_collection_modifyitems(config: pytest.Config, items: list[pytest.Item
items[:] = [item for item in items if not _needs_unset_opt_in(item)]
for item in items:
attach_result_properties(item)
if os.environ.get("LITELLM_MIGRATION_TESTS") != "1":
deselected = [item for item in items if item.get_closest_marker("migration_startup") is not None]
items[:] = [item for item in items if item.get_closest_marker("migration_startup") is None]
if deselected:
deselected[0].config.hook.pytest_deselected(items=deselected)
items.sort(key=lambda item: item.get_closest_marker("load") is not None)
@ -219,7 +227,7 @@ def pytest_runtest_setup(item: pytest.Item) -> None:
run even when none is up. Never skip for a missing proxy. Replay mode needs
the proxy too: only provider-bound traffic replays from the bundle."""
LIVE_PROVIDER_REQUIRED.set(item.get_closest_marker("provider_live") is not None)
if item.get_closest_marker("e2e") is None:
if item.get_closest_marker("e2e") is None or item.get_closest_marker("migration_startup") is not None:
return
if isinstance(item, pytest.Function) and "oauth_gateway" in item.fixturenames:
return
@ -234,7 +242,7 @@ def pytest_runtest_call(item: pytest.Item) -> None:
guard before truncating the spend-log DB. Tests under `tests/e2e/` without the
`e2e` marker (pure unit coverage for the harness itself) never hit the proxy,
so they must not arm the destructive DB truncate."""
if item.get_closest_marker("e2e") is None:
if item.get_closest_marker("e2e") is None or item.get_closest_marker("migration_startup") is not None:
return
item.session.stash[_E2E_TEST_RAN] = True

View file

@ -0,0 +1,134 @@
import hashlib
from contextlib import ExitStack
from typing import Final
from uuid import uuid4
from psycopg import sql
from .containers import Containers, Replica, failed, until
from .database import GATE_KEY, Database
from .startup_models import Migration
COMPLETE_SQL: Final = "CREATE TABLE migration_effect (id int PRIMARY KEY); INSERT INTO migration_effect VALUES (1);"
COMPLETE: Final = Migration("20990101000000_startup_test", COMPLETE_SQL)
NEXT: Final = Migration(
"20990102000000_next_test",
"CREATE TABLE migration_next (id int PRIMARY KEY); INSERT INTO migration_next VALUES (2);",
)
FATAL: Final = Migration(COMPLETE.name, "DO $$ BEGIN RAISE EXCEPTION 'MIGRATION_TEST_FATAL'; END $$;")
GATED: Final = Migration(
COMPLETE.name, f"SELECT pg_advisory_lock({GATE_KEY}); {COMPLETE.script} SELECT pg_advisory_unlock({GATE_KEY});"
)
def start_replicas(
stack: ExitStack, containers: Containers, database: Database, migrations: tuple[Migration, ...] = (), count: int = 3
) -> tuple[Replica, ...]:
return tuple(stack.enter_context(containers.start(database, migrations)) for _ in range(count))
def assert_completed(database: Database, migration: Migration = COMPLETE) -> None:
assert database.query(
'SELECT finished_at IS NOT NULL, rolled_back_at IS NULL, applied_steps_count FROM '
'_prisma_migrations WHERE migration_name = %s',
(migration.name,),
) == ((True, True, 1),), "Expected exactly one successful SQL execution"
assert database.query("SELECT id FROM migration_effect") == ((1,),)
def confirmed_history(database: Database) -> str:
database.execute(COMPLETE_SQL)
row_id: Final = str(uuid4())
database.execute(
"INSERT INTO _prisma_migrations (id, migration_name, checksum, applied_steps_count) VALUES (%s, %s, %s, 1)",
(row_id, COMPLETE.name, hashlib.sha256(COMPLETE.script.encode()).hexdigest()),
)
return row_id
def assert_original_proof(database: Database, row_id: str, finished: bool) -> None:
assert database.query(
'SELECT id, applied_steps_count, finished_at IS NOT NULL, rolled_back_at IS NULL FROM '
'_prisma_migrations WHERE migration_name = %s',
(COMPLETE.name,),
) == ((row_id, 1, finished, True),), "Recovery lost or replaced the original durable SQL proof"
assert database.query("SELECT id FROM migration_effect") == ((1,),)
def pause_completion(database: Database) -> None:
database.execute(
sql.SQL(
"CREATE FUNCTION migration_pause() RETURNS trigger LANGUAGE plpgsql AS $$ BEGIN "
"IF NEW.migration_name = {name} AND NEW.finished_at IS NOT NULL THEN "
"PERFORM pg_advisory_lock({gate}); PERFORM pg_advisory_unlock({gate}); END IF; RETURN NEW; END $$; "
'CREATE TRIGGER migration_pause BEFORE UPDATE ON _prisma_migrations FOR EACH ROW '
'EXECUTE FUNCTION migration_pause()'
).format(name=sql.Literal(COMPLETE.name), gate=sql.Literal(GATE_KEY))
)
def interrupt_owner(
containers: Containers, database: Database, after_commit: bool, *, stop_database_session: bool = True
) -> None:
if after_commit:
pause_completion(database)
with database.lock():
with containers.start(database, (COMPLETE if after_commit else GATED,)) as owner:
until("migration at the intended crash boundary", lambda: bool(database.blocked()))
assert database.exists("migration_effect") == after_commit
assert database.query(
"SELECT finished_at IS NULL FROM _prisma_migrations WHERE migration_name = %s", (COMPLETE.name,)
) == ((True,),)
blocked: Final = database.blocked()
assert len(blocked) == 1
backend: Final = blocked[0][0]
assert owner.state().Running
owner.kill()
assert owner.state().ExitCode == 137
if stop_database_session:
database.query("SELECT pg_terminate_backend(%s)", (backend,))
until(
"terminated migration backend released",
lambda: not database.query("SELECT pid FROM pg_stat_activity WHERE pid = %s", (backend,)),
)
assert database.query(
"SELECT finished_at IS NULL, applied_steps_count FROM _prisma_migrations WHERE migration_name = %s",
(COMPLETE.name,),
) == ((True, int(after_commit)),)
assert database.exists("migration_effect") == after_commit
if not stop_database_session:
until(
"database backend noticed container death",
lambda: not database.query("SELECT pid FROM pg_stat_activity WHERE pid = %s", (backend,)),
60,
)
def unconfirmed(replicas: tuple[Replica, ...], database: Database) -> None:
failed(replicas, "Migration completion could not be verified")
started: Final = str(
database.query(
"SELECT to_char(started_at AT TIME ZONE 'UTC', 'YYYY-MM-DD HH24:MI:SS') FROM "
'_prisma_migrations WHERE migration_name = %s',
(COMPLETE.name,),
)[0][0]
)
for replica in replicas:
assert_guidance(replica.logs(), started)
def assert_guidance(log: str, started: str) -> None:
for detail in (
COMPLETE.name,
started,
"cannot determine whether its SQL committed",
"_prisma_migrations",
"migration.sql",
"Only after verifying every migration change is present",
"prisma migrate resolve --applied <migration_name>",
"Only after verifying no migration changes remain",
"prisma migrate resolve --rolled-back <migration_name>",
"leave migration history unchanged",
"Repeated restarts alone",
):
assert detail in log, f"Missing recovery guidance: {detail}"

View file

@ -0,0 +1,62 @@
import json
import os
from collections.abc import Iterator
from pathlib import Path
from typing import Final
from urllib.parse import urlsplit
import pytest
from _pytest.fixtures import SubRequest
from .containers import Containers, docker, ready
from .database import Database, Databases
@pytest.fixture(scope="session")
def migration_image(tmp_path_factory: pytest.TempPathFactory) -> str:
configured: Final = os.environ.get("LITELLM_MIGRATION_TEST_IMAGE")
assert configured, "LITELLM_MIGRATION_TEST_IMAGE must name the built candidate image"
image: Final = docker("image", "inspect", configured, "--format", "{{.Id}}")
assert image.startswith("sha256:"), "Unable to identify the candidate image"
output: Final = Path(os.environ.get("MIGRATION_TEST_OUTPUT", str(tmp_path_factory.getbasetemp())))
output.mkdir(parents=True, exist_ok=True)
(output / "image.json").write_text(json.dumps({"requested": configured, "image_id": image}))
return image
@pytest.fixture(scope="session")
def databases() -> Databases:
admin: Final = os.environ.get("MIGRATION_TEST_ADMIN_URL", "")
parsed: Final = urlsplit(admin)
assert parsed.hostname in ("127.0.0.1", "localhost"), "Use an isolated loopback PostgreSQL test cluster"
assert parsed.port and parsed.path and not parsed.query, "Supply the test cluster port and admin database"
container_admin: Final = os.environ.get(
"MIGRATION_TEST_CONTAINER_ADMIN_URL",
admin.replace("127.0.0.1", "host.docker.internal").replace("localhost", "host.docker.internal"),
)
return Databases(admin, container_admin)
@pytest.fixture(scope="session")
def migrated_template(
databases: Databases, migration_image: str, tmp_path_factory: pytest.TempPathFactory
) -> Iterator[Database]:
output: Final = Path(os.environ.get("MIGRATION_TEST_OUTPUT", str(tmp_path_factory.getbasetemp()))) / "seed"
with databases.create() as database:
with Containers(migration_image, output).start(database) as replica:
ready((replica,), database)
yield database
@pytest.fixture
def database(databases: Databases, migrated_template: Database) -> Iterator[Database]:
with databases.create(migrated_template) as database:
yield database
@pytest.fixture
def containers(migration_image: str, tmp_path: Path, request: SubRequest) -> Containers:
configured: Final = os.environ.get("MIGRATION_TEST_OUTPUT")
output: Final = Path(configured) / request.node.name if configured else tmp_path
output.mkdir(parents=True, exist_ok=True)
return Containers(migration_image, output)

View file

@ -0,0 +1,203 @@
from __future__ import annotations
import hashlib
import subprocess
import time
from collections.abc import Callable, Generator, Mapping
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from typing import Final
from uuid import uuid4
from e2e_http import NoBody, Success, unwrap
from models import KeyGenerateBody, KeyGenerateResponse, KeyInfoParams, KeyInfoResponse
from transport import HttpTransport
from .database import Database, prisma_url
from .startup_models import ContainerState, Migration, Observation, Readiness
MASTER_KEY: Final = "sk-migration-ci-fixture"
def docker(*args: str) -> str:
result: Final = subprocess.run(("docker", *args), capture_output=True, text=True, timeout=90)
assert result.returncode == 0, f"Docker operation failed: {result.stderr}"
return result.stdout.strip()
def until(description: str, condition: Callable[[], bool], seconds: float = 150) -> None:
deadline: Final = time.monotonic() + seconds
while time.monotonic() < deadline:
if condition():
return
time.sleep(0.25)
raise AssertionError(f"Timed out waiting for {description}")
@dataclass(frozen=True, slots=True)
class Replica:
name: str
transport: HttpTransport
output: Path
def state(self) -> ContainerState:
return ContainerState.model_validate_json(docker("inspect", "--format", "{{json .State}}", self.name))
def observe(self) -> Observation:
state: Final = self.state()
result: Final = self.transport.get(
"/health/readiness", headers=self.transport.master, params=NoBody(), response_type=Readiness, timeout=1
)
ready: Final = isinstance(result, Success) and result.data.status == "healthy" and result.data.db == "connected"
return Observation(None if state.Running else state.ExitCode, ready)
def logs(self) -> str:
result: Final = subprocess.run(("docker", "logs", self.name), capture_output=True, text=True, timeout=30)
assert result.returncode == 0, result.stderr
return result.stdout + result.stderr
def kill(self) -> None:
if self.state().Running:
docker("kill", self.name)
def usable(self, database: Database) -> None:
alias: Final = f"migration-{uuid4().hex}"
key: Final = unwrap(
self.transport.post(
"/key/generate",
headers=self.transport.master,
json=KeyGenerateBody(key_alias=alias),
response_type=KeyGenerateResponse,
)
).key
info: Final = unwrap(
self.transport.get(
"/key/info",
headers=self.transport.master,
params=KeyInfoParams(key=key),
response_type=KeyInfoResponse,
)
)
assert info.info.key_alias == alias
assert database.query(
'SELECT key_alias FROM "LiteLLM_VerificationToken" WHERE token = %s',
(hashlib.sha256(key.encode()).hexdigest(),),
) == ((alias,),)
def ready(replicas: tuple[Replica, ...], database: Database) -> None:
def all_ready() -> bool:
observations: Final = tuple(replica.observe() for replica in replicas)
assert all(item.exit_code is None for item in observations), "Replica exited before readiness"
return all(item.ready for item in observations)
until("every replica ready", all_ready)
for replica in replicas:
replica.usable(database)
def failed(replicas: tuple[Replica, ...], marker: str) -> None:
def all_stopped() -> bool:
observations: Final = tuple(replica.observe() for replica in replicas)
assert not any(item.ready for item in observations), "Failed migration exposed a ready proxy"
return all(item.exit_code is not None for item in observations)
until("every replica to reject startup", all_stopped)
for replica in replicas:
assert replica.state().ExitCode != 0, "Failed startup returned success"
assert marker in replica.logs(), f"Startup failed outside the expected migration: {marker}"
def waiting(replicas: tuple[Replica, ...], seconds: float) -> None:
deadline: Final = time.monotonic() + seconds
while time.monotonic() < deadline:
assert all(item.exit_code is None and not item.ready for item in (replica.observe() for replica in replicas)), (
"Contending replica exited or served early"
)
time.sleep(0.25)
@dataclass(frozen=True, slots=True)
class Containers:
image: str
output: Path
@contextmanager
def start(
self,
database: Database,
migrations: tuple[Migration, ...] = (),
*,
v2: bool = True,
disabled: bool = False,
environment: Mapping[str, str] | None = None,
) -> Generator[Replica]:
name: Final = f"litellm-migration-{uuid4().hex[:16]}"
directory: Final = self.output / name
directory.mkdir(parents=True)
for migration in migrations:
write_migration(directory, migration)
(directory / "config.yaml").write_text(
"model_list: []\ngeneral_settings:\n master_key: os.environ/LITELLM_MASTER_KEY\n"
)
env: Final = {
"DATABASE_URL": prisma_url(database.container_url, database.schema),
"LITELLM_MASTER_KEY": MASTER_KEY,
"LITELLM_SALT_KEY": MASTER_KEY,
"LITELLM_LOCAL_MODEL_COST_MAP": "True",
"LITELLM_TELEMETRY": "False",
"LITELLM_LOG": "INFO",
"DATABASE_CONNECTION_POOL_LIMIT": "2",
"DEFAULT_NUM_WORKERS_LITELLM_PROXY": "1",
"USE_V2_MIGRATION_RESOLVER": str(v2).lower(),
"DISABLE_SCHEMA_UPDATE": str(disabled).lower(),
"LITELLM_MIGRATION_DIR": "/migration-test/prisma",
"LITELLM_PRISMA_MIGRATE_DEPLOY_TIMEOUT": "180",
**(environment or {}),
}
try:
docker(
"run",
"-d",
"--name",
name,
"--label",
"litellm-migration-test=true",
"--add-host",
"host.docker.internal:host-gateway",
"-p",
"127.0.0.1::4000",
"-v",
f"{directory}:/migration-test",
*(arg for key, value in env.items() for arg in ("-e", f"{key}={value}")),
self.image,
"--config",
"/migration-test/config.yaml",
"--host",
"0.0.0.0",
"--port",
"4000",
)
port: Final = int(docker("port", name, "4000/tcp").rsplit(":", 1)[1])
replica: Final = Replica(name, HttpTransport(f"http://127.0.0.1:{port}", MASTER_KEY, 15), directory)
yield replica
finally:
try:
state: Final = subprocess.run(
("docker", "inspect", "--format", "{{json .State}}", name),
capture_output=True,
text=True,
timeout=30,
)
(directory / "state.json").write_text(state.stdout or state.stderr)
logs: Final = subprocess.run(("docker", "logs", name), capture_output=True, text=True, timeout=30)
(directory / "proxy.log").write_text(logs.stdout + logs.stderr)
finally:
subprocess.run(("docker", "rm", "-f", name), capture_output=True, text=True, timeout=30, check=True)
def write_migration(directory: Path, migration: Migration) -> None:
path: Final = directory / "prisma" / "migrations" / migration.name
path.mkdir(parents=True)
(path / "migration.sql").write_text(migration.script)

View file

@ -0,0 +1,135 @@
from __future__ import annotations
from collections.abc import Generator
from contextlib import contextmanager
from dataclasses import dataclass
from typing import Final, LiteralString
from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit
from uuid import uuid4
import psycopg
from psycopg import sql
from pydantic import TypeAdapter
Scalar = str | int | bool | None
ROWS: Final = TypeAdapter(tuple[tuple[Scalar, ...], ...])
GATE_KEY: Final = 39178002
PRISMA_LOCK: Final = 72707369
COORDINATOR_LOCK: Final = int.from_bytes(b"llm_mig2", "big")
def connect_url(url: str, name: str) -> str:
return urlunsplit(urlsplit(url)._replace(path=f"/{name}", query=""))
def prisma_url(url: str, schema: str) -> str:
parsed: Final = urlsplit(url)
query: Final = tuple((key, value) for key, value in parse_qsl(parsed.query) if key != "schema")
return urlunsplit(parsed._replace(query=urlencode((*query, ("schema", schema)))))
@dataclass(frozen=True, slots=True)
class Database:
name: str
url: str
container_url: str
schema: str = "public"
@contextmanager
def connection(self) -> Generator[psycopg.Connection[tuple[object, ...]]]:
with psycopg.connect(self.url, autocommit=True, connect_timeout=5) as connection:
connection.execute(sql.SQL("SET search_path TO {}").format(sql.Identifier(self.schema)))
connection.execute("SET statement_timeout = '15s'")
yield connection
def execute(self, statement: LiteralString | sql.Composed, params: tuple[Scalar, ...] = ()) -> None:
with self.connection() as connection:
connection.execute(statement, params or None)
def query(
self, statement: LiteralString | sql.Composed, params: tuple[Scalar, ...] = ()
) -> tuple[tuple[Scalar, ...], ...]:
with self.connection() as connection:
return ROWS.validate_python(connection.execute(statement, params or None).fetchall())
def exists(self, name: str) -> bool:
return self.query("SELECT to_regclass(%s) IS NOT NULL", (name,)) == ((True,),)
def history(self) -> tuple[tuple[Scalar, ...], ...]:
if not self.exists("_prisma_migrations"):
return ()
return self.query(
"SELECT id, migration_name, checksum, started_at::text, finished_at::text, rolled_back_at::text, "
"applied_steps_count, logs FROM _prisma_migrations ORDER BY id"
)
def blocked(self, key: int = GATE_KEY) -> tuple[tuple[Scalar, ...], ...]:
return self.query(
"SELECT pid FROM pg_locks WHERE locktype = 'advisory' AND NOT granted "
"AND database = (SELECT oid FROM pg_database WHERE datname = current_database()) "
"AND classid = %s AND objid = %s ORDER BY pid",
(key >> 32, key & 0xFFFFFFFF),
)
@contextmanager
def lock(self, key: int = GATE_KEY) -> Generator[None]:
with self.connection() as connection:
connection.execute("SELECT pg_advisory_lock(%s)", (key,))
try:
yield
finally:
connection.execute("SELECT pg_advisory_unlock(%s)", (key,))
@dataclass(frozen=True, slots=True)
class Databases:
admin_url: str
container_admin_url: str
@contextmanager
def create(self, template: Database | None = None, schema: str = "public") -> Generator[Database]:
name: Final = f"litellm_migration_test_{uuid4().hex[:20]}"
database: Final = Database(
name, connect_url(self.admin_url, name), connect_url(self.container_admin_url, name), schema
)
with psycopg.connect(self.admin_url, autocommit=True, connect_timeout=5) as connection:
connection.execute(
sql.SQL("CREATE DATABASE {} TEMPLATE {}").format(
sql.Identifier(name), sql.Identifier(template.name if template else "template0")
)
)
try:
yield database
finally:
with psycopg.connect(self.admin_url, autocommit=True, connect_timeout=5) as connection:
connection.execute(sql.SQL("DROP DATABASE {} WITH (FORCE)").format(sql.Identifier(name)))
@contextmanager
def restricted_user(database: Database) -> Generator[Database]:
role: Final = f"migration_reader_{uuid4().hex[:16]}"
password: Final = "migration-test-password"
with database.connection() as connection:
connection.execute(
sql.SQL("CREATE ROLE {} LOGIN PASSWORD {}").format(sql.Identifier(role), sql.Literal(password))
)
try:
database.execute(
sql.SQL("GRANT USAGE ON SCHEMA {} TO {}").format(sql.Identifier(database.schema), sql.Identifier(role))
)
database.execute(
sql.SQL("GRANT SELECT ON ALL TABLES IN SCHEMA {} TO {}").format(
sql.Identifier(database.schema), sql.Identifier(role)
)
)
local: Final = urlsplit(database.url)
remote: Final = urlsplit(database.container_url)
yield Database(
database.name,
urlunsplit(local._replace(netloc=f"{role}:{password}@{local.hostname}:{local.port}")),
urlunsplit(remote._replace(netloc=f"{role}:{password}@{remote.hostname}:{remote.port}")),
database.schema,
)
finally:
database.execute(sql.SQL("DROP OWNED BY {}").format(sql.Identifier(role)))
database.execute(sql.SQL("DROP ROLE {}").format(sql.Identifier(role)))

View file

@ -0,0 +1,25 @@
from dataclasses import dataclass
from pydantic import BaseModel
class Readiness(BaseModel):
status: str = ""
db: str = ""
class ContainerState(BaseModel):
Running: bool
ExitCode: int
@dataclass(frozen=True, slots=True)
class Observation:
exit_code: int | None
ready: bool
@dataclass(frozen=True, slots=True)
class Migration:
name: str
script: str

View file

@ -0,0 +1,87 @@
from contextlib import ExitStack
from dataclasses import replace
from typing import Final, Literal
import pytest
from .checks import COMPLETE, assert_completed, confirmed_history, assert_original_proof, start_replicas
from .containers import Containers, failed, ready
from .database import Database, Databases
pytestmark: Final = [pytest.mark.e2e, pytest.mark.migration_startup]
def adopt_legacy(containers: Containers, database: Database) -> None:
count: Final = database.query("SELECT count(*) FROM _prisma_migrations")[0][0]
existing_keys: Final = database.query('SELECT token FROM "LiteLLM_VerificationToken" ORDER BY token')
database.execute(
'INSERT INTO "LiteLLM_ShadowEvalJob" (id, group_id, target_id, router_name, judge_model, '
"shadow_percentage, max_turns, ends_at, stopped_at) VALUES ('migration-legacy', "
"'migration-legacy', 'target', 'router', 'judge', 1, 1, now(), now())"
)
database.execute("DROP TABLE _prisma_migrations")
with ExitStack() as stack:
replicas: Final = start_replicas(stack, containers, database)
ready(replicas, database)
logs: Final = "\n".join(replica.logs() for replica in replicas)
for detail in (
"Legacy migration history was missing",
"historical data backfills were not replayed or verified",
"Continuing startup",
):
assert detail in logs
assert database.query("SELECT count(*) FROM _prisma_migrations") == ((count,),)
assert database.query(
'SELECT count(*) FROM _prisma_migrations WHERE finished_at IS NULL OR rolled_back_at IS '
'NOT NULL OR applied_steps_count <> 0'
) == ((0,),)
assert set(existing_keys).issubset(database.query('SELECT token FROM "LiteLLM_VerificationToken" ORDER BY token'))
assert database.query("SELECT stopped_by FROM \"LiteLLM_ShadowEvalJob\" WHERE id = 'migration-legacy'") == (
(None,),
)
class TestLegacyMigrations:
def test_matching_schema_warns_and_starts(self, containers: Containers, database: Database) -> None:
adopt_legacy(containers, database)
@pytest.mark.parametrize("fault", ("schema_drift", "custom_migrations", "empty_ledger"))
def test_unrecognized_legacy_state_is_not_baselined(
self, containers: Containers, database: Database, fault: str
) -> None:
if fault == "empty_ledger":
database.execute("TRUNCATE _prisma_migrations")
else:
database.execute("DROP TABLE _prisma_migrations")
if fault == "schema_drift":
database.execute('ALTER TABLE "LiteLLM_VerificationToken" DROP COLUMN key_alias CASCADE')
with containers.start(database, (COMPLETE,) if fault == "custom_migrations" else ()) as replica:
failed((replica,), "Cannot automatically baseline" if fault != "empty_ledger" else "migration")
assert not database.exists("migration_effect")
if database.exists("_prisma_migrations"):
assert database.query(
"SELECT count(*) FROM _prisma_migrations WHERE finished_at IS NOT NULL AND applied_steps_count <> 1"
) == ((0,),)
@pytest.mark.parametrize("scenario", ("upgrade", "recovery", "legacy"))
def test_non_default_schema(
self, containers: Containers, databases: Databases, scenario: Literal["upgrade", "recovery", "legacy"]
) -> None:
with databases.create(schema="migration tenant") as database:
with containers.start(database) as seed:
ready((seed,), database)
match scenario:
case "upgrade":
with ExitStack() as stack:
ready(start_replicas(stack, containers, database, (COMPLETE,)), database)
assert_completed(database)
case "recovery":
original: Final = confirmed_history(database)
with ExitStack() as stack:
ready(start_replicas(stack, containers, database, (COMPLETE,)), database)
assert_original_proof(database, original, True)
case "legacy":
adopt_legacy(containers, database)
public: Final = replace(database, schema="public")
assert not public.exists("_prisma_migrations")
assert not public.exists('"LiteLLM_VerificationToken"')

View file

@ -0,0 +1,136 @@
import subprocess
from collections.abc import Generator
from contextlib import ExitStack, contextmanager
from pathlib import Path
from typing import Final
from urllib.parse import urlsplit, urlunsplit
from uuid import uuid4
import psycopg
import pytest
from psycopg import sql
from .checks import COMPLETE, assert_completed
from .containers import Containers, docker, ready, until
from .database import Database, Databases, prisma_url, restricted_user
POOL_IMAGE: Final = (
"ghcr.io/cloudnative-pg/pgbouncer@sha256:e6ddfe22d845e603825e235dd8334b21ecd125abea2a2172478f556b8dee2bb8"
)
pytestmark: Final = [pytest.mark.e2e, pytest.mark.migration_startup]
@contextmanager
def application_user(database: Database) -> Generator[Database]:
with restricted_user(database) as application:
role: Final = sql.Identifier(str(urlsplit(application.url).username))
schema: Final = sql.Identifier(database.schema)
database.execute(sql.SQL("REVOKE CREATE ON SCHEMA {} FROM PUBLIC").format(schema))
for statement in (
"GRANT SELECT, INSERT, UPDATE, DELETE ON ALL TABLES IN SCHEMA {} TO {}",
"GRANT USAGE, SELECT ON ALL SEQUENCES IN SCHEMA {} TO {}",
"ALTER DEFAULT PRIVILEGES IN SCHEMA {} GRANT SELECT, INSERT, UPDATE, DELETE ON TABLES TO {}",
"ALTER DEFAULT PRIVILEGES IN SCHEMA {} GRANT USAGE, SELECT ON SEQUENCES TO {}",
):
database.execute(sql.SQL(statement).format(schema, role))
assert application.query("SELECT has_schema_privilege(current_user, %s, 'CREATE')", (database.schema,)) == (
(False,),
)
yield application
@contextmanager
def pool(database: Database, output: Path) -> Generator[str]:
name: Final = f"litellm-migration-pool-{uuid4().hex[:12]}"
url: Final = urlsplit(database.container_url)
directory: Final = output / name
directory.mkdir(parents=True)
(directory / "users.txt").write_text(f'"{url.username}" "{url.password}"\n')
(directory / "pgbouncer.ini").write_text(
f"[databases]\n* = host={url.hostname} port={url.port} user={url.username} password={url.password}\n"
"[pgbouncer]\nlisten_addr = 0.0.0.0\nlisten_port = 6432\nauth_type = trust\nauth_file = /pool/users.txt\n"
"pool_mode = transaction\ndefault_pool_size = 1\nreserve_pool_size = 0\nmax_client_conn = 100\n"
'max_prepared_statements = 100\nquery_wait_timeout = 8\nignore_startup_parameters = '
'extra_float_digits,options\n'
)
try:
docker(
"run",
"-d",
"--name",
name,
"--label",
"litellm-migration-test=true",
"--add-host",
"host.docker.internal:host-gateway",
"-p",
"0.0.0.0::6432",
"-v",
f"{directory}:/pool:ro",
"--entrypoint",
"/usr/bin/pgbouncer",
POOL_IMAGE,
"/pool/pgbouncer.ini",
)
port: Final = int(docker("port", name, "6432/tcp").splitlines()[0].rsplit(":", 1)[1])
local_url: Final = urlunsplit(url._replace(netloc=f"{url.username}:{url.password}@127.0.0.1:{port}"))
def connected() -> bool:
try:
with psycopg.connect(local_url, autocommit=True, connect_timeout=2) as connection:
return connection.execute("SELECT 1").fetchone() == (1,)
except psycopg.Error:
return False
until("PgBouncer ready", connected, 30)
yield local_url.replace("127.0.0.1", "host.docker.internal") + "?pgbouncer=true"
finally:
try:
logs: Final = subprocess.run(("docker", "logs", name), text=True, capture_output=True, timeout=30)
(directory / "pool.log").write_text(logs.stdout + logs.stderr)
finally:
subprocess.run(("docker", "rm", "-f", name), capture_output=True, text=True, timeout=30, check=True)
class TestMigrationPooling:
@pytest.mark.parametrize("scenario,replica_count", (("fresh", 3), ("upgrade", 3), ("legacy", 3), ("upgrade", 6)))
def test_direct_migrations_with_one_application_backend(
self,
containers: Containers,
databases: Databases,
migrated_template: Database,
scenario: str,
replica_count: int,
) -> None:
with databases.create(None if scenario == "fresh" else migrated_template) as database:
if scenario == "legacy":
database.execute("DROP TABLE _prisma_migrations")
with (
application_user(database) as application,
pool(application, containers.output) as pooled_url,
ExitStack() as stack,
):
replicas: Final = tuple(
stack.enter_context(
containers.start(
database,
(COMPLETE,) if scenario == "upgrade" else (),
environment={
"DATABASE_URL": prisma_url(pooled_url, database.schema),
"DIRECT_URL": database.container_url,
},
)
)
for _ in range(replica_count)
)
ready(replicas, database)
if scenario == "upgrade":
assert_completed(database)
if scenario == "legacy":
assert any(
"historical data backfills were not replayed or verified" in replica.logs()
for replica in replicas
)
assert database.query("SELECT count(*) FROM _prisma_migrations WHERE applied_steps_count <> 0") == (
(0,),
)

View file

@ -0,0 +1,185 @@
from contextlib import ExitStack
from typing import Final, Literal
from uuid import uuid4
import pytest
from .checks import (
COMPLETE,
FATAL,
GATED,
NEXT,
assert_completed,
confirmed_history,
interrupt_owner,
assert_original_proof,
pause_completion,
start_replicas,
unconfirmed,
)
from .containers import Containers, failed, ready, until, waiting
from .database import COORDINATOR_LOCK, GATE_KEY, Database
from .startup_models import Migration
pytestmark: Final = [pytest.mark.e2e, pytest.mark.migration_startup]
class TestMigrationRecovery:
@pytest.mark.parametrize("after_commit", (False, True))
def test_container_owner_crash(self, containers: Containers, database: Database, after_commit: bool) -> None:
interrupt_owner(containers, database, after_commit, stop_database_session=False)
history: Final = database.history()
assert database.query(
"SELECT applied_steps_count FROM _prisma_migrations WHERE migration_name = %s", (COMPLETE.name,)
) == ((int(after_commit),),)
with ExitStack() as stack:
successors: Final = start_replicas(stack, containers, database, (COMPLETE if after_commit else GATED,))
if after_commit:
ready(successors, database)
assert_completed(database)
else:
unconfirmed(successors, database)
assert database.history() == history
@pytest.mark.parametrize("after_commit", (False, True))
def test_owner_and_database_session_crash(
self, containers: Containers, database: Database, after_commit: bool
) -> None:
interrupt_owner(containers, database, after_commit)
history: Final = database.history()
with ExitStack() as stack:
successors: Final = start_replicas(stack, containers, database, (COMPLETE,))
if after_commit:
ready(successors, database)
assert_completed(database)
return
unconfirmed(successors, database)
assert database.history() == history
with containers.start(database, (COMPLETE,)) as restarted:
unconfirmed((restarted,), database)
assert database.history() == history
@pytest.mark.parametrize("later_failure", (False, True))
def test_remaining_migrations_after_recovery(
self, containers: Containers, database: Database, later_failure: bool
) -> None:
original: Final = confirmed_history(database)
next_migration: Final = Migration(
NEXT.name,
f"SELECT pg_advisory_lock({GATE_KEY}); "
+ (FATAL.script if later_failure else NEXT.script)
+ f" SELECT pg_advisory_unlock({GATE_KEY});",
)
with ExitStack() as stack:
with database.lock():
owner: Final = stack.enter_context(containers.start(database, (COMPLETE, next_migration)))
def pending() -> bool:
observation: Final = owner.observe()
assert observation.exit_code is None and not observation.ready, (
"Recovered owner served before pending SQL completed"
)
return bool(database.blocked())
until("recovering owner reached the next migration", pending)
assert_original_proof(database, original, True)
assert not database.exists("migration_next")
followers: Final = start_replicas(stack, containers, database, (COMPLETE, next_migration), count=2)
replicas: Final = (owner, *followers)
waiting(replicas, 1)
if later_failure:
failed(replicas, NEXT.name)
assert database.query(
"SELECT finished_at IS NULL, logs LIKE %s FROM _prisma_migrations WHERE migration_name = %s",
("%MIGRATION_TEST_FATAL%", NEXT.name),
) == ((True, True),)
else:
ready(replicas, database)
assert database.query("SELECT id FROM migration_next") == ((2,),)
assert_original_proof(database, original, True)
def test_second_crash_during_recovery_is_atomic(self, containers: Containers, database: Database) -> None:
original: Final = confirmed_history(database)
pause_completion(database)
with database.lock():
with containers.start(database, (COMPLETE,)) as recovering:
until("history update blocked before commit", lambda: bool(database.blocked()))
assert_original_proof(database, original, False)
blocked: Final = database.blocked()
assert len(blocked) == 1
assert database.query("SELECT pg_terminate_backend(%s)", (blocked[0][0],)) == ((True,),)
failed((recovering,), "Lost or could not establish v2 migration coordination")
assert_original_proof(database, original, False)
with ExitStack() as stack:
ready(start_replicas(stack, containers, database, (COMPLETE,)), database)
assert_original_proof(database, original, True)
def test_competing_recovery_rechecks_stale_failures(self, containers: Containers, database: Database) -> None:
original: Final = confirmed_history(database)
with ExitStack() as stack:
with database.lock(COORDINATOR_LOCK):
replicas: Final = start_replicas(stack, containers, database, (COMPLETE,))
until(
"all replicas observed the unfinished migration",
lambda: all(
"Waiting for the v2 migration coordinator lock" in replica.logs() for replica in replicas
),
)
assert_original_proof(database, original, False)
ready(replicas, database)
assert_original_proof(database, original, True)
@pytest.mark.parametrize(
"fault", ("no_steps", "extra_steps", "failure_logs", "checksum", "duplicate_history", "missing_script")
)
def test_unproven_history_is_never_repaired(
self,
containers: Containers,
database: Database,
fault: Literal["no_steps", "extra_steps", "failure_logs", "checksum", "duplicate_history", "missing_script"],
) -> None:
confirmed_history(database)
match fault:
case "no_steps":
database.execute(
"UPDATE _prisma_migrations SET applied_steps_count = 0 WHERE migration_name = %s", (COMPLETE.name,)
)
case "extra_steps":
database.execute(
"UPDATE _prisma_migrations SET applied_steps_count = 2 WHERE migration_name = %s", (COMPLETE.name,)
)
case "failure_logs":
database.execute(
"UPDATE _prisma_migrations SET logs = 'permission denied' WHERE migration_name = %s",
(COMPLETE.name,),
)
case "checksum":
database.execute(
"UPDATE _prisma_migrations SET checksum = %s WHERE migration_name = %s", ("0" * 64, COMPLETE.name)
)
case "duplicate_history":
database.execute(
'INSERT INTO _prisma_migrations (id, migration_name, checksum, '
'applied_steps_count) SELECT %s, migration_name, checksum, '
'applied_steps_count FROM _prisma_migrations WHERE migration_name = %s',
(str(uuid4()), COMPLETE.name),
)
case "missing_script":
pass
history: Final = database.history()
with containers.start(database, () if fault == "missing_script" else (COMPLETE,)) as replica:
unconfirmed((replica,), database)
assert database.history() == history
assert database.query("SELECT id FROM migration_effect") == ((1,),)
def test_coordinator_timeout_preserves_proof(self, containers: Containers, database: Database) -> None:
original: Final = confirmed_history(database)
with database.lock(COORDINATOR_LOCK):
with containers.start(
database, (COMPLETE,), environment={"LITELLM_MIGRATION_LOCK_TIMEOUT": "3"}
) as replica:
failed((replica,), "Timed out waiting for another v2 migration resolver")
assert_original_proof(database, original, False)
with containers.start(database, (COMPLETE,)) as replica:
ready((replica,), database)
assert_original_proof(database, original, True)

View file

@ -0,0 +1,101 @@
from contextlib import ExitStack
from typing import Final
import pytest
from .checks import COMPLETE, FATAL, GATED, assert_completed, start_replicas
from .containers import Containers, failed, ready, until, waiting
from .database import PRISMA_LOCK, Database, Databases, restricted_user
from .startup_models import Migration
pytestmark: Final = [pytest.mark.e2e, pytest.mark.migration_startup]
class TestMigrationStartup:
@pytest.mark.parametrize("replicas,v2", ((1, True), (3, True), (1, False)))
def test_fresh_database(self, containers: Containers, databases: Databases, replicas: int, v2: bool) -> None:
with databases.create() as database, ExitStack() as stack:
ready(tuple(stack.enter_context(containers.start(database, v2=v2)) for _ in range(replicas)), database)
assert database.query(
"SELECT count(*) FROM _prisma_migrations WHERE finished_at IS NULL AND rolled_back_at IS NULL"
) == ((0,),)
assert database.query("SELECT count(*) > 0 FROM _prisma_migrations") == ((True,),)
def test_concurrent_upgrade(self, containers: Containers, database: Database) -> None:
with ExitStack() as stack:
ready(start_replicas(stack, containers, database, (COMPLETE,)), database)
assert_completed(database)
def test_waiters_survive_prolonged_contention(self, containers: Containers, database: Database) -> None:
with ExitStack() as stack:
with database.lock():
owner: Final = stack.enter_context(containers.start(database, (GATED,)))
until("owner blocked in migration SQL", lambda: bool(database.blocked()))
followers: Final = start_replicas(stack, containers, database, (GATED,), count=2)
until("both followers attempted Prisma locking", lambda: len(database.blocked(PRISMA_LOCK)) == 2)
waiting((owner, *followers), 120)
ready((owner, *followers), database)
assert_completed(database, GATED)
def test_lock_deadline_then_restart(self, containers: Containers, database: Database) -> None:
history: Final = database.history()
with database.lock(PRISMA_LOCK):
with containers.start(
database, (COMPLETE,), environment={"LITELLM_MIGRATION_LOCK_TIMEOUT": "12"}
) as replica:
until("Prisma lock contention", lambda: bool(database.blocked(PRISMA_LOCK)))
failed((replica,), "Timed out waiting for")
assert database.history() == history
assert not database.exists("migration_effect")
with containers.start(database, (COMPLETE,)) as restarted:
ready((restarted,), database)
assert_completed(database)
def test_fatal_sql(self, containers: Containers, database: Database) -> None:
with ExitStack() as stack:
replicas: Final = start_replicas(stack, containers, database, (FATAL,))
failed(replicas, COMPLETE.name)
assert database.query(
'SELECT count(*) FROM _prisma_migrations WHERE migration_name = %s AND logs LIKE '
'%s AND finished_at IS NULL',
(COMPLETE.name, "%MIGRATION_TEST_FATAL%"),
) == ((1,),)
def test_duplicate_object_does_not_hide_incomplete_sql(self, containers: Containers, database: Database) -> None:
database.execute(
"CREATE TABLE migration_existing (id int PRIMARY KEY); INSERT INTO migration_existing VALUES (42)"
)
migration: Final = Migration(
COMPLETE.name, "CREATE TABLE migration_existing (id int PRIMARY KEY); " + COMPLETE.script
)
with ExitStack() as stack:
failed(start_replicas(stack, containers, database, (migration,)), COMPLETE.name)
assert not database.exists("migration_effect")
assert database.query("SELECT id FROM migration_existing") == ((42,),)
assert database.query(
"SELECT finished_at IS NULL FROM _prisma_migrations WHERE migration_name = %s", (COMPLETE.name,)
) == ((True,),)
@pytest.mark.parametrize("v2", (True, False))
def test_restart_preserves_history_and_data(self, containers: Containers, database: Database, v2: bool) -> None:
history: Final = database.history()
before: Final = database.query('SELECT token FROM "LiteLLM_VerificationToken" ORDER BY token')
for _ in range(2):
with containers.start(database, v2=v2) as replica:
ready((replica,), database)
assert database.history() == history
assert set(before).issubset(database.query('SELECT token FROM "LiteLLM_VerificationToken" ORDER BY token'))
def test_disabled_migrations(self, containers: Containers, database: Database) -> None:
history: Final = database.history()
with containers.start(database, (FATAL,), disabled=True) as replica:
ready((replica,), database)
assert database.history() == history
def test_insufficient_privileges(self, containers: Containers, database: Database) -> None:
history: Final = database.history()
with restricted_user(database) as limited:
with containers.start(limited, (COMPLETE,)) as replica:
failed((replica,), "permission denied")
assert database.history() == history
assert not database.exists("migration_effect")

View file

@ -0,0 +1,25 @@
# tests/integration
Real proxy, Postgres, Redis, scripted upstream. `README.md` has shards and CI wiring
## What good looks like
The root example is from here. The spend row lands async: poll, never sleep
```python
rows = eventually(
lambda: read_rows('SELECT spend FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (body["id"],)),
lambda values: len(values) == 1,
seconds=70,
)
assert float(rows[0]["spend"]) == pytest.approx(10 * 0.001 + 5 * 0.0001 + 7 * 0.002 + 4 * 0.002)
```
`sleep(3)` fails on a slow runner and taxes every fast one. Assert the outbound body in the upstream
handler; a leaked field is invisible from the response. `monkeypatch.setenv` is fine; patching our own
function in a full stack is not
## Where it goes
By the domain a user would name: `pricing`, `spend`, `routing`. Add the node and its `covers` ids to
`contracts.json` or collection fails. Needs no proxy, DB or Redis: `tests/unit`

View file

@ -74,8 +74,6 @@ def owned_proxy(gateway: Gateway, directory: Path, overrides: Mapping[str, str],
str(port),
"--num_workers",
"1",
"--telemetry",
"False",
"--use_prisma_db_push",
"--enforce_prisma_migration_check",
],

View file

@ -728,12 +728,36 @@ ERROR: relation "SomeTable" already exists
"""
@pytest.mark.parametrize(
"pooled,direct,expected",
(
("postgresql://pool/db?pgbouncer=true", None, "postgresql://pool/db?pgbouncer=true"),
("postgresql://pool/db?pgbouncer=true", "postgresql://writer/db", "postgresql://writer/db?schema=public"),
(
"postgresql://pool/db?schema=tenant%20one&pgbouncer=true",
"postgresql://writer/db?sslmode=require&schema=wrong",
"postgresql://writer/db?sslmode=require&schema=tenant+one",
),
),
)
def test_v2_migrations_use_the_direct_connection_with_the_runtime_schema(pooled, direct, expected):
from litellm_proxy_extras.migration_lock import migration_environment
environment = {"DATABASE_URL": pooled, "PRISMA_OFFLINE_MODE": "true"}
configured = {**environment, **({"DIRECT_URL": direct} if direct else {})}
migrated = migration_environment(configured)
assert migrated["DATABASE_URL"] == expected
assert migrated["PRISMA_OFFLINE_MODE"] == "true"
assert configured["DATABASE_URL"] == pooled
class _MigrateDeployHarness:
"""Drives _setup_database_v2 with a scripted sequence of
`prisma migrate deploy` outcomes, with every recovery command faked out so
nothing touches a database or the packaged migrations directory."""
def __init__(self, monkeypatch, tmp_path, outcomes, repeat_last=False):
def __init__(self, monkeypatch, tmp_path, outcomes, repeat_last=False, confirmed_migrations=()):
import subprocess as subprocess_module
import litellm_proxy_extras.utils as utils_module
@ -744,34 +768,20 @@ class _MigrateDeployHarness:
self._outcomes = list(outcomes)
self._repeat_last = repeat_last
self._subprocess_module = subprocess_module
self.confirmed_migrations = set(confirmed_migrations)
monkeypatch.delenv("DATABASE_URL", raising=False)
monkeypatch.setattr(
ProxyExtrasDBManager, "_get_prisma_dir", staticmethod(lambda: str(tmp_path))
)
monkeypatch.setattr(
ProxyExtrasDBManager,
"_create_baseline_migration",
staticmethod(self._fake_baseline),
)
monkeypatch.setattr(
ProxyExtrasDBManager,
"_roll_back_migration",
staticmethod(lambda name: None),
)
monkeypatch.setattr(
ProxyExtrasDBManager,
"_resolve_specific_migration",
staticmethod(self.resolved.append),
)
monkeypatch.setenv("LITELLM_MIGRATION_DIR", str(tmp_path))
monkeypatch.setattr(utils_module.prisma_toolchain, "run_prisma", self._fake_run)
monkeypatch.setattr(utils_module, "_get_prisma_env", lambda: {})
monkeypatch.setattr(utils_module.time, "sleep", lambda seconds: None)
self.baseline_succeeds = True
def _fake_baseline(self, *args, **kwargs):
self.baselines += 1
return self.baseline_succeeds
if not self.baseline_succeeds:
raise RuntimeError("The existing schema was not verified")
def _next_outcome(self):
if self._outcomes:
@ -791,79 +801,126 @@ class _MigrateDeployHarness:
raise self._subprocess_module.CalledProcessError(1, cmd, stderr=outcome)
def run(self):
return ProxyExtrasDBManager._setup_database_v2(use_migrate=True)
while not ProxyExtrasDBManager._run_database_v2(
use_migrate=True,
recover_completed=self._fake_recovery,
baseline_existing=self._fake_baseline,
):
continue
return True
def _fake_recovery(self, name):
if name not in self.confirmed_migrations:
return False
self.confirmed_migrations.remove(name)
self.resolved.append(name)
return True
class TestMigrateDeployAttemptAccounting:
"""A `prisma db push` database has a full schema and no ledger, so the v2
resolver baselines it and then works through every migration whose objects
already exist. Those recoveries make progress, so they must not spend the
retry budget, which is there to stop a run that is getting nowhere."""
def test_a_push_created_database_finishes_bootstrapping(
self, monkeypatch, tmp_path
):
already_there = [
"20250329084805_new_cron_job_table",
"20250806095134_rename_alias_to_server_name_mcp_table",
"20260224203854_add_agent_object_permissions_table",
"20260301120000_fourth_table",
"20260302120000_fifth_table",
"20260303120000_sixth_table",
]
def test_a_push_created_database_finishes_bootstrapping(self, monkeypatch, tmp_path):
harness = _MigrateDeployHarness(
monkeypatch,
tmp_path,
[_P3005_STDERR]
+ [_p3018_stderr(name) for name in already_there]
+ ["ok"],
[_P3005_STDERR, "ok"],
)
assert harness.run() is True
assert harness.baselines == 1
assert harness.resolved == already_there
assert len(harness.deploy_calls) == len(already_there) + 2
assert harness.resolved == []
assert len(harness.deploy_calls) == 2
def test_repeated_recovery_of_one_migration_still_gives_up(
self, monkeypatch, tmp_path
):
def test_repeated_recovery_of_one_migration_still_gives_up(self, monkeypatch, tmp_path):
harness = _MigrateDeployHarness(
monkeypatch,
tmp_path,
[_p3018_stderr("20250329084805_new_cron_job_table")],
repeat_last=True,
confirmed_migrations=("20250329084805_new_cron_job_table",),
)
with pytest.raises(RuntimeError):
harness.run()
assert len(harness.deploy_calls) <= _ATTEMPT_BUDGET + 1
assert len(harness.deploy_calls) == 2
assert harness.resolved == ["20250329084805_new_cron_job_table"]
def test_timeouts_still_spend_the_budget(self, monkeypatch, tmp_path):
harness = _MigrateDeployHarness(
monkeypatch, tmp_path, ["timeout"], repeat_last=True
)
harness = _MigrateDeployHarness(monkeypatch, tmp_path, ["timeout"], repeat_last=True)
with pytest.raises(RuntimeError):
harness.run()
assert len(harness.deploy_calls) == _ATTEMPT_BUDGET
def test_a_baseline_that_never_lands_stops_after_the_budget(
self, monkeypatch, tmp_path
):
harness = _MigrateDeployHarness(
monkeypatch, tmp_path, [_P3005_STDERR], repeat_last=True
)
def test_an_unverified_baseline_stops_without_replaying_migrations(self, monkeypatch, tmp_path):
harness = _MigrateDeployHarness(monkeypatch, tmp_path, [_P3005_STDERR], repeat_last=True)
harness.baseline_succeeds = False
with pytest.raises(RuntimeError):
harness.run()
assert len(harness.deploy_calls) == _ATTEMPT_BUDGET
assert len(harness.deploy_calls) == 1
def test_lock_contention_does_not_spend_the_failure_budget(self, monkeypatch, tmp_path):
harness = _MigrateDeployHarness(
monkeypatch,
tmp_path,
["Error: P1002\nTimed out waiting for the advisory lock"] * 6 + ["ok"],
)
assert harness.run() is True
assert len(harness.deploy_calls) == 7
def test_duplicate_object_error_without_completion_proof_is_fatal(self, monkeypatch, tmp_path):
harness = _MigrateDeployHarness(monkeypatch, tmp_path, [_p3018_stderr("20260101000000_x")])
with pytest.raises(RuntimeError, match="cannot be auto-recovered"):
harness.run()
assert harness.resolved == []
assert len(harness.deploy_calls) == 1
@pytest.mark.parametrize("name", ("20260101000000_x", "20260101000000_migration with spaces"))
def test_an_interrupted_migration_with_confirmed_sql_can_finish(self, monkeypatch, tmp_path, name):
harness = _MigrateDeployHarness(
monkeypatch,
tmp_path,
[f"Error: P3009\nThe `{name}` migration failed", "ok"],
confirmed_migrations=(name,),
)
assert harness.run() is True
assert harness.resolved == [name]
def test_an_interrupted_migration_without_confirmation_stops(self, monkeypatch, tmp_path):
name = "20260101000000_x"
started = "2026-09-12 20:15:06.694553 UTC"
report = f"Error: P3009\nThe `{name}` migration started at {started} failed"
harness = _MigrateDeployHarness(
monkeypatch,
tmp_path,
[report],
)
with pytest.raises(RuntimeError, match="Migration completion could not be verified") as failure:
harness.run()
message = str(failure.value)
assert name in message
assert started in message
assert "start record but no successful completion record" in message
assert "cannot determine whether its SQL committed" in message
assert "avoid repeating or skipping database changes" in message
assert "_prisma_migrations" in message
assert "migration.sql" in message
assert "same database" in message
assert "Only after verifying every migration change is present" in message
assert "prisma migrate resolve --applied <migration_name>" in message
assert "Only after verifying no migration changes remain" in message
assert "prisma migrate resolve --rolled-back <migration_name>" in message
assert "leave migration history unchanged" in message
assert "Repeated restarts alone" in message
assert report in message
assert len(harness.deploy_calls) == 1
assert harness.resolved == []
def test_an_unrecoverable_error_is_not_retried(self, monkeypatch, tmp_path):
harness = _MigrateDeployHarness(
monkeypatch,
tmp_path,
["Error: P3018\n\nMigration name: 20260101000000_x\n\nERROR: syntax error at or near \"SLECT\"\n"],
['Error: P3018\n\nMigration name: 20260101000000_x\n\nERROR: syntax error at or near "SLECT"\n'],
repeat_last=True,
)
@ -873,6 +930,36 @@ class TestMigrateDeployAttemptAccounting:
assert harness.resolved == []
@pytest.mark.parametrize(
"steps,logs,script,expected",
(
(1, "", b"CREATE TABLE item (id int);", True),
(0, "", b"CREATE TABLE item (id int);", False),
(0, "already exists", b"CREATE TABLE item (id int);", False),
(1, "permission denied", b"CREATE TABLE item (id int);", False),
(1, "", b"CREATE TABLE item (id text);", False),
(2, "", b"CREATE TABLE item (id int);", False),
),
)
def test_migration_completion_requires_a_matching_successful_script(steps, logs, script, expected):
import hashlib
from litellm_proxy_extras.migration_recovery import MigrationProgress
progress = MigrationProgress(hashlib.sha256(b"CREATE TABLE item (id int);").hexdigest(), steps, logs)
assert progress.confirms_completion(script) is expected
def test_prisma_lock_waiting_has_its_own_deadline():
from litellm_proxy_extras.utils import _MigrateAttemptBudget
budget = _MigrateAttemptBudget(attempts_left=4, contention_seconds_left=2)
waiting = budget.after_contention(1)
assert waiting.attempts_left == 4
with pytest.raises(RuntimeError, match="advisory lock"):
waiting.after_contention(2)
class TestJWTKeyMappingCascade:
"""Regression tests for issue #33702.

View file

@ -883,7 +883,6 @@ async def test_provider_specific_fields_in_proxy_http_response(
max_tokens=None,
request_timeout=600,
max_budget=None,
telemetry=False,
drop_params=True,
add_function_to_prompt=False,
headers=None,

View file

@ -0,0 +1,36 @@
import importlib.util
from pathlib import Path
from typing import Final
import pytest
SCRIPT: Final = Path(__file__).resolve().parents[2] / ".circleci/scripts/run_migration_tests.py"
SPEC: Final = importlib.util.spec_from_file_location("migration_ci", SCRIPT)
assert SPEC is not None and SPEC.loader is not None
MODULE: Final = importlib.util.module_from_spec(SPEC)
SPEC.loader.exec_module(MODULE)
@pytest.mark.parametrize(
"xml,expected,exit_code,passed",
(
('<testsuites><testsuite><testcase name="a"/><testcase name="b"/></testsuite></testsuites>', 2, 0, True),
('<testsuites><testsuite><testcase name="a"/></testsuite></testsuites>', 2, 0, False),
("<testsuite><testcase><failure/></testcase></testsuite>", 1, 0, False),
("<testsuite><testcase><error/></testcase></testsuite>", 1, 0, False),
("<testsuite><testcase><skipped/></testcase></testsuite>", 1, 0, False),
("<testsuite><testcase/></testsuite>", 1, 1, False),
("<testsuite><testcase/></testsuite>", 1, 5, False),
("<testsuite/>", 1, 0, False),
("<broken", 1, 0, False),
(None, 1, 0, False),
('<testsuite><testcase name="a"/><testcase name="a"/></testsuite>', 2, 0, False),
),
)
def test_only_a_complete_passing_suite_can_certify_an_image(
tmp_path: Path, xml: str | None, expected: int, exit_code: int, passed: bool
) -> None:
path: Final = tmp_path / "results.xml"
if xml is not None:
path.write_text(xml)
assert MODULE.successful_junit(path, expected, exit_code) is passed

View file

@ -1112,7 +1112,6 @@ def test_settings_store_preserves_yaml_team_configuration_when_db_value_is_null(
},
"param_name": "litellm_settings",
"db_param_value": {
"telemetry": False,
"drop_params": True,
"num_retries": 5,
"request_timeout": 600,

View file

@ -1,195 +0,0 @@
import os
import pytest
# Ensure the project root is on the import path
from litellm import completion
from litellm.types.utils import ModelResponse, Usage, Choices, Message
def _has_api_key() -> bool:
"""Check if Amazon Nova API key is available"""
return (
"AMAZON_NOVA_API_KEY" in os.environ
and os.environ["AMAZON_NOVA_API_KEY"] is not None
)
def _create_mock_nova_response():
"""Helper function to create mock Amazon Nova response for testing"""
return ModelResponse(
id="chatcmpl-test-nova-micro",
choices=[
Choices(
finish_reason="stop",
index=0,
message=Message(
content="I am Amazon Nova Micro. 777 times 9 equals 6993.",
role="assistant",
),
)
],
created=1234567890,
model="amazon-nova/nova-micro-v1",
object="chat.completion",
usage=Usage(prompt_tokens=25, completion_tokens=15, total_tokens=40),
)
def test_amazon_nova_chat_completion_nova_micro():
if _has_api_key():
response: ModelResponse = completion(
model="amazon-nova/nova-micro-v1",
messages=[
{"role": "system", "content": "You are a helpful assistant"},
{
"role": "user",
"content": "What model are you? Can you calculate 777 times 9?",
},
],
api_key=os.environ["AMAZON_NOVA_API_KEY"],
)
else:
# Use mock response when API key is not available
response = _create_mock_nova_response()
# Additional mock-specific assertions for code review reference
assert (
response.choices[0].message.content
== "I am Amazon Nova Micro. 777 times 9 equals 6993."
)
assert response.model == "amazon-nova/nova-micro-v1"
assert response.usage.prompt_tokens == 25
assert response.usage.completion_tokens == 15
assert response.object == "chat.completion"
assert response.choices[0].finish_reason == "stop"
assert response.choices[0].message.role == "assistant"
# Common assertions for both real and mock responses
assert response is not None
assert hasattr(response, "choices")
assert len(response.choices) > 0
assert response.choices[0].message.content is not None
assert response.usage.total_tokens > 0
@pytest.mark.skipif(not _has_api_key(), reason="Amazon Nova API key not available")
def test_amazon_nova_chat_completion_nova_lite():
response: ModelResponse = completion(
model="amazon-nova/nova-lite-v1",
messages=[
{"role": "system", "content": "You are a helpful assistant"},
{
"role": "user",
"content": "What model are you? Please tell me a poem on rain",
},
],
api_key=os.environ["AMAZON_NOVA_API_KEY"],
)
assert response is not None
assert hasattr(response, "choices")
assert len(response.choices) > 0
assert response.choices[0].message.content is not None
assert response.usage.total_tokens > 0
@pytest.mark.skipif(not _has_api_key(), reason="Amazon Nova API key not available")
def test_amazon_nova_chat_completion_nova_pro():
response: ModelResponse = completion(
model="amazon-nova/nova-pro-v1",
messages=[
{"role": "system", "content": "You are a helpful assistant"},
{
"role": "user",
"content": "What model are you? What is MCP server and how does that help in building GenAI applications?",
},
],
timeout=30,
api_key=os.environ["AMAZON_NOVA_API_KEY"],
)
assert response is not None
assert hasattr(response, "choices")
assert len(response.choices) > 0
assert response.choices[0].message.content is not None
assert response.usage.total_tokens > 0
@pytest.mark.skipif(not _has_api_key(), reason="Amazon Nova API key not available")
def test_amazon_nova_chat_completion_nova_premier():
response: ModelResponse = completion(
model="amazon-nova/nova-premier-v1",
messages=[
{"role": "system", "content": "You are a helpful assistant"},
{
"role": "user",
"content": "What model are you? Can you help me understand what Trigonometry is?",
},
],
timeout=60,
api_key=os.environ["AMAZON_NOVA_API_KEY"],
)
assert response is not None
print(response.choices[0].message.content)
assert hasattr(response, "choices")
assert len(response.choices) > 0
assert response.choices[0].message.content is not None
assert response.usage.total_tokens > 0
@pytest.mark.skipif(not _has_api_key(), reason="Amazon Nova API key not available")
def test_amazon_nova_chat_completion_with_tool_usage():
response: ModelResponse = completion(
model="amazon-nova/nova-micro-v1",
messages=[
{"role": "system", "content": "You are a helpful assistant"},
{"role": "user", "content": "What is the temperature in SFO?"},
],
tools=[
{
"type": "function",
"function": {
"name": "getCurrentWeather",
"description": "Get the current weather in a given city",
"parameters": {
"type": "object",
"properties": {
"location": {
"type": "string",
"description": "City and country e.g. Bogotá, Colombia",
}
},
"required": ["location"],
},
},
}
],
api_key=os.environ["AMAZON_NOVA_API_KEY"],
)
assert response is not None
assert hasattr(response, "choices")
assert len(response.choices) > 0
assert response.choices[0].message is not None
@pytest.mark.skipif(not _has_api_key(), reason="Amazon Nova API key not available")
def test_amazon_nova_chat_completion_with_stream_response():
response = completion(
model="amazon-nova/nova-micro-v1",
stream=True,
messages=[
{"role": "system", "content": "You are a helpful assistant"},
{
"role": "user",
"content": "What are MMO games? Can you give me some sample references?",
},
],
api_key=os.environ["AMAZON_NOVA_API_KEY"],
)
assert response is not None
chunks = list(response)
assert chunks is not None
assert len(chunks) > 0

View file

@ -1,115 +0,0 @@
"""
Test Bedrock files integration with main files API
"""
import base64
from unittest.mock import MagicMock, patch
import pytest
import litellm
from litellm.types.llms.openai import HttpxBinaryResponseContent
from litellm.types.utils import SpecialEnums
class TestBedrockFilesIntegration:
"""Test integration of Bedrock files with main litellm API"""
@pytest.mark.asyncio
async def test_litellm_afile_content_bedrock_provider_with_s3_uri(self):
"""Test litellm.afile_content with bedrock provider using direct S3 URI"""
file_id = "s3://test-bucket/test-file.jsonl"
expected_content = (
b'{"recordId": "request-1", "modelInput": {}, "modelOutput": {}}'
)
# Create a mock HttpxBinaryResponseContent response
import httpx
mock_response = httpx.Response(
status_code=200,
content=expected_content,
headers={"content-type": "application/octet-stream"},
request=httpx.Request(method="GET", url="s3://test-bucket/test-file.jsonl"),
)
mock_result = HttpxBinaryResponseContent(response=mock_response)
# Mock the base_llm_http_handler.retrieve_file_content since the code
# now routes through ProviderConfigManager -> base_llm_http_handler
with patch(
"litellm.files.main.base_llm_http_handler.retrieve_file_content",
new_callable=MagicMock,
) as mock_retrieve:
mock_retrieve.return_value = mock_result
# Call litellm.afile_content
result = await litellm.afile_content(
file_id=file_id,
custom_llm_provider="bedrock",
aws_region_name="us-west-2",
)
# Verify the result
assert isinstance(result, HttpxBinaryResponseContent)
assert result.response.content == expected_content
assert result.response.status_code == 200
# Verify the mock was called with correct parameters
mock_retrieve.assert_called_once()
call_kwargs = mock_retrieve.call_args.kwargs
assert call_kwargs["_is_async"] is True
assert call_kwargs["file_content_request"]["file_id"] == file_id
@pytest.mark.asyncio
async def test_litellm_afile_content_bedrock_provider_with_unified_file_id(self):
"""Test litellm.afile_content with bedrock provider using unified file ID"""
# Create a unified file ID
s3_uri = "s3://test-bucket/batch-outputs/output.jsonl"
unified_id = "test-unified-id-123"
model_id = "test-model-id-456"
unified_file_id_str = f"litellm_proxy:application/json;unified_id,{unified_id};target_model_names,;llm_output_file_id,{s3_uri};llm_output_file_model_id,{model_id}"
encoded_file_id = (
base64.urlsafe_b64encode(unified_file_id_str.encode()).decode().rstrip("=")
)
expected_content = (
b'{"recordId": "request-1", "modelInput": {}, "modelOutput": {}}'
)
# Create a mock HttpxBinaryResponseContent response
import httpx
mock_response = httpx.Response(
status_code=200,
content=expected_content,
headers={"content-type": "application/octet-stream"},
request=httpx.Request(method="GET", url=s3_uri),
)
mock_result = HttpxBinaryResponseContent(response=mock_response)
# Mock the base_llm_http_handler.retrieve_file_content
with patch(
"litellm.files.main.base_llm_http_handler.retrieve_file_content",
new_callable=MagicMock,
) as mock_retrieve:
mock_retrieve.return_value = mock_result
# Call litellm.afile_content with unified file ID
result = await litellm.afile_content(
file_id=encoded_file_id,
custom_llm_provider="bedrock",
aws_region_name="us-west-2",
)
# Verify the result
assert isinstance(result, HttpxBinaryResponseContent)
assert result.response.content == expected_content
assert result.response.status_code == 200
# Verify the mock was called
mock_retrieve.assert_called_once()
call_kwargs = mock_retrieve.call_args.kwargs
assert call_kwargs["_is_async"] is True
# The handler passes the encoded file_id as-is
assert call_kwargs["file_content_request"]["file_id"] == encoded_file_id

View file

@ -47,7 +47,6 @@ WANDB_REASONING_MODELS: Final = (
@pytest.fixture
def wandb_test_config(local_model_cost_map, monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(litellm, "disable_aiohttp_transport", True)
monkeypatch.setattr(litellm, "telemetry", False)
monkeypatch.setattr(litellm, "drop_params", False)

View file

@ -43,6 +43,7 @@ from litellm.proxy.proxy_server import (
cost_tracking,
get_litellm_model_info,
initialize,
initialize_from_worker_config,
load_from_azure_key_vault,
proxy_shutdown_event,
proxy_startup_event,
@ -376,7 +377,7 @@ def _lit4152_worker_config_dict():
"master_key": _LIT4152_SECRETS[0],
"database_url": _LIT4152_SECRETS[3],
"api_key": _LIT4152_SECRETS[2],
"telemetry": True,
"drop_params": True,
}
@ -394,7 +395,7 @@ def test__redact_worker_config_for_logging_dict_masks_all_secret_shapes():
assert secret not in rendered, f"leak: {secret} in {rendered!r}"
assert isinstance(redacted, dict)
assert redacted["model"] == "openai/gpt-4o-mini"
assert redacted["telemetry"] is True
assert redacted["drop_params"] is True
def test__redact_worker_config_for_logging_json_string_round_trips_masked():
@ -501,7 +502,7 @@ def test__redact_worker_config_for_logging_masks_nested_secret_fields():
def test_initialize_signature_is_async_with_expected_params():
sig = inspect.signature(initialize)
# Hard-coded so a signature change (param added/removed) trips the gate.
expected_param_count = 17
expected_param_count = 16
observed = {
"is_async": inspect.iscoroutinefunction(initialize),
"param_count": len(sig.parameters),
@ -522,6 +523,16 @@ async def test_initialize_invalid_unexpected_kwarg_raises_type_error():
await initialize(this_is_not_a_real_kwarg=True)
@pytest.mark.asyncio
async def test_initialize_from_worker_config_drops_legacy_telemetry_key():
with pytest.raises(TypeError):
await initialize(telemetry=True)
await initialize_from_worker_config({"telemetry": True, "request_timeout": 77})
assert ps.user_request_timeout == 77
with pytest.raises(TypeError):
await initialize_from_worker_config({"this_is_not_a_real_kwarg": True})
# ---------------------------------------------------------------------------
# load_from_azure_key_vault
# ---------------------------------------------------------------------------

View file

@ -617,6 +617,15 @@ class TestProxyInitializationHelpers:
assert "Skipping server startup" in result.output
mock_uvicorn_run.assert_not_called()
result = runner.invoke(
run_server, ["--local", "--skip_server_startup", "--telemetry", "False"]
)
assert (
result.exit_code == 0
), f"exit_code={result.exit_code}, output={result.output}"
assert "Skipping server startup" in result.output
assert "telemetry" not in runner.invoke(run_server, ["--help"]).output
# --- normal startup ---
mock_uvicorn_run.reset_mock()

View file

@ -9,6 +9,7 @@ import importlib.util
import subprocess
import sys
from pathlib import Path
from typing import Final
_MODULE_PATH = (
Path(__file__).resolve().parents[2] / "scripts" / "budget_ratchet_check.py"
@ -92,6 +93,36 @@ def test_graduation_never_excuses_a_raised_limit():
assert "0 -> 7" in regs[0].detail
def test_dropped_rule_the_checker_retired_is_clean():
base: Final = {"TQ008": _spec_of(10993)}
assert ratchet.regressions_for("b.json", base, {}, retired=frozenset({"TQ008"})) == []
def test_dropped_rule_the_checker_still_emits_is_a_regression():
base: Final = {"TQ001": _spec_of(5), "TQ008": _spec_of(10993)}
regs: Final = ratchet.regressions_for("b.json", base, {}, retired=frozenset({"TQ008"}))
assert [r.rule for r in regs] == ["TQ001"]
assert "dropped" in regs[0].detail
def test_retirement_never_excuses_a_raised_limit():
base: Final = {"TQ008": _spec_of(0)}
regs: Final = ratchet.regressions_for("b.json", base, {"TQ008": _spec_of(7)}, retired=frozenset({"TQ008"}))
assert [r.rule for r in regs] == ["TQ008"]
assert "0 -> 7" in regs[0].detail
def test_retired_rules_come_from_the_paired_checker():
base: Final = {"TQ001": _spec_of(5), "TQ008": _spec_of(10993)}
assert ratchet.retired_rules("test-quality-budget.json", base) == frozenset({"TQ008"})
def test_budgets_without_a_paired_checker_never_retire():
base: Final = {"TQ008": _spec_of(1)}
for rel in ("ruff-strict-budget.json", "type-discipline-budget.json", "basedpyright-code-budget.json"):
assert ratchet.retired_rules(rel, base) == frozenset()
def test_graduated_selectors_come_from_the_paired_ruff_config():
selectors = ratchet.graduated_selectors("ruff-strict-budget.json")
assert "UP006" in selectors

View file

@ -12,6 +12,8 @@ import os
import subprocess
import sys
from pathlib import Path
from types import MappingProxyType
from typing import Final
import pytest
@ -187,7 +189,7 @@ def test_mock_echo_is_flagged(tmp_path):
" run()\n"
" mock_completion.assert_called_once()\n"
)
assert _codes(tmp_path, source) == ["TQ002", "TQ008"]
assert _codes(tmp_path, source) == ["TQ002"]
def test_call_args_inspection_is_mock_echo(tmp_path):
@ -200,7 +202,7 @@ def test_call_args_inspection_is_mock_echo(tmp_path):
" run()\n"
" assert mock_completion.call_args[1]['model'] == 'gpt-4o'\n"
)
assert _codes(tmp_path, source) == ["TQ002", "TQ008"]
assert _codes(tmp_path, source) == ["TQ002"]
def test_patch_decorator_counts_as_installing_a_patch(tmp_path):
@ -213,7 +215,7 @@ def test_patch_decorator_counts_as_installing_a_patch(tmp_path):
" run()\n"
" mock_completion.assert_called_once()\n"
)
assert _codes(tmp_path, source) == ["TQ002", "TQ008"]
assert _codes(tmp_path, source) == ["TQ002"]
def test_patching_but_asserting_the_output_is_not_mock_echo(tmp_path):
@ -227,7 +229,7 @@ def test_patching_but_asserting_the_output_is_not_mock_echo(tmp_path):
" mock_completion.assert_called_once()\n"
" assert result.choices[0].message.content == 'pong'\n"
)
assert _codes(tmp_path, source) == ["TQ008"]
assert _codes(tmp_path, source) == []
def test_asserting_without_patching_is_not_mock_echo(tmp_path):
@ -244,7 +246,7 @@ def test_a_test_with_no_assertions_is_tq001_not_tq002(tmp_path):
" with patch('litellm.completion'):\n"
" run()\n"
)
assert _codes(tmp_path, source) == ["TQ001", "TQ008"]
assert _codes(tmp_path, source) == ["TQ001"]
def test_sys_path_insert_is_flagged(tmp_path):
@ -554,133 +556,6 @@ def test_a_loop_storing_under_a_key_that_is_not_the_loop_variable_is_not_an_inve
assert [v.code for v in checker.check_file(_written(tmp_path, source))] == []
def test_patching_an_sdk_function_by_string_is_flagged(tmp_path):
source = 'from unittest.mock import patch\n\n\n@patch("litellm.completion")\ndef test_x(m):\n assert m\n'
assert "TQ008" in _codes(tmp_path, source)
def test_patching_a_deep_sdk_path_is_flagged(tmp_path):
source = (
"from unittest.mock import patch\n\n\n"
"def test_x():\n"
' with patch("litellm.llms.openai.chat.handler.OpenAIChatCompletion.completion"):\n'
" assert True\n"
)
assert "TQ008" in _codes(tmp_path, source)
def test_patch_object_rooted_at_the_sdk_is_flagged(tmp_path):
source = (
"import litellm\nfrom unittest.mock import patch\n\n\n"
"def test_x():\n"
' with patch.object(litellm, "api_key", "x"):\n'
" assert True\n"
)
assert "TQ008" in _codes(tmp_path, source)
def test_patch_object_on_a_from_imported_sdk_module_is_flagged(tmp_path):
source = (
"from litellm.llms.openai.chat import handler\nfrom unittest.mock import patch\n\n\n"
"def test_x():\n"
' with patch.object(handler.OpenAIChatCompletion, "completion"):\n'
" assert True\n"
)
assert "TQ008" in _codes(tmp_path, source)
def test_patch_object_on_an_aliased_sdk_module_is_flagged(tmp_path):
source = (
"import litellm.llms.openai.chat.handler as oai\nfrom unittest.mock import patch\n\n\n"
"def test_x():\n"
' with patch.object(oai.OpenAIChatCompletion, "completion"):\n'
" assert True\n"
)
assert "TQ008" in _codes(tmp_path, source)
def test_patch_object_on_a_renamed_sdk_symbol_is_flagged(tmp_path):
source = (
"from litellm.utils import get_llm_provider as glp\nfrom unittest.mock import patch\n\n\n"
"def test_x():\n"
' with patch.object(glp, "__wrapped__"):\n'
" assert True\n"
)
assert "TQ008" in _codes(tmp_path, source)
def test_the_reported_target_is_the_resolved_sdk_path(tmp_path):
source = (
"from litellm.llms.openai.chat import handler\nfrom unittest.mock import patch\n\n\n"
"def test_x():\n"
' with patch.object(handler.OpenAIChatCompletion, "completion"):\n'
" assert True\n"
)
reported = [v.message for v in checker.check_file(_written(tmp_path, source)) if v.code == "TQ008"]
assert reported
assert "litellm.llms.openai.chat.handler.OpenAIChatCompletion" in reported[0]
def test_patch_object_on_a_from_imported_third_party_is_not_flagged(tmp_path):
source = (
"from openai import OpenAI\nfrom unittest.mock import patch\n\n\n"
"def test_x():\n"
' with patch.object(OpenAI, "chat"):\n'
" assert True\n"
)
assert "TQ008" not in _codes(tmp_path, source)
def test_a_local_name_with_no_sdk_import_behind_it_is_not_flagged(tmp_path):
source = (
"from unittest.mock import patch\n\n\n"
"def test_x(handler):\n"
' with patch.object(handler, "completion"):\n'
" assert True\n"
)
assert "TQ008" not in _codes(tmp_path, source)
def test_mocking_a_third_party_client_is_not_flagged(tmp_path):
source = (
"from unittest.mock import patch\n\n\n"
"def test_x():\n"
' with patch("openai.OpenAI.chat"):\n'
" assert True\n"
)
assert "TQ008" not in _codes(tmp_path, source)
def test_mocking_the_http_transport_is_not_flagged(tmp_path):
source = (
"from unittest.mock import patch\n\n\n"
"def test_x():\n"
' with patch("httpx.AsyncClient.send"):\n'
" assert True\n"
)
assert "TQ008" not in _codes(tmp_path, source)
def test_a_name_merely_starting_with_litellm_is_not_the_sdk(tmp_path):
source = (
"from unittest.mock import patch\n\n\n"
"def test_x():\n"
' with patch("litellm_enterprise.thing.go"):\n'
" assert True\n"
)
assert "TQ008" not in _codes(tmp_path, source)
def test_an_sdk_patch_can_be_suppressed(tmp_path):
source = (
"from unittest.mock import patch\n\n\n"
"def test_x():\n"
' with patch("litellm.completion"): # test-quality-ok: pinning the router seam\n'
" assert True\n"
)
assert "TQ008" not in _codes(tmp_path, source)
_FANS_OUT = checker._worker_count(checker.PARALLEL_MIN_PATHS) > 1
_SERIAL_ONLY = "one usable core, so scan_paths stays serial and there is no fan-out to compare"
@ -739,6 +614,44 @@ def test_a_fanned_out_run_reports_each_generated_file_exactly_once(tmp_path):
assert all(" TQ001 " in line for line in reported)
_VIOLATING_SNIPPETS: Final = MappingProxyType(
{
"TQ000": ("test_snippet.py", "def test_broken(:\n pass\n"),
"TQ001": ("test_snippet.py", "def test_nothing():\n compute()\n"),
"TQ002": (
"test_snippet.py",
"from unittest.mock import patch\n"
"\n"
"\n"
"def test_echo():\n"
" with patch('litellm.completion') as mock_completion:\n"
" run()\n"
" mock_completion.assert_called_once()\n",
),
"TQ003": ("test_snippet.py", "import sys\n\nsys.path.insert(0, '..')\n"),
"TQ004": ("test_snippet.py", "import os\n\nos.environ['KEY'] = 'v'\n"),
"TQ005": ("test_snippet.py", "import litellm\n\nlitellm.drop_params = True\n"),
"TQ006": ("test_snippet.py", _DIRECT_GATE),
"TQ007": ("conftest.py", _SNAPSHOT_CONFTEST),
"TQ009": (
"test_snippet.py",
'import subprocess, sys\nsubprocess.run([sys.executable, "-c", "pass"])\n',
),
}
)
def test_rule_codes_match_every_code_the_checker_emits(tmp_path):
emitted: Final = frozenset(
v.code
for name, source in _VIOLATING_SNIPPETS.values()
for v in checker.check_file(_written(tmp_path, source, name))
)
for code, (name, source) in _VIOLATING_SNIPPETS.items():
assert code in [v.code for v in checker.check_file(_written(tmp_path, source, name))], code
assert emitted == checker.RULE_CODES
def test_sys_executable_child_without_isolation_flag_is_flagged(tmp_path):
source = 'import subprocess, sys\nsubprocess.run([sys.executable, "-c", "pass"])\n'
assert _codes(tmp_path, source) == ["TQ009"]

View file

@ -137,7 +137,7 @@ def test_the_shipped_budget_covers_every_rule_the_checker_can_emit():
budget = json.loads((_REPO_ROOT / "test-quality-budget.json").read_text())
assert set(budget) == {
"TQ001", "TQ002", "TQ003", "TQ004", "TQ005", "TQ006", "TQ007", "TQ008", "TQ009"
"TQ001", "TQ002", "TQ003", "TQ004", "TQ005", "TQ006", "TQ007", "TQ009"
}
assert all(spec["limit"] >= 0 for spec in budget.values())

37
tests/unit/AGENTS.md Normal file
View file

@ -0,0 +1,37 @@
# tests/unit
In-process. No network, clock or subprocess
## What good looks like
```python
def test_send_result_same_version_is_identity_passthrough():
rpc = _rpc(V03_MESSAGE)
out = normalize_jsonrpc_response(rpc, "0.3", method="message/send")
assert out is rpc
```
`is`, because `==` passes on a copy. Many inputs: parametrize
(`test_an_incomplete_reservation_accrues_nothing`, fifteen cases, fifteen results)
No doubles on our own code
```python
with patch.object(streamer, "_group_by_date") as mock_group, patch.object(streamer, "_send_daily_batch") as mock_send:
mock_group.return_value = {"2025-01-19": pl.DataFrame({"test": ["data1"]}), "2025-01-20": pl.DataFrame({"test": ["data2"]})}
streamer.send_batched(pl.DataFrame({"test": ["data"]}), "replace_hourly")
assert mock_send.call_count == 2
```
Green if `send_batched` drops every row. pydantic doubles in 12 of 203 files, fastapi 11 of 594;
`tests/test_litellm` 59 percent. Exception: the count is the behaviour
(`test_dual_cache_async_batch_get_cache_coalesces_concurrent_redis_reads`, fifty readers, `call_count == 1`)
## Where it goes
`tests/unit/<path>` mirrors `litellm/<path>`, so a changed file selects its tests by path, not a mapping
file. Empty today; new unit tests go here. The examples above live in `tests/test_litellm`
## Writing it so a human can read it
A class only when tests share an arrange

25
tests/unit/conftest.py Normal file
View file

@ -0,0 +1,25 @@
import os
from typing import Final
import pytest
from pytest_socket import enable_socket, socket_allow_hosts
os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True"
LOOPBACK_HOSTS: Final = ["127.0.0.1", "::1"]
def _allow_loopback_only() -> None:
socket_allow_hosts(LOOPBACK_HOSTS, allow_unix_socket=True)
_allow_loopback_only()
@pytest.hookimpl(trylast=True)
def pytest_runtest_setup() -> None:
_allow_loopback_only()
def pytest_sessionfinish() -> None:
enable_socket()

View file

@ -151,48 +151,6 @@ class TestLevoConfig(unittest.TestCase):
class TestLevoIntegration(unittest.TestCase):
"""Integration tests for LevoLogger."""
@patch.dict(
"os.environ",
{
"LEVOAI_API_KEY": "test-api-key",
"LEVOAI_ORG_ID": "test-org-id",
"LEVOAI_WORKSPACE_ID": "test-workspace-id",
"LEVOAI_COLLECTOR_URL": "https://collector.levo.ai",
},
)
@pytest.mark.skipif(
not OPENTELEMETRY_AVAILABLE, reason="OpenTelemetry packages not installed"
)
@patch(
"litellm.integrations.opentelemetry.OpenTelemetry._init_otel_logger_on_litellm_proxy"
)
@pytest.mark.asyncio
async def test_levo_logger_health_check_healthy(self, mock_init_proxy):
"""Test health check returns healthy status when config is valid."""
# Mock the proxy initialization to avoid importing proxy code
mock_init_proxy.return_value = None
config = LevoLogger.get_levo_config()
otel_config = OpenTelemetryConfig(
exporter=config.protocol,
endpoint=config.endpoint,
headers=config.otlp_auth_headers,
)
# Create tracer provider with in-memory exporter
tracer_provider = TracerProvider()
tracer_provider.add_span_processor(SimpleSpanProcessor(InMemorySpanExporter()))
levo_logger = LevoLogger(
config=otel_config, callback_name="levo", tracer_provider=tracer_provider
)
# Run health check
result = await levo_logger.async_health_check()
self.assertEqual(result["status"], "healthy")
self.assertIn("message", result)
@patch.dict("os.environ", {}, clear=True)
def test_levo_logger_health_check_unhealthy(self):
"""Test health check returns unhealthy status when required vars are missing."""

View file

@ -129,22 +129,6 @@ def test_subclass_missing_any_abstract_member_cannot_instantiate(missing_member)
Incomplete()
def test_concrete_instance_methods_run():
"""Sanity: the trivial overrides actually execute through the base contract."""
instance = _ConcreteBatchesConfig()
assert instance.custom_llm_provider == LlmProviders.OPENAI
assert instance.validate_environment(
headers={"x": "1"},
model="m",
messages=[],
optional_params={},
litellm_params={},
) == {"x": "1"}
assert instance.transform_retrieve_batch_request(
batch_id="b-1", optional_params={}, litellm_params={}
) == {"batch_id": "b-1"}
# =========================================================================== #
# get_config()
# =========================================================================== #

View file

@ -1,5 +1,8 @@
import json
import pytest
import litellm
from litellm.llms.bedrock.chat.invoke_transformations.amazon_nova_transformation import (
AmazonInvokeNovaConfig,
)
@ -13,6 +16,25 @@ TOOL_CALL = {"id": "call_1", "type": "function", "function": {"name": "f", "argu
PNG_DATA_URL = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNkYPhfDwAChwGA60e6kgAAAABJRU5ErkJggg=="
@pytest.fixture
def local_model_cost_map(monkeypatch):
"""Force the bundled in-repo cost map so capability and pricing assertions do not
depend on the network-fetched ``main`` copy, which lags this branch until merge.
``get_model_info`` is lru_cached, so swapping ``model_cost`` is not enough on its
own; clear on the way in and out so entries warmed against either map never leak
across tests."""
original_model_cost = litellm.model_cost
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
litellm.model_cost = litellm.get_model_cost_map(url="")
litellm.get_model_info.cache_clear()
try:
yield
finally:
litellm.model_cost = original_model_cost
litellm.get_model_info.cache_clear()
def _transform_request(messages, optional_params, litellm_params=None):
return AmazonInvokeNovaConfig().transform_request(
model=MODEL,

View file

@ -1,6 +1,8 @@
import asyncio
import base64
import json
import uuid
from types import SimpleNamespace
from typing import Final
from unittest.mock import patch
@ -17,6 +19,77 @@ from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transfor
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
ONE_PIXEL_PNG = base64.b64decode(
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNkYPhfDwAChwGA60e6kgAAAABJRU5ErkJggg=="
)
@pytest.fixture
def async_only_image_fetch(monkeypatch):
from litellm.litellm_core_utils.prompt_templates import factory, image_handling
from litellm.llms.gemini.chat import transformation as gemini_chat_transformation
fetch = SimpleNamespace(
fetched=[],
base64_png=base64.b64encode(ONE_PIXEL_PNG).decode(),
data_url="data:image/png;base64," + base64.b64encode(ONE_PIXEL_PNG).decode(),
)
def forbid_sync_fetch(client, url, **kwargs):
raise litellm.ImageFetchError(f"sync image fetch ran on the event loop: {url}")
async def serve_png(client, url, **kwargs):
fetch.fetched.append(url)
return httpx.Response(
200,
content=ONE_PIXEL_PNG,
headers={"content-type": "image/png"},
request=httpx.Request("GET", url),
)
def forbid_sync_convert(url, *args, **kwargs):
if url.startswith(("http://", "https://")):
raise litellm.ImageFetchError(f"sync convert_url_to_base64 ran on the request path: {url}")
return url
monkeypatch.setattr(image_handling, "safe_get", forbid_sync_fetch)
monkeypatch.setattr(image_handling, "async_safe_get", serve_png)
for module in (image_handling, factory, gemini_chat_transformation):
monkeypatch.setattr(module, "convert_url_to_base64", forbid_sync_convert)
return fetch
@pytest.fixture
def local_model_cost_map(monkeypatch):
"""Force the bundled in-repo cost map so capability and pricing assertions do not
depend on the network-fetched ``main`` copy, which lags this branch until merge.
``get_model_info`` is lru_cached, so swapping ``model_cost`` is not enough on its
own; clear on the way in and out so entries warmed against either map never leak
across tests."""
original_model_cost = litellm.model_cost
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
litellm.model_cost = litellm.get_model_cost_map(url="")
litellm.get_model_info.cache_clear()
try:
yield
finally:
litellm.model_cost = original_model_cost
litellm.get_model_info.cache_clear()
@pytest.fixture
def local_beta_headers_config(monkeypatch):
"""Pin the bundled ``anthropic_beta_headers_config.json`` so beta header assertions
do not depend on the network-fetched copy or on what earlier tests left cached."""
from litellm.anthropic_beta_headers_manager import reload_beta_headers_config
monkeypatch.setenv("LITELLM_LOCAL_ANTHROPIC_BETA_HEADERS", "True")
reload_beta_headers_config()
yield
reload_beta_headers_config()
def test_get_supported_params_thinking():
config = AmazonAnthropicClaudeConfig()
params = config.get_supported_openai_params(

View file

@ -1,12 +1,55 @@
import base64
import json
import uuid
from types import SimpleNamespace
import httpx
import pytest
import litellm
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
ONE_PIXEL_PNG = base64.b64decode(
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNkYPhfDwAChwGA60e6kgAAAABJRU5ErkJggg=="
)
@pytest.fixture
def async_only_image_fetch(monkeypatch):
from litellm.litellm_core_utils.prompt_templates import factory, image_handling
from litellm.llms.gemini.chat import transformation as gemini_chat_transformation
fetch = SimpleNamespace(
fetched=[],
base64_png=base64.b64encode(ONE_PIXEL_PNG).decode(),
data_url="data:image/png;base64," + base64.b64encode(ONE_PIXEL_PNG).decode(),
)
def forbid_sync_fetch(client, url, **kwargs):
raise litellm.ImageFetchError(f"sync image fetch ran on the event loop: {url}")
async def serve_png(client, url, **kwargs):
fetch.fetched.append(url)
return httpx.Response(
200,
content=ONE_PIXEL_PNG,
headers={"content-type": "image/png"},
request=httpx.Request("GET", url),
)
def forbid_sync_convert(url, *args, **kwargs):
if url.startswith(("http://", "https://")):
raise litellm.ImageFetchError(f"sync convert_url_to_base64 ran on the request path: {url}")
return url
monkeypatch.setattr(image_handling, "safe_get", forbid_sync_fetch)
monkeypatch.setattr(image_handling, "async_safe_get", serve_png)
for module in (image_handling, factory, gemini_chat_transformation):
monkeypatch.setattr(module, "convert_url_to_base64", forbid_sync_convert)
return fetch
async def test_bedrock_mantle_claude_async_completion_inlines_remote_images_off_the_event_loop(async_only_image_fetch):
image_url = f"http://img.example/{uuid.uuid4()}.png"
captured = {}

View file

@ -5,6 +5,7 @@ Source: litellm/llms/chatgpt/responses/transformation.py
"""
import json
from collections.abc import Generator
from unittest.mock import MagicMock, patch
import httpx
@ -19,6 +20,15 @@ from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
@pytest.fixture
def local_model_cost_map(monkeypatch: pytest.MonkeyPatch) -> Generator[None, None, None]:
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
litellm.get_model_info.cache_clear()
yield
litellm.get_model_info.cache_clear()
class TestChatGPTResponsesAPITransformation:
@pytest.mark.parametrize(
"model_name",

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