diff --git a/.circleci/scripts/run_integration.sh b/.circleci/scripts/run_integration.sh index ba24e66ba1c..f3166e5db39 100644 --- a/.circleci/scripts/run_integration.sh +++ b/.circleci/scripts/run_integration.sh @@ -173,7 +173,7 @@ start_proxy() { AWS_EC2_METADATA_DISABLED=true DO_NOT_TRACK=1 COVERAGE_FILE="$coverage_data" \ "${proxy_command[@]}" --config tests/integration/proxy_config.yaml \ --host 127.0.0.1 --port "$port" --num_workers 1 --telemetry False \ - --use_prisma_db_push --enforce_prisma_migration_check \ + --use_prisma_db_push \ > "$results/$log_name" 2>&1 & launched_pid=$! } diff --git a/.circleci/scripts/unit_selection.sh b/.circleci/scripts/unit_selection.sh index bfaa27c3ef3..542984dd2e0 100755 --- a/.circleci/scripts/unit_selection.sh +++ b/.circleci/scripts/unit_selection.sh @@ -77,6 +77,7 @@ legacy_paths() { echo tests/unit/embeddings echo tests/unit/endpoints echo tests/unit/files + echo tests/unit/harness echo tests/unit/images echo tests/unit/interactions echo tests/unit/messages @@ -145,6 +146,7 @@ legacy_paths() { echo tests/unit/proxy/test_proxy_token_counter.py echo tests/unit/proxy/test_server_root_path.py ;; proxy-db-proxy-server-core) + echo tests/unit/proxy/test__lazy_features.py echo tests/unit/proxy/test_aproxy_startup.py echo tests/unit/proxy/test_proxy_server.py ;; proxy-db-proxy-utils) echo tests/unit/proxy/test_proxy_utils.py ;; diff --git a/.github/ci-coverage-allowlist.yml b/.github/ci-coverage-allowlist.yml index 445a8519436..eea25e8e285 100644 --- a/.github/ci-coverage-allowlist.yml +++ b/.github/ci-coverage-allowlist.yml @@ -4,6 +4,14 @@ description: >- by a job nor listed here, so every entry below is a decision on the record. test_paths: + - reason: >- + litellm.agent() end-to-end suite. It drives the real claude, codex and opencode CLIs and + deepagents against a live LiteLLM AI Gateway, so it needs those binaries on PATH plus + LITELLM_PROXY_API_BASE / LITELLM_PROXY_API_KEY, and skips without them. Run manually + before changing litellm/harness; the mocked coverage runs in tests/unit/harness and + tests/unit/llms/*/harness + paths: + - tests/harness_e2e - reason: >- The Rust/Python parity harness is run manually through its local CLI. Recorded replay, fixture generation, and harness checks are intentionally outside pull request CI diff --git a/README.md b/README.md index 98c5343daee..4004e6474ee 100644 --- a/README.md +++ b/README.md @@ -268,6 +268,31 @@ For MCP OAuth, an upstream may advertise dynamic client registration but refuse +
+Agents - Run Claude Code, Codex, OpenCode or Deep Agents on any model (Python SDK) + +### Python SDK - Agents + +```python +import litellm +from litellm import Harness, sandbox + +result = litellm.agent( + Harness.CLAUDE_CODE, # or Harness.CODEX, Harness.OPENCODE, Harness.DEEPAGENTS + "Find why tests/test_router.py is flaky and fix it.", + sandbox=sandbox.local("./repo"), + model="litellm_proxy/claude-sonnet-4-5", # a model group on your AI Gateway +) + +print(result.text, result.cost, [f.path for f in result.files]) +``` + +Set `LITELLM_PROXY_API_BASE` and `LITELLM_PROXY_API_KEY` and every model call the agent makes goes through your AI Gateway, tagged `harness,claude_code`. Drop the `litellm_proxy/` prefix to call a provider directly. Install `starlette uvicorn` plus the agent's CLI (`claude`, `codex` or `opencode`), or `deepagents langchain-litellm` for Deep Agents. + +[**Docs: Agent Harnesses**](https://docs.litellm.ai/docs/harness) + +
+ ### Supported Providers ([Website Supported Models](https://models.litellm.ai/) | [Docs](https://docs.litellm.ai/docs/providers)) | Provider | `/chat/completions` | `/messages` | `/responses` | `/embeddings` | `/image/generations` | `/audio/transcriptions` | `/audio/speech` | `/moderations` | `/batches` | `/rerank` | diff --git a/backend/routes/allowlist.py b/backend/routes/allowlist.py index 80ca0ef22bb..d7f3e615c67 100644 --- a/backend/routes/allowlist.py +++ b/backend/routes/allowlist.py @@ -60,6 +60,7 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = ( # Tools / agents (registry & policy admin) "/v1/tool/", "/v1/agents", + "/agent/daily/activity/", # Guardrails admin "/v2/guardrails/", # MCP server admin + BYOK OAuth flow (UI-initiated) + dynamic per-server endpoints diff --git a/deploy/lens/compose.yaml b/deploy/lens/compose.yaml index d41cb8eb203..4d1224fd41e 100644 --- a/deploy/lens/compose.yaml +++ b/deploy/lens/compose.yaml @@ -1,6 +1,6 @@ services: lens-worker: - image: ${LENS_WORKER_IMAGE:-ghcr.io/berriai/litellm-lens-worker@sha256:a8e8731d954916594eea462969946b9292fb771681ff515a9fd296b53f856c77} + image: ${LENS_WORKER_IMAGE:-ghcr.io/berriai/litellm-lens-worker@sha256:67eba741c1b97c749975c5c38e2370a603e1105babc908d613c1b79d7b995393} environment: LITELLM_URL: ${LITELLM_URL:?Set the URL reachable from this container} LENS_WORKER_TOKEN: ${LENS_WORKER_TOKEN:?Create a worker credential in the Lens UI} diff --git a/docker-compose.hardened.yml b/docker-compose.hardened.yml index 31d0c2e9ef2..84a23faa054 100644 --- a/docker-compose.hardened.yml +++ b/docker-compose.hardened.yml @@ -6,8 +6,6 @@ services: context: . dockerfile: docker/Dockerfile.non_root target: runtime - args: - PROXY_EXTRAS_SOURCE: "local" depends_on: - squid user: "101:101" diff --git a/docker/Dockerfile.non_root b/docker/Dockerfile.non_root index eca12855afa..ca526e06834 100644 --- a/docker/Dockerfile.non_root +++ b/docker/Dockerfile.non_root @@ -3,7 +3,6 @@ # Base images ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d -ARG PROXY_EXTRAS_SOURCE=published ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a # Pinned by digest like the other base images; bump explicitly on Node upgrades. ARG UI_BUILD_IMAGE=node:24.19-alpine3.24@sha256:d32cdf619f63fe0471182d08996dd516c6275bb5fd31ae06e55a570bd9e1ad43 @@ -44,7 +43,6 @@ COPY ui/litellm-dashboard/ ./ RUN npm run build FROM $LITELLM_BUILD_IMAGE AS builder -ARG PROXY_EXTRAS_SOURCE WORKDIR /app USER root @@ -107,26 +105,14 @@ RUN mkdir -p /var/lib/litellm/ui /var/lib/litellm/assets && \ touch /var/lib/litellm/ui/.litellm_ui_ready RUN --mount=type=cache,target=/app/.cache/uv,id=litellm-uv-cache \ - if [ "$PROXY_EXTRAS_SOURCE" = "published" ]; then \ - uv sync --frozen --no-default-groups --no-editable \ - --extra proxy \ - --extra proxy-runtime \ - --extra extra_proxy \ - --extra semantic-router \ - --extra saml \ - --extra bedrock-realtime \ - --python python3.13 \ - --no-sources-package litellm-proxy-extras; \ - else \ - uv sync --frozen --no-default-groups --no-editable \ - --extra proxy \ - --extra proxy-runtime \ - --extra extra_proxy \ - --extra semantic-router \ - --extra saml \ - --extra bedrock-realtime \ - --python python3.13; \ - fi + uv sync --frozen --no-default-groups --no-editable \ + --extra proxy \ + --extra proxy-runtime \ + --extra extra_proxy \ + --extra semantic-router \ + --extra saml \ + --extra bedrock-realtime \ + --python python3.13 RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/opt/prisma/binaries \ npm_config_cache=/root/.npm \ @@ -136,7 +122,6 @@ RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh && \ sed -i 's/\r$//' docker/prod_entrypoint.sh && chmod +x docker/prod_entrypoint.sh FROM $LITELLM_RUNTIME_IMAGE AS runtime -ARG PROXY_EXTRAS_SOURCE WORKDIR /app USER root diff --git a/enterprise/pyproject.toml b/enterprise/pyproject.toml index 74cedb9d84d..43aa5a1f728 100644 --- a/enterprise/pyproject.toml +++ b/enterprise/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "litellm-enterprise" -version = "0.1.72" +version = "0.1.73" description = "Package for LiteLLM Enterprise features" readme = "README.md" requires-python = ">=3.9" @@ -26,7 +26,7 @@ required-version = ">=0.10.9" module-root = "" [tool.commitizen] -version = "0.1.72" +version = "0.1.73" version_files = [ "pyproject.toml:^version", "../pyproject.toml:litellm-enterprise==", diff --git a/litellm-proxy-extras/litellm_proxy_extras/migration_lock.py b/litellm-proxy-extras/litellm_proxy_extras/migration_lock.py index e4ccbe585a9..bea5e36fd18 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/migration_lock.py +++ b/litellm-proxy-extras/litellm_proxy_extras/migration_lock.py @@ -87,3 +87,21 @@ def migration_lock(database_url: str) -> Generator[MigrationCoordinator, None, N 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}." ) + + +@contextmanager +def held_migration_lock(connection: "psycopg.Connection[tuple[object, ...]]") -> Generator[bool, None, None]: + """A session-level, non-blocking hold of the migration coordinator lock on an autocommit + connection, for DDL that cannot run inside a transaction (`CREATE INDEX CONCURRENTLY`). + Yields whether the lock was acquired; a v2 resolver or another migration job's index build + holding it yields False. Released on exit.""" + from psycopg.rows import class_row + + with connection.cursor(row_factory=class_row(_LockResult)) as cursor: + row: Final = cursor.execute("SELECT pg_try_advisory_lock(%s) AS acquired", (MIGRATION_LOCK_KEY,)).fetchone() + acquired: Final = row is not None and row.acquired + try: + yield acquired + finally: + if acquired: + connection.execute("SELECT pg_advisory_unlock(%s)", (MIGRATION_LOCK_KEY,)) diff --git a/litellm-proxy-extras/litellm_proxy_extras/migration_recovery.py b/litellm-proxy-extras/litellm_proxy_extras/migration_recovery.py index 9202317c776..5a55b35b255 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/migration_recovery.py +++ b/litellm-proxy-extras/litellm_proxy_extras/migration_recovery.py @@ -1,4 +1,5 @@ import hashlib +import re import subprocess from collections.abc import Mapping from dataclasses import dataclass @@ -156,3 +157,48 @@ def baseline_current_schema( "review any feature-specific backfill requirements.", len(migrations), ) + + +_LINE_COMMENT_RE: Final = re.compile(r"--[^\n]*") +_BLOCK_COMMENT_RE: Final = re.compile(r"/\*.*?\*/", re.DOTALL) +_NO_OP_STATEMENT_RE: Final = re.compile(r"^\s*SELECT\s+1\s*$", re.IGNORECASE) + + +def is_inert_migration(script: str) -> bool: + """Whether a migration file changes nothing: only comments and `SELECT 1`, so + applying it can neither repeat nor skip a database change.""" + stripped: Final = _LINE_COMMENT_RE.sub("", _BLOCK_COMMENT_RE.sub("", script)) + return all(not part.strip() or _NO_OP_STATEMENT_RE.match(part) for part in stripped.split(";")) + + +def roll_back_failed_inert_migration(coordinator: MigrationCoordinator, schema: str, migration: Path) -> bool: + """Roll back the failed ledger row of a migration whose file in this build is inert, + so `migrate deploy` applies the inert file on its next pass. The row records an + earlier build's attempt at SQL this build no longer ships (an index now built by the + migration job), so no database change can be repeated or skipped by replaying + the empty file. The caller commits this checkpoint before the next Prisma command. + """ + from psycopg import sql + + if not is_inert_migration(migration.read_text(encoding="utf-8")): + return False + coordinator.acquire_prisma_lock() + records: Final = _migration_records(coordinator.connection, schema, migration) + unfinished: Final = tuple(record for record in records if not record.finished) + if len(unfinished) != 1: + return False + result: Final = coordinator.connection.execute( + sql.SQL( + "UPDATE {} SET rolled_back_at = current_timestamp " + "WHERE id = %s AND finished_at IS NULL AND rolled_back_at IS NULL" + ).format(sql.Identifier(schema, "_prisma_migrations")), + (unfinished[0].id,), + ) + if result.rowcount != 1: + raise RuntimeError("Could not roll back the failed inert migration history row; rerun the database setup.") + logger.info( + "Rolled back the failed history row of %s: this build ships it as an inert migration, " + "its index is built by the migration job", + migration.parent.name, + ) + return True diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260823000000_add_spend_logs_api_key_starttime_index/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260823000000_add_spend_logs_api_key_starttime_index/migration.sql index 9a061aaed43..a2bec81ca00 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260823000000_add_spend_logs_api_key_starttime_index/migration.sql +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260823000000_add_spend_logs_api_key_starttime_index/migration.sql @@ -1,2 +1,6 @@ --- CreateIndex -CREATE INDEX IF NOT EXISTS "LiteLLM_SpendLogs_api_key_startTime_idx" ON "LiteLLM_SpendLogs"("api_key", "startTime"); +-- The (api_key, startTime) index on LiteLLM_SpendLogs is built after migrate deploy, +-- through litellm_proxy_extras/request_log_indexes.py: concurrently on a plain table and +-- per partition on a partitioned one. The migration job builds it; a serving proxy that +-- ran the migrations itself builds it in the background once it serves. A migration +-- cannot do either without blocking spend-log writes or failing on a partitioned table. +SELECT 1; diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260831120001_spend_logs_litellm_call_id_index/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260831120001_spend_logs_litellm_call_id_index/migration.sql index 62ad5c42ba7..7eba7fc9b97 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260831120001_spend_logs_litellm_call_id_index/migration.sql +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260831120001_spend_logs_litellm_call_id_index/migration.sql @@ -1,12 +1,6 @@ --- CreateIndex (CONCURRENTLY) --- --- Disclaimer: --- - CREATE INDEX CONCURRENTLY cannot run inside a transaction. This migration must stay a --- single statement so Prisma Migrate on PostgreSQL can apply it outside a transaction. --- - Builds are slower and use more I/O than a blocking CREATE INDEX; if the build is --- interrupted, Postgres may leave an INVALID index that must be dropped and recreated. --- - Do not edit this file after it has been applied to any database: Prisma checksums --- migrations; add a new migration instead. --- - Requires PostgreSQL that supports CONCURRENTLY with IF NOT EXISTS (use a new migration --- without IF NOT EXISTS if you must support older versions). -CREATE INDEX CONCURRENTLY IF NOT EXISTS "LiteLLM_SpendLogs_litellm_call_id_idx" ON "LiteLLM_SpendLogs"("litellm_call_id"); +-- The litellm_call_id index on LiteLLM_SpendLogs is built after migrate deploy, through +-- litellm_proxy_extras/request_log_indexes.py: concurrently on a plain table and per +-- partition on a partitioned one. The migration job builds it; a serving proxy that ran +-- the migrations itself builds it in the background once it serves. Postgres refuses +-- CREATE INDEX CONCURRENTLY on a partitioned parent, so this migration no longer runs it. +SELECT 1; diff --git a/litellm-proxy-extras/litellm_proxy_extras/request_log_indexes.py b/litellm-proxy-extras/litellm_proxy_extras/request_log_indexes.py new file mode 100644 index 00000000000..6c31e8364a9 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/request_log_indexes.py @@ -0,0 +1,463 @@ +"""The request-log indexes built after `prisma migrate deploy` instead of by a migration: +by the migration job, or by a serving proxy that ran the migrations itself (in the +background, once it serves). + +A migration cannot build them: a plain `CREATE INDEX` blocks spend-log inserts for the +whole build, and `CREATE INDEX CONCURRENTLY` is refused on a partitioned parent +(db_scripts/partition_spend_logs.sql). `REQUEST_LOG_INDEXES` is the one list to extend; +names match what Prisma derives from the `@@index` declarations in schema.prisma, so an +index a database already has is recognized and never rebuilt. +""" + +import hashlib +import random +import re +import time +from collections.abc import Callable +from dataclasses import dataclass +from typing import TYPE_CHECKING, Final + +from litellm_proxy_extras._logging import logger +from litellm_proxy_extras.migration_lock import held_migration_lock + +if TYPE_CHECKING: + import psycopg + from psycopg import sql + + +@dataclass(frozen=True, slots=True) +class RequestLogIndex: + """One index the migration job owns: the table, the exact Prisma index name and the + column list as it would be written after `ON `.""" + + table: str + name: str + definition: str + + @property + def columns(self) -> tuple[str, ...]: + return tuple(re.findall(r'"([^"]+)"', self.definition)) + + def partition_index_name(self, partition: str) -> str: + """The child index name for one partition, built the way Postgres names the + children of a partitioned index, and kept within the 63 byte identifier limit.""" + name: Final = f"{partition}_{self.name.removeprefix(f'{self.table}_')}" + if len(name.encode()) <= _IDENTIFIER_MAX_BYTES: + return name + digest: Final = hashlib.sha256(name.encode()).hexdigest()[:_DIGEST_LENGTH] + budget: Final = _IDENTIFIER_MAX_BYTES - _DIGEST_LENGTH - 1 + kept: Final = next(name[:length] for length in range(len(name), 0, -1) if len(name[:length].encode()) <= budget) + return f"{kept}_{digest}" + + +REQUEST_LOG_INDEXES: Final = ( + RequestLogIndex("LiteLLM_SpendLogs", "LiteLLM_SpendLogs_api_key_startTime_idx", '("api_key", "startTime")'), + RequestLogIndex("LiteLLM_SpendLogs", "LiteLLM_SpendLogs_litellm_call_id_idx", '("litellm_call_id")'), +) + +_IDENTIFIER_MAX_BYTES: Final = 63 +_DDL_LOCK_TIMEOUT: Final = "200ms" +_DDL_LOCK_ATTEMPTS: Final = 10 +_DDL_RETRY_BASE_SECONDS: Final = 0.25 +_DDL_RETRY_MAX_SECONDS: Final = 8.0 +_LOCK_HANDOVER_SECONDS: Final = 2.0 +_DIGEST_LENGTH: Final = 8 +_CREATE_INDEX_STATEMENT: Final = re.compile( + r'^\s*CREATE\s+(?:UNIQUE\s+)?INDEX\s+(?:CONCURRENTLY\s+)?(?:IF\s+NOT\s+EXISTS\s+)?"(?P[^"]+)"\s+ON\b', + re.IGNORECASE, +) +_TABLE_KIND_SQL: Final = "SELECT c.relkind = 'p' AS partitioned FROM pg_class c WHERE c.oid = to_regclass(%s)" +_CHILDREN_WITHOUT_THE_INDEX_SQL: Final = ( + "SELECT child.relname AS name, n.nspname AS schema, child.relkind = 'p' AS partitioned " + "FROM pg_inherits i JOIN pg_class child ON child.oid = i.inhrelid " + "JOIN pg_namespace n ON n.oid = child.relnamespace " + "WHERE i.inhparent = to_regclass(%s) AND NOT EXISTS (" + "SELECT 1 FROM pg_inherits attached JOIN pg_index x ON x.indexrelid = attached.inhrelid " + "WHERE attached.inhparent = to_regclass(%s) AND x.indrelid = child.oid) " + "ORDER BY child.relname" +) +_EQUIVALENT_INDEXES_SQL: Final = ( + "SELECT i.relname AS name, x.indisvalid AS valid " + "FROM pg_index x JOIN pg_class i ON i.oid = x.indexrelid JOIN pg_am am ON am.oid = i.relam " + "WHERE x.indrelid = to_regclass(%s) AND i.relname <> %s AND am.amname = 'btree' AND NOT x.indisunique " + "AND x.indexprs IS NULL AND x.indpred IS NULL AND x.indnkeyatts = x.indnatts " + "AND NOT EXISTS (SELECT 1 FROM unnest(x.indoption::int2[]) o WHERE o <> 0) " + "AND NOT EXISTS (SELECT 1 FROM unnest(x.indclass::oid[]) c JOIN pg_opclass oc ON oc.oid = c WHERE NOT oc.opcdefault) " + "AND NOT EXISTS (SELECT 1 FROM unnest(x.indcollation::oid[]) WITH ORDINALITY c(coll, ord) " + "JOIN unnest(x.indkey::int2[]) WITH ORDINALITY k(attnum, ord) ON k.ord = c.ord " + "JOIN pg_attribute a ON a.attrelid = x.indrelid AND a.attnum = k.attnum " + "WHERE c.coll <> 0 AND c.coll <> a.attcollation) " + "AND (SELECT array_agg(a.attname::text ORDER BY k.ord) FROM unnest(x.indkey::int2[]) WITH ORDINALITY k(attnum, ord) " + "JOIN pg_attribute a ON a.attrelid = x.indrelid AND a.attnum = k.attnum) = %s::text[] " + "AND NOT EXISTS (SELECT 1 FROM pg_inherits WHERE inhrelid = x.indexrelid) " + "ORDER BY x.indisvalid DESC, i.relname" +) +_INDEX_STATE_SQL: Final = ( + 'SELECT x.indisvalid AS valid, t.relname AS "table" ' + "FROM pg_index x JOIN pg_class t ON t.oid = x.indrelid WHERE x.indexrelid = to_regclass(%s)" +) + + +@dataclass(frozen=True, slots=True) +class _Relation: + name: str + schema: str + partitioned: bool + + +@dataclass(frozen=True, slots=True) +class _IndexState: + valid: bool + table: str + + +@dataclass(frozen=True, slots=True) +class _EquivalentIndex: + name: str + valid: bool + + +@dataclass(frozen=True, slots=True) +class _TableKind: + partitioned: bool + + +def filter_request_log_index_diff(diff_sql: str, indexes: tuple[RequestLogIndex, ...] = REQUEST_LOG_INDEXES) -> str: + """The `prisma migrate diff` script without the statements that create a migration-job-owned + index, which the schema declares and the migrations deliberately do not build.""" + names: Final = frozenset(index.name for index in indexes) + statements: Final = diff_sql.split(";") + kept: Final = tuple(statement for statement in statements if not _creates_one_of(statement, names)) + return ";".join(kept) if any(part.strip() for part in kept) else "" + + +def _creates_one_of(statement: str, names: frozenset[str]) -> bool: + match: Final = _CREATE_INDEX_STATEMENT.match(_without_comments(statement)) + return match is not None and match["index"] in names + + +def _without_comments(statement: str) -> str: + return "\n".join(line for line in statement.splitlines() if not line.lstrip().startswith("--")) + + +def _connect(database_url: str) -> "psycopg.Connection[tuple[object, ...]]": + import psycopg + + return psycopg.connect(database_url, connect_timeout=10, autocommit=True) + + +def ensure_request_log_indexes( + database_url: str, + schema: str, + indexes: tuple[RequestLogIndex, ...] = REQUEST_LOG_INDEXES, + connect: "Callable[[str], psycopg.Connection[tuple[object, ...]]]" = _connect, +) -> bool: + """Build every listed index that is missing or invalid. Each build step runs under + the migration coordinator lock, held per statement so a resolver booting on another + replica gets in between partitions rather than waiting for the whole table. Any + failure is logged and left for the next index build; the result says whether + every index ended up valid. Never raises.""" + import psycopg + + try: + with connect(database_url) as connection: + connection.execute("SET statement_timeout = 0") + results: Final = tuple(_ensure_index(connection, schema, index) for index in indexes) + except psycopg.Error as exc: + logger.warning("Could not build the request-log indexes, leaving them for the next index build: %s", exc) + return False + if not all(results): + logger.warning("Some request-log indexes are not in place yet, leaving them for the next index build") + return False + logger.info("Request-log indexes are all in place") + return True + + +def _under_migration_lock(connection: "psycopg.Connection[tuple[object, ...]]", step: Callable[[], bool]) -> bool: + with held_migration_lock(connection) as held: + if not held: + logger.info( + "Another process holds the migration lock, leaving the request-log indexes to the next index build" + ) + return False + return step() + + +def _with_bounded_lock( + connection: "psycopg.Connection[tuple[object, ...]]", step: Callable[[], bool], what: str +) -> bool: + """Run `step` under the migration lock with a short lock_timeout, so a DDL statement that has to wait for open + transactions holds new writes back for at most that long; retry with capped exponential backoff, holding the + migration lock per attempt only and releasing it while sleeping. False when another process holds the migration + lock or every attempt timed out.""" + import psycopg + from psycopg import sql + + for attempt in range(_DDL_LOCK_ATTEMPTS): + if attempt: + time.sleep(min(_DDL_RETRY_MAX_SECONDS, _DDL_RETRY_BASE_SECONDS * 2.0**attempt) * random.uniform(0.5, 1.0)) + connection.execute(sql.SQL("SET lock_timeout = {}").format(sql.Literal(_DDL_LOCK_TIMEOUT))) + try: + return _under_migration_lock(connection, step) + except psycopg.errors.LockNotAvailable: + logger.info("Waiting for open transactions before %s", what) + finally: + connection.execute("SET lock_timeout = 0") + logger.warning( + "Could not get the lock for %s without holding writes back, leaving it for the next index build", what + ) + return False + + +def _ensure_index(connection: "psycopg.Connection[tuple[object, ...]]", schema: str, index: RequestLogIndex) -> bool: + from psycopg.rows import class_row + + with connection.cursor(row_factory=class_row(_TableKind)) as cursor: + table: Final = cursor.execute(_TABLE_KIND_SQL, (_regclass_name(connection, schema, index.table),)).fetchone() + if table is None: + logger.info("Table %s does not exist yet, skipping index %s", index.table, index.name) + return True + if table.partitioned: + return build_index_on_partitioned_table(connection, schema, index) + return _build_leaf_index(connection, schema, index.table, index.name, index) + + +def _regclass_name(connection: "psycopg.Connection[tuple[object, ...]]", schema: str, name: str) -> str: + from psycopg import sql + + return sql.Identifier(schema, name).as_string(connection) + + +def _create_index_statement( + connection: "psycopg.Connection[tuple[object, ...]]", prefix: "sql.Composed", definition: str +) -> bytes: + return (prefix.as_string(connection) + definition).encode() + + +def _index_state(connection: "psycopg.Connection[tuple[object, ...]]", schema: str, index: str) -> "_IndexState | None": + from psycopg.rows import class_row + + with connection.cursor(row_factory=class_row(_IndexState)) as cursor: + return cursor.execute(_INDEX_STATE_SQL, (_regclass_name(connection, schema, index),)).fetchone() + + +def _equivalent_indexes( + connection: "psycopg.Connection[tuple[object, ...]]", + schema: str, + table: str, + name: str, + index: RequestLogIndex, +) -> tuple[_EquivalentIndex, ...]: + """The indexes on `table` other than `name` with the same definition: default btree + over the same columns in the same order, no expression, predicate, DESC or custom + opclass or collation, and not attached under a partitioned index. Valid ones first.""" + from psycopg.rows import class_row + + with connection.cursor(row_factory=class_row(_EquivalentIndex)) as cursor: + return tuple( + cursor.execute( + _EQUIVALENT_INDEXES_SQL, (_regclass_name(connection, schema, table), name, list(index.columns)) + ).fetchall() + ) + + +def _adopt_equivalent_index( + connection: "psycopg.Connection[tuple[object, ...]]", + schema: str, + table: str, + name: str, + index: RequestLogIndex, +) -> bool: + """Rename a valid index of the same definition under another name (an operator's + hand-built copy, say) to the name this code expects, instead of building a second + one. RENAME on an index is a catalog change that lets writes through.""" + from psycopg import sql + + equivalent: Final = next( + (found for found in _equivalent_indexes(connection, schema, table, name, index) if found.valid), None + ) + if equivalent is None: + return False + logger.info( + "Renaming the equivalent index %s on %s to %s instead of building a second one", equivalent.name, table, name + ) + connection.execute( + sql.SQL("ALTER INDEX {} RENAME TO {}").format(sql.Identifier(schema, equivalent.name), sql.Identifier(name)) + ) + return True + + +def _report_second_copies( + connection: "psycopg.Connection[tuple[object, ...]]", + schema: str, + table: str, + name: str, + index: RequestLogIndex, + concurrently: bool, +) -> None: + """Log every other index of the same definition with the statement that removes it. + Dropping is the operator's call: a second copy costs writes and disk, never results.""" + from psycopg import sql + + drop: Final = "DROP INDEX CONCURRENTLY" if concurrently else "DROP INDEX" + for copy in _equivalent_indexes(connection, schema, table, name, index): + logger.warning( + "Index %s on %s is a second copy of %s and only costs writes and disk; remove it with: %s %s", + copy.name, + table, + name, + drop, + sql.Identifier(schema, copy.name).as_string(connection), + ) + + +def _children_without_the_index( + connection: "psycopg.Connection[tuple[object, ...]]", schema: str, table: str, index: str +) -> tuple[_Relation, ...]: + from psycopg.rows import class_row + + with connection.cursor(row_factory=class_row(_Relation)) as cursor: + return tuple( + cursor.execute( + _CHILDREN_WITHOUT_THE_INDEX_SQL, + (_regclass_name(connection, schema, table), _regclass_name(connection, schema, index)), + ).fetchall() + ) + + +def _build_leaf_index( + connection: "psycopg.Connection[tuple[object, ...]]", + schema: str, + table: str, + name: str, + index: RequestLogIndex, +) -> bool: + """Build one plain table's or partition's index with CONCURRENTLY so writes keep + flowing. The catalog is read under the migration lock, so a replica that saw an + invalid index before the lock finds the valid one another replica just built and + leaves it. An invalid index left by an interrupted build is dropped and rebuilt; a + valid index of the same definition under another name is renamed rather than + duplicated; an index of that name on another table is a collision this code will + not touch.""" + from psycopg import sql + + def build() -> bool: + existing: Final = _index_state(connection, schema, name) + if existing is not None and existing.table != table: + logger.warning( + "Index %s already exists on %s rather than %s, leaving it alone", name, existing.table, table + ) + return False + if existing is not None and existing.valid: + return True + if existing is not None: + logger.info("Dropping the invalid index %s left by an interrupted build on %s", name, table) + connection.execute(sql.SQL("DROP INDEX CONCURRENTLY {}").format(sql.Identifier(schema, name))) + elif _adopt_equivalent_index(connection, schema, table, name, index): + return True + logger.info("Building index %s on %s concurrently", name, table) + prefix: Final = sql.SQL("CREATE INDEX CONCURRENTLY IF NOT EXISTS {} ON {} ").format( + sql.Identifier(name), sql.Identifier(schema, table) + ) + connection.execute(_create_index_statement(connection, prefix, index.definition)) + built: Final = _index_state(connection, schema, name) + return built is not None and built.valid + + current: Final = _index_state(connection, schema, name) + if current is None or not current.valid or current.table != table: + if not _under_migration_lock(connection, build): + return False + time.sleep(_LOCK_HANDOVER_SECONDS) + _report_second_copies(connection, schema, table, name, index, concurrently=True) + return True + + +def build_index_on_partitioned_table( + connection: "psycopg.Connection[tuple[object, ...]]", + schema: str, + index: RequestLogIndex, + table: "str | None" = None, + name: "str | None" = None, +) -> bool: + """Build the index the way Postgres allows on a partitioned parent: a metadata-only + parent index ON ONLY the parent, one CONCURRENTLY build per partition, and ATTACH + PARTITION for each child. Partitions that are themselves partitioned get the same + treatment one level down. Every step checks the catalog before acting, so an + interrupted run resumes where it stopped and a second run finds nothing to do; a + parent or child index of the same definition under another name is renamed and + used rather than duplicated. The connection must be in autocommit mode. True when + the parent index ends up valid.""" + + parent_table: Final = index.table if table is None else table + parent_index: Final = index.name if name is None else name + existing: Final = _index_state(connection, schema, parent_index) + if existing is not None and existing.table != parent_table: + logger.warning( + "Index %s already exists on %s rather than %s, leaving it alone", parent_index, existing.table, parent_table + ) + return False + if existing is None and not _with_bounded_lock( + connection, + lambda: ( + _adopt_equivalent_index(connection, schema, parent_table, parent_index, index) + or _create_parent_index(connection, schema, parent_index, parent_table, index) + ), + f"creating the parent index {parent_index}", + ): + return False + children: Final = _children_without_the_index(connection, schema, parent_table, parent_index) + if not all(_attach_child_index(connection, schema, parent_index, child, index) for child in children): + return False + final: Final = _index_state(connection, schema, parent_index) + if final is None or not final.valid: + return False + _report_second_copies(connection, schema, parent_table, parent_index, index, concurrently=False) + return True + + +def _create_parent_index( + connection: "psycopg.Connection[tuple[object, ...]]", + schema: str, + name: str, + table: str, + index: RequestLogIndex, +) -> bool: + """Create the metadata-only parent index. The caller bounds Postgres's SHARE lock wait on the parent.""" + from psycopg import sql + + prefix: Final = sql.SQL("CREATE INDEX IF NOT EXISTS {} ON ONLY {} ").format( + sql.Identifier(name), sql.Identifier(schema, table) + ) + statement: Final = _create_index_statement(connection, prefix, index.definition) + connection.execute(statement) + return True + + +def _attach_child_index( + connection: "psycopg.Connection[tuple[object, ...]]", + schema: str, + parent_index: str, + child: _Relation, + index: RequestLogIndex, +) -> bool: + from psycopg import sql + + child_index: Final = index.partition_index_name(child.name) + built: Final = ( + build_index_on_partitioned_table(connection, child.schema, index, child.name, child_index) + if child.partitioned + else _build_leaf_index(connection, child.schema, child.name, child_index, index) + ) + if not built: + return False + + def attach() -> bool: + connection.execute( + sql.SQL("ALTER INDEX {} ATTACH PARTITION {}").format( + sql.Identifier(schema, parent_index), sql.Identifier(child.schema, child_index) + ) + ) + logger.info("Attached index %s on partition %s to %s", child_index, child.name, parent_index) + return True + + return _with_bounded_lock(connection, attach, f"attaching {child_index}") diff --git a/litellm-proxy-extras/litellm_proxy_extras/utils.py b/litellm-proxy-extras/litellm_proxy_extras/utils.py index 2f74df63367..3acc19d397d 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/utils.py +++ b/litellm-proxy-extras/litellm_proxy_extras/utils.py @@ -5,6 +5,7 @@ import re import shutil import subprocess import tempfile +import threading import time from collections.abc import Callable from dataclasses import dataclass, replace @@ -13,6 +14,7 @@ from typing import TYPE_CHECKING, Final, Optional from litellm_proxy_extras import prisma_toolchain from litellm_proxy_extras._logging import logger +from litellm_proxy_extras.migration_lock import held_migration_lock from litellm_proxy_extras.prisma_toolchain import ( PRISMA_COMMAND_TIMEOUT_ENV_VAR, PRISMA_MIGRATE_DEPLOY_TIMEOUT_ENV_VAR, @@ -24,6 +26,7 @@ from litellm_proxy_extras.replica_identity import ( REPLICA_IDENTITY_FULL_ENV_VAR, apply_replica_identity_full, ) +from litellm_proxy_extras.request_log_indexes import ensure_request_log_indexes, filter_request_log_index_diff if TYPE_CHECKING: import psycopg @@ -433,6 +436,21 @@ class ProxyExtrasDBManager: return True return False + @staticmethod + def _filter_migration_job_owned_drift(diff_sql: str, partitioned: bool | None = None) -> str: + """The drift script without the indexes the migration job builds (the schema + declares them, the migrations deliberately do not) and, when LiteLLM_SpendLogs + is partitioned, without its primary-key rewrite and partitioning artifacts.""" + without_indexes: Final = filter_request_log_index_diff(diff_sql) + is_partitioned: Final = ProxyExtrasDBManager.spend_logs_is_partitioned() if partitioned is None else partitioned + if not is_partitioned: + return without_indexes + logger.info( + "LiteLLM_SpendLogs is partitioned; removed its primary-key " + "rewrite and partitioning artifacts from the drift script" + ) + return filter_partitioned_spend_logs_diff(without_indexes) + @staticmethod def _resolve_all_migrations( migrations_dir: str, schema_path: str, mark_all_applied: bool = True @@ -513,21 +531,14 @@ class ProxyExtrasDBManager: return logger.info(f"Migration diff created at {diff_sql_path}") - if ProxyExtrasDBManager.spend_logs_is_partitioned(): - filtered_sql = filter_partitioned_spend_logs_diff( - diff_sql_path.read_text() - ) - diff_sql_path.write_text(filtered_sql) - logger.info( - "LiteLLM_SpendLogs is partitioned; removed its primary-key " - "rewrite and partitioning artifacts from the drift script" - ) - if not filtered_sql.strip(): - logger.info("Drift script is empty after filtering; nothing to apply") - if not mark_all_applied: - return - ProxyExtrasDBManager._mark_migrations_applied(migrations_dir) + filtered_sql: Final = ProxyExtrasDBManager._filter_migration_job_owned_drift(diff_sql_path.read_text()) + diff_sql_path.write_text(filtered_sql) + if not filtered_sql.strip(): + logger.info("Drift script is empty after filtering; nothing to apply") + if not mark_all_applied: return + ProxyExtrasDBManager._mark_migrations_applied(migrations_dir) + return # 2. Run prisma db execute to apply the migration applied_ok = False @@ -800,7 +811,7 @@ class ProxyExtrasDBManager: conn.execute(statement) except psycopg.Error as e: logger.warning( - "Could not repair invalid index %s.%s, will retry on the next startup. " + "Could not repair invalid index %s.%s, will retry on the next database setup run. " "If this keeps happening, run `%s` by hand as the index owner. Error: %s", index.schema, index.name, @@ -811,16 +822,21 @@ class ProxyExtrasDBManager: logger.info("%s invalid index %s.%s", action, index.schema, index.name) @staticmethod - def repair_invalid_indexes(lock_timeout: str = "30s") -> bool: + def repair_invalid_indexes( + lock_timeout: str = "30s", + repair: "Callable[[psycopg.Connection[tuple[str, str, str]], _InvalidIndex], None] | None" = None, + ) -> bool: """Rebuild LiteLLM indexes an interrupted CREATE INDEX CONCURRENTLY left INVALID (a migration deadlock between replicas is the usual cause; the retried migration skips them because of IF NOT EXISTS). Never raises: returns True when no invalid index remains, False when the repair was - skipped or failed and will be retried on the next startup. Looks in the + skipped or failed and will be retried on the next database setup run. Looks in the schema DATABASE_URL names, the only URL Prisma migrates through, but connects over DIRECT_URL when set: the session settings, the advisory lock and REINDEX CONCURRENTLY all need one server session, which a - transaction pooler does not give.""" + transaction pooler does not give. Each rebuild holds the migration + coordinator lock on its own, like the migration job's index build, so a resolver + booting on another replica waits for one index at most.""" prisma_url: Final = os.getenv("DATABASE_URL") if not prisma_url: return False @@ -856,20 +872,53 @@ class ProxyExtrasDBManager: if lock_row is None or not lock_row[0]: logger.info("Another replica is already rebuilding the invalid indexes, skipping") return False - for index in ProxyExtrasDBManager._invalid_litellm_indexes(conn, schema): - ProxyExtrasDBManager._repair_index(conn, index) + repair_one: Final = repair or ProxyExtrasDBManager._repair_index + repaired: Final = all( + ProxyExtrasDBManager._repair_under_migration_lock(conn, schema, index, repair_one) + for index in found + ) + if not repaired: + return False remaining: Final = ProxyExtrasDBManager._invalid_litellm_indexes(conn, schema) except psycopg.Error as e: - logger.warning("Could not check for invalid indexes, will retry on the next startup. Error: %s", e) + logger.warning( + "Could not check for invalid indexes, will retry on the next database setup run. Error: %s", e + ) return False return not remaining + @staticmethod + def _repair_under_migration_lock( + conn: "psycopg.Connection[tuple[str, str, str]]", + schema: str, + index: _InvalidIndex, + repair: "Callable[[psycopg.Connection[tuple[str, str, str]], _InvalidIndex], None]", + ) -> bool: + """Rebuild one index under the migration coordinator lock, skipping it when a + migration job finished or dropped it in the meantime. False when another process + holds the lock, so the check waits for the next database setup run.""" + with held_migration_lock(conn) as held: + if not held: + logger.info( + "Another process is building indexes under the migration lock, leaving the " + "invalid index check to the next database setup run" + ) + return False + still_invalid: Final = ProxyExtrasDBManager._invalid_litellm_indexes(conn, schema) + if any(found.schema == index.schema and found.name == index.name for found in still_invalid): + repair(conn, index) + return True + @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 + from litellm_proxy_extras.migration_recovery import ( + baseline_current_schema, + recover_completed_migration, + roll_back_failed_inert_migration, + ) database_url: Final = os.environ.get("DATABASE_URL") if not database_url: @@ -884,7 +933,9 @@ class ProxyExtrasDBManager: if not migration.is_file(): return False with migration_lock(lock_url) as coordinator: - return recover_completed_migration(coordinator, schema, migration) + return recover_completed_migration(coordinator, schema, migration) or roll_back_failed_inert_migration( + coordinator, schema, migration + ) def baseline_existing(migrations_dir: str) -> None: with migration_lock(lock_url) as coordinator: @@ -1177,13 +1228,16 @@ class ProxyExtrasDBManager: ) @staticmethod - def setup_database( - use_migrate: bool = False, use_v2_resolver: bool = False - ) -> bool: + def setup_database(use_migrate: bool = False, use_v2_resolver: bool = False) -> bool: """ Set up the database using either prisma migrate or prisma db push Uses migrations from litellm-proxy-extras package + The request-log indexes in `REQUEST_LOG_INDEXES` are not built here: the + migration job builds them through `run_migration_job`, and a serving proxy that + ran the migrations itself starts them through `start_request_log_index_build` + once it is ready to serve. + Args: use_migrate: Whether to use prisma migrate instead of db push use_v2_resolver: Opt into the v2 migration resolver (safer during @@ -1200,10 +1254,48 @@ class ProxyExtrasDBManager: migrated = ProxyExtrasDBManager._run_migrations( use_migrate=use_migrate, use_v2_resolver=use_v2_resolver ) - if migrated: - ProxyExtrasDBManager.repair_invalid_indexes() - ProxyExtrasDBManager.apply_replica_identity_full_if_requested() - return migrated + if not migrated: + return False + ProxyExtrasDBManager.repair_invalid_indexes() + ProxyExtrasDBManager.apply_replica_identity_full_if_requested() + return True + + @staticmethod + def build_request_log_indexes(build: Callable[[str, str], bool] = ensure_request_log_indexes) -> bool: + """Build the indexes in `REQUEST_LOG_INDEXES` on the writer, in the schema the + migrations target. Idempotent and never raises; False when an index is still + missing or invalid, so the migration job reports it and gets rerun instead of + leaving the table unindexed until the next deploy.""" + database_url: Final = os.environ.get("DATABASE_URL") + if not database_url: + return True + direct_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" + return build(direct_url, schema) + + @staticmethod + def run_migration_job( + use_migrate: bool = False, + use_v2_resolver: bool = False, + setup: Callable[[bool, bool], bool] = setup_database, + build: Callable[[], bool] = build_request_log_indexes, + ) -> bool: + """The migration job's whole run: `setup_database`, then the request-log indexes, + built synchronously so the job exits only once they are in place. False when the + migrations failed or an index could not be built, so the Job is rerun.""" + return setup(use_migrate, use_v2_resolver) and build() + + @staticmethod + def start_request_log_index_build(build: Callable[[], bool] = build_request_log_indexes) -> threading.Thread: + """A serving proxy that ran the migrations itself (schema updates not disabled) + builds the request-log indexes on a daemon thread, so a long build never delays + readiness. A build that could not finish is logged and picked up by the next boot + or the migration job.""" + thread: Final = threading.Thread(target=build, name="litellm-request-log-indexes", daemon=True) + thread.start() + return thread @staticmethod def _run_migrations(use_migrate: bool, use_v2_resolver: bool) -> bool: @@ -1247,15 +1339,16 @@ class ProxyExtrasDBManager: logger.info("✅ Post-migration sanity check completed") return True except subprocess.CalledProcessError as e: - logger.info(f"prisma db error: {e.stderr}, e: {e.stdout}") - if "P3009" in e.stderr: + stderr: Final = str(e.stderr or "") + logger.info(f"prisma db error: {stderr}, e: {e.stdout}") + if "P3009" in stderr: # Extract the failed migration name from the error message migration_match = re.search( - r"`(\d+_.*)` migration", e.stderr + r"`(\d+_.*)` migration", stderr ) if migration_match: failed_migration = migration_match.group(1) - if ProxyExtrasDBManager._is_idempotent_error(e.stderr): + if ProxyExtrasDBManager._is_idempotent_error(stderr): logger.info( f"Migration {failed_migration} failed due to idempotent error (e.g., column already exists), resolving as applied" ) @@ -1311,8 +1404,8 @@ class ProxyExtrasDBManager: f"✅ Migration {failed_migration} marked as rolled back... retrying" ) elif ( - "P3005" in e.stderr - and "database schema is not empty" in e.stderr + "P3005" in stderr + and "database schema is not empty" in stderr ): logger.info( "Database schema is not empty, creating baseline migration. In read-only file system, please set an environment variable `LITELLM_MIGRATION_DIR` to a writable directory to enable migrations. Learn more - https://docs.litellm.ai/docs/proxy/prod#read-only-file-system" @@ -1326,13 +1419,13 @@ class ProxyExtrasDBManager: ) logger.info("✅ All migrations resolved.") return True - elif "P3018" in e.stderr: + elif "P3018" in stderr: # Check if this is a permission error or idempotent error - if ProxyExtrasDBManager._is_permission_error(e.stderr): + if ProxyExtrasDBManager._is_permission_error(stderr): # Permission errors should NOT be marked as applied # Extract migration name for logging migration_match = re.search( - r"Migration name: (\d+_.*)", e.stderr + r"Migration name: (\d+_.*)", stderr ) migration_name = ( migration_match.group(1) @@ -1342,7 +1435,7 @@ class ProxyExtrasDBManager: logger.error( f"❌ Migration {migration_name} failed due to insufficient permissions. " - f"Please check database user privileges. Error: {e.stderr}" + f"Please check database user privileges. Error: {stderr}" ) # Mark as rolled back and exit with error @@ -1365,7 +1458,7 @@ class ProxyExtrasDBManager: f"was NOT applied. Please grant necessary database permissions and retry." ) from e - elif ProxyExtrasDBManager._is_idempotent_error(e.stderr): + elif ProxyExtrasDBManager._is_idempotent_error(stderr): # Idempotent errors mean the migration has effectively been applied logger.info( "Migration failed due to idempotent error (e.g., column already exists), " @@ -1373,7 +1466,7 @@ class ProxyExtrasDBManager: ) # Extract the migration name from the error message migration_match = re.search( - r"Migration name: (\d+_.*)", e.stderr + r"Migration name: (\d+_.*)", stderr ) if migration_match: migration_name = migration_match.group(1) @@ -1422,7 +1515,7 @@ class ProxyExtrasDBManager: logger.warning( f"P3018 error encountered but could not classify " f"as permission or idempotent error. " - f"Error: {e.stderr}" + f"Error: {stderr}" ) raise else: diff --git a/litellm-proxy-extras/pyproject.toml b/litellm-proxy-extras/pyproject.toml index 2e2f3f2ce5a..c6d6060acfa 100644 --- a/litellm-proxy-extras/pyproject.toml +++ b/litellm-proxy-extras/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "litellm-proxy-extras" -version = "0.4.103" +version = "0.4.104" description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package." readme = "README.md" requires-python = ">=3.9" @@ -30,7 +30,7 @@ required-version = ">=0.10.9" module-root = "" [tool.commitizen] -version = "0.4.103" +version = "0.4.104" version_files = [ "pyproject.toml:^version", "../pyproject.toml:litellm-proxy-extras==", diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index ff0eafee47e..40552b19e43 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -97,6 +97,53 @@ version = "1.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "03918c3dbd7701a85c6b9887732e2921175f26c350b4563841d0958c21d57e6d" +[[package]] +name = "askama" +version = "0.16.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6024d73179f43f15ccd2b881bfea6fee7f3a46ec53f33b52210dea749ebebaa4" +dependencies = [ + "askama_macros", + "itoa", + "percent-encoding", + "serde", + "serde_json", +] + +[[package]] +name = "askama_derive" +version = "0.16.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "071ee5ebf2138e3ad180e0aacf6940c2cab5e6d8333741d9925c7bee2b153f39" +dependencies = [ + "askama_parser", + "memchr", + "proc-macro2", + "quote", + "rustc-hash", + "syn 3.0.6", +] + +[[package]] +name = "askama_macros" +version = "0.16.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "643e1c7cbb6aec1d920332fe51a7c0d8219e273dcb8602db03f5263e4d16487b" +dependencies = [ + "askama_derive", +] + +[[package]] +name = "askama_parser" +version = "0.16.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2c5ae75772275d268b03ab8bdccdd12117b6169ee23256942b34e46c9f476583" +dependencies = [ + "rustc-hash", + "unicode-ident", + "winnow 1.0.4", +] + [[package]] name = "asn1-rs" version = "0.7.2" @@ -4038,6 +4085,26 @@ dependencies = [ "strum", ] +[[package]] +name = "litellm-migrate" +version = "0.1.0" +dependencies = [ + "litellm-migrate-macros", + "rstest", +] + +[[package]] +name = "litellm-migrate-macros" +version = "0.1.0" +dependencies = [ + "proc-macro2", + "quote", + "rstest", + "syn 2.0.119", + "tempfile", + "thiserror 2.0.19", +] + [[package]] name = "litellm-model-catalog" version = "0.1.0" @@ -4384,11 +4451,17 @@ dependencies = [ name = "litellm-traces" version = "0.1.0" dependencies = [ + "askama", "base64 0.22.1", "criterion", "flate2", + "futures-util", + "hmac 0.12.1", + "indexmap 2.14.0", "litellm-http", + "litellm-migrate", "litellm-storage-clickhouse", + "moka", "opentelemetry-proto", "prost", "rstest", @@ -4400,6 +4473,7 @@ dependencies = [ "thiserror 2.0.19", "time", "tokio", + "url", "wiremock", ] diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index 8d837c2d31b..f4cb2ecbe59 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -14,6 +14,8 @@ litellm-router = { path = "crates/router" } litellm-tracing = { path = "crates/tracing" } litellm-traces = { path = "crates/traces" } litellm-storage-clickhouse = { path = "crates/storage-clickhouse" } +litellm-migrate = { path = "crates/migrate" } +litellm-migrate-macros = { path = "crates/migrate-macros" } litellm-core = { path = "crates/core" } litellm-gateway-mcp = { path = "crates/gateway-mcp" } litellm-gateway = { path = "crates/gateway" } @@ -63,6 +65,7 @@ litellm-token-counter-tiktoken = { path = "crates/token-counter-tiktoken" } litellm-host-python = { path = "crates/host-python" } litellm-python-compat = { path = "crates/python-compat" } +askama = { version = "0.16.1", default-features = false, features = ["derive", "std"] } tracing = "0.1" axum = { version = "0.8.9", default-features = false, features = ["http1", "tokio", "multipart"] } axum-login = "0.18.0" @@ -93,7 +96,10 @@ serde = { version = "1.0", features = ["derive"] } serde_json = { version = "1.0", features = ["float_roundtrip"] } serde_with = { version = "=3.16.1", default-features = false, features = ["std", "macros"] } sha2 = "0.10" +syn = { version = "2", default-features = false } sqlx = { version = "0.9.0", default-features = false, features = ["json", "macros", "postgres", "runtime-tokio", "chrono", "tls-rustls-ring-native-roots"] } +proc-macro2 = "1" +quote = "1" subtle = "2" thiserror = "2.0" tokenizers = { version = "0.23.1", default-features = false, features = ["onig"] } diff --git a/litellm-rust/crates/migrate-macros/Cargo.toml b/litellm-rust/crates/migrate-macros/Cargo.toml new file mode 100644 index 00000000000..5cd68415ca2 --- /dev/null +++ b/litellm-rust/crates/migrate-macros/Cargo.toml @@ -0,0 +1,19 @@ +[package] +name = "litellm-migrate-macros" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[lib] +proc-macro = true + +[dependencies] +proc-macro2.workspace = true +quote.workspace = true +syn = { workspace = true, features = ["parsing", "printing", "proc-macro"] } +thiserror.workspace = true + +[dev-dependencies] +rstest.workspace = true +tempfile.workspace = true diff --git a/litellm-rust/crates/migrate-macros/src/error.rs b/litellm-rust/crates/migrate-macros/src/error.rs new file mode 100644 index 00000000000..9833009517b --- /dev/null +++ b/litellm-rust/crates/migrate-macros/src/error.rs @@ -0,0 +1,21 @@ +use std::io; + +#[derive(Debug, thiserror::Error)] +pub enum Error { + #[error("could not read migrations directory `{path}`")] + ReadDirectory { + path: String, + #[source] + source: io::Error, + }, + #[error( + "migration name `{name}` must be `_.sql` with a `[a-z0-9_]` description" + )] + InvalidName { name: String }, + #[error("migration version `{version}` is declared more than once")] + DuplicateVersion { version: u64 }, + #[error("migrations directory `{path}` contains no migrations")] + Empty { path: String }, + #[error("migration path `{path}` is not valid UTF-8")] + NonUtf8Path { path: String }, +} diff --git a/litellm-rust/crates/migrate-macros/src/lib.rs b/litellm-rust/crates/migrate-macros/src/lib.rs new file mode 100644 index 00000000000..501f59e6fc2 --- /dev/null +++ b/litellm-rust/crates/migrate-macros/src/lib.rs @@ -0,0 +1,199 @@ +mod error; + +use std::path::{Path, PathBuf}; + +use error::Error; +use proc_macro::TokenStream; +use quote::quote; +use syn::LitStr; + +struct Entry { + version: u64, + description: String, + path: PathBuf, +} + +fn resolve(dir: &Path) -> Result, Error> { + let mut entries = Vec::new(); + let files = std::fs::read_dir(dir).map_err(|source| Error::ReadDirectory { + path: dir.display().to_string(), + source, + })?; + for file in files { + let file = file.map_err(|source| Error::ReadDirectory { + path: dir.display().to_string(), + source, + })?; + let path = file.path(); + let name = path + .file_name() + .and_then(|name| name.to_str()) + .ok_or_else(|| Error::NonUtf8Path { + path: path.display().to_string(), + })? + .to_owned(); + let invalid = || Error::InvalidName { name: name.clone() }; + let stem = name + .strip_suffix(".sql") + .filter(|_| file.file_type().is_ok_and(|kind| kind.is_file())) + .and_then(|stem| stem.split_once('_')) + .filter(|(version, description)| { + !version.is_empty() + && version.bytes().all(|b| b.is_ascii_digit()) + && !description.is_empty() + && description + .bytes() + .all(|b| b.is_ascii_lowercase() || b.is_ascii_digit() || b == b'_') + }) + .ok_or_else(invalid)?; + let version = stem.0.parse::().map_err(|_| invalid())?; + entries.push(Entry { + version, + description: stem.1.to_owned(), + path, + }); + } + if entries.is_empty() { + return Err(Error::Empty { + path: dir.display().to_string(), + }); + } + entries.sort_by_key(|entry| entry.version); + for pair in entries.windows(2) { + if pair[0].version == pair[1].version { + return Err(Error::DuplicateVersion { + version: pair[0].version, + }); + } + } + Ok(entries) +} + +fn resolve_input(lit: &LitStr) -> Result, Error> { + let root = std::env::var("CARGO_MANIFEST_DIR") + .map(PathBuf::from) + .unwrap_or_default(); + let dir = root.join(lit.value()); + let dir = dir.canonicalize().map_err(|source| Error::ReadDirectory { + path: dir.display().to_string(), + source, + })?; + if dir.to_str().is_none() { + return Err(Error::NonUtf8Path { + path: dir.display().to_string(), + }); + } + resolve(&dir) +} + +#[proc_macro] +pub fn migrate(input: TokenStream) -> TokenStream { + let lit = syn::parse_macro_input!(input as LitStr); + match resolve_input(&lit) { + Ok(entries) => { + let migrations = entries.iter().map(|entry| { + let version = entry.version; + let description = &entry.description; + let path = entry + .path + .to_str() + .expect("canonical migration path is UTF-8"); + quote! { + ::litellm_migrate::Migration { + version: #version, + description: #description, + sql: ::core::include_str!(#path), + } + } + }); + quote! { &[#(#migrations),*] }.into() + } + Err(err) => syn::Error::new(lit.span(), err).to_compile_error().into(), + } +} + +#[cfg(test)] +mod tests { + use std::fs; + + use rstest::rstest; + use tempfile::TempDir; + + use super::{Error, resolve}; + + fn migrations_dir(files: &[&str]) -> TempDir { + let dir = TempDir::new().expect("tempdir"); + for file in files { + fs::write(dir.path().join(file), "SELECT 1").expect("write fixture"); + } + dir + } + + #[rstest] + fn orders_versions_numerically() { + let dir = migrations_dir(&["10_tenth.sql", "2_second.sql", "1_first.sql"]); + let entries = resolve(dir.path()).expect("resolves"); + let versions: Vec = entries.iter().map(|entry| entry.version).collect(); + let descriptions: Vec<&str> = entries + .iter() + .map(|entry| entry.description.as_str()) + .collect(); + assert_eq!(versions, [1, 2, 10]); + assert_eq!(descriptions, ["first", "second", "tenth"]); + } + + #[rstest] + #[case::dash_in_version(&["0001-dash.sql"])] + #[case::not_sql(&["notes.txt"])] + #[case::empty_description(&["0001_.sql"])] + #[case::non_digit_version(&["x_name.sql"])] + #[case::uppercase_description(&["0001_Upper.sql"])] + #[case::no_underscore(&["0001.sql"])] + #[case::plus_sign_version(&["+10_add.sql"])] + fn rejects_invalid_names(#[case] files: &[&str]) { + let dir = migrations_dir(files); + assert!(matches!( + resolve(dir.path()), + Err(Error::InvalidName { .. }) + )); + } + + #[rstest] + fn rejects_subdirectories() { + let dir = migrations_dir(&["0001_a.sql"]); + fs::create_dir(dir.path().join("0002_b.sql")).expect("subdir"); + assert!(matches!( + resolve(dir.path()), + Err(Error::InvalidName { .. }) + )); + } + + #[cfg(unix)] + #[rstest] + fn rejects_symlinks() { + let dir = migrations_dir(&["0001_a.sql"]); + let target = TempDir::new().expect("tempdir"); + let target_file = target.path().join("real.sql"); + fs::write(&target_file, "SELECT 2").expect("write fixture"); + std::os::unix::fs::symlink(&target_file, dir.path().join("0002_b.sql")).expect("symlink"); + assert!(matches!( + resolve(dir.path()), + Err(Error::InvalidName { .. }) + )); + } + + #[rstest] + fn rejects_duplicate_versions() { + let dir = migrations_dir(&["0001_a.sql", "1_b.sql"]); + assert!(matches!( + resolve(dir.path()), + Err(Error::DuplicateVersion { version: 1 }) + )); + } + + #[rstest] + fn rejects_empty_directory() { + let dir = migrations_dir(&[]); + assert!(matches!(resolve(dir.path()), Err(Error::Empty { .. }))); + } +} diff --git a/litellm-rust/crates/migrate/Cargo.toml b/litellm-rust/crates/migrate/Cargo.toml new file mode 100644 index 00000000000..bb1ecaa3128 --- /dev/null +++ b/litellm-rust/crates/migrate/Cargo.toml @@ -0,0 +1,12 @@ +[package] +name = "litellm-migrate" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +litellm-migrate-macros.workspace = true + +[dev-dependencies] +rstest.workspace = true diff --git a/litellm-rust/crates/migrate/README.md b/litellm-rust/crates/migrate/README.md new file mode 100644 index 00000000000..4817029451c --- /dev/null +++ b/litellm-rust/crates/migrate/README.md @@ -0,0 +1,5 @@ +# Migrations + +`litellm-migrate` exports the `Migration` struct and the `migrate!` macro that embeds a directory of `_.sql` files at compile time, sorted by numeric version + +The crate does not apply or track migrations; callers decide how and when the embedded SQL runs diff --git a/litellm-rust/crates/migrate/src/lib.rs b/litellm-rust/crates/migrate/src/lib.rs new file mode 100644 index 00000000000..f4e065e1b53 --- /dev/null +++ b/litellm-rust/crates/migrate/src/lib.rs @@ -0,0 +1,8 @@ +pub use litellm_migrate_macros::migrate; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct Migration { + pub version: u64, + pub description: &'static str, + pub sql: &'static str, +} diff --git a/litellm-rust/crates/migrate/tests/fixtures/migrations/10_tenth.sql b/litellm-rust/crates/migrate/tests/fixtures/migrations/10_tenth.sql new file mode 100644 index 00000000000..31807719e9c --- /dev/null +++ b/litellm-rust/crates/migrate/tests/fixtures/migrations/10_tenth.sql @@ -0,0 +1 @@ +SELECT 10; diff --git a/litellm-rust/crates/migrate/tests/fixtures/migrations/1_first.sql b/litellm-rust/crates/migrate/tests/fixtures/migrations/1_first.sql new file mode 100644 index 00000000000..e0ac49d1ecf --- /dev/null +++ b/litellm-rust/crates/migrate/tests/fixtures/migrations/1_first.sql @@ -0,0 +1 @@ +SELECT 1; diff --git a/litellm-rust/crates/migrate/tests/fixtures/migrations/2_second.sql b/litellm-rust/crates/migrate/tests/fixtures/migrations/2_second.sql new file mode 100644 index 00000000000..e7f8100648d --- /dev/null +++ b/litellm-rust/crates/migrate/tests/fixtures/migrations/2_second.sql @@ -0,0 +1 @@ +SELECT 2; diff --git a/litellm-rust/crates/migrate/tests/migrate.rs b/litellm-rust/crates/migrate/tests/migrate.rs new file mode 100644 index 00000000000..61c80351cf4 --- /dev/null +++ b/litellm-rust/crates/migrate/tests/migrate.rs @@ -0,0 +1,21 @@ +use litellm_migrate::Migration; +use rstest::rstest; + +const MIGRATIONS: &[Migration] = litellm_migrate::migrate!("tests/fixtures/migrations"); + +#[rstest] +#[case::first(0, 1, "first", include_str!("fixtures/migrations/1_first.sql"))] +#[case::second(1, 2, "second", include_str!("fixtures/migrations/2_second.sql"))] +#[case::tenth(2, 10, "tenth", include_str!("fixtures/migrations/10_tenth.sql"))] +fn embeds_every_file_sorted_by_numeric_version( + #[case] index: usize, + #[case] version: u64, + #[case] description: &str, + #[case] sql: &str, +) { + assert_eq!(MIGRATIONS.len(), 3); + let migration = &MIGRATIONS[index]; + assert_eq!(migration.version, version); + assert_eq!(migration.description, description); + assert_eq!(migration.sql, sql); +} diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index d269fa4015f..6e091f0594f 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -44,7 +44,10 @@ mod _native { #[pymodule_export] use crate::routes::token_counter::TokenCounter; #[pymodule_export] - use crate::routes::traces::{NativeTraceStorage, trace_decode_otlp, trace_encode_error}; + use crate::routes::traces::{ + NativeTraceStorage, trace_decode_otlp, trace_encode_error, + trace_normalized_field_definitions, + }; #[cfg(feature = "huggingface")] #[pymodule_export] use crate::tokenizer::HuggingFaceEncoding; @@ -112,6 +115,7 @@ mod tests { "NativeTraceStorage", "trace_decode_otlp", "trace_encode_error", + "trace_normalized_field_definitions", "TokenCounter", "Tokenizer", "gil_stats", diff --git a/litellm-rust/crates/python-bridge/src/routes/traces.rs b/litellm-rust/crates/python-bridge/src/routes/traces.rs index ca66e2e46be..0ff5fe98f55 100644 --- a/litellm-rust/crates/python-bridge/src/routes/traces.rs +++ b/litellm-rust/crates/python-bridge/src/routes/traces.rs @@ -3,7 +3,9 @@ use std::collections::BTreeMap; use litellm_host_python::{FromPythonCache, ToPythonCache}; use litellm_http::ClientVariant; use litellm_storage_clickhouse::Storage; -use litellm_traces::{Error, InsertTable, Parameter, ReadQuery, Shared}; +use litellm_traces::{ + Error, InsertTable, Parameter, QueryAccessError, QueryReaders, QueryScope, ReadQuery, Shared, +}; use prost::Message; use pyo3::{ exceptions::{PyOverflowError, PyRuntimeError, PyValueError}, @@ -46,9 +48,25 @@ fn map_error(error: Error) -> PyErr { } } +fn map_sql_error(error: Error) -> PyErr { + match error { + Error::QueryFailed(400 | 404) => PyValueError::new_err(error.to_string()), + error => map_error(error), + } +} + +fn map_query_access_error(error: QueryAccessError) -> PyErr { + match error { + QueryAccessError::Storage(error) => map_sql_error(error), + QueryAccessError::InvalidScope => PyValueError::new_err(error.to_string()), + error => PyRuntimeError::new_err(error.to_string()), + } +} + #[pyclass] pub struct NativeTraceStorage { storage: Storage, + query_readers: QueryReaders, } #[pymethods] @@ -57,8 +75,13 @@ impl NativeTraceStorage { #[pyo3(signature = (database, url, reader_url = None))] fn new(database: String, url: &str, reader_url: Option<&str>) -> PyResult { litellm_traces::schema_statements(&database, 1, 1).map_err(map_error)?; + let storage = Storage::new(database, url, reader_url).map_err(map_error)?; Ok(Self { - storage: Storage::new(database, url, reader_url).map_err(map_error)?, + query_readers: QueryReaders::new( + storage.writer().clone(), + storage.database().to_owned(), + ), + storage, }) } @@ -107,6 +130,52 @@ impl NativeTraceStorage { ) } + fn query_sql<'py>( + &self, + py: Python<'py>, + sql: String, + #[pyo3(from_py_with = litellm_host_python::from_py_argument)] scope: QueryScope, + secret: String, + ) -> PyResult> { + if sql.trim().is_empty() { + return Err(map_error(Error::EmptySql)); + } + let readers = self.query_readers.clone(); + let client = crate::http::host_client(py, ClientVariant::NoRedirect)?; + crate::execution::run_async( + py, + async move { + let _permit = readers.acquire()?; + let connection = readers.connection(&client, &scope, &secret).await?; + litellm_traces::query_sql(&client, &connection, &sql) + .await + .map_err(QueryAccessError::Storage) + }, + map_query_access_error, + ) + } + + fn query_help<'py>( + &self, + py: Python<'py>, + #[pyo3(from_py_with = litellm_host_python::from_py_argument)] scope: QueryScope, + secret: String, + ) -> PyResult> { + let readers = self.query_readers.clone(); + let client = crate::http::host_client(py, ClientVariant::NoRedirect)?; + crate::execution::run_async( + py, + async move { + let _permit = readers.acquire()?; + let connection = readers.connection(&client, &scope, &secret).await?; + litellm_traces::query_help(&client, &connection) + .await + .map_err(QueryAccessError::Storage) + }, + map_query_access_error, + ) + } + fn lens_query<'py>( &self, py: Python<'py>, @@ -236,6 +305,14 @@ fn spans_to_py<'py>( "events", litellm_host_python::Pythonized(&span.events).into_pyobject(py)?, )?; + row.set_item( + "normalized", + litellm_host_python::Pythonized(&span.normalized).into_pyobject(py)?, + )?; + row.set_item( + "consumed_attributes", + litellm_host_python::Pythonized(&span.consumed_attributes).into_pyobject(py)?, + )?; result.append(row)?; } Ok(result) @@ -288,3 +365,8 @@ mod tests { }); } } + +#[pyfunction] +pub fn trace_normalized_field_definitions<'py>(py: Python<'py>) -> PyResult> { + litellm_host_python::Pythonized(litellm_traces::NORMALIZED_FIELD_DEFINITIONS).into_pyobject(py) +} diff --git a/litellm-rust/crates/traces/AGENTS.md b/litellm-rust/crates/traces/AGENTS.md index 645e88dfae1..d0181c53308 100644 --- a/litellm-rust/crates/traces/AGENTS.md +++ b/litellm-rust/crates/traces/AGENTS.md @@ -1,6 +1,6 @@ - Keep OTLP decoding, trace schema, row encoding and named query selection here. Generic ClickHouse connections and HTTP execution belong in `litellm-storage-clickhouse` - Keep this crate independent of Python; PyO3 conversion and public Python exceptions belong in `python-bridge` -- Keep the SQL migrations here as the only ClickHouse schema definition +- Keep the SQL migrations here as the only ClickHouse schema definition, as `migrations/NNNN_description.sql` files embedded by `litellm_migrate::migrate!`; adding a file is the only step - Use typed query parameters and a dedicated SELECT-only reader with server-side limits - Keep `config/reader.xml` grants on the database the schema is created in (CLICKHOUSE_DATABASE, default `litellm`) - Bound insert time and encoded bytes; make retry deduplication behavior explicit for supported ClickHouse versions diff --git a/litellm-rust/crates/traces/Cargo.toml b/litellm-rust/crates/traces/Cargo.toml index 74de400764c..5e3b41e719a 100644 --- a/litellm-rust/crates/traces/Cargo.toml +++ b/litellm-rust/crates/traces/Cargo.toml @@ -6,18 +6,26 @@ license.workspace = true repository.workspace = true [dependencies] +askama.workspace = true base64.workspace = true flate2.workspace = true +futures-util.workspace = true +hmac = "0.12.1" +indexmap = { version = "2", features = ["serde"] } +moka.workspace = true opentelemetry-proto = { workspace = true, features = ["gen-tonic-messages", "trace", "with-serde"] } prost.workspace = true time = { workspace = true, features = ["formatting"] } litellm-http.workspace = true +litellm-migrate.workspace = true litellm-storage-clickhouse.workspace = true sha2.workspace = true serde = { workspace = true, features = ["rc"] } serde_json.workspace = true strum.workspace = true thiserror.workspace = true +tokio.workspace = true +url.workspace = true [dev-dependencies] criterion.workspace = true diff --git a/litellm-rust/crates/traces/build.rs b/litellm-rust/crates/traces/build.rs new file mode 100644 index 00000000000..3a8149ef075 --- /dev/null +++ b/litellm-rust/crates/traces/build.rs @@ -0,0 +1,3 @@ +fn main() { + println!("cargo:rerun-if-changed=migrations"); +} diff --git a/litellm-rust/crates/traces/migrations/0001_otel_traces.sql b/litellm-rust/crates/traces/migrations/0001_otel_traces.sql index d8e0184b5a3..fb5eaa367d7 100644 --- a/litellm-rust/crates/traces/migrations/0001_otel_traces.sql +++ b/litellm-rust/crates/traces/migrations/0001_otel_traces.sql @@ -38,10 +38,11 @@ CREATE TABLE IF NOT EXISTS {database}.otel_traces Input String CODEC(ZSTD(3)), Output String CODEC(ZSTD(3)), InputPreview String DEFAULT substring(Input, 1, 240), + EngineReceivedMs UInt64 DEFAULT 0, INDEX idx_trace_id TraceId TYPE bloom_filter(0.001) GRANULARITY 1, INDEX idx_req_id LiteLLMRequestId TYPE bloom_filter(0.01) GRANULARITY 1 ) ENGINE = MergeTree PARTITION BY toDate(Timestamp) ORDER BY (TeamId, ServiceName, toDateTime(Timestamp), TraceId) -SETTINGS ttl_only_drop_parts = 1, non_replicated_deduplication_window = 1000 +SETTINGS ttl_only_drop_parts = 1, materialize_ttl_recalculate_only = 1, non_replicated_deduplication_window = 1000 diff --git a/litellm-rust/crates/traces/migrations/0005_otel_traces_ttl.sql b/litellm-rust/crates/traces/migrations/0002_otel_traces_ttl.sql similarity index 100% rename from litellm-rust/crates/traces/migrations/0005_otel_traces_ttl.sql rename to litellm-rust/crates/traces/migrations/0002_otel_traces_ttl.sql diff --git a/litellm-rust/crates/traces/migrations/0002_agent_traces.sql b/litellm-rust/crates/traces/migrations/0003_agent_traces.sql similarity index 93% rename from litellm-rust/crates/traces/migrations/0002_agent_traces.sql rename to litellm-rust/crates/traces/migrations/0003_agent_traces.sql index 0c3547872bb..821cc2f3723 100644 --- a/litellm-rust/crates/traces/migrations/0002_agent_traces.sql +++ b/litellm-rust/crates/traces/migrations/0003_agent_traces.sql @@ -22,4 +22,4 @@ CREATE TABLE IF NOT EXISTS {database}.agent_traces_by_key ) ENGINE = AggregatingMergeTree ORDER BY (TeamId, ApiKeyHash, TraceId) -SETTINGS non_replicated_deduplication_window = 1000 +SETTINGS materialize_ttl_recalculate_only = 1, non_replicated_deduplication_window = 1000 diff --git a/litellm-rust/crates/traces/migrations/0006_agent_traces_ttl.sql b/litellm-rust/crates/traces/migrations/0004_agent_traces_ttl.sql similarity index 100% rename from litellm-rust/crates/traces/migrations/0006_agent_traces_ttl.sql rename to litellm-rust/crates/traces/migrations/0004_agent_traces_ttl.sql diff --git a/litellm-rust/crates/traces/migrations/0003_agent_traces_mv.sql b/litellm-rust/crates/traces/migrations/0005_agent_traces_mv.sql similarity index 100% rename from litellm-rust/crates/traces/migrations/0003_agent_traces_mv.sql rename to litellm-rust/crates/traces/migrations/0005_agent_traces_mv.sql diff --git a/litellm-rust/crates/traces/migrations/0004_spend_logs.sql b/litellm-rust/crates/traces/migrations/0006_spend_logs.sql similarity index 94% rename from litellm-rust/crates/traces/migrations/0004_spend_logs.sql rename to litellm-rust/crates/traces/migrations/0006_spend_logs.sql index a14930f438f..44f7959b2bf 100644 --- a/litellm-rust/crates/traces/migrations/0004_spend_logs.sql +++ b/litellm-rust/crates/traces/migrations/0006_spend_logs.sql @@ -34,9 +34,11 @@ CREATE TABLE IF NOT EXISTS {database}.spend_logs metadata String CODEC(ZSTD(3)), messages String CODEC(ZSTD(3)), response String CODEC(ZSTD(3)), + EngineReceivedMs UInt64 DEFAULT 0, INDEX idx_response_id response_id TYPE bloom_filter(0.001) GRANULARITY 1, INDEX idx_trace_id trace_id TYPE bloom_filter(0.001) GRANULARITY 1 ) ENGINE = ReplacingMergeTree(end_time) PARTITION BY toYYYYMM(start_time) ORDER BY (team_id, start_time, request_id) +SETTINGS materialize_ttl_recalculate_only = 1 diff --git a/litellm-rust/crates/traces/migrations/0008_trace_received.sql b/litellm-rust/crates/traces/migrations/0008_trace_received.sql deleted file mode 100644 index 9d8113b2430..00000000000 --- a/litellm-rust/crates/traces/migrations/0008_trace_received.sql +++ /dev/null @@ -1 +0,0 @@ -ALTER TABLE {database}.otel_traces ADD COLUMN IF NOT EXISTS EngineReceivedMs UInt64 DEFAULT 0 diff --git a/litellm-rust/crates/traces/migrations/0009_spend_received.sql b/litellm-rust/crates/traces/migrations/0009_spend_received.sql deleted file mode 100644 index 2b2d2c7e5d7..00000000000 --- a/litellm-rust/crates/traces/migrations/0009_spend_received.sql +++ /dev/null @@ -1 +0,0 @@ -ALTER TABLE {database}.spend_logs ADD COLUMN IF NOT EXISTS EngineReceivedMs UInt64 DEFAULT 0 diff --git a/litellm-rust/crates/traces/query/lens_agents.sql b/litellm-rust/crates/traces/query/lens_agents.sql new file mode 100644 index 00000000000..fbdd578f8e7 --- /dev/null +++ b/litellm-rust/crates/traces/query/lens_agents.sql @@ -0,0 +1,6 @@ +SELECT DISTINCT AgentName AS agent_name +FROM otel_traces +WHERE AgentName != '' + AND ({all_teams:UInt8}=1 OR TeamId={team:String}) + AND ({key_hash:String}='' OR ApiKeyHash={key_hash:String}) +ORDER BY agent_name diff --git a/litellm-rust/crates/traces/query/lens_availability.sql b/litellm-rust/crates/traces/query/lens_availability.sql new file mode 100644 index 00000000000..8d350dd1779 --- /dev/null +++ b/litellm-rust/crates/traces/query/lens_availability.sql @@ -0,0 +1,8 @@ +SELECT + EXISTS(SELECT 1 FROM otel_traces + WHERE ({all_teams:UInt8}=1 OR TeamId={team:String}) + AND ({key_hash:String}='' OR ApiKeyHash={key_hash:String})) AS traces, + EXISTS(SELECT 1 FROM spend_logs + WHERE ({all_teams:UInt8}=1 OR team_id={team:String}) + AND ({key_hash:String}='' OR api_key={key_hash:String}) + AND NOT JSONExtractBool(metadata,'litellm_lens_internal')) AS requests diff --git a/litellm-rust/crates/traces/query/lens_sample.sql b/litellm-rust/crates/traces/query/lens_sample.sql index 1fc9c964a6f..92086c33c13 100644 --- a/litellm-rust/crates/traces/query/lens_sample.sql +++ b/litellm-rust/crates/traces/query/lens_sample.sql @@ -28,6 +28,7 @@ SELECT *, selection_key FROM ( GROUP BY TeamId,ApiKeyHash,TraceId HAVING max(EngineReceivedMs) < {end:UInt64} AND max(toUnixTimestamp64Milli(Timestamp)+toInt64(intDiv(Duration,1000000))) < {end:UInt64} + AND ({agent_name:String}='' OR countIf(AgentName={agent_name:String}) > 0) AND countIf(arrayAll((k,v) -> ResourceAttributes[k]=v OR SpanAttributes[k]=v, {filter_keys:Array(String)},{filter_values:Array(String)}) AND ({service:String}='' OR ServiceName={service:String})) > 0 @@ -48,6 +49,7 @@ SELECT *, selection_key FROM ( OR JSONExtractString(metadata,'requester_metadata',k)=v OR (k='tag' AND has(request_tags,v)), {filter_keys:Array(String)},{filter_values:Array(String)}) AND ({service:String}='' OR model_group={service:String}) + AND {agent_name:String}='' AND NOT JSONExtractBool(metadata,'litellm_lens_internal') AND ({source:String}!='both' OR (team_id,api_key,response_id) NOT IN ( SELECT TeamId,ApiKeyHash,LiteLLMRequestId FROM otel_traces diff --git a/litellm-rust/crates/traces/src/error.rs b/litellm-rust/crates/traces/src/error.rs index 18fa4af9b53..c677e73e96f 100644 --- a/litellm-rust/crates/traces/src/error.rs +++ b/litellm-rust/crates/traces/src/error.rs @@ -4,4 +4,26 @@ pub enum DecodeError { InvalidPayload, #[error("OTLP trace payload exceeds the decoding budget")] TooLarge, + #[error("OTLP token count is outside the storage range")] + TokenCountOutOfRange, +} + +#[derive(Debug, thiserror::Error)] +pub enum QueryAccessError { + #[error("trace SQL queries require a configured proxy master key")] + MissingSecret, + #[error("invalid trace query scope")] + InvalidScope, + #[error("trace SQL query concurrency limit exceeded")] + Busy, + #[error( + "ClickHouse reader provisioning failed with HTTP status {0}; the configured connection must be allowed to manage users, row policies, and SELECT grants on the trace tables" + )] + ProvisionFailed(u16), + #[error("ClickHouse reader provisioning transport failed")] + ProvisionTransport, + #[error(transparent)] + Storage(#[from] litellm_storage_clickhouse::Error), + #[error(transparent)] + Cached(#[from] std::sync::Arc), } diff --git a/litellm-rust/crates/traces/src/lib.rs b/litellm-rust/crates/traces/src/lib.rs index 1489b44c118..9e01d10f73e 100644 --- a/litellm-rust/crates/traces/src/lib.rs +++ b/litellm-rust/crates/traces/src/lib.rs @@ -1,14 +1,23 @@ mod error; mod insert; +mod normalize; mod otlp; +mod query; +mod query_access; mod schema; mod shared; mod sql; -pub use error::DecodeError; +pub use error::{DecodeError, QueryAccessError}; pub use insert::{InsertRow, InsertTable, encode_rows, insert_rows, insert_shared_rows}; pub use litellm_storage_clickhouse::{Connection, Error, Parameter, execute_read}; +pub use normalize::{ + NORMALIZED_FIELD_DEFINITIONS, NormalizedFieldDefinition, NormalizedSpan, ObservationType, +}; pub use otlp::{DecodedSpan, decode_otlp}; +pub use query_access::{QueryReaders, QueryScope}; pub use schema::{ensure_schema, schema_statements}; pub use shared::{Shared, SharedIdentity}; pub use sql::{LensQuery, ReadQuery, execute_named_read}; + +pub use query::{query_help, query_sql}; diff --git a/litellm-rust/crates/traces/src/normalize/genai.rs b/litellm-rust/crates/traces/src/normalize/genai.rs new file mode 100644 index 00000000000..15a334ab0fd --- /dev/null +++ b/litellm-rust/crates/traces/src/normalize/genai.rs @@ -0,0 +1,63 @@ +use std::collections::BTreeMap; + +use super::{NormalizedSpan, ObservationType, SpanNormalizer, attr, first, usage_tokens}; +use crate::DecodeError; + +pub(super) struct GenAiNormalizer; + +impl SpanNormalizer for GenAiNormalizer { + fn matches(&self, _scope_name: &str, _attributes: &BTreeMap) -> bool { + true + } + + fn consumed_attributes(&self, attributes: &BTreeMap) -> [&'static str; 2] { + [ + if attr(attributes, "gen_ai.input.messages").is_empty() { + "gen_ai.tool.call.arguments" + } else { + "gen_ai.input.messages" + }, + if attr(attributes, "gen_ai.output.messages").is_empty() { + "gen_ai.tool.call.result" + } else { + "gen_ai.output.messages" + }, + ] + } + + fn normalize( + &self, + _name: &str, + parent_span_id: &str, + attributes: &BTreeMap, + ) -> Result { + let (input_tokens, output_tokens) = usage_tokens(attributes)?; + let observation_type = match attr(attributes, "gen_ai.operation.name") { + "invoke_agent" => ObservationType::Agent, + "chat" | "text_completion" | "generate_content" => ObservationType::Llm, + "execute_tool" => ObservationType::Tool, + _ if parent_span_id.is_empty() => ObservationType::Agent, + _ => ObservationType::Chain, + }; + Ok(NormalizedSpan { + observation_type, + agent_name: attr(attributes, "gen_ai.agent.name").to_owned(), + litellm_request_id: attr(attributes, "gen_ai.response.id").to_owned(), + model: first(attributes, "gen_ai.request.model", "gen_ai.response.model").to_owned(), + input_tokens, + output_tokens, + input: first( + attributes, + "gen_ai.input.messages", + "gen_ai.tool.call.arguments", + ) + .to_owned(), + output: first( + attributes, + "gen_ai.output.messages", + "gen_ai.tool.call.result", + ) + .to_owned(), + }) + } +} diff --git a/litellm-rust/crates/traces/src/normalize/langsmith.rs b/litellm-rust/crates/traces/src/normalize/langsmith.rs new file mode 100644 index 00000000000..bfe7a2c1796 --- /dev/null +++ b/litellm-rust/crates/traces/src/normalize/langsmith.rs @@ -0,0 +1,468 @@ +use std::{collections::BTreeMap, io}; + +use indexmap::IndexMap; +use serde::{Deserialize, Deserializer, Serialize, de::DeserializeOwned}; +use serde_json::{Value, ser::Formatter}; + +use super::{NormalizedSpan, ObservationType, SpanNormalizer, attr, usage_tokens}; +use crate::DecodeError; + +pub(super) struct LangSmithNormalizer; + +#[derive(Deserialize)] +#[serde(untagged)] +enum MessageContent { + Text(String), + Blocks(Vec), + Other(Value), +} + +impl MessageContent { + fn display_text(&self) -> String { + match self { + Self::Text(text) => text.clone(), + Self::Blocks(blocks) => blocks + .iter() + .filter_map(|block| match block { + ContentBlock::Text { text } => Some(text.as_str()), + ContentBlock::Hidden(kind) => match kind { + HiddenBlock::Reasoning + | HiddenBlock::Thinking + | HiddenBlock::RedactedThinking + | HiddenBlock::FunctionCall + | HiddenBlock::ToolUse + | HiddenBlock::ToolCall => None, + }, + }) + .collect::>() + .join("\n\n"), + Self::Other(value) => encode(value), + } + } +} + +#[derive(Deserialize)] +#[serde(untagged)] +enum ContentBlock { + Text { text: String }, + Hidden(HiddenBlock), +} + +#[derive(Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +enum HiddenBlock { + Reasoning, + Thinking, + RedactedThinking, + FunctionCall, + ToolUse, + ToolCall, +} + +#[derive(Deserialize, Serialize)] +#[serde(transparent)] +struct RawToolCall(IndexMap); + +#[derive(Deserialize)] +struct ResponseMetadata { + id: Option, +} + +#[derive(Deserialize)] +struct RawMessage { + kwargs: Option>, + #[serde(rename = "type")] + kind: Option, + role: Option, + content: Option, + tool_calls: Option>, + name: Option, + response_metadata: Option, +} + +impl RawMessage { + fn unwrapped(&self) -> &Self { + self.kwargs.as_deref().unwrap_or(self) + } + + fn normalized(&self) -> NormalizedMessage<'_> { + let fields = self.unwrapped(); + let raw_role = fields + .kind + .as_deref() + .filter(|role| !role.is_empty()) + .or_else(|| fields.role.as_deref().filter(|role| !role.is_empty())) + .unwrap_or_default(); + let role = match raw_role { + "human" => "user", + "ai" => "assistant", + other => other, + }; + NormalizedMessage { + role, + content: fields + .content + .as_ref() + .map_or_else(String::new, MessageContent::display_text), + tool_calls: fields + .tool_calls + .as_deref() + .filter(|calls| !calls.is_empty()), + name: (role == "tool") + .then_some(fields.name.as_ref()) + .flatten() + .filter(|name| !name.is_null() && name != &&Value::String(String::new())), + } + } +} + +#[derive(Serialize)] +struct NormalizedMessage<'a> { + role: &'a str, + content: String, + #[serde(skip_serializing_if = "Option::is_none")] + tool_calls: Option<&'a [RawToolCall]>, + #[serde(skip_serializing_if = "Option::is_none")] + name: Option<&'a Value>, +} + +enum MessageBatch { + Flat(Vec), + Nested(Vec>), +} + +impl<'de> Deserialize<'de> for MessageBatch { + fn deserialize>(deserializer: D) -> Result { + let value = Value::deserialize(deserializer)?; + let Value::Array(items) = value else { + return Err(serde::de::Error::custom("messages must be an array")); + }; + let parse = |items: Vec| { + items + .into_iter() + .filter_map(|item| serde_json::from_value(item).ok()) + .collect() + }; + Ok(if items.first().is_some_and(Value::is_array) { + Self::Nested( + items + .into_iter() + .filter_map(|item| item.as_array().cloned()) + .map(parse) + .collect(), + ) + } else { + Self::Flat(parse(items)) + }) + } +} + +fn lenient<'de, D: Deserializer<'de>, T: DeserializeOwned>( + deserializer: D, +) -> Result, D::Error> { + let value = Value::deserialize(deserializer)?; + Ok(serde_json::from_value(value).ok()) +} + +impl MessageBatch { + fn first_batch(&self) -> &[RawMessage] { + match self { + Self::Flat(messages) => messages, + Self::Nested(batches) => batches.first().map(Vec::as_slice).unwrap_or_default(), + } + } + + fn agent_messages(&self) -> &[RawMessage] { + match self { + Self::Flat(messages) => messages, + Self::Nested(_) => &[], + } + } +} + +#[derive(Deserialize)] +struct GenerationMessage { + kwargs: Option, +} + +#[derive(Deserialize)] +struct Generation { + message: Option, +} + +#[derive(Default, Deserialize)] +struct Payload { + #[serde(default, deserialize_with = "lenient")] + messages: Option, + #[serde(default, deserialize_with = "lenient")] + generations: Option>>, +} + +#[derive(Deserialize)] +struct Command { + update: CommandUpdate, +} + +#[derive(Deserialize)] +struct CommandUpdate { + messages: Vec, +} + +#[derive(Deserialize)] +struct ContentValue { + content: Value, +} + +struct SpanIo { + input: String, + output: String, + request_id: String, +} + +struct PythonJsonFormatter; + +impl Formatter for PythonJsonFormatter { + fn begin_array_value( + &mut self, + writer: &mut W, + first: bool, + ) -> io::Result<()> { + if first { + Ok(()) + } else { + writer.write_all(b", ") + } + } + + fn begin_object_key( + &mut self, + writer: &mut W, + first: bool, + ) -> io::Result<()> { + if first { + Ok(()) + } else { + writer.write_all(b", ") + } + } + + fn begin_object_value(&mut self, writer: &mut W) -> io::Result<()> { + writer.write_all(b": ") + } +} + +fn encode(value: &T) -> String { + let mut output = Vec::new(); + let mut serializer = serde_json::Serializer::with_formatter(&mut output, PythonJsonFormatter); + if value.serialize(&mut serializer).is_err() { + return String::new(); + } + String::from_utf8(output).unwrap_or_default() +} + +fn normalized_messages(messages: &[RawMessage]) -> String { + encode( + &messages + .iter() + .map(RawMessage::normalized) + .collect::>(), + ) +} + +fn span_type( + name: &str, + parent_span_id: &str, + attributes: &BTreeMap, +) -> ObservationType { + match attr(attributes, "langsmith.span.kind") { + "llm" => ObservationType::Llm, + "tool" => ObservationType::Tool, + _ if parent_span_id.is_empty() + || name == attr(attributes, "langsmith.metadata.lc_agent_name") => + { + ObservationType::Agent + } + _ if [ + ".wrap_model_call", + ".wrap_tool_call", + ".before_agent", + ".after_agent", + ".before_model", + ".after_model", + ] + .iter() + .any(|suffix| name.ends_with(suffix)) => + { + ObservationType::Framework + } + _ => ObservationType::Chain, + } +} + +fn tool_output(raw_completion: &str) -> String { + let completion = serde_json::from_str::(raw_completion).unwrap_or(Value::Null); + let raw = completion.get("output").cloned().unwrap_or(completion); + let selected = serde_json::from_value::(raw.clone()) + .ok() + .and_then(|command| command.update.messages.into_iter().last()) + .unwrap_or(raw); + let output = serde_json::from_value::(selected.clone()) + .map(|message| message.content) + .unwrap_or(selected); + output + .as_str() + .map(str::to_owned) + .unwrap_or_else(|| encode(&output)) +} + +fn span_io(kind: ObservationType, attributes: &BTreeMap) -> SpanIo { + let raw_prompt = attr(attributes, "gen_ai.prompt"); + let raw_completion = attr(attributes, "gen_ai.completion"); + let prompt = serde_json::from_str::(raw_prompt).unwrap_or_default(); + let completion = serde_json::from_str::(raw_completion).unwrap_or_default(); + if kind == ObservationType::Llm + && serde_json::from_str::(raw_completion).is_ok_and(|value| value.is_object()) + { + let input = prompt.messages.as_ref().map_or_else( + || "[]".to_owned(), + |messages| normalized_messages(messages.first_batch()), + ); + let generation = completion + .generations + .as_ref() + .and_then(|batches| batches.first()) + .and_then(|batch| batch.first()) + .and_then(|generation| generation.message.as_ref()) + .and_then(|message| message.kwargs.as_ref()); + if let Some(generation) = generation { + let id = generation + .response_metadata + .as_ref() + .and_then(|metadata| metadata.id.as_deref()) + .unwrap_or_default() + .to_owned(); + return SpanIo { + input, + output: encode(&generation.normalized()), + request_id: id, + }; + } + return SpanIo { + input, + output: raw_completion.to_owned(), + request_id: String::new(), + }; + } + if kind == ObservationType::Tool { + return SpanIo { + input: raw_prompt.to_owned(), + output: tool_output(raw_completion), + request_id: String::new(), + }; + } + if kind == ObservationType::Agent { + let input = prompt + .messages + .as_ref() + .filter(|messages| !messages.agent_messages().is_empty()) + .map_or_else( + || raw_prompt.to_owned(), + |messages| normalized_messages(messages.agent_messages()), + ); + let output = completion + .messages + .as_ref() + .and_then(|messages| messages.agent_messages().last()) + .map_or_else( + || raw_completion.to_owned(), + |message| encode(&message.normalized()), + ); + return SpanIo { + input, + output, + request_id: String::new(), + }; + } + SpanIo { + input: raw_prompt.to_owned(), + output: raw_completion.to_owned(), + request_id: String::new(), + } +} + +impl SpanNormalizer for LangSmithNormalizer { + fn matches(&self, scope_name: &str, attributes: &BTreeMap) -> bool { + scope_name == "langsmith" || attributes.contains_key("langsmith.span.kind") + } + + fn consumed_attributes(&self, _attributes: &BTreeMap) -> [&'static str; 2] { + ["gen_ai.prompt", "gen_ai.completion"] + } + + fn normalize( + &self, + name: &str, + parent_span_id: &str, + attributes: &BTreeMap, + ) -> Result { + let (input_tokens, output_tokens) = usage_tokens(attributes)?; + let observation_type = span_type(name, parent_span_id, attributes); + let io = span_io(observation_type, attributes); + Ok(NormalizedSpan { + observation_type, + agent_name: attr(attributes, "langsmith.metadata.lc_agent_name").to_owned(), + litellm_request_id: io.request_id, + model: attr(attributes, "gen_ai.request.model").to_owned(), + input_tokens, + output_tokens, + input: io.input, + output: io.output, + }) + } +} + +#[cfg(test)] +mod tests { + use std::collections::BTreeMap; + + use rstest::rstest; + use serde_json::Value; + + use super::{ObservationType, span_io}; + + #[rstest] + fn malformed_messages_preserve_valid_input_and_response_id() { + let attributes = BTreeMap::from([ + ( + "gen_ai.prompt".to_owned(), + r#"{"messages":[[{"kwargs":{"type":"human","content":"hello"}},null]]}"#.to_owned(), + ), + ( + "gen_ai.completion".to_owned(), + r#"{"messages":"unexpected","generations":[[{"message":{"kwargs":{"type":"ai","content":"hi","response_metadata":{"id":"response-1"}}}}]]}"#.to_owned(), + ), + ]); + let io = span_io(ObservationType::Llm, &attributes); + let input: Value = serde_json::from_str(&io.input).expect("normalized input"); + assert_eq!(input.as_array().expect("messages").len(), 1); + assert_eq!(input[0]["content"], "hello"); + assert_eq!(io.request_id, "response-1"); + } + + #[rstest] + fn explicit_null_tool_output_is_preserved() { + let attributes = BTreeMap::from([( + "gen_ai.completion".to_owned(), + r#"{"output":null}"#.to_owned(), + )]); + let io = span_io(ObservationType::Tool, &attributes); + assert_eq!(io.output, "null"); + } + + #[rstest] + fn absent_llm_messages_render_as_an_empty_list() { + let attributes = BTreeMap::from([("gen_ai.completion".to_owned(), "{}".to_owned())]); + let io = span_io(ObservationType::Llm, &attributes); + assert_eq!(io.input, "[]"); + } +} diff --git a/litellm-rust/crates/traces/src/normalize/mod.rs b/litellm-rust/crates/traces/src/normalize/mod.rs new file mode 100644 index 00000000000..b4d5a4a9d06 --- /dev/null +++ b/litellm-rust/crates/traces/src/normalize/mod.rs @@ -0,0 +1,239 @@ +use std::collections::BTreeMap; + +use crate::DecodeError; +use serde::Serialize; + +#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize)] +#[serde(rename_all = "lowercase")] +pub enum ObservationType { + Agent, + Llm, + Tool, + Chain, + Framework, +} + +#[derive(Debug, Serialize)] +pub struct NormalizedSpan { + pub observation_type: ObservationType, + pub agent_name: String, + pub litellm_request_id: String, + pub model: String, + pub input_tokens: u32, + pub output_tokens: u32, + pub input: String, + pub output: String, +} + +pub(crate) struct Normalization { + pub span: NormalizedSpan, + pub consumed_attributes: [&'static str; 2], +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq, Serialize)] +pub struct NormalizedFieldDefinition { + pub name: &'static str, + pub clickhouse_column: &'static str, + pub clickhouse_type: &'static str, + pub meaning: &'static str, +} + +pub const NORMALIZED_FIELD_DEFINITIONS: [NormalizedFieldDefinition; 8] = [ + NormalizedFieldDefinition { + name: "observation_type", + clickhouse_column: "ObservationType", + clickhouse_type: "LowCardinality(String)", + meaning: "Agent, LLM, tool, chain, or framework span", + }, + NormalizedFieldDefinition { + name: "agent_name", + clickhouse_column: "AgentName", + clickhouse_type: "LowCardinality(String)", + meaning: "Agent associated with this span", + }, + NormalizedFieldDefinition { + name: "litellm_request_id", + clickhouse_column: "LiteLLMRequestId", + clickhouse_type: "String", + meaning: "LiteLLM response ID used to link a span to a spend log", + }, + NormalizedFieldDefinition { + name: "model", + clickhouse_column: "Model", + clickhouse_type: "LowCardinality(String)", + meaning: "Model used by this span", + }, + NormalizedFieldDefinition { + name: "input_tokens", + clickhouse_column: "InputTokens", + clickhouse_type: "UInt32", + meaning: "Input token count", + }, + NormalizedFieldDefinition { + name: "output_tokens", + clickhouse_column: "OutputTokens", + clickhouse_type: "UInt32", + meaning: "Output token count", + }, + NormalizedFieldDefinition { + name: "input", + clickhouse_column: "Input", + clickhouse_type: "String", + meaning: "Normalized input payload", + }, + NormalizedFieldDefinition { + name: "output", + clickhouse_column: "Output", + clickhouse_type: "String", + meaning: "Normalized output payload", + }, +]; + +trait SpanNormalizer { + fn matches(&self, scope_name: &str, attributes: &BTreeMap) -> bool; + fn consumed_attributes(&self, attributes: &BTreeMap) -> [&'static str; 2]; + fn normalize( + &self, + name: &str, + parent_span_id: &str, + attributes: &BTreeMap, + ) -> Result; +} + +mod genai; +mod langsmith; +mod openinference; + +use genai::GenAiNormalizer; +use langsmith::LangSmithNormalizer; +use openinference::OpenInferenceNormalizer; + +fn attr<'a>(attributes: &'a BTreeMap, key: &str) -> &'a str { + attributes.get(key).map(String::as_str).unwrap_or_default() +} + +fn first<'a>(attributes: &'a BTreeMap, left: &str, right: &str) -> &'a str { + let value = attr(attributes, left); + if value.is_empty() { + attr(attributes, right) + } else { + value + } +} + +fn tokens(attributes: &BTreeMap, key: &str) -> Result { + let value = attr(attributes, key).trim(); + if value.is_empty() { + return Ok(0); + } + match value.parse::() { + Ok(number) if (0..=u32::MAX as i128).contains(&number) => Ok(number as u32), + Ok(_) => Err(DecodeError::TokenCountOutOfRange), + Err(_) + if value + .trim_start_matches(['+', '-']) + .bytes() + .all(|byte| byte.is_ascii_digit()) => + { + Err(DecodeError::TokenCountOutOfRange) + } + Err(_) => Ok(0), + } +} + +fn usage_tokens(attributes: &BTreeMap) -> Result<(u32, u32), DecodeError> { + Ok(( + tokens(attributes, "gen_ai.usage.input_tokens")?, + tokens(attributes, "gen_ai.usage.output_tokens")?, + )) +} + +pub fn normalize( + scope_name: &str, + name: &str, + parent_span_id: &str, + attributes: &BTreeMap, +) -> Result { + let normalizers: [&dyn SpanNormalizer; 3] = [ + &LangSmithNormalizer, + &OpenInferenceNormalizer, + &GenAiNormalizer, + ]; + let normalizer = normalizers + .into_iter() + .find(|normalizer| normalizer.matches(scope_name, attributes)) + .expect("GenAI fallback always matches"); + Ok(Normalization { + span: normalizer.normalize(name, parent_span_id, attributes)?, + consumed_attributes: normalizer.consumed_attributes(attributes), + }) +} + +#[cfg(test)] +mod tests { + use std::collections::{BTreeMap, BTreeSet}; + + use rstest::rstest; + + use super::{NORMALIZED_FIELD_DEFINITIONS, ObservationType, normalize}; + + #[rstest] + #[case::langsmith("langsmith", [("langsmith.span.kind", "llm"), ("openinference.span.kind", "TOOL")], ObservationType::Llm)] + #[case::openinference("other", [("openinference.span.kind", "LLM"), ("gen_ai.operation.name", "execute_tool")], ObservationType::Llm)] + #[case::genai("other", [("gen_ai.operation.name", "execute_tool"), ("gen_ai.usage.input_tokens", "7")], ObservationType::Tool)] + fn convention_dispatch_preserves_precedence( + #[case] scope: &str, + #[case] attributes: [(&str, &str); 2], + #[case] expected: ObservationType, + ) { + let attributes = attributes + .into_iter() + .map(|(key, value)| (key.to_owned(), value.to_owned())) + .collect(); + let fields = normalize(scope, "step", "parent", &attributes) + .expect("valid tokens") + .span; + assert_eq!(fields.observation_type, expected); + if expected == ObservationType::Tool { + assert_eq!(fields.input_tokens, 7); + } + } + + #[rstest] + fn field_definitions_match_serialized_normalized_span() { + let fields = normalize("", "root", "", &BTreeMap::new()) + .expect("valid tokens") + .span; + let serialized = serde_json::to_value(fields).expect("serializable fields"); + let keys: BTreeSet<_> = serialized + .as_object() + .expect("field object") + .keys() + .map(String::as_str) + .collect(); + let mapped: BTreeSet<_> = NORMALIZED_FIELD_DEFINITIONS + .iter() + .map(|field| field.name) + .collect(); + assert_eq!(keys, mapped); + } + + #[rstest] + fn token_counts_accept_surrounding_whitespace() { + let attributes = + BTreeMap::from([("gen_ai.usage.input_tokens".to_owned(), " 7 ".to_owned())]); + let fields = normalize("", "root", "", &attributes) + .expect("valid tokens") + .span; + assert_eq!(fields.input_tokens, 7); + } + + #[rstest] + #[case::negative("-1")] + #[case::overflow("4294967296")] + fn token_counts_outside_storage_range_are_rejected(#[case] value: &str) { + let attributes = + BTreeMap::from([("gen_ai.usage.input_tokens".to_owned(), value.to_owned())]); + assert!(normalize("", "root", "", &attributes).is_err()); + } +} diff --git a/litellm-rust/crates/traces/src/normalize/openinference.rs b/litellm-rust/crates/traces/src/normalize/openinference.rs new file mode 100644 index 00000000000..b6156d8f24f --- /dev/null +++ b/litellm-rust/crates/traces/src/normalize/openinference.rs @@ -0,0 +1,53 @@ +use std::collections::BTreeMap; + +use super::{NormalizedSpan, ObservationType, SpanNormalizer, attr, tokens, usage_tokens}; +use crate::DecodeError; + +pub(super) struct OpenInferenceNormalizer; + +impl SpanNormalizer for OpenInferenceNormalizer { + fn matches(&self, _scope_name: &str, attributes: &BTreeMap) -> bool { + attributes.contains_key("openinference.span.kind") + } + + fn consumed_attributes(&self, _attributes: &BTreeMap) -> [&'static str; 2] { + ["input.value", "output.value"] + } + + fn normalize( + &self, + _name: &str, + parent_span_id: &str, + attributes: &BTreeMap, + ) -> Result { + let (usage_input, usage_output) = usage_tokens(attributes)?; + let observation_type = match attr(attributes, "openinference.span.kind") + .to_ascii_uppercase() + .as_str() + { + "AGENT" => ObservationType::Agent, + "LLM" => ObservationType::Llm, + "TOOL" => ObservationType::Tool, + _ if parent_span_id.is_empty() => ObservationType::Agent, + _ => ObservationType::Chain, + }; + Ok(NormalizedSpan { + observation_type, + agent_name: attr(attributes, "agent.name").to_owned(), + litellm_request_id: String::new(), + model: attr(attributes, "llm.model_name").to_owned(), + input_tokens: if attributes.contains_key("llm.token_count.prompt") { + tokens(attributes, "llm.token_count.prompt")? + } else { + usage_input + }, + output_tokens: if attributes.contains_key("llm.token_count.completion") { + tokens(attributes, "llm.token_count.completion")? + } else { + usage_output + }, + input: attr(attributes, "input.value").to_owned(), + output: attr(attributes, "output.value").to_owned(), + }) + } +} diff --git a/litellm-rust/crates/traces/src/otlp/mod.rs b/litellm-rust/crates/traces/src/otlp/mod.rs index fcc42082151..1beef48fe2c 100644 --- a/litellm-rust/crates/traces/src/otlp/mod.rs +++ b/litellm-rust/crates/traces/src/otlp/mod.rs @@ -6,7 +6,7 @@ mod wire; use serde::Serialize; use std::collections::BTreeMap; -use crate::{DecodeError, Shared}; +use crate::{DecodeError, NormalizedSpan, Shared}; #[derive(Serialize)] pub struct DecodedEvent { @@ -31,6 +31,8 @@ pub struct DecodedSpan { pub status_code: String, pub status_message: String, pub events: Vec, + pub normalized: NormalizedSpan, + pub consumed_attributes: [&'static str; 2], } pub fn decode_otlp( diff --git a/litellm-rust/crates/traces/src/otlp/span.rs b/litellm-rust/crates/traces/src/otlp/span.rs index fa993f71e3c..1c1e53e756c 100644 --- a/litellm-rust/crates/traces/src/otlp/span.rs +++ b/litellm-rust/crates/traces/src/otlp/span.rs @@ -10,7 +10,7 @@ use super::{ attributes::attributes, limits::{Budget, MAX_ATTRIBUTES, MAX_DECODED_SPAN_BYTES, MAX_EVENTS, MAX_SPANS}, }; -use crate::{DecodeError, Shared}; +use crate::{DecodeError, Shared, normalize::normalize}; pub(super) fn flatten(request: ExportTraceServiceRequest) -> Result, DecodeError> { let mut budget = Budget::new(MAX_DECODED_SPAN_BYTES); @@ -125,10 +125,26 @@ fn decoded_span( budget: &mut Budget, ) -> Result { let status = span.status.unwrap_or_default(); + let parent_span_id = hex_bytes(&span.parent_span_id); + let span_attributes = attributes(span.attributes, budget)?; + let normalization = normalize( + scope_name.as_ref(), + &span.name, + &parent_span_id, + &span_attributes, + )?; + let normalized = normalization.span; + budget.consume( + normalized.input.len() + + normalized.output.len() + + normalized.agent_name.len() + + normalized.litellm_request_id.len() + + normalized.model.len(), + )?; Ok(DecodedSpan { trace_id: hex_bytes(&span.trace_id), span_id: hex_bytes(&span.span_id), - parent_span_id: hex_bytes(&span.parent_span_id), + parent_span_id, trace_state: span.trace_state, name: span.name, kind: SpanKind::try_from(span.kind) @@ -143,7 +159,7 @@ fn decoded_span( })?, scope_name: budget.clone_shared(scope_name, String::len)?, scope_version: budget.clone_shared(scope_version, String::len)?, - attributes: attributes(span.attributes, budget)?, + attributes: span_attributes, start_ns: span.start_time_unix_nano, end_ns: span.end_time_unix_nano, status_code: StatusCode::try_from(status.code) @@ -162,5 +178,7 @@ fn decoded_span( }) }) .collect::, DecodeError>>()?, + normalized, + consumed_attributes: normalization.consumed_attributes, }) } diff --git a/litellm-rust/crates/traces/src/query.rs b/litellm-rust/crates/traces/src/query.rs new file mode 100644 index 00000000000..b2e8ddd8702 --- /dev/null +++ b/litellm-rust/crates/traces/src/query.rs @@ -0,0 +1,346 @@ +use std::collections::{BTreeMap, BTreeSet}; + +use futures_util::{ + StreamExt, + stream::{self, TryStreamExt}, +}; +use litellm_http::Client; +use serde::{Deserialize, Serialize}; +use serde_json::{Value, json}; + +use crate::{Connection, Error, NORMALIZED_FIELD_DEFINITIONS, execute_read}; + +mod guide; + +const SAMPLE_ROWS: usize = 200; +const MAX_FIELDS: usize = 200; +const MAX_DEPTH: usize = 16; +const METADATA_SQL: &str = "SELECT metadata FROM spend_logs FINAL \ + WHERE start_time >= now() - INTERVAL 7 DAY AND length(metadata) <= 8192 \ + LIMIT 201"; +const METADATA_SCOPE: &str = "Up to 200 unordered rows from the last 7 days, excluding metadata larger than 8192 bytes; up to 200 paths and 16 levels. Missing paths may exist outside this sample. Array indexes are 1-based and describe sampled positions, not a fixed schema"; +const ATTRIBUTE_SCOPE: &str = "Distinct keys from up to 200 unordered spans in the last 7 days; up to 200 keys per map. Missing keys may exist outside this sample"; + +#[derive(Deserialize)] +struct Rows { + data: Vec, +} + +#[derive(Deserialize)] +struct MetadataRow { + metadata: String, +} + +#[derive(Deserialize)] +struct AttributeRow { + key: String, +} + +#[derive(Clone, Debug, Eq, Ord, PartialEq, PartialOrd, Serialize)] +#[serde(untagged)] +enum PathPart { + Key(String), + Index(usize), +} + +#[derive(Serialize)] +struct MetadataField { + path: Vec, + types: BTreeSet<&'static str>, + expression: String, +} + +#[derive(Deserialize, Serialize)] +struct ColumnSchema { + name: String, + #[serde(rename = "type")] + kind: String, + #[serde(flatten)] + details: BTreeMap, +} + +#[derive(Serialize)] +struct TableSchema { + name: &'static str, + columns: Vec, +} + +#[derive(Serialize)] +struct MetadataCatalog { + table: &'static str, + column: &'static str, + fields: Vec, + sampled_rows: usize, + invalid_json_rows: usize, + truncated: bool, + sample_sql: &'static str, + scope: &'static str, + #[serde(skip_serializing_if = "Option::is_none")] + error: Option, +} + +#[derive(Serialize)] +struct AttributeField { + key: String, + #[serde(rename = "type")] + kind: &'static str, + expression: String, +} + +#[derive(Serialize)] +struct AttributeCatalog { + table: &'static str, + column: &'static str, + fields: Vec, + truncated: bool, + discovery_sql: String, + scope: &'static str, + #[serde(skip_serializing_if = "Option::is_none")] + error: Option, +} + +pub async fn query_sql( + client: &Client, + connection: &Connection, + sql: &str, +) -> Result { + execute_read(client, connection, sql, &BTreeMap::new()).await +} + +async fn rows( + client: &Client, + connection: &Connection, + sql: &str, +) -> Result, Error> { + let body = query_sql(client, connection, sql).await?; + serde_json::from_str::>(&body) + .map(|result| result.data) + .map_err(|_| Error::InvalidResponse) +} + +fn literal(value: &str) -> String { + format!("'{}'", value.replace('\\', "\\\\").replace('\'', "\\'")) +} + +fn metadata_expression(path: &[PathPart]) -> String { + let arguments = path + .iter() + .map(|part| match part { + PathPart::Key(key) => literal(key), + PathPart::Index(index) => index.to_string(), + }) + .collect::>() + .join(", "); + format!("JSONExtractRaw(metadata, {arguments})") +} + +fn discover( + value: &Value, + path: Vec, + fields: &mut BTreeMap, BTreeSet<&'static str>>, +) -> bool { + if path.len() > MAX_DEPTH || (fields.len() >= MAX_FIELDS && !fields.contains_key(&path)) { + return true; + } + if !path.is_empty() { + let kind = match value { + Value::Null => "null", + Value::Bool(_) => "boolean", + Value::Number(number) if number.is_i64() || number.is_u64() => "integer", + Value::Number(_) => "number", + Value::String(_) => "string", + Value::Array(_) => "array", + Value::Object(_) => "object", + }; + fields.entry(path.clone()).or_default().insert(kind); + } + match value { + Value::Object(object) => object.iter().fold(false, |limited, (key, value)| { + let child = path + .iter() + .cloned() + .chain([PathPart::Key(key.clone())]) + .collect(); + discover(value, child, fields) | limited + }), + Value::Array(array) => array + .iter() + .enumerate() + .fold(false, |limited, (index, value)| { + let child = path + .iter() + .cloned() + .chain([PathPart::Index(index + 1)]) + .collect(); + discover(value, child, fields) | limited + }), + _ => false, + } +} + +fn metadata_catalog(sample: &[MetadataRow]) -> MetadataCatalog { + let (fields, limited, invalid_rows) = sample.iter().take(SAMPLE_ROWS).fold( + (BTreeMap::new(), sample.len() > SAMPLE_ROWS, 0), + |(fields, limited, invalid_rows), row| match serde_json::from_str::(&row.metadata) { + Ok(value) => { + let mut fields = fields; + let limited = limited | discover(&value, Vec::new(), &mut fields); + (fields, limited, invalid_rows) + } + Err(_) => (fields, limited, invalid_rows + 1), + }, + ); + let fields: Vec<_> = fields + .into_iter() + .map(|(path, types)| MetadataField { + expression: metadata_expression(&path), + path, + types, + }) + .collect(); + MetadataCatalog { + table: "spend_logs", + column: "metadata", + fields, + sampled_rows: sample.len().min(SAMPLE_ROWS), + invalid_json_rows: invalid_rows, + truncated: limited, + sample_sql: METADATA_SQL, + error: None, + scope: METADATA_SCOPE, + } +} + +pub async fn query_help(client: &Client, connection: &Connection) -> Result { + let tables = stream::iter(["otel_traces", "agent_traces_by_key", "spend_logs"]) + .then(|table| async move { + Ok::<_, Error>(TableSchema { + name: table, + columns: rows::( + client, + connection, + &format!("DESCRIBE TABLE {table}"), + ) + .await?, + }) + }) + .try_collect::>() + .await?; + let metadata = match rows::(client, connection, METADATA_SQL).await { + Ok(sample) => metadata_catalog(&sample), + Err(error) => MetadataCatalog { + error: Some(error.to_string()), + truncated: true, + ..metadata_catalog(&[]) + }, + }; + let attributes = stream::iter(["SpanAttributes", "ResourceAttributes"]) + .then(|column| async move { + let sql = format!( + "SELECT DISTINCT arrayJoin(mapKeys({column})) AS key FROM \ + (SELECT {column} FROM otel_traces WHERE Timestamp >= now() - INTERVAL 7 DAY \ + LIMIT 200) ORDER BY key LIMIT 201" + ); + let (keys, error) = match rows::(client, connection, &sql).await { + Ok(keys) => (keys, None), + Err(error) => (Vec::new(), Some(error.to_string())), + }; + let fields = keys + .iter() + .take(MAX_FIELDS) + .map(|row| AttributeField { + key: row.key.clone(), + kind: "String", + expression: format!("{column}[{}]", literal(&row.key)), + }) + .collect(); + AttributeCatalog { + table: "otel_traces", + column, + fields, + truncated: error.is_some() || keys.len() > MAX_FIELDS, + discovery_sql: sql, + scope: ATTRIBUTE_SCOPE, + error, + } + }) + .collect::>() + .await; + let guide = guide::QueryGuide { + tables: &tables, + normalized_fields: &NORMALIZED_FIELD_DEFINITIONS, + metadata: &metadata, + attributes: &attributes, + }; + Ok(json!({ + "dialect": "ClickHouse SQL", + "access": "Authenticated team scope enforced by ClickHouse row policies; proxy admins can read all teams, while project-bound and teamless keys can read only their own rows", + "response": "ClickHouse JSON envelope: meta, data, rows, statistics; 64-bit integers may be strings", + "tables": tables, + "normalized_fields": NORMALIZED_FIELD_DEFINITIONS.iter().map(|field| json!({ + "table": "otel_traces", "name": field.name, "column": field.clickhouse_column, + "type": field.clickhouse_type, "meaning": field.meaning + })).collect::>(), + "metadata": metadata, + "attributes": attributes, + "relationships": [{ + "left": "otel_traces.LiteLLMRequestId", "right": "spend_logs.response_id", + "additional_predicates": "otel_traces.TeamId = spend_logs.team_id AND otel_traces.ApiKeyHash = spend_logs.api_key", + "meaning": "The normalized ID is the response ID, not request_id. Cached requests can share response_id; joins may return multiple spend rows" + }], + "examples": guide.examples()?, + "gotchas": guide.gotchas()?, + "guide": guide::render(&guide)?, + }).to_string()) +} + +#[cfg(test)] +mod tests { + use super::*; + use rstest::rstest; + + #[rstest] + fn metadata_discovery_preserves_mixed_types_and_reports_invalid_rows() { + let sample = [ + MetadataRow { + metadata: r#"{"x": 1}"#.into(), + }, + MetadataRow { + metadata: r#"{"x": "one"}"#.into(), + }, + MetadataRow { + metadata: "invalid".into(), + }, + ]; + let catalog = json!(metadata_catalog(&sample)); + assert_eq!( + catalog["fields"], + json!([{ + "path": ["x"], "types": ["integer", "string"], "expression": "JSONExtractRaw(metadata, 'x')" + }]) + ); + assert_eq!(catalog["invalid_json_rows"], 1); + assert_eq!(catalog["sampled_rows"], sample.len()); + } + + #[rstest] + #[case::rows(SAMPLE_ROWS + 1, 1)] + #[case::paths(1, MAX_FIELDS + 1)] + fn metadata_discovery_reports_truncation(#[case] row_count: usize, #[case] field_count: usize) { + let metadata: BTreeMap<_, _> = (0..field_count) + .map(|index| (format!("field{index}"), index)) + .collect(); + let sample: Vec<_> = (0..row_count) + .map(|_| MetadataRow { + metadata: json!(metadata).to_string(), + }) + .collect(); + let catalog = json!(metadata_catalog(&sample)); + assert_eq!(catalog["truncated"], true); + assert_eq!(catalog["sampled_rows"], row_count.min(SAMPLE_ROWS)); + assert_eq!( + catalog["fields"].as_array().unwrap().len(), + field_count.min(MAX_FIELDS) + ); + } +} diff --git a/litellm-rust/crates/traces/src/query/guide.rs b/litellm-rust/crates/traces/src/query/guide.rs new file mode 100644 index 00000000000..3bf7336648d --- /dev/null +++ b/litellm-rust/crates/traces/src/query/guide.rs @@ -0,0 +1,89 @@ +use askama::Template; +use serde::Serialize; + +use super::{AttributeCatalog, MetadataCatalog, TableSchema}; +use crate::{Error, NormalizedFieldDefinition}; + +#[derive(Template)] +#[template(path = "query_help.jinja", escape = "none", blocks = [ + "recent_spans_name", + "recent_spans_sql", + "custom_metadata_name", + "custom_metadata_sql", + "nested_metadata_name", + "nested_metadata_sql", + "correlated_calls_name", + "correlated_calls_sql", + "discover_keys_name", + "discover_keys_sql", + "time_window", + "reader_limits", + "reader_profile", + "output_format", + "json_values", + "map_values", + "literal_keys", + "time_units", + "spend_totals", + "trace_rollups", + "sampling", +])] +pub(super) struct QueryGuide<'a> { + pub tables: &'a [TableSchema], + pub normalized_fields: &'a [NormalizedFieldDefinition], + pub metadata: &'a MetadataCatalog, + pub attributes: &'a [AttributeCatalog], +} + +#[derive(Serialize)] +pub(super) struct Example { + name: String, + sql: String, +} + +impl QueryGuide<'_> { + pub fn examples(&self) -> Result<[Example; 5], Error> { + Ok([ + Example { + name: render(&self.as_recent_spans_name())?, + sql: render(&self.as_recent_spans_sql())?, + }, + Example { + name: render(&self.as_custom_metadata_name())?, + sql: render(&self.as_custom_metadata_sql())?, + }, + Example { + name: render(&self.as_nested_metadata_name())?, + sql: render(&self.as_nested_metadata_sql())?, + }, + Example { + name: render(&self.as_correlated_calls_name())?, + sql: render(&self.as_correlated_calls_sql())?, + }, + Example { + name: render(&self.as_discover_keys_name())?, + sql: render(&self.as_discover_keys_sql())?, + }, + ]) + } + + pub fn gotchas(&self) -> Result<[String; 11], Error> { + Ok([ + render(&self.as_time_window())?, + render(&self.as_reader_limits())?, + render(&self.as_reader_profile())?, + render(&self.as_output_format())?, + render(&self.as_json_values())?, + render(&self.as_map_values())?, + render(&self.as_literal_keys())?, + render(&self.as_time_units())?, + render(&self.as_spend_totals())?, + render(&self.as_trace_rollups())?, + render(&self.as_sampling())?, + ]) + } +} + +pub(super) fn render(template: &impl Template) -> Result { + template.render().map_err(|_| Error::InvalidResponse) +} diff --git a/litellm-rust/crates/traces/src/query_access.rs b/litellm-rust/crates/traces/src/query_access.rs new file mode 100644 index 00000000000..881b609177d --- /dev/null +++ b/litellm-rust/crates/traces/src/query_access.rs @@ -0,0 +1,200 @@ +use std::{sync::Arc, time::Duration}; + +use hmac::{Hmac, Mac}; +use litellm_http::Client; +use moka::future::Cache; +use serde::{Deserialize, Serialize}; +use sha2::{Digest, Sha256}; +use tokio::sync::{OwnedSemaphorePermit, Semaphore}; + +use crate::{Connection, QueryAccessError}; + +const TABLES: [&str; 3] = ["otel_traces", "agent_traces_by_key", "spend_logs"]; + +#[derive(Clone, Debug, Deserialize, Serialize)] +#[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)] +pub enum QueryScope { + Admin, + Team { + team_id: String, + }, + Key { + team_id: String, + api_key_hash: String, + }, +} + +impl QueryScope { + fn validate(&self) -> Result<(), QueryAccessError> { + match self { + Self::Admin => Ok(()), + Self::Team { team_id } if !team_id.is_empty() => Ok(()), + Self::Key { api_key_hash, .. } if !api_key_hash.is_empty() => Ok(()), + _ => Err(QueryAccessError::InvalidScope), + } + } + + fn predicate(&self, table: &str) -> String { + let (team, key) = if table == "spend_logs" { + ("team_id", "api_key") + } else { + ("TeamId", "ApiKeyHash") + }; + match self { + Self::Admin => "1".to_owned(), + Self::Team { team_id } => format!("{team} = {}", literal(team_id)), + Self::Key { + team_id, + api_key_hash, + } => format!( + "{team} = {} AND {key} = {}", + literal(team_id), + literal(api_key_hash) + ), + } + } +} + +#[derive(Clone)] +pub struct QueryReaders { + writer: Connection, + database: String, + readers: Cache, + slots: Arc, +} + +impl QueryReaders { + pub fn new(writer: Connection, database: String) -> Self { + Self { + writer, + database, + readers: Cache::builder().max_capacity(1024).build(), + slots: Arc::new(Semaphore::new(8)), + } + } + + pub fn acquire(&self) -> Result { + self.slots + .clone() + .try_acquire_owned() + .map_err(|_| QueryAccessError::Busy) + } + + pub async fn connection( + &self, + client: &Client, + scope: &QueryScope, + secret: &str, + ) -> Result { + scope.validate()?; + if secret.is_empty() { + return Err(QueryAccessError::MissingSecret); + } + let identity = serde_json::to_vec(&("litellm_trace_reader_v1", &self.database, scope)) + .map_err(|_| QueryAccessError::InvalidScope)?; + let user = format!("litellm_traces_{:x}", Sha256::digest(&identity)); + let password = credential(secret, b"password", &identity)?; + self.readers + .try_get_with( + user.clone(), + self.provision(client, scope, &user, &password), + ) + .await + .map_err(QueryAccessError::Cached) + } + + async fn provision( + &self, + client: &Client, + scope: &QueryScope, + user: &str, + password: &str, + ) -> Result { + let database = &self.database; + if database.is_empty() + || !database + .bytes() + .all(|c| c.is_ascii_alphanumeric() || c == b'_') + { + return Err(QueryAccessError::InvalidScope); + } + let password_hash = format!("{:x}", Sha256::digest(password)); + self.execute( + client, + format!( + "CREATE USER IF NOT EXISTS {user} IDENTIFIED WITH sha256_hash BY '{password_hash}' \ + SETTINGS readonly = 1 CONST, max_execution_time = 10 CONST, \ + max_result_rows = 1000 CONST, max_result_bytes = 4194304 CONST, \ + result_overflow_mode = 'throw' CONST, max_memory_usage = 268435456 CONST, \ + max_threads = 2 CONST, max_concurrent_queries_for_user = 8 CONST" + ), + ) + .await?; + self.execute( + client, + format!("ALTER USER {user} IDENTIFIED WITH sha256_hash BY '{password_hash}'"), + ) + .await?; + for table in TABLES { + let predicate = scope.predicate(table); + self.execute( + client, + format!( + "CREATE ROW POLICY IF NOT EXISTS {user}_allow ON `{database}`.{table} \ + USING 1 TO {user}" + ), + ) + .await?; + self.execute( + client, + format!( + "CREATE ROW POLICY IF NOT EXISTS {user}_scope ON `{database}`.{table} \ + AS RESTRICTIVE USING {predicate} TO {user}" + ), + ) + .await?; + } + for table in TABLES { + self.execute( + client, + format!("GRANT SELECT ON `{database}`.{table} TO {user}"), + ) + .await?; + } + Connection::configured( + &self.writer.url()[..url::Position::AfterPath], + database, + user, + password, + ) + .map_err(QueryAccessError::Storage) + } + + async fn execute(&self, client: &Client, sql: String) -> Result<(), QueryAccessError> { + let response = client + .post(self.writer.url().clone()) + .timeout(Duration::from_secs(15)) + .body(sql) + .send() + .await + .map_err(|_| QueryAccessError::ProvisionTransport)?; + if !response.status().is_success() { + return Err(QueryAccessError::ProvisionFailed( + response.status().as_u16(), + )); + } + Ok(()) + } +} + +fn credential(secret: &str, purpose: &[u8], identity: &[u8]) -> Result { + let mut mac = Hmac::::new_from_slice(secret.as_bytes()) + .map_err(|_| QueryAccessError::MissingSecret)?; + mac.update(purpose); + mac.update(identity); + Ok(format!("{:x}", mac.finalize().into_bytes())) +} + +fn literal(value: &str) -> String { + format!("'{}'", value.replace('\\', "\\\\").replace('\'', "\\'")) +} diff --git a/litellm-rust/crates/traces/src/schema.rs b/litellm-rust/crates/traces/src/schema.rs index 4943f00f7c9..cf11c564931 100644 --- a/litellm-rust/crates/traces/src/schema.rs +++ b/litellm-rust/crates/traces/src/schema.rs @@ -1,4 +1,5 @@ use litellm_http::Client; +use litellm_migrate::Migration; use std::time::Duration; use crate::Connection; @@ -6,17 +7,7 @@ use crate::Error; const SCHEMA_REQUEST_TIMEOUT: Duration = Duration::from_secs(30); -const MIGRATIONS: [&str; 9] = [ - include_str!("../migrations/0001_otel_traces.sql"), - include_str!("../migrations/0002_agent_traces.sql"), - include_str!("../migrations/0003_agent_traces_mv.sql"), - include_str!("../migrations/0004_spend_logs.sql"), - include_str!("../migrations/0005_otel_traces_ttl.sql"), - include_str!("../migrations/0006_agent_traces_ttl.sql"), - include_str!("../migrations/0007_spend_logs_ttl.sql"), - include_str!("../migrations/0008_trace_received.sql"), - include_str!("../migrations/0009_spend_received.sql"), -]; +const MIGRATIONS: &[Migration] = litellm_migrate::migrate!("migrations"); pub fn schema_statements( database: &str, @@ -35,8 +26,10 @@ pub fn schema_statements( let database = format!("`{database}`"); Ok( std::iter::once(format!("CREATE DATABASE IF NOT EXISTS {database}")) - .chain(MIGRATIONS.iter().map(|sql| { - sql.replace("{database}", &database) + .chain(MIGRATIONS.iter().map(|migration| { + migration + .sql + .replace("{database}", &database) .replace("{trace_retention_days}", &trace_retention_days.to_string()) .replace( "{spend_log_retention_days}", diff --git a/litellm-rust/crates/traces/src/sql.rs b/litellm-rust/crates/traces/src/sql.rs index 36d6e3b4521..51632066fdd 100644 --- a/litellm-rust/crates/traces/src/sql.rs +++ b/litellm-rust/crates/traces/src/sql.rs @@ -37,6 +37,8 @@ impl ReadQuery { #[derive(Clone, Copy)] pub enum LensQuery { + Availability, + Agents, Sample, Content, Evidence, @@ -45,6 +47,8 @@ pub enum LensQuery { impl LensQuery { pub fn parse(name: &str) -> Result { match name { + "availability" => Ok(Self::Availability), + "agents" => Ok(Self::Agents), "sample" => Ok(Self::Sample), "content" => Ok(Self::Content), "evidence" => Ok(Self::Evidence), @@ -53,6 +57,8 @@ impl LensQuery { } pub fn sql(self) -> &'static str { match self { + Self::Availability => include_str!("../query/lens_availability.sql"), + Self::Agents => include_str!("../query/lens_agents.sql"), Self::Sample => include_str!("../query/lens_sample.sql"), Self::Content => include_str!("../query/lens_content.sql"), Self::Evidence => include_str!("../query/lens_evidence.sql"), diff --git a/litellm-rust/crates/traces/templates/query_help.jinja b/litellm-rust/crates/traces/templates/query_help.jinja new file mode 100644 index 00000000000..4bded74588c --- /dev/null +++ b/litellm-rust/crates/traces/templates/query_help.jinja @@ -0,0 +1,65 @@ +Trace SQL query guide + +Live ClickHouse schema +{% for table in tables %} +{{ table.name }} +{% for column in table.columns %}{{ column.name }}: {{ column.kind }} +{% endfor %}{% endfor %} +Normalized span fields +{% for field in normalized_fields %}{{ field.name }}: otel_traces.{{ field.clickhouse_column }} ({{ field.clickhouse_type }}) +{{ field.meaning }} +{% endfor %} +Observed LLM call metadata +{{ metadata.scope }} +Sampled rows: {{ metadata.sampled_rows }}; invalid JSON rows: {{ metadata.invalid_json_rows }}; truncated: {{ metadata.truncated }} +{% if let Some(error) = metadata.error %}Metadata discovery unavailable: {{ error }} +{% else if metadata.fields.is_empty() %}No metadata paths found in the sampled rows +{% else %}{% for field in metadata.fields %}{{ field.expression }}: {% for kind in field.types %}{{ kind }} {% endfor %} +{% endfor %}{% endif %} +Observed span and resource attributes +{% for catalog in attributes %}{{ catalog.table }}.{{ catalog.column }} +{{ catalog.scope }} +{% if let Some(error) = catalog.error %}Attribute discovery unavailable: {{ error }} +{% else if catalog.fields.is_empty() %}No attribute keys found in the sampled spans +{% else %}{% for field in catalog.fields %}{{ field.expression }}: {{ field.kind }} +{% endfor %}{% endif %}{% endfor %} +Examples + +{% block recent_spans_name %}Recent normalized LLM spans{% endblock %} +{% block recent_spans_sql %}SELECT TraceId, SpanId, Model, InputTokens, OutputTokens, Duration / 1000000 AS duration_ms FROM otel_traces WHERE Timestamp >= now() - INTERVAL 1 DAY AND ObservationType = 'llm' ORDER BY Timestamp DESC LIMIT 100{% endblock %} + +{% block custom_metadata_name %}Find calls by custom metadata{% endblock %} +{% block custom_metadata_sql %}SELECT request_id, response_id, model, spend, JSONExtractString(metadata, 'project') AS project FROM spend_logs FINAL WHERE start_time >= now() - INTERVAL 1 DAY AND JSONHas(metadata, 'project') AND JSONExtractString(metadata, 'project') = 'example' ORDER BY start_time DESC LIMIT 100{% endblock %} + +{% block nested_metadata_name %}Nested metadata with unknown types{% endblock %} +{% block nested_metadata_sql %}SELECT request_id, JSONType(metadata, 'labels', 'priority') AS type, JSONExtractRaw(metadata, 'labels', 'priority') AS value FROM spend_logs FINAL WHERE start_time >= now() - INTERVAL 1 DAY AND JSONHas(metadata, 'labels', 'priority') LIMIT 100{% endblock %} + +{% block correlated_calls_name %}Traces correlated with LLM call metadata{% endblock %} +{% block correlated_calls_sql %}SELECT t.TraceId, t.SpanId, s.request_id, s.spend, s.metadata FROM otel_traces AS t INNER JOIN (SELECT * FROM spend_logs FINAL WHERE start_time >= now() - INTERVAL 1 DAY) AS s ON t.LiteLLMRequestId = s.response_id AND t.TeamId = s.team_id AND t.ApiKeyHash = s.api_key WHERE t.Timestamp >= now() - INTERVAL 1 DAY AND t.LiteLLMRequestId != '' AND JSONExtractString(s.metadata, 'project') = 'example' LIMIT 100{% endblock %} + +{% block discover_keys_name %}Discover metadata keys over a different window{% endblock %} +{% block discover_keys_sql %}SELECT DISTINCT arrayJoin(JSONExtractKeys(metadata)) AS key FROM spend_logs FINAL WHERE start_time >= now() - INTERVAL 30 DAY ORDER BY key LIMIT 200{% endblock %} + +Gotchas + +{% block time_window %}Always bound Timestamp or start_time and use LIMIT; add TeamId/ApiKeyHash or team_id/api_key filters when investigating one tenant{% endblock %} + +{% block reader_limits %}The reader enforces 1000 result rows, 4 MiB response bytes, 256 MiB memory and a 10 second query limit; exceeding limits fails instead of returning partial results{% endblock %} + +{% block reader_profile %}LiteLLM provisions SELECT-only readers from the configured ClickHouse connection and enforces authenticated team scope through row policies. Project-bound and teamless keys see only their own rows. Provisioning requires CREATE USER, ALTER USER, CREATE ROW POLICY, and GRANT SELECT permissions{% endblock %} + +{% block output_format %}Do not add FORMAT clauses; the endpoint requires ClickHouse JSON output{% endblock %} + +{% block json_values %}metadata is a JSON-encoded String; use JSONHas before typed extraction to distinguish missing values from empty strings, zero and false{% endblock %} + +{% block map_values %}SpanAttributes and ResourceAttributes are Map(String, String); missing map keys return an empty string, so use mapContains for existence checks{% endblock %} + +{% block literal_keys %}Use the discovered path components as separate JSONExtract arguments; a dot inside a key is literal, not a path separator{% endblock %} + +{% block time_units %}Duration is nanoseconds; Timestamp has nanosecond precision, spend start_time has millisecond precision{% endblock %} + +{% block spend_totals %}Use spend_logs FINAL to collapse replacement rows before totals. Shared response IDs and multiple spans can multiply costs in joins; aggregate spend separately{% endblock %} + +{% block trace_rollups %}agent_traces_by_key uses SimpleAggregateFunction columns; group by TeamId, ApiKeyHash and TraceId, using min(StartTs), max(EndTs), sum(SpanCount) and groupUniqArrayArray(Models). Do not use Merge combinators{% endblock %} + +{% block sampling %}Discovery is sampled, contains no metadata values, and is not an exhaustive schema. Edit the supplied discovery SQL for older data or nested JSONExtractKeys(metadata, 'parent'){% endblock %} diff --git a/litellm-rust/crates/traces/tests/migrations.rs b/litellm-rust/crates/traces/tests/migrations.rs index 01e8982423f..d93083b052c 100644 --- a/litellm-rust/crates/traces/tests/migrations.rs +++ b/litellm-rust/crates/traces/tests/migrations.rs @@ -2,8 +2,8 @@ use std::{collections::BTreeMap, time::Duration}; use litellm_http::Client; use litellm_traces::{ - Connection, Error, InsertTable, Parameter, ReadQuery, encode_rows, ensure_schema, - execute_named_read, execute_read, schema_statements, + Connection, Error, InsertTable, NORMALIZED_FIELD_DEFINITIONS, Parameter, ReadQuery, + encode_rows, ensure_schema, execute_named_read, execute_read, schema_statements, }; use rstest::{fixture, rstest}; use testcontainers_modules::{ @@ -210,6 +210,43 @@ async fn schema_supports_span_rollups_and_spend_joins( Ok(()) } +#[rstest] +#[tokio::test] +async fn normalized_fields_match_clickhouse_catalog( + #[future(awt)] database: TestResult, +) -> TestResult { + let database = database?; + ensure_schema( + &database.client, + &Connection::writer(&database.url)?, + "trace_test", + 7, + 14, + ) + .await?; + let catalog = read_json(&database, "SELECT name, type FROM system.columns WHERE database = 'trace_test' AND table = 'otel_traces'").await?; + let columns: BTreeMap<&str, &str> = catalog["data"] + .as_array() + .expect("catalog rows") + .iter() + .map(|row| { + ( + row["name"].as_str().expect("column name"), + row["type"].as_str().expect("column type"), + ) + }) + .collect(); + for field in NORMALIZED_FIELD_DEFINITIONS { + assert_eq!( + columns.get(field.clickhouse_column).copied(), + Some(field.clickhouse_type), + "{}", + field.name + ); + } + Ok(()) +} + #[rstest] #[tokio::test] async fn insert_rejects_unknown_columns_even_if_url_requests_skipping_them( @@ -433,6 +470,20 @@ async fn retention_changes_materialize_existing_rows_and_remain_idempotent( let database = database?; let writer = Connection::writer(&database.url)?; ensure_schema(&database.client, &writer, "trace_test", 30, 30).await?; + let tables = read_json( + &database, + "SELECT name FROM system.tables WHERE database = 'trace_test' \ + AND match(engine_full, 'materialize_ttl_recalculate_only = 1') ORDER BY name", + ) + .await?; + assert_eq!( + tables["data"], + serde_json::json!([ + {"name": "agent_traces_by_key"}, + {"name": "otel_traces"}, + {"name": "spend_logs"} + ]) + ); let old_time = time::OffsetDateTime::now_utc() - time::Duration::days(20); let old_timestamp_ns = old_time.unix_timestamp_nanos() as i64; let old_timestamp_ms = old_timestamp_ns / 1_000_000; @@ -551,6 +602,7 @@ async fn lens_filters_reads_and_evidence_keep_reused_trace_ids_separate( "end".into(), Parameter::Integer(timestamp / 1_000_000 + 1000), ), + ("agent_name".into(), Parameter::Text(String::new())), ("service".into(), Parameter::Text("review".into())), ( "filter_keys".into(), @@ -652,6 +704,7 @@ async fn lens_request_sample_does_not_trust_caller_tags( ("key_hash".into(), Parameter::Text(String::new())), ("start".into(), Parameter::Integer(timestamp - 1000)), ("end".into(), Parameter::Integer(timestamp + 60000)), + ("agent_name".into(), Parameter::Text(String::new())), ("service".into(), Parameter::Text(String::new())), ("filter_keys".into(), Parameter::Strings(vec![])), ("filter_values".into(), Parameter::Strings(vec![])), @@ -719,6 +772,7 @@ async fn lens_selection_pages_without_losing_or_repeating_runs( ("key_hash".into(), Parameter::Text(String::new())), ("start".into(), Parameter::Integer(0)), ("end".into(), Parameter::Integer(end)), + ("agent_name".into(), Parameter::Text(String::new())), ("service".into(), Parameter::Text(String::new())), ("filter_keys".into(), Parameter::Strings(vec![])), ("filter_values".into(), Parameter::Strings(vec![])), @@ -980,3 +1034,404 @@ async fn duplicate_span_preview_matches_diagnostic( assert_eq!(diagnostic["data"][0]["message"], message); Ok(()) } + +#[rstest] +fn schema_includes_every_migration_file() -> TestResult { + let files = std::fs::read_dir(concat!(env!("CARGO_MANIFEST_DIR"), "/migrations"))? + .filter_map(|entry| entry.ok()) + .filter(|entry| entry.path().extension().is_some_and(|ext| ext == "sql")) + .count(); + assert_eq!(schema_statements("trace_test", 7, 14)?.len(), 1 + files); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn lens_agent_discovery_and_selection_preserve_scope( + #[future] database: TestResult, +) -> TestResult { + use litellm_traces::LensQuery; + let database = database.await?; + let writer = Connection::writer(&database.url)?; + ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64; + for (team, key, trace, agent, span, parent) in [ + ("alpha", "one", "research", "research_agent", "root", ""), + ("alpha", "one", "research", "", "tool", "root"), + ("alpha", "one", "support", "support_agent", "root", ""), + ("alpha", "two", "hidden-key", "private_agent", "root", ""), + ("beta", "one", "hidden-team", "other_agent", "root", ""), + ] { + insert_rows( + &database, + "otel_traces", + vec![serde_json::from_value(serde_json::json!({ + "Timestamp": timestamp, "TraceId": trace, "SpanId": span, "ParentSpanId": parent, + "ServiceName": "shared-app", "SpanName": "run", "Input": "test", + "SpanAttributes": {"gen_ai.agent.name": agent}, + "ResourceAttributes": {"litellm.team_id": team, "litellm.api_key_hash": key} + }))?], + ) + .await?; + } + let connection = Connection::configured(&database.url, "trace_test", "default", "")?; + let scope_parameters = BTreeMap::from([ + ("all_teams".into(), Parameter::Integer(0)), + ("team".into(), Parameter::Text("alpha".into())), + ("key_hash".into(), Parameter::Text("one".into())), + ]); + let agents: serde_json::Value = serde_json::from_str( + &execute_read( + &database.client, + &connection, + LensQuery::Agents.sql(), + &scope_parameters, + ) + .await?, + )?; + assert_eq!( + agents["data"], + serde_json::json!([ + {"agent_name": "research_agent"}, {"agent_name": "support_agent"} + ]) + ); + let parameters = scope_parameters + .into_iter() + .chain([ + ("source".into(), Parameter::Text("traces".into())), + ( + "start".into(), + Parameter::Integer(timestamp / 1_000_000 - 1000), + ), + ( + "end".into(), + Parameter::Integer(timestamp / 1_000_000 + 1000), + ), + ("service".into(), Parameter::Text("shared-app".into())), + ( + "agent_name".into(), + Parameter::Text("research_agent".into()), + ), + ("filter_keys".into(), Parameter::Strings(vec![])), + ("filter_values".into(), Parameter::Strings(vec![])), + ("limit".into(), Parameter::Integer(100)), + ("offset".into(), Parameter::Integer(0)), + ("after".into(), Parameter::Text(String::new())), + ("sample_percent".into(), Parameter::Text("100".into())), + ("sample_cap".into(), Parameter::Integer(0)), + ("preview".into(), Parameter::Integer(1)), + ("selected_team".into(), Parameter::Text(String::new())), + ("execution_ids".into(), Parameter::Strings(vec![])), + ]) + .collect::>(); + let sample: serde_json::Value = serde_json::from_str( + &execute_read( + &database.client, + &connection, + LensQuery::Sample.sql(), + ¶meters, + ) + .await?, + )?; + assert_eq!(sample["data"].as_array().expect("rows").len(), 1); + assert_eq!(sample["data"][0]["trace_id"], "research"); + assert_eq!(sample["data"][0]["span_count"], 2); + let available: serde_json::Value = serde_json::from_str( + &execute_read( + &database.client, + &connection, + LensQuery::Availability.sql(), + ¶meters, + ) + .await?, + )?; + assert_eq!(available["data"][0]["traces"], 1); + assert_eq!(available["data"][0]["requests"], 0); + Ok(()) +} + +#[rstest] +#[case::empty(false)] +#[case::custom_metadata(true)] +#[tokio::test] +async fn query_help_discovers_live_schema_and_runs_its_examples( + #[future(awt)] database: TestResult, + #[case] populated: bool, +) -> TestResult { + let database = database?; + let writer = Connection::writer(&database.url)?; + ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + execute_write(&database, "CREATE USER help_reader").await?; + for table in ["otel_traces", "agent_traces_by_key", "spend_logs"] { + execute_write( + &database, + &format!("GRANT SELECT ON trace_test.{table} TO help_reader"), + ) + .await?; + } + let reader = Connection::configured(&database.url, "trace_test", "help_reader", "")?; + let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64; + if populated { + execute_write(&database, "SYSTEM STOP MERGES trace_test.spend_logs").await?; + insert_rows( + &database, + "spend_logs", + vec![serde_json::from_value(serde_json::json!({ + "request_id": "request-1", "response_id": "response-1", "team_id": "team-1", + "api_key": "key-1", "metadata": r#"{"obsolete":true,"labels":{"priority":"old"}}"#, + "start_time": timestamp / 1_000_000, "end_time": timestamp / 1_000_000 + }))?], + ) + .await?; + let metadata = serde_json::json!({ + "project": "example", "labels": {"priority": 3, "enabled": true}, + "dotted.key": "literal", "quote'\\key": null, "items": [{"name": "first"}], + "&{{key}}": {"nested.key": true} + }); + insert_rows( + &database, + "spend_logs", + vec![serde_json::from_value(serde_json::json!({ + "request_id": "request-1", "response_id": "response-1", "team_id": "team-1", + "api_key": "key-1", "metadata": metadata.to_string(), "spend": 0.25, + "start_time": timestamp / 1_000_000, "end_time": timestamp / 1_000_000 + 100 + }))?], + ) + .await?; + insert_rows( + &database, + "otel_traces", + vec![serde_json::from_value(serde_json::json!({ + "Timestamp": timestamp, "TraceId": "trace-1", "SpanId": "span-1", + "TeamId": "team-1", "ApiKeyHash": "key-1", "ObservationType": "llm", + "LiteLLMRequestId": "response-1", "SpanAttributes": {"custom.tag": "value"}, + "ResourceAttributes": {"custom.resource": "value"} + }))?], + ) + .await?; + execute_write( + &database, + "ALTER TABLE trace_test.otel_traces ADD COLUMN CustomColumn String", + ) + .await?; + } + let help: serde_json::Value = + serde_json::from_str(&litellm_traces::query_help(&database.client, &reader).await?)?; + let keys: std::collections::BTreeSet<_> = help + .as_object() + .ok_or("missing help object")? + .keys() + .map(String::as_str) + .collect(); + assert_eq!( + keys, + std::collections::BTreeSet::from([ + "access", + "attributes", + "dialect", + "examples", + "gotchas", + "guide", + "metadata", + "normalized_fields", + "relationships", + "response", + "tables", + ]) + ); + let guide = help["guide"].as_str().ok_or("missing rendered guide")?; + assert!(guide.starts_with("Trace SQL query guide")); + for table in ["otel_traces", "agent_traces_by_key", "spend_logs"] { + let described = read_json(&database, &format!("DESCRIBE TABLE {table}")).await?; + let schema = help["tables"] + .as_array() + .ok_or("missing tables")? + .iter() + .find(|schema| schema["name"] == table) + .ok_or("missing table")?; + assert_eq!(schema["columns"], described["data"]); + for column in described["data"].as_array().ok_or("missing live columns")? { + assert!(guide.contains(&format!( + "{}: {}", + column["name"].as_str().ok_or("column name")?, + column["type"].as_str().ok_or("column type")? + ))); + } + } + for gotcha in help["gotchas"].as_array().ok_or("missing gotchas")? { + assert!(guide.contains(gotcha.as_str().ok_or("gotcha text")?)); + } + let tables = help["tables"].as_array().ok_or("missing tables")?; + assert_eq!(tables.len(), 3); + let columns = tables[0]["columns"].as_array().ok_or("missing columns")?; + for field in NORMALIZED_FIELD_DEFINITIONS { + assert!( + columns + .iter() + .any(|column| column["name"] == field.clickhouse_column + && column["type"] == field.clickhouse_type) + ); + assert!( + help["normalized_fields"] + .as_array() + .ok_or("missing mappings")? + .iter() + .any(|mapped| { + mapped["name"] == field.name && mapped["column"] == field.clickhouse_column + }) + ); + } + let fields = help["metadata"]["fields"] + .as_array() + .ok_or("missing metadata fields")?; + assert_eq!(fields.is_empty(), !populated); + assert_eq!(help["metadata"]["truncated"], false); + assert!(guide.contains(help["metadata"]["scope"].as_str().ok_or("missing scope")?)); + assert_eq!( + guide.contains("No metadata paths found in the sampled rows"), + !populated + ); + if populated { + let versions = read_json(&database, "SELECT count() AS count FROM spend_logs").await?; + assert_eq!(versions["data"][0]["count"], 2); + assert_eq!(help["metadata"]["sampled_rows"], 1); + assert!( + !fields + .iter() + .any(|field| field["path"] == serde_json::json!(["obsolete"])) + ); + assert!( + columns + .iter() + .any(|column| column["name"] == "CustomColumn") + ); + assert!(fields.iter().any(|field| field["path"] + == serde_json::json!(["labels", "priority"]) + && field["types"] == serde_json::json!(["integer"]))); + assert!( + fields + .iter() + .any(|field| field["path"] == serde_json::json!(["items", 1, "name"])) + ); + assert!(guide.contains("CustomColumn: String")); + assert!(guide.contains("JSONExtractRaw(metadata, '&{{key}}', 'nested.key')")); + assert!(guide.contains("SpanAttributes['custom.tag']")); + assert!(guide.contains("ResourceAttributes['custom.resource']")); + assert_eq!(help["attributes"][0]["fields"][0]["key"], "custom.tag"); + assert_eq!(help["attributes"][1]["fields"][0]["key"], "custom.resource"); + for field in fields { + let expression = field["expression"].as_str().ok_or("missing expression")?; + assert!( + guide.contains(expression), + "missing plain-text expression: {expression}" + ); + let sql = format!("SELECT {expression} AS value FROM spend_logs FINAL"); + let body = litellm_traces::query_sql(&database.client, &reader, &sql).await?; + let values: serde_json::Value = serde_json::from_str(&body)?; + assert_ne!(values["data"][0]["value"], ""); + } + } + for example in help["examples"].as_array().ok_or("missing examples")? { + let sql = example["sql"].as_str().ok_or("missing example SQL")?; + assert!(guide.contains(example["name"].as_str().ok_or("missing example name")?)); + assert!(guide.contains(sql)); + assert_eq!( + example + .as_object() + .ok_or("example object")? + .keys() + .map(String::as_str) + .collect::>(), + std::collections::BTreeSet::from(["name", "sql"]) + ); + let body = litellm_traces::query_sql(&database.client, &reader, sql).await?; + let values: serde_json::Value = serde_json::from_str(&body)?; + assert_eq!( + values["data"].as_array().ok_or("missing data")?.is_empty(), + !populated, + "{sql}" + ); + if populated && example["name"] == "Traces correlated with LLM call metadata" { + assert_eq!(values["data"][0]["TraceId"], "trace-1"); + assert_eq!(values["data"][0]["spend"], 0.25); + } + } + Ok(()) +} + +#[rstest] +#[case::metadata(2, 1)] +#[case::attributes(1, 2)] +#[case::all(2, 2)] +#[tokio::test] +async fn query_help_preserves_schema_and_guide_when_discovery_hits_reader_limits( + #[future(awt)] database: TestResult, + #[case] spend_rows: usize, + #[case] span_rows: usize, +) -> TestResult { + let database = database?; + let writer = Connection::writer(&database.url)?; + ensure_schema(&database.client, &writer, "trace_test", 7, 14).await?; + execute_write( + &database, + "CREATE USER help_reader SETTINGS max_rows_to_read = 1", + ) + .await?; + for table in ["otel_traces", "agent_traces_by_key", "spend_logs"] { + execute_write( + &database, + &format!("GRANT SELECT ON trace_test.{table} TO help_reader"), + ) + .await?; + } + let timestamp = time::OffsetDateTime::now_utc().unix_timestamp_nanos() as i64; + let spend = (0..spend_rows) + .map(|index| { + serde_json::from_value(serde_json::json!({ + "request_id": format!("request-{index}"), "start_time": timestamp / 1_000_000, + "end_time": timestamp / 1_000_000, "metadata": r#"{"custom":{"enabled":true}}"# + })) + }) + .collect::, _>>()?; + insert_rows(&database, "spend_logs", spend).await?; + let spans = (0..span_rows).map(|index| serde_json::from_value(serde_json::json!({ + "Timestamp": timestamp, "TraceId": "trace", "SpanId": format!("span-{index}"), + "SpanAttributes": {"custom.span": "value"}, "ResourceAttributes": {"custom.resource": "value"} + }))).collect::, _>>()?; + insert_rows(&database, "otel_traces", spans).await?; + let reader = Connection::configured(&database.url, "trace_test", "help_reader", "")?; + let help: serde_json::Value = + serde_json::from_str(&litellm_traces::query_help(&database.client, &reader).await?)?; + assert_eq!(help["tables"].as_array().ok_or("tables")?.len(), 3); + assert!(!help["examples"].as_array().ok_or("examples")?.is_empty()); + assert_eq!( + help["normalized_fields"] + .as_array() + .ok_or("normalized fields")? + .len(), + NORMALIZED_FIELD_DEFINITIONS.len() + ); + let guide = help["guide"].as_str().ok_or("guide")?; + assert!(guide.contains("TraceId: String")); + assert_eq!( + guide.contains("Metadata discovery unavailable:"), + spend_rows > 1 + ); + assert_eq!( + guide.contains("Attribute discovery unavailable:"), + span_rows > 1 + ); + for (catalog, unavailable) in [ + (&help["metadata"], spend_rows > 1), + (&help["attributes"][0], span_rows > 1), + (&help["attributes"][1], span_rows > 1), + ] { + assert_eq!(catalog.get("error").is_some(), unavailable); + assert_eq!(catalog["truncated"], unavailable); + assert_eq!( + catalog["fields"].as_array().ok_or("fields")?.is_empty(), + unavailable + ); + } + Ok(()) +} diff --git a/litellm-rust/crates/traces/tests/otlp.rs b/litellm-rust/crates/traces/tests/otlp.rs index 8aa2cbedeb3..3baa820d312 100644 --- a/litellm-rust/crates/traces/tests/otlp.rs +++ b/litellm-rust/crates/traces/tests/otlp.rs @@ -1,5 +1,5 @@ -use litellm_traces::Shared; use litellm_traces::decode_otlp; +use litellm_traces::{ObservationType, Shared}; use rstest::rstest; const FIXTURE: &[u8] = include_bytes!( @@ -341,3 +341,48 @@ fn escaped_attribute_expansion_is_bounded_below_four_mib( Err(litellm_traces::DecodeError::TooLarge) )); } + +#[rstest] +fn normalizes_langsmith_fixture() { + let spans = decode_otlp(FIXTURE, Some("application/json")).expect("valid OTLP export"); + let llm = spans + .iter() + .find(|span| span.name == "ChatOpenAI") + .expect("LLM span"); + assert_eq!(llm.normalized.observation_type, ObservationType::Llm); + assert_eq!(llm.normalized.agent_name, "deep_research_agent"); + assert_eq!(llm.normalized.model, "claude-sonnet-4-5"); + assert_eq!( + (llm.normalized.input_tokens, llm.normalized.output_tokens), + (3332, 467) + ); + assert_eq!( + llm.normalized.litellm_request_id, + "chatcmpl-4077bb36-9380-4a3b-9481-245700cef09a" + ); + let input: serde_json::Value = + serde_json::from_str(&llm.normalized.input).expect("message input"); + assert_eq!(input[0]["role"], "system"); + assert_eq!(input[1]["role"], "user"); + let output: serde_json::Value = + serde_json::from_str(&llm.normalized.output).expect("message output"); + assert_eq!(output["role"], "assistant"); + assert!(output["tool_calls"][0]["name"].is_string()); + assert!(output["tool_calls"][0]["id"].is_string()); + assert_eq!(output["tool_calls"][0]["type"], "tool_call"); + let root = spans + .iter() + .find(|span| span.name == "deep_research_agent") + .expect("root span"); + assert_eq!(root.normalized.observation_type, ObservationType::Agent); + assert_eq!( + root.normalized.input, + "[{\"role\": \"user\", \"content\": \"Should we store OTEL agent spans in ClickHouse or Postgres at 50k spans/sec?\"}]" + ); + let tool = spans + .iter() + .find(|span| span.name == "task") + .expect("tool span"); + assert_eq!(tool.normalized.observation_type, ObservationType::Tool); + assert!(tool.normalized.output.starts_with("Based on my research")); +} diff --git a/litellm-rust/crates/traces/tests/query_access.rs b/litellm-rust/crates/traces/tests/query_access.rs new file mode 100644 index 00000000000..d6f44075e3b --- /dev/null +++ b/litellm-rust/crates/traces/tests/query_access.rs @@ -0,0 +1,268 @@ +use std::collections::BTreeMap; + +use litellm_http::Client; +use litellm_traces::{ + Connection, Error, QueryReaders, QueryScope, ensure_schema, query_help, query_sql, +}; +use rstest::{fixture, rstest}; +use serde_json::{Value, json}; +use testcontainers_modules::{ + clickhouse::ClickHouse, + testcontainers::{ContainerAsync, ImageExt, runners::AsyncRunner}, +}; + +struct Database { + _container: ContainerAsync, + client: Client, + writer: Connection, + readers: QueryReaders, +} + +#[fixture] +async fn database() -> Result> { + let container = ClickHouse::default() + .with_tag( + "26.9.6.6@sha256:eb4870e7ca7ed70c259eebfcfbee6cf797017f6b5436c2926bbbfe3d4d28486e", + ) + .with_env_var("CLICKHOUSE_SKIP_USER_SETUP", "1") + .start() + .await?; + let writer = Connection::parse(&format!( + "http://{}:{}", + container.get_host().await?, + container.get_host_port_ipv4(8123).await? + ))?; + let client = Client::no_redirect_for_test(); + ensure_schema(&client, &writer, "trace_test", 7, 7).await?; + for sql in [ + "INSERT INTO trace_test.otel_traces (TeamId, ApiKeyHash, TraceId, SpanId, Timestamp, SpanAttributes) VALUES ('team-a', 'key-a1', 'shared-trace', 'a1', now(), map('visible', 'a')), ('team-a', 'key-a2', 'shared-trace', 'a2', now(), map('visible', 'a')), ('team-b', 'key-b', 'shared-trace', 'b', now(), map('secret-b', 'b'))", + "INSERT INTO trace_test.spend_logs (team_id, api_key, request_id, start_time, end_time, metadata) VALUES ('team-a', 'key-a1', 'a1', now(), now(), '{\"visible\":1}'), ('team-a', 'key-a2', 'a2', now(), now(), '{\"visible\":1}'), ('team-b', 'key-b', 'b', now(), now(), '{\"secret_b\":1}')", + "CREATE TABLE trace_test.private_data (secret String) ENGINE = Memory", + "INSERT INTO trace_test.private_data VALUES ('hidden')", + ] { + let response = client.post(writer.url().clone()).body(sql).send().await?; + assert!(response.status().is_success(), "{}", response.text().await?); + } + let readers = QueryReaders::new(writer.clone(), "trace_test".to_owned()); + Ok(Database { + _container: container, + client, + writer, + readers, + }) +} + +#[rstest] +#[case::team(QueryScope::Team { team_id: "team-a".to_owned() }, vec!["a1", "a2"])] +#[case::project_key(QueryScope::Key { team_id: "team-a".to_owned(), api_key_hash: "key-a1".to_owned() }, vec!["a1"])] +#[case::admin(QueryScope::Admin, vec!["a1", "a2", "b"])] +#[case::quoted_team(QueryScope::Team { team_id: "team-a' OR 1=1 --\\".to_owned() }, vec![])] +#[tokio::test] +async fn queries_and_help_are_scoped_by_the_database( + #[future(awt)] database: Result>, + #[case] scope: QueryScope, + #[case] expected: Vec<&str>, +) -> Result<(), Box> { + let database = database?; + let reader = database + .readers + .connection(&database.client, &scope, "test-master-secret") + .await?; + let queries = [ + "SELECT SpanId AS id FROM otel_traces ORDER BY id", + "SELECT SpanId AS id FROM trace_test.otel_traces WHERE 1 = 1 ORDER BY id", + "SELECT SpanId AS id FROM merge('trace_test', '^otel_traces$') ORDER BY id", + "WITH source AS (SELECT * FROM trace_test.otel_traces) SELECT SpanId AS id FROM source ORDER BY id", + "SELECT SpanId AS id FROM otel_traces UNION DISTINCT SELECT SpanId AS id FROM trace_test.otel_traces ORDER BY id", + "SELECT t.SpanId AS id FROM otel_traces t INNER JOIN spend_logs s ON t.SpanId = s.request_id ORDER BY id", + "SELECT request_id AS id FROM spend_logs FINAL ORDER BY id", + ]; + for sql in queries { + let body: Value = serde_json::from_str(&query_sql(&database.client, &reader, sql).await?)?; + assert_eq!( + body["data"], + json!( + expected + .iter() + .map(|id| json!({"id": id})) + .collect::>() + ), + "{sql}" + ); + } + let summary: Value = serde_json::from_str( + &query_sql( + &database.client, + &reader, + "SELECT sum(SpanCount) AS count FROM agent_traces_by_key", + ) + .await?, + )?; + assert_eq!(summary["data"][0]["count"], json!(expected.len())); + let help = query_help(&database.client, &reader).await?; + assert_eq!(help.contains("secret_b"), expected.contains(&"b")); + assert_eq!(help.contains("secret-b"), expected.contains(&"b")); + let recreated = QueryReaders::new(database.writer.clone(), "trace_test".to_owned()); + let repeated = recreated + .connection(&database.client, &scope, "test-master-secret") + .await?; + assert_eq!(reader.url(), repeated.url()); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn rotating_master_secret_revokes_previous_reader_credentials( + #[future(awt)] database: Result>, +) -> Result<(), Box> { + let database = database?; + let scope = QueryScope::Team { + team_id: "team-a".to_owned(), + }; + let old_reader = database + .readers + .connection(&database.client, &scope, "old-master-secret") + .await?; + let old_result = query_sql( + &database.client, + &old_reader, + "SELECT SpanId AS id FROM otel_traces ORDER BY id", + ) + .await?; + let old_rows: Value = serde_json::from_str(&old_result)?; + assert_eq!(old_rows["data"], json!([{ "id": "a1" }, { "id": "a2" }])); + + let rotated_readers = QueryReaders::new(database.writer.clone(), "trace_test".into()); + let new_reader = rotated_readers + .connection(&database.client, &scope, "new-master-secret") + .await?; + assert!( + query_sql( + &database.client, + &old_reader, + "SELECT SpanId AS id FROM otel_traces ORDER BY id", + ) + .await + .is_err() + ); + let new_result = query_sql( + &database.client, + &new_reader, + "SELECT SpanId AS id FROM otel_traces ORDER BY id", + ) + .await?; + let new_rows: Value = serde_json::from_str(&new_result)?; + assert_eq!(new_rows["data"], json!([{ "id": "a1" }, { "id": "a2" }])); + assert_eq!(old_reader.url().username(), new_reader.url().username()); + assert_ne!(old_reader.url().password(), new_reader.url().password()); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn managed_reader_rejects_privilege_and_scope_bypasses( + #[future(awt)] database: Result>, +) -> Result<(), Box> { + let database = database?; + let scope = QueryScope::Team { + team_id: "team-a".to_owned(), + }; + let reader = database + .readers + .connection(&database.client, &scope, "test-master-secret") + .await?; + for sql in [ + "INSERT INTO otel_traces (TraceId) VALUES ('injected')", + "DROP TABLE otel_traces", + "SELECT * FROM private_data", + "SELECT * FROM otel_traces SETTINGS readonly = 0", + "SELECT * FROM otel_traces SETTINGS max_memory_usage = 0", + "SELECT * FROM otel_traces SETTINGS max_execution_time = 0", + "CREATE USER scope_bypass", + "CREATE NAMED COLLECTION scope_bypass AS host = 'localhost'", + "BACKUP TABLE otel_traces TO Disk('default', 'scope-bypass')", + "SELECT * FROM url('http://127.0.0.1:1/', 'LineAsString', 'line String')", + "SELECT * FROM remote('127.0.0.1', 'trace_test', 'otel_traces')", + ] { + assert!( + matches!( + query_sql(&database.client, &reader, sql).await, + Err(Error::QueryFailed(_)) + ), + "{sql}" + ); + } + let roles: Value = serde_json::from_str( + &query_sql(&database.client, &reader, "SELECT enabledRoles() AS roles").await?, + )?; + assert_eq!(roles["data"], json!([{ "roles": [] }])); + let rows: Value = serde_json::from_str( + &query_sql( + &database.client, + &reader, + "SELECT DISTINCT TeamId FROM otel_traces", + ) + .await?, + )?; + assert_eq!(rows["data"], json!([{ "TeamId": "team-a" }])); + Ok(()) +} + +#[rstest] +#[tokio::test] +async fn provisioning_failure_never_returns_a_writer_connection( + #[future(awt)] database: Result>, +) -> Result<(), Box> { + let database = database?; + let reader = database + .readers + .connection(&database.client, &QueryScope::Admin, "test-master-secret") + .await?; + let no_provision_privileges = QueryReaders::new(reader, "trace_test".to_owned()); + let result = no_provision_privileges + .connection( + &database.client, + &QueryScope::Team { + team_id: "team-a".to_owned(), + }, + "other-secret", + ) + .await; + assert!(result.is_err()); + assert!( + database + .readers + .connection(&database.client, &QueryScope::Admin, "") + .await + .is_err() + ); + assert!( + database + .readers + .connection( + &database.client, + &QueryScope::Team { + team_id: String::new() + }, + "test-master-secret" + ) + .await + .is_err() + ); + let permits = (0..8) + .map(|_| database.readers.acquire()) + .collect::, _>>()?; + assert!(database.readers.acquire().is_err()); + drop(permits); + assert!(database.readers.acquire().is_ok()); + let rows = litellm_traces::execute_read( + &database.client, + &database.writer, + "SELECT count() AS count FROM trace_test.otel_traces", + &BTreeMap::new(), + ) + .await?; + let rows: Value = serde_json::from_str(&rows)?; + assert_eq!(rows["data"][0]["count"], 3); + Ok(()) +} diff --git a/litellm/__init__.py b/litellm/__init__.py index d79086a90f3..9a4f4605519 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -2282,6 +2282,24 @@ if TYPE_CHECKING: # Track if async client cleanup has been registered (for lazy loading) _async_client_cleanup_registered = False +# litellm.agent() entrypoints, resolved lazily from litellm.harness by __getattr__. +_AGENT_EXPORTS: Final = frozenset( + { + "agent", + "aagent", + "agent_session", + "aagent_session", + "agent_resume", + "aagent_resume", + "agent_capabilities", + "Harness", + "ClaudeCodeOptions", + "CodexOptions", + "OpenCodeOptions", + "DeepAgentsOptions", + } +) + # Eager loading for backwards compatibility with VCR and other HTTP recording tools # When LITELLM_DISABLE_LAZY_LOADING is set, lazy-loaded attributes are loaded at import time # For now, this only affects encoding (tiktoken) as it was the only reported issue @@ -2315,6 +2333,13 @@ def __getattr__(name: str) -> Any: handler_func: Final = registry[name] return handler_func(name) + # litellm.agent() and friends: imported on first access (not needed for completion calls) + if name == "harness" or name in _AGENT_EXPORTS: + import importlib + + harness_module = importlib.import_module("litellm.harness") + return harness_module if name == "harness" else getattr(harness_module, name) + # Lazy load encoding from main.py to avoid heavy tiktoken import if name == "encoding": from ._lazy_imports import get_litellm_globals diff --git a/litellm/anthropic_beta_headers_config.json b/litellm/anthropic_beta_headers_config.json index b57239f8699..332d9b3ad9d 100644 --- a/litellm/anthropic_beta_headers_config.json +++ b/litellm/anthropic_beta_headers_config.json @@ -67,6 +67,7 @@ "text_editor_20241022": null, "text_editor_20250124": null, "thinking-binding-controls-2026-08-01": null, + "thinking-display-updates-2026-08-18": "thinking-display-updates-2026-08-18", "token-efficient-tools-2025-02-19": null, "web-fetch-2025-09-10": "web-fetch-2025-09-10", "web-search-2025-03-05": "web-search-2025-03-05" @@ -100,6 +101,7 @@ "text_editor_20241022": null, "text_editor_20250124": null, "thinking-binding-controls-2026-08-01": "thinking-binding-controls-2026-08-01", + "thinking-display-updates-2026-08-18": "thinking-display-updates-2026-08-18", "token-efficient-tools-2025-02-19": null, "tool-search-tool-2025-10-19": null, "web-fetch-2025-09-10": null, @@ -172,6 +174,7 @@ "text_editor_20241022": null, "text_editor_20250124": null, "thinking-binding-controls-2026-08-01": "thinking-binding-controls-2026-08-01", + "thinking-display-updates-2026-08-18": "thinking-display-updates-2026-08-18", "token-efficient-tools-2025-02-19": "token-efficient-tools-2025-02-19", "tool-examples-2025-10-29": "tool-examples-2025-10-29", "tool-search-tool-2025-10-19": "tool-search-tool-2025-10-19", @@ -207,6 +210,7 @@ "text_editor_20241022": null, "text_editor_20250124": null, "thinking-binding-controls-2026-08-01": "thinking-binding-controls-2026-08-01", + "thinking-display-updates-2026-08-18": "thinking-display-updates-2026-08-18", "token-efficient-tools-2025-02-19": null, "tool-search-tool-2025-10-19": "tool-search-tool-2025-10-19", "web-fetch-2025-09-10": null, diff --git a/litellm/constants.py b/litellm/constants.py index 76ab419272f..af4d1268c03 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -2162,6 +2162,17 @@ MCP_SPEND_LOG_MODEL_PREFIX: Final[str] = "MCP: " PTU_SENTINEL_API_KEY: Final[str] = "__ptu_flat_cost__" PTU_ROLLUP_JOB_ID: Final[str] = "ptu_flat_cost_rollup_job" PTU_ROLLUP_LOCK_TTL_SECONDS: Final[int] = 900 +USAGE_TOP_API_KEYS_DEFAULT: Final[int] = 100 +USAGE_TOP_API_KEYS_MAX: Final[int] = 1000 +USAGE_KEY_PAGE_DEFAULT: Final[int] = 50 +USAGE_KEY_PAGE_MAX: Final[int] = 100 +USAGE_KEY_SEARCH_DEFAULT: Final[int] = 100 +USAGE_KEY_SEARCH_MAX: Final[int] = 100 +USAGE_MODEL_TOP_KEYS_DEFAULT: Final[int] = 5 +USAGE_MODEL_TOP_KEYS_MAX: Final[int] = 100 +USAGE_CACHE_LEAKAGE_KEYS_DEFAULT: Final[int] = 20 +USAGE_CACHE_LEAKAGE_KEYS_MAX: Final[int] = 100 +USAGE_EXPORT_BATCH_SIZE: Final[int] = 1000 # Furthest back the catch-up pass looks for unpriced PTU days when a deployment # declares no ptu_effective_from, bounding the scan for an open-ended window. PTU_ROLLUP_MAX_BACKFILL_DAYS: Final[int] = 90 @@ -2199,3 +2210,26 @@ EMPTY_MAPPING: Final = MappingProxyType({}) # API endpoint for breached password k-anonymity search HIBP_RANGE_API_BASE: Final = "https://api.pwnedpasswords.com/range" + +# litellm.harness defaults +HARNESS_ENDPOINT_HOST: Final = "127.0.0.1" +HARNESS_ENDPOINT_STARTUP_TIMEOUT_SECONDS: Final = 10.0 +HARNESS_ENDPOINT_REQUEST_TIMEOUT_SECONDS: Final = 600.0 +HARNESS_SESSION_TOKEN_BYTES: Final = 32 +HARNESS_MAX_DIFF_BYTES: Final = 256 * 1024 +HARNESS_STDERR_TAIL_LINES: Final = 40 +HARNESS_STREAM_READ_CHUNK_BYTES: Final = 64 * 1024 +HARNESS_EVENT_QUEUE_MAX_SIZE: Final = 1024 +HARNESS_PROCESS_KILL_GRACE_SECONDS: Final = 5.0 +HARNESS_SNAPSHOT_SKIP_DIRS: Final = frozenset( + { + ".git", + "node_modules", + ".venv", + "venv", + "__pycache__", + ".mypy_cache", + ".pytest_cache", + ".ruff_cache", + } +) diff --git a/litellm/harness/__init__.py b/litellm/harness/__init__.py new file mode 100644 index 00000000000..79322cdc3e5 --- /dev/null +++ b/litellm/harness/__init__.py @@ -0,0 +1,98 @@ +"""Agent harnesses: run Claude Code, Codex, OpenCode or Deep Agents on any LiteLLM model. + +The entrypoints live on the top-level package: + + import litellm + from litellm import Harness, sandbox + + result = litellm.agent( + Harness.CLAUDE_CODE, + "fix the failing test", + sandbox=sandbox.local("."), + model="litellm_proxy/claude-sonnet-4-5", # a model group on your AI Gateway + ) + +This module holds the types you get back: events, Result, State, errors. +""" + +from litellm.harness.errors import ( + CapabilityUnsupported, + HarnessError, + HarnessInstallFailed, + OptionsMismatch, + OutputInvalid, + SandboxError, + SessionClosed, + StateIncompatible, +) +from litellm.harness.options import ( + ClaudeCodeOptions, + CodexOptions, + DeepAgentsOptions, + OpenCodeOptions, +) +from litellm.harness.runtime import ( + AsyncEventStream, + AsyncSession, + aagent, + aagent_resume, + aagent_session, + agent_capabilities, +) +from litellm.harness.sync import EventStream, Session, agent, agent_resume, agent_session +from litellm.harness.types import ( + Approval, + Capabilities, + Compaction, + Done, + Event, + FileChange, + Harness, + Reasoning, + Result, + State, + Text, + ToolCall, + ToolResult, + Usage, +) + +__all__ = ( + "Approval", + "AsyncEventStream", + "AsyncSession", + "Capabilities", + "CapabilityUnsupported", + "ClaudeCodeOptions", + "CodexOptions", + "Compaction", + "DeepAgentsOptions", + "Done", + "Event", + "EventStream", + "FileChange", + "Harness", + "HarnessError", + "HarnessInstallFailed", + "OpenCodeOptions", + "OptionsMismatch", + "OutputInvalid", + "Reasoning", + "Result", + "SandboxError", + "Session", + "SessionClosed", + "State", + "StateIncompatible", + "Text", + "ToolCall", + "ToolResult", + "Usage", + "aagent", + "aagent_resume", + "aagent_session", + "agent", + "agent_capabilities", + "agent_resume", + "agent_session", +) diff --git a/litellm/harness/context.py b/litellm/harness/context.py new file mode 100644 index 00000000000..eaae9aafe1f --- /dev/null +++ b/litellm/harness/context.py @@ -0,0 +1,62 @@ +"""Per-session state shared by the runtime, handlers and harness configs.""" + +from __future__ import annotations + +from collections.abc import Awaitable, Callable, Mapping, Sequence +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any, TypeAlias + +from pydantic import BaseModel + +from litellm.harness.options import HarnessOptions +from litellm.harness.sandbox.base import Sandbox +from litellm.harness.types import Approval, Harness, PermissionMode + +if TYPE_CHECKING: + from litellm.harness.endpoint import ModelEndpoint + +ApprovalHandler: TypeAlias = Callable[ + [Approval], bool | Awaitable[bool] # mutable-ok: Callable parameter list in a type alias, not a runtime collection +] + + +@dataclass(frozen=True) +class GatewayTarget: + """Resolved LiteLLM AI Gateway for `litellm_proxy/` models. Internal, not exported.""" + + api_base: str + api_key: str + + +@dataclass +class SessionContext: + """Everything a handler and config need for a session. Owned by the runtime.""" + + harness: Harness + sandbox: Sandbox + session_id: str + # Model name as sent to the runtime (litellm_proxy/ prefix already stripped). + model: str | None = None + gateway: GatewayTarget | None = None + api_key: str | None = None + api_base: str | None = None + endpoint: ModelEndpoint | None = None + instructions: str | None = None + tools: Sequence[Callable[..., Any]] = () + skills: Sequence[str] = () + disable_tools: Sequence[str] = () + permissions: PermissionMode = "full" + on_approval: ApprovalHandler | None = None + output: type[BaseModel] | None = None + max_turns: int | None = None + timeout: float | None = None + metadata: Mapping[str, Any] = field(default_factory=dict) + options: HarnessOptions | None = None + # Set by the handler after each turn. + final_text: str = "" + output_json: str | None = None + # Usage for in-process harnesses that call LiteLLM directly (no model endpoint). + input_tokens: int = 0 + output_tokens: int = 0 + cost: float = 0.0 + calls: int = 0 diff --git a/litellm/harness/endpoint.py b/litellm/harness/endpoint.py new file mode 100644 index 00000000000..c492e4794ba --- /dev/null +++ b/litellm/harness/endpoint.py @@ -0,0 +1,689 @@ +"""Per-session local model endpoint every CLI harness talks to. + +The runtime inside the sandbox points its Anthropic / OpenAI base URL at this endpoint and +authenticates with a random per-session token. The endpoint either reverse-proxies to a LiteLLM +AI Gateway (gateway mode) or calls the LiteLLM SDK directly (SDK mode), and counts usage + cost. + +starlette and uvicorn are optional: they are imported only when an endpoint starts. +""" + +from __future__ import annotations + +import asyncio +import contextlib +import itertools +import json +import logging +import secrets +from collections.abc import AsyncIterable, AsyncIterator, Mapping +from dataclasses import dataclass +from types import MappingProxyType, ModuleType +from typing import TYPE_CHECKING, Any, Final + +import httpx +import openai + +import litellm +from litellm.constants import ( + DEFAULT_POLLING_INTERVAL, + HARNESS_ENDPOINT_HOST, + HARNESS_ENDPOINT_REQUEST_TIMEOUT_SECONDS, + HARNESS_ENDPOINT_STARTUP_TIMEOUT_SECONDS, + HARNESS_PROCESS_KILL_GRACE_SECONDS, + HARNESS_SESSION_TOKEN_BYTES, +) +from litellm.harness.context import GatewayTarget +from litellm.harness.errors import HarnessError, HarnessInstallFailed +from litellm.harness.types import Harness, Usage +from litellm.llms.custom_httpx.http_handler import get_async_httpx_client +from litellm.types.llms.custom_http import httpxSpecialProvider + +if TYPE_CHECKING: + from starlette.applications import Starlette + from starlette.requests import Request + from starlette.responses import Response + from uvicorn import Server + +verbose_logger: Final = logging.getLogger("LiteLLM") + +MISSING_DEPS_MESSAGE = "litellm.harness needs starlette and uvicorn: pip install starlette uvicorn" + +ROUTE_MESSAGES = "messages" +ROUTE_CHAT = "chat/completions" +ROUTE_RESPONSES = "responses" +POST_ROUTES = (ROUTE_MESSAGES, ROUTE_CHAT, ROUTE_RESPONSES) +ROUTE_PREFIXES: Final = ("", "/v1") + +# What an SDK call or its stream can raise: LiteLLM maps provider failures onto openai's +# exception hierarchy; transport errors, bad request kwargs and unserializable chunks remain. +SDK_ERRORS: Final = (openai.OpenAIError, httpx.HTTPError, HarnessError, ValueError, TypeError) + +HOP_BY_HOP_HEADERS = frozenset( + { + "connection", + "keep-alive", + "proxy-authenticate", + "proxy-authorization", + "te", + "trailer", + "trailers", + "transfer-encoding", + "upgrade", + "host", + "content-length", + } +) +DROPPED_REQUEST_HEADERS = HOP_BY_HOP_HEADERS | frozenset( + ( + "authorization", + "x-api-key", + "accept-encoding", + ) +) +DROPPED_RESPONSE_HEADERS = HOP_BY_HOP_HEADERS | frozenset(("content-encoding",)) +COST_HEADER = "x-litellm-response-cost" +SSE_MEDIA_TYPE = "text/event-stream" + + +@dataclass(frozen=True) +class _ServerDeps: + uvicorn: ModuleType + applications: ModuleType + routing: ModuleType + responses: ModuleType + + +def _load_server_deps() -> _ServerDeps: + """Import starlette + uvicorn on demand; they are not litellm dependencies.""" + try: + import uvicorn + from starlette import applications, responses, routing + except ImportError as e: + raise HarnessInstallFailed(MISSING_DEPS_MESSAGE) from e + return _ServerDeps( + uvicorn=uvicorn, + applications=applications, + routing=routing, + responses=responses, + ) + + +@dataclass +class UsageTracker: + """Running token + cost totals for one session.""" + + input_tokens: int = 0 + output_tokens: int = 0 + cost: float = 0.0 + calls: int = 0 + + def add(self, input_tokens: int = 0, output_tokens: int = 0, cost: float = 0.0) -> None: + self.input_tokens += input_tokens + self.output_tokens += output_tokens + self.cost += cost + self.calls += 1 + + def snapshot(self) -> Usage: + return Usage( + input_tokens=self.input_tokens, + output_tokens=self.output_tokens, + calls=self.calls, + ) + + +# --------------------------------------------------------------------------- +# Usage parsing +# --------------------------------------------------------------------------- + + +def _as_int(value: object) -> int: + if isinstance(value, bool): + return 0 + if isinstance(value, (int, float)): + return int(value) + return 0 + + +def usage_from_mapping(usage: object) -> tuple[int, int]: + """(input, output) from a usage dict using OpenAI or Anthropic/Responses field names.""" + if not isinstance(usage, Mapping): + return 0, 0 + input_tokens = usage.get("input_tokens", usage.get("prompt_tokens")) + output_tokens = usage.get("output_tokens", usage.get("completion_tokens")) + return _as_int(input_tokens), _as_int(output_tokens) + + +def usage_from_body(body: object) -> tuple[int, int]: + """Usage from a non-streaming JSON response body.""" + if not isinstance(body, Mapping): + return 0, 0 + if isinstance(body.get("usage"), Mapping): + return usage_from_mapping(body["usage"]) + response = body.get("response") + if isinstance(response, Mapping): + return usage_from_mapping(response.get("usage")) + return 0, 0 + + +class SSEUsageParser: + """Collects token usage from an SSE byte stream as it passes through.""" + + def __init__(self) -> None: + self.input_tokens = 0 + self.output_tokens = 0 + self._buffer = b"" + + def feed(self, chunk: bytes) -> None: + self._buffer += chunk + *lines, self._buffer = self._buffer.split(b"\n") + for line in lines: + self._feed_line(line) + + def close(self) -> None: + if self._buffer: + self._feed_line(self._buffer) + self._buffer = b"" + + def _feed_line(self, line: bytes) -> None: + text = line.strip() + if not text.startswith(b"data:"): + return + payload = text[len(b"data:") :].strip() + if not payload or payload == b"[DONE]": + return + try: + event = json.loads(payload) + except ValueError: + return + if isinstance(event, Mapping): + self.absorb(event) + + def absorb(self, event: Mapping[str, Any]) -> None: + event_type = event.get("type") + if event_type == "message_start": + self._absorb_message_start(event) + elif event_type == "message_delta": + self._absorb_message_delta(event) + elif event_type == "response.completed": + self._absorb_response_completed(event) + elif isinstance(event.get("usage"), Mapping): + self._set(*usage_from_mapping(event["usage"])) + + def _absorb_message_start(self, event: Mapping[str, Any]) -> None: + message = event.get("message") + if isinstance(message, Mapping): + self._set(*usage_from_mapping(message.get("usage"))) + + def _absorb_message_delta(self, event: Mapping[str, Any]) -> None: + # message_delta output_tokens is cumulative for the whole message. + self._set(*usage_from_mapping(event.get("usage"))) + + def _absorb_response_completed(self, event: Mapping[str, Any]) -> None: + response = event.get("response") + if isinstance(response, Mapping): + self._set(*usage_from_mapping(response.get("usage"))) + + def _set(self, input_tokens: int, output_tokens: int) -> None: + if input_tokens: + self.input_tokens = input_tokens + if output_tokens: + self.output_tokens = output_tokens + + +# --------------------------------------------------------------------------- +# Cost + helpers +# --------------------------------------------------------------------------- + + +def compute_cost(model: str | None, input_tokens: int, output_tokens: int) -> float: + """Cost from LiteLLM's price map. Never raises; unknown models cost 0.0.""" + if not model or not (input_tokens or output_tokens): + return 0.0 + try: + prompt_cost, completion_cost = litellm.cost_per_token( + model=model, prompt_tokens=input_tokens, completion_tokens=output_tokens + ) + return float(prompt_cost) + float(completion_cost) + except Exception: # accounting must never break a call; the price-map lookup raises bare Exception + verbose_logger.debug("harness endpoint: cost lookup failed for %s", model, exc_info=True) + return 0.0 + + +def header_cost(headers: Mapping[str, str]) -> float | None: + raw = headers.get(COST_HEADER) + if raw is None: + return None + try: + return float(raw) + except (TypeError, ValueError): + return None + + +def hidden_cost(response: object) -> float | None: + hidden = getattr(response, "_hidden_params", None) + if not isinstance(hidden, Mapping): + return None + try: + cost = hidden.get("response_cost") + return None if cost is None else float(cost) + except (TypeError, ValueError): + return None + + +def extract_token(headers: Mapping[str, str]) -> str | None: + auth = headers.get("authorization") or "" + if auth.lower().startswith("bearer "): + return auth[len("bearer ") :].strip() + return headers.get("x-api-key") + + +def gateway_headers( + incoming: Mapping[str, str], + gateway: GatewayTarget, + harness: Harness, + metadata: Mapping[str, Any] | None, +) -> Mapping[str, str]: + """Incoming headers minus hop-by-hop/auth/x-litellm-*, plus gateway auth, tags, metadata.""" + kept = ( + (name, value) + for name, value in incoming.items() + if name.lower() not in DROPPED_REQUEST_HEADERS and not name.lower().startswith("x-litellm-") + ) + metadata_json = json.dumps(dict(metadata), default=str) if metadata else None # mutable-ok: for json.dumps + metadata_header = (("x-litellm-spend-logs-metadata", metadata_json),) if metadata_json is not None else () + added = ( + ("authorization", f"Bearer {gateway.api_key}"), + ("x-litellm-tags", f"harness,{harness.value}"), + *metadata_header, + ) + return MappingProxyType(dict(itertools.chain(kept, added))) + + +def response_headers(upstream: Mapping[str, str]) -> Mapping[str, str]: + return MappingProxyType( + {name: value for name, value in upstream.items() if name.lower() not in DROPPED_RESPONSE_HEADERS} + ) + + +def sanitize(message: str, secret_values: tuple[str | None, ...]) -> str: + for value in secret_values: + if value: + message = message.replace(value, "***") + return message + + +def error_status(exc: BaseException) -> int: + status = getattr(exc, "status_code", None) + if isinstance(status, int) and 400 <= status <= 599: + return status + return 500 + + +def error_body(exc: BaseException, message: str) -> dict[str, Any]: # mutable-ok: JSONResponse body + return {"error": {"type": type(exc).__name__, "message": message}} # mutable-ok: JSONResponse body + + +def to_jsonable(obj: object) -> object: + if hasattr(obj, "model_dump"): + return obj.model_dump(mode="json", exclude_none=True) + if isinstance(obj, Mapping): + return dict(obj) # mutable-ok: plain-dict copy so json.dumps can serialize any Mapping + return obj + + +def encode_anthropic_chunk(chunk: object) -> bytes: + if isinstance(chunk, bytes): + return chunk + if isinstance(chunk, str): + return chunk.encode() + data = to_jsonable(chunk) + event_type = data.get("type", "message") if isinstance(data, Mapping) else "message" + return f"event: {event_type}\ndata: {json.dumps(data)}\n\n".encode() + + +def encode_chat_chunk(chunk: object) -> bytes: + if hasattr(chunk, "model_dump_json"): + return f"data: {chunk.model_dump_json()}\n\n".encode() + return f"data: {json.dumps(to_jsonable(chunk))}\n\n".encode() + + +def encode_responses_chunk(chunk: object) -> bytes: + data = to_jsonable(chunk) + event_type = data.get("type", "message") if isinstance(data, Mapping) else "message" + return f"event: {event_type}\ndata: {json.dumps(data)}\n\n".encode() + + +STREAM_ENCODERS = MappingProxyType( + { + ROUTE_MESSAGES: encode_anthropic_chunk, + ROUTE_CHAT: encode_chat_chunk, + ROUTE_RESPONSES: encode_responses_chunk, + } +) +STREAM_TRAILERS = MappingProxyType({ROUTE_CHAT: b"data: [DONE]\n\n"}) + + +def route_of(path: str) -> str: + stripped = path.strip("/") + stripped = stripped.removeprefix("v1/") + return stripped + + +def _noop() -> None: + return None + + +# --------------------------------------------------------------------------- +# ModelEndpoint +# --------------------------------------------------------------------------- + + +class ModelEndpoint: + """Local HTTP endpoint for one harness session. Use as an async context manager.""" + + def __init__( + self, + harness: Harness, + model: str | None, + gateway: GatewayTarget | None, + api_key: str | None = None, + api_base: str | None = None, + metadata: Mapping[str, Any] | None = None, + *, + client: httpx.AsyncClient | None = None, + ) -> None: + self.harness = harness + self.model = model + self.gateway = gateway + self.api_key = api_key + self.api_base = api_base + self.metadata: Mapping[str, Any] = MappingProxyType(dict(metadata or ())) + self.token = secrets.token_urlsafe(HARNESS_SESSION_TOKEN_BYTES) + self.usage = UsageTracker() + self.port = 0 + # Injected client (tests); production uses LiteLLM's shared cached client. + self._injected_client = client + self._deps: _ServerDeps | None = None + self._client: httpx.AsyncClient | None = None + self._server: Any = None + self._task: asyncio.Task[None] | None = None + + @property + def url(self) -> str: + return f"http://{HARNESS_ENDPOINT_HOST}:{self.port}" + + # -- lifecycle ---------------------------------------------------------- + + async def __aenter__(self) -> ModelEndpoint: + await self.start() + return self + + async def __aexit__(self, *exc_info: object) -> None: + await self.stop() + + async def start(self) -> None: + self._deps = _load_server_deps() + if self.gateway is not None: + self._client = self._gateway_client() + self._server = self._build_server(self._deps) + self._task = asyncio.create_task(self._server.serve()) + try: + await asyncio.wait_for(self._wait_started(), HARNESS_ENDPOINT_STARTUP_TIMEOUT_SECONDS) + except BaseException: + await self.stop() + raise + self.port = self._server.servers[0].sockets[0].getsockname()[1] + + async def stop(self) -> None: + if self._server is not None: + self._server.should_exit = True + if self._task is not None: + with contextlib.suppress(BaseException): + await self._task + self._task = None + # Never close the client: the shared cached one may still serve other requests, + # and an injected one belongs to its caller. + self._client = None + + def _gateway_client(self) -> httpx.AsyncClient: + """LiteLLM's shared cached async client, unless one was injected.""" + if self._injected_client is not None: + return self._injected_client + handler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.AgentHarness, + params={ # mutable-ok: get_async_httpx_client takes a dict params argument + "timeout": HARNESS_ENDPOINT_REQUEST_TIMEOUT_SECONDS + }, + ) + return handler.client + + async def _wait_started(self) -> None: + while not self._server.started: + if self._task is not None and self._task.done(): + raise HarnessError("harness model endpoint failed to start") + await asyncio.sleep(DEFAULT_POLLING_INTERVAL) + + def _build_server(self, deps: _ServerDeps) -> Server: + config = deps.uvicorn.Config( + self._build_app(deps), + host=HARNESS_ENDPOINT_HOST, + port=0, + log_config=None, + log_level="warning", + access_log=False, + lifespan="off", + timeout_graceful_shutdown=HARNESS_PROCESS_KILL_GRACE_SECONDS, + ) + server = deps.uvicorn.Server(config) + # Never touch the host process's signal handlers. + if hasattr(server, "capture_signals"): + server.capture_signals = contextlib.nullcontext + if hasattr(server, "install_signal_handlers"): + server.install_signal_handlers = _noop + return server + + def _build_app(self, deps: _ServerDeps) -> Starlette: + Route = deps.routing.Route + post_routes = tuple( + Route( + f"{prefix}/{route}", + self._handle, + methods=["POST"], # mutable-ok: Starlette Route takes a methods list + ) + for prefix, route in itertools.product(ROUTE_PREFIXES, POST_ROUTES) + ) + get_routes = tuple( + Route(f"{prefix}/models", self._models, methods=["GET"]) # mutable-ok: Starlette Route takes a methods list + for prefix in ROUTE_PREFIXES + ) + return deps.applications.Starlette( + routes=[*post_routes, *get_routes] # mutable-ok: Starlette takes a routes list + ) + + # -- request handling --------------------------------------------------- + + @property + def _responses(self) -> ModuleType: + if self._deps is None: + raise HarnessError("harness model endpoint is not started") + return self._deps.responses + + def _authorized(self, request: Request) -> bool: + token = extract_token(request.headers) + return token is not None and secrets.compare_digest(token.encode(), self.token.encode()) + + def _json(self, body: object, status_code: int = 200) -> Response: + return self._responses.JSONResponse(body, status_code=status_code) + + def _unauthorized(self) -> Response: + return self._json( + {"error": {"type": "authentication_error", "message": "invalid token"}}, # mutable-ok: JSONResponse body + 401, + ) + + def _error(self, exc: BaseException, status_code: int | None = None) -> Response: + message = sanitize(str(exc), self._secrets()) + return self._json(error_body(exc, message), status_code or error_status(exc)) + + def _secrets(self) -> tuple[str | None, ...]: + gateway_key = self.gateway.api_key if self.gateway else None + return (gateway_key, self.api_key, self.token) + + async def _models(self, request: Request) -> Response: + if not self._authorized(request): + return self._unauthorized() + entry = {"id": self.model, "object": "model", "created": 0, "owned_by": "litellm"} # mutable-ok: JSON body + data = (entry,) if self.model else () + return self._json({"object": "list", "data": data}) # mutable-ok: JSON response body for Starlette JSONResponse + + async def _handle(self, request: Request) -> Response: + if not self._authorized(request): + return self._unauthorized() + try: + body = json.loads(await request.body()) + except ValueError as e: + return self._error(e, 400) + if not isinstance(body, dict): + return self._error(ValueError("request body must be a JSON object"), 400) + route = route_of(request.url.path) + if self.gateway is not None: + return await self._forward(request, route, body) + return await self._call_sdk(route, body) + + def _cost_model(self, body: Mapping[str, Any]) -> str | None: + model = self.model or body.get("model") + return model if isinstance(model, str) else None + + def _record( + self, + model: str | None, + input_tokens: int, + output_tokens: int, + cost: float | None, + ) -> None: + if cost is None: + cost = compute_cost(model, input_tokens, output_tokens) + self.usage.add(input_tokens, output_tokens, cost) + + # -- gateway mode ------------------------------------------------------- + + async def _forward(self, request: Request, route: str, body: Mapping[str, Any]) -> Response: + if self._client is None or self.gateway is None: + raise HarnessError("gateway client is not started") + if self.model: + body = {**body, "model": self.model} # mutable-ok: JSON request body re-sent upstream via httpx json= + upstream_request = self._client.build_request( + "POST", + f"{self.gateway.api_base}/v1/{route}", + json=body, + headers=gateway_headers(request.headers, self.gateway, self.harness, self.metadata), + ) + try: + upstream = await self._client.send(upstream_request, stream=True) + except httpx.HTTPError as e: + return self._error(e, 502) + return self._responses.StreamingResponse( + self._relay(upstream, self._cost_model(body)), + status_code=upstream.status_code, + headers=response_headers(upstream.headers), + ) + + async def _relay(self, upstream: httpx.Response, model: str | None) -> AsyncIterator[bytes]: + is_sse = SSE_MEDIA_TYPE in upstream.headers.get("content-type", "") + parser = SSEUsageParser() + collected = bytearray() + try: + async for chunk in upstream.aiter_bytes(): + if is_sse: + parser.feed(chunk) + else: + collected.extend(chunk) + yield chunk + finally: + await upstream.aclose() + if upstream.status_code < 400: + self._record_relayed(upstream, model, parser, is_sse, bytes(collected)) + + def _record_relayed( + self, + upstream: httpx.Response, + model: str | None, + parser: SSEUsageParser, + is_sse: bool, + collected: bytes, + ) -> None: + if is_sse: + parser.close() + tokens = (parser.input_tokens, parser.output_tokens) + else: + try: + tokens = usage_from_body(json.loads(collected)) + except ValueError: + tokens = (0, 0) + self._record(model, tokens[0], tokens[1], header_cost(upstream.headers)) + + # -- SDK mode ----------------------------------------------------------- + + def _sdk_kwargs( + self, body: Mapping[str, Any] + ) -> dict[str, Any]: # mutable-ok: SDK call kwargs, mutated by _invoke_sdk then splatted + kwargs: dict[str, Any] = {**body} # mutable-ok: SDK call kwargs built from the JSON body, then overridden + if self.model: + kwargs["model"] = self.model + if self.api_key: + kwargs["api_key"] = self.api_key + if self.api_base: + kwargs["api_base"] = self.api_base + return kwargs + + async def _invoke_sdk( + self, + route: str, + kwargs: dict[str, Any], # mutable-ok: injects stream_options into the SDK kwargs + ) -> object: + if route == ROUTE_MESSAGES: + return await litellm.anthropic.messages.acreate(**kwargs) + if route == ROUTE_CHAT: + if kwargs.get("stream"): + stream_options = kwargs.get("stream_options") or {} # mutable-ok: empty default for a JSON field + kwargs["stream_options"] = { # mutable-ok: JSON field sent to litellm.acompletion + "include_usage": True, + **stream_options, + } + return await litellm.acompletion(**kwargs) + return await litellm.aresponses(**kwargs) + + async def _call_sdk(self, route: str, body: Mapping[str, Any]) -> Response: + kwargs = self._sdk_kwargs(body) + model = self._cost_model(kwargs) + try: + response = await self._invoke_sdk(route, kwargs) + except SDK_ERRORS as e: + verbose_logger.debug("harness endpoint: SDK call failed: %s", type(e).__name__) + return self._error(e) + if kwargs.get("stream") and isinstance(response, AsyncIterable): + return self._responses.StreamingResponse( + self._sdk_stream(route, response, model), media_type=SSE_MEDIA_TYPE + ) + data = to_jsonable(response) + input_tokens, output_tokens = usage_from_body(data) + self._record(model, input_tokens, output_tokens, hidden_cost(response)) + return self._json(data) + + async def _sdk_stream(self, route: str, iterator: AsyncIterable[object], model: str | None) -> AsyncIterator[bytes]: + encode = STREAM_ENCODERS[route] + parser = SSEUsageParser() + try: + async for chunk in iterator: + encoded = encode(chunk) + parser.feed(encoded) + yield encoded + trailer = STREAM_TRAILERS.get(route) + if trailer: + yield trailer + except SDK_ERRORS as e: + message = sanitize(str(e), self._secrets()) + yield f"event: error\ndata: {json.dumps(error_body(e, message))}\n\n".encode() + finally: + parser.close() + self._record(model, parser.input_tokens, parser.output_tokens, None) diff --git a/litellm/harness/errors.py b/litellm/harness/errors.py new file mode 100644 index 00000000000..efa155fdffc --- /dev/null +++ b/litellm/harness/errors.py @@ -0,0 +1,45 @@ +"""Exceptions raised by litellm.harness.""" + +from __future__ import annotations + +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from litellm.harness.types import Result + + +class HarnessError(Exception): + """Base class for every litellm.harness error.""" + + +class CapabilityUnsupported(HarnessError): + """The harness cannot do what was asked. Raised before the runtime starts.""" + + +class OptionsMismatch(HarnessError): + """Options for a different harness, or a native option LiteLLM manages itself.""" + + +class HarnessInstallFailed(HarnessError): + """The runtime is missing from the sandbox or failed to start.""" + + +class SandboxError(HarnessError): + """The sandbox failed to start, run a command, or reach the host.""" + + +class SessionClosed(HarnessError): + """A turn was started on a session that is closed or detached.""" + + +class StateIncompatible(HarnessError): + """resume() was given a State from another harness or an unreadable version.""" + + +class OutputInvalid(HarnessError): + """The final answer did not validate against output=.""" + + def __init__(self, message: str, raw: str, result: Result | None = None) -> None: + super().__init__(message) + self.raw = raw + self.result = result diff --git a/litellm/harness/handlers/__init__.py b/litellm/harness/handlers/__init__.py new file mode 100644 index 00000000000..4b32d7ca649 --- /dev/null +++ b/litellm/harness/handlers/__init__.py @@ -0,0 +1,35 @@ +"""Handlers run a harness config: CLI runtimes as subprocesses, Deep Agents in-process.""" + +from __future__ import annotations + +from litellm.harness.errors import HarnessError +from litellm.harness.handlers.base import BaseHarnessHandler +from litellm.harness.types import Harness, require_harness +from litellm.llms.base_llm.harness.transformation import ( + BaseCLIHarnessConfig, + BaseHarnessConfig, +) +from litellm.utils import ProviderConfigManager + + +def get_harness_config(harness: Harness) -> BaseHarnessConfig: + config = ProviderConfigManager.get_provider_harness_config(require_harness(harness)) + if config is None: + raise HarnessError(f"No harness config registered for Harness.{harness.name}") + return config + + +def get_harness_handler(config: BaseHarnessConfig) -> BaseHarnessHandler: + """The handler that knows how to run this kind of config.""" + if isinstance(config, BaseCLIHarnessConfig): + from litellm.harness.handlers.cli_handler import CLIHarnessHandler + + return CLIHarnessHandler(config) + if config.harness is Harness.DEEPAGENTS: + from litellm.harness.handlers.deepagents_handler import DeepAgentsHandler + + return DeepAgentsHandler(config) + raise HarnessError(f"No handler for Harness.{config.harness.name}") + + +__all__ = ("BaseHarnessHandler", "get_harness_config", "get_harness_handler") diff --git a/litellm/harness/handlers/base.py b/litellm/harness/handlers/base.py new file mode 100644 index 00000000000..18b7391e988 --- /dev/null +++ b/litellm/harness/handlers/base.py @@ -0,0 +1,42 @@ +"""The handler interface the runtime drives. A handler owns I/O for one session.""" + +from __future__ import annotations + +from abc import ABC, abstractmethod +from collections.abc import AsyncIterator +from typing import Any + +from litellm.harness.context import SessionContext +from litellm.harness.errors import CapabilityUnsupported +from litellm.harness.types import Event +from litellm.llms.base_llm.harness.transformation import BaseHarnessConfig + + +class BaseHarnessHandler(ABC): + """Runs one harness session. The config decides what to run; the handler runs it.""" + + def __init__(self, config: BaseHarnessConfig) -> None: + self.config = config + + @abstractmethod + async def start(self, ctx: SessionContext) -> None: + """Prepare the runtime (config files, skills, agent build). Called again after an interrupt.""" + + @abstractmethod + def turn(self, ctx: SessionContext, prompt: str) -> AsyncIterator[Event]: + """Run one turn and yield events (never Done). Sets ctx.final_text / ctx.output_json.""" + + @abstractmethod + async def stop(self, ctx: SessionContext) -> None: + """Stop anything this handler started. Safe to call twice.""" + + @abstractmethod + def native_session_id(self) -> str | None: + """The runtime's own session id, for State / resume.""" + + @abstractmethod + async def resume(self, ctx: SessionContext, native_session_id: str) -> None: + """Continue the runtime's own session on the next turn.""" + + async def history(self, ctx: SessionContext) -> list[dict[str, Any]]: # mutable-ok: public history() API shape + raise CapabilityUnsupported(f"Harness.{self.config.harness.name} does not expose history") diff --git a/litellm/harness/handlers/cli_handler.py b/litellm/harness/handlers/cli_handler.py new file mode 100644 index 00000000000..0fda16c34fe --- /dev/null +++ b/litellm/harness/handlers/cli_handler.py @@ -0,0 +1,161 @@ +""" +Generic handler for CLI harnesses (Claude Code, Codex, OpenCode). + +The config (`litellm/llms//harness/transformation.py`) says what to run and how to +read it; this handler does every sandbox and process operation: binary check, private dir, +config files, persisted dirs, skills, spawning the turn, streaming stdout lines into the +config's parser, collecting stderr, and killing the process on early exit. +""" + +from __future__ import annotations + +import asyncio +import os +from collections import deque +from collections.abc import AsyncIterator, Sequence +from typing import Any, Final + +from litellm._logging import verbose_logger +from litellm.constants import HARNESS_STDERR_TAIL_LINES, HARNESS_STREAM_READ_CHUNK_BYTES +from litellm.harness.context import SessionContext +from litellm.harness.errors import HarnessInstallFailed, SandboxError +from litellm.harness.handlers.base import BaseHarnessHandler +from litellm.harness.sandbox.base import Process, Sandbox +from litellm.harness.types import Event +from litellm.llms.base_llm.harness.transformation import BaseCLIHarnessConfig, HarnessSessionSetup +from litellm.llms.base_llm.harness.utils import decode_json_line, read_skill_files + +# Link / to a LiteLLM-owned cache dir so a later session can resume. +PERSIST_DIR_SCRIPT: Final = ( + 'd="${HOME:-/tmp}/.cache/litellm-harness/$2"; mkdir -p "$d" && mkdir -p "$(dirname "$1")" && ln -sfn "$d" "$1"' +) + + +async def iter_stream_lines(stream: asyncio.StreamReader) -> AsyncIterator[bytes]: + """Newline-delimited lines without StreamReader's 64KiB readline limit.""" + buffer = b"" + while True: + chunk = await stream.read(HARNESS_STREAM_READ_CHUNK_BYTES) + if not chunk: + break + buffer += chunk + *lines, buffer = buffer.split(b"\n") + for line in lines: + yield line + if buffer: + yield buffer + + +async def drain_stderr(stream: asyncio.StreamReader, tail: deque[str]) -> None: # mutable-ok: stderr ring + async for line in iter_stream_lines(stream): + tail.append(line.decode("utf-8", errors="replace")) + + +async def send_stdin(proc: Process, data: str) -> None: + if proc.stdin is None: + raise SandboxError("harness process has no stdin") + proc.stdin.write(data.encode("utf-8")) + await proc.stdin.drain() + proc.stdin.close() + + +async def private_dir_for(sandbox: Sandbox) -> str: + tempdir = getattr(sandbox, "tempdir", None) + if tempdir is None: + raise SandboxError(f"{type(sandbox).__name__} has no tempdir(); CLI harnesses need a private config dir") + path: str = await tempdir() + return path + + +def sandbox_path(private_dir: str, path: str) -> str: + return path if path.startswith("/") else f"{private_dir}/{path}" + + +async def persist_dir(sandbox: Sandbox, link_path: str, cache_subpath: str) -> None: + script_args: Final = ("-c", PERSIST_DIR_SCRIPT, "sh", link_path, cache_subpath) + cmd: Final = ["sh", *script_args] # mutable-ok: Sandbox.run takes list[str] + run = await sandbox.run(cmd) + if run.exit_code != 0: + verbose_logger.debug( + "harness: could not persist %s, resume across sessions disabled: %s", cache_subpath, run.stderr.strip() + ) + + +async def copy_skills(sandbox: Sandbox, skills: Sequence[str], skills_root: str) -> None: + for skill in skills: + name = os.path.basename(os.path.realpath(os.fspath(skill))) + for rel, data in await asyncio.to_thread(read_skill_files, skill): + await sandbox.write(f"{skills_root}/{name}/{rel.replace(os.sep, '/')}", data) + + +class CLIHarnessHandler(BaseHarnessHandler): + config: BaseCLIHarnessConfig + + def __init__(self, config: BaseCLIHarnessConfig) -> None: + super().__init__(config) + self._private_dir: str | None = None + self._setup: HarnessSessionSetup | None = None + self._native_id: str | None = None + self._proc: Process | None = None + + async def start(self, ctx: SessionContext) -> None: + self.config.validate_environment(ctx) + binary = self.config.get_binary() + if not await ctx.sandbox.which(binary): + raise HarnessInstallFailed( + f"`{binary}` was not found on PATH in the sandbox. Install it with: {self.config.get_install_hint()}" + ) + private_dir = await private_dir_for(ctx.sandbox) + setup = self.config.transform_session_setup(ctx, private_dir) + for link, cache_subpath in setup.persisted_dirs: + await persist_dir(ctx.sandbox, sandbox_path(private_dir, link), cache_subpath) + for rel_path, data in setup.files.items(): + await ctx.sandbox.write(sandbox_path(private_dir, rel_path), data) + if ctx.skills and setup.skills_dir: + await copy_skills(ctx.sandbox, tuple(ctx.skills), sandbox_path(private_dir, setup.skills_dir)) + self._private_dir = private_dir + self._setup = setup + + async def turn(self, ctx: SessionContext, prompt: str) -> AsyncIterator[Event]: + if self._setup is None or self._private_dir is None: + raise RuntimeError("CLIHarnessHandler.turn() called before start()") + request = self.config.transform_turn_request(ctx, self._setup, self._private_dir, prompt, self._native_id) + argv: Final = list(request.argv) # mutable-ok: Sandbox.exec takes list[str] + proc = await ctx.sandbox.exec(argv, env=request.env, cwd=request.cwd) + self._proc = proc + tail: Final[deque[str]] = deque(maxlen=HARNESS_STDERR_TAIL_LINES) # mutable-ok: bounded stderr ring buffer + stderr_task = asyncio.ensure_future(drain_stderr(proc.stderr, tail)) + state: Any = self.config.create_stream_state() + exit_code: int | None = None + try: + await send_stdin(proc, request.stdin) + async for raw in iter_stream_lines(proc.stdout): + line = decode_json_line(raw) + if line is None: + continue + for event in self.config.transform_stream_line(line, state): + yield event + self._native_id = self.config.get_native_session_id(state) or self._native_id + exit_code = await proc.wait() + await stderr_task + finally: + self._proc = None + if exit_code is None: + # Consumer stopped early, timed out or errored: don't leave the runtime running. + await proc.kill() + if not stderr_task.done(): + stderr_task.cancel() + response = self.config.transform_turn_response(ctx, state, exit_code, tuple(tail)) + ctx.final_text = response.final_text + ctx.output_json = response.output_json + + async def stop(self, ctx: SessionContext) -> None: + proc, self._proc = self._proc, None + if proc is not None: + await proc.kill() + + def native_session_id(self) -> str | None: + return self._native_id + + async def resume(self, ctx: SessionContext, native_session_id: str) -> None: + self._native_id = native_session_id diff --git a/litellm/harness/handlers/deepagents_handler.py b/litellm/harness/handlers/deepagents_handler.py new file mode 100644 index 00000000000..3bdc158a0d1 --- /dev/null +++ b/litellm/harness/handlers/deepagents_handler.py @@ -0,0 +1,261 @@ +""" +In-process handler for Deep Agents. + +Deep Agents is a Python library, so there is no process or model endpoint: the handler +builds the agent with a LiteLLM chat model, streams the LangGraph run, turns interrupts into +Approval events and counts usage. Translation lives in +`litellm/llms/deepagents/harness/transformation.py`. +""" + +from __future__ import annotations + +import asyncio +import importlib +from collections.abc import AsyncIterator, Mapping +from dataclasses import dataclass +from types import MappingProxyType, ModuleType +from typing import TYPE_CHECKING, Any + +from litellm.harness.context import SessionContext +from litellm.harness.errors import HarnessError, HarnessInstallFailed +from litellm.harness.handlers.base import BaseHarnessHandler +from litellm.harness.handlers.cli_handler import copy_skills +from litellm.harness.options import DeepAgentsOptions +from litellm.harness.types import Approval, Event +from litellm.llms.deepagents.harness.transformation import ( + EXECUTE_TOOLS, + INSTALL_HINT, + SKILLS_DIR, + WRITE_TOOLS, + TurnState, + approval_requests, + blocked_tools, + chat_model_kwargs, + decision, + final_ai_text, + interrupt_config, + interrupts_in, + normalized_tool_name, + recursion_limit, + stream_events, + structured_json, + update_events, +) + +if TYPE_CHECKING: + from langchain_core.callbacks import BaseCallbackHandler + from langchain_core.language_models import BaseChatModel + from langchain_core.runnables import RunnableConfig + from langgraph.checkpoint.base import BaseCheckpointSaver + from langgraph.graph.state import CompiledStateGraph + from langgraph.types import Command + + from litellm.llms.base_llm.harness.transformation import BaseHarnessConfig + +_MODEL_NODE = "model" + + +@dataclass(frozen=True) +class DeepAgentsDeps: + """The optional-dependency entrypoints this handler uses.""" + + create_deep_agent: Any + chat_litellm: Any + checkpointer_cls: Any + command_cls: Any + subagent_defaults: Mapping[str, Any] + convert_to_openai_messages: Any + backend: ModuleType + + +def load_deps() -> DeepAgentsDeps: + """Import deepagents + langchain-litellm, or raise HarnessInstallFailed.""" + try: + deepagents = importlib.import_module("deepagents") + subagents = importlib.import_module("deepagents.middleware.subagents") + chat = importlib.import_module("langchain_litellm") + memory = importlib.import_module("langgraph.checkpoint.memory") + lg_types = importlib.import_module("langgraph.types") + messages = importlib.import_module("langchain_core.messages") + backend = importlib.import_module("litellm.llms.deepagents.harness.sandbox_backend") + except ImportError as e: + raise HarnessInstallFailed(f"{INSTALL_HINT} ({e})") from e + return DeepAgentsDeps( + create_deep_agent=deepagents.create_deep_agent, + chat_litellm=chat.ChatLiteLLM, + checkpointer_cls=memory.InMemorySaver, + command_cls=lg_types.Command, + subagent_defaults=subagents.GENERAL_PURPOSE_SUBAGENT, + convert_to_openai_messages=messages.convert_to_openai_messages, + backend=backend, + ) + + +_SHARED_CHECKPOINTER: dict[str, Any] = {} # mutable-ok: process-wide lazy singleton slot for the in-memory checkpointer + + +def shared_checkpointer(deps: DeepAgentsDeps) -> BaseCheckpointSaver: + """One in-memory checkpointer per process, so resume() works across sessions in-process.""" + saver = _SHARED_CHECKPOINTER.get("saver") + if saver is None: + saver = deps.checkpointer_cls() + _SHARED_CHECKPOINTER["saver"] = saver + return saver + + +def build_chat_model(ctx: SessionContext, deps: DeepAgentsDeps) -> BaseChatModel: + """The LangChain chat model for this session. Tests monkeypatch this.""" + return deps.chat_litellm(**chat_model_kwargs(ctx)) + + +class DeepAgentsHandler(BaseHarnessHandler): + def __init__(self, config: BaseHarnessConfig) -> None: + super().__init__(config) + self._deps: DeepAgentsDeps | None = None + self._agent: Any = None + self._thread_id: str | None = None + self._skip_tools: frozenset[str] = frozenset() + + async def start(self, ctx: SessionContext) -> None: + self.config.validate_environment(ctx) + deps = load_deps() + self._deps = deps + blocked = blocked_tools(ctx.permissions, ctx.disable_tools) + backend = deps.backend.SandboxBackend( + ctx.sandbox, + loop=asyncio.get_running_loop(), + writable=WRITE_TOOLS.isdisjoint(blocked), + allow_execute=EXECUTE_TOOLS.isdisjoint(blocked), + ) + self._agent = deps.create_deep_agent( + model=build_chat_model(ctx, deps), + tools=list(ctx.tools), # mutable-ok: deepagents create_deep_agent(tools=) takes a list + system_prompt=ctx.instructions, + middleware=self._middleware(deps, blocked), + subagents=self._subagents(ctx, deps, blocked), + skills=await self._install_skills(ctx), + backend=backend, + interrupt_on=interrupt_config(ctx.permissions, blocked), + response_format=ctx.output, + checkpointer=shared_checkpointer(deps), + ) + self._skip_tools = frozenset({ctx.output.__name__}) if ctx.output is not None else frozenset() + if self._thread_id is None: + self._thread_id = ctx.session_id + + async def stop(self, ctx: SessionContext) -> None: + self._agent = None + + def native_session_id(self) -> str | None: + return self._thread_id + + async def resume(self, ctx: SessionContext, native_session_id: str) -> None: + self._thread_id = native_session_id + + async def history( + self, ctx: SessionContext + ) -> list[dict[str, Any]]: # mutable-ok: BaseHarnessHandler.history API returns OpenAI message dicts + agent, deps = self._require_agent() + snapshot = await agent.aget_state(self._run_config(ctx, None)) + messages = (snapshot.values or MappingProxyType({})).get("messages") or () + converted: list[dict[str, Any]] = deps.convert_to_openai_messages( # mutable-ok: LangChain returns a list + messages + ) + return converted + + async def turn(self, ctx: SessionContext, prompt: str) -> AsyncIterator[Event]: + agent, deps = self._require_agent() + run_config = self._run_config(ctx, deps.backend.UsageCallback(ctx, ctx.model)) + user_message = {"role": "user", "content": prompt} # mutable-ok: LangGraph input message dict + payload: dict[str, object] | Command = {"messages": [user_message]} # mutable-ok: LangGraph input state + while True: + state = TurnState() + async for event in self._stream_pass(agent, payload, run_config, state): + yield event + if not state.interrupts: + break + resume: dict[str, Any] = {} # mutable-ok: Command(resume=) payload, filled per answered approval + for interrupt in state.interrupts: + decisions: list[dict[str, Any]] = [] # mutable-ok: HITL decisions collected across awaited approvals + for request in approval_requests(getattr(interrupt, "value", None)): + approval = Approval( + tool=normalized_tool_name(str(request.get("name") or "")), + input=dict(request.get("args") or ()), # mutable-ok: Approval.input is a public dict field + ) + yield approval + decisions.append(decision(*await approval.wait())) + resume[interrupt.id] = {"decisions": decisions} # mutable-ok: LangGraph HITL resume payload + payload = deps.command_cls(resume=resume) + await self._finish_turn(ctx, agent, run_config) + + async def _stream_pass( + self, + agent: CompiledStateGraph, + payload: dict[str, object] | Command, # mutable-ok: LangGraph astream input type + run_config: RunnableConfig, + state: TurnState, + ) -> AsyncIterator[Event]: + async for part in agent.astream( + payload, + run_config, + stream_mode=["messages", "updates"], # mutable-ok: LangGraph stream_mode takes a list + ): + # A list stream_mode yields (mode, chunk) tuples; LangGraph's overloads don't say so. + if not isinstance(part, tuple) or len(part) != 2: + continue + mode, chunk = part + if mode == "messages": + message, meta = chunk + if isinstance(meta, Mapping) and meta.get("langgraph_node") == _MODEL_NODE: + for event in stream_events(message): + yield event + elif mode == "updates": + state.interrupts = (*state.interrupts, *interrupts_in(chunk)) + for event in update_events(chunk, self._skip_tools): + yield event + + async def _finish_turn(self, ctx: SessionContext, agent: CompiledStateGraph, run_config: RunnableConfig) -> None: + snapshot = await agent.aget_state(run_config) + values = snapshot.values or MappingProxyType({}) + ctx.final_text = final_ai_text(values.get("messages") or ()) + if ctx.output is not None: + ctx.output_json = structured_json(values.get("structured_response")) + + def _require_agent(self) -> tuple[Any, DeepAgentsDeps]: + if self._agent is None or self._deps is None: + raise HarnessError("Deep Agents session is not started") + return self._agent, self._deps + + def _run_config(self, ctx: SessionContext, usage_callback: BaseCallbackHandler | None) -> RunnableConfig: + run_config: RunnableConfig = { + "configurable": {"thread_id": self._thread_id or ctx.session_id}, + "recursion_limit": recursion_limit(ctx), + } + if usage_callback is not None: + run_config["callbacks"] = [usage_callback] # mutable-ok: LangChain RunnableConfig.callbacks is a list + return run_config + + @staticmethod + def _middleware(deps: DeepAgentsDeps, blocked: frozenset[str]) -> list[Any]: # mutable-ok: deepagents API + filters = (deps.backend.ToolFilterMiddleware(blocked),) if blocked else () + return list(filters) # mutable-ok: deepagents create_deep_agent(middleware=) takes a list + + def _subagents( + self, ctx: SessionContext, deps: DeepAgentsDeps, blocked: frozenset[str] + ) -> list[Any]: # mutable-ok: deepagents create_deep_agent(subagents=) takes a list + """User subagents, plus a general-purpose one that honours disable_tools when set.""" + options = ctx.options if isinstance(ctx.options, DeepAgentsOptions) else None + user_subagents = tuple(options.subagents) if options is not None else () + has_general = any( + isinstance(s, Mapping) and s.get("name") == deps.subagent_defaults["name"] for s in user_subagents + ) + spec = {**deps.subagent_defaults, "middleware": self._middleware(deps, blocked)} # mutable-ok: SubAgent dict + general = (spec,) if blocked and not has_general else () + return [*general, *user_subagents] # mutable-ok: deepagents create_deep_agent(subagents=) takes a list + + @staticmethod + async def _install_skills(ctx: SessionContext) -> list[str] | None: # mutable-ok: deepagents skills= takes a list + if not ctx.skills: + return None + await copy_skills(ctx.sandbox, ctx.skills, f"{ctx.sandbox.workdir}/{SKILLS_DIR}") + return [f"/{SKILLS_DIR}/"] # mutable-ok: deepagents create_deep_agent(skills=) takes a list diff --git a/litellm/harness/options.py b/litellm/harness/options.py new file mode 100644 index 00000000000..2e3ec5d0fd5 --- /dev/null +++ b/litellm/harness/options.py @@ -0,0 +1,37 @@ +"""Typed per-harness options. Settings that only make sense for one runtime live here.""" + +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from dataclasses import dataclass, field +from typing import Any, Literal + + +@dataclass(frozen=True) +class ClaudeCodeOptions: + config: Mapping[str, Any] = field(default_factory=dict) + env: Mapping[str, str] = field(default_factory=dict) + + +@dataclass(frozen=True) +class CodexOptions: + reasoning_effort: Literal["low", "medium", "high", "xhigh"] | None = None + web_search: bool = False + config: Mapping[str, Any] = field(default_factory=dict) + env: Mapping[str, str] = field(default_factory=dict) + + +@dataclass(frozen=True) +class OpenCodeOptions: + agent: str = "build" + config: Mapping[str, Any] = field(default_factory=dict) + env: Mapping[str, str] = field(default_factory=dict) + + +@dataclass(frozen=True) +class DeepAgentsOptions: + subagents: Sequence[Any] = () + recursion_limit: int | None = None + + +HarnessOptions = ClaudeCodeOptions | CodexOptions | OpenCodeOptions | DeepAgentsOptions diff --git a/litellm/harness/runtime.py b/litellm/harness/runtime.py new file mode 100644 index 00000000000..53a06e7d50d --- /dev/null +++ b/litellm/harness/runtime.py @@ -0,0 +1,1070 @@ +"""The harness engine: validation, sessions, turns, approvals, files, usage and results. + +Adapters only translate a runtime's native protocol into events. Everything that must behave +the same across harnesses (timeouts, max_turns, approvals, FileChange, structured output, +usage and cost) lives here. +""" + +from __future__ import annotations + +import asyncio +import inspect +import logging +import os +import uuid +from collections.abc import AsyncIterator, Callable, Coroutine, Generator, Mapping, Sequence +from dataclasses import dataclass, field +from types import MappingProxyType +from typing import ( + Any, + Final, + get_args, +) + +from pydantic import BaseModel, ValidationError + +import litellm +from litellm.constants import HARNESS_EVENT_QUEUE_MAX_SIZE +from litellm.harness.context import ApprovalHandler, GatewayTarget, SessionContext +from litellm.harness.endpoint import ModelEndpoint +from litellm.harness.errors import ( + CapabilityUnsupported, + HarnessError, + HarnessInstallFailed, + OptionsMismatch, + OutputInvalid, + SessionClosed, + StateIncompatible, +) +from litellm.harness.handlers import get_harness_config, get_harness_handler +from litellm.harness.handlers.base import BaseHarnessHandler +from litellm.harness.options import HarnessOptions +from litellm.harness.sandbox.base import Sandbox +from litellm.harness.sandbox.snapshot import build_file_changes, capture_text_contents +from litellm.harness.types import ( + Approval, + Capabilities, + Done, + Event, + FileChange, + Harness, + PermissionMode, + Result, + State, + StopReason, + Text, + ToolCall, + Usage, + require_harness, +) +from litellm.llms.base_llm.harness.transformation import BaseHarnessConfig +from litellm.llms.base_llm.harness.utils import last_json_object + +PERMISSION_MODES: Final = frozenset(get_args(PermissionMode)) +SKILL_FILE: Final = "SKILL.md" +# Adapter errors that mean "misconfigured", not "the runtime crashed": re-raised to the caller. +verbose_logger: Final = logging.getLogger("LiteLLM") + +PROPAGATED_ERRORS: Final = (HarnessInstallFailed, CapabilityUnsupported) + + +# --------------------------------------------------------------------------- +# Configuration + validation +# --------------------------------------------------------------------------- + + +@dataclass(frozen=True) +class SessionConfig: + """Every per-session parameter a caller can pass, already normalized.""" + + harness: Harness + sandbox: Sandbox + model: str | None = None + gateway: GatewayTarget | None = None + api_key: str | None = None + api_base: str | None = None + instructions: str | None = None + tools: Sequence[Callable[..., Any]] = () + skills: Sequence[str] = () + disable_tools: Sequence[str] = () + permissions: PermissionMode = "full" + on_approval: ApprovalHandler | None = None + output: type[BaseModel] | None = None + max_turns: int | None = None + timeout: float | None = None + metadata: Mapping[str, Any] = field(default_factory=dict) + options: HarnessOptions | None = None + install: bool = False + + +LITELLM_PROXY_PREFIX: Final = "litellm_proxy/" + + +def resolve_model_route( + model: str | None, api_key: str | None, api_base: str | None +) -> tuple[str | None, GatewayTarget | None]: + """(model sent to the runtime, gateway or None). + + `litellm_proxy/` (or `litellm.use_litellm_proxy = True`) routes every model call + through the LiteLLM AI Gateway, using api_base/api_key or LITELLM_PROXY_API_BASE / + LITELLM_PROXY_API_KEY. Anything else is called directly through the LiteLLM SDK. + """ + prefixed = model is not None and model.startswith(LITELLM_PROXY_PREFIX) + if not prefixed and not litellm.use_litellm_proxy: + return model, None + group = model[len(LITELLM_PROXY_PREFIX) :] if prefixed and model is not None else model + base = (api_base or os.environ.get("LITELLM_PROXY_API_BASE") or "").strip() + key = (api_key or os.environ.get("LITELLM_PROXY_API_KEY") or "").strip() + if not base: + raise ValueError("litellm_proxy/ models need the gateway URL: pass api_base= or set LITELLM_PROXY_API_BASE") + if not key: + raise ValueError("litellm_proxy/ models need a gateway virtual key: pass api_key= or set LITELLM_PROXY_API_KEY") + return group, GatewayTarget(api_base=base.rstrip("/"), api_key=key) + + +def _normalize_skill(skill: str | os.PathLike[str]) -> str: + path = os.path.abspath(os.fspath(skill)) + if not os.path.isfile(os.path.join(path, SKILL_FILE)): + raise ValueError(f"Skill folder {path!r} has no {SKILL_FILE}") + return path + + +def _normalize_skills(skills: Sequence[str | os.PathLike[str]]) -> tuple[str, ...]: + return tuple(_normalize_skill(skill) for skill in skills) + + +def _check_basic(config: SessionConfig) -> None: + if config.permissions not in PERMISSION_MODES: + raise ValueError(f"permissions must be one of {sorted(PERMISSION_MODES)}, got {config.permissions!r}") + if config.max_turns is not None and config.max_turns < 1: + raise ValueError("max_turns must be >= 1") + if config.timeout is not None and config.timeout <= 0: + raise ValueError("timeout must be > 0") + if config.install: + raise CapabilityUnsupported("install=True is not supported yet; put the runtime binary on PATH in the sandbox") + + +def _check_options(config: SessionConfig, harness_config: BaseHarnessConfig) -> None: + if config.options is None or isinstance(config.options, harness_config.options_type): + return + raise OptionsMismatch( + f"{type(config.options).__name__} cannot be used with Harness.{config.harness.name}; " + f"use {harness_config.options_type.__name__}" + ) + + +def _check_capabilities(config: SessionConfig, caps: Capabilities, interactive: bool) -> None: + name = f"Harness.{config.harness.name}" + if config.permissions not in caps.permission_modes: + raise CapabilityUnsupported( + f"{name} does not support permissions={config.permissions!r}; supported: {sorted(caps.permission_modes)}" + ) + if config.permissions == "ask": + if not caps.tool_approval: + raise CapabilityUnsupported(f"{name} does not support tool approvals") + if config.on_approval is None and not interactive: + raise ValueError("permissions='ask' needs on_approval=, or use stream() and answer Approval events") + if config.output is not None and not caps.structured_output: + raise CapabilityUnsupported(f"{name} does not support output=") + if config.tools and not caps.custom_tools: + raise CapabilityUnsupported(f"{name} does not support custom tools=") + if config.skills and not caps.skills: + raise CapabilityUnsupported(f"{name} does not support skills=") + if config.disable_tools and not caps.tool_filtering: + raise CapabilityUnsupported(f"{name} does not support disable_tools=") + + +def validate(config: SessionConfig, harness_config: BaseHarnessConfig, interactive: bool) -> None: + """Raise before anything starts if the request cannot be served.""" + _check_basic(config) + _check_options(config, harness_config) + _check_capabilities(config, harness_config.capabilities, interactive) + + +def build_config( + harness: Harness, + *, + sandbox: Sandbox, + model: str | None = None, + api_key: str | None = None, + api_base: str | None = None, + instructions: str | None = None, + tools: Sequence[Callable[..., Any]] = (), + skills: Sequence[str | os.PathLike[str]] = (), + disable_tools: Sequence[str] = (), + permissions: PermissionMode = "full", + on_approval: ApprovalHandler | None = None, + output: type[BaseModel] | None = None, + max_turns: int | None = None, + timeout: float | None = None, + metadata: Mapping[str, Any] | None = None, + options: HarnessOptions | None = None, + install: bool = False, +) -> SessionConfig: + """Normalize public keyword arguments into a SessionConfig.""" + resolved_harness = require_harness(harness) + routed_model, gateway = resolve_model_route(model, api_key, api_base) + return SessionConfig( + harness=resolved_harness, + sandbox=sandbox, + model=routed_model, + gateway=gateway, + api_key=api_key, + api_base=api_base, + instructions=instructions, + tools=tuple(tools), + skills=tuple(_normalize_skills(skills)), + disable_tools=tuple(disable_tools), + permissions=permissions, + on_approval=on_approval, + output=output, + max_turns=max_turns, + timeout=timeout, + metadata=MappingProxyType(dict(metadata or ())), + options=options, + install=install, + ) + + +def _context_for(config: SessionConfig) -> SessionContext: + return SessionContext( + harness=config.harness, + sandbox=config.sandbox, + session_id=uuid.uuid4().hex, + model=config.model, + gateway=config.gateway, + api_key=config.api_key, + api_base=config.api_base, + instructions=config.instructions, + tools=config.tools, + skills=config.skills, + disable_tools=config.disable_tools, + permissions=config.permissions, + on_approval=config.on_approval, + output=config.output, + max_turns=config.max_turns, + timeout=config.timeout, + metadata=config.metadata, + options=config.options, + ) + + +# --------------------------------------------------------------------------- +# Structured output +# --------------------------------------------------------------------------- + + +def parse_output(output: type[BaseModel], output_json: str | None, text: str) -> tuple[BaseModel | None, str | None]: + """Return (model, None) on success or (None, error message) on failure.""" + raw = output_json or last_json_object(text) + if raw is None: + return None, "no JSON object found in the final answer" + try: + return output.model_validate_json(raw), None + except ValidationError as e: + return None, str(e) + + +# --------------------------------------------------------------------------- +# Turn machinery +# --------------------------------------------------------------------------- + + +@dataclass +class _End: + """Sentinel the producer puts on the queue when the handler turn is over.""" + + reason: StopReason | None = None + error: BaseException | None = None + + +class TurnControl: + """Lets a stream consumer cancel the running turn.""" + + def __init__(self) -> None: + self.cancelled = False + self.producer: asyncio.Task[None] | None = None + + def cancel(self) -> None: + self.cancelled = True + if self.producer is not None and not self.producer.done(): + self.producer.cancel() + + +async def _aclose(events: AsyncIterator[Event]) -> None: + closer = getattr(events, "aclose", None) + if closer is None: + return + try: + await closer() + except Exception: # closing must not mask the turn's own outcome + verbose_logger.debug("harness: error closing handler turn", exc_info=True) + + +async def pump_events( + events: AsyncIterator[Event], + queue: asyncio.Queue[Event | _End], + max_turns: int | None, +) -> None: + """Drive the handler turn in one task, enforcing max_turns on ToolCall events.""" + end = _End() + tool_calls = 0 + try: + async for event in events: + if isinstance(event, ToolCall): + tool_calls += 1 + if max_turns is not None and tool_calls > max_turns: + end = _End(reason="max_turns") + break + # Backpressure: a runtime that streams faster than the consumer waits here. + await queue.put(event) + except asyncio.CancelledError: + end = _End(reason="cancelled") + raise + except Exception as e: # any runtime failure becomes stop_reason="runtime_error" (see _Turn._finish) + verbose_logger.debug("harness: handler turn raised", exc_info=True) + end = _End(error=e) + finally: + await _aclose(events) + await _put_end(queue, end) + + +async def _put_end(queue: asyncio.Queue[Event | _End], end: _End) -> None: + """Queue the end marker behind every event, waiting for room so no event is dropped. + + A cancelled turn has no consumer left to drain the queue, so only then is space made + by discarding queued events. + """ + if end.reason != "cancelled": + try: + await queue.put(end) + return + except asyncio.CancelledError: + pass + while queue.full(): + queue.get_nowait() + queue.put_nowait(end) + + +async def call_approval_handler(handler: ApprovalHandler, approval: Approval) -> None: + """Run on_approval (sync in a worker thread, or async) and resolve approval.""" + try: + if inspect.iscoroutinefunction(handler): + decision: Any = await handler(approval) + else: + decision = await asyncio.to_thread(handler, approval) + if inspect.isawaitable(decision): + decision = await decision + except Exception as e: # a failing user callback denies the tool instead of crashing the turn + verbose_logger.warning("harness: on_approval raised for tool %s; denying", approval.tool, exc_info=True) + approval.deny(f"on_approval raised: {e}") + return + if decision: + approval.allow() + else: + approval.deny("denied by on_approval") + + +class _Turn: + """One prompt -> events -> Done cycle on a started session.""" + + def __init__( + self, + session: AsyncSession, + prompt: str, + control: TurnControl, + interactive: bool, + ) -> None: + self.session = session + self.ctx = session.ctx + self.prompt = prompt + self.control = control + self.interactive = interactive + self.queue: asyncio.Queue[Event | _End] = asyncio.Queue(maxsize=HARNESS_EVENT_QUEUE_MAX_SIZE) + self.events: list[Event] = [] # mutable-ok: per-turn accumulator the runtime appends events to + self.text_parts: list[str] = [] # mutable-ok: per-turn accumulator of streamed text deltas + self.emitted_files: set[tuple[str, str]] = set() # mutable-ok: per-turn record of emitted FileChanges + self.approval_tasks: list[asyncio.Future[None]] = [] # mutable-ok: per-turn in-flight approval tasks + self.stop_reason: StopReason = "done" + self.error_text: str | None = None + self.before: Mapping[str, str] = MappingProxyType({}) + self.before_contents: Mapping[str, bytes] = MappingProxyType({}) + self.usage_before: tuple[int, int, int, float] = (0, 0, 0, 0.0) + self.deadline: float | None = None + + # -- setup / teardown --------------------------------------------------- + + async def _begin(self) -> None: + sandbox = self.ctx.sandbox + self.before = await sandbox.snapshot() + self.before_contents = await capture_text_contents(sandbox, self.before) + self.usage_before = self.session.usage_counters() + self.ctx.final_text = "" + self.ctx.output_json = None + if self.ctx.timeout is not None: + self.deadline = asyncio.get_running_loop().time() + self.ctx.timeout + + def _start_producer(self) -> None: + events = self.session.handler.turn(self.ctx, self.prompt) + self.control.producer = asyncio.ensure_future(pump_events(events, self.queue, self.ctx.max_turns)) + if self.control.cancelled: + self.control.producer.cancel() + + async def _stop_producer(self) -> None: + producer = self.control.producer + live_producer = (producer,) if producer is not None and not producer.done() else () + pending = (*(task for task in self.approval_tasks if not task.done()), *live_producer) + for task in pending: + task.cancel() + if pending: + await asyncio.wait(pending) + + # -- event loop --------------------------------------------------------- + + async def _next_item(self) -> Event | _End: + if self.deadline is None: + return await self.queue.get() + remaining = self.deadline - asyncio.get_running_loop().time() + try: + if remaining <= 0: + raise asyncio.TimeoutError + return await asyncio.wait_for(self.queue.get(), remaining) + except asyncio.TimeoutError: + await self._stop_producer() + return _End(reason="timeout") + + def _finish(self, end: _End) -> None: + if end.error is not None: + if isinstance(end.error, PROPAGATED_ERRORS): + raise end.error + self.stop_reason = "runtime_error" + self.error_text = f"{type(end.error).__name__}: {end.error}" + verbose_logger.warning("harness %s runtime error: %s", self.ctx.harness.value, self.error_text) + return + if self.control.cancelled: + self.stop_reason = "cancelled" + elif end.reason is not None: + self.stop_reason = end.reason + + async def _on_approval(self, approval: Approval) -> None: + handler = self.ctx.on_approval + if handler is not None: + self.approval_tasks.append(asyncio.ensure_future(call_approval_handler(handler, approval))) + elif not self.interactive: + approval.deny("no approval handler") + + async def _record(self, event: Event) -> None: + if isinstance(event, Text): + self.text_parts.append(event.delta) + elif isinstance(event, FileChange): + self.emitted_files.add((event.path, event.kind)) + elif isinstance(event, Approval): + await self._on_approval(event) + self.events.append(event) + + async def _drain(self) -> AsyncIterator[Event]: + while True: + item = await self._next_item() + if isinstance(item, _End): + self._finish(item) + return + if isinstance(item, Done): + continue + await self._record(item) + yield item + if isinstance(item, Approval) and self.ctx.on_approval is None: + # The consumer asked for the next event without answering. + item.deny("approval not answered") + + # -- results ------------------------------------------------------------ + + async def _file_changes(self) -> list[FileChange]: # mutable-ok: becomes the public Result.files list + sandbox = self.ctx.sandbox + after = await sandbox.snapshot() + files = await build_file_changes(sandbox, self.before, after, self.before_contents) + seen = { # mutable-ok: dedupe set grown while merging streamed FileChange events + change.path for change in files + } + for event in self.events: + if isinstance(event, FileChange) and event.path not in seen: + files.append(event) + seen.add(event.path) + return files + + def _text(self) -> str: + text = self.ctx.final_text or "".join(self.text_parts) + if self.error_text is None: + return text + return f"{text}\n\n{self.error_text}" if text else self.error_text + + def _usage(self) -> tuple[Usage, float]: + now = self.session.usage_counters() + before = self.usage_before + usage = Usage( + input_tokens=now[0] - before[0], + output_tokens=now[1] - before[1], + calls=now[2] - before[2], + ) + return usage, max(now[3] - before[3], 0.0) + + def _result( + self, + files: list[FileChange], # mutable-ok: Result.files is a public list field + output: BaseModel | None, + ) -> Result: + usage, cost = self._usage() + return Result( + text=self._text(), + output=output, + files=files, + events=list( # mutable-ok: Result.events is a public list field; copy detaches it from the accumulator + self.events + ), + usage=usage, + cost=cost, + stop_reason=self.stop_reason, + session_id=self.ctx.session_id, + ) + + def _output(self) -> tuple[BaseModel | None, str | None, str | None]: + """(parsed output, raw text, error) for the structured-output check.""" + output_type = self.ctx.output + if output_type is None or self.stop_reason != "done": + return None, None, None + text = self._text() + parsed, error = parse_output(output_type, self.ctx.output_json, text) + return parsed, self.ctx.output_json or text, error + + # -- entry -------------------------------------------------------------- + + async def run(self) -> AsyncIterator[Event]: + await self._begin() + self._start_producer() + try: + async for event in self._drain(): + yield event + finally: + await self._stop_producer() + if self.stop_reason != "done": + await self.session.interrupt() + files = await self._file_changes() + for change in files: + if (change.path, change.kind) not in self.emitted_files: + self.events.append(change) + yield change + parsed, raw, error = self._output() + result = self._result(files, parsed) + self.session.record(result) + yield Done(result) + if error is not None: + raise OutputInvalid( + f"Final answer did not match {self.ctx.output.__name__ if self.ctx.output else 'output'}: {error}", + raw=raw or "", + result=result, + ) + + +# --------------------------------------------------------------------------- +# Streams +# --------------------------------------------------------------------------- + + +class AsyncEventStream: + """Async iterator of events for one turn. `.result` is set once Done is seen.""" + + def __init__(self, source: AsyncIterator[Event], control: TurnControl) -> None: + self._source = source + self._control = control + self._result: Result | None = None + + def __aiter__(self) -> AsyncEventStream: + return self + + async def __anext__(self) -> Event: + event = await self._source.__anext__() + if isinstance(event, Done): + self._result = event.result + return event + + @property + def result(self) -> Result | None: + return self._result + + def cancel(self) -> None: + """Stop the turn. The stream still ends with Done(stop_reason='cancelled').""" + self._control.cancel() + + async def aclose(self) -> None: + await _aclose(self._source) + + +async def _one_shot(session: AsyncSession, prompt: str, control: TurnControl) -> AsyncIterator[Event]: + """Stream one turn on a fresh session and close it before Done is handed out.""" + try: + async for event in session.turn_events(prompt, control, interactive=True): + if isinstance(event, Done): + await session.aclose() + yield event + finally: + await session.aclose() + + +# --------------------------------------------------------------------------- +# Sessions +# --------------------------------------------------------------------------- + + +class AsyncSession: + """A multi-turn conversation with one harness. Use `async with` or `await`.""" + + def __init__( + self, + config: SessionConfig, + *, + resume_from: str | None = None, + interactive: bool = True, + ) -> None: + self.config = config + self.harness_config = get_harness_config(config.harness) + validate(config, self.harness_config, interactive=interactive) + if resume_from is not None and not self.harness_config.capabilities.resume: + raise CapabilityUnsupported(f"Harness.{config.harness.name} does not support resume") + self.ctx = _context_for(config) + # Config-specific static checks (managed option keys, required model) before any I/O. + self.harness_config.validate_environment(self.ctx) + self.handler: BaseHarnessHandler = get_harness_handler(self.harness_config) + self.results: list[Result] = [] # mutable-ok: session accumulator; each turn's Result is appended + self._resume_from = resume_from + self._native_id: str | None = resume_from + self._started = False + self._closed = False + self._busy = False + self._restart_needed = False + + # -- lifecycle ---------------------------------------------------------- + + def __await__(self) -> Generator[object, None, AsyncSession]: + return self.start().__await__() + + async def __aenter__(self) -> AsyncSession: + return await self.start() + + async def __aexit__(self, *exc_info: object) -> None: + await self.aclose() + + async def _open_endpoint(self) -> None: + if not self.harness_config.uses_model_endpoint or self.ctx.endpoint is not None: + return + endpoint = ModelEndpoint( + self.config.harness, + self.config.model, + self.config.gateway, + api_key=self.config.api_key, + api_base=self.config.api_base, + metadata=self.config.metadata, + ) + await endpoint.__aenter__() + self.ctx.endpoint = endpoint + + async def _launch(self) -> None: + await self.handler.start(self.ctx) + if self._native_id is not None and (self._resume_from is not None or self._restart_needed): + await self.handler.resume(self.ctx, self._native_id) + + async def start(self) -> AsyncSession: + if self._closed: + raise SessionClosed("session is closed") + if self._started: + return self + await self._open_endpoint() + try: + await self._launch() + except BaseException: + await self._close_endpoint() + raise + self._started = True + return self + + async def interrupt(self) -> None: + """Stop the runtime after a timeout / max_turns / cancel; next turn restarts it.""" + self._native_id = self.handler.native_session_id() or self._native_id + try: + await self.handler.stop(self.ctx) + except Exception: # the next turn restarts the runtime regardless + verbose_logger.warning("harness: handler stop failed", exc_info=True) + self._restart_needed = True + + async def _ensure_ready(self) -> None: + if self._closed: + raise SessionClosed("session is closed") + if not self._started: + await self.start() + elif self._restart_needed: + await self._launch() + self._restart_needed = False + + async def _close_endpoint(self) -> None: + endpoint = self.ctx.endpoint + self.ctx.endpoint = None + if endpoint is None: + return + try: + await endpoint.__aexit__(None, None, None) + except Exception: # shutdown is best-effort cleanup + verbose_logger.warning("harness: endpoint shutdown failed", exc_info=True) + + async def aclose(self) -> None: + """Stop the runtime and the endpoint. Safe to call twice.""" + if self._closed: + return + self._closed = True + if self._started: + self._native_id = self.handler.native_session_id() or self._native_id + try: + await self.handler.stop(self.ctx) + except Exception: # still close the endpoint below + verbose_logger.warning("harness: handler stop failed", exc_info=True) + await self._close_endpoint() + + close = aclose + + def state(self) -> State: + native = self._native_id + if self._started and not self._closed: + native = self.handler.native_session_id() or native + return State( + harness=self.config.harness, + native_session_id=native, + workdir=self.config.sandbox.workdir, + model=self.config.model, + ) + + async def adetach(self) -> State: + """Release local resources and return State to resume() later.""" + await self.aclose() + return self.state() + + async def astop(self) -> State: + """Stop the session for good and return its final State.""" + await self.aclose() + return self.state() + + detach = adetach + stop = astop + + # -- turns -------------------------------------------------------------- + + def usage_counters(self) -> tuple[int, int, int, float]: + """(input_tokens, output_tokens, calls, cost) so far, from endpoint or handler.""" + endpoint = self.ctx.endpoint + if endpoint is not None: + usage = endpoint.usage + return (usage.input_tokens, usage.output_tokens, usage.calls, usage.cost) + ctx = self.ctx + return (ctx.input_tokens, ctx.output_tokens, ctx.calls, ctx.cost) + + def record(self, result: Result) -> None: + self.results.append(result) + + async def turn_events(self, prompt: str, control: TurnControl, interactive: bool) -> AsyncIterator[Event]: + if self._busy: + raise HarnessError("a turn is already running on this session") + self._busy = True + try: + await self._ensure_ready() + async for event in _Turn(self, prompt, control, interactive).run(): + yield event + finally: + self._busy = False + + def astream(self, prompt: str) -> AsyncEventStream: + control = TurnControl() + return AsyncEventStream(self.turn_events(prompt, control, interactive=True), control) + + async def arun(self, prompt: str) -> Result: + return await _collect(self.turn_events(prompt, TurnControl(), False)) + + async def history( + self, + ) -> list[dict[str, Any]]: # mutable-ok: public API returns OpenAI-format message dicts from the handler + if not self.harness_config.capabilities.history: + raise CapabilityUnsupported(f"Harness.{self.config.harness.name} does not expose history") + await self._ensure_ready() + return await self.handler.history(self.ctx) + + @property + def cost(self) -> float: + return sum(result.cost for result in self.results) + + @property + def usage(self) -> Usage: + return Usage( + input_tokens=sum(r.usage.input_tokens for r in self.results), + output_tokens=sum(r.usage.output_tokens for r in self.results), + calls=sum(r.usage.calls for r in self.results), + ) + + @property + def session_id(self) -> str: + return self.ctx.session_id + + @property + def closed(self) -> bool: + return self._closed + + +async def _collect(events: AsyncIterator[Event]) -> Result: + result: Result | None = None + async for event in events: + if isinstance(event, Done): + result = event.result + if result is None: + raise HarnessError("turn ended without a result") + return result + + +# --------------------------------------------------------------------------- +# Public async API +# --------------------------------------------------------------------------- + + +def aagent_session( + harness: Harness, + *, + sandbox: Sandbox, + model: str | None = None, + api_key: str | None = None, + api_base: str | None = None, + instructions: str | None = None, + tools: Sequence[Callable[..., Any]] = (), + skills: Sequence[str | os.PathLike[str]] = (), + disable_tools: Sequence[str] = (), + permissions: PermissionMode = "full", + on_approval: ApprovalHandler | None = None, + output: type[BaseModel] | None = None, + max_turns: int | None = None, + timeout: float | None = None, + metadata: Mapping[str, Any] | None = None, + options: HarnessOptions | None = None, + install: bool = False, +) -> AsyncSession: + """A multi-turn agent session: `async with litellm.aagent_session(...) as s:`.""" + config = build_config( + harness, + sandbox=sandbox, + model=model, + api_key=api_key, + api_base=api_base, + instructions=instructions, + tools=tools, + skills=skills, + disable_tools=disable_tools, + permissions=permissions, + on_approval=on_approval, + output=output, + max_turns=max_turns, + timeout=timeout, + metadata=metadata, + options=options, + install=install, + ) + return AsyncSession(config) + + +async def arun_agent( + harness: Harness, + prompt: str, + *, + sandbox: Sandbox, + model: str | None = None, + api_key: str | None = None, + api_base: str | None = None, + instructions: str | None = None, + tools: Sequence[Callable[..., Any]] = (), + skills: Sequence[str | os.PathLike[str]] = (), + disable_tools: Sequence[str] = (), + permissions: PermissionMode = "full", + on_approval: ApprovalHandler | None = None, + output: type[BaseModel] | None = None, + max_turns: int | None = None, + timeout: float | None = None, + metadata: Mapping[str, Any] | None = None, + options: HarnessOptions | None = None, + install: bool = False, +) -> Result: + """Run one prompt to completion and return the Result.""" + config = build_config( + harness, + sandbox=sandbox, + model=model, + api_key=api_key, + api_base=api_base, + instructions=instructions, + tools=tools, + skills=skills, + disable_tools=disable_tools, + permissions=permissions, + on_approval=on_approval, + output=output, + max_turns=max_turns, + timeout=timeout, + metadata=metadata, + options=options, + install=install, + ) + async with AsyncSession(config, interactive=False) as session: + return await session.arun(prompt) + + +def astream_agent( + harness: Harness, + prompt: str, + *, + sandbox: Sandbox, + model: str | None = None, + api_key: str | None = None, + api_base: str | None = None, + instructions: str | None = None, + tools: Sequence[Callable[..., Any]] = (), + skills: Sequence[str | os.PathLike[str]] = (), + disable_tools: Sequence[str] = (), + permissions: PermissionMode = "full", + on_approval: ApprovalHandler | None = None, + output: type[BaseModel] | None = None, + max_turns: int | None = None, + timeout: float | None = None, + metadata: Mapping[str, Any] | None = None, + options: HarnessOptions | None = None, + install: bool = False, +) -> AsyncEventStream: + """Stream events for one prompt. Validation errors raise here, before iteration.""" + session = aagent_session( + harness, + sandbox=sandbox, + model=model, + api_key=api_key, + api_base=api_base, + instructions=instructions, + tools=tools, + skills=skills, + disable_tools=disable_tools, + permissions=permissions, + on_approval=on_approval, + output=output, + max_turns=max_turns, + timeout=timeout, + metadata=metadata, + options=options, + install=install, + ) + control = TurnControl() + return AsyncEventStream(_one_shot(session, prompt, control), control) + + +def _coerce_state(state: State | bytes) -> State: + if isinstance(state, (bytes, bytearray)): + return State.loads(bytes(state)) + if not isinstance(state, State): + raise TypeError(f"state must be a State or bytes, got {type(state).__name__}") + return state + + +def aagent_resume( + state: State | bytes, + *, + sandbox: Sandbox, + model: str | None = None, + api_key: str | None = None, + api_base: str | None = None, + instructions: str | None = None, + tools: Sequence[Callable[..., Any]] = (), + skills: Sequence[str | os.PathLike[str]] = (), + disable_tools: Sequence[str] = (), + permissions: PermissionMode = "full", + on_approval: ApprovalHandler | None = None, + output: type[BaseModel] | None = None, + max_turns: int | None = None, + timeout: float | None = None, + metadata: Mapping[str, Any] | None = None, + options: HarnessOptions | None = None, + install: bool = False, +) -> AsyncSession: + """Continue a detached/stopped session from its State.""" + resolved = _coerce_state(state) + if not resolved.native_session_id: + raise StateIncompatible("State has no native session id to resume") + config = build_config( + resolved.harness, + sandbox=sandbox, + model=model if model is not None else resolved.model, + api_key=api_key, + api_base=api_base, + instructions=instructions, + tools=tools, + skills=skills, + disable_tools=disable_tools, + permissions=permissions, + on_approval=on_approval, + output=output, + max_turns=max_turns, + timeout=timeout, + metadata=metadata, + options=options, + install=install, + ) + return AsyncSession(config, resume_from=resolved.native_session_id) + + +def agent_capabilities(harness: Harness) -> Capabilities: + """What a harness supports (permission modes, structured output, tools...).""" + return get_harness_config(require_harness(harness)).capabilities + + +def aagent( + harness: Harness, + prompt: str, + *, + sandbox: Sandbox, + stream: bool = False, + model: str | None = None, + api_key: str | None = None, + api_base: str | None = None, + instructions: str | None = None, + tools: Sequence[Callable[..., Any]] = (), + skills: Sequence[str | os.PathLike[str]] = (), + disable_tools: Sequence[str] = (), + permissions: PermissionMode = "full", + on_approval: ApprovalHandler | None = None, + output: type[BaseModel] | None = None, + max_turns: int | None = None, + timeout: float | None = None, + metadata: Mapping[str, Any] | None = None, + options: HarnessOptions | None = None, + install: bool = False, +) -> Coroutine[Any, Any, Result] | AsyncEventStream: + """Run an agent harness on one prompt. + + `await litellm.aagent(...)` returns a Result. With stream=True it returns an async + iterator of events instead: `async for event in litellm.aagent(..., stream=True)`. + """ + kwargs: dict[str, Any] = { # mutable-ok: forwarded as **kwargs to arun_agent/astream_agent + "sandbox": sandbox, + "model": model, + "api_key": api_key, + "api_base": api_base, + "instructions": instructions, + "tools": tools, + "skills": skills, + "disable_tools": disable_tools, + "permissions": permissions, + "on_approval": on_approval, + "output": output, + "max_turns": max_turns, + "timeout": timeout, + "metadata": metadata, + "options": options, + "install": install, + } + if stream: + return astream_agent(harness, prompt, **kwargs) + return arun_agent(harness, prompt, **kwargs) diff --git a/litellm/harness/sandbox/__init__.py b/litellm/harness/sandbox/__init__.py new file mode 100644 index 00000000000..434917d625a --- /dev/null +++ b/litellm/harness/sandbox/__init__.py @@ -0,0 +1,25 @@ +"""Sandboxes for litellm.harness: where the runtime runs and which files it can touch.""" + +from litellm.harness.sandbox.base import CompletedRun, Process, Sandbox +from litellm.harness.sandbox.docker import DockerSandbox, docker +from litellm.harness.sandbox.local import LocalSandbox, local +from litellm.harness.sandbox.snapshot import ( + build_file_changes, + capture_text_contents, + diff_snapshots, + snapshot_local, +) + +__all__ = ( + "CompletedRun", + "DockerSandbox", + "LocalSandbox", + "Process", + "Sandbox", + "build_file_changes", + "capture_text_contents", + "diff_snapshots", + "docker", + "local", + "snapshot_local", +) diff --git a/litellm/harness/sandbox/base.py b/litellm/harness/sandbox/base.py new file mode 100644 index 00000000000..18bdeac58a1 --- /dev/null +++ b/litellm/harness/sandbox/base.py @@ -0,0 +1,62 @@ +"""The Sandbox protocol: where a harness runtime runs and which files it can touch.""" + +from __future__ import annotations + +import asyncio +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from typing import Protocol, runtime_checkable + + +@dataclass(frozen=True) +class CompletedRun: + stdout: str + stderr: str + exit_code: int + + +@runtime_checkable +class Process(Protocol): + stdin: asyncio.StreamWriter | None + stdout: asyncio.StreamReader + stderr: asyncio.StreamReader + + async def wait(self) -> int: ... + + async def kill(self) -> None: ... + + +@runtime_checkable +class Sandbox(Protocol): + workdir: str + + async def exec( + self, + cmd: Sequence[str], + *, + env: Mapping[str, str] | None = None, + cwd: str | None = None, + ) -> Process: ... + + async def run( + self, + cmd: Sequence[str], + *, + env: Mapping[str, str] | None = None, + cwd: str | None = None, + timeout: float | None = None, + ) -> CompletedRun: ... + + async def read(self, path: str) -> bytes: ... + + async def write(self, path: str, data: bytes) -> None: ... + + def host_url(self, port: int) -> str: ... + + async def which(self, binary: str) -> str | None: ... + + async def snapshot(self) -> Mapping[str, str]: ... + + async def tempdir(self) -> str: ... + + async def close(self) -> None: ... diff --git a/litellm/harness/sandbox/docker.py b/litellm/harness/sandbox/docker.py new file mode 100644 index 00000000000..32dbfc27ddc --- /dev/null +++ b/litellm/harness/sandbox/docker.py @@ -0,0 +1,298 @@ +"""DockerSandbox: run the harness runtime inside a container via the docker CLI.""" + +from __future__ import annotations + +import asyncio +import os +import posixpath +import shutil +from collections.abc import Mapping, Sequence +from types import MappingProxyType +from typing import Final + +from litellm.constants import HARNESS_SNAPSHOT_SKIP_DIRS +from litellm.harness.errors import SandboxError +from litellm.harness.sandbox.base import CompletedRun +from litellm.harness.sandbox.local import SubprocessHandle, collect_output +from litellm.harness.sandbox.snapshot import HARNESS_SNAPSHOT_MAX_FILE_BYTES + +DOCKER_HOST_ALIAS: Final = "host.docker.internal" +_SHA256_HEX_LEN: Final = 64 +_WRITE_SCRIPT: Final = 'mkdir -p "$(dirname "$1")" && cat > "$1"' +_WHICH_SCRIPT: Final = 'command -v "$1"' + + +def _snapshot_script() -> str: + prune = " -o ".join(f"-name '{name}'" for name in sorted(HARNESS_SNAPSHOT_SKIP_DIRS)) + return ( + 'cd "$1" && find . -type d \\( ' + + prune + + " \\) -prune -o -type f -size -" + + f"{HARNESS_SNAPSHOT_MAX_FILE_BYTES + 1}c" + + " -exec sha256sum {} +" + ) + + +def parse_sha256sum(output: str) -> Mapping[str, str]: + """Parse `sha256sum` lines (" ./rel/path") into {rel/path: hex}.""" + return MappingProxyType( + { + line[_SHA256_HEX_LEN + 2 :].removeprefix("./"): line[:_SHA256_HEX_LEN] + for line in output.splitlines() + if len(line) > _SHA256_HEX_LEN + 2 + } + ) + + +class DockerSandbox: + """Sandbox backed by a long-lived `sleep infinity` container.""" + + # Harness configs read this to skip a runtime's own nested OS sandbox. + is_container = True + + def __init__( + self, + image: str, + mounts: Mapping[str | os.PathLike[str], str] | None = None, + workdir: str = "/workspace", + env: Mapping[str, str] | None = None, + name: str | None = None, + ) -> None: + if not image: + raise SandboxError("docker sandbox needs an image") + if not posixpath.isabs(workdir): + raise SandboxError(f"docker workdir must be absolute: {workdir}") + self.image = image + self.workdir: str = posixpath.normpath(workdir) + self.mounts: Mapping[str, str] = MappingProxyType( + {os.path.abspath(os.fspath(host)): container for host, container in (mounts.items() if mounts else ())} + ) + self.env: Mapping[str, str] = MappingProxyType(dict(env or ())) + self.name = name + self.container_id: str | None = None + self._start_lock = asyncio.Lock() + self._processes: set[SubprocessHandle] = set() # mutable-ok: live-process registry (add/discard) + self._closed = False + + def __repr__(self) -> str: + return f"DockerSandbox({self.image!r}, workdir={self.workdir!r})" + + # -- docker CLI plumbing (tests monkeypatch these two) --------------------- + + def _docker_binary(self) -> str: + binary = shutil.which("docker") + if binary is None: + raise SandboxError( + "docker sandbox requires the `docker` CLI on PATH; install Docker or use sandbox.local(path)" + ) + return binary + + async def _spawn(self, args: Sequence[str]) -> SubprocessHandle: + """Start `docker ` with stdin/stdout/stderr pipes.""" + try: + proc = await asyncio.create_subprocess_exec( + self._docker_binary(), + *args, + stdin=asyncio.subprocess.PIPE, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + ) + except (FileNotFoundError, PermissionError) as exc: + raise SandboxError(f"could not run docker: {exc}") from exc + return SubprocessHandle(proc) + + async def _docker( + self, + args: Sequence[str], + *, + input: bytes | None = None, + timeout: float | None = None, + ) -> tuple[int, bytes, bytes]: + """Run `docker ` to completion; returns (exit_code, stdout, stderr).""" + handle = await self._spawn(args) + try: + return await asyncio.wait_for(_communicate(handle, input), timeout) + except asyncio.TimeoutError: + await handle.kill() + raise SandboxError(f"docker {args[0]} timed out after {timeout}s") + + # -- command construction -------------------------------------------------- + + def run_args( + self, + ) -> list[str]: # mutable-ok: argv is returned as a list, the shape callers and tests compare against + name_args = ("--name", self.name) if self.name else () + mount_args = tuple( + arg for host, container in self.mounts.items() for arg in ("-v", f"{host}:{container}") + ) # comprehension-ok: flattens (flag, value) pairs into argv + env_args = tuple( + arg for key, value in self.env.items() for arg in ("-e", f"{key}={value}") + ) # comprehension-ok: flattens (flag, value) pairs into argv + return [ # mutable-ok: argv is returned as a list, the shape callers and tests compare against + "run", + "-d", + "--rm", + f"--add-host={DOCKER_HOST_ALIAS}:host-gateway", + *name_args, + *mount_args, + *env_args, + "-w", + self.workdir, + self.image, + "sleep", + "infinity", + ] + + def exec_args( + self, + container_id: str, + cmd: Sequence[str], + *, + env: Mapping[str, str] | None = None, + cwd: str | None = None, + ) -> list[str]: # mutable-ok: argv is returned as a list, the shape callers and tests compare against + env_args = tuple( + arg for key, value in (env.items() if env else ()) for arg in ("-e", f"{key}={value}") + ) # comprehension-ok: flattens (flag, value) pairs into argv + return [ # mutable-ok: argv is returned as a list, the shape callers and tests compare against + "exec", + "-i", + "-w", + self.container_path(cwd or self.workdir), + *env_args, + container_id, + *cmd, + ] + + def container_path(self, path: str) -> str: + """Absolute container path; relative paths resolve against workdir.""" + joined = path if posixpath.isabs(path) else posixpath.join(self.workdir, path) + return posixpath.normpath(joined) + + # -- lifecycle --------------------------------------------------------------- + + async def start(self) -> str: + """Start the container if needed and return its id.""" + if self._closed: + raise SandboxError("sandbox is closed") + async with self._start_lock: + if self.container_id is not None: + return self.container_id + code, out, err = await self._docker(self.run_args()) + if code != 0: + raise SandboxError(f"docker run {self.image} failed ({code}): {err.decode(errors='replace').strip()}") + container_id = out.decode().strip() + if not container_id: + raise SandboxError("docker run returned no container id") + self.container_id = container_id + return container_id + + async def _exec_capture(self, cmd: Sequence[str], *, input: bytes | None = None) -> tuple[int, bytes, bytes]: + container_id = await self.start() + return await self._docker(self.exec_args(container_id, cmd), input=input) + + # -- Sandbox protocol -------------------------------------------------------- + + async def exec( + self, + cmd: Sequence[str], + *, + env: Mapping[str, str] | None = None, + cwd: str | None = None, + ) -> SubprocessHandle: + if not cmd: + raise SandboxError("exec() needs a non-empty command") + container_id = await self.start() + handle = await self._spawn(self.exec_args(container_id, cmd, env=env, cwd=cwd)) + self._processes.add(handle) + return handle + + async def run( + self, + cmd: Sequence[str], + *, + env: Mapping[str, str] | None = None, + cwd: str | None = None, + timeout: float | None = None, + ) -> CompletedRun: + handle = await self.exec(cmd, env=env, cwd=cwd) + try: + return await collect_output(handle, cmd, timeout) + finally: + self._processes.discard(handle) + + async def read(self, path: str) -> bytes: + target = self.container_path(path) + code, out, err = await self._exec_capture(("cat", target)) + if code != 0: + raise SandboxError(f"could not read {target}: {err.decode(errors='replace').strip()}") + return out + + async def write(self, path: str, data: bytes) -> None: + target = self.container_path(path) + code, _, err = await self._exec_capture(("sh", "-c", _WRITE_SCRIPT, "sh", target), input=data) + if code != 0: + raise SandboxError(f"could not write {target}: {err.decode(errors='replace').strip()}") + + def host_url(self, port: int) -> str: + return f"http://{DOCKER_HOST_ALIAS}:{port}" + + async def which(self, binary: str) -> str | None: + code, out, _ = await self._exec_capture(("sh", "-lc", _WHICH_SCRIPT, "sh", binary)) + found = out.decode(errors="replace").strip() + return found if code == 0 and found else None + + async def tempdir(self) -> str: + """A fresh `mktemp -d` directory inside the container.""" + code, out, err = await self._exec_capture(("mktemp", "-d")) + path = out.decode(errors="replace").strip() + if code != 0 or not path: + raise SandboxError(f"mktemp -d failed: {err.decode(errors='replace').strip()}") + return path + + async def snapshot(self) -> Mapping[str, str]: + code, out, err = await self._exec_capture(("sh", "-c", _snapshot_script(), "sh", self.workdir)) + if code != 0: + raise SandboxError(f"snapshot failed: {err.decode(errors='replace').strip()}") + return parse_sha256sum(out.decode("utf-8", errors="replace")) + + async def close(self) -> None: + if self._closed: + return + self._closed = True + live = tuple(h for h in self._processes if h.returncode is None) + await asyncio.gather(*(h.kill() for h in live), return_exceptions=True) + self._processes.clear() + if self.container_id is not None: + container_id, self.container_id = self.container_id, None + await self._docker( + ["rm", "-f", container_id] # mutable-ok: argv list, the shape _spawn records and tests assert on + ) + + async def __aenter__(self) -> DockerSandbox: + await self.start() + return self + + async def __aexit__(self, *exc_info: object) -> None: + await self.close() + + +async def _communicate(handle: SubprocessHandle, data: bytes | None) -> tuple[int, bytes, bytes]: + if handle.stdin is not None: + if data: + handle.stdin.write(data) + await handle.stdin.drain() + handle.stdin.close() + stdout, stderr = await asyncio.gather(handle.stdout.read(), handle.stderr.read()) + return await handle.wait(), stdout, stderr + + +def docker( + image: str, + mounts: Mapping[str | os.PathLike[str], str] | None = None, + workdir: str = "/workspace", + env: Mapping[str, str] | None = None, + name: str | None = None, +) -> DockerSandbox: + """Sandbox in a new container of `image`, started lazily on first use.""" + return DockerSandbox(image, mounts=mounts, workdir=workdir, env=env, name=name) diff --git a/litellm/harness/sandbox/local.py b/litellm/harness/sandbox/local.py new file mode 100644 index 00000000000..10c7303ef27 --- /dev/null +++ b/litellm/harness/sandbox/local.py @@ -0,0 +1,277 @@ +"""LocalSandbox: run the harness runtime as a subprocess on this machine.""" + +from __future__ import annotations + +import asyncio +import itertools +import os +import shutil +import signal +import tempfile +from collections.abc import Mapping, Sequence +from types import MappingProxyType +from typing import Final + +from litellm.constants import HARNESS_PROCESS_KILL_GRACE_SECONDS +from litellm.harness.errors import SandboxError +from litellm.harness.sandbox.base import CompletedRun +from litellm.harness.sandbox.snapshot import snapshot_local + +_SECRET_PREFIXES: Final = ( + "ANTHROPIC_", + "OPENAI_", + "LITELLM_", + "AZURE_", + "AWS_", + "GEMINI_", + "CODEX_", + "CURSOR_", + "VERTEX", + # A parent Claude Code session's socket/session vars make a child `claude` attach to + # the parent's login instead of the harness token. + "CLAUDE_CODE_", + "CLAUDE_PID", + "CLAUDECODE", +) +_SECRET_NAMES: Final = frozenset({"GOOGLE_API_KEY", "GOOGLE_APPLICATION_CREDENTIALS"}) +_SECRET_SUBSTRINGS: Final = ("API_KEY", "TOKEN", "SECRET") +_TEMPDIR_PREFIX: Final = "litellm-harness-" + + +def is_secret_env_name(name: str) -> bool: + """True if an env var name looks like a provider credential.""" + upper = name.upper() + if upper in _SECRET_NAMES or upper.startswith(_SECRET_PREFIXES): + return True + return any(part in upper for part in _SECRET_SUBSTRINGS) + + +def filtered_environ( + base: Mapping[str, str] | None = None, + extra: Mapping[str, str] | None = None, +) -> Mapping[str, str]: + """base (default os.environ) without provider secrets, then extra on top.""" + source = os.environ if base is None else base + kept = ((k, v) for k, v in source.items() if not is_secret_env_name(k)) + overlay = extra.items() if extra else () + return MappingProxyType(dict(itertools.chain(kept, overlay))) + + +def _signal_process(proc: asyncio.subprocess.Process, sig: int) -> None: + try: + os.killpg(proc.pid, sig) + except (ProcessLookupError, PermissionError, OSError): + try: + proc.send_signal(sig) + except ProcessLookupError: + pass + + +class SubprocessHandle: + """Process-protocol wrapper around an asyncio subprocess.""" + + def __init__(self, proc: asyncio.subprocess.Process) -> None: + if proc.stdout is None or proc.stderr is None: + raise SandboxError("subprocess was started without stdout/stderr pipes") + self._proc = proc + self.stdin: asyncio.StreamWriter | None = proc.stdin + self.stdout: asyncio.StreamReader = proc.stdout + self.stderr: asyncio.StreamReader = proc.stderr + + @property + def pid(self) -> int: + return self._proc.pid + + @property + def returncode(self) -> int | None: + return self._proc.returncode + + async def wait(self) -> int: + return await self._proc.wait() + + async def kill(self) -> None: + """SIGTERM, wait HARNESS_PROCESS_KILL_GRACE_SECONDS, then SIGKILL.""" + if self._proc.returncode is not None: + return + _signal_process(self._proc, signal.SIGTERM) + try: + await asyncio.wait_for(self._proc.wait(), timeout=HARNESS_PROCESS_KILL_GRACE_SECONDS) + return + except asyncio.TimeoutError: + pass + _signal_process(self._proc, signal.SIGKILL) + await self._proc.wait() + + +async def _read_all(handle: SubprocessHandle) -> tuple[bytes, bytes, int]: + if handle.stdin is not None: + handle.stdin.close() + stdout, stderr = await asyncio.gather(handle.stdout.read(), handle.stderr.read()) + exit_code = await handle.wait() + return stdout, stderr, exit_code + + +async def collect_output(handle: SubprocessHandle, cmd: Sequence[str], timeout: float | None) -> CompletedRun: + """Close stdin, read stdout/stderr to EOF; kill and raise SandboxError on timeout.""" + try: + stdout, stderr, code = await asyncio.wait_for(_read_all(handle), timeout) + except asyncio.TimeoutError: + await handle.kill() + raise SandboxError(f"command timed out after {timeout}s: {cmd[0]}") + return CompletedRun( + stdout=stdout.decode("utf-8", errors="replace"), + stderr=stderr.decode("utf-8", errors="replace"), + exit_code=code, + ) + + +def _is_within(path: str, root: str) -> bool: + return path == root or path.startswith(root.rstrip(os.sep) + os.sep) + + +class LocalSandbox: + """Sandbox backed by the local filesystem and asyncio subprocesses.""" + + def __init__(self, path: str | os.PathLike[str]) -> None: + resolved = os.path.realpath(os.path.abspath(os.fspath(path))) + if not os.path.isdir(resolved): + raise SandboxError(f"sandbox path does not exist or is not a directory: {resolved}") + self.workdir: str = resolved + self._processes: set[SubprocessHandle] = set() # mutable-ok: live-process registry (add/discard) + self._tempdirs: list[str] = [] # mutable-ok: tempdirs created on demand by tempdir(), removed on close() + self._closed = False + + def __repr__(self) -> str: + return f"LocalSandbox({self.workdir!r})" + + def _check_open(self) -> None: + if self._closed: + raise SandboxError("sandbox is closed") + + def _allowed_roots(self) -> tuple[str, ...]: + return (self.workdir, *self._tempdirs) + + def resolve_path(self, path: str) -> str: + """Absolute real path for path; SandboxError if it escapes the sandbox.""" + joined = path if os.path.isabs(path) else os.path.join(self.workdir, path) + real = os.path.realpath(joined) + if not any(_is_within(real, root) for root in self._allowed_roots()): + raise SandboxError(f"path escapes the sandbox: {path}") + return real + + def _resolve_cwd(self, cwd: str | None) -> str: + if cwd is None: + return self.workdir + resolved = self.resolve_path(cwd) + if not os.path.isdir(resolved): + raise SandboxError(f"cwd is not a directory: {cwd}") + return resolved + + def child_env(self, env: Mapping[str, str] | None = None) -> Mapping[str, str]: + return filtered_environ(extra=env) + + async def exec( + self, + cmd: Sequence[str], + *, + env: Mapping[str, str] | None = None, + cwd: str | None = None, + ) -> SubprocessHandle: + self._check_open() + if not cmd: + raise SandboxError("exec() needs a non-empty command") + try: + proc = await asyncio.create_subprocess_exec( + *cmd, + stdin=asyncio.subprocess.PIPE, + stdout=asyncio.subprocess.PIPE, + stderr=asyncio.subprocess.PIPE, + cwd=self._resolve_cwd(cwd), + env=self.child_env(env), + start_new_session=True, + ) + except (FileNotFoundError, PermissionError) as exc: + raise SandboxError(f"could not start {cmd[0]}: {exc}") from exc + handle = SubprocessHandle(proc) + self._processes.add(handle) + return handle + + async def run( + self, + cmd: Sequence[str], + *, + env: Mapping[str, str] | None = None, + cwd: str | None = None, + timeout: float | None = None, + ) -> CompletedRun: + handle = await self.exec(cmd, env=env, cwd=cwd) + try: + return await collect_output(handle, cmd, timeout) + finally: + self._processes.discard(handle) + + async def read(self, path: str) -> bytes: + self._check_open() + resolved = self.resolve_path(path) + try: + return await asyncio.to_thread(_read_bytes, resolved) + except OSError as exc: + raise SandboxError(f"could not read {path}: {exc}") from exc + + async def write(self, path: str, data: bytes) -> None: + self._check_open() + resolved = self.resolve_path(path) + try: + await asyncio.to_thread(_write_bytes, resolved, data) + except OSError as exc: + raise SandboxError(f"could not write {path}: {exc}") from exc + + def host_url(self, port: int) -> str: + return f"http://127.0.0.1:{port}" + + async def which(self, binary: str) -> str | None: + return shutil.which(binary, path=self.child_env().get("PATH")) + + async def tempdir(self) -> str: + """A private temp dir (e.g. for CODEX_HOME), removed on close().""" + self._check_open() + path = os.path.realpath(tempfile.mkdtemp(prefix=_TEMPDIR_PREFIX)) + self._tempdirs.append(path) + return path + + async def snapshot(self) -> Mapping[str, str]: + self._check_open() + return await snapshot_local(self.workdir) + + async def close(self) -> None: + if self._closed: + return + self._closed = True + live = tuple(h for h in self._processes if h.returncode is None) + await asyncio.gather(*(h.kill() for h in live), return_exceptions=True) + self._processes.clear() + for path in self._tempdirs: + shutil.rmtree(path, ignore_errors=True) + self._tempdirs.clear() + + async def __aenter__(self) -> LocalSandbox: + return self + + async def __aexit__(self, *exc_info: object) -> None: + await self.close() + + +def _read_bytes(path: str) -> bytes: + with open(path, "rb") as fh: + return fh.read() + + +def _write_bytes(path: str, data: bytes) -> None: + os.makedirs(os.path.dirname(path), exist_ok=True) + with open(path, "wb") as fh: + fh.write(data) + + +def local(path: str | os.PathLike[str]) -> LocalSandbox: + """Sandbox rooted at an existing local directory.""" + return LocalSandbox(path) diff --git a/litellm/harness/sandbox/snapshot.py b/litellm/harness/sandbox/snapshot.py new file mode 100644 index 00000000000..61407c6c063 --- /dev/null +++ b/litellm/harness/sandbox/snapshot.py @@ -0,0 +1,183 @@ +"""Workspace snapshots and FileChange construction. + +A snapshot maps a workspace-relative POSIX path to the sha256 of its contents. +Diffing two snapshots tells us which files a turn created, modified or deleted; +`build_file_changes` turns that into `FileChange` events with unified diffs for +small text files. +""" + +from __future__ import annotations + +import asyncio +import difflib +import functools +import hashlib +import os +from collections.abc import Iterator, Mapping +from types import MappingProxyType +from typing import TYPE_CHECKING, Final + +from litellm.constants import HARNESS_MAX_DIFF_BYTES, HARNESS_SNAPSHOT_SKIP_DIRS +from litellm.harness.errors import HarnessError +from litellm.harness.types import FileChange, FileChangeKind + +if TYPE_CHECKING: + from litellm.harness.sandbox.base import Sandbox + +# Files larger than this are left out of snapshots entirely. +HARNESS_SNAPSHOT_MAX_FILE_BYTES: Final = 50 * 1024 * 1024 +# Upper bound on bytes read by capture_text_contents() for one turn. +HARNESS_SNAPSHOT_MAX_TOTAL_BYTES: Final = 16 * 1024 * 1024 +_HASH_CHUNK_BYTES: Final = 1024 * 1024 +_NO_NEWLINE_MARKER: Final = "\\ No newline at end of file\n" + + +def _hash_file(path: str) -> str: + digest = hashlib.sha256() + with open(path, "rb") as fh: + for chunk in iter(functools.partial(fh.read, _HASH_CHUNK_BYTES), b""): + digest.update(chunk) + return digest.hexdigest() + + +def _hash_entry(root: str, dirpath: str, filename: str) -> tuple[str, str] | None: + full = os.path.join(dirpath, filename) + try: + info = os.lstat(full) + except OSError: + return None + if not os.path.isfile(full) or os.path.islink(full): + return None + if info.st_size > HARNESS_SNAPSHOT_MAX_FILE_BYTES: + return None + try: + digest = _hash_file(full) + except OSError: + return None + rel = os.path.relpath(full, root).replace(os.sep, "/") + return rel, digest + + +def _walk_entries(root: str) -> Iterator[tuple[str, str]]: + for dirpath, dirnames, filenames in os.walk(root, followlinks=False): + dirnames[:] = [ # mutable-ok: os.walk prunes only via in-place mutation of its dirnames list + d for d in dirnames if d not in HARNESS_SNAPSHOT_SKIP_DIRS + ] + for filename in filenames: + entry = _hash_entry(root, dirpath, filename) + if entry is not None: + yield entry + + +def snapshot_local_sync(root: str) -> Mapping[str, str]: + """Hash every regular file under root. Symlinks are never followed.""" + return MappingProxyType(dict(_walk_entries(root))) + + +async def snapshot_local(root: str) -> Mapping[str, str]: + """Async wrapper around snapshot_local_sync (runs in a worker thread).""" + return await asyncio.to_thread(snapshot_local_sync, root) + + +def _change_kind(path: str, before: Mapping[str, str], after: Mapping[str, str]) -> FileChangeKind | None: + if path not in before: + return "created" + if path not in after: + return "deleted" + if before[path] != after[path]: + return "modified" + return None + + +def diff_snapshots( + before: Mapping[str, str], after: Mapping[str, str] +) -> list[tuple[str, FileChangeKind]]: # mutable-ok: public sandbox helper; callers compare against a list + """Return (path, kind) for every changed file, sorted by path.""" + kinds = ((path, _change_kind(path, before, after)) for path in sorted(frozenset(before) | frozenset(after))) + return [ # mutable-ok: public sandbox helper returns a list + (path, kind) for path, kind in kinds if kind is not None + ] + + +def _as_text(data: bytes) -> str | None: + if len(data) > HARNESS_MAX_DIFF_BYTES or b"\0" in data: + return None + try: + return data.decode("utf-8") + except UnicodeDecodeError: + return None + + +def unified_diff(path: str, old: str | None, new: str | None) -> str: + """Unified diff between two versions of path; None means the file is absent.""" + from_file = "/dev/null" if old is None else f"a/{path}" + to_file = "/dev/null" if new is None else f"b/{path}" + lines = difflib.unified_diff( + (old or "").splitlines(keepends=True), + (new or "").splitlines(keepends=True), + fromfile=from_file, + tofile=to_file, + ) + return "".join(line if line.endswith("\n") else line + "\n" + _NO_NEWLINE_MARKER for line in lines) + + +async def _read_or_none(sandbox: Sandbox, path: str) -> bytes | None: + try: + return await sandbox.read(path) + except (HarnessError, OSError): + return None + + +async def capture_text_contents(sandbox: Sandbox, paths_hashes: Mapping[str, str]) -> Mapping[str, bytes]: + """Read small text files before a turn so "modified"/"deleted" diffs can be built. + + Each kept file is <= HARNESS_MAX_DIFF_BYTES; every byte read (kept or not) counts + toward HARNESS_SNAPSHOT_MAX_TOTAL_BYTES, after which capture stops. + """ + captured: dict[str, bytes] = {} # mutable-ok: async accumulator (awaits per read), frozen on return + total = 0 + for path in sorted(paths_hashes): + if total >= HARNESS_SNAPSHOT_MAX_TOTAL_BYTES: + break + data = await _read_or_none(sandbox, path) + if data is None: + continue + total += len(data) + if _as_text(data) is not None: + captured[path] = data + return MappingProxyType(captured) + + +async def _change_for( + sandbox: Sandbox, + path: str, + kind: FileChangeKind, + before_contents: Mapping[str, bytes], +) -> FileChange: + old_bytes = before_contents.get(path) + old = _as_text(old_bytes) if old_bytes is not None else None + if kind == "deleted": + diff = unified_diff(path, old, None) if old is not None else None + return FileChange(path=path, kind=kind, diff=diff) + new_bytes = await _read_or_none(sandbox, path) + new = _as_text(new_bytes) if new_bytes is not None else None + if new is None or (kind == "modified" and old is None): + return FileChange(path=path, kind=kind, diff=None) + return FileChange( + path=path, + kind=kind, + diff=unified_diff(path, old if kind == "modified" else None, new), + ) + + +async def build_file_changes( + sandbox: Sandbox, + before: Mapping[str, str], + after: Mapping[str, str], + before_contents: Mapping[str, bytes] | None = None, +) -> list[FileChange]: # mutable-ok: feeds the public Result.files list + """FileChange per changed path. diff is None when it cannot be built as text.""" + contents: Mapping[str, bytes] = before_contents or MappingProxyType({}) + return [ # mutable-ok: feeds the public Result.files list + await _change_for(sandbox, path, kind, contents) for path, kind in diff_snapshots(before, after) + ] diff --git a/litellm/harness/sync.py b/litellm/harness/sync.py new file mode 100644 index 00000000000..4583a98bf54 --- /dev/null +++ b/litellm/harness/sync.py @@ -0,0 +1,443 @@ +"""Sync API for litellm.harness: one daemon event-loop thread runs every async call.""" + +from __future__ import annotations + +import asyncio +import os +import threading +from collections.abc import AsyncIterator, Callable, Coroutine, Mapping, Sequence +from concurrent.futures import Future +from typing import ( + Any, + TypeVar, +) + +from pydantic import BaseModel + +from litellm.harness.context import ApprovalHandler +from litellm.harness.options import HarnessOptions +from litellm.harness.runtime import ( + AsyncEventStream, + AsyncSession, + aagent_resume, + aagent_session, + arun_agent, + astream_agent, +) +from litellm.harness.sandbox.base import Sandbox +from litellm.harness.types import ( + Done, + Event, + Harness, + PermissionMode, + Result, + State, + Usage, +) + +T = TypeVar("T") + +IN_LOOP_MESSAGE = ( + "litellm.{name}() cannot be called from a running event loop; use `await litellm.a{name}(...)` instead" +) + + +class _LoopThread: + """A single background event loop shared by every sync call in the process.""" + + def __init__(self) -> None: + self._lock = threading.Lock() + self._loop: asyncio.AbstractEventLoop | None = None + self._thread: threading.Thread | None = None + + def loop(self) -> asyncio.AbstractEventLoop: + with self._lock: + if self._loop is None or self._thread is None or not self._thread.is_alive(): + self._loop = asyncio.new_event_loop() + self._thread = threading.Thread( + target=self._loop.run_forever, + name="litellm-harness-loop", + daemon=True, + ) + self._thread.start() + return self._loop + + def in_loop_thread(self) -> bool: + return self._thread is not None and threading.current_thread() is self._thread + + def submit(self, coro: Coroutine[Any, Any, T]) -> Future[T]: + return asyncio.run_coroutine_threadsafe(coro, self.loop()) + + +_LOOP = _LoopThread() + + +def _ensure_sync_context(name: str) -> None: + try: + asyncio.get_running_loop() + except RuntimeError: + return + raise RuntimeError(IN_LOOP_MESSAGE.format(name=name)) + + +def run_sync(coro: Coroutine[Any, Any, T], name: str) -> T: + """Run coro on the harness loop thread and block for its result.""" + try: + _ensure_sync_context(name) + except RuntimeError: + coro.close() + raise + future = _LOOP.submit(coro) + try: + return future.result() + except KeyboardInterrupt: + future.cancel() + raise + + +async def _anext(iterator: AsyncIterator[Event]) -> Event | None: + try: + return await iterator.__anext__() + except StopAsyncIteration: + return None + + +async def _aclose_stream(stream: AsyncEventStream) -> None: + await stream.aclose() + + +class EventStream: + """Sync iterator of events for one turn. `.result` is set once Done is seen.""" + + def __init__(self, stream: AsyncEventStream, name: str = "stream") -> None: + self._stream = stream + self._name = name + self._result: Result | None = None + self._finished = False + + def __iter__(self) -> EventStream: + return self + + def __next__(self) -> Event: + if self._finished: + raise StopIteration + event = run_sync(_anext(self._stream), self._name) + if event is None: + self._finished = True + raise StopIteration + if isinstance(event, Done): + self._result = event.result + return event + + def __enter__(self) -> EventStream: + return self + + def __exit__(self, *exc_info: object) -> None: + self.close() + + @property + def result(self) -> Result | None: + return self._result + + def cancel(self) -> None: + """Stop the turn. Iteration still ends with Done(stop_reason='cancelled').""" + _LOOP.loop().call_soon_threadsafe(self._stream.cancel) + + def close(self) -> None: + """Abandon the stream and release the session behind it.""" + if self._finished: + return + self._finished = True + run_sync(_aclose_stream(self._stream), self._name) + + +class Session: + """Sync multi-turn session. Use as a context manager.""" + + def __init__(self, inner: AsyncSession) -> None: + self._inner = inner + + @property + def aio(self) -> AsyncSession: + """The underlying AsyncSession (runs on the harness loop thread).""" + return self._inner + + def start(self) -> Session: + run_sync(self._inner.start(), "session") + return self + + def __enter__(self) -> Session: + return self.start() + + def __exit__(self, *exc_info: object) -> None: + self.close() + + def run(self, prompt: str) -> Result: + return run_sync(self._inner.arun(prompt), "run") + + def stream(self, prompt: str) -> EventStream: + return EventStream(self._inner.astream(prompt)) + + def close(self) -> None: + run_sync(self._inner.aclose(), "close") + + def detach(self) -> State: + return run_sync(self._inner.adetach(), "detach") + + def stop(self) -> State: + return run_sync(self._inner.astop(), "stop") + + def history( + self, + ) -> list[dict[str, Any]]: # mutable-ok: public API returns OpenAI-format message dicts from the handler + return run_sync(self._inner.history(), "history") + + @property + def cost(self) -> float: + return self._inner.cost + + @property + def usage(self) -> Usage: + return self._inner.usage + + @property + def results(self) -> list[Result]: # mutable-ok: public property; returns a detached copy of the session's results + return list(self._inner.results) # mutable-ok: detached copy so callers cannot mutate the session's accumulator + + @property + def session_id(self) -> str: + return self._inner.session_id + + +def _run( + harness: Harness, + prompt: str, + *, + sandbox: Sandbox, + model: str | None = None, + api_key: str | None = None, + api_base: str | None = None, + instructions: str | None = None, + tools: Sequence[Callable[..., Any]] = (), + skills: Sequence[str | os.PathLike[str]] = (), + disable_tools: Sequence[str] = (), + permissions: PermissionMode = "full", + on_approval: ApprovalHandler | None = None, + output: type[BaseModel] | None = None, + max_turns: int | None = None, + timeout: float | None = None, + metadata: Mapping[str, Any] | None = None, + options: HarnessOptions | None = None, + install: bool = False, +) -> Result: + """Run one prompt to completion (blocking) and return the Result.""" + return run_sync( + arun_agent( + harness, + prompt, + sandbox=sandbox, + model=model, + api_key=api_key, + api_base=api_base, + instructions=instructions, + tools=tools, + skills=skills, + disable_tools=disable_tools, + permissions=permissions, + on_approval=on_approval, + output=output, + max_turns=max_turns, + timeout=timeout, + metadata=metadata, + options=options, + install=install, + ), + "agent", + ) + + +def _stream( + harness: Harness, + prompt: str, + *, + sandbox: Sandbox, + model: str | None = None, + api_key: str | None = None, + api_base: str | None = None, + instructions: str | None = None, + tools: Sequence[Callable[..., Any]] = (), + skills: Sequence[str | os.PathLike[str]] = (), + disable_tools: Sequence[str] = (), + permissions: PermissionMode = "full", + on_approval: ApprovalHandler | None = None, + output: type[BaseModel] | None = None, + max_turns: int | None = None, + timeout: float | None = None, + metadata: Mapping[str, Any] | None = None, + options: HarnessOptions | None = None, + install: bool = False, +) -> EventStream: + """Stream events for one prompt (sync iterator). Validation errors raise here.""" + _ensure_sync_context("agent") + inner = astream_agent( + harness, + prompt, + sandbox=sandbox, + model=model, + api_key=api_key, + api_base=api_base, + instructions=instructions, + tools=tools, + skills=skills, + disable_tools=disable_tools, + permissions=permissions, + on_approval=on_approval, + output=output, + max_turns=max_turns, + timeout=timeout, + metadata=metadata, + options=options, + install=install, + ) + return EventStream(inner) + + +def agent_session( + harness: Harness, + *, + sandbox: Sandbox, + model: str | None = None, + api_key: str | None = None, + api_base: str | None = None, + instructions: str | None = None, + tools: Sequence[Callable[..., Any]] = (), + skills: Sequence[str | os.PathLike[str]] = (), + disable_tools: Sequence[str] = (), + permissions: PermissionMode = "full", + on_approval: ApprovalHandler | None = None, + output: type[BaseModel] | None = None, + max_turns: int | None = None, + timeout: float | None = None, + metadata: Mapping[str, Any] | None = None, + options: HarnessOptions | None = None, + install: bool = False, +) -> Session: + """A multi-turn agent session: `with litellm.agent_session(...) as s: s.run(...)`.""" + _ensure_sync_context("agent_session") + return Session( + aagent_session( + harness, + sandbox=sandbox, + model=model, + api_key=api_key, + api_base=api_base, + instructions=instructions, + tools=tools, + skills=skills, + disable_tools=disable_tools, + permissions=permissions, + on_approval=on_approval, + output=output, + max_turns=max_turns, + timeout=timeout, + metadata=metadata, + options=options, + install=install, + ) + ) + + +def agent_resume( + state: State | bytes, + *, + sandbox: Sandbox, + model: str | None = None, + api_key: str | None = None, + api_base: str | None = None, + instructions: str | None = None, + tools: Sequence[Callable[..., Any]] = (), + skills: Sequence[str | os.PathLike[str]] = (), + disable_tools: Sequence[str] = (), + permissions: PermissionMode = "full", + on_approval: ApprovalHandler | None = None, + output: type[BaseModel] | None = None, + max_turns: int | None = None, + timeout: float | None = None, + metadata: Mapping[str, Any] | None = None, + options: HarnessOptions | None = None, + install: bool = False, +) -> Session: + """Continue a detached or stopped agent session from its State.""" + _ensure_sync_context("agent_resume") + return Session( + aagent_resume( + state, + sandbox=sandbox, + model=model, + api_key=api_key, + api_base=api_base, + instructions=instructions, + tools=tools, + skills=skills, + disable_tools=disable_tools, + permissions=permissions, + on_approval=on_approval, + output=output, + max_turns=max_turns, + timeout=timeout, + metadata=metadata, + options=options, + install=install, + ) + ) + + +def agent( + harness: Harness, + prompt: str, + *, + sandbox: Sandbox, + stream: bool = False, + model: str | None = None, + api_key: str | None = None, + api_base: str | None = None, + instructions: str | None = None, + tools: Sequence[Callable[..., Any]] = (), + skills: Sequence[str | os.PathLike[str]] = (), + disable_tools: Sequence[str] = (), + permissions: PermissionMode = "full", + on_approval: ApprovalHandler | None = None, + output: type[BaseModel] | None = None, + max_turns: int | None = None, + timeout: float | None = None, + metadata: Mapping[str, Any] | None = None, + options: HarnessOptions | None = None, + install: bool = False, +) -> Result | EventStream: + """Run an agent harness (Claude Code, Codex, OpenCode, Deep Agents) on one prompt. + + Returns a Result. With stream=True it returns an iterator of events instead. + Prefix the model with `litellm_proxy/` to route every model call through your + LiteLLM AI Gateway. + """ + kwargs: dict[str, Any] = { # mutable-ok: forwarded as **kwargs to _run/_stream + "sandbox": sandbox, + "model": model, + "api_key": api_key, + "api_base": api_base, + "instructions": instructions, + "tools": tools, + "skills": skills, + "disable_tools": disable_tools, + "permissions": permissions, + "on_approval": on_approval, + "output": output, + "max_turns": max_turns, + "timeout": timeout, + "metadata": metadata, + "options": options, + "install": install, + } + if stream: + return _stream(harness, prompt, **kwargs) + return _run(harness, prompt, **kwargs) diff --git a/litellm/harness/types.py b/litellm/harness/types.py new file mode 100644 index 00000000000..88ad7f711ea --- /dev/null +++ b/litellm/harness/types.py @@ -0,0 +1,214 @@ +"""Public types for litellm.harness: the Harness enum, events, results and state.""" + +from __future__ import annotations + +import asyncio +import json +from collections.abc import Mapping +from dataclasses import dataclass, field +from enum import Enum +from typing import Any, Literal + +from pydantic import BaseModel + +from litellm.harness.errors import StateIncompatible + +StopReason = Literal["done", "max_turns", "timeout", "cancelled", "runtime_error"] +PermissionMode = Literal["read-only", "ask", "edit", "full"] +FileChangeKind = Literal["created", "modified", "deleted"] + +STATE_VERSION = 1 + + +class Harness(Enum): + """Supported agent runtimes. A plain Enum on purpose: strings are rejected.""" + + CLAUDE_CODE = "claude_code" + CODEX = "codex" + OPENCODE = "opencode" + DEEPAGENTS = "deepagents" + + +def require_harness(harness: object) -> Harness: + """Return harness if it is a Harness member, else raise TypeError with a hint.""" + if isinstance(harness, Harness): + return harness + hint = "" + if isinstance(harness, str): + normalized = harness.strip().lower().replace("-", "_") + for member in Harness: + if normalized in (member.value, member.name.lower()): + hint = f" Did you mean Harness.{member.name}?" + raise TypeError( + f"harness must be a litellm.harness.Harness member, got {type(harness).__name__} {harness!r}.{hint}" + ) + + +@dataclass(frozen=True) +class Usage: + input_tokens: int = 0 + output_tokens: int = 0 + calls: int = 0 + + @property + def total_tokens(self) -> int: + return self.input_tokens + self.output_tokens + + +@dataclass(frozen=True) +class Text: + delta: str + + +@dataclass(frozen=True) +class Reasoning: + delta: str + + +@dataclass(frozen=True) +class ToolCall: + id: str + name: str + native_name: str + input: Mapping[str, Any] + builtin: bool = True + + +@dataclass(frozen=True) +class ToolResult: + id: str + output: str + is_error: bool = False + + +@dataclass(frozen=True) +class FileChange: + path: str + kind: FileChangeKind + diff: str | None = None + + +@dataclass(frozen=True) +class Compaction: + tokens_before: int | None = None + tokens_after: int | None = None + + +@dataclass(frozen=True) +class Approval: + """A request to run a tool. The turn waits until allow() or deny() is called.""" + + tool: str + input: Mapping[str, Any] + _decision: asyncio.Future[tuple[bool, str]] = field( + default_factory=lambda: asyncio.get_event_loop().create_future(), + compare=False, + repr=False, + ) + + def allow(self) -> None: + self._resolve(True, "") + + def deny(self, reason: str = "") -> None: + self._resolve(False, reason) + + @property + def answered(self) -> bool: + return self._decision.done() + + async def wait(self) -> tuple[bool, str]: + return await self._decision + + def _resolve(self, allowed: bool, reason: str) -> None: + if self._decision.done(): + return + loop = self._decision.get_loop() + loop.call_soon_threadsafe(self._set_result, allowed, reason) + + def _set_result(self, allowed: bool, reason: str) -> None: + if not self._decision.done(): + self._decision.set_result((allowed, reason)) + + +@dataclass(frozen=True) +class Result: + text: str + output: BaseModel | None + files: list[FileChange] # mutable-ok: public Result field; users index/iterate it as a list + events: list[Event] # mutable-ok: public Result field; users index/iterate it as a list + usage: Usage + cost: float + stop_reason: StopReason + session_id: str + + +@dataclass(frozen=True) +class Done: + result: Result + + @property + def usage(self) -> Usage: + return self.result.usage + + @property + def cost(self) -> float: + return self.result.cost + + @property + def stop_reason(self) -> StopReason: + return self.result.stop_reason + + +Event = Text | Reasoning | ToolCall | ToolResult | FileChange | Compaction | Approval | Done + + +@dataclass(frozen=True) +class Capabilities: + structured_output: bool + tool_approval: bool + tool_filtering: bool + history: bool + custom_tools: bool + skills: bool + resume: bool + permission_modes: frozenset[str] + + +@dataclass(frozen=True) +class State: + """Resume state for a detached or stopped session. Contains no credentials.""" + + harness: Harness + native_session_id: str | None + workdir: str + model: str | None = None + version: int = STATE_VERSION + + def dumps(self) -> bytes: + return json.dumps( + { # mutable-ok: JSON payload serialized immediately by json.dumps + "harness": self.harness.value, + "native_session_id": self.native_session_id, + "workdir": self.workdir, + "model": self.model, + "version": self.version, + } + ).encode("utf-8") + + @classmethod + def loads(cls, data: bytes) -> State: + try: + raw = json.loads(data.decode("utf-8")) + harness = Harness(raw["harness"]) + version = int(raw["version"]) + except (ValueError, KeyError, TypeError, UnicodeDecodeError) as e: + raise StateIncompatible(f"Unreadable harness state: {e}") from e + if version != STATE_VERSION: + raise StateIncompatible(f"State version {version} is not supported (expected {STATE_VERSION})") + return cls( + harness=harness, + native_session_id=raw.get("native_session_id"), + workdir=raw["workdir"], + model=raw.get("model"), + version=version, + ) diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index ae4ac29de9e..0c48b7c1c2a 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -3826,7 +3826,7 @@ def _build_bedrock_tool_result_content_blocks( if tool_result_content_blocks: return tool_result_content_blocks, True - message_content: Final = message["content"] + message_content: Final = message.get("content") if isinstance(message_content, str): return [BedrockToolResultContentBlock(text=message_content)], False if isinstance(message_content, list): @@ -4095,23 +4095,20 @@ def get_user_message_block_or_continue_message( ) -> ChatCompletionUserMessage: """ Returns the user content block - if content block is an empty string, then return the default continue message + if content block is missing or an empty string, then return the default continue message Relevant Issue: https://github.com/BerriAI/litellm/issues/7169 """ content_block: Final = message.get("content", None) - # Handle None case - if content_block is None or (user_continue_message is None and litellm.modify_params is False): + if user_continue_message is None and litellm.modify_params is False: return skip_empty_text_blocks(message=message) - # Handle string case + if content_block is None or (isinstance(content_block, str) and not content_block.strip()): + return ChatCompletionUserMessage(**(user_continue_message or DEFAULT_USER_CONTINUE_MESSAGE)) + if isinstance(content_block, str): - # check if content is empty - if content_block.strip(): - return message - else: - return ChatCompletionUserMessage(**(user_continue_message or DEFAULT_USER_CONTINUE_MESSAGE)) + return message # Handle list case if isinstance(content_block, list): @@ -4374,9 +4371,10 @@ class BedrockConverseMessagesProcessor: message=messages[msg_i], user_continue_message=user_continue_message, ) - if isinstance(message_block["content"], list): + message_content = message_block.get("content") + if isinstance(message_content, list): _parts: list[BedrockContentBlock] = [] - for element in message_block["content"]: + for element in message_content: if isinstance(element, dict): if element["type"] == "text": _part = BedrockContentBlock(text=element["text"]) @@ -4418,8 +4416,8 @@ class BedrockConverseMessagesProcessor: if _cache_point_block is not None: _parts.append(_cache_point_block) user_content.extend(_parts) - elif message_block["content"] and isinstance(message_block["content"], str): - _part = BedrockContentBlock(text=messages[msg_i]["content"]) + elif message_content and isinstance(message_content, str): + _part = BedrockContentBlock(text=message_content) _cache_point_block = litellm.AmazonConverseConfig().get_cache_point_block( message_block, block_type="content_block", model=model ) @@ -4746,9 +4744,10 @@ def _bedrock_converse_messages_pt( message=messages[msg_i], user_continue_message=user_continue_message, ) - if isinstance(message_block["content"], list): + message_content = message_block.get("content") + if isinstance(message_content, list): _parts: list[BedrockContentBlock] = [] - for element in message_block["content"]: + for element in message_content: if isinstance(element, dict): if element["type"] == "text": _part = BedrockContentBlock(text=element["text"]) @@ -4791,8 +4790,8 @@ def _bedrock_converse_messages_pt( if _cache_point_block is not None: _parts.append(_cache_point_block) user_content.extend(_parts) - elif message_block["content"] and isinstance(message_block["content"], str): - _part = BedrockContentBlock(text=messages[msg_i]["content"]) + elif message_content and isinstance(message_content, str): + _part = BedrockContentBlock(text=message_content) _cache_point_block = litellm.AmazonConverseConfig().get_cache_point_block( message_block, block_type="content_block", model=model ) diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index 30e521b7671..65c2fccceeb 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -27,11 +27,13 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( from litellm.litellm_core_utils.prompt_templates.factory import ( THOUGHT_SIGNATURE_SEPARATOR, ) +from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import message_field, parts_of from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.types.llms.anthropic import ( ANTHROPIC_HOSTED_TOOLS, ANTHROPIC_MID_CONVERSATION_OUTPUT_CONFIG_BETA_HEADER, + ANTHROPIC_MID_CONVERSATION_TOOL_CHANGES_BETA_HEADER, ANTHROPIC_OAUTH_BETA_HEADER, ANTHROPIC_OAUTH_TOKEN_PREFIX, ANTHROPIC_THINKING_DISPLAY_UPDATES_BETA_HEADER, @@ -344,6 +346,18 @@ class AnthropicModelInfo(BaseLLMModelInfo): return False return thinking.get("type") in ("adaptive", "enabled") and thinking.get("display") == "updates" + def is_mid_conversation_tool_change_used(self, messages: Sequence[object]) -> bool: + for message in messages: + if message_field(message, "role") != "system": + continue + for block in parts_of(message_field(message, "content")): + if ( + message_field(block, "type") in ("tool_addition", "tool_removal") + and message_field(message_field(block, "tool"), "type") == "tool_reference" + ): + return True + return False + def is_mid_conversation_output_config_used(self, messages: list[AllMessageValues]) -> bool: """ Return if "output_config" is in a message @@ -881,6 +895,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): custom_llm_provider: str, is_mid_conversation_output_config_used: bool = False, is_thinking_display_updates_used: bool = False, + is_mid_conversation_tool_change_used: bool = False, ) -> list[str]: """ Get list of common beta headers based on the features that are active. @@ -919,7 +934,10 @@ class AnthropicModelInfo(BaseLLMModelInfo): thinking_display_betas: Final = ( (ANTHROPIC_THINKING_DISPLAY_UPDATES_BETA_HEADER,) if is_thinking_display_updates_used else () ) - return list(set(betas).union(thinking_display_betas)) + tool_change_betas: Final = ( + (ANTHROPIC_MID_CONVERSATION_TOOL_CHANGES_BETA_HEADER,) if is_mid_conversation_tool_change_used else () + ) + return list(set(betas).union(thinking_display_betas, tool_change_betas)) @staticmethod def _make_api_key_auth_header(api_key: str, api_base: str | None, use_bearer_for_custom_base: bool = False) -> dict: @@ -953,6 +971,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): use_bearer_for_custom_base: bool = False, is_mid_conversation_output_config_used: bool = False, is_thinking_display_updates_used: bool = False, + is_mid_conversation_tool_change_used: bool = False, ) -> dict: betas: Final = set() # Anthropic no longer requires the prompt-caching beta header @@ -1010,7 +1029,8 @@ class AnthropicModelInfo(BaseLLMModelInfo): betas.update(user_anthropic_beta_headers) all_betas: Final = betas.union( - (ANTHROPIC_THINKING_DISPLAY_UPDATES_BETA_HEADER,) if is_thinking_display_updates_used else () + (ANTHROPIC_THINKING_DISPLAY_UPDATES_BETA_HEADER,) if is_thinking_display_updates_used else (), + (ANTHROPIC_MID_CONVERSATION_TOOL_CHANGES_BETA_HEADER,) if is_mid_conversation_tool_change_used else (), ) # Don't send any beta headers to Vertex, except web search which is required @@ -1080,6 +1100,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): file_id_used=file_id_used, is_mid_conversation_output_config_used=is_mid_conversation_output_config_used, is_thinking_display_updates_used=self.is_thinking_display_updates_used(optional_params.get("thinking")), + is_mid_conversation_tool_change_used=self.is_mid_conversation_tool_change_used(messages), web_search_tool_used=web_search_tool_used, is_vertex_request=optional_params.get("is_vertex_request", False), user_anthropic_beta_headers=user_anthropic_beta_headers, diff --git a/litellm/llms/anthropic/pass_through/messages/transformation.py b/litellm/llms/anthropic/pass_through/messages/transformation.py index b1be92e49b6..2fbb51ec949 100644 --- a/litellm/llms/anthropic/pass_through/messages/transformation.py +++ b/litellm/llms/anthropic/pass_through/messages/transformation.py @@ -12,6 +12,7 @@ from litellm.llms.base_llm.anthropic_messages.transformation import ( from litellm.types.llms.anthropic import ( ANTHROPIC_ADVISOR_TOOL_TYPE, ANTHROPIC_BETA_HEADER_VALUES, + ANTHROPIC_MID_CONVERSATION_TOOL_CHANGES_BETA_HEADER, ANTHROPIC_THINKING_DISPLAY_UPDATES_BETA_HEADER, AnthropicMessagesRequest, ) @@ -694,7 +695,12 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): if AnthropicModelInfo().is_thinking_display_updates_used(optional_params.get("thinking")) else () ) - all_beta_values: Final = beta_values.union(thinking_display_betas) + tool_change_betas: Final = ( + (ANTHROPIC_MID_CONVERSATION_TOOL_CHANGES_BETA_HEADER,) + if AnthropicModelInfo().is_mid_conversation_tool_change_used(messages) + else () + ) + all_beta_values: Final = beta_values.union(thinking_display_betas, tool_change_betas) if not all_beta_values: return headers diff --git a/litellm/llms/base_llm/harness/__init__.py b/litellm/llms/base_llm/harness/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/base_llm/harness/transformation.py b/litellm/llms/base_llm/harness/transformation.py new file mode 100644 index 00000000000..643f808d4d8 --- /dev/null +++ b/litellm/llms/base_llm/harness/transformation.py @@ -0,0 +1,151 @@ +""" +Base agent-harness transformation configuration. + +A harness is a complete agent runtime (Claude Code, Codex, OpenCode, Deep Agents). +Like the LLM provider configs in `litellm/llms/base_llm/chat/transformation.py`, a +harness config only translates: LiteLLM's session parameters in, the runtime's native +command / config / event stream out. It never does I/O. A handler in +`litellm/harness/handlers/` owns the sandbox, the process and the per-session model +endpoint, and calls these transforms. + +Adding a CLI harness is one subclass of `BaseCLIHarnessConfig` in +`litellm/llms//harness/transformation.py`, plus one line in +`ProviderConfigManager.get_provider_harness_config`. +""" + +from __future__ import annotations + +from abc import ABC, abstractmethod +from collections.abc import Mapping, Sequence +from dataclasses import dataclass, field +from typing import TYPE_CHECKING, Any, ClassVar, Generic, TypeVar + +from litellm.harness.errors import HarnessError, OptionsMismatch +from litellm.harness.types import Capabilities, Event, Harness + +if TYPE_CHECKING: + from litellm.harness.context import SessionContext + +# A config's typed options (ClaudeCodeOptions, CodexOptions, ...) and its per-turn parser state. +OptionsT = TypeVar("OptionsT") +StreamStateT = TypeVar("StreamStateT") + + +def event_list(*events: Event) -> Sequence[Event]: + """A transform_stream_line result. One place builds it so every parser returns the same shape.""" + return list(events) # mutable-ok: stream-line results are list-shaped; callers and tests compare with list literals + + +class HarnessTurnError(HarnessError): + """The runtime reported a failed turn. The runtime maps this to stop_reason='runtime_error'.""" + + +@dataclass(frozen=True) +class HarnessSessionSetup: + """What the handler must prepare in the sandbox before the first turn. + + Paths are relative to `private_dir` (a per-session temp dir inside the sandbox) + unless they are absolute. + """ + + files: Mapping[str, bytes] = field(default_factory=dict) + # (dir inside private_dir, cache subpath under ~/.cache/litellm-harness) linked so a + # later session can resume the runtime's own conversation. + persisted_dirs: Sequence[tuple[str, str]] = () + # Where skill folders are copied, relative to private_dir, or absolute. + skills_dir: str | None = None + # Env passed on every turn. Values may contain `{private_dir}`. + env: Mapping[str, str] = field(default_factory=dict) + + +@dataclass(frozen=True) +class HarnessTurnRequest: + """One turn of a CLI runtime: the process to run and what to send on stdin.""" + + argv: Sequence[str] + env: Mapping[str, str] + stdin: str + cwd: str | None = None + + +@dataclass(frozen=True) +class HarnessTurnResponse: + """What the runtime produced for one turn, after the process exited.""" + + final_text: str + output_json: str | None = None + + +class BaseHarnessConfig(ABC, Generic[OptionsT]): + """Declares what a harness is and validates a session before anything starts.""" + + harness: ClassVar[Harness] + options_type: type[OptionsT] + capabilities: ClassVar[Capabilities] + # CLI runtimes call a per-session model endpoint; in-process ones call LiteLLM directly. + uses_model_endpoint: ClassVar[bool] = True + + def get_options(self, ctx: SessionContext) -> OptionsT: + """ctx.options, or this harness's default options.""" + options = ctx.options + if options is None: + return self.options_type() + if not isinstance(options, self.options_type): + raise OptionsMismatch( + f"{type(options).__name__} cannot be used with Harness.{self.harness.name}; " + f"use {self.options_type.__name__}" + ) + return options + + def validate_environment(self, ctx: SessionContext) -> None: + """Static checks on the session. Raise OptionsMismatch / ValueError early.""" + self.get_options(ctx) + + +class BaseCLIHarnessConfig(BaseHarnessConfig[OptionsT], Generic[OptionsT, StreamStateT]): + """A runtime driven as a subprocess that prints one JSON event per line.""" + + @abstractmethod + def get_binary(self) -> str: + """Executable that must be on the sandbox's PATH.""" + + @abstractmethod + def get_install_hint(self) -> str: + """How to install the binary; shown in HarnessInstallFailed.""" + + @abstractmethod + def transform_session_setup(self, ctx: SessionContext, private_dir: str) -> HarnessSessionSetup: + """Config files, env and persisted dirs for the session.""" + + @abstractmethod + def transform_turn_request( + self, + ctx: SessionContext, + setup: HarnessSessionSetup, + private_dir: str, + prompt: str, + native_session_id: str | None, + ) -> HarnessTurnRequest: + """argv / env / stdin for one turn. native_session_id is set after the first turn.""" + + @abstractmethod + def create_stream_state(self) -> StreamStateT: + """Fresh per-turn parser state.""" + + @abstractmethod + def transform_stream_line(self, line: Mapping[str, Any], state: StreamStateT) -> Sequence[Event]: + """One decoded JSON line from stdout to zero or more events. Pure.""" + + @abstractmethod + def get_native_session_id(self, state: StreamStateT) -> str | None: + """The runtime's own session / thread id, once the stream has reported it.""" + + @abstractmethod + def transform_turn_response( + self, + ctx: SessionContext, + state: StreamStateT, + exit_code: int, + stderr_tail: Sequence[str], + ) -> HarnessTurnResponse: + """Final text and structured output, or raise HarnessTurnError.""" diff --git a/litellm/llms/base_llm/harness/utils.py b/litellm/llms/base_llm/harness/utils.py new file mode 100644 index 00000000000..7c0278d5e36 --- /dev/null +++ b/litellm/llms/base_llm/harness/utils.py @@ -0,0 +1,116 @@ +"""Pure helpers shared by harness configs.""" + +from __future__ import annotations + +import itertools +import json +import os +from collections.abc import Iterator, Mapping, Sequence +from types import MappingProxyType +from typing import Any, Final, TypeAlias + +from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH + +# A decoded JSON document: what json.loads / model_json_schema() produce. +JSONValue: TypeAlias = "dict[str, JSONValue] | list[JSONValue] | str | int | float | bool | None" + +SKILL_MANIFEST: Final = "SKILL.md" +_JSON_DECODER: Final = json.JSONDecoder() + + +def normalize_tool_name(native_name: str, mapping: Mapping[str, str]) -> str: + """Normalized tool name (read, write, edit, bash, ...) or the native name if unmapped.""" + return mapping.get(native_name, native_name) + + +def native_tool_names(normalized: Sequence[str], mapping: Mapping[str, Sequence[str]]) -> Sequence[str]: + """Native names for normalized tool names, de-duplicated, order kept.""" + expanded: Final = itertools.chain.from_iterable(mapping.get(name, (name,)) for name in normalized) + return list(dict.fromkeys(expanded)) # mutable-ok: public helper whose callers/tests compare against list literals + + +def last_json_object(text: str) -> str | None: + """The last top-level `{...}` in text that parses as a JSON object, re-serialized.""" + last: str | None = None + index = text.find("{") + while index != -1: + try: + obj, end = _JSON_DECODER.raw_decode(text, index) + except json.JSONDecodeError: + index = text.find("{", index + 1) + continue + if isinstance(obj, dict): + last = json.dumps(obj) + index = text.find("{", end) + return last + + +def structured_output_instruction(schema: Mapping[str, Any]) -> str: + return ( + "When you have finished, your final message must be a single JSON object that " + "matches this JSON schema, with no other text before or after it:\n" + f"{json.dumps(schema)}" + ) + + +def strict_json_schema(schema: JSONValue, depth: int = 0) -> JSONValue: + """Make a JSON schema acceptable to OpenAI strict structured outputs. + + Every object gets `additionalProperties: false` and all of its properties required, + recursively. Keywords strict mode rejects next to $ref are dropped. Nesting deeper than + DEFAULT_MAX_RECURSE_DEPTH raises instead of recursing further. + """ + if depth > DEFAULT_MAX_RECURSE_DEPTH: + raise ValueError(f"output schema is nested deeper than {DEFAULT_MAX_RECURSE_DEPTH} levels") + if isinstance(schema, list): + return [strict_json_schema(entry, depth + 1) for entry in schema] # mutable-ok: JSON document output + if not isinstance(schema, dict): + return schema + entries: Final = ((key, strict_json_schema(value, depth + 1)) for key, value in schema.items()) + result = dict(entries) # mutable-ok: JSON document; "default" is popped below + if "$ref" in result: + return {"$ref": result["$ref"]} # mutable-ok: JSONValue output is a plain JSON document + result.pop("default", None) + properties = result.get("properties") + if result.get("type") == "object" or isinstance(properties, dict): + props = properties if isinstance(properties, dict) else {} # mutable-ok: JSONValue object member + required: Final[list[JSONValue]] = list(props) # mutable-ok: JSON array in the output schema + strict: Final[Mapping[str, JSONValue]] = MappingProxyType( + {"properties": props, "required": required, "additionalProperties": False} + ) + result = {**result, **strict} # mutable-ok: JSON document output + return result + + +def decode_json_line(line: bytes | str) -> Mapping[str, Any] | None: + """One JSONL line as a dict, or None for blank / non-JSON / non-object lines.""" + text = line.strip() + if not text: + return None + try: + obj = json.loads(text) + except json.JSONDecodeError: + return None + return obj if isinstance(obj, dict) else None + + +def stderr_tail_text(stderr_tail: Sequence[str]) -> str: + return "\n".join(line for line in stderr_tail if line.strip()) + + +def _read_bytes(path: str) -> bytes: + with open(path, "rb") as fh: + return fh.read() + + +def _walk_files(root: str) -> Iterator[str]: + for dirpath, _dirnames, filenames in os.walk(root): + yield from (os.path.join(dirpath, filename) for filename in sorted(filenames)) + + +def read_skill_files(skill_dir: str) -> tuple[tuple[str, bytes], ...]: + """(relative path, bytes) for every file under a local skill folder.""" + root: Final = os.path.realpath(os.fspath(skill_dir)) + if not os.path.isfile(os.path.join(root, SKILL_MANIFEST)): + raise ValueError(f"skill folder {skill_dir!r} has no {SKILL_MANIFEST}") + return tuple((os.path.relpath(path, root), _read_bytes(path)) for path in _walk_files(root)) diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index c9fec2db1d7..48a8b1b44bb 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -1448,7 +1448,7 @@ class AmazonConverseConfig(BaseConfig): ) def _converted_text_blocks(self, message: ChatCompletionSystemMessage) -> tuple[ChatCompletionTextObject, ...]: - content: Final = message["content"] + content: Final = message.get("content") if isinstance(content, str): return (self._converted_text_block(content, message.get("cache_control")),) if content else () parts: Final[Sequence[object]] = content or () @@ -1483,13 +1483,14 @@ class AmazonConverseConfig(BaseConfig): for message in hoisted: if message["role"] != "system": continue - if isinstance(message["content"], str) and message["content"]: - system_content_blocks.append(SystemContentBlock(text=message["content"])) + content = message.get("content") + if isinstance(content, str) and content: + system_content_blocks.append(SystemContentBlock(text=content)) cache_block = self.get_cache_point_block(message, block_type="system", model=model) if cache_block: system_content_blocks.append(cache_block) - elif isinstance(message["content"], list): - for m in message["content"]: + elif isinstance(content, list): + for m in content: if m.get("type") == "text" and m.get("text"): system_content_blocks.append(SystemContentBlock(text=m["text"])) cache_block = self.get_cache_point_block(m, block_type="system", model=model) diff --git a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py index 73da7c41a09..6234ca3a9c3 100644 --- a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py @@ -549,6 +549,9 @@ class AmazonAnthropicClaudeMessagesConfig( is_thinking_display_updates_used=anthropic_model_info.is_thinking_display_updates_used( anthropic_messages_request.get("thinking") ), + is_mid_conversation_tool_change_used=anthropic_model_info.is_mid_conversation_tool_change_used( + outgoing_messages_typed + ), ) beta_set.update(auto_betas) diff --git a/litellm/llms/claude_code/__init__.py b/litellm/llms/claude_code/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/claude_code/harness/__init__.py b/litellm/llms/claude_code/harness/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/claude_code/harness/transformation.py b/litellm/llms/claude_code/harness/transformation.py new file mode 100644 index 00000000000..a85968f78be --- /dev/null +++ b/litellm/llms/claude_code/harness/transformation.py @@ -0,0 +1,389 @@ +""" +Claude Code harness config: `claude -p --output-format stream-json`, once per turn. + +Every model call goes to the per-session endpoint with the per-session token. The CLI gets +a private CLAUDE_CONFIG_DIR and only the `user` setting source (that private dir), so +neither the user's login, keychain, nor a repo's `.claude/settings.json` can swap the base +URL or credentials. Verified against Claude Code 2.1.285. +""" + +from __future__ import annotations + +import itertools +import json +from collections.abc import Mapping, Sequence +from dataclasses import dataclass, field +from types import MappingProxyType +from typing import TYPE_CHECKING, Any, Final + +from litellm.harness.errors import HarnessError, OptionsMismatch +from litellm.harness.options import ClaudeCodeOptions +from litellm.harness.types import ( + Capabilities, + Compaction, + Event, + Harness, + Reasoning, + Text, + ToolCall, + ToolResult, +) +from litellm.llms.base_llm.harness.transformation import ( + BaseCLIHarnessConfig, + HarnessSessionSetup, + HarnessTurnError, + HarnessTurnRequest, + HarnessTurnResponse, + event_list, +) +from litellm.llms.base_llm.harness.utils import ( + last_json_object, + native_tool_names, + normalize_tool_name, + stderr_tail_text, +) + +if TYPE_CHECKING: + from litellm.harness.context import SessionContext + +CLAUDE_BINARY: Final = "claude" +SYNTHETIC_MODEL: Final = "" + +BASE_COMMAND: Final = ("-p", "--output-format", "stream-json", "--verbose", "--input-format", "text") + +PERMISSION_MODES: Final[Mapping[str, str]] = MappingProxyType( + { + "read-only": "plan", + "ask": "default", + "edit": "acceptEdits", + "full": "bypassPermissions", + } +) + +NATIVE_TO_NORMALIZED: Final[Mapping[str, str]] = MappingProxyType( + { + "Read": "read", + "Write": "write", + "Edit": "edit", + "MultiEdit": "edit", + "Bash": "bash", + "Glob": "glob", + "Grep": "grep", + "WebSearch": "web_search", + "LS": "ls", + } +) + +NORMALIZED_TO_NATIVE: Final[Mapping[str, tuple[str, ...]]] = MappingProxyType( + { + "read": ("Read",), + "write": ("Write",), + "edit": ("Edit", "MultiEdit"), + "bash": ("Bash",), + "glob": ("Glob",), + "grep": ("Grep",), + "web_search": ("WebSearch",), + "ls": ("LS",), + } +) + +# Env the config owns; ClaudeCodeOptions.env may not override these. +MANAGED_ENV_KEYS: Final = frozenset( + { + "ANTHROPIC_BASE_URL", + "ANTHROPIC_AUTH_TOKEN", + "ANTHROPIC_API_KEY", + "ANTHROPIC_MODEL", + "ANTHROPIC_SMALL_FAST_MODEL", + "CLAUDE_CONFIG_DIR", + } +) + +# Claude Code settings.json keys LiteLLM manages (or that could reroute model calls or credentials). +MANAGED_CONFIG_KEYS: Final = frozenset( + {"env", "apiKeyHelper", "model", "permissions", "awsAuthRefresh", "awsCredentialExport", "forceLoginMethod"} +) + +STATIC_ENV: Final[Mapping[str, str]] = MappingProxyType( + { + "DISABLE_TELEMETRY": "1", + "DISABLE_ERROR_REPORTING": "1", + "DISABLE_AUTOUPDATER": "1", + "CLAUDE_CODE_DISABLE_NONESSENTIAL_TRAFFIC": "1", + } +) + +STRUCTURED_OUTPUT_INSTRUCTION: Final = ( + "When you have finished the task, end your final reply with only a single JSON " + "object (no code fences, no prose after it) that matches this JSON schema:\n{schema}" +) + + +@dataclass +class ClaudeCodeStreamState: + """What the parser has learned from one turn's stream-json output.""" + + session_id: str | None = None + text_parts: list[str] = field(default_factory=list) # mutable-ok: parser appends text deltas + result_seen: bool = False + result_text: str | None = None + is_error: bool = False + errors: Sequence[str] = () + structured_output: Any | None = None + + @property + def final_text(self) -> str: + if self.result_text is not None: + return self.result_text + return "".join(self.text_parts) + + +def stringify_tool_output(content: object) -> str: + """tool_result content is a string or a list of content blocks.""" + if content is None: + return "" + if isinstance(content, str): + return content + if isinstance(content, list): + return "\n".join(_stringify_block(block) for block in content) + return json.dumps(content, ensure_ascii=False) + + +def _stringify_block(block: object) -> str: + if isinstance(block, dict) and block.get("type") == "text": + return str(block.get("text", "")) + if isinstance(block, str): + return block + return json.dumps(block, ensure_ascii=False) + + +def _message_blocks(event: Mapping[str, Any]) -> Sequence[Any]: + message: Final = event.get("message") + content: Final = message.get("content") if isinstance(message, Mapping) else None + if isinstance(content, str): + return ({"type": "text", "text": content},) # mutable-ok: JSON content block, like the stream's + return content if isinstance(content, list) else () + + +def _assistant_block_events(block: Mapping[str, Any], state: ClaudeCodeStreamState) -> tuple[Event, ...]: + kind = block.get("type") + if kind == "text" and block.get("text"): + state.text_parts.append(block["text"]) + return (Text(delta=block["text"]),) + if kind == "thinking" and block.get("thinking"): + return (Reasoning(delta=block["thinking"]),) + if kind == "tool_use": + native = str(block.get("name", "")) + return ( + ToolCall( + id=str(block.get("id", "")), + name=normalize_tool_name(native, NATIVE_TO_NORMALIZED), + native_name=native, + input=block.get("input") + or {}, # mutable-ok: ToolCall.input is a dict field; empty default for a missing input + builtin=not native.startswith("mcp__"), + ), + ) + return () + + +def _assistant_events(event: Mapping[str, Any], state: ClaudeCodeStreamState) -> Sequence[Event]: + if event.get("parent_tool_use_id"): + return event_list() # subagent traffic + message: Final = event.get("message") + if isinstance(message, Mapping) and message.get("model") == SYNTHETIC_MODEL: + return event_list() # CLI-generated error text; surfaced via the result event + blocks: Final = (block for block in _message_blocks(event) if isinstance(block, dict)) + return event_list(*itertools.chain.from_iterable(_assistant_block_events(block, state) for block in blocks)) + + +def _is_tool_result(block: object) -> bool: + return isinstance(block, dict) and block.get("type") == "tool_result" + + +def _user_events(event: Mapping[str, Any]) -> Sequence[Event]: + if event.get("parent_tool_use_id"): + return event_list() + return event_list( + *( + ToolResult( + id=str(block.get("tool_use_id", "")), + output=stringify_tool_output(block.get("content")), + is_error=bool(block.get("is_error", False)), + ) + for block in _message_blocks(event) + if _is_tool_result(block) + ) + ) + + +def _system_events(event: Mapping[str, Any], state: ClaudeCodeStreamState) -> Sequence[Event]: + subtype = event.get("subtype") + if subtype == "init" and event.get("session_id"): + state.session_id = str(event["session_id"]) + return event_list() + if subtype == "compact_boundary": + meta: Final = event.get("compact_metadata") + pre_tokens: Final = meta.get("pre_tokens") if isinstance(meta, Mapping) else None + return event_list(Compaction(tokens_before=pre_tokens, tokens_after=None)) + return event_list() + + +def _record_result(event: Mapping[str, Any], state: ClaudeCodeStreamState) -> Sequence[Event]: + state.result_seen = True + state.is_error = bool(event.get("is_error", False)) + result = event.get("result") + state.result_text = result if isinstance(result, str) else None + state.errors = [str(e) for e in event.get("errors") or ()] # mutable-ok: mirrors the JSON errors array + state.structured_output = event.get("structured_output") + if event.get("session_id"): + state.session_id = str(event["session_id"]) + return event_list() + + +def turn_error_message(state: ClaudeCodeStreamState, exit_code: int, stderr_tail: Sequence[str]) -> str | None: + """None if the turn succeeded, else the message for HarnessTurnError.""" + if exit_code == 0 and state.result_seen and not state.is_error: + return None + reason = state.result_text or "; ".join(state.errors) + if not reason: + reason = "no result event" if not state.result_seen else "unknown error" + message = f"claude exited with code {exit_code}: {reason}" + tail = stderr_tail_text(stderr_tail) + return f"{message}\nstderr:\n{tail}" if tail else message + + +def build_system_prompt(instructions: str | None, output_schema: Mapping[str, Any] | None) -> str | None: + schema_part: Final = ( + STRUCTURED_OUTPUT_INSTRUCTION.format(schema=json.dumps(output_schema)) if output_schema is not None else None + ) + parts: Final = tuple(part for part in (instructions, schema_part) if part) + return "\n\n".join(parts) if parts else None + + +class ClaudeCodeHarnessConfig(BaseCLIHarnessConfig): + harness = Harness.CLAUDE_CODE + options_type = ClaudeCodeOptions + capabilities = Capabilities( + structured_output=True, + tool_approval=False, + tool_filtering=True, + history=False, + custom_tools=False, + skills=True, + resume=True, + permission_modes=frozenset({"read-only", "edit", "full"}), + ) + + def get_binary(self) -> str: + return CLAUDE_BINARY + + def get_install_hint(self) -> str: + return "npm install -g @anthropic-ai/claude-code" + + def validate_environment(self, ctx: SessionContext) -> None: + options: ClaudeCodeOptions = self.get_options(ctx) + clashing = sorted(MANAGED_ENV_KEYS.intersection(options.env)) + if clashing: + raise OptionsMismatch(f"ClaudeCodeOptions.env may not set {', '.join(clashing)}; LiteLLM manages it") + managed = sorted(MANAGED_CONFIG_KEYS.intersection(options.config)) + if managed: + raise OptionsMismatch( + f"ClaudeCodeOptions.config may not set {', '.join(managed)}; " + "use the matching agent() argument (model=, permissions=) instead" + ) + + def transform_session_setup(self, ctx: SessionContext, private_dir: str) -> HarnessSessionSetup: + if ctx.endpoint is None or not ctx.endpoint.token: + raise HarnessError("Claude Code needs the session model endpoint") + options: ClaudeCodeOptions = self.get_options(ctx) + model = ctx.model + # Background calls (titles, summaries) use the same model group, like OpenCode. + model_env: Final = ( + MappingProxyType({"ANTHROPIC_MODEL": model, "ANTHROPIC_SMALL_FAST_MODEL": model}) + if model + else MappingProxyType({}) + ) + env: Final = MappingProxyType( + { + **options.env, + **STATIC_ENV, + "ANTHROPIC_BASE_URL": ctx.sandbox.host_url(ctx.endpoint.port), + "ANTHROPIC_AUTH_TOKEN": ctx.endpoint.token, + "ANTHROPIC_API_KEY": "", + "CLAUDE_CONFIG_DIR": private_dir, + **model_env, + } + ) + return HarnessSessionSetup( + persisted_dirs=[("projects", "claude_code/projects")], # mutable-ok: tests compare to a list + skills_dir="skills", + env=env, + ) + + def transform_turn_request( + self, + ctx: SessionContext, + setup: HarnessSessionSetup, + private_dir: str, + prompt: str, + native_session_id: str | None, + ) -> HarnessTurnRequest: + options: ClaudeCodeOptions = self.get_options(ctx) + schema = ctx.output.model_json_schema() if ctx.output is not None else None + system_prompt: Final = build_system_prompt(ctx.instructions, schema) + disallowed: Final = native_tool_names(ctx.disable_tools, NORMALIZED_TO_NATIVE) + config: Final = dict(options.config) # mutable-ok: json.dumps needs a plain dict + settings: Final = json.dumps(config) if config else None + argv: Final = ( + CLAUDE_BINARY, + *BASE_COMMAND, + "--permission-mode", + PERMISSION_MODES[ctx.permissions], + *(("--model", ctx.model) if ctx.model else ()), + # Only read settings from the private CLAUDE_CONFIG_DIR, never the repo's .claude/. + "--setting-sources", + "user", + *(("--settings", settings) if settings else ()), + *(("--append-system-prompt", system_prompt) if system_prompt else ()), + *(("--max-turns", str(ctx.max_turns)) if ctx.max_turns is not None else ()), + *(("--disallowedTools", ",".join(disallowed)) if disallowed else ()), + *(("--resume", native_session_id) if native_session_id else ()), + ) + return HarnessTurnRequest(argv=argv, env=setup.env, stdin=prompt) + + def create_stream_state(self) -> ClaudeCodeStreamState: + return ClaudeCodeStreamState() + + def transform_stream_line(self, line: Mapping[str, Any], state: ClaudeCodeStreamState) -> Sequence[Event]: + kind = line.get("type") + if kind == "assistant": + return _assistant_events(line, state) + if kind == "user": + return _user_events(line) + if kind == "system": + return _system_events(line, state) + if kind == "result": + return _record_result(line, state) + return event_list() + + def get_native_session_id(self, state: ClaudeCodeStreamState) -> str | None: + return state.session_id + + def transform_turn_response( + self, + ctx: SessionContext, + state: ClaudeCodeStreamState, + exit_code: int, + stderr_tail: Sequence[str], + ) -> HarnessTurnResponse: + error = turn_error_message(state, exit_code, stderr_tail) + if error is not None: + raise HarnessTurnError(error) + output_json: str | None = None + if ctx.output is not None: + if isinstance(state.structured_output, dict): + output_json = json.dumps(state.structured_output) + else: + output_json = last_json_object(state.final_text) + return HarnessTurnResponse(final_text=state.final_text, output_json=output_json) diff --git a/litellm/llms/codex/__init__.py b/litellm/llms/codex/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/codex/harness/__init__.py b/litellm/llms/codex/harness/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/codex/harness/transformation.py b/litellm/llms/codex/harness/transformation.py new file mode 100644 index 00000000000..ab17cf3d869 --- /dev/null +++ b/litellm/llms/codex/harness/transformation.py @@ -0,0 +1,349 @@ +""" +Codex harness config: `codex exec --json` (JSONL events), once per turn. + +Every model call goes to one custom provider (`litellm`, wire_api=responses) pointing at the +per-session endpoint. The bearer token only travels in the LITELLM_HARNESS_TOKEN env var, +never in argv. CODEX_HOME is the private session dir so the user's own Codex config and +auth are never read. Verified against codex-cli 0.135.0. +""" + +from __future__ import annotations + +import itertools +import json +import re +from collections.abc import Iterator, Mapping, Sequence +from dataclasses import dataclass, field +from types import MappingProxyType +from typing import TYPE_CHECKING, Any, Final + +from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH +from litellm.harness.errors import HarnessError, OptionsMismatch +from litellm.harness.options import CodexOptions +from litellm.harness.types import ( + Capabilities, + Event, + Harness, + Reasoning, + Text, + ToolCall, + ToolResult, +) +from litellm.llms.base_llm.harness.transformation import ( + BaseCLIHarnessConfig, + HarnessSessionSetup, + HarnessTurnError, + HarnessTurnRequest, + HarnessTurnResponse, + event_list, +) +from litellm.llms.base_llm.harness.utils import stderr_tail_text, strict_json_schema + +if TYPE_CHECKING: + from litellm.harness.context import SessionContext + +CODEX_BINARY: Final = "codex" +CODEX_PROVIDER_ID: Final = "litellm" +CODEX_TOKEN_ENV: Final = "LITELLM_HARNESS_TOKEN" +CODEX_SCHEMA_FILENAME: Final = "output_schema.json" +# Top-level config keys LiteLLM sets itself; users may not override them via options.config. +MANAGED_CONFIG_KEYS: Final = frozenset( + { + "model", + "model_provider", + "model_providers", + "approval_policy", + "sandbox_mode", + "mcp_servers", + "developer_instructions", + "web_search", + } +) +_BARE_TOML_KEY: Final = re.compile(r"^[A-Za-z0-9_-]+$") +_TOOL_ITEM_TYPES: Final = frozenset({"command_execution", "file_change", "web_search", "mcp_tool_call"}) + + +@dataclass +class CodexStreamState: + """What the parser has learned from one turn's JSONL events.""" + + thread_id: str | None = None + final_text: str = "" + error: str | None = None + failed: bool = False + started: set[str] = field(default_factory=set) # mutable-ok: parser records announced tool items + + +def _tool_input(item: Mapping[str, Any]) -> tuple[str, str, Mapping[str, Any], bool]: + """(normalized name, native name, input, builtin) for a tool-like item.""" + item_type = item.get("type") + if item_type == "command_execution": + return "bash", "command_execution", MappingProxyType({"command": item.get("command", "")}), True + if item_type == "file_change": + changes: Final = list(item.get("changes") or ()) # mutable-ok: JSON array, as codex reports it + return "edit", "apply_patch", MappingProxyType({"changes": changes}), True + if item_type == "web_search": + return "web_search", "web_search", MappingProxyType({"query": item.get("query", "")}), True + server = str(item.get("server") or "") + tool = str(item.get("tool") or "") + arguments = item.get("arguments") + tool_args = arguments if isinstance(arguments, dict) else MappingProxyType({"arguments": arguments}) + name = f"{server}.{tool}" if server else tool + return name, tool, tool_args, False + + +def _tool_output(item: Mapping[str, Any]) -> tuple[str, bool]: + """(output text, is_error) for a completed tool-like item.""" + item_type = item.get("type") + status = item.get("status") + if item_type == "command_execution": + exit_code = item.get("exit_code") + is_error = status == "failed" or (exit_code is not None and exit_code != 0) + return str(item.get("aggregated_output") or ""), is_error + if item_type == "file_change": + lines = (f"{c.get('kind', '')} {c.get('path', '')}".strip() for c in item.get("changes") or ()) + return "\n".join(lines), status == "failed" + if item_type == "web_search": + return "", status == "failed" + error = item.get("error") + if error: + message = error.get("message") if isinstance(error, dict) else error + return str(message), True + result = item.get("result") + if result is None: + return "", status == "failed" + if isinstance(result, str): + return result, status == "failed" + return json.dumps(result), status == "failed" + + +def _tool_item_events( + item_id: str, item: Mapping[str, Any], completed: bool, state: CodexStreamState +) -> Iterator[Event]: + if item_id not in state.started: + state.started.add(item_id) + name, native_name, tool_input, builtin = _tool_input(item) + yield ToolCall(id=item_id, name=name, native_name=native_name, input=tool_input, builtin=builtin) + if completed: + output, is_error = _tool_output(item) + yield ToolResult(id=item_id, output=output, is_error=is_error) + + +def _item_events(event_type: str, item: Mapping[str, Any], state: CodexStreamState) -> Sequence[Event]: + item_type = item.get("type") + item_id = str(item.get("id") or "") + completed = event_type == "item.completed" + if item_type == "agent_message": + if not completed: + return event_list() + text = str(item.get("text") or "") + state.final_text = text + return event_list(Text(delta=text)) if text else event_list() + if item_type == "reasoning": + text = str(item.get("text") or "") + return event_list(Reasoning(delta=text)) if completed and text else event_list() + if item_type not in _TOOL_ITEM_TYPES: + return event_list() + return event_list(*_tool_item_events(item_id, item, completed, state)) + + +def toml_value(value: object, depth: int = 0) -> str: + """Encode a Python value as a TOML value for `codex -c key=value`.""" + if depth > DEFAULT_MAX_RECURSE_DEPTH: + raise OptionsMismatch(f"CodexOptions.config is nested deeper than {DEFAULT_MAX_RECURSE_DEPTH} levels") + if isinstance(value, bool): + return "true" if value else "false" + if isinstance(value, (int, float)): + return repr(value) + if isinstance(value, str): + return json.dumps(value) + if isinstance(value, Mapping): + pairs = ", ".join(f"{toml_key(k)} = {toml_value(v, depth + 1)}" for k, v in value.items()) + return "{" + pairs + "}" + if isinstance(value, (list, tuple)): + return "[" + ", ".join(toml_value(v, depth + 1) for v in value) + "]" + raise OptionsMismatch(f"CodexOptions.config value of type {type(value).__name__} cannot be passed to codex") + + +def toml_key(key: object) -> str: + text = str(key) + return text if _BARE_TOML_KEY.match(text) else json.dumps(text) + + +def _config_override(key: object, value: object) -> str: + dotted = str(key) + if not dotted or "=" in dotted: + raise OptionsMismatch(f"Invalid CodexOptions.config key: {dotted!r}") + if dotted.split(".", 1)[0] in MANAGED_CONFIG_KEYS: + raise OptionsMismatch( + f"CodexOptions.config[{dotted!r}] is managed by LiteLLM; use the matching agent() argument instead" + ) + return f"{dotted}={toml_value(value)}" + + +def config_overrides(config: Mapping[str, Any]) -> Sequence[str]: + """`-c` override strings for CodexOptions.config, rejecting managed keys.""" + overrides: Final = (_config_override(key, value) for key, value in config.items()) + return list(overrides) # mutable-ok: public helper; tests compare to a list + + +def _flag_pairs(flag: str, values: Sequence[str]) -> tuple[str, ...]: + return tuple(itertools.chain.from_iterable((flag, value) for value in values)) + + +class CodexHarnessConfig(BaseCLIHarnessConfig): + harness = Harness.CODEX + options_type = CodexOptions + capabilities = Capabilities( + structured_output=True, + tool_approval=False, + tool_filtering=False, + history=False, + custom_tools=False, + skills=True, + resume=True, + permission_modes=frozenset({"read-only", "full"}), + ) + + def get_binary(self) -> str: + return CODEX_BINARY + + def get_install_hint(self) -> str: + return "npm install -g @openai/codex (or brew install codex)" + + def validate_environment(self, ctx: SessionContext) -> None: + options: CodexOptions = self.get_options(ctx) + config_overrides(options.config) + + def transform_session_setup(self, ctx: SessionContext, private_dir: str) -> HarnessSessionSetup: + if ctx.endpoint is None: + raise HarnessError("Codex needs the session model endpoint") + options: CodexOptions = self.get_options(ctx) + files: Final = ( + MappingProxyType( + {CODEX_SCHEMA_FILENAME: json.dumps(strict_json_schema(ctx.output.model_json_schema())).encode("utf-8")} + ) + if ctx.output is not None + else MappingProxyType({}) + ) + return HarnessSessionSetup( + files=files, + persisted_dirs=[("sessions", "codex/sessions")], # mutable-ok: tests compare to a list + skills_dir="skills", + env=MappingProxyType({**options.env, CODEX_TOKEN_ENV: ctx.endpoint.token, "CODEX_HOME": private_dir}), + ) + + def transform_turn_request( + self, + ctx: SessionContext, + setup: HarnessSessionSetup, + private_dir: str, + prompt: str, + native_session_id: str | None, + ) -> HarnessTurnRequest: + if ctx.endpoint is None: + raise HarnessError("Codex needs the session model endpoint") + options: CodexOptions = self.get_options(ctx) + head: Final = ( + (CODEX_BINARY, "exec", "resume", native_session_id) if native_session_id else (CODEX_BINARY, "exec") + ) + argv: Final = ( + *head, + "--json", + "--skip-git-repo-check", + *(("-m", ctx.model) if ctx.model else ()), + *_flag_pairs("-c", self._provider_overrides(ctx)), + *self._permission_args(ctx, native_session_id), + *_flag_pairs("-c", self._feature_overrides(ctx, options)), + *_flag_pairs("-c", config_overrides(options.config)), + *( + ("--output-schema", f"{private_dir}/{CODEX_SCHEMA_FILENAME}") + if CODEX_SCHEMA_FILENAME in setup.files + else () + ), + *(() if native_session_id else ("-C", ctx.sandbox.workdir)), + "-", + ) + return HarnessTurnRequest(argv=argv, env=setup.env, stdin=prompt, cwd=ctx.sandbox.workdir) + + def _provider_overrides(self, ctx: SessionContext) -> tuple[str, ...]: + assert ctx.endpoint is not None + base_url = ctx.sandbox.host_url(ctx.endpoint.port).rstrip("/") + "/v1" + prefix = f"model_providers.{CODEX_PROVIDER_ID}" + return ( + f"model_provider={CODEX_PROVIDER_ID}", + f"{prefix}.name={CODEX_PROVIDER_ID}", + f"{prefix}.base_url={toml_value(base_url)}", + f"{prefix}.env_key={CODEX_TOKEN_ENV}", + f"{prefix}.wire_api=responses", + "approval_policy=never", + ) + + @staticmethod + def _permission_args(ctx: SessionContext, native_session_id: str | None) -> tuple[str, ...]: + if ctx.permissions == "read-only": + mode = "read-only" + elif getattr(ctx.sandbox, "is_container", False): + # The container is already the boundary; nested sandboxing fails in containers. + return ("--dangerously-bypass-approvals-and-sandbox",) + else: + mode = "workspace-write" + # `codex exec resume` has no --sandbox flag; the config key works for both. + if native_session_id: + return ("-c", f"sandbox_mode={toml_value(mode)}") + return ("--sandbox", mode) + + @staticmethod + def _feature_overrides(ctx: SessionContext, options: CodexOptions) -> tuple[str, ...]: + reasoning: Final = ( + ( + f"model_reasoning_effort={options.reasoning_effort}", + "model_reasoning_summary=auto", + "model_supports_reasoning_summaries=true", + ) + if options.reasoning_effort + else () + ) + instructions: Final = (f"developer_instructions={toml_value(ctx.instructions)}",) if ctx.instructions else () + return (f"web_search={'live' if options.web_search else 'disabled'}", *reasoning, *instructions) + + def create_stream_state(self) -> CodexStreamState: + return CodexStreamState() + + def transform_stream_line(self, line: Mapping[str, Any], state: CodexStreamState) -> Sequence[Event]: + """turn.completed usage is ignored on purpose: the session endpoint accounts it.""" + event_type = line.get("type") + if event_type == "thread.started": + if line.get("thread_id"): + state.thread_id = str(line["thread_id"]) + return event_list() + if event_type in ("item.started", "item.updated", "item.completed"): + item = line.get("item") + return _item_events(str(event_type), item, state) if isinstance(item, dict) else event_list() + if event_type == "error": + state.error = str(line.get("message") or "codex reported an error") + return event_list() + if event_type == "turn.failed": + error = line.get("error") + message = error.get("message") if isinstance(error, dict) else error + state.error = str(message or state.error or "codex turn failed") + state.failed = True + return event_list() + + def get_native_session_id(self, state: CodexStreamState) -> str | None: + return state.thread_id + + def transform_turn_response( + self, + ctx: SessionContext, + state: CodexStreamState, + exit_code: int, + stderr_tail: Sequence[str], + ) -> HarnessTurnResponse: + if state.failed: + raise HarnessTurnError(f"codex turn failed: {state.error}") + if exit_code != 0: + detail = stderr_tail_text(stderr_tail) or state.error or "no output" + raise HarnessTurnError(f"codex exited with code {exit_code}: {detail}") + output_json = state.final_text if ctx.output is not None else None + return HarnessTurnResponse(final_text=state.final_text, output_json=output_json) diff --git a/litellm/llms/deepagents/__init__.py b/litellm/llms/deepagents/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/deepagents/harness/__init__.py b/litellm/llms/deepagents/harness/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/deepagents/harness/sandbox_backend.py b/litellm/llms/deepagents/harness/sandbox_backend.py new file mode 100644 index 00000000000..49610014e21 --- /dev/null +++ b/litellm/llms/deepagents/harness/sandbox_backend.py @@ -0,0 +1,559 @@ +"""Deep Agents pieces that subclass optional-dependency bases. + +Only imported by `litellm.harness.handlers.deepagents_handler.load_deps()`, so `deepagents`, +`langchain` and `langchain-core` are never imported unless Harness.DEEPAGENTS is used. + +`SandboxBackend` implements deepagents' `SandboxBackendProtocol` on top of a litellm +`Sandbox`. The agent sees virtual paths rooted at the sandbox workdir (`/src/a.py` is +`/src/a.py`); file bytes move through `Sandbox.read/write`, and ls/glob/grep/ +delete/execute run plain POSIX commands through `Sandbox.run`, so the same code serves the +local and docker sandboxes (no python3 needed inside the sandbox) and inherits the sandbox's +env scrubbing and path confinement. +""" + +from __future__ import annotations + +import asyncio +import base64 +import itertools +import posixpath +import re +import shlex +import uuid +from collections.abc import Awaitable, Callable, Coroutine, Iterator, Mapping, Sequence +from types import MappingProxyType +from typing import Any, Final, TypeVar + +from deepagents.backends.protocol import ( + DeleteResult, + EditResult, + ExecuteResponse, + FileData, + FileDownloadResponse, + FileInfo, + FileUploadResponse, + GlobResult, + GrepMatch, + GrepResult, + LsResult, + ReadResult, + SandboxBackendProtocol, + WriteResult, +) +from deepagents.backends.utils import ( + InvalidGlobPatternError, + compile_grep_include_glob, + perform_string_replacement, + slice_read_response, +) +from langchain.agents.middleware import AgentMiddleware, ModelRequest, ModelResponse, ToolCallRequest +from langchain.agents.middleware.types import ModelCallResult +from langchain_core.callbacks import AsyncCallbackHandler +from langchain_core.messages import ToolMessage +from langchain_core.outputs import LLMResult +from langchain_core.tools import BaseTool +from langgraph.types import Command + +import litellm +from litellm._logging import verbose_logger +from litellm.constants import HARNESS_SNAPSHOT_SKIP_DIRS +from litellm.harness.context import SessionContext +from litellm.harness.errors import SandboxError +from litellm.harness.sandbox.base import CompletedRun, Sandbox + +T = TypeVar("T") + +# Module alias so ruff recognises `.exception()` as logging the swallowed error (BLE001). +_logger = verbose_logger +# Models cost_per_token could not price: logged once, then skipped (always 0.0). +_UNPRICED_MODELS: set[str] = set() # mutable-ok: process-wide log-once memo, grown as unpriced models are seen + +DEEPAGENTS_EXECUTE_TIMEOUT_SECONDS: Final = 120.0 +DEEPAGENTS_FS_TIMEOUT_SECONDS: Final = 60.0 +DEEPAGENTS_MAX_OUTPUT_BYTES: Final = 100_000 +_EXIT_NOT_FOUND: Final = 3 +_EXIT_NOT_DIR: Final = 4 +_EXIT_TIMEOUT: Final = 124 +_READ_ONLY_ERROR: Final = "Error: this session is read-only; files cannot be changed" +_NO_EXECUTE_ERROR: Final = "Error: shell execution is disabled for this session" +# $1 = directory. Prints "d/" or "f/" per entry ("/" never appears in a name). +_LS_SCRIPT: Final = ( + '[ -e "$1" ] || exit 3; [ -d "$1" ] || exit 4; cd "$1" || exit 5; ' + 'for f in * .[!.]* ..?*; do if [ -e "$f" ] || [ -L "$f" ]; then ' + 'if [ -d "$f" ]; then printf "d/%s\\n" "$f"; else printf "f/%s\\n" "$f"; fi; fi; done' +) +_DELETE_SCRIPT: Final = '[ -e "$1" ] || [ -L "$1" ] || exit 3; rm -rf -- "$1"' +# $1 = path. Prints the symlink-resolved absolute path of its nearest existing ancestor (the +# path itself when it exists). New files and new directories (`a/b/new.py`) resolve through +# whatever part already exists, so a symlinked ancestor is still caught. +_REALPATH_SCRIPT: Final = ( + 'p="$1"; while [ ! -e "$p" ] && [ ! -L "$p" ]; do q=$(dirname -- "$p"); ' + '[ "$q" = "$p" ] && exit 3; p="$q"; done; realpath -- "$p"' +) +_GREP_LINE: Final = re.compile(r"^(.+?):(\d+):(.*)$") +_FILTER_MIDDLEWARE_NAME: Final = "LiteLLMHarnessToolFilter" + + +def _decode(data: bytes) -> str: + return data.decode("utf-8", errors="replace") + + +def _ls_entries(base: str, stdout: str) -> Iterator[FileInfo]: + for line in stdout.splitlines(): + kind, _, name = line.partition("/") + if name: + is_dir = kind == "d" + yield FileInfo(path=f"{base}/{name}" + ("/" if is_dir else ""), is_dir=is_dir) + + +class SandboxBackend(SandboxBackendProtocol): # pyright: ignore[reportUntypedBaseClass] # deepagents/langchain are optional and not installed for type checking + """deepagents backend whose files and shell live in a litellm Sandbox.""" + + def __init__( + self, + sandbox: Sandbox, + *, + loop: asyncio.AbstractEventLoop, + writable: bool = True, + allow_execute: bool = True, + ) -> None: + self._sandbox = sandbox + self._loop = loop + self._root = posixpath.normpath(sandbox.workdir) + self._real_root: str | None = None + self._writable = writable + self._allow_execute = allow_execute + self._id = f"litellm-harness-{uuid.uuid4().hex[:8]}" + + @property + def id(self) -> str: + return self._id + + # -- paths -------------------------------------------------------------- + + def to_real(self, path: str) -> str: + """Sandbox path for a virtual path (or an absolute path already under workdir).""" + normalized = posixpath.normpath("/" + path.lstrip("/")) + if ".." in normalized.split("/"): + raise ValueError(f"path traversal not allowed: {path}") + if normalized == self._root or normalized.startswith(self._root + "/"): + return normalized + if normalized == "/": + return self._root + return self._root + normalized + + async def to_confined(self, path: str) -> str: + """to_real, then resolve symlinks inside the sandbox and refuse anything outside workdir. + + A repo can contain `link -> ~/.aws/credentials`; without this, read/grep/glob would + follow it and read host secrets even in read-only mode. + """ + real = self.to_real(path) + done = await self._run(("sh", "-c", _REALPATH_SCRIPT, "sh", real)) + resolved = done.stdout.strip() + if done.exit_code != 0 or not resolved: + raise ValueError(f"path not found: {path}") + root = await self._resolved_root() + if resolved != root and not resolved.startswith(root + "/"): + raise ValueError(f"path resolves outside the workspace: {path}") + return real + + async def _resolved_root(self) -> str: + if self._real_root is None: + done = await self._run(("realpath", "--", self._root)) + self._real_root = done.stdout.strip() if done.exit_code == 0 and done.stdout.strip() else self._root + return self._real_root + + def to_virtual(self, real: str) -> str: + if real == self._root: + return "/" + if real.startswith(self._root + "/"): + return real[len(self._root) :] + return real + + # -- sync bridge (deepagents only calls these outside the event loop) --- + + def _sync(self, coro: Coroutine[Any, Any, T]) -> T: + try: + running = asyncio.get_running_loop() + except RuntimeError: + running = None + if running is self._loop: + coro.close() + raise RuntimeError("SandboxBackend sync methods cannot run on the event loop thread") + return asyncio.run_coroutine_threadsafe(coro, self._loop).result() + + async def _run(self, cmd: Sequence[str], timeout: float | None = DEEPAGENTS_FS_TIMEOUT_SECONDS) -> CompletedRun: + return await self._sandbox.run(cmd, timeout=timeout) + + # -- ls ----------------------------------------------------------------- + + async def als(self, path: str) -> LsResult: + try: + real = await self.to_confined(path) + done = await self._run(("sh", "-c", _LS_SCRIPT, "sh", real)) + except (ValueError, SandboxError) as e: + return LsResult(error=f"Path '{path}': {e}") + if done.exit_code == _EXIT_NOT_FOUND: + return LsResult(error=f"Path '{path}': path_not_found") + if done.exit_code == _EXIT_NOT_DIR: + return LsResult(error=f"Path '{path}': not_a_directory") + if done.exit_code != 0: + return LsResult(error=f"Path '{path}': {done.stderr.strip() or 'ls failed'}") + base = self.to_virtual(real).rstrip("/") + entries = sorted(_ls_entries(base, done.stdout), key=lambda e: e["path"]) + return LsResult(entries=entries) + + def ls(self, path: str) -> LsResult: + return self._sync(self.als(path)) + + # -- read / write / edit ------------------------------------------------ + + async def _read_bytes(self, path: str) -> bytes: + return await self._sandbox.read(await self.to_confined(path)) + + async def aread(self, file_path: str, offset: int = 0, limit: int = 2000) -> ReadResult: + try: + data = await self._read_bytes(file_path) + except ValueError as e: + return ReadResult(error=f"Error reading file '{file_path}': {e}") + except SandboxError: + return ReadResult(error=f"File '{file_path}' not found") + try: + text = data.decode("utf-8") + except UnicodeDecodeError: + encoded = base64.standard_b64encode(data).decode("ascii") + return ReadResult(file_data=FileData(content=encoded, encoding="base64")) + return slice_read_response(FileData(content=text, encoding="utf-8"), offset, limit) + + def read(self, file_path: str, offset: int = 0, limit: int = 2000) -> ReadResult: + return self._sync(self.aread(file_path, offset, limit)) + + async def awrite(self, file_path: str, content: str) -> WriteResult: + if not self._writable: + return WriteResult(error=_READ_ONLY_ERROR) + try: + await self._sandbox.write(await self.to_confined(file_path), content.encode("utf-8")) + except (ValueError, SandboxError) as e: + return WriteResult(error=f"Error writing file '{file_path}': {e}") + return WriteResult(path=file_path) + + def write(self, file_path: str, content: str) -> WriteResult: + return self._sync(self.awrite(file_path, content)) + + async def aedit( + self, + file_path: str, + old_string: str, + new_string: str, + replace_all: bool = False, + ) -> EditResult: + if not self._writable: + return EditResult(error=_READ_ONLY_ERROR) + try: + content = _decode(await self._read_bytes(file_path)) + except ValueError as e: + return EditResult(error=f"Error editing file '{file_path}': {e}") + except SandboxError: + return EditResult(error=f"Error: File '{file_path}' not found") + old = old_string.replace("\r\n", "\n") + new = new_string.replace("\r\n", "\n") + replaced = perform_string_replacement(content.replace("\r\n", "\n"), old, new, replace_all) + if isinstance(replaced, str): + return EditResult(error=replaced) + new_content, occurrences = replaced + try: + await self._sandbox.write(await self.to_confined(file_path), new_content.encode("utf-8")) + except SandboxError as e: + return EditResult(error=f"Error editing file '{file_path}': {e}") + return EditResult(path=file_path, occurrences=int(occurrences)) + + def edit( + self, + file_path: str, + old_string: str, + new_string: str, + replace_all: bool = False, + ) -> EditResult: + return self._sync(self.aedit(file_path, old_string, new_string, replace_all)) + + async def adelete(self, file_path: str) -> DeleteResult: + if not self._writable: + return DeleteResult(error=_READ_ONLY_ERROR) + try: + real = await self.to_confined(file_path) + if real == self._root: + return DeleteResult(error="Error: refusing to delete the workspace root") + done = await self._run(("sh", "-c", _DELETE_SCRIPT, "sh", real)) + except (ValueError, SandboxError) as e: + return DeleteResult(error=f"Error deleting '{file_path}': {e}") + if done.exit_code == _EXIT_NOT_FOUND: + return DeleteResult(error=f"Error: '{file_path}' not found") + if done.exit_code != 0: + return DeleteResult(error=f"Error deleting '{file_path}': {done.stderr.strip()}") + return DeleteResult(path=file_path) + + def delete(self, file_path: str) -> DeleteResult: + return self._sync(self.adelete(file_path)) + + # -- glob / grep -------------------------------------------------------- + + def _find_cmd(self, root: str) -> tuple[str, ...]: + prune = tuple( + itertools.chain.from_iterable( + ("-o", "-name", name) if index else ("-name", name) + for index, name in enumerate(sorted(HARNESS_SNAPSHOT_SKIP_DIRS)) + ) + ) + # -P: never follow symlinks, so a repo link to ~/.aws cannot pull host files in. + return ("find", "-P", root, "(", *prune, ")", "-prune", "-o", "-type", "f", "-print") + + def _grep_cmd(self, pattern: str, root: str) -> tuple[str, ...]: + # grep only the regular files `find -P -type f` lists: symlinks are never followed, + # whatever grep implementation (GNU -R vs BSD -r) the sandbox has. + find_cmd = " ".join(shlex.quote(part) for part in self._find_cmd(root)) + return ("sh", "-c", f'{find_cmd} | tr "\\n" "\\0" | xargs -0 grep -nHFI -e "$1" --', "sh", pattern) + + async def aglob(self, pattern: str, path: str | None = None) -> GlobResult: + try: + matcher = compile_grep_include_glob(pattern) + root = await self.to_confined(path or "/") + done = await self._run(self._find_cmd(root)) + except (InvalidGlobPatternError, ValueError, SandboxError) as e: + return GlobResult(error=str(e), matches=None) + if done.exit_code != 0 and not done.stdout: + return GlobResult(matches=[]) # mutable-ok: deepagents GlobResult.matches is typed list[FileInfo] + matches = sorted( + ( + FileInfo(path=self.to_virtual(real), is_dir=False) + for real in done.stdout.splitlines() + if matcher(posixpath.relpath(real, root)) + ), + key=lambda m: m["path"], + ) + return GlobResult(matches=matches, truncated=done.exit_code != 0) + + def _grep_matches(self, stdout: str, root: str, include: Callable[[str], bool] | None) -> Iterator[GrepMatch]: + for line in stdout.splitlines(): + parsed = _GREP_LINE.match(line) + if parsed is None: + continue + real = parsed.group(1) + if include is None or include(posixpath.relpath(real, root)): + yield GrepMatch(path=self.to_virtual(real), line=int(parsed.group(2)), text=parsed.group(3)) + + def glob(self, pattern: str, path: str | None = None) -> GlobResult: + return self._sync(self.aglob(pattern, path)) + + async def agrep( + self, + pattern: str, + path: str | None = None, + glob: str | None = None, + *, + max_count: int | None = None, + ) -> GrepResult: + try: + include = compile_grep_include_glob(glob) if glob else None + root = await self.to_confined(path or "/") + done = await self._run(self._grep_cmd(pattern, root)) + except (InvalidGlobPatternError, ValueError, SandboxError) as e: + return GrepResult(error=f"Path '{path or '/'}': {e}") + if done.exit_code not in (0, 1) and not done.stdout: + return GrepResult(error=f"Path '{path or '/'}': {done.stderr.strip() or 'grep failed'}") + matches = list( # mutable-ok: GrepResult.matches is list[GrepMatch] + self._grep_matches(done.stdout, root, include) + ) + if max_count is not None and len(matches) > max_count: + return GrepResult(matches=matches[:max_count], truncated=True) + return GrepResult(matches=matches) + + def grep( + self, + pattern: str, + path: str | None = None, + glob: str | None = None, + *, + max_count: int | None = None, + ) -> GrepResult: + return self._sync(self.agrep(pattern, path, glob, max_count=max_count)) + + # -- upload / download -------------------------------------------------- + + async def _upload_one(self, path: str, data: bytes) -> FileUploadResponse: + if not self._writable: + return FileUploadResponse(path=path, error="permission_denied") + try: + await self._sandbox.write(await self.to_confined(path), data) + except ValueError: + return FileUploadResponse(path=path, error="invalid_path") + except SandboxError as e: + return FileUploadResponse(path=path, error=str(e)) + return FileUploadResponse(path=path) + + async def aupload_files( + self, + files: list[tuple[str, bytes]], # mutable-ok: signature fixed by deepagents BackendProtocol + ) -> list[FileUploadResponse]: # mutable-ok: return type fixed by deepagents BackendProtocol + return [ # mutable-ok: BackendProtocol returns a list + await self._upload_one(path, data) for path, data in files + ] + + def upload_files( + self, + files: list[tuple[str, bytes]], # mutable-ok: signature fixed by deepagents BackendProtocol + ) -> list[FileUploadResponse]: # mutable-ok: return type fixed by deepagents BackendProtocol + return self._sync(self.aupload_files(files)) + + async def _download_one(self, path: str) -> FileDownloadResponse: + try: + return FileDownloadResponse(path=path, content=await self._read_bytes(path)) + except ValueError: + return FileDownloadResponse(path=path, error="invalid_path") + except SandboxError: + return FileDownloadResponse(path=path, error="file_not_found") + + async def adownload_files( + self, + paths: list[str], # mutable-ok: signature fixed by deepagents BackendProtocol + ) -> list[FileDownloadResponse]: # mutable-ok: return type fixed by deepagents BackendProtocol + return [await self._download_one(path) for path in paths] # mutable-ok: BackendProtocol returns a list + + def download_files( + self, + paths: list[str], # mutable-ok: signature fixed by deepagents BackendProtocol + ) -> list[FileDownloadResponse]: # mutable-ok: return type fixed by deepagents BackendProtocol + return self._sync(self.adownload_files(paths)) + + # -- execute ------------------------------------------------------------ + + async def aexecute(self, command: str, *, timeout: int | None = None) -> ExecuteResponse: + if not self._allow_execute: + return ExecuteResponse(output=_NO_EXECUTE_ERROR, exit_code=1) + limit = float(timeout) if timeout else DEEPAGENTS_EXECUTE_TIMEOUT_SECONDS + try: + done = await self._run(("sh", "-c", command), timeout=limit) + except SandboxError as e: + return ExecuteResponse(output=f"Error: {e}", exit_code=_EXIT_TIMEOUT) + output = done.stdout + if done.stderr: + output = f"{output}\n{done.stderr}" if output else done.stderr + truncated = len(output.encode("utf-8")) > DEEPAGENTS_MAX_OUTPUT_BYTES + if truncated: + output = output.encode("utf-8")[:DEEPAGENTS_MAX_OUTPUT_BYTES].decode("utf-8", errors="ignore") + return ExecuteResponse(output=output, exit_code=done.exit_code, truncated=truncated) + + def execute(self, command: str, *, timeout: int | None = None) -> ExecuteResponse: + return self._sync(self.aexecute(command, timeout=timeout)) + + +def _tool_name(tool: BaseTool | Mapping[str, object]) -> str | None: + name = tool.get("name") if isinstance(tool, Mapping) else tool.name + return name if isinstance(name, str) else None + + +def _blocked_message(request: ToolCallRequest, blocked: frozenset[str]) -> ToolMessage | None: + name = request.tool_call["name"] + if name not in blocked: + return None + return ToolMessage( + content=f"Error: {name} is disabled for this session.", + tool_call_id=request.tool_call["id"] or "", + name=name, + status="error", + ) + + +class ToolFilterMiddleware(AgentMiddleware): # pyright: ignore[reportUntypedBaseClass] # deepagents/langchain are optional and not installed for type checking + """Hide tools from the model and refuse calls to them (disable_tools / permissions).""" + + def __init__(self, blocked: frozenset[str]) -> None: + super().__init__() + self._blocked = blocked + + @property + def name(self) -> str: + return _FILTER_MIDDLEWARE_NAME + + def _filtered(self, request: ModelRequest) -> ModelRequest: + return request.override( + tools=[ # mutable-ok: ModelRequest.tools is a list + t for t in request.tools if _tool_name(t) not in self._blocked + ] + ) + + def wrap_model_call( + self, request: ModelRequest, handler: Callable[[ModelRequest], ModelResponse] + ) -> ModelCallResult: + return handler(self._filtered(request)) + + async def awrap_model_call( + self, request: ModelRequest, handler: Callable[[ModelRequest], Awaitable[ModelResponse]] + ) -> ModelCallResult: + return await handler(self._filtered(request)) + + def wrap_tool_call( + self, request: ToolCallRequest, handler: Callable[[ToolCallRequest], ToolMessage | Command] + ) -> ToolMessage | Command: + return _blocked_message(request, self._blocked) or handler(request) + + async def awrap_tool_call( + self, request: ToolCallRequest, handler: Callable[[ToolCallRequest], Awaitable[ToolMessage | Command]] + ) -> ToolMessage | Command: + return _blocked_message(request, self._blocked) or await handler(request) + + +def _number(value: object) -> float | None: + if isinstance(value, bool) or not isinstance(value, (int, float)): + return None + return float(value) + + +def message_cost(message: object, cost_model: str | None, input_tokens: int, output_tokens: int) -> float: + """Cost of one model call: response_cost reported by litellm, else cost_per_token, else 0.""" + metadata = getattr(message, "response_metadata", None) or MappingProxyType({}) + reported = _number(metadata.get("response_cost")) + if reported is not None: + return reported + if not cost_model or cost_model in _UNPRICED_MODELS: + return 0.0 + try: + prompt_cost, completion_cost = litellm.cost_per_token( + model=cost_model, + prompt_tokens=input_tokens, + completion_tokens=output_tokens, + ) + return float(prompt_cost) + float(completion_cost) + except Exception: # litellm raises plain Exception for unmapped models + _UNPRICED_MODELS.add(cost_model) + _logger.exception("harness deepagents: no cost for %s; counting its calls as 0", cost_model) + return 0.0 + + +def record_llm_usage(ctx: SessionContext, cost_model: str | None, response: LLMResult) -> None: + """Add one model call's tokens and cost to the session counters. Never raises.""" + try: + ctx.calls += 1 + for generations in response.generations: + for generation in generations: + message = getattr(generation, "message", None) + usage = getattr(message, "usage_metadata", None) or MappingProxyType({}) + input_tokens = int(usage.get("input_tokens") or 0) + output_tokens = int(usage.get("output_tokens") or 0) + ctx.input_tokens += input_tokens + ctx.output_tokens += output_tokens + ctx.cost += message_cost(message, cost_model, input_tokens, output_tokens) + except Exception: + # Usage accounting must never fail a turn; log with traceback and move on. + _logger.exception("harness deepagents: usage accounting failed") + + +class UsageCallback(AsyncCallbackHandler): # pyright: ignore[reportUntypedBaseClass] # deepagents/langchain are optional and not installed for type checking + """Counts every model call in the graph, subagents and summarization included.""" + + def __init__(self, ctx: SessionContext, cost_model: str | None) -> None: + self._ctx = ctx + self._cost_model = cost_model + + async def on_llm_end(self, response: LLMResult, **kwargs: object) -> None: + record_llm_usage(self._ctx, self._cost_model, response) diff --git a/litellm/llms/deepagents/harness/transformation.py b/litellm/llms/deepagents/harness/transformation.py new file mode 100644 index 00000000000..339a6274d33 --- /dev/null +++ b/litellm/llms/deepagents/harness/transformation.py @@ -0,0 +1,313 @@ +""" +Deep Agents harness config: LangChain `deepagents` running in your Python process. + +Pure translation only: model kwargs (gateway mode uses `litellm_proxy/` with the same +attribution headers the CLI endpoint adds), permission and tool filtering, and LangGraph +stream chunks to events. `litellm/harness/handlers/deepagents_handler.py` builds the agent, +streams it, answers approvals and counts usage. +""" + +from __future__ import annotations + +import itertools +import json +from collections.abc import Iterator, Mapping, Sequence +from dataclasses import dataclass +from types import MappingProxyType +from typing import TYPE_CHECKING, Any, Final + +from litellm.harness.options import DeepAgentsOptions +from litellm.harness.types import ( + Capabilities, + Event, + Harness, + Reasoning, + Text, + ToolCall, + ToolResult, +) +from litellm.llms.base_llm.harness.transformation import BaseHarnessConfig + +if TYPE_CHECKING: + from litellm.harness.context import SessionContext + +INSTALL_HINT: Final = "Deep Agents is not installed. Run: pip install deepagents langchain-litellm" +SKILLS_DIR: Final = ".deepagents/skills" +# Graph supersteps per agent turn (model node, tools node, middleware hooks) for recursion_limit. +DEEPAGENTS_STEPS_PER_TURN: Final = 6 +DEEPAGENTS_BASE_RECURSION_LIMIT: Final = 25 +DEEPAGENTS_DEFAULT_RECURSION_LIMIT: Final = 1000 +_MODEL_NODE: Final = "model" +# Only these graph nodes produce new messages; middleware hooks may re-emit history. +_EVENT_NODES: Final = frozenset({"model", "tools"}) + +NORMALIZED_TO_NATIVE: Final[Mapping[str, str]] = MappingProxyType( + { + "read": "read_file", + "write": "write_file", + "edit": "edit_file", + "bash": "execute", + } +) +NATIVE_TO_NORMALIZED: Final[Mapping[str, str]] = MappingProxyType({v: k for k, v in NORMALIZED_TO_NATIVE.items()}) +BUILTIN_TOOLS: Final = frozenset( + { + "ls", + "read_file", + "write_file", + "edit_file", + "delete", + "glob", + "grep", + "execute", + "write_todos", + "task", + } +) +WRITE_TOOLS: Final = frozenset({"write_file", "edit_file", "delete"}) +EXECUTE_TOOLS: Final = frozenset({"execute"}) +APPROVAL_TOOLS: Final = WRITE_TOOLS | EXECUTE_TOOLS +_APPROVAL_DECISIONS: Final = ("approve", "reject") + + +# --------------------------------------------------------------------------- +# Pure helpers (unit tested directly; kept module-level so they port cleanly) +# --------------------------------------------------------------------------- + + +def gateway_headers( + ctx: SessionContext, +) -> dict[str, str]: # mutable-ok: ChatLiteLLM.extra_headers is a pydantic dict field + """Same attribution headers the session endpoint adds for CLI harnesses.""" + metadata = ctx.metadata + metadata_json = json.dumps(dict(metadata), default=str) if metadata else None # mutable-ok: for json.dumps + metadata_header = (("x-litellm-spend-logs-metadata", metadata_json),) if metadata_json is not None else () + return dict( # mutable-ok: ChatLiteLLM.extra_headers is a pydantic dict field + (("x-litellm-tags", f"harness,{ctx.harness.value}"), *metadata_header) + ) + + +def chat_model_kwargs( + ctx: SessionContext, +) -> dict[str, Any]: # mutable-ok: ChatLiteLLM constructor kwargs, splatted as **kwargs + """ChatLiteLLM constructor kwargs for gateway or SDK mode.""" + if not ctx.model: + raise ValueError("Harness.DEEPAGENTS needs model=") + if ctx.gateway is not None: + return { # mutable-ok: ChatLiteLLM constructor kwargs, splatted as **kwargs + "model": f"litellm_proxy/{ctx.model}", + "api_base": ctx.gateway.api_base, + "api_key": ctx.gateway.api_key, + "extra_headers": gateway_headers(ctx), + } + return {"model": ctx.model, "api_key": ctx.api_key, "api_base": ctx.api_base} # mutable-ok: ChatLiteLLM kwargs + + +def native_tool_name(name: str) -> str: + return NORMALIZED_TO_NATIVE.get(name, name) + + +def normalized_tool_name(native: str) -> str: + return NATIVE_TO_NORMALIZED.get(native, native) + + +def blocked_tools(permissions: str, disable_tools: Sequence[str]) -> frozenset[str]: + """Native tool names the model must not see or call.""" + disabled = frozenset(native_tool_name(name) for name in disable_tools) + if permissions == "read-only": + return disabled | WRITE_TOOLS | EXECUTE_TOOLS + if permissions == "edit": + return disabled | EXECUTE_TOOLS + return disabled + + +def interrupt_config( + permissions: str, blocked: frozenset[str] +) -> dict[str, Any] | None: # mutable-ok: deepagents create_deep_agent(interrupt_on=) takes a dict + """interrupt_on for permissions='ask': approve/reject every mutating built-in.""" + if permissions != "ask": + return None + return { # mutable-ok: deepagents interrupt_on config (dict of InterruptOnConfig with list allowed_decisions) + name: {"allowed_decisions": list(_APPROVAL_DECISIONS)} # mutable-ok: deepagents InterruptOnConfig shape + for name in sorted(APPROVAL_TOOLS - blocked) + } + + +def recursion_limit(ctx: SessionContext) -> int: + options = ctx.options if isinstance(ctx.options, DeepAgentsOptions) else None + if options is not None and options.recursion_limit is not None: + return options.recursion_limit + if ctx.max_turns is not None: + return DEEPAGENTS_BASE_RECURSION_LIMIT + ctx.max_turns * DEEPAGENTS_STEPS_PER_TURN + return DEEPAGENTS_DEFAULT_RECURSION_LIMIT + + +def content_text(content: object) -> str: + """Plain text of a LangChain message content (str or content blocks).""" + if isinstance(content, str): + return content + if not isinstance(content, list): + return "" + return "".join( + block if isinstance(block, str) else block.get("text", "") for block in content if _is_text_block(block) + ) + + +def _is_text_block(block: object) -> bool: + return isinstance(block, str) or (isinstance(block, dict) and block.get("type") == "text") + + +def reasoning_text(message: object) -> str: + """Reasoning deltas from additional_kwargs or reasoning/thinking content blocks.""" + extra = getattr(message, "additional_kwargs", None) or MappingProxyType({}) + reasoning = extra.get("reasoning_content") + if isinstance(reasoning, str) and reasoning: + return reasoning + content = getattr(message, "content", None) + if not isinstance(content, list): + return "" + return "".join( + str(block.get("reasoning") or block.get("thinking") or "") + for block in content + if isinstance(block, dict) and block.get("type") in ("reasoning", "thinking") + ) + + +def stream_events( + message: object, +) -> list[Event]: # mutable-ok: returns a list; existing callers/tests compare it to list literals + """Text / Reasoning deltas for one streamed message chunk.""" + if getattr(message, "type", None) not in ("AIMessageChunk", "ai"): + return [] # mutable-ok: returns a list; existing callers/tests compare it to list literals + reasoning = reasoning_text(message) + text = content_text(getattr(message, "content", "")) + reasoning_events: tuple[Event, ...] = (Reasoning(delta=reasoning),) if reasoning else () + text_events: tuple[Event, ...] = (Text(delta=text),) if text else () + return [ # mutable-ok: returns a list; existing callers/tests compare it to list literals + *reasoning_events, + *text_events, + ] + + +def tool_call_event(call: Mapping[str, Any]) -> ToolCall: + native = str(call.get("name") or "") + args = call.get("args") + return ToolCall( + id=str(call.get("id") or ""), + name=normalized_tool_name(native), + native_name=native, + input=dict(args) if isinstance(args, Mapping) else {"args": args}, # mutable-ok: ToolCall.input is a dict + builtin=native in BUILTIN_TOOLS, + ) + + +def _node_messages(update: Mapping[Any, Any]) -> Iterator[object]: + for node, delta in update.items(): + if node not in _EVENT_NODES or not isinstance(delta, Mapping): + continue + messages = delta.get("messages") + if isinstance(messages, list): + yield from messages + + +def update_events( + update: object, skip_tools: frozenset[str] +) -> list[Event]: # mutable-ok: returns a list; existing callers/tests compare it to list literals + """ToolCall / ToolResult events from one `updates` stream chunk (node -> state delta).""" + if not isinstance(update, Mapping): + return [] # mutable-ok: returns a list; existing callers/tests compare it to list literals + return list( # mutable-ok: returns a list; existing callers/tests compare it to list literals + itertools.chain.from_iterable(_message_events(message, skip_tools) for message in _node_messages(update)) + ) + + +def _message_events(message: object, skip_tools: frozenset[str]) -> tuple[Event, ...]: + kind = getattr(message, "type", None) + if kind == "ai": + calls = getattr(message, "tool_calls", None) or () + return tuple(tool_call_event(call) for call in calls if call.get("name") not in skip_tools) + if kind == "tool" and getattr(message, "name", None) not in skip_tools: + return ( + ToolResult( + id=str(getattr(message, "tool_call_id", "") or ""), + output=content_text(getattr(message, "content", "")), + is_error=getattr(message, "status", None) == "error", + ), + ) + return () + + +def interrupts_in( + update: object, +) -> list[Any]: # mutable-ok: returns a list; existing callers/tests compare it to list literals + if not isinstance(update, Mapping): + return [] # mutable-ok: returns a list; existing callers/tests compare it to list literals + found = update.get("__interrupt__") + items = tuple(found) if isinstance(found, (list, tuple)) else () + return list(items) # mutable-ok: list return; callers/tests compare to lists + + +def final_ai_text(messages: Sequence[Any]) -> str: + for message in reversed(messages): + if getattr(message, "type", None) == "ai": + text = content_text(getattr(message, "content", "")) + if text: + return text + return "" + + +def structured_json(value: object) -> str | None: + if value is None: + return None + dump = getattr(value, "model_dump_json", None) + if callable(dump): + return str(dump()) + return json.dumps(value, default=str) + + +def approval_requests( + interrupt_value: object, +) -> list[Mapping[str, Any]]: # mutable-ok: returns a list; existing callers/tests compare it to list literals + """action_requests of a HumanInTheLoopMiddleware interrupt payload.""" + if not isinstance(interrupt_value, Mapping): + return [] # mutable-ok: returns a list; existing callers/tests compare it to list literals + requests = interrupt_value.get("action_requests") + kept = tuple(r for r in requests if isinstance(r, Mapping)) if isinstance(requests, list) else () + return list(kept) # mutable-ok: list return; callers/tests compare to lists + + +def decision(allowed: bool, reason: str) -> dict[str, Any]: # mutable-ok: LangGraph resume payload (HITL decision dict) + if allowed: + return {"type": "approve"} # mutable-ok: LangGraph resume payload (HITL decision dict) + return { # mutable-ok: LangGraph HITL decision + "type": "reject", + "message": reason or "The user denied this tool call.", + } + + +@dataclass +class TurnState: + """Mutable state across the stream passes of one turn.""" + + interrupts: tuple[Any, ...] = () + + +class DeepAgentsHarnessConfig(BaseHarnessConfig): + harness = Harness.DEEPAGENTS + options_type = DeepAgentsOptions + uses_model_endpoint = False + capabilities = Capabilities( + structured_output=True, + tool_approval=True, + tool_filtering=True, + history=True, + custom_tools=True, + skills=True, + resume=True, + permission_modes=frozenset({"read-only", "ask", "edit", "full"}), + ) + + def validate_environment(self, ctx: SessionContext) -> None: + self.get_options(ctx) + if not ctx.model: + raise ValueError("Harness.DEEPAGENTS needs model=") diff --git a/litellm/llms/opencode/__init__.py b/litellm/llms/opencode/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/opencode/harness/__init__.py b/litellm/llms/opencode/harness/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/opencode/harness/transformation.py b/litellm/llms/opencode/harness/transformation.py new file mode 100644 index 00000000000..af5fa1ae71d --- /dev/null +++ b/litellm/llms/opencode/harness/transformation.py @@ -0,0 +1,427 @@ +""" +OpenCode harness config: `opencode run --format json`, once per turn. + +Every model call goes to one custom provider (`litellm`, `@ai-sdk/openai-compatible`, +bundled in the binary) whose baseURL is the per-session endpoint. The config travels in +OPENCODE_CONFIG_CONTENT, which opencode applies after global and project config, so a +repo's own opencode.json cannot redirect model calls. The token is never in argv or env: +the config references it with `{file:/token}`. XDG dirs point at a persisted +LiteLLM-owned root so the user's opencode config and auth are never read, and the session +DB outlives a session for resume. Verified against opencode 1.14.41. +""" + +from __future__ import annotations + +import itertools +import json +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from types import MappingProxyType +from typing import TYPE_CHECKING, Any, Final + +from litellm.harness.errors import CapabilityUnsupported, HarnessError, OptionsMismatch +from litellm.harness.options import OpenCodeOptions +from litellm.harness.types import ( + Capabilities, + Event, + Harness, + PermissionMode, + Reasoning, + Text, + ToolCall, + ToolResult, +) +from litellm.llms.base_llm.harness.transformation import ( + BaseCLIHarnessConfig, + HarnessSessionSetup, + HarnessTurnError, + HarnessTurnRequest, + HarnessTurnResponse, + event_list, +) +from litellm.llms.base_llm.harness.utils import ( + last_json_object, + native_tool_names, + normalize_tool_name, + stderr_tail_text, + structured_output_instruction, +) + +if TYPE_CHECKING: + from litellm.harness.context import SessionContext + +OPENCODE_BINARY: Final = "opencode" +OPENCODE_PROVIDER_ID: Final = "litellm" +OPENCODE_PROVIDER_NPM: Final = "@ai-sdk/openai-compatible" +# A fixed title skips opencode's extra title-generation model call on the first turn. +OPENCODE_SESSION_TITLE: Final = "litellm-harness" +TOKEN_FILENAME: Final = "token" +INSTRUCTIONS_FILENAME: Final = "instructions.md" +XDG_DIRNAME: Final = "xdg" +XDG_SUBDIRS: Final = ("config", "data", "state", "cache") + +# Env that keeps opencode off the network (except the endpoint) and away from ~/.claude. +OPENCODE_ISOLATION_ENV: Final[Mapping[str, str]] = MappingProxyType( + { + "OPENCODE_DISABLE_AUTOUPDATE": "1", + "OPENCODE_DISABLE_MODELS_FETCH": "1", + "OPENCODE_DISABLE_LSP_DOWNLOAD": "1", + "OPENCODE_DISABLE_SHARE": "1", + "OPENCODE_DISABLE_DEFAULT_PLUGINS": "1", + "OPENCODE_DISABLE_CLAUDE_CODE": "1", + "OPENCODE_DISABLE_EXTERNAL_SKILLS": "1", + # Blank (falsy to opencode) so an inherited value can't add config, auth or rules. + "OPENCODE_CONFIG": "", + "OPENCODE_CONFIG_DIR": "", + "OPENCODE_PERMISSION": "", + "OPENCODE_AUTH_CONTENT": "", + } +) + +MANAGED_CONFIG_KEYS: Final = frozenset( + { + "provider", + "model", + "small_model", + "permission", + "tools", + "enabled_providers", + "disabled_providers", + # plugins run arbitrary code as the host user; runs always use --pure + "plugin", + } +) +AGENT_MANAGED_KEYS: Final = frozenset({"permission", "tools", "model"}) + +# Later keys win in opencode, so disable_tools denies go last. `opencode run` auto-rejects +# anything left at "ask", so no mode leaves a tool on ask. +PERMISSION_RULES: Final[Mapping[str, Mapping[str, str]]] = MappingProxyType( + { + "read-only": MappingProxyType({"edit": "deny", "bash": "deny", "webfetch": "deny"}), + "edit": MappingProxyType({"edit": "allow", "bash": "deny", "webfetch": "allow"}), + "full": MappingProxyType({"*": "allow"}), + } +) + +# opencode gates write, edit and apply_patch with the single `edit` permission. +NORMALIZED_TO_NATIVE: Final[Mapping[str, tuple[str, ...]]] = MappingProxyType( + { + "read": ("read",), + "write": ("edit",), + "edit": ("edit",), + "bash": ("bash",), + "glob": ("glob",), + "grep": ("grep",), + "ls": ("list",), + "web_search": ("webfetch", "websearch"), + } +) + +NATIVE_TO_NORMALIZED: Final[Mapping[str, str]] = MappingProxyType( + { + "read": "read", + "write": "write", + "edit": "edit", + "multiedit": "edit", + "patch": "edit", + "apply_patch": "edit", + "bash": "bash", + "glob": "glob", + "grep": "grep", + "list": "ls", + "webfetch": "web_search", + "websearch": "web_search", + } +) + +OPENCODE_BUILTIN_TOOLS: Final = frozenset( + { + *NATIVE_TO_NORMALIZED, + "task", + "todowrite", + "todoread", + "skill", + "invalid", + "question", + "lsp", + "codesearch", + "plan_enter", + "plan_exit", + } +) + + +@dataclass +class OpenCodeStreamState: + """What the parser has learned from one `opencode run`.""" + + session_id: str | None = None + final_text: str = "" + error: str | None = None + step_texts: Sequence[str] = () + + +def _as_dict(value: object) -> Mapping[str, Any]: + return value if isinstance(value, dict) else MappingProxyType({}) + + +def _tool_events(part: Mapping[str, Any]) -> Sequence[Event]: + native = str(part.get("tool") or "") + call_id = str(part.get("callID") or part.get("id") or "") + state = _as_dict(part.get("state")) + tool_input = state.get("input") + call = ToolCall( + id=call_id, + name=normalize_tool_name(native, NATIVE_TO_NORMALIZED), + native_name=native, + input=tool_input if isinstance(tool_input, dict) else MappingProxyType({"input": tool_input}), + builtin=native in OPENCODE_BUILTIN_TOOLS, + ) + if state.get("status") == "error": + message = str(state.get("error") or state.get("output") or "tool failed") + return event_list(call, ToolResult(id=call_id, output=message, is_error=True)) + output = state.get("output") + text = output if isinstance(output, str) else json.dumps(output) + # The `invalid` pseudo-tool is how opencode reports a call to an unavailable tool. + return event_list(call, ToolResult(id=call_id, output=text, is_error=native == "invalid")) + + +def _error_message(error: object) -> str: + if not isinstance(error, dict): + return str(error or "opencode reported an error") + data = error.get("data") + if isinstance(data, dict) and data.get("message"): + return str(data["message"]) + return str(error.get("name") or "opencode reported an error") + + +def validate_user_config(config: Mapping[str, Any]) -> None: + """Reject OpenCodeOptions.config keys LiteLLM manages (or that bypass permissions).""" + for key in config: + if key in MANAGED_CONFIG_KEYS: + raise OptionsMismatch( + f"OpenCodeOptions.config[{key!r}] is managed by LiteLLM; use the matching " + "agent() argument (model=, permissions=, disable_tools=) instead" + ) + for section in ("agent", "mode"): + entries = config.get(section) + if entries is None: + continue + if not isinstance(entries, Mapping): + raise OptionsMismatch(f"OpenCodeOptions.config[{section!r}] must be a mapping") + for name, agent in entries.items(): + managed = AGENT_MANAGED_KEYS & frozenset(agent or ()) + if managed: + raise OptionsMismatch( + f"OpenCodeOptions.config[{section!r}][{name!r}] sets {sorted(managed)}, " + "which LiteLLM manages; use permissions=/disable_tools=/model= instead" + ) + + +def permission_rules(permissions: PermissionMode, disable_tools: Sequence[str]) -> Mapping[str, str]: + """opencode `permission` config for a mode plus denies for disable_tools.""" + if permissions not in PERMISSION_RULES: + raise CapabilityUnsupported( + f"Harness.OPENCODE does not support permissions={permissions!r} (supported: {sorted(PERMISSION_RULES)})" + ) + denied: Final = native_tool_names(disable_tools, NORMALIZED_TO_NATIVE) + # Denies go last (later keys win in opencode), so drop them from the mode rules first. + kept: Final = ((key, value) for key, value in PERMISSION_RULES[permissions].items() if key not in denied) + rules: Final = itertools.chain(kept, ((native, "deny") for native in denied)) + return dict(rules) # mutable-ok: opencode config JSON + + +def build_opencode_config( + *, + model: str, + base_url: str, + token_path: str, + permissions: PermissionMode, + disable_tools: Sequence[str] = (), + user_config: Mapping[str, Any] | None = None, + instructions_path: str | None = None, + skills_path: str | None = None, +) -> Mapping[str, Any]: + """The full opencode config: user config underneath, LiteLLM-managed keys on top.""" + user: Final = user_config or MappingProxyType({}) + validate_user_config(user) + qualified = f"{OPENCODE_PROVIDER_ID}/{model}" + extra_instructions: Final = (instructions_path,) if instructions_path else () + instructions: Final = [*(user.get("instructions") or ()), *extra_instructions] # mutable-ok: opencode config JSON + user_skills = _as_dict(user.get("skills")) + extra_skills: Final = (skills_path,) if skills_path else () + skill_paths: Final = [*(user_skills.get("paths") or ()), *extra_skills] # mutable-ok: opencode config JSON + options: Final = {"baseURL": base_url, "apiKey": "{file:" + token_path + "}"} # mutable-ok: opencode config JSON + models: Final[dict[str, Any]] = {model: {}} # mutable-ok: opencode config JSON + provider: Final = { # mutable-ok: opencode config JSON + "npm": OPENCODE_PROVIDER_NPM, + "name": "LiteLLM", + "options": options, + "models": models, + } + managed: Final = { # mutable-ok: opencode config JSON + "provider": {OPENCODE_PROVIDER_ID: provider}, # mutable-ok: opencode config JSON + "enabled_providers": [OPENCODE_PROVIDER_ID], # mutable-ok: opencode config JSON + "model": qualified, + "small_model": qualified, + "permission": permission_rules(permissions, disable_tools), + "autoupdate": False, + "share": "disabled", + } + skills: Final = {**user_skills, "paths": skill_paths} # mutable-ok: opencode config JSON + optional: Final = (("instructions", instructions), ("skills", skills if skill_paths else None)) + present: Final = ((key, value) for key, value in optional if value) + return {**user, **managed, **dict(present)} # mutable-ok: opencode config JSON + + +def build_instructions(ctx: SessionContext) -> str | None: + schema_part: Final = ( + structured_output_instruction(ctx.output.model_json_schema()) if ctx.output is not None else None + ) + sections: Final = tuple(section for section in (ctx.instructions, schema_part) if section) + return "\n\n".join(sections) if sections else None + + +def turn_prompt(ctx: SessionContext, prompt: str) -> str: + """Repeat the schema instruction in the user turn; system instructions alone are too weak.""" + if ctx.output is None: + return prompt + return f"{prompt}\n\n{structured_output_instruction(ctx.output.model_json_schema())}" + + +class OpenCodeHarnessConfig(BaseCLIHarnessConfig): + harness = Harness.OPENCODE + options_type = OpenCodeOptions + capabilities = Capabilities( + structured_output=True, + tool_approval=False, + tool_filtering=True, + history=False, + custom_tools=False, + skills=True, + resume=True, + permission_modes=frozenset({"read-only", "edit", "full"}), + ) + + def get_binary(self) -> str: + return OPENCODE_BINARY + + def get_install_hint(self) -> str: + return "npm install -g opencode-ai (or brew install sst/tap/opencode)" + + def validate_environment(self, ctx: SessionContext) -> None: + options: OpenCodeOptions = self.get_options(ctx) + validate_user_config(options.config) + + def transform_session_setup(self, ctx: SessionContext, private_dir: str) -> HarnessSessionSetup: + if ctx.endpoint is None: + raise HarnessError("OpenCode needs the session model endpoint") + model = ctx.model or ctx.endpoint.model + if not model: + raise ValueError("Harness.OPENCODE needs model= (a gateway model group or litellm model)") + options: OpenCodeOptions = self.get_options(ctx) + instructions = build_instructions(ctx) + token: Final = ctx.endpoint.token.encode("utf-8") + files: Final = ( + MappingProxyType({TOKEN_FILENAME: token, INSTRUCTIONS_FILENAME: instructions.encode("utf-8")}) + if instructions is not None + else MappingProxyType({TOKEN_FILENAME: token}) + ) + config = build_opencode_config( + model=model, + base_url=ctx.sandbox.host_url(ctx.endpoint.port).rstrip("/") + "/v1", + token_path=f"{private_dir}/{TOKEN_FILENAME}", + permissions=ctx.permissions, + disable_tools=ctx.disable_tools, + user_config=options.config, + instructions_path=f"{private_dir}/{INSTRUCTIONS_FILENAME}" if instructions is not None else None, + skills_path=f"{private_dir}/skills" if ctx.skills else None, + ) + xdg: Final = MappingProxyType( + {f"XDG_{sub.upper()}_HOME": f"{private_dir}/{XDG_DIRNAME}/{sub}" for sub in XDG_SUBDIRS} + ) + return HarnessSessionSetup( + files=files, + persisted_dirs=((XDG_DIRNAME, "opencode"),), + skills_dir="skills", + env=MappingProxyType( + {**OPENCODE_ISOLATION_ENV, **options.env, **xdg, "OPENCODE_CONFIG_CONTENT": json.dumps(config)} + ), + ) + + def transform_turn_request( + self, + ctx: SessionContext, + setup: HarnessSessionSetup, + private_dir: str, + prompt: str, + native_session_id: str | None, + ) -> HarnessTurnRequest: + options: OpenCodeOptions = self.get_options(ctx) + model = ctx.model or (ctx.endpoint.model if ctx.endpoint else None) + # --pure: never load plugins. A repo's .opencode/plugin/*.js would otherwise run as the + # host user at startup, before any tool permission applies. + argv: Final = ( + OPENCODE_BINARY, + "run", + "--pure", + "--format", + "json", + "--thinking", + "-m", + f"{OPENCODE_PROVIDER_ID}/{model}", + *(("--agent", options.agent) if options.agent else ()), + *(("--session", native_session_id) if native_session_id else ("--title", OPENCODE_SESSION_TITLE)), + ) + # The prompt goes on stdin; opencode appends non-TTY stdin to the message. + return HarnessTurnRequest(argv=argv, env=setup.env, stdin=turn_prompt(ctx, prompt), cwd=ctx.sandbox.workdir) + + def create_stream_state(self) -> OpenCodeStreamState: + return OpenCodeStreamState() + + def transform_stream_line(self, line: Mapping[str, Any], state: OpenCodeStreamState) -> Sequence[Event]: + """step_finish token counts are ignored on purpose: the session endpoint accounts usage.""" + session_id = line.get("sessionID") + if session_id and state.session_id is None: + state.session_id = str(session_id) + event_type = line.get("type") + part = _as_dict(line.get("part")) + if event_type == "step_start": + state.step_texts = () + return event_list() + if event_type == "text": + text = str(part.get("text") or "") + if not text: + return event_list() + state.step_texts = (*state.step_texts, text) + state.final_text = "\n\n".join(state.step_texts) + return event_list(Text(delta=text)) + if event_type == "reasoning": + text = str(part.get("text") or "") + return event_list(Reasoning(delta=text)) if text else event_list() + if event_type == "tool_use": + return _tool_events(part) + if event_type == "error": + message = _error_message(line.get("error")) + state.error = f"{state.error}\n{message}" if state.error else message + return event_list() + + def get_native_session_id(self, state: OpenCodeStreamState) -> str | None: + return state.session_id + + def transform_turn_response( + self, + ctx: SessionContext, + state: OpenCodeStreamState, + exit_code: int, + stderr_tail: Sequence[str], + ) -> HarnessTurnResponse: + # opencode exits 0 after an `error` event, so check state first. + if state.error: + raise HarnessTurnError(f"opencode turn failed: {state.error}") + if exit_code != 0: + raise HarnessTurnError( + f"opencode exited with code {exit_code}: {stderr_tail_text(stderr_tail) or 'no output'}" + ) + output_json = last_json_object(state.final_text) if ctx.output is not None else None + return HarnessTurnResponse(final_text=state.final_text, output_json=output_json) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 40fdf083cf8..02e658ca779 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -5516,7 +5516,7 @@ "supports_web_search": false }, "azure/gpt-4.1-nano": { - "deprecation_date": "2026-10-14", + "deprecation_date": "2027-04-14", "cache_read_input_token_cost": 2.5e-08, "input_cost_per_token": 1e-07, "input_cost_per_token_batches": 5e-08, @@ -5550,7 +5550,7 @@ "supports_vision": true }, "azure/gpt-4.1-nano-2025-04-14": { - "deprecation_date": "2026-10-14", + "deprecation_date": "2027-04-14", "cache_read_input_token_cost": 2.5e-08, "input_cost_per_token": 1e-07, "input_cost_per_token_batches": 5e-08, @@ -6133,7 +6133,7 @@ "cache_creation_input_audio_token_cost": 4e-07, "cache_read_input_audio_token_cost": 4e-07, "cache_read_input_token_cost": 4e-07, - "deprecation_date": "2027-06-25", + "deprecation_date": "2027-07-31", "input_cost_per_audio_token": 3.2e-05, "input_cost_per_image_token": 5e-06, "input_cost_per_token": 4e-06, @@ -6168,7 +6168,7 @@ "cache_creation_input_audio_token_cost": 3e-07, "cache_read_input_audio_token_cost": 3e-07, "cache_read_input_token_cost": 6e-08, - "deprecation_date": "2027-06-25", + "deprecation_date": "2027-07-31", "input_cost_per_audio_token": 1e-05, "input_cost_per_image_token": 8e-07, "input_cost_per_token": 6e-07, @@ -6346,7 +6346,7 @@ "supports_tool_choice": true }, "azure/gpt-4o-transcribe": { - "deprecation_date": "2026-10-15", + "deprecation_date": "2026-12-31", "input_cost_per_audio_token": 2.5e-06, "input_cost_per_token": 2.5e-06, "litellm_provider": "azure", @@ -6354,6 +6354,7 @@ "max_output_tokens": 2000, "mode": "audio_transcription", "output_cost_per_token": 1e-05, + "source": "https://management.azure.com/subscriptions/c873328e-b572-4770-8dff-aaeb6f1f0e79/providers/Microsoft.CognitiveServices/locations/eastus2/models?api-version=2024-10-01", "supported_endpoints": [ "/v1/audio/transcriptions" ] @@ -8571,6 +8572,7 @@ "cache_creation_input_token_cost_above_272k_tokens": 5e-06, "cache_read_input_token_cost": 1e-07, "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "deprecation_date": "2028-03-11", "input_cost_per_token": 2e-06, "input_cost_per_token_above_272k_tokens": 4e-06, "litellm_provider": "azure", @@ -8619,6 +8621,7 @@ "cache_creation_input_token_cost_above_272k_tokens": 5e-06, "cache_read_input_token_cost": 1e-07, "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "deprecation_date": "2028-03-11", "input_cost_per_token": 2e-06, "input_cost_per_token_above_272k_tokens": 4e-06, "litellm_provider": "azure", @@ -10972,7 +10975,7 @@ "supports_web_search": false }, "azure/us/gpt-4.1-nano-2025-04-14": { - "deprecation_date": "2026-10-14", + "deprecation_date": "2027-04-14", "cache_read_input_token_cost": 2.8e-08, "input_cost_per_token": 1.1e-07, "input_cost_per_token_batches": 5.5e-08, @@ -11958,6 +11961,7 @@ "supports_vision": true }, "azure_ai/Meta-Llama-3-70B-Instruct": { + "deprecation_date": "2025-06-30", "input_cost_per_token": 1.1e-06, "litellm_provider": "azure_ai", "max_input_tokens": 8192, @@ -11965,9 +11969,11 @@ "max_tokens": 2048, "mode": "chat", "output_cost_per_token": 3.7e-07, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true }, "azure_ai/Meta-Llama-3.1-70B-Instruct": { + "deprecation_date": "2025-06-30", "input_cost_per_token": 2.68e-06, "litellm_provider": "azure_ai", "max_input_tokens": 128000, @@ -11975,10 +11981,11 @@ "max_tokens": 2048, "mode": "chat", "output_cost_per_token": 3.54e-06, - "source": "https://marketplace.microsoft.com/en-us/marketplace/apps/metagenai.meta-llama-3-1-70b-instruct-offer?tab=PlansAndPrice", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true }, "azure_ai/Phi-3-medium-128k-instruct": { + "deprecation_date": "2025-08-30", "input_cost_per_token": 1.7e-07, "litellm_provider": "azure_ai", "max_input_tokens": 128000, @@ -11986,11 +11993,12 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6.8e-07, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true, "supports_vision": false }, "azure_ai/Phi-3-medium-4k-instruct": { + "deprecation_date": "2025-08-30", "input_cost_per_token": 1.7e-07, "litellm_provider": "azure_ai", "max_input_tokens": 4096, @@ -11998,11 +12006,12 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6.8e-07, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true, "supports_vision": false }, "azure_ai/Phi-3-mini-128k-instruct": { + "deprecation_date": "2025-08-30", "input_cost_per_token": 1.3e-07, "litellm_provider": "azure_ai", "max_input_tokens": 128000, @@ -12010,11 +12019,12 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 5.2e-07, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true, "supports_vision": false }, "azure_ai/Phi-3-mini-4k-instruct": { + "deprecation_date": "2025-08-30", "input_cost_per_token": 1.3e-07, "litellm_provider": "azure_ai", "max_input_tokens": 4096, @@ -12022,11 +12032,12 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 5.2e-07, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true, "supports_vision": false }, "azure_ai/Phi-3-small-128k-instruct": { + "deprecation_date": "2025-08-30", "input_cost_per_token": 1.5e-07, "litellm_provider": "azure_ai", "max_input_tokens": 128000, @@ -12034,11 +12045,12 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-07, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true, "supports_vision": false }, "azure_ai/Phi-3-small-8k-instruct": { + "deprecation_date": "2025-08-30", "input_cost_per_token": 1.5e-07, "litellm_provider": "azure_ai", "max_input_tokens": 8192, @@ -12046,11 +12058,12 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-07, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true, "supports_vision": false }, "azure_ai/Phi-3.5-MoE-instruct": { + "deprecation_date": "2025-08-30", "input_cost_per_token": 1.6e-07, "litellm_provider": "azure_ai", "max_input_tokens": 128000, @@ -12058,11 +12071,12 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6.4e-07, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true, "supports_vision": false }, "azure_ai/Phi-3.5-mini-instruct": { + "deprecation_date": "2025-08-30", "input_cost_per_token": 1.3e-07, "litellm_provider": "azure_ai", "max_input_tokens": 128000, @@ -12070,11 +12084,12 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 5.2e-07, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true, "supports_vision": false }, "azure_ai/Phi-3.5-vision-instruct": { + "deprecation_date": "2025-08-30", "input_cost_per_token": 1.3e-07, "litellm_provider": "azure_ai", "max_input_tokens": 128000, @@ -12082,7 +12097,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 5.2e-07, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true, "supports_vision": true }, @@ -12217,6 +12232,7 @@ "source": "https://azure.microsoft.com/en-us/pricing/details/ai-document-intelligence/" }, "azure_ai/MAI-DS-R1": { + "deprecation_date": "2026-02-27", "input_cost_per_token": 1.35e-06, "litellm_provider": "azure_ai", "max_input_tokens": 128000, @@ -12224,11 +12240,12 @@ "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 5.4e-06, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_reasoning": true, "supports_tool_choice": true }, "azure_ai/cohere-rerank-v3-english": { + "deprecation_date": "2025-06-30", "input_cost_per_query": 0.002, "input_cost_per_token": 0.0, "litellm_provider": "azure_ai", @@ -12236,9 +12253,11 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "rerank", - "output_cost_per_token": 0.0 + "output_cost_per_token": 0.0, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models" }, "azure_ai/cohere-rerank-v3-multilingual": { + "deprecation_date": "2025-06-30", "input_cost_per_query": 0.002, "input_cost_per_token": 0.0, "litellm_provider": "azure_ai", @@ -12246,7 +12265,8 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "rerank", - "output_cost_per_token": 0.0 + "output_cost_per_token": 0.0, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models" }, "azure_ai/cohere-rerank-v4.0-pro": { "input_cost_per_query": 0.0025, @@ -12301,6 +12321,7 @@ "supports_tool_choice": true }, "azure_ai/deepseek-v3": { + "deprecation_date": "2025-08-31", "input_cost_per_token": 1.14e-06, "litellm_provider": "azure_ai", "max_input_tokens": 128000, @@ -12308,7 +12329,7 @@ "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 4.56e-06, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true }, "azure_ai/deepseek-v4-pro": { @@ -12515,6 +12536,7 @@ "supports_web_search": true }, "azure_ai/jais-30b-chat": { + "deprecation_date": "2026-01-30", "input_cost_per_token": 0.0032, "litellm_provider": "azure_ai", "max_input_tokens": 8192, @@ -12522,7 +12544,7 @@ "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 0.00971, - "source": "https://ai.azure.com/catalog/models/jais-30b-chat" + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models" }, "azure_ai/jamba-instruct": { "input_cost_per_token": 5e-07, @@ -12588,6 +12610,7 @@ "supports_tool_choice": true }, "azure_ai/mistral-large": { + "deprecation_date": "2025-04-15", "input_cost_per_token": 4e-06, "litellm_provider": "azure_ai", "max_input_tokens": 32000, @@ -12595,10 +12618,12 @@ "max_tokens": 8191, "mode": "chat", "output_cost_per_token": 1.2e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_function_calling": true, "supports_tool_choice": true }, "azure_ai/mistral-large-2407": { + "deprecation_date": "2025-05-13", "input_cost_per_token": 2e-06, "litellm_provider": "azure_ai", "max_input_tokens": 128000, @@ -12606,7 +12631,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-06, - "source": "https://marketplace.microsoft.com/en/marketplace/apps/000-000.mistral-ai-large-2407-offer?tab=Overview", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_function_calling": true, "supports_tool_choice": true }, @@ -12648,6 +12673,7 @@ "supports_tool_choice": true }, "azure_ai/mistral-nemo": { + "deprecation_date": "2026-01-30", "input_cost_per_token": 1.5e-07, "litellm_provider": "azure_ai", "max_input_tokens": 131072, @@ -12655,10 +12681,11 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.5e-07, - "source": "https://marketplace.microsoft.com/en/marketplace/apps/000-000.mistral-nemo-12b-2407?tab=PlansAndPrice", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_function_calling": true }, "azure_ai/mistral-small": { + "deprecation_date": "2025-07-31", "input_cost_per_token": 1e-06, "litellm_provider": "azure_ai", "max_input_tokens": 32000, @@ -12666,6 +12693,7 @@ "max_tokens": 8191, "mode": "chat", "output_cost_per_token": 3e-06, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_function_calling": true, "supports_tool_choice": true }, @@ -42206,14 +42234,14 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro": { - "cache_read_input_token_cost": 6.525e-08, - "input_cost_per_token": 7.83e-07, + "cache_read_input_token_cost": 1.74e-08, + "input_cost_per_token": 2.088e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 1.566e-06, + "output_cost_per_token": 4.176e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42226,14 +42254,14 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4.1-flash": { - "cache_read_input_token_cost": 2.91e-09, - "input_cost_per_token": 1.98e-08, + "cache_read_input_token_cost": 1e-08, + "input_cost_per_token": 3e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 3.96e-07, + "output_cost_per_token": 5e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42820,14 +42848,14 @@ "supports_web_search": false }, "openrouter/nvidia/nemotron-3.5-lightning": { - "cache_read_input_token_cost": 3e-08, - "input_cost_per_token": 6e-08, + "cache_read_input_token_cost": 2.975e-08, + "input_cost_per_token": 5.95e-08, "litellm_provider": "openrouter", "max_input_tokens": 262144, - "max_output_tokens": 32768, - "max_tokens": 32768, + "max_output_tokens": 131072, + "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 1.6e-07, + "output_cost_per_token": 1.7e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -43829,6 +43857,7 @@ "openrouter/z-ai/glm-4.7": { "cache_creation_input_token_cost": 0.0, "cache_read_input_token_cost": 1.1e-07, + "deprecation_date": "2026-12-31", "input_cost_per_token": 6e-07, "litellm_provider": "openrouter", "max_input_tokens": 204800, @@ -67382,13 +67411,13 @@ "supports_web_search": false }, "openrouter/z-ai/glm-5.3": { - "input_cost_per_token": 1.4e-06, - "output_cost_per_token": 4.4e-06, - "cache_read_input_token_cost": 2.6e-07, + "input_cost_per_token": 2.219e-07, + "output_cost_per_token": 3.39e-06, + "cache_read_input_token_cost": 1.775e-07, "litellm_provider": "openrouter", - "max_input_tokens": 1310720, - "max_output_tokens": 943717, - "max_tokens": 943717, + "max_input_tokens": 1048576, + "max_output_tokens": 943718, + "max_tokens": 943718, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, @@ -67519,8 +67548,8 @@ "supports_prompt_caching": true }, "openrouter/deepseek/deepseek-v4-flash-0731": { - "cache_read_input_token_cost": 8.9e-09, - "input_cost_per_token": 8.9e-09, + "cache_read_input_token_cost": 1.08e-08, + "input_cost_per_token": 1.08e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, @@ -67608,8 +67637,8 @@ "supports_web_search": false }, "openrouter/moonshotai/kimi-k3": { - "cache_read_input_token_cost": 2.7e-07, - "input_cost_per_token": 2.8e-07, + "cache_read_input_token_cost": 4.357e-07, + "input_cost_per_token": 4.357e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, @@ -67732,7 +67761,7 @@ }, "openrouter/z-ai/glm-5.2": { "cache_read_input_token_cost": 2.6e-07, - "input_cost_per_token": 3.249e-07, + "input_cost_per_token": 4.1e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, @@ -68094,14 +68123,14 @@ "supports_web_search": true }, "openrouter/deepseek/deepseek-v4-flash": { - "cache_read_input_token_cost": 1.5708e-08, - "input_cost_per_token": 7.854e-08, + "cache_read_input_token_cost": 8.372e-09, + "input_cost_per_token": 4.186e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, - "max_output_tokens": 384000, - "max_tokens": 384000, + "max_output_tokens": 131072, + "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 1.5708e-07, + "output_cost_per_token": 8.372e-08, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -68114,9 +68143,9 @@ "supports_web_search": false }, "openrouter/moonshotai/kimi-k2.6": { - "input_cost_per_token": 6.5e-07, - "output_cost_per_token": 3.41e-06, - "cache_read_input_token_cost": 1.5e-07, + "input_cost_per_token": 4.3415e-07, + "output_cost_per_token": 1.828e-06, + "cache_read_input_token_cost": 7.312e-08, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 235929, @@ -68563,7 +68592,7 @@ "openrouter/z-ai/glm-4.6v": { "input_cost_per_token": 3e-07, "output_cost_per_token": 9e-07, - "cache_read_input_token_cost": 5.5e-08, + "cache_read_input_token_cost": 5e-08, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 32768, @@ -69229,7 +69258,7 @@ "supports_web_search": false }, "openrouter/minimax/minimax-m1": { - "input_cost_per_token": 4e-07, + "input_cost_per_token": 5.5e-07, "output_cost_per_token": 2.2e-06, "litellm_provider": "openrouter", "max_input_tokens": 1000000, @@ -70653,7 +70682,7 @@ "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" }, "azure/eu/gpt-4.1-nano": { - "deprecation_date": "2026-10-14", + "deprecation_date": "2027-04-14", "cache_read_input_token_cost": 2.8e-08, "input_cost_per_token": 1.1e-07, "input_cost_per_token_batches": 5.5e-08, @@ -71102,7 +71131,7 @@ "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" }, "azure/us/gpt-4.1-nano": { - "deprecation_date": "2026-10-14", + "deprecation_date": "2027-04-14", "cache_read_input_token_cost": 2.8e-08, "input_cost_per_token": 1.1e-07, "input_cost_per_token_batches": 5.5e-08, diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 2c63d957e75..9c7778e2b77 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -1947,10 +1947,41 @@ class MCPRequestHandler: scope, the tool union and billing attribution are all just different reads of this one answer — computing it separately per consumer is how they drift (a throttle map scoped by roster instead of by grant charged unrelated teams' buckets).""" - return [ + grants: Final = [ (source, set(await MCPRequestHandler.get_allowed_mcp_servers(source, keyless_source=True))) for source in await MCPRequestHandler.admitted_subject_sources(auth, allowed_team_ids=allowed_team_ids) ] + scope: Final = await MCPRequestHandler._toolset_scope(auth) + if scope is None: + return grants + return [(source, granted & frozenset(scope)) for source, granted in grants] + + @staticmethod + async def _toolset_scope(auth: UserAPIKeyAuth) -> dict[str, list[str]] | None: + """The ``server_id -> tools`` a namespaced toolset route pinned this subject to via + ``mcp_toolset_id``, or None on the aggregate scope. Every source's servers and tools are + intersected with it, so the route narrows a team grant exactly as it narrows the user's own.""" + if auth.mcp_toolset_id is None: + return None + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + + return await global_mcp_server_manager.resolve_toolset_tool_permissions( + toolset_ids=[auth.mcp_toolset_id], requires_fresh_policy=auth.requires_fresh_policy + ) + + @staticmethod + async def _narrow_tools_to_toolset( + tools: list[str] | None, + server_id: str, + auth: UserAPIKeyAuth, + ) -> list[str] | None: + scope: Final = await MCPRequestHandler._toolset_scope(auth) + if scope is None: + return tools + scoped: Final = frozenset(scope.get(server_id, ())) + return sorted(scoped if tools is None else scoped & frozenset(tools)) @staticmethod async def resolve_admitted_subject_servers( @@ -2048,9 +2079,9 @@ class MCPRequestHandler: continue tools = await MCPRequestHandler.get_allowed_tools_for_server(server_id, source, keyless_source=True) if tools is None: - return None + return await MCPRequestHandler._narrow_tools_to_toolset(None, server_id, auth) allowed.update(tools) - return sorted(allowed) + return await MCPRequestHandler._narrow_tools_to_toolset(sorted(allowed), server_id, auth) @staticmethod def _get_key_object_permission( @@ -2067,6 +2098,16 @@ class MCPRequestHandler: return user_api_key_auth.object_permission + @staticmethod + async def team_object_permission(user_api_key_auth: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable | None: + return await MCPRequestHandler._get_team_object_permission(user_api_key_auth) + + @staticmethod + async def key_object_permission_hydrated( + user_api_key_auth: UserAPIKeyAuth, + ) -> LiteLLM_ObjectPermissionTable | None: + return await MCPRequestHandler._key_object_permission_hydrated(user_api_key_auth) + @staticmethod async def _get_team_object_permission( user_api_key_auth: UserAPIKeyAuth | None = None, diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 44523f5cf45..d17b849750b 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -915,11 +915,9 @@ if MCP_AVAILABLE: ) -> UserAPIKeyAuth: """The one credential this tools request acts as. - A toolset name narrows the caller's own credential to that toolset; otherwise a dashboard - session is swapped for its admitted subject. The two are mutually exclusive by construction, - which is why they share an owner: the admitted subject resolves per grant source and a team - source deliberately carries none of the caller's ``object_permission``, so a toolset - narrowing layered on top would evaporate on every team-granted server.""" + A toolset name pins the acting principal to that toolset through ``_apply_toolset_scope``, + which itself swaps a dashboard session for its admitted subject; otherwise the swap happens + here so both shapes resolve as the same identity.""" if not toolset_name: return await acting_user_auth(user_api_key_dict) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 44979d12a00..2090a0c7421 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -66,7 +66,13 @@ from litellm.proxy._experimental.mcp_server.oauth_utils import ( get_route_relative_request_path, well_known_root_suffix, ) -from litellm.proxy._experimental.mcp_server.ui_session_utils import is_ui_session_credential +from litellm.proxy._experimental.mcp_server.ui_session_utils import ( + ActingUser, + GrantedToolsetIds, + acting_user_auth, + granted_toolset_ids, + is_ui_session_credential, +) from litellm.proxy._experimental.mcp_server.utils import ( LITELLM_MCP_SERVER_DESCRIPTION, LITELLM_MCP_SERVER_NAME, @@ -1515,16 +1521,21 @@ if MCP_AVAILABLE: async def _apply_toolset_scope( user_api_key_auth: UserAPIKeyAuth, toolset_id: str, + acting_user: ActingUser = acting_user_auth, + granted: GrantedToolsetIds = granted_toolset_ids, ) -> UserAPIKeyAuth: """ - Restrict a key's MCP permissions to a single toolset. + Pin a principal's MCP permissions to a single toolset for /toolset/{name}/mcp. - When a request arrives via /toolset/{name}/mcp we override the key's - object_permission so that only the toolset's tools are visible. + A virtual key (and an admin session) has its object_permission rewritten to + the toolset's servers and tools. A keyless subject resolves per grant source, + so a non-admin dashboard session first becomes its admitted user and the + toolset rides along as ``mcp_toolset_id``, which every source's grant is + intersected with; a team-granted toolset is served without the user's own + row capping it. - Raises HTTPException(403) if the key has an explicit toolset grant list - that does not include toolset_id (i.e. mcp_toolsets is set but empty, - or set to a list that omits this toolset). Admin keys always pass. + Raises HTTPException(403) unless the principal holds toolset_id through one + of its grant sources. Admins always pass. """ from litellm.proxy._types import LiteLLM_ObjectPermissionTable from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view @@ -1540,26 +1551,31 @@ if MCP_AVAILABLE: detail="API key is scoped to no MCP servers; toolset access is denied.", ) - # Access control: non-admin keys must have this toolset in their grant list. - # Use _user_has_admin_view so that PROXY_ADMIN_VIEW_ONLY is also treated as admin. - is_admin: Final = _user_has_admin_view(user_api_key_auth) - if not is_admin: - op: Final = user_api_key_auth.object_permission - granted: Final = getattr(op, "mcp_toolsets", None) if op else None - # granted=None → key has no explicit toolset grants → deny (same semantics as - # fetch_mcp_toolsets which returns [] for non-admin keys with no grants configured). - # granted=[] or list without toolset_id → also deny. - if granted is None or toolset_id not in granted: + acting: Final = await acting_user(user_api_key_auth) + is_admin: Final = _user_has_admin_view(acting) + if not is_admin and toolset_id not in await granted(acting): + raise HTTPException( + status_code=403, + detail=f"API key does not have access to toolset '{toolset_id}'.", + ) + if _is_mcp_admitted_user_subject(acting): + resource_server_id: Final = acting.mcp_session_resource_server_id + if resource_server_id is not None and resource_server_id not in ( + await operations.global_mcp_server_manager.resolve_toolset_tool_permissions( + toolset_ids=[toolset_id], requires_fresh_policy=acting.requires_fresh_policy + ) + ): raise HTTPException( status_code=403, detail=f"API key does not have access to toolset '{toolset_id}'.", ) + return acting.model_copy(update={"mcp_toolset_id": toolset_id}) tool_permissions = await operations.global_mcp_server_manager.resolve_toolset_tool_permissions( toolset_ids=[toolset_id] ) server_ids: Final = list(tool_permissions.keys()) - existing_op: Final = user_api_key_auth.object_permission + existing_op: Final = acting.object_permission if existing_op is not None: updated_op = existing_op.model_copy( update={ @@ -1576,7 +1592,12 @@ if MCP_AVAILABLE: mcp_servers=server_ids, mcp_tool_permissions=tool_permissions, ) - return user_api_key_auth.model_copy(update={"object_permission": updated_op, "mcp_toolset_id": toolset_id}) + return acting.model_copy(update={"object_permission": updated_op, "mcp_toolset_id": toolset_id}) + + async def _toolset_server_ids(toolset_id: str) -> set[str]: + return set( + await operations.global_mcp_server_manager.resolve_toolset_tool_permissions(toolset_ids=[toolset_id]) + ) async def _raise_preemptive_401_for_unauthenticated_servers( scope: Scope, @@ -2007,8 +2028,7 @@ if MCP_AVAILABLE: toolset_allowed_server_ids: set[str] | None = None if active_toolset_id and user_api_key_auth is not None: user_api_key_auth = await _apply_toolset_scope(user_api_key_auth, active_toolset_id) - op: Final = user_api_key_auth.object_permission - toolset_allowed_server_ids = set(op.mcp_servers or []) if op else set() + toolset_allowed_server_ids = await _toolset_server_ids(active_toolset_id) # https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response # Must run after toolset scoping so the challenge set is derived @@ -2355,8 +2375,7 @@ if MCP_AVAILABLE: toolset_allowed_server_ids: set[str] | None = None if active_toolset_id and user_api_key_auth is not None: user_api_key_auth = await _apply_toolset_scope(user_api_key_auth, active_toolset_id) - op: Final = user_api_key_auth.object_permission - toolset_allowed_server_ids = set(op.mcp_servers or []) if op else set() + toolset_allowed_server_ids = await _toolset_server_ids(active_toolset_id) # https://datatracker.ietf.org/doc/html/rfc9728#name-www-authenticate-response # Must run after toolset scoping so the challenge set is derived diff --git a/litellm/proxy/_experimental/mcp_server/ui_session_utils.py b/litellm/proxy/_experimental/mcp_server/ui_session_utils.py index 901259c18ad..b9e25259868 100644 --- a/litellm/proxy/_experimental/mcp_server/ui_session_utils.py +++ b/litellm/proxy/_experimental/mcp_server/ui_session_utils.py @@ -2,14 +2,23 @@ from __future__ import annotations -from collections.abc import Awaitable, Callable -from typing import Final +import asyncio +from collections.abc import Awaitable, Callable, Sequence +from itertools import chain +from typing import Final, TypeAlias from fastapi import HTTPException from litellm._logging import verbose_logger from litellm.constants import UI_SESSION_TOKEN_TEAM_ID -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth + +EffectiveAuthContexts: TypeAlias = Callable[[UserAPIKeyAuth], Awaitable[Sequence[UserAPIKeyAuth]]] +TeamObjectPermission: TypeAlias = Callable[[UserAPIKeyAuth], Awaitable[LiteLLM_ObjectPermissionTable | None]] +OwnObjectPermission: TypeAlias = Callable[[UserAPIKeyAuth], Awaitable[LiteLLM_ObjectPermissionTable | None]] +AdmittedContext: TypeAlias = Callable[[UserAPIKeyAuth], Awaitable[UserAPIKeyAuth | None]] +ActingUser: TypeAlias = Callable[[UserAPIKeyAuth], Awaitable[UserAPIKeyAuth]] +GrantedToolsetIds: TypeAlias = Callable[[UserAPIKeyAuth], Awaitable[frozenset[str]]] def clone_user_api_key_auth_with_team( @@ -109,10 +118,10 @@ async def acting_user_auth(user_api_key_auth: UserAPIKeyAuth) -> UserAPIKeyAuth: both surfaces. An admin session keeps its operator view and any caller-passed credential is returned unchanged, never widened. - Do not combine this with a narrowing that rewrites a single credential's ``object_permission`` - (toolset scope): the admitted subject resolves per grant source and a team source deliberately - carries none of the caller's own grants, so the narrowing would silently evaporate on every - team-granted server. A request carrying such a scope keeps the caller's own credential.""" + A toolset narrowing is never applied to the admitted subject by rewriting its ``object_permission``: + it resolves per grant source and a team source deliberately carries none of the caller's own grants, + so the rewrite would evaporate on every team-granted server. The route pins ``mcp_toolset_id`` + instead, which every source's grant is intersected with.""" if not is_ui_session_credential(user_api_key_auth): return user_api_key_auth @@ -153,3 +162,135 @@ async def can_access_mcp_server( if server_id in await allowed_servers(context): return True return False + + +def _restricts_mcp(permission: LiteLLM_ObjectPermissionTable | None) -> bool: + return permission is not None and bool( + permission.mcp_servers + or permission.mcp_toolsets + or permission.mcp_tool_permissions + or permission.mcp_access_groups + ) + + +def is_keyless_mcp_subject(user_api_key_auth: UserAPIKeyAuth) -> bool: + """A principal with no virtual key to declare MCP access on: the dashboard's own session token or a + gateway-admitted user. Its grants are resolved per source, never through a key row.""" + + return is_ui_session_credential(user_api_key_auth) or user_api_key_auth.mcp_admitted_user_subject is True + + +async def toolset_grant_contexts( + user_api_key_auth: UserAPIKeyAuth, + admitted_context: AdmittedContext = admitted_user_context, + admitted_sources: EffectiveAuthContexts | None = None, +) -> Sequence[UserAPIKeyAuth]: + """The grant sources a toolset is looked up through. A virtual key is its own single source. A keyless + subject fans out exactly as the aggregate /mcp resolution does: its own user row plus every team whose + live roster still lists it, so a membership that only survives in the user's cached team list grants + nothing.""" + + if not is_keyless_mcp_subject(user_api_key_auth): + return (user_api_key_auth,) + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + + load_sources: Final = admitted_sources or MCPRequestHandler.admitted_subject_sources + acting: Final = await admitted_context(user_api_key_auth) + return tuple(await load_sources(acting if acting is not None else user_api_key_auth)) + + +async def _own_toolset_ids( + context: UserAPIKeyAuth, + load_own_permission: OwnObjectPermission, +) -> Sequence[str] | None: + """The source's own toolsets, or None when it declares no MCP grant of its own. A source that names an + ``object_permission_id`` is a known restriction even when the row is unhydrated, unreadable or gone, + so it is loaded rather than read as unrestricted, and grants nothing when it cannot be read.""" + if context.object_permission is None and not context.object_permission_id: + return None + try: + own: Final = await load_own_permission(context) + except Exception as exc: # noqa: BLE001 # a named but unreadable own grant must deny, not widen to the team + verbose_logger.warning( + "MCP toolset grants: object permission %s unreadable, granting nothing through it: %s", + context.object_permission_id, + exc, + ) + return () + if own is None: + return () + if not _restricts_mcp(own): + return None + return own.mcp_toolsets or () + + +async def _inherited_toolset_ids( + context: UserAPIKeyAuth, + load_team_permission: TeamObjectPermission, +) -> Sequence[str]: + try: + team: Final = await load_team_permission(context) + except Exception as exc: # noqa: BLE001 # an unreadable team grants nothing through this source and must not fail the caller's other sources + verbose_logger.warning( + "MCP toolset grants: team %s unreadable, inheriting nothing from it: %s", + context.team_id, + exc, + ) + return () + return () if team is None else (team.mcp_toolsets or ()) + + +async def _context_toolset_ids( + context: UserAPIKeyAuth, + inherits_team: bool, + load_team_permission: TeamObjectPermission, + load_own_permission: OwnObjectPermission, +) -> Sequence[str]: + own: Final = await _own_toolset_ids(context, load_own_permission) + if own is not None: + return own + if not inherits_team or not context.team_id: + return () + return await _inherited_toolset_ids(context, load_team_permission) + + +async def granted_toolset_ids( + user_api_key_auth: UserAPIKeyAuth, + effective_contexts: EffectiveAuthContexts = toolset_grant_contexts, + team_object_permission: TeamObjectPermission | None = None, + require_key_access: bool | None = None, + own_object_permission: OwnObjectPermission | None = None, +) -> frozenset[str]: + """Toolset ids the principal holds, resolved per grant source with the key/team rule the aggregate + /mcp listing applies: a source that declares any MCP grant of its own is scoped to its own toolsets and + never reads its team, one that declares none inherits its team's, except a virtual key under + ``require_key_mcp_access_defined``, which inherits nothing. A keyless subject's team sources always + inherit. A team that cannot be read contributes nothing while every other source still counts, and an + own grant that is named but cannot be read grants nothing. No grant anywhere yields the empty set.""" + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + from litellm.proxy.proxy_server import general_settings + + require: Final = ( + require_key_access + if require_key_access is not None + else bool( + general_settings.get( # pyright: ignore[reportUnknownArgumentType] # general_settings is an untyped dict; truthiness must match the /mcp path's read of this flag + "require_key_mcp_access_defined", False + ) + ) + ) + inherits_team: Final = is_keyless_mcp_subject(user_api_key_auth) or not require + load_team_permission: Final = team_object_permission or MCPRequestHandler.team_object_permission + load_own_permission: Final = own_object_permission or MCPRequestHandler.key_object_permission_hydrated + contexts: Final = await effective_contexts(user_api_key_auth) + per_context: Final = await asyncio.gather( + *( + _context_toolset_ids(context, inherits_team, load_team_permission, load_own_permission) + for context in contexts + ) + ) + return frozenset(chain.from_iterable(per_context)) diff --git a/litellm/proxy/_experimental/out/assets/logos/microsoft_365.svg b/litellm/proxy/_experimental/out/assets/logos/microsoft_365.svg new file mode 100644 index 00000000000..e053ac831fb --- /dev/null +++ b/litellm/proxy/_experimental/out/assets/logos/microsoft_365.svg @@ -0,0 +1 @@ + diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index 0b687340ea5..0b470cf7bda 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -3,19 +3,24 @@ Lazy registration for optional feature routers. Each LAZY_FEATURES entry imports its module only on the first request matching its path prefix, saving ~700 MB at idle for deployments that don't use these features. First hit pays the import cost (1-3 s for heavy modules); /openapi.json -omits each feature's routes until the feature is warmed. +omits each feature's routes until the feature is warmed. Setting +LITELLM_DISABLE_LAZY_ROUTES registers every feature at worker startup +instead, so the route table is complete before the first request. """ import asyncio import importlib -from collections.abc import Callable, Mapping, Sequence +import os +from collections.abc import AsyncGenerator, Callable, Mapping, Sequence from collections.abc import Set as AbstractSet +from contextlib import asynccontextmanager from dataclasses import dataclass, field +from functools import partial from types import MappingProxyType from typing import TYPE_CHECKING, Final from starlette.routing import BaseRoute, Match -from starlette.types import ASGIApp, Receive, Scope, Send +from starlette.types import ASGIApp, Lifespan, Receive, Scope, Send from litellm._logging import verbose_proxy_logger from litellm.proxy.route_priority import hot_routes_first @@ -428,57 +433,127 @@ def _in_registry_order( async def _force_load(app: "FastAPI", feat: LazyFeature, features: tuple[LazyFeature, ...] = LAZY_FEATURES) -> bool: """Import + register a lazy feature exactly once per (app, module). Shared by the middleware and the /lazy/warm endpoint.""" + async with _lazy_lock(app, feat.module_path): + if feat.module_path in _lazy_loaded(app): + return False + # Import on a thread (heavy modules take 1-3 s). register_fn + # mutates app.router.routes, so it stays on the loop thread. + imported: Final = asyncio.get_running_loop().run_in_executor(None, importlib.import_module, feat.module_path) + await asyncio.wait((imported,)) + return _install(app, feat, imported.result, features) + + +def _install( + app: "FastAPI", feat: LazyFeature, import_module: Callable[[], object], features: tuple[LazyFeature, ...] +) -> bool: + try: + _register_feature(app, feat, import_module(), features) + return True + except Exception as exc: + _mark_failed(app, feat, exc) + return False + + +def _lazy_loaded(app: "FastAPI") -> set[str]: if not hasattr(app.state, "lazy_loaded"): - app.state.lazy_loaded = set() - app.state.lazy_locks = {} - lock: Final = app.state.lazy_locks.setdefault(feat.module_path, asyncio.Lock()) - async with lock: - if feat.module_path in app.state.lazy_loaded: - return False - try: - # Import on a thread (heavy modules take 1-3 s). register_fn - # mutates app.router.routes, so it stays on the loop thread. - loop: Final = asyncio.get_running_loop() - module: Final = await loop.run_in_executor(None, importlib.import_module, feat.module_path) - before: Final = len(app.router.routes) - feat.register_fn(app, module) - previous: Final[Mapping[str, tuple[BaseRoute, ...]]] = ( - app.state.lazy_routes if hasattr(app.state, "lazy_routes") else MappingProxyType({}) - ) - lazy_routes: Final[Mapping[str, tuple[BaseRoute, ...]]] = MappingProxyType( - {**previous, feat.module_path: tuple(app.router.routes[before:])} - ) - app.state.lazy_routes = lazy_routes # rebind-ok: the app owns the record of which routes each feature added - app.router.routes[:] = hot_routes_first( # rebind-ok: the app owns its route table - _in_registry_order(app.router.routes, lazy_routes, features, _lazy_slots(app)) - ) - app.state.lazy_loaded.add(feat.module_path) - app.openapi_schema = None - verbose_proxy_logger.info( - "Lazy-loaded optional feature %r (module: %s)", - feat.name, - feat.module_path, - ) - return True - except Exception as exc: - # Mark loaded anyway so we don't retry on every request. - app.state.lazy_loaded.add(feat.module_path) - verbose_proxy_logger.warning( - "Failed to lazy-load optional feature %r (module: %s): %s. " - "This feature's endpoints will return 404 until restart.", - feat.name, - feat.module_path, - exc, - ) - return False + app.state.lazy_loaded = set[str]() + app.state.lazy_locks = dict[str, asyncio.Lock]() + loaded: Final[set[str]] = app.state.lazy_loaded + return loaded -def attach_lazy_features(app: "FastAPI") -> None: - app.include_router(_make_warmup_router(app)) - app.add_middleware(LazyFeatureMiddleware, fastapi_app=app) +def _lazy_lock(app: "FastAPI", module_path: str) -> asyncio.Lock: + if not hasattr(app.state, "lazy_locks"): + app.state.lazy_locks = dict[str, asyncio.Lock]() + locks: Final[dict[str, asyncio.Lock]] = app.state.lazy_locks + return locks.setdefault(module_path, asyncio.Lock()) -def _make_warmup_router(app: "FastAPI") -> "APIRouter": +def _register_feature(app: "FastAPI", feat: LazyFeature, module: object, features: tuple[LazyFeature, ...]) -> None: + before: Final = len(app.router.routes) + feat.register_fn(app, module) + previous: Final[Mapping[str, tuple[BaseRoute, ...]]] = ( + app.state.lazy_routes if hasattr(app.state, "lazy_routes") else MappingProxyType({}) + ) + lazy_routes: Final[Mapping[str, tuple[BaseRoute, ...]]] = MappingProxyType( + {**previous, feat.module_path: tuple(app.router.routes[before:])} + ) + app.state.lazy_routes = lazy_routes # rebind-ok: the app owns the record of which routes each feature added + app.router.routes[:] = hot_routes_first( # rebind-ok: the app owns its route table + _in_registry_order(app.router.routes, lazy_routes, features, _lazy_slots(app)) + ) + _lazy_loaded(app).add(feat.module_path) + app.openapi_schema = None + verbose_proxy_logger.info( + "Lazy-loaded optional feature %r (module: %s)", + feat.name, + feat.module_path, + ) + + +def _mark_failed(app: "FastAPI", feat: LazyFeature, exc: Exception) -> None: + # Mark loaded anyway so we don't retry on every request. + _lazy_loaded(app).add(feat.module_path) + verbose_proxy_logger.warning( + "Failed to lazy-load optional feature %r (module: %s): %s. " + "This feature's endpoints will return 404 until restart.", + feat.name, + feat.module_path, + exc, + ) + + +def lazy_routes_disabled() -> bool: + return os.getenv("LITELLM_DISABLE_LAZY_ROUTES", "").lower() in ("1", "true", "yes", "on") + + +def register_all_features(app: "FastAPI", features: tuple[LazyFeature, ...] = LAZY_FEATURES) -> None: + """Register every feature router now, in registry order, so app.routes is + complete before the app serves its first request.""" + for feat in features: + _install(app, feat, partial(importlib.import_module, feat.module_path), features) + + +def attach_lazy_features(app: "FastAPI", features: tuple[LazyFeature, ...] = LAZY_FEATURES) -> None: + if lazy_routes_disabled(): + app.router.lifespan_context = _register_all_on_startup(app.router.lifespan_context, features) + return + app.include_router(_make_warmup_router(app, features)) + app.add_middleware(LazyFeatureMiddleware, fastapi_app=app, features=features) + + +def _register_all_on_startup(inner: "Lifespan[FastAPI]", features: tuple[LazyFeature, ...]) -> "Lifespan[FastAPI]": + """Registering at startup, once every route the app defines exists, lands the features + where lazy mode splices them: after every eager route (so /mcp/proxy, defined after + attach_lazy_features(), still beats the /mcp mount) and before LITELLM_WORKER_STARTUP_HOOKS + or an outer lifespan can filter the table. The inner lifespan then adds routes of its own + (config pass-through endpoints), so the table is put back in lazy mode's order once it is up.""" + + @asynccontextmanager + async def lifespan(app: "FastAPI") -> AsyncGenerator[None]: + register_all_features(app, features) + async with inner(app): + _restore_registry_order(app, features) + yield + + return lifespan + + +def _restore_registry_order(app: "FastAPI", features: tuple[LazyFeature, ...]) -> None: + present: Final = frozenset(id(route) for route in app.router.routes) + registered: Final[Mapping[str, tuple[BaseRoute, ...]]] = ( + app.state.lazy_routes if hasattr(app.state, "lazy_routes") else MappingProxyType({}) + ) + still_routed: Final = MappingProxyType( + {module_path: tuple(r for r in routes if id(r) in present) for module_path, routes in registered.items()} + ) + app.router.routes[:] = hot_routes_first( # rebind-ok: the app owns its route table + _in_registry_order(app.router.routes, still_routed, features, _lazy_slots(app)) + ) + app.openapi_schema = None + + +def _make_warmup_router(app: "FastAPI", features: tuple[LazyFeature, ...] = LAZY_FEATURES) -> "APIRouter": """POST /lazy/warm/{name}: load a feature and return its partial openapi so the Swagger plugin can merge in-place without a full /openapi.json refetch. Requires auth — anyone who can hit the proxy can already trigger the same @@ -497,13 +572,13 @@ def _make_warmup_router(app: "FastAPI") -> "APIRouter": dependencies=[Depends(user_api_key_auth)], ) async def warm(name: str): - feat: Final = next((f for f in LAZY_FEATURES if f.name == name), None) + feat: Final = next((f for f in features if f.name == name), None) if feat is None: raise HTTPException(404, f"unknown lazy feature: {name}") if feat.persistent_swagger_stub: return {"stub_path": None, "paths": {}, "components": {"schemas": {}}} - await _force_load(app, feat) + await _force_load(app, feat, features) feat_routes: Final = [r for r in app.routes if feat.matches(getattr(r, "path", ""))] full: Final = get_openapi(title=app.title, version=app.version, routes=feat_routes) @@ -524,7 +599,7 @@ def _make_warmup_router(app: "FastAPI") -> "APIRouter": def loaded_lazy_modules(app: "FastAPI") -> frozenset[str]: """The set of lazy feature modules whose routers are actually registered - on this app (tracked by _force_load), empty before the middleware ever ran. + on this app (tracked by _install), empty until a feature loads or eager startup runs. sys.modules is the wrong signal: boot code imports several feature modules (mcp_management, cloudzero, vantage, config_overrides) without mounting their routers, and their stubs must still be injected.""" @@ -583,6 +658,6 @@ def lazy_tag_to_prefix() -> dict[str, str]: because /openapi.json already has full route info.""" from litellm.proxy._lazy_openapi_snapshot import load_snapshot - if load_snapshot(): + if lazy_routes_disabled() or load_snapshot(): return {} return {feat.name: feat.path_prefixes[0] for feat in LAZY_FEATURES if not feat.persistent_swagger_stub} diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 663a8e0d6d8..5de6e91a0f2 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -3432,6 +3432,53 @@ "title": "BreakdownMetrics", "type": "object" }, + "DailyActivityKeyPageResponse": { + "properties": { + "api_keys": { + "items": { + "$ref": "#/components/schemas/KeySpendActivityRow" + }, + "title": "Api Keys", + "type": "array" + }, + "limit": { + "title": "Limit", + "type": "integer" + }, + "offset": { + "title": "Offset", + "type": "integer" + }, + "total_api_keys": { + "title": "Total Api Keys", + "type": "integer" + } + }, + "required": [ + "api_keys", + "total_api_keys", + "offset", + "limit" + ], + "title": "DailyActivityKeyPageResponse", + "type": "object" + }, + "DailyActivityKeySearchResponse": { + "properties": { + "api_keys": { + "items": { + "$ref": "#/components/schemas/KeyActivityRow" + }, + "title": "Api Keys", + "type": "array" + } + }, + "required": [ + "api_keys" + ], + "title": "DailyActivityKeySearchResponse", + "type": "object" + }, "DailySpendData": { "properties": { "breakdown": { @@ -3455,6 +3502,33 @@ }, "DailySpendMetadata": { "properties": { + "api_key_limit": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "description": "When set, api_keys and every api_key_breakdown list at most this many keys, ranked by spend. Totals and the model, provider, mcp and endpoint rollups still cover every key.", + "title": "Api Key Limit" + }, + "entity_total_api_keys": { + "anyOf": [ + { + "additionalProperties": { + "type": "integer" + }, + "type": "object" + }, + { + "type": "null" + } + ], + "description": "Distinct API keys per entity over the requested range, set when the entity breakdown is included. When an entity's count exceeds api_key_limit, its api_key_breakdown lists only its keys among the top api_key_limit keys overall.", + "title": "Entity Total Api Keys" + }, "has_more": { "default": false, "title": "Has More", @@ -3465,6 +3539,18 @@ "title": "Page", "type": "integer" }, + "total_api_keys": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "description": "Distinct API keys matching the filters. When this exceeds api_key_limit, the per-key lists are truncated to the highest-spend keys.", + "title": "Total Api Keys" + }, "total_api_requests": { "default": 0, "title": "Total Api Requests", @@ -3614,6 +3700,16 @@ "title": "EntraIdentityConfig", "type": "object" }, + "ExportType": { + "enum": [ + "daily", + "daily_with_keys", + "daily_with_models", + "daily_with_users" + ], + "title": "ExportType", + "type": "string" + }, "HTTPAuthSecurityScheme": { "description": "Defines a security scheme using HTTP authentication.", "properties": { @@ -3669,6 +3765,27 @@ "title": "HTTPValidationError", "type": "object" }, + "KeyActivityRow": { + "properties": { + "api_key": { + "title": "Api Key", + "type": "string" + }, + "metadata": { + "$ref": "#/components/schemas/KeyMetadata" + }, + "metrics": { + "$ref": "#/components/schemas/SpendMetrics" + } + }, + "required": [ + "api_key", + "metrics", + "metadata" + ], + "title": "KeyActivityRow", + "type": "object" + }, "KeyMetadata": { "description": "Metadata for a key", "properties": { @@ -3747,6 +3864,78 @@ "title": "KeyMetricWithMetadata", "type": "object" }, + "KeySpendActivityRow": { + "properties": { + "api_key": { + "title": "Api Key", + "type": "string" + }, + "metadata": { + "$ref": "#/components/schemas/KeyMetadata" + }, + "metrics": { + "$ref": "#/components/schemas/KeySpendMetrics" + } + }, + "required": [ + "api_key", + "metrics", + "metadata" + ], + "title": "KeySpendActivityRow", + "type": "object" + }, + "KeySpendMetrics": { + "properties": { + "api_requests": { + "default": 0, + "title": "Api Requests", + "type": "integer" + }, + "cache_creation_input_tokens": { + "default": 0, + "title": "Cache Creation Input Tokens", + "type": "integer" + }, + "cache_read_input_tokens": { + "default": 0, + "title": "Cache Read Input Tokens", + "type": "integer" + }, + "completion_tokens": { + "default": 0, + "title": "Completion Tokens", + "type": "integer" + }, + "failed_requests": { + "default": 0, + "title": "Failed Requests", + "type": "integer" + }, + "prompt_tokens": { + "default": 0, + "title": "Prompt Tokens", + "type": "integer" + }, + "spend": { + "default": 0.0, + "title": "Spend", + "type": "number" + }, + "successful_requests": { + "default": 0, + "title": "Successful Requests", + "type": "integer" + }, + "total_tokens": { + "default": 0, + "title": "Total Tokens", + "type": "integer" + } + }, + "title": "KeySpendMetrics", + "type": "object" + }, "MakeAgentsPublicRequest": { "properties": { "agent_ids": { @@ -3835,6 +4024,32 @@ "title": "MetricWithMetadata", "type": "object" }, + "ModelTopKeysResponse": { + "properties": { + "api_keys": { + "items": { + "$ref": "#/components/schemas/KeySpendActivityRow" + }, + "title": "Api Keys", + "type": "array" + }, + "by_model_group": { + "title": "By Model Group", + "type": "boolean" + }, + "model": { + "title": "Model", + "type": "string" + } + }, + "required": [ + "model", + "by_model_group", + "api_keys" + ], + "title": "ModelTopKeysResponse", + "type": "object" + }, "MutualTLSSecurityScheme": { "description": "Defines a security scheme using mTLS authentication.", "properties": { @@ -4436,6 +4651,876 @@ ] } }, + "/agent/daily/activity/aggregated": { + "get": { + "operationId": "get_agent_daily_activity_aggregated_agent_daily_activity_aggregated_get", + "parameters": [ + { + "in": "query", + "name": "api_key_limit", + "required": false, + "schema": { + "default": 100, + "maximum": 1000, + "minimum": 1, + "title": "Api Key Limit", + "type": "integer" + } + }, + { + "in": "query", + "name": "agent_ids", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Agent Ids" + } + }, + { + "in": "query", + "name": "start_date", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Start Date" + } + }, + { + "in": "query", + "name": "end_date", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "End Date" + } + }, + { + "in": "query", + "name": "model", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Model" + } + }, + { + "in": "query", + "name": "api_key", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Api Key" + } + }, + { + "in": "query", + "name": "exclude_agent_ids", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Exclude Agent Ids" + } + }, + { + "in": "query", + "name": "timezone", + "required": false, + "schema": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "title": "Timezone" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/SpendAnalyticsPaginatedResponse" + } + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Get Agent Daily Activity Aggregated", + "tags": [ + "agents" + ] + } + }, + "/agent/daily/activity/aggregated/keys": { + "get": { + "operationId": "get_agent_daily_activity_aggregated_keys_agent_daily_activity_aggregated_keys_get", + "parameters": [ + { + "in": "query", + "name": "offset", + "required": false, + "schema": { + "default": 0, + "minimum": 0, + "title": "Offset", + "type": "integer" + } + }, + { + "in": "query", + "name": "limit", + "required": false, + "schema": { + "default": 50, + "maximum": 100, + "minimum": 1, + "title": "Limit", + "type": "integer" + } + }, + { + "in": "query", + "name": "agent_ids", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Agent Ids" + } + }, + { + "in": "query", + "name": "start_date", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Start Date" + } + }, + { + "in": "query", + "name": "end_date", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "End Date" + } + }, + { + "in": "query", + "name": "model", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Model" + } + }, + { + "in": "query", + "name": "api_key", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Api Key" + } + }, + { + "in": "query", + "name": "exclude_agent_ids", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Exclude Agent Ids" + } + }, + { + "in": "query", + "name": "timezone", + "required": false, + "schema": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "title": "Timezone" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/DailyActivityKeyPageResponse" + } + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Get Agent Daily Activity Aggregated Keys", + "tags": [ + "agents" + ] + } + }, + "/agent/daily/activity/aggregated/model_top_keys": { + "get": { + "operationId": "get_agent_daily_activity_model_top_keys_agent_daily_activity_aggregated_model_top_keys_get", + "parameters": [ + { + "in": "query", + "name": "model_group", + "required": true, + "schema": { + "minLength": 1, + "title": "Model Group", + "type": "string" + } + }, + { + "in": "query", + "name": "by_model_group", + "required": false, + "schema": { + "default": true, + "title": "By Model Group", + "type": "boolean" + } + }, + { + "in": "query", + "name": "limit", + "required": false, + "schema": { + "default": 5, + "maximum": 100, + "minimum": 1, + "title": "Limit", + "type": "integer" + } + }, + { + "in": "query", + "name": "agent_ids", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Agent Ids" + } + }, + { + "in": "query", + "name": "start_date", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Start Date" + } + }, + { + "in": "query", + "name": "end_date", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "End Date" + } + }, + { + "in": "query", + "name": "model", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Model" + } + }, + { + "in": "query", + "name": "api_key", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Api Key" + } + }, + { + "in": "query", + "name": "exclude_agent_ids", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Exclude Agent Ids" + } + }, + { + "in": "query", + "name": "timezone", + "required": false, + "schema": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "title": "Timezone" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ModelTopKeysResponse" + } + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Get Agent Daily Activity Model Top Keys", + "tags": [ + "agents" + ] + } + }, + "/agent/daily/activity/aggregated/search": { + "get": { + "operationId": "get_agent_daily_activity_aggregated_search_agent_daily_activity_aggregated_search_get", + "parameters": [ + { + "in": "query", + "name": "search", + "required": true, + "schema": { + "minLength": 1, + "title": "Search", + "type": "string" + } + }, + { + "in": "query", + "name": "limit", + "required": false, + "schema": { + "default": 100, + "maximum": 100, + "minimum": 1, + "title": "Limit", + "type": "integer" + } + }, + { + "in": "query", + "name": "agent_ids", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Agent Ids" + } + }, + { + "in": "query", + "name": "start_date", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Start Date" + } + }, + { + "in": "query", + "name": "end_date", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "End Date" + } + }, + { + "in": "query", + "name": "model", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Model" + } + }, + { + "in": "query", + "name": "api_key", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Api Key" + } + }, + { + "in": "query", + "name": "exclude_agent_ids", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Exclude Agent Ids" + } + }, + { + "in": "query", + "name": "timezone", + "required": false, + "schema": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "title": "Timezone" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/DailyActivityKeySearchResponse" + } + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Get Agent Daily Activity Aggregated Search", + "tags": [ + "agents" + ] + } + }, + "/agent/daily/activity/export": { + "get": { + "operationId": "get_agent_daily_activity_export_agent_daily_activity_export_get", + "parameters": [ + { + "in": "query", + "name": "export_type", + "required": false, + "schema": { + "$ref": "#/components/schemas/ExportType", + "default": "daily" + } + }, + { + "in": "query", + "name": "format", + "required": false, + "schema": { + "default": "csv", + "enum": [ + "csv", + "json" + ], + "title": "Format", + "type": "string" + } + }, + { + "in": "query", + "name": "agent_ids", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Agent Ids" + } + }, + { + "in": "query", + "name": "start_date", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Start Date" + } + }, + { + "in": "query", + "name": "end_date", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "End Date" + } + }, + { + "in": "query", + "name": "model", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Model" + } + }, + { + "in": "query", + "name": "api_key", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Api Key" + } + }, + { + "in": "query", + "name": "exclude_agent_ids", + "required": false, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Exclude Agent Ids" + } + }, + { + "in": "query", + "name": "timezone", + "required": false, + "schema": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "title": "Timezone" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "items": { + "type": "object" + }, + "type": "array" + } + }, + "text/csv": { + "schema": { + "type": "string" + } + } + }, + "description": "Streamed daily activity export" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Get Agent Daily Activity Export", + "tags": [ + "agents" + ] + } + }, "/v1/agents": { "get": { "description": "Example usage:\n```\ncurl -X GET \"http://localhost:4000/v1/agents\" -H \"Content-Type: application/json\" -H \"Authorization: Bearer your-key\" ```\n\nPass `?health_check=true` to filter out agents whose URL is unreachable:\n```\ncurl -X GET \"http://localhost:4000/v1/agents?health_check=true\" -H \"Content-Type: application/json\" -H \"Authorization: Bearer your-key\" ```\n\nPass `?query=` to get the best matching agents ranked by semantic similarity:\n```\ncurl -X GET \"http://localhost:4000/v1/agents?query=translate+a+PDF+document&top_k=5\" -H \"Content-Type: application/json\" -H \"Authorization: Bearer your-key\" ```\n\nReturns: List[AgentResponse]", diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index a471fb6f6f8..bb113715d2c 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -537,6 +537,8 @@ class LiteLLMRoutes(enum.Enum): "/lens/workers/register", "/lens/workers/{worker_id}", "/v1/traces", + "/v1/traces/query", + "/v1/traces/query/help", "/v1/traces/{trace_id}", "/v1/traces/{trace_id}/spans/{span_id}", ] @@ -699,6 +701,10 @@ class LiteLLMRoutes(enum.Enum): KeyManagementRoutes.TEAM_KEY_BULK_UPDATE.value, KeyManagementRoutes.TEAM_DAILY_ACTIVITY.value, KeyManagementRoutes.TEAM_DAILY_ACTIVITY_AGGREGATED.value, + "/team/daily/activity/aggregated/keys", + "/team/daily/activity/aggregated/search", + "/team/daily/activity/aggregated/model_top_keys", + "/team/daily/activity/export", KeyManagementRoutes.SPEND_LOGS.value, KeyManagementRoutes.SPEND_LOGS_V2.value, KeyManagementRoutes.KEY_RESET_SPEND.value, @@ -725,6 +731,11 @@ class LiteLLMRoutes(enum.Enum): "/user/list", "/user/daily/activity", "/user/daily/activity/aggregated", + "/user/daily/activity/aggregated/keys", + "/user/daily/activity/aggregated/search", + "/user/daily/activity/aggregated/model_top_keys", + "/user/daily/activity/export", + "/user/daily/activity/aggregated/cache_leakage_keys", # team "/team/new", "/team/update", @@ -742,6 +753,10 @@ class LiteLLMRoutes(enum.Enum): "/team/permissions_bulk_update", "/team/daily/activity", "/team/daily/activity/aggregated", + "/team/daily/activity/aggregated/keys", + "/team/daily/activity/aggregated/search", + "/team/daily/activity/aggregated/model_top_keys", + "/team/daily/activity/export", "/team/spend/by_user", # gateway request counts (SGR); deployment-wide, admin-only "/gateway/daily/activity", @@ -870,6 +885,11 @@ class LiteLLMRoutes(enum.Enum): # Tag usage endpoints scope internal users to tags produced by # their own keys in tag_management_endpoints.py. "/tag/daily/activity", + "/tag/daily/activity/aggregated", + "/tag/daily/activity/aggregated/keys", + "/tag/daily/activity/aggregated/search", + "/tag/daily/activity/aggregated/model_top_keys", + "/tag/daily/activity/export", "/tag/list", "/v1/models/{model_id}", "/models/{model_id}", @@ -894,6 +914,11 @@ class LiteLLMRoutes(enum.Enum): # Tag usage endpoints scope internal viewers to tags produced by # their own keys in tag_management_endpoints.py. "/tag/daily/activity", + "/tag/daily/activity/aggregated", + "/tag/daily/activity/aggregated/keys", + "/tag/daily/activity/aggregated/search", + "/tag/daily/activity/aggregated/model_top_keys", + "/tag/daily/activity/export", "/tag/list", ] ) @@ -913,6 +938,10 @@ class LiteLLMRoutes(enum.Enum): "/team/permissions_update", "/team/daily/activity", "/team/daily/activity/aggregated", + "/team/daily/activity/aggregated/keys", + "/team/daily/activity/aggregated/search", + "/team/daily/activity/aggregated/model_top_keys", + "/team/daily/activity/export", "/team/spend/by_user", "/team/{team_id}/members/me", # POST/GET the team's logging callbacks, and DELETE one of them. Every @@ -928,9 +957,19 @@ class LiteLLMRoutes(enum.Enum): "/model/delete", "/user/daily/activity", "/user/daily/activity/aggregated", + "/user/daily/activity/aggregated/keys", + "/user/daily/activity/aggregated/search", + "/user/daily/activity/aggregated/model_top_keys", + "/user/daily/activity/export", + "/user/daily/activity/aggregated/cache_leakage_keys", # Endpoint restricts results to organizations the caller is ORG_ADMIN # of; a caller who administers none gets an empty result set. "/organization/daily/activity", + "/organization/daily/activity/aggregated", + "/organization/daily/activity/aggregated/keys", + "/organization/daily/activity/aggregated/search", + "/organization/daily/activity/aggregated/model_top_keys", + "/organization/daily/activity/export", "/user/available_roles", # read-only role metadata; any authenticated user may read # Claude Code gateway: the signed-in CLI fetches its managed settings and posts its own telemetry "/claude_code_gateway/managed/settings", @@ -1009,9 +1048,24 @@ class LiteLLMRoutes(enum.Enum): "/user/available_users", "/user/available_roles", "/user/daily/activity", + "/user/daily/activity/aggregated", + "/user/daily/activity/aggregated/keys", + "/user/daily/activity/aggregated/search", + "/user/daily/activity/aggregated/model_top_keys", + "/user/daily/activity/export", + "/user/daily/activity/aggregated/cache_leakage_keys", "/team/daily/activity", "/team/daily/activity/aggregated", + "/team/daily/activity/aggregated/keys", + "/team/daily/activity/aggregated/search", + "/team/daily/activity/aggregated/model_top_keys", + "/team/daily/activity/export", "/tag/daily/activity", + "/tag/daily/activity/aggregated", + "/tag/daily/activity/aggregated/keys", + "/tag/daily/activity/aggregated/search", + "/tag/daily/activity/aggregated/model_top_keys", + "/tag/daily/activity/export", "/tag/list", "/audit", "/audit/{id}", diff --git a/litellm/proxy/agent_endpoints/endpoints.py b/litellm/proxy/agent_endpoints/endpoints.py index 875b62103b9..c0f34a24a0a 100644 --- a/litellm/proxy/agent_endpoints/endpoints.py +++ b/litellm/proxy/agent_endpoints/endpoints.py @@ -13,7 +13,7 @@ import os import uuid from collections.abc import Mapping, Sequence from types import MappingProxyType -from typing import Annotated, Final, TypedDict +from typing import Annotated, Final, NamedTuple, TypedDict from fastapi import APIRouter, Depends, HTTPException, Query, Request from pydantic import ValidationError @@ -47,7 +47,11 @@ from litellm.proxy.agent_endpoints.agent_search import ( global_agent_search_index, search_agents, ) -from litellm.proxy.agent_endpoints.auth.agent_permission_handler import accessible_agents +from litellm.proxy.agent_endpoints.auth.agent_permission_handler import ( + AgentRequestHandler, + UnrestrictedAgentAccess, + accessible_agents, +) from litellm.proxy.agent_endpoints.identity import reject_legacy_identity from litellm.proxy.agent_endpoints.identity_store import AgentIdentityStore from litellm.proxy.agent_endpoints.kill_switch import ( @@ -63,7 +67,8 @@ from litellm.proxy.agent_endpoints.managed_identity import raise_identity_failur from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.rbac_utils import check_feature_access_for_user from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity -from litellm.proxy.utils import get_custom_url +from litellm.proxy.utils import PrismaClient, get_custom_url +from litellm.repositories.chunked_in import find_many_in from litellm.types.agents import ( AgentCard, AgentConfig, @@ -303,6 +308,73 @@ async def _rank_agents_by_query( assert_never(outcome) +class _AgentDailyActivityScope(NamedTuple): + agent_ids: tuple[str, ...] | None + agent_metadata: Mapping[str, dict[str, object]] + + +async def _owned_agent_ids(*, user_id: str | None, prisma_client: PrismaClient) -> frozenset[str]: + if user_id is None: + return frozenset() + owned_records: Final = await agents_table(prisma_client).find_many(where={"created_by": user_id}) + return frozenset(agent.agent_id for agent in owned_records) + + +async def _permitted_daily_activity_agent_ids( + *, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient +) -> frozenset[str]: + access: Final = await AgentRequestHandler.resolve_agent_access(user_api_key_auth=user_api_key_dict) + if isinstance(access, UnrestrictedAgentAccess): + return await _owned_agent_ids(user_id=user_api_key_dict.user_id, prisma_client=prisma_client) + return access.agent_ids + + +async def _resolve_daily_activity_agent_ids( + *, + agent_ids: tuple[str, ...] | None, + user_api_key_dict: UserAPIKeyAuth, + prisma_client: PrismaClient, +) -> tuple[str, ...] | None: + from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view + + if _user_has_admin_view(user_api_key_dict): + return agent_ids + permitted_agent_ids: Final = await _permitted_daily_activity_agent_ids( + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + ) + return ( + tuple(agent_id for agent_id in agent_ids if agent_id in permitted_agent_ids) + if agent_ids + else tuple(permitted_agent_ids) + ) + + +async def resolve_agent_daily_activity_scope( + *, + agent_ids: tuple[str, ...] | None, + user_api_key_dict: UserAPIKeyAuth, + prisma_client: PrismaClient, +) -> _AgentDailyActivityScope: + await check_feature_access_for_user(user_api_key_dict, "agents") + + resolved_agent_ids: Final = await _resolve_daily_activity_agent_ids( + agent_ids=agent_ids, + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + ) + + agent_records: Final = ( + await agents_table(prisma_client).find_many(where={}) + if resolved_agent_ids is None + else await find_many_in(agents_table(prisma_client), "agent_id", resolved_agent_ids) + ) + agent_metadata: Final[Mapping[str, dict[str, object]]] = MappingProxyType( + {agent.agent_id: {"agent_name": agent.agent_name} for agent in agent_records} + ) + return _AgentDailyActivityScope(resolved_agent_ids, agent_metadata) + + @router.get( "/v1/agents", tags=["[beta] A2A Agents"], @@ -1288,82 +1360,39 @@ async def get_agent_daily_activity( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - agent_ids_list = agent_ids.split(",") if agent_ids else None - exclude_agent_ids_list: list[str] | None = None - if exclude_agent_ids: - exclude_agent_ids_list = exclude_agent_ids.split(",") if exclude_agent_ids else None - - # Without scoping, an empty `agent_ids` query returned every agent's - # spend/token rows on the proxy. Restrict non-admin callers to the - # agents they're permitted to invoke (or that they created), and - # intersect their explicit `agent_ids` filter with the same allowlist. - from litellm.proxy.agent_endpoints.auth.agent_permission_handler import ( - AgentRequestHandler, - RestrictedAgentAccess, - UnrestrictedAgentAccess, + requested_agent_ids: Final = tuple(agent_ids.split(",")) if agent_ids else None + exclude_agent_ids_list: Final[list[str] | None] = exclude_agent_ids.split(",") if exclude_agent_ids else None + agent_scope: Final = await resolve_agent_daily_activity_scope( + agent_ids=requested_agent_ids, + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, ) - from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view - - where_condition: Final[dict[str, object]] = {} - if not _user_has_admin_view(user_api_key_dict): - permitted_agent_ids: list[str] = [] - # An unrestricted caller is not "see everything" for activity scoping. Fall - # back to the agents the caller created so they cannot enumerate other - # tenants' agents. - # Guard against `user_id is None`: a literal None in Prisma - # `where={"created_by": None}` resolves to ``created_by IS NULL`` - # and would expose every ownerless agent's rows. - match await AgentRequestHandler.resolve_agent_access(user_api_key_auth=user_api_key_dict): - case RestrictedAgentAccess(allowed_agent_ids): - permitted_agent_ids = list(allowed_agent_ids) - case UnrestrictedAgentAccess(): - if user_api_key_dict.user_id is not None: - owned_records: Final = await agents_table(prisma_client).find_many( - where={"created_by": user_api_key_dict.user_id} - ) - permitted_agent_ids = [a.agent_id for a in owned_records] - - if agent_ids_list: - permitted_agent_id_set: Final = set(permitted_agent_ids) - agent_ids_list = [aid for aid in agent_ids_list if aid in permitted_agent_id_set] - else: - agent_ids_list = list(permitted_agent_ids) - - # No accessible agents → return an empty page without querying. - if not agent_ids_list: - return SpendAnalyticsPaginatedResponse( - results=[], - metadata=DailySpendMetadata( - total_spend=0.0, - total_prompt_tokens=0, - total_completion_tokens=0, - total_tokens=0, - total_api_requests=0, - total_successful_requests=0, - total_failed_requests=0, - total_cache_read_input_tokens=0, - total_cache_creation_input_tokens=0, - total_compression_saved_tokens=0, - page=page, - total_pages=0, - has_more=False, - ), - ) - - if agent_ids_list: - where_condition["agent_id"] = {"in": list(agent_ids_list)} - - agent_records: Final = await agents_table(prisma_client).find_many(where=where_condition) - agent_metadata: Final[Mapping[str, dict[str, object]]] = { - agent.agent_id: {"agent_name": agent.agent_name} for agent in agent_records - } + if agent_scope.agent_ids == (): + return SpendAnalyticsPaginatedResponse( + results=[], + metadata=DailySpendMetadata( + total_spend=0.0, + total_prompt_tokens=0, + total_completion_tokens=0, + total_tokens=0, + total_api_requests=0, + total_successful_requests=0, + total_failed_requests=0, + total_cache_read_input_tokens=0, + total_cache_creation_input_tokens=0, + total_compression_saved_tokens=0, + page=page, + total_pages=0, + has_more=False, + ), + ) return await get_daily_activity( prisma_client=prisma_client, table_name="litellm_dailyagentspend", entity_id_field="agent_id", - entity_id=agent_ids_list, - entity_metadata_field=agent_metadata, + entity_id=None if agent_scope.agent_ids is None else list(agent_scope.agent_ids), + entity_metadata_field=agent_scope.agent_metadata, exclude_entity_ids=exclude_agent_ids_list, start_date=start_date, end_date=end_date, diff --git a/litellm/proxy/db/prisma_client.py b/litellm/proxy/db/prisma_client.py index afaa41e97a5..e7c7102c98f 100644 --- a/litellm/proxy/db/prisma_client.py +++ b/litellm/proxy/db/prisma_client.py @@ -984,6 +984,30 @@ class PrismaManager: os.chdir(original_dir) return False + @staticmethod + def build_request_log_indexes() -> bool: + """Build the request-log indexes the migrations leave out and wait for them, for the + migration job (`--skip_server_startup`) after `setup_database` succeeds. False when + an index could not be built, so the job exits non-zero and is rerun.""" + try: + from litellm_proxy_extras.utils import ProxyExtrasDBManager + except ImportError as e: + verbose_proxy_logger.error("\x1b[1;31mLiteLLM: Failed to import proxy extras. Got %s\x1b[0m", e) + return False + return ProxyExtrasDBManager.build_request_log_indexes() + + @staticmethod + def start_request_log_index_build() -> None: + """Build the request-log indexes on a daemon thread, for a serving proxy that ran the + migrations itself (`DISABLE_SCHEMA_UPDATE` unset), so a long build never delays + readiness. A build that could not finish is logged and retried on the next boot.""" + try: + from litellm_proxy_extras.utils import ProxyExtrasDBManager + except ImportError as e: + verbose_proxy_logger.error("\x1b[1;31mLiteLLM: Failed to import proxy extras. Got %s\x1b[0m", e) + return + ProxyExtrasDBManager.start_request_log_index_build() + def should_update_prisma_schema( disable_updates: bool | str | None = None, diff --git a/litellm/proxy/lens/endpoints.py b/litellm/proxy/lens/endpoints.py index 0349c594adf..349ecb9a353 100644 --- a/litellm/proxy/lens/endpoints.py +++ b/litellm/proxy/lens/endpoints.py @@ -2,6 +2,7 @@ import hashlib import secrets from datetime import datetime, timedelta, timezone from functools import reduce +from itertools import chain from types import MappingProxyType from typing import Annotated, Final, TypeAlias from uuid import uuid4 @@ -35,7 +36,7 @@ from litellm.proxy.lens.models import ( WorkerCreated, ) from litellm.proxy.lens.repository import LensRepository, WriterDatabase -from litellm.proxy.lens.sources import SourceReader, Storage, parse_execution +from litellm.proxy.lens.sources import ActivityAvailability, SourceReader, Storage, parse_execution from litellm.proxy.lens.state import ( can_access, claim_job, @@ -168,6 +169,18 @@ async def create_lens(settings: LensSettings, auth: Auth) -> Lens: return await repository().create(queue_job(lens, now, str(uuid4()))) +@router.get("/activity/available", response_model=ActivityAvailability) +async def activity_available(auth: Auth, storage: StorageDep) -> ActivityAvailability: + scope: Final = user_scope(auth) + return await source_reader(storage).availability(scope) if storage is not None else ActivityAvailability() + + +@router.get("/agents", response_model=tuple[str, ...]) +async def list_agents(auth: Auth, storage: StorageDep) -> tuple[str, ...]: + scope: Final = user_scope(auth) + return await source_reader(storage).agents(scope) if storage is not None else () + + @router.put("/{lens_id}", response_model=Lens) async def update_lens(lens_id: str, settings: LensSettings, auth: Auth) -> Lens: await get_lens(lens_id, user_scope(auth, write=True)) @@ -264,7 +277,7 @@ class Preview(BaseModel): as_of: AwareDatetime | None = None offset: int = Field(default=0, ge=0) settings: LensSettings - lookback_hours: int = Field(default=24, ge=1, le=720) + lookback_hours: int = Field(default=24, ge=1, le=8760) @router.post("/preview/sample", response_model=Sample) @@ -326,6 +339,9 @@ async def revoke_worker(worker_id: str, auth: Auth) -> bool: worker: Final = next((w for w in await repository().workers() if w.id == worker_id), None) if worker is None or not can_access(scope, worker.scope): raise HTTPException(404, "Worker not found") + jobs: Final = chain.from_iterable(lens.jobs for lens in await repository().lenses()) + if any(job.status == "running" and job.worker_id == worker.id for job in jobs): + raise HTTPException(409, "Wait for this worker's investigation to finish or cancel it before revoking access") await repository().revoke_worker(worker.id) return True diff --git a/litellm/proxy/lens/models.py b/litellm/proxy/lens/models.py index eb88801d065..91f0ad582bf 100644 --- a/litellm/proxy/lens/models.py +++ b/litellm/proxy/lens/models.py @@ -29,8 +29,9 @@ class LensSettings(Record): name: str = Field(min_length=1, max_length=100) context: str = Field(default="", max_length=6000) source: Literal["traces", "requests", "both"] = "traces" - lookback_hours: int = Field(default=24, ge=1, le=720) + lookback_hours: int = Field(default=24, ge=1, le=8760) service: str = Field(default="", max_length=200) + agent_name: str = Field(default="", max_length=200) filters: tuple[MetadataFilter, ...] = Field(default=(), max_length=8) checks: tuple[Check, ...] = () model: str = Field(min_length=1, max_length=200) @@ -41,7 +42,7 @@ class LensSettings(Record): concurrency: int = Field(default=8, ge=1) team_id: str = "" execution_ids: tuple[str, ...] = () - monthly_budget: float = Field(default=20, gt=0, le=100000, allow_inf_nan=False) + monthly_budget: float = Field(default=100, gt=0, le=100000, allow_inf_nan=False) @model_validator(mode="after") def unique_checks(self) -> "LensSettings": @@ -214,7 +215,7 @@ class LensList(Record): class RunRequest(Record): settings: LensSettings | None = None - lookback_hours: int | None = Field(default=None, ge=1, le=720) + lookback_hours: int | None = Field(default=None, ge=1, le=8760) class FindingUpdate(Record): diff --git a/litellm/proxy/lens/sources.py b/litellm/proxy/lens/sources.py index 12d26cd4974..17550cc4aab 100644 --- a/litellm/proxy/lens/sources.py +++ b/litellm/proxy/lens/sources.py @@ -18,7 +18,14 @@ from litellm.proxy.lens.models import ( ) +class ActivityAvailability(BaseModel): + traces: bool = False + requests: bool = False + + class Storage(Protocol): + def lens_availability(self, parameters: Mapping[str, object]) -> Awaitable[object]: ... + def lens_agents(self, parameters: Mapping[str, object]) -> Awaitable[object]: ... def lens_sample(self, parameters: Mapping[str, object]) -> Awaitable[object]: ... def lens_content(self, parameters: Mapping[str, object]) -> Awaitable[object]: ... def lens_evidence(self, parameters: Mapping[str, object]) -> Awaitable[object]: ... @@ -53,6 +60,12 @@ class CountRow(BaseModel): count: int +class AgentRow(BaseModel): + agent_name: str + + +_AVAILABILITY: Final = TypeAdapter(tuple[ActivityAvailability, ...]) +_AGENTS: Final = TypeAdapter(tuple[AgentRow, ...]) _ROWS: Final = TypeAdapter(tuple[ExecutionRow, ...]) _PARTS: Final = TypeAdapter(tuple[PartRow, ...]) _COUNTS: Final = TypeAdapter(tuple[CountRow, ...]) @@ -90,6 +103,14 @@ class SourceReader: def __init__(self, storage: Storage) -> None: self.storage: Final = storage + async def availability(self, scope: Scope) -> ActivityAvailability: + rows: Final = _AVAILABILITY.validate_python(await self.storage.lens_availability(parameters(scope, ()))) + return rows[0] if rows else ActivityAvailability() + + async def agents(self, scope: Scope) -> tuple[str, ...]: + rows: Final = _AGENTS.validate_python(await self.storage.lens_agents(parameters(scope, ()))) + return tuple(row.agent_name for row in rows) + async def sample( self, scope: Scope, @@ -108,6 +129,7 @@ class SourceReader: "start": start, "end": end, "service": settings.service, + "agent_name": settings.agent_name, "limit": page_size, "offset": offset, "after": cursor, diff --git a/litellm/proxy/lens/worker.py b/litellm/proxy/lens/worker.py index 2980f62deed..62f8295e7d3 100644 --- a/litellm/proxy/lens/worker.py +++ b/litellm/proxy/lens/worker.py @@ -15,6 +15,40 @@ from .models import Claim, Coverage, ExecutionContent, ModelRequest, ModelResult logger: Final = logging.getLogger("litellm.lens.worker") +def failure_message(error: Exception) -> str: + if isinstance(error, (OSError, sqlite3.Error)): + return "Worker temporary storage failed. Increase its capacity or reduce analysis parallelism." + if isinstance(error, httpx.TimeoutException): + return "The worker timed out waiting for the proxy. Check proxy availability and model response times." + if isinstance(error, httpx.TransportError): + return "The worker could not connect to the proxy. Check the proxy URL, network access, and TLS configuration." + if isinstance(error, httpx.HTTPStatusError): + path: Final = error.request.url.path + action: Final = ( + "Model request" + if path.endswith("/model") + else "Reading trace data" + if path.endswith(("/sample", "/content")) + else "Saving results" + if path.endswith("/result") + else "Worker request" + ) + status: Final = error.response.status_code + guidance: Final = MappingProxyType( + { + 400: "Check the configured model and whether the worker's billing key is enabled.", + 401: "Check the worker credential and its assigned billing key.", + 402: "Check the investigation's monthly limit and the worker key's remaining budget.", + 403: "Check the worker key's model permissions and access restrictions.", + 404: "Check that the proxy and worker versions match and the requested model is configured.", + 409: "This worker no longer owns the run. Check whether it was cancelled or claimed again.", + 429: "The request was rate limited. Retry later or check the worker key's rate limits.", + } + ).get(status, "Check proxy and model availability, then retry the investigation.") + return f"{action} failed (HTTP {status}). {guidance}" + return "The worker could not read an analysis response. Check structured JSON support and matching proxy/worker versions." + + class LensWorker: def __init__(self, client: httpx.AsyncClient, sleep: Callable[[float], Awaitable[None]] = asyncio.sleep) -> None: self.client: Final = client @@ -82,14 +116,7 @@ class LensWorker: saved: Final = await self.client.post(prefix + "/result", json=result.model_dump(mode="json")) saved.raise_for_status() except (httpx.HTTPError, ValueError, OSError, sqlite3.Error) as exc: - status: Final = exc.response.status_code if isinstance(exc, httpx.HTTPStatusError) else None - message: Final = ( - "Worker temporary storage failed. Increase its capacity or reduce analysis parallelism." - if isinstance(exc, (OSError, sqlite3.Error)) - else "Monthly budget reached" - if status == 402 - else "Analysis interrupted. Check worker connectivity, model configuration, and trace storage." - ) + message: Final = failure_message(exc) logger.warning("Analysis %s interrupted (%s)", claim.job.id, type(exc).__name__) failed: Final = await self.client.post( prefix + "/result", json=Result(coverage=Coverage(), error=message).model_dump() diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index f3564fcbc56..27b0960c823 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -1,13 +1,15 @@ import asyncio from collections.abc import Awaitable, Callable, Mapping, Sequence from collections.abc import Set as AbstractSet -from datetime import datetime, timedelta, timezone -from types import MappingProxyType, SimpleNamespace -from typing import TYPE_CHECKING, Final, Protocol +from dataclasses import dataclass, replace +from datetime import date, datetime, timedelta +from types import MappingProxyType +from typing import Final, Literal, NoReturn, Protocol from fastapi import HTTPException, status -from typing_extensions import ReadOnly, TypedDict +from typing_extensions import ReadOnly, TypedDict, assert_never +from litellm import constants from litellm._logging import verbose_proxy_logger from litellm.constants import PTU_SENTINEL_API_KEY from litellm.proxy._types import CommonProxyErrors @@ -20,11 +22,7 @@ from litellm.proxy.spend_tracking.key_metadata_recovery import ( ) from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled from litellm.proxy.utils import PrismaClient -from litellm.repositories.prisma_protocols import TableActions -from litellm.repositories.table_repositories import DeletedVerificationTokenRepository -from litellm.repositories.verification_token_repository import ( - VerificationTokenRepository, -) +from litellm.repositories.daily_activity_repository import DailyActivityRepository from litellm.types.proxy.management_endpoints.common_daily_activity import ( BreakdownMetrics, DailySpendData, @@ -36,24 +34,63 @@ from litellm.types.proxy.management_endpoints.common_daily_activity import ( SpendAnalyticsPaginatedResponse, SpendMetrics, ) +from litellm.types.repositories.daily_activity import ( + DailyActivityScope, + DailyActivityTable, + EntityRollupRow, + GroupingSetsRow, + KeyMetadataRow, + RollupMetricsRow, + SpendLogsWindow, +) -if TYPE_CHECKING: - from prisma.models import ( - LiteLLM_DeletedVerificationToken as PrismaDeletedVerificationToken, - ) - from prisma.models import ( - LiteLLM_VerificationToken as PrismaVerificationToken, - ) -# Mapping from Prisma accessor names to actual PostgreSQL table names. -_PRISMA_TO_PG_TABLE: Final[Mapping[str, str]] = { - "litellm_dailyuserspend": "LiteLLM_DailyUserSpend", - "litellm_dailyteamspend": "LiteLLM_DailyTeamSpend", - "litellm_dailyorganizationspend": "LiteLLM_DailyOrganizationSpend", - "litellm_dailyenduserspend": "LiteLLM_DailyEndUserSpend", - "litellm_dailyagentspend": "LiteLLM_DailyAgentSpend", - "litellm_dailytagspend": "LiteLLM_DailyTagSpend", -} +@dataclass(frozen=True, slots=True) +class ScopeDenied: + status_code: Literal[403, 404] + reason: str + + +@dataclass(frozen=True, slots=True) +class InvalidDateRange: + reason: str + + +def raise_public(error: ScopeDenied | InvalidDateRange) -> NoReturn: + match error: + case ScopeDenied(): + raise HTTPException(status_code=error.status_code, detail={"error": error.reason}) + case InvalidDateRange(): + raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail={"error": error.reason}) + case _: + assert_never(error) + + +@dataclass(frozen=True, slots=True) +class CanonicalDateRange: + start: date + end: date + + +def parse_canonical_date(value: str) -> date | None: + """The daily spend tables store ``date`` as text and compare it against the raw request + string, so only the exact ``YYYY-MM-DD`` spelling can match a row. Spellings the parser + would normalise (``2026-9-24``, ``20260924``, full-width digits) are rejected instead.""" + try: + parsed: Final = date.fromisoformat(value) + except ValueError: + return None + return parsed if parsed.isoformat() == value else None + + +def parse_canonical_date_range(start_date: str | None, end_date: str | None) -> CanonicalDateRange | InvalidDateRange: + if start_date is None or end_date is None: + return InvalidDateRange(reason="Please provide start_date and end_date") + start: Final = parse_canonical_date(start_date) + end: Final = parse_canonical_date(end_date) + if start is None or end is None: + return InvalidDateRange(reason="start_date and end_date must be valid YYYY-MM-DD dates") + return CanonicalDateRange(start=start, end=end) class DailySpendRecord(Protocol): @@ -132,57 +169,23 @@ class _KeyMetadataDict(TypedDict, total=False): key_exists: ReadOnly[bool] -def _key_metadata(api_key_metadata: Mapping[str, _KeyMetadataDict], api_key: str) -> KeyMetadata: - meta: Final = api_key_metadata.get(api_key, {}) +class _AggregatedSpendData(TypedDict): + results: ReadOnly[list[DailySpendData]] + totals: ReadOnly[SpendMetrics] + + +def _key_metadata(api_key_metadata: Mapping[str, KeyMetadataRow], api_key: str) -> KeyMetadata: + meta: Final = api_key_metadata.get(api_key) return KeyMetadata( - key_alias=meta.get("key_alias"), - team_id=meta.get("team_id"), - user_id=meta.get("user_id"), - user_email=meta.get("user_email"), - key_exists=meta.get("key_exists", False), + key_alias=meta.key_alias if meta is not None else None, + team_id=meta.team_id if meta is not None else None, + user_id=meta.user_id if meta is not None else None, + user_email=meta.user_email if meta is not None else None, + key_exists=meta.key_exists if meta is not None else False, ) -_WhereValue = str | dict[str, object] - - -class _AggregatedSpendData(TypedDict): - results: list[DailySpendData] - totals: SpendMetrics - - -class _GroupingSetsRow(SimpleNamespace): - date: str - api_key: str | None - model: str | None - model_group: str | None - custom_llm_provider: str | None - mcp_namespaced_tool_name: str | None - endpoint: str | None - group_level: int - spend: float | None - prompt_tokens: int | None - completion_tokens: int | None - cache_read_input_tokens: int | None - cache_creation_input_tokens: int | None - compression_saved_tokens: int | None - compression_savings_spend: float | None - prompt_caching_savings_spend: float | None - gateway_injected_caching_savings_spend: float | None - autorouter_savings_spend: float | None - api_requests: int | None - successful_requests: int | None - failed_requests: int | None - total_response_time_ms: int | None - timed_requests: int | None - - -class _EntityRollupRow(_GroupingSetsRow): - entity_id: str | None - api_key_rolled: int - - -def _reported_flat_cost(record: DailySpendRecord | _GroupingSetsRow) -> float: +def _reported_flat_cost(record: DailySpendRecord | RollupMetricsRow) -> float: """Flat cost a daily row reports, which is zero unless PTU cost attribution is enabled. Both read paths funnel through here: the paginated path reads the ``ptu_flat_cost`` @@ -223,9 +226,7 @@ def update_metrics(existing_metrics: SpendMetrics, record: DailySpendRecord) -> existing_metrics.compression_saved_tokens += record.compression_saved_tokens or 0 existing_metrics.compression_savings_spend += record.compression_savings_spend or 0 existing_metrics.prompt_caching_savings_spend += record.prompt_caching_savings_spend or 0 - existing_metrics.gateway_injected_caching_savings_spend += ( # rebind-ok: this accumulator mutates its target in place for every metric on the row - record.gateway_injected_caching_savings_spend or 0 - ) + existing_metrics.gateway_injected_caching_savings_spend += record.gateway_injected_caching_savings_spend or 0 existing_metrics.autorouter_savings_spend += record.autorouter_savings_spend or 0 existing_metrics.api_requests += record.api_requests or 0 existing_metrics.successful_requests += record.successful_requests or 0 @@ -255,7 +256,7 @@ def compute_tag_metadata_totals(records: Sequence[DailySpendRecord]) -> SpendMet if not request_id: continue - tag_value = getattr(record, "tag", None) + tag_value: str | None = getattr(record, "tag", None) if _is_user_agent_tag(tag_value): continue @@ -283,7 +284,7 @@ def update_breakdown_metrics( record: DailySpendRecord, model_metadata: Mapping[str, dict[str, object]], provider_metadata: Mapping[str, dict[str, object]], - api_key_metadata: Mapping[str, _KeyMetadataDict], + api_key_metadata: Mapping[str, KeyMetadataRow], entity_id_field: str | None = None, entity_metadata_field: Mapping[str, dict[str, object]] | None = None, ) -> BreakdownMetrics: @@ -426,8 +427,7 @@ def update_breakdown_metrics( # Update entity-specific metrics if entity_id_field is provided if entity_id_field: - entity_value = getattr(record, entity_id_field, None) - entity_value = entity_value if entity_value else "Unassigned" # allow for null entity_id_field + entity_value: Final[str] = getattr(record, entity_id_field, None) or "Unassigned" if entity_value not in breakdown.entities: breakdown.entities[entity_value] = MetricWithMetadata( metrics=SpendMetrics(), @@ -450,7 +450,7 @@ def update_breakdown_metrics( return breakdown -def _spend_logs_window(dates: AbstractSet[str | None]) -> tuple[datetime, datetime] | None: +def spend_logs_window(dates: AbstractSet[str | None]) -> tuple[datetime, datetime] | None: parsed: Final = sorted(day for day in (_parse_spend_date(raw) for raw in dates) if day is not None) if not parsed: return None @@ -466,9 +466,6 @@ def _parse_spend_date(raw: str | None) -> datetime | None: return None -_EMPTY_KEY_METADATA: Final[Mapping[str, _KeyMetadataDict]] = MappingProxyType({}) - - def _metadata_with_recovered_owner( metadata: Mapping[str, _KeyMetadataDict], key: str, @@ -480,432 +477,135 @@ def _metadata_with_recovered_owner( return {**current, "user_id": owner} -async def get_api_key_metadata( - prisma_client: PrismaClient, - api_keys: AbstractSet[str], - spend_logs_window: tuple[datetime, datetime] | None = None, -) -> Mapping[str, _KeyMetadataDict]: - """Get api key metadata, falling back to deleted keys table for keys not found in active table. +@dataclass(frozen=True, slots=True) +class _ProxyDailyActivityReads: + prisma_client: PrismaClient - This ensures that key_alias and team_id are preserved in historical activity logs - even after a key is deleted or regenerated. Also recovers aliases for api_key - values that were double-hashed by the v1.99 spend-log provenance gate. - """ - key_records: Sequence[PrismaVerificationToken] = await VerificationTokenRepository(prisma_client).table.find_many( - where={"token": {"in": list(api_keys)}} - ) - result: Final[dict[str, _KeyMetadataDict]] = { - k.token: { - "key_alias": k.key_alias, - "team_id": k.team_id, - "user_id": getattr(k, "user_id", None), - "key_exists": True, + async def recover_key_metadata( + self, resolved: Mapping[str, KeyMetadataRow], api_keys: frozenset[str], window: SpendLogsWindow | None + ) -> Mapping[str, KeyMetadataRow]: + result: Final[dict[str, _KeyMetadataDict]] = { + key: { + "key_alias": row.key_alias, + "team_id": row.team_id, + "user_id": row.user_id, + "user_email": row.user_email, + "key_exists": row.key_exists, + } + for key, row in resolved.items() } - for k in key_records - } - - # For any keys not found in the active table, check the deleted keys table - missing_keys: Final = api_keys - set(result.keys()) - if missing_keys: - try: - deleted_key_records: Final[ - Sequence[PrismaDeletedVerificationToken] - ] = await DeletedVerificationTokenRepository(prisma_client).table.find_many( - where={"token": {"in": list(missing_keys)}}, - order={"deleted_at": "desc"}, - ) - # Use the most recent deleted record for each token (ordered by deleted_at desc) - for k in deleted_key_records: - if k.token not in result: - result[k.token] = { - "key_alias": k.key_alias, - "team_id": k.team_id, - "user_id": getattr(k, "user_id", None), - } - except Exception as e: - verbose_proxy_logger.warning( - "Failed to fetch deleted key metadata for %d missing keys: %s", - len(missing_keys), - e, - ) - - from_session_keys: Final = await recover_cli_session_key_metadata(prisma_client, api_keys - frozenset(result)) - still_missing: Final = api_keys - frozenset(result) - frozenset(from_session_keys) - from_reverse_hash: Final = ( - await recover_double_hashed_key_metadata(prisma_client, still_missing) if still_missing else _EMPTY_KEY_METADATA - ) - after_token_recovery: Final = MappingProxyType({**result, **from_session_keys, **from_reverse_hash}) - unresolved: Final = api_keys - frozenset(after_token_recovery) - from_spend_logs: Final = ( - await recover_key_metadata_from_spend_logs(prisma_client, unresolved, spend_logs_window) - if unresolved and spend_logs_window is not None - else _EMPTY_KEY_METADATA - ) - combined: Final = MappingProxyType({**after_token_recovery, **from_spend_logs}) - ownerless: Final = frozenset( - key - for key in api_keys - if not combined.get(key, {}).get("user_id") and not combined.get(key, {}).get("key_exists") - ) - owners: Final = await recover_key_owner_from_daily_spend(prisma_client, ownerless) - metadata_with_owners: Final[Mapping[str, _KeyMetadataDict]] = MappingProxyType( - { - **combined, - **{key: _metadata_with_recovered_owner(combined, key, owner) for key, owner in owners.items()}, - } - ) - return await attach_user_details(prisma_client, metadata_with_owners) + from_session_keys: Final = await recover_cli_session_key_metadata( + self.prisma_client, api_keys - frozenset(result) + ) + still_missing: Final = api_keys - frozenset(result) - frozenset(from_session_keys) + from_reverse_hash: Final = ( + await recover_double_hashed_key_metadata(self.prisma_client, still_missing) + if still_missing + else MappingProxyType({}) + ) + after_token_recovery: Final = MappingProxyType({**result, **from_session_keys, **from_reverse_hash}) + unresolved: Final = api_keys - frozenset(after_token_recovery) + from_spend_logs: Final = ( + await recover_key_metadata_from_spend_logs(self.prisma_client, unresolved, window) + if unresolved and window is not None + else MappingProxyType({}) + ) + combined: Final = MappingProxyType({**after_token_recovery, **from_spend_logs}) + ownerless: Final = frozenset( + key + for key in api_keys + if not combined.get(key, {}).get("user_id") and not combined.get(key, {}).get("key_exists") + ) + owners: Final = await recover_key_owner_from_daily_spend(self.prisma_client, ownerless) + with_owners: Final[Mapping[str, _KeyMetadataDict]] = MappingProxyType( + { + **combined, + **{key: _metadata_with_recovered_owner(combined, key, owner) for key, owner in owners.items()}, + } + ) + attached: Final = await attach_user_details(self.prisma_client, with_owners) + return MappingProxyType( + { + key: replace( + resolved[key], + key_alias=value.get("key_alias"), + team_id=value.get("team_id"), + user_id=value.get("user_id"), + user_email=value.get("user_email"), + key_exists=value.get("key_exists", False), + ) + if key in resolved + else KeyMetadataRow( + api_key=key, + key_alias=value.get("key_alias"), + team_id=value.get("team_id"), + user_id=value.get("user_id"), + user_email=value.get("user_email"), + key_exists=value.get("key_exists", False), + tags=(), + ) + for key, value in attached.items() + } + ) -def _adjust_dates_for_timezone( +def daily_activity_repository(prisma_client: PrismaClient) -> DailyActivityRepository: + return DailyActivityRepository(prisma_client, proxy_reads=_ProxyDailyActivityReads(prisma_client)) + + +def daily_activity_scope( + table: str, + entity_id_field: str, + entity_id: str | list[str] | None, + exclude_entity_ids: list[str] | None, + api_key: str | list[str] | None, start_date: str, end_date: str, + model: str | None, timezone_offset_minutes: int | None, include_current_utc_day: bool = False, - utc_now: datetime | None = None, -) -> tuple[str, str]: - """ - Map a caller-local date range onto UTC bucket keys, extending only the live end. - - The aggregation table (e.g. LiteLLM_DailyUserSpend) stores spend in whole-UTC-day - buckets keyed on date as YYYY-MM-DD. Any conversion of an interior local-day - boundary using only date arithmetic must round to whole UTC days, allowing up to - 24h of slop at each boundary. A previous implementation expanded the SQL range by - an extra full UTC day on whichever side the offset pointed, which pulled in 24h of - unrelated bucket data per boundary and produced approximately 100% over-counting on - single-day queries (e.g. IST May 29 returning UTC May 28 + UTC May 29 in full). - Sums of single-day queries then exceeded the equivalent multi-day aggregate, which - is mathematically impossible. Historical dates therefore stay a pass-through: the - local date is the UTC bucket key, trading boundary slop for monotonic, additive - results. Hour-level buckets or pro-rata weighting would fix that properly; both - require data the current schema does not store. - - The end boundary is different when the range reaches the caller's current day. A - caller west of UTC asking for a range ending "today" is asking for data up to now, - but once UTC has rolled past their local midnight, everything they sent since then - sits in the next UTC bucket, which the pass-through excludes: a PT dashboard goes - stale every evening from 5pm until local midnight, showing $0 for anything that - only started accruing that evening. Extending such a range to today's UTC bucket - cannot over-count, because the only part of that bucket outside the caller's range - is the future, and the future is empty. ``timezone_offset_minutes`` follows the - JS ``Date.getTimezoneOffset`` convention: UTC minus local, positive west of UTC. - - The extension is strictly opt-in via ``include_current_utc_day`` so a consumer - whose axis or reconciliation expects the range to stop at the requested end date - keeps today's byte-for-byte behaviour; the cost optimization dashboard opts in. - """ - if not include_current_utc_day or timezone_offset_minutes is None: - return start_date, end_date - now: Final = utc_now if utc_now is not None else datetime.now(timezone.utc) - caller_local_today: Final = (now - timedelta(minutes=timezone_offset_minutes)).date().isoformat() - if end_date < caller_local_today: - return start_date, end_date - return start_date, max(end_date, now.date().isoformat()) - - -def _build_where_conditions( - *, - entity_id_field: str, - entity_id: str | list[str] | None, - start_date: str, - end_date: str, - model: str | None, - api_key: str | list[str] | None, - exclude_entity_ids: list[str] | None = None, - timezone_offset_minutes: int | None = None, - include_current_utc_day: bool = False, -) -> dict[str, "_WhereValue"]: - """Build prisma where clause for daily activity queries.""" - # Adjust dates for timezone if provided - adjusted_start, adjusted_end = _adjust_dates_for_timezone( - start_date, end_date, timezone_offset_minutes, include_current_utc_day +) -> DailyActivityScope: + table_value: Final = DailyActivityTable(table) + entity_ids: tuple[str, ...] | None = ( + (entity_id,) if isinstance(entity_id, str) else tuple(entity_id) if entity_id is not None else None + ) + api_keys: tuple[str, ...] | None = ( + None if api_key in (None, "") else (api_key,) if isinstance(api_key, str) else tuple(api_key) + ) + return DailyActivityScope( + table=table_value, + entity_id_field=entity_id_field, + entity_ids=entity_ids, + exclude_entity_ids=tuple(exclude_entity_ids or ()), + api_keys=api_keys, + start_date=start_date, + end_date=end_date, + model=model, + timezone_offset_minutes=timezone_offset_minutes, + include_current_utc_day=include_current_utc_day, ) - where_conditions: Final[dict[str, _WhereValue]] = { - "date": { - "gte": adjusted_start, - "lte": adjusted_end, + +async def get_api_key_metadata( + prisma_client: PrismaClient, api_keys: AbstractSet[str], spend_logs_window: SpendLogsWindow | None = None +) -> Mapping[str, _KeyMetadataDict]: + rows: Final = await daily_activity_repository(prisma_client).key_metadata(frozenset(api_keys), spend_logs_window) + return { + key: { + "key_alias": value.key_alias, + "team_id": value.team_id, + "user_id": value.user_id, + "user_email": value.user_email, + "key_exists": value.key_exists, } + for key, value in rows.items() } - if model: - where_conditions["model"] = model - if api_key: - if isinstance(api_key, list): - where_conditions["api_key"] = {"in": api_key} - else: - where_conditions["api_key"] = api_key - - if entity_id is not None: - if isinstance(entity_id, list): - where_conditions[entity_id_field] = {"in": entity_id} - else: - where_conditions[entity_id_field] = {"equals": entity_id} - - if exclude_entity_ids: - current: _WhereValue = where_conditions.get(entity_id_field, {}) - if isinstance(current, str): - current = {"equals": current} - current["not"] = {"in": exclude_entity_ids} - where_conditions[entity_id_field] = current - - return where_conditions - - -def _build_aggregated_where_clause( - *, - entity_id_field: str, - entity_id: str | list[str] | None, - adjusted_start: str, - adjusted_end: str, - model: str | None, - api_key: str | list[str] | None, # mutable-ok: filter union shared with the paginated path - exclude_entity_ids: list[str] | None, # mutable-ok: filter union shared with the paginated path -) -> tuple[str, list[str]]: - """Build the WHERE clause and $N params shared by the aggregated queries.""" - sql_conditions: Final[list[str]] = [] - sql_params: Final[list[str]] = [] - p = 1 # parameter index (1-based for PostgreSQL $N placeholders) - - # Date range (always present) - sql_conditions.append(f"date >= ${p}") - sql_params.append(adjusted_start) - p += 1 - - sql_conditions.append(f"date <= ${p}") - sql_params.append(adjusted_end) - p += 1 - - # Optional entity filter; an empty list must match nothing, not everything - if entity_id is not None: - if isinstance(entity_id, list): - if entity_id: - placeholders = ", ".join(f"${p + i}" for i in range(len(entity_id))) - sql_conditions.append(f'"{entity_id_field}" IN ({placeholders})') - sql_params.extend(entity_id) - p += len(entity_id) - else: - sql_conditions.append("FALSE") - else: - sql_conditions.append(f'"{entity_id_field}" = ${p}') - sql_params.append(entity_id) - p += 1 - - # Exclude specific entities - if exclude_entity_ids: - placeholders = ", ".join(f"${p + i}" for i in range(len(exclude_entity_ids))) - sql_conditions.append(f'"{entity_id_field}" NOT IN ({placeholders})') - sql_params.extend(exclude_entity_ids) - p += len(exclude_entity_ids) - - # Optional model filter - if model: - sql_conditions.append(f"model = ${p}") - sql_params.append(model) - p += 1 - - # Optional api_key filter; an empty list must match nothing, not everything - if isinstance(api_key, list): - if api_key: - placeholders = ", ".join(f"${p + i}" for i in range(len(api_key))) - sql_conditions.append(f"api_key IN ({placeholders})") - sql_params.extend(api_key) - p += len(api_key) - else: - sql_conditions.append("FALSE") - elif api_key: - sql_conditions.append(f"api_key = ${p}") - sql_params.append(api_key) - p += 1 - - return " AND ".join(sql_conditions), sql_params - - -def _ptu_flat_cost_select(table_name: str) -> str: - """Only LiteLLM_DailyTeamSpend carries ptu_flat_cost; other daily tables emit a - constant zero so the SpendMetrics.flat_cost response shape stays uniform.""" - if table_name == "litellm_dailyteamspend": - return "SUM(ptu_flat_cost)::float AS ptu_flat_cost" - return "0::float AS ptu_flat_cost" - - -def _build_aggregated_sql_query( - *, - table_name: str, - entity_id_field: str, - entity_id: str | list[str] | None, # mutable-ok: filter union shared with the paginated path - start_date: str, - end_date: str, - model: str | None, - api_key: str | list[str] | None, # mutable-ok: filter union shared with the paginated path - exclude_entity_ids: list[str] | None = None, # mutable-ok: filter union shared with the paginated path - timezone_offset_minutes: int | None = None, - include_current_utc_day: bool = False, -) -> tuple[str, list[str]]: # mutable-ok: SQL text plus its ordered $N params - """Build a parameterized SQL GROUP BY query for aggregated daily activity. - - Groups by (date, api_key, model, model_group, custom_llm_provider, - mcp_namespaced_tool_name, endpoint) with SUMs on all metric columns. - The entity_id column is intentionally omitted from GROUP BY to collapse - rows across entities — this is where the biggest row reduction comes from. - - Returns: - Tuple of (sql_query, params_list) ready for prisma_client.db.query_raw(). - """ - pg_table: Final = _PRISMA_TO_PG_TABLE.get(table_name) - if pg_table is None: - raise ValueError(f"Unknown table name: {table_name}") - - adjusted_start, adjusted_end = _adjust_dates_for_timezone( - start_date, end_date, timezone_offset_minutes, include_current_utc_day - ) - - where_clause, sql_params = _build_aggregated_where_clause( - entity_id_field=entity_id_field, - entity_id=entity_id, - adjusted_start=adjusted_start, - adjusted_end=adjusted_end, - model=model, - api_key=api_key, - exclude_entity_ids=exclude_entity_ids, - ) - - # Postgres computes every rollup level the response needs — per-date - # totals, per-(date, model), per-(date, model, api_key), per-provider, - # etc. — in a single pass via GROUPING SETS. The GROUPING() bitmask - # encodes which level a row belongs to so Python can dispatch rows - # straight into their buckets without re-summing. The leaf grouping - # is omitted on purpose: nothing in the response shape needs it once - # all the rollups are present. - # - # TODO: drop the successful_requests/failed_requests aggregates (and the - # total_successful_requests metadata they feed) once the admin UI reads SGR - # only from LiteLLM_DailyGatewayRequests. The remaining spend, token and - # api_requests rollups are still served from here. - sql_query: Final = f""" - SELECT - date, - api_key, - model, - COALESCE(NULLIF(model_group, ''), model) AS model_group, - custom_llm_provider, - mcp_namespaced_tool_name, - endpoint, - GROUPING(date, api_key, model, COALESCE(NULLIF(model_group, ''), model), - custom_llm_provider, mcp_namespaced_tool_name, - endpoint) AS group_level, - SUM(spend)::float AS spend, - {_ptu_flat_cost_select(table_name)}, - SUM(prompt_tokens)::bigint AS prompt_tokens, - SUM(completion_tokens)::bigint AS completion_tokens, - SUM(cache_read_input_tokens)::bigint AS cache_read_input_tokens, - SUM(cache_creation_input_tokens)::bigint AS cache_creation_input_tokens, - SUM(compression_saved_tokens)::bigint AS compression_saved_tokens, - SUM(compression_savings_spend)::float AS compression_savings_spend, - SUM(prompt_caching_savings_spend)::float AS prompt_caching_savings_spend, - SUM(gateway_injected_caching_savings_spend)::float AS gateway_injected_caching_savings_spend, - SUM(autorouter_savings_spend)::float AS autorouter_savings_spend, - SUM(api_requests)::bigint AS api_requests, - SUM(successful_requests)::bigint AS successful_requests, - SUM(failed_requests)::bigint AS failed_requests, - SUM(total_response_time_ms)::bigint AS total_response_time_ms, - SUM(timed_requests)::bigint AS timed_requests - FROM "{pg_table}" - WHERE {where_clause} - GROUP BY GROUPING SETS ( - (date), - (date, api_key), - (date, model), - (date, model, api_key), - (date, COALESCE(NULLIF(model_group, ''), model)), - (date, COALESCE(NULLIF(model_group, ''), model), api_key), - (date, custom_llm_provider), - (date, custom_llm_provider, api_key), - (date, mcp_namespaced_tool_name), - (date, mcp_namespaced_tool_name, api_key), - (date, endpoint), - (date, endpoint, api_key), - () - ) - """ - - return sql_query, sql_params - - -def _build_entity_rollup_sql_query( - *, - table_name: str, - entity_id_field: str, - entity_id: str | list[str] | None, # mutable-ok: filter union shared with the paginated path - start_date: str, - end_date: str, - model: str | None, - api_key: str | list[str] | None, # mutable-ok: filter union shared with the paginated path - exclude_entity_ids: list[str] | None = None, # mutable-ok: filter union shared with the paginated path - timezone_offset_minutes: int | None = None, - include_current_utc_day: bool = False, -) -> tuple[str, list[str]]: # mutable-ok: SQL text plus its ordered $N params - """Per-entity companion to _build_aggregated_sql_query. - - Two rollup levels over the same WHERE clause — (date, entity) and - (date, entity, api_key) — told apart by GROUPING(api_key): 1 when the - api_key column is rolled up, 0 when it is part of the key. - """ - pg_table: Final = _PRISMA_TO_PG_TABLE.get(table_name) - if pg_table is None: - raise ValueError(f"Unknown table name: {table_name}") - - adjusted_start, adjusted_end = _adjust_dates_for_timezone( - start_date, end_date, timezone_offset_minutes, include_current_utc_day - ) - - where_clause, sql_params = _build_aggregated_where_clause( - entity_id_field=entity_id_field, - entity_id=entity_id, - adjusted_start=adjusted_start, - adjusted_end=adjusted_end, - model=model, - api_key=api_key, - exclude_entity_ids=exclude_entity_ids, - ) - - sql_query: Final = f""" - SELECT - "{entity_id_field}" AS entity_id, - date, - api_key, - GROUPING(api_key) AS api_key_rolled, - SUM(spend)::float AS spend, - {_ptu_flat_cost_select(table_name)}, - SUM(prompt_tokens)::bigint AS prompt_tokens, - SUM(completion_tokens)::bigint AS completion_tokens, - SUM(cache_read_input_tokens)::bigint AS cache_read_input_tokens, - SUM(cache_creation_input_tokens)::bigint AS cache_creation_input_tokens, - SUM(compression_saved_tokens)::bigint AS compression_saved_tokens, - SUM(compression_savings_spend)::float AS compression_savings_spend, - SUM(prompt_caching_savings_spend)::float AS prompt_caching_savings_spend, - SUM(gateway_injected_caching_savings_spend)::float AS gateway_injected_caching_savings_spend, - SUM(autorouter_savings_spend)::float AS autorouter_savings_spend, - SUM(api_requests)::bigint AS api_requests, - SUM(successful_requests)::bigint AS successful_requests, - SUM(failed_requests)::bigint AS failed_requests, - SUM(total_response_time_ms)::bigint AS total_response_time_ms, - SUM(timed_requests)::bigint AS timed_requests - FROM "{pg_table}" - WHERE {where_clause} - GROUP BY GROUPING SETS ( - (date, "{entity_id_field}"), - (date, "{entity_id_field}", api_key) - ) - """ - - return sql_query, sql_params - def _aggregate_spend_records_sync( *, records: Sequence[DailySpendRecord], - api_key_metadata: Mapping[str, _KeyMetadataDict], + api_key_metadata: Mapping[str, KeyMetadataRow], entity_id_field: str | None, entity_metadata_field: Mapping[str, dict[str, object]] | None, ) -> _AggregatedSpendData: @@ -954,7 +654,7 @@ def _aggregate_spend_records_sync( async def _aggregate_spend_records( *, - prisma_client: PrismaClient, + repository: DailyActivityRepository, records: Sequence[DailySpendRecord], entity_id_field: str | None, entity_metadata_field: Mapping[str, dict[str, object]] | None, @@ -968,11 +668,13 @@ async def _aggregate_spend_records( record.api_key for record in records if record.api_key and record.api_key != PTU_SENTINEL_API_KEY } - api_key_metadata: Mapping[str, _KeyMetadataDict] = MappingProxyType({}) - if api_keys: - api_key_metadata = await get_api_key_metadata( - prisma_client, api_keys, _spend_logs_window(frozenset(record.date for record in records)) + api_key_metadata: Final[Mapping[str, KeyMetadataRow]] = ( + await repository.key_metadata( + frozenset(api_keys), spend_logs_window(frozenset(record.date for record in records)) ) + if api_keys + else MappingProxyType({}) + ) return await asyncio.to_thread( _aggregate_spend_records_sync, @@ -983,8 +685,7 @@ async def _aggregate_spend_records( ) -# GROUPING() bitmask values for each grouping set emitted by -# _build_aggregated_sql_query. Per Postgres semantics, the rightmost argument +# GROUPING() bitmask values returned by the daily activity repository. Per Postgres semantics, the rightmost argument # is the least-significant bit. Argument order: # date, api_key, model, model_group, custom_llm_provider, # mcp_namespaced_tool_name, endpoint @@ -992,6 +693,7 @@ async def _aggregate_spend_records( # current grouping set's key), 0 when the column is part of the key. _GROUP_GRAND_TOTAL: Final = 127 # 0b1111111 — all rolled up _GROUP_DATE: Final = 63 # 0b0111111 — only date kept +_API_KEY_ROLLED_UP_BIT: Final = 32 # 0b0100000 _GROUP_DATE_API_KEY: Final = 31 # 0b0011111 _GROUP_DATE_MODEL: Final = 47 # 0b0101111 _GROUP_DATE_MODEL_API_KEY: Final = 15 # 0b0001111 @@ -1005,7 +707,7 @@ _GROUP_DATE_ENDPOINT: Final = 62 # 0b0111110 _GROUP_DATE_ENDPOINT_API_KEY: Final = 30 # 0b0011110 -def _record_to_spend_metrics(record: _GroupingSetsRow) -> SpendMetrics: +def _record_to_spend_metrics(record: RollupMetricsRow) -> SpendMetrics: """Build a SpendMetrics directly from one already-aggregated rollup row. SUM() over zero rows is SQL NULL, so rollup rows (notably the grand-total @@ -1036,8 +738,8 @@ def _record_to_spend_metrics(record: _GroupingSetsRow) -> SpendMetrics: def _aggregate_grouping_sets_records_sync( *, - records: Sequence[_GroupingSetsRow], - api_key_metadata: Mapping[str, _KeyMetadataDict], + records: Sequence[GroupingSetsRow], + api_key_metadata: Mapping[str, KeyMetadataRow], ) -> _AggregatedSpendData: """Build the response from rollup rows produced by the GROUPING SETS query. @@ -1162,17 +864,17 @@ def _aggregate_grouping_sets_records_sync( async def _aggregate_grouping_sets_records( *, - prisma_client: PrismaClient, - records: Sequence[_GroupingSetsRow], + repository: DailyActivityRepository, + records: Sequence[GroupingSetsRow], ) -> _AggregatedSpendData: """Async wrapper: fetch api_key_metadata, then dispatch on a worker thread.""" api_keys: Final[set[str]] = {r.api_key for r in records if r.api_key and r.api_key != PTU_SENTINEL_API_KEY} - api_key_metadata: Mapping[str, _KeyMetadataDict] = MappingProxyType({}) - if api_keys: - api_key_metadata = await get_api_key_metadata( - prisma_client, api_keys, _spend_logs_window(frozenset(r.date for r in records)) - ) + api_key_metadata: Final[Mapping[str, KeyMetadataRow]] = ( + await repository.key_metadata(frozenset(api_keys), spend_logs_window(frozenset(r.date for r in records))) + if api_keys + else MappingProxyType({}) + ) return await asyncio.to_thread( _aggregate_grouping_sets_records_sync, @@ -1200,82 +902,42 @@ async def get_daily_activity( resolve_entity_metadata: Callable[[Sequence[DailySpendRecord]], Awaitable[dict[str, dict[str, object]]]] | None = None, ) -> SpendAnalyticsPaginatedResponse: - """Common function to get daily activity for any entity type. - - ``resolve_entity_metadata`` lets a caller resolve entity metadata from the - rows actually on the page (e.g. user_id -> user_email) instead of fetching - the whole entity table upfront, which matters when the entity set is - unbounded. - """ - if prisma_client is None: - raise HTTPException( - status_code=500, - detail={"error": CommonProxyErrors.db_not_connected_error.value}, - ) - - if start_date is None or end_date is None: - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail={"error": "Please provide start_date and end_date"}, - ) - + raise HTTPException(status_code=500, detail={"error": CommonProxyErrors.db_not_connected_error.value}) + date_range: Final = parse_canonical_date_range(start_date, end_date) + if isinstance(date_range, InvalidDateRange): + raise_public(date_range) try: - where_conditions: Final = _build_where_conditions( - entity_id_field=entity_id_field, - entity_id=entity_id, - start_date=start_date, - end_date=end_date, - model=model, - api_key=api_key, - exclude_entity_ids=exclude_entity_ids, - timezone_offset_minutes=timezone_offset_minutes, - include_current_utc_day=include_current_utc_day, + scope: Final = daily_activity_scope( + table_name, + entity_id_field, + entity_id, + exclude_entity_ids, + api_key, + date_range.start.isoformat(), + date_range.end.isoformat(), + model, + timezone_offset_minutes, + include_current_utc_day, ) - - spend_table: Final[TableActions[DailySpendRecord]] = getattr(prisma_client.db, table_name) - - # Get total count for pagination - total_count: Final[int] = await spend_table.count(where=where_conditions) - - # Fetch paginated results. - # ``date`` alone is not a unique sort key -- a busy tenant has many - # rows per date (one per api_key, model, model_group, provider, - # endpoint, ...), so offset pagination over ``date desc`` lands on - # arbitrary boundaries and the same row can be skipped on one page - # and returned on another. A client that pages through and sums the - # per-page metrics (the Usage dashboard) then gets a non-deterministic - # total. Adding ``id`` (the row's UUID primary key, present on both - # LiteLLM_DailyUserSpend and LiteLLM_DailyTeamSpend) as a tiebreaker - # gives every page a stable cursor (#30164). - daily_spend_data: Final[Sequence[DailySpendRecord]] = await spend_table.find_many( - where=where_conditions, - order=[ - {"date": "desc"}, - {"id": "asc"}, - ], - skip=(page - 1) * page_size, - take=page_size, - ) - + repository: Final = daily_activity_repository(prisma_client) + page_data: Final = await repository.daily_rows(scope, page=page, page_size=page_size) + daily_spend_data: Final = page_data.rows resolved_entity_metadata = entity_metadata_field if resolve_entity_metadata is not None: resolved_entity_metadata = { **(entity_metadata_field or {}), **(await resolve_entity_metadata(daily_spend_data)), } - aggregated: Final = await _aggregate_spend_records( - prisma_client=prisma_client, + repository=repository, records=daily_spend_data, entity_id_field=entity_id_field, entity_metadata_field=resolved_entity_metadata, ) - metadata_metrics = aggregated["totals"] if metadata_metrics_func: metadata_metrics = metadata_metrics_func(daily_spend_data) - return SpendAnalyticsPaginatedResponse( results=aggregated["results"], metadata=DailySpendMetadata( @@ -1297,24 +959,22 @@ async def get_daily_activity( total_response_time_ms=metadata_metrics.total_response_time_ms, total_timed_requests=metadata_metrics.timed_requests, page=page, - total_pages=-(-total_count // page_size), # Ceiling division - has_more=(page * page_size) < total_count, + total_pages=-(-page_data.total_count // page_size), + has_more=(page * page_size) < page_data.total_count, ), ) - - except Exception as e: - verbose_proxy_logger.exception("Error fetching daily activity: %s", e) + except Exception as exc: + verbose_proxy_logger.exception("Error fetching daily activity: %s", exc) raise HTTPException( - status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, - detail={"error": f"Failed to fetch analytics: {e}"}, + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail={"error": f"Failed to fetch analytics: {exc}"} ) def _fold_entity_rollups_sync( *, results: Sequence[DailySpendData], - entity_rows: Sequence[_EntityRollupRow], - api_key_metadata: Mapping[str, _KeyMetadataDict], + entity_rows: Sequence[EntityRollupRow], + api_key_metadata: Mapping[str, KeyMetadataRow], entity_metadata_field: Mapping[str, dict[str, object]] | None, # mutable-ok: shared field shape ) -> None: """Write breakdown.entities onto the already-built per-day results.""" @@ -1345,105 +1005,42 @@ def _fold_entity_rollups_sync( async def get_daily_activity_aggregated( - prisma_client: PrismaClient | None, - table_name: str, - entity_id_field: str, - entity_id: str | list[str] | None, - entity_metadata_field: Mapping[str, dict[str, object]] | None, - start_date: str | None, - end_date: str | None, - model: str | None, - api_key: str | list[str] | None, # mutable-ok: filter union shared with the paginated path - exclude_entity_ids: list[str] | None = None, - timezone_offset_minutes: int | None = None, + repository: DailyActivityRepository, + scope: DailyActivityScope, + *, + entity_metadata_field: Mapping[str, dict[str, object]] | None = None, include_entity_breakdown: bool = False, - include_current_utc_day: bool = False, + api_key_limit: int = constants.USAGE_TOP_API_KEYS_DEFAULT, ) -> SpendAnalyticsPaginatedResponse: - """Aggregated variant that returns the full result set (no pagination). - - Uses SQL GROUP BY to aggregate rows in the database rather than fetching - all individual rows into Python. This collapses rows across entities - (users/teams/orgs), reducing ~150k rows to ~2-3k grouped rows. - - include_entity_breakdown runs a small companion rollup query and folds - `breakdown.entities` onto the response, as entity-scoped views like Team Usage need. - - Matches the response model of the paginated endpoint so the UI does not need to transform. - """ - if prisma_client is None: - raise HTTPException( - status_code=500, - detail={"error": CommonProxyErrors.db_not_connected_error.value}, - ) - - if start_date is None or end_date is None: - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail={"error": "Please provide start_date and end_date"}, - ) - try: - sql_query, sql_params = _build_aggregated_sql_query( - table_name=table_name, - entity_id_field=entity_id_field, - entity_id=entity_id, - start_date=start_date, - end_date=end_date, - model=model, - api_key=api_key, - exclude_entity_ids=exclude_entity_ids, - timezone_offset_minutes=timezone_offset_minutes, - include_current_utc_day=include_current_utc_day, + aggregated_rows: Final = await repository.aggregated( + scope, include_entity_breakdown=include_entity_breakdown, api_key_limit=api_key_limit ) - - entity_query: Final = ( - _build_entity_rollup_sql_query( - table_name=table_name, - entity_id_field=entity_id_field, - entity_id=entity_id, - start_date=start_date, - end_date=end_date, - model=model, - api_key=api_key, - exclude_entity_ids=exclude_entity_ids, - timezone_offset_minutes=timezone_offset_minutes, - include_current_utc_day=include_current_utc_day, - ) + records: Final = aggregated_rows.grouping_rows + aggregated: Final = await _aggregate_grouping_sets_records( + repository=repository, + records=records, + ) + entity_total_api_keys: Final[dict[str, int] | None] = ( + { + row.entity_id or "Unassigned": row.distinct_api_keys + for row in aggregated_rows.entity_rows or () + if row.api_key_rolled and row.distinct_api_keys is not None + } if include_entity_breakdown else None ) - - # Execute the GROUPING SETS query (one row per rollup level), alongside - # the per-entity companion rollup when the caller wants entities. - raw_rows, raw_entity_rows = ( - await asyncio.gather( - prisma_client.db.query_raw(sql_query, *sql_params), - prisma_client.db.query_raw(entity_query[0], *entity_query[1]), - ) - if entity_query is not None - else (await prisma_client.db.query_raw(sql_query, *sql_params), None) - ) - - records: Final = [_GroupingSetsRow(**row) for row in (raw_rows or [])] - - # The grouping-sets dispatcher places each row directly in its bucket - # using the row's GROUPING() bitmask. No Python-side summing needed. - aggregated: Final = await _aggregate_grouping_sets_records( - prisma_client=prisma_client, - records=records, - ) - - if raw_entity_rows: - entity_records: Final = tuple(_EntityRollupRow(**row) for row in raw_entity_rows) + if aggregated_rows.entity_rows: + entity_records: Final = aggregated_rows.entity_rows entity_api_keys: Final = frozenset( - r.api_key for r in entity_records if r.api_key and r.api_key != PTU_SENTINEL_API_KEY + row.api_key for row in entity_records if row.api_key and row.api_key != PTU_SENTINEL_API_KEY ) - entity_key_metadata: Final = ( - await get_api_key_metadata( - prisma_client, entity_api_keys, _spend_logs_window(frozenset(r.date for r in entity_records)) + entity_key_metadata: Final[Mapping[str, KeyMetadataRow]] = ( + await repository.key_metadata( + entity_api_keys, spend_logs_window(frozenset(row.date for row in entity_records)) ) if entity_api_keys - else {} + else MappingProxyType({}) ) await asyncio.to_thread( _fold_entity_rollups_sync, @@ -1452,7 +1049,6 @@ async def get_daily_activity_aggregated( api_key_metadata=entity_key_metadata, entity_metadata_field=entity_metadata_field, ) - return SpendAnalyticsPaginatedResponse( results=aggregated["results"], metadata=DailySpendMetadata( @@ -1478,12 +1074,13 @@ async def get_daily_activity_aggregated( page=1, total_pages=1, has_more=False, + api_key_limit=api_key_limit, + total_api_keys=aggregated_rows.distinct_api_keys, + entity_total_api_keys=entity_total_api_keys, ), ) - - except Exception as e: - verbose_proxy_logger.exception("Error fetching aggregated daily activity: %s", e) + except Exception as exc: + verbose_proxy_logger.exception("Error fetching aggregated daily activity: %s", exc) raise HTTPException( - status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, - detail={"error": f"Failed to fetch analytics: {e}"}, + status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail={"error": f"Failed to fetch analytics: {exc}"} ) diff --git a/litellm/proxy/management_endpoints/customer_endpoints.py b/litellm/proxy/management_endpoints/customer_endpoints.py index b35bc01b4d0..7236fd12e9d 100644 --- a/litellm/proxy/management_endpoints/customer_endpoints.py +++ b/litellm/proxy/management_endpoints/customer_endpoints.py @@ -12,7 +12,8 @@ All /customer management endpoints #### END-USER/CUSTOMER MANAGEMENT #### from collections.abc import Mapping, Sequence from datetime import datetime, timedelta -from typing import TYPE_CHECKING, Final, Protocol, TypeVar, overload +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, NamedTuple, Protocol, TypeVar, overload import fastapi from fastapi import APIRouter, Depends, HTTPException, Request @@ -41,6 +42,7 @@ from litellm.proxy.management_helpers.object_permission_utils import ( ) from litellm.proxy.utils import handle_exception_on_proxy from litellm.repositories.budget_repository import BudgetRepository +from litellm.repositories.chunked_in import find_many_in from litellm.repositories.table_repositories import EndUserRepository from litellm.types.proxy.management_endpoints.common_daily_activity import ( SpendAnalyticsPaginatedResponse, @@ -490,6 +492,33 @@ async def new_end_user( raise handle_exception_on_proxy(e) +class _CustomerDailyActivityScope(NamedTuple): + end_user_ids: tuple[str, ...] | None + end_user_metadata: Mapping[str, dict[str, object]] + + +async def resolve_customer_daily_activity_scope( + *, + end_user_ids: tuple[str, ...] | None, + prisma_client: "PrismaClient", +) -> _CustomerDailyActivityScope: + end_user_table: Final = _typed_table(EndUserRepository(prisma_client)) + end_user_aliases: Final = ( + await find_many_in(end_user_table, "user_id", end_user_ids) + if end_user_ids is not None + else await end_user_table.find_many(where={}) + ) + metadata: Final = MappingProxyType({end_user.user_id: {"alias": end_user.alias} for end_user in end_user_aliases}) + return _CustomerDailyActivityScope(end_user_ids, metadata) + + +def customer_daily_activity_is_admin(user_api_key_dict: UserAPIKeyAuth) -> bool: + return user_api_key_dict.user_role in ( + LitellmUserRoles.PROXY_ADMIN, + LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, + ) + + @router.get( "/customer/info", tags=["Customer Management"], @@ -888,10 +917,7 @@ async def get_customer_daily_activity( """ Get daily activity for specific organizations or all accessible organizations. """ - if ( - user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN - and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY - ): + if not customer_daily_activity_is_admin(user_api_key_dict): raise HTTPException( status_code=401, detail={"error": f"Admin-only endpoint. Your user role={user_api_key_dict.user_role}"}, @@ -906,24 +932,22 @@ async def get_customer_daily_activity( ) # Parse comma-separated ids - end_user_ids_list: Final = end_user_ids.split(",") if end_user_ids else None + end_user_ids_list: Final = tuple(end_user_ids.split(",")) if end_user_ids else None exclude_end_user_ids_list: list[str] | None = None if exclude_end_user_ids: exclude_end_user_ids_list = exclude_end_user_ids.split(",") if exclude_end_user_ids else None - # Fetch organization aliases for metadata - where_condition: Final = dict[str, object]() - if end_user_ids_list: - where_condition["user_id"] = {"in": list(end_user_ids_list)} - end_user_aliases: Final = await _typed_table(EndUserRepository(prisma_client)).find_many(where=where_condition) + customer_scope: Final = await resolve_customer_daily_activity_scope( + end_user_ids=end_user_ids_list, + prisma_client=prisma_client, + ) - # Query daily activity for organizations return await get_daily_activity( prisma_client=prisma_client, table_name="litellm_dailyenduserspend", entity_id_field="end_user_id", - entity_id=end_user_ids_list, - entity_metadata_field={e.user_id: {"alias": e.alias} for e in end_user_aliases}, + entity_id=None if customer_scope.end_user_ids is None else list(customer_scope.end_user_ids), + entity_metadata_field=customer_scope.end_user_metadata, exclude_entity_ids=exclude_end_user_ids_list, start_date=start_date, end_date=end_date, diff --git a/litellm/proxy/management_endpoints/daily_activity_routes.py b/litellm/proxy/management_endpoints/daily_activity_routes.py new file mode 100644 index 00000000000..fa3c745c88e --- /dev/null +++ b/litellm/proxy/management_endpoints/daily_activity_routes.py @@ -0,0 +1,579 @@ +import csv +import io +import json +from collections.abc import AsyncIterator, Mapping, Sequence +from dataclasses import asdict, fields, replace +from datetime import date, datetime +from typing import Annotated, Final, Literal + +from fastapi import APIRouter, Depends, HTTPException, Query +from fastapi.encoders import jsonable_encoder +from fastapi.responses import StreamingResponse + +from litellm import constants +from litellm._logging import verbose_proxy_logger +from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.management_endpoints.common_daily_activity import ( + InvalidDateRange, + ScopeDenied, + daily_activity_repository, + get_daily_activity_aggregated, + parse_canonical_date_range, + raise_public, + spend_logs_window, +) +from litellm.proxy.management_endpoints.daily_activity_scopes import ( + AGENT_RESOLVER, + CUSTOMER_RESOLVER, + ORGANIZATION_RESOLVER, + TAG_RESOLVER, + TEAM_RESOLVER, + USER_RESOLVER, + EntityQuery, + EntityScopeResolver, + ResolvedScope, +) +from litellm.proxy.management_endpoints.team_endpoints import aggregated_date_range_error +from litellm.proxy.management_helpers.utils import management_endpoint_wrapper +from litellm.proxy.utils import PrismaClient, get_prisma_client_or_throw +from litellm.repositories.daily_activity_repository import DailyActivityRepository +from litellm.types.proxy.management_endpoints.common_daily_activity import ( + CacheLeakageKeysResponse, + DailyActivityKeyPageResponse, + DailyActivityKeySearchResponse, + KeyActivityRow, + KeyMetadata, + KeySpendActivityRow, + KeySpendMetrics, + ModelTopKeysResponse, + SpendAnalyticsPaginatedResponse, + SpendMetrics, +) +from litellm.types.repositories.daily_activity import ( + ExportRow, + ExportType, + KeyMetadataRow, + KeySpendRow, +) + +router = APIRouter() + + +def get_daily_activity_prisma_client() -> PrismaClient: + return get_prisma_client_or_throw(CommonProxyErrors.db_not_connected_error.value) + + +def get_daily_activity_repository() -> DailyActivityRepository: + return daily_activity_repository(get_daily_activity_prisma_client()) + + +def _date_range_error(query: EntityQuery, *, user_aggregated: bool) -> InvalidDateRange | None: + if user_aggregated: + date_range: Final = parse_canonical_date_range(query.start_date, query.end_date) + return date_range if isinstance(date_range, InvalidDateRange) else None + + range_error: Final[str | None] = aggregated_date_range_error(query.start_date, query.end_date) + return None if range_error is None else InvalidDateRange(reason=range_error) + + +async def _resolved_scope( + resolver: EntityScopeResolver, + query: EntityQuery, + user_api_key_dict: UserAPIKeyAuth, + prisma_client: PrismaClient, + *, + user_aggregated: bool, +) -> ResolvedScope: + date_error: Final[InvalidDateRange | None] = _date_range_error(query, user_aggregated=user_aggregated) + if date_error is not None: + raise_public(date_error) + result: ResolvedScope | ScopeDenied = await resolver.resolve(user_api_key_dict, query, prisma_client) + if isinstance(result, ScopeDenied): + raise_public(result) + return result + + +def _sum_metrics(metrics: Sequence[SpendMetrics]) -> SpendMetrics: + return SpendMetrics( + spend=sum(metric.spend for metric in metrics), + flat_cost=sum(metric.flat_cost for metric in metrics), + prompt_tokens=sum(metric.prompt_tokens for metric in metrics), + completion_tokens=sum(metric.completion_tokens for metric in metrics), + cache_read_input_tokens=sum(metric.cache_read_input_tokens for metric in metrics), + cache_creation_input_tokens=sum(metric.cache_creation_input_tokens for metric in metrics), + compression_saved_tokens=sum(metric.compression_saved_tokens for metric in metrics), + compression_savings_spend=sum(metric.compression_savings_spend for metric in metrics), + prompt_caching_savings_spend=sum(metric.prompt_caching_savings_spend for metric in metrics), + gateway_injected_caching_savings_spend=sum(metric.gateway_injected_caching_savings_spend for metric in metrics), + autorouter_savings_spend=sum(metric.autorouter_savings_spend for metric in metrics), + total_tokens=sum(metric.total_tokens for metric in metrics), + successful_requests=sum(metric.successful_requests for metric in metrics), + failed_requests=sum(metric.failed_requests for metric in metrics), + api_requests=sum(metric.api_requests for metric in metrics), + total_response_time_ms=sum(metric.total_response_time_ms for metric in metrics), + timed_requests=sum(metric.timed_requests for metric in metrics), + ) + + +def _key_metadata(api_key: str, metadata: Mapping[str, KeyMetadataRow]) -> KeyMetadata: + row: Final[KeyMetadataRow | None] = metadata.get(api_key) + if row is None: + return KeyMetadata() + return KeyMetadata( + key_alias=row.key_alias, + team_id=row.team_id, + user_id=row.user_id, + user_email=row.user_email, + key_exists=row.key_exists, + ) + + +def _key_activity_row(row: KeySpendRow, metadata: Mapping[str, KeyMetadataRow]) -> KeySpendActivityRow: + return KeySpendActivityRow( + api_key=row.api_key, + metrics=KeySpendMetrics( + spend=row.spend, + prompt_tokens=row.prompt_tokens, + completion_tokens=row.completion_tokens, + total_tokens=row.total_tokens, + api_requests=row.api_requests, + successful_requests=row.successful_requests, + failed_requests=row.failed_requests, + cache_read_input_tokens=row.cache_read_input_tokens, + cache_creation_input_tokens=row.cache_creation_input_tokens, + ), + metadata=_key_metadata(row.api_key, metadata), + ) + + +async def _key_activity_rows( + repository: DailyActivityRepository, + rows: Sequence[KeySpendRow], + resolved_scope: ResolvedScope, +) -> list[KeySpendActivityRow]: + spend_window: Final[tuple[datetime, datetime] | None] = spend_logs_window( + frozenset((resolved_scope.scope.start_date, resolved_scope.scope.end_date)) + ) + metadata: Final[Mapping[str, KeyMetadataRow]] = await repository.key_metadata( + frozenset(row.api_key for row in rows), + spend_window, + ) + return [_key_activity_row(row, metadata) for row in rows] + + +def _export_filename( + entity: str, + start_date: date, + end_date: date, + export_type: ExportType, + file_format: Literal["csv", "json"], +) -> str: + extension: Final[str] = "csv" if file_format == "csv" else "json" + return f"{entity}-usage-{start_date.isoformat()}-{end_date.isoformat()}-{export_type.value}.{extension}" + + +def _content_disposition( + entity: str, + start_date: date, + end_date: date, + export_type: ExportType, + file_format: Literal["csv", "json"], +) -> str: + filename: Final = _export_filename(entity, start_date, end_date, export_type, file_format) + return f'attachment; filename="{filename}"' + + +def _fold_key_metrics(api_key: str, response: SpendAnalyticsPaginatedResponse) -> KeyActivityRow | None: + metrics: Final[tuple[SpendMetrics, ...]] = tuple( + day.breakdown.api_keys[api_key].metrics for day in response.results if api_key in day.breakdown.api_keys + ) + if not metrics: + return None + metadata: Final[KeyMetadata] = next( + day.breakdown.api_keys[api_key].metadata for day in response.results if api_key in day.breakdown.api_keys + ) + return KeyActivityRow(api_key=api_key, metrics=_sum_metrics(metrics), metadata=metadata) + + +def _search_rows(keys: Sequence[str], response: SpendAnalyticsPaginatedResponse) -> list[KeyActivityRow]: + return [row for key in keys if (row := _fold_key_metrics(key, response)) is not None] + + +def _csv_cell(value: object) -> object: + if isinstance(value, str) and value.startswith(("=", "+", "-", "@", "\t", "\r")): + return f"'{value}" + return value + + +def _csv_row(values: Sequence[object]) -> bytes: + output: Final[io.StringIO] = io.StringIO(newline="") + csv.writer(output, lineterminator="\r\n").writerow(tuple(_csv_cell(value) for value in values)) + return output.getvalue().encode() + + +def _stream_export_rows( + first_row: ExportRow | None, + rows: AsyncIterator[ExportRow], + file_format: Literal["csv", "json"], +) -> AsyncIterator[bytes]: + async def stream() -> AsyncIterator[bytes]: + if file_format == "csv": + yield _csv_row(tuple(field.name for field in fields(ExportRow))) + if first_row is not None: + yield _csv_row(tuple(asdict(first_row).values())) + async for row in rows: + yield _csv_row(tuple(asdict(row).values())) + return + + if first_row is None: + yield b"[]" + return + yield b"[" + json.dumps(jsonable_encoder(first_row), separators=(",", ":")).encode() + async for row in rows: + yield b"," + json.dumps(jsonable_encoder(row), separators=(",", ":")).encode() + yield b"]" + + return stream() + + +def _register_aggregated_route(router: APIRouter, resolver: EntityScopeResolver, prefix: str) -> None: + @management_endpoint_wrapper + async def aggregated( + entity_query: Annotated[EntityQuery, Depends(resolver.query)], + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + repository: Annotated[DailyActivityRepository, Depends(get_daily_activity_repository)], + prisma_client: Annotated[PrismaClient, Depends(get_daily_activity_prisma_client)], + api_key_limit: Annotated[ + int, Query(ge=1, le=constants.USAGE_TOP_API_KEYS_MAX) + ] = constants.USAGE_TOP_API_KEYS_DEFAULT, + ) -> SpendAnalyticsPaginatedResponse: + try: + resolved: ResolvedScope = await _resolved_scope( + resolver, + entity_query, + user_api_key_dict, + prisma_client, + user_aggregated=resolver.entity == "user", + ) + return await get_daily_activity_aggregated( + repository, + resolved.scope, + entity_metadata_field=resolved.entity_metadata, + include_entity_breakdown=resolver.include_entity_breakdown, + api_key_limit=api_key_limit, + ) + except HTTPException: + raise + except Exception as exc: + verbose_proxy_logger.exception("Daily activity aggregation failed: %s", exc) + raise HTTPException(status_code=500, detail={"error": f"Failed to fetch analytics: {exc}"}) + + router.add_api_route( + f"{prefix}/daily/activity/aggregated", + aggregated, + methods=["GET"], + name=resolver.operation_names["aggregated"], + tags=list(resolver.tags), + dependencies=(Depends(user_api_key_auth),), + response_model=SpendAnalyticsPaginatedResponse, + include_in_schema=prefix != "/end_user", + ) + + if resolver.entity == "user": + aggregated.__doc__ = ( + "Aggregated analytics for a user's daily activity without pagination.\n" + "Returns the same response shape as the paginated endpoint with page metadata set to single-page.\n\n" + "Reads daily spend records that only ever accumulate and are never affected by budget\n" + "resets. Their total can legitimately exceed the `spend` field returned by\n" + "`/v2/user/info`, which is a running budget counter that every budget reset sets back\n" + "to zero (or to the overage above `max_budget` when `budget_rollover` is enabled)." + ) + elif resolver.entity == "team": + aggregated.__doc__ = ( + "Aggregated daily activity for teams without pagination, including per-team breakdown.\n\n" + "One SQL GROUPING SETS pass returns every day in the range regardless of row\n" + "volume, so callers never reassemble pages. Same response shape as the\n" + "paginated endpoint with page metadata pinned to a single page.\n\n" + "Args:\n" + " team_ids (Optional[str]): Comma-separated list of team IDs to filter by. If not provided, " + "returns data for all teams.\n" + " start_date (Optional[str]): Start date for the activity period (YYYY-MM-DD).\n" + " end_date (Optional[str]): End date for the activity period (YYYY-MM-DD).\n" + " model (Optional[str]): Filter by model name.\n" + " api_key (Optional[str]): Filter by API key.\n" + " exclude_team_ids (Optional[str]): Comma-separated list of team IDs to exclude.\n" + " timezone (Optional[int]): Timezone offset in minutes from UTC, matching JavaScript's " + "Date.getTimezoneOffset() convention.\n" + "Returns:\n" + " SpendAnalyticsPaginatedResponse: Response containing all daily activity data for the range." + ) + + +def _register_key_page_route(router: APIRouter, resolver: EntityScopeResolver, prefix: str) -> None: + @management_endpoint_wrapper + async def key_page( + entity_query: Annotated[EntityQuery, Depends(resolver.query)], + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + repository: Annotated[DailyActivityRepository, Depends(get_daily_activity_repository)], + prisma_client: Annotated[PrismaClient, Depends(get_daily_activity_prisma_client)], + offset: Annotated[int, Query(ge=0)] = 0, + limit: Annotated[int, Query(ge=1, le=constants.USAGE_KEY_PAGE_MAX)] = constants.USAGE_KEY_PAGE_DEFAULT, + ) -> DailyActivityKeyPageResponse: + try: + resolved: Final = await _resolved_scope( + resolver, + entity_query, + user_api_key_dict, + prisma_client, + user_aggregated=False, + ) + page: Final = await repository.key_page(resolved.scope, offset=offset, limit=limit) + return DailyActivityKeyPageResponse( + api_keys=await _key_activity_rows(repository, page.rows, resolved), + total_api_keys=page.total_api_keys, + offset=offset, + limit=limit, + ) + except HTTPException: + raise + except Exception as exc: + verbose_proxy_logger.exception("Daily activity key page failed: %s", exc) + raise HTTPException(status_code=500, detail={"error": f"Failed to fetch analytics: {exc}"}) + + router.add_api_route( + f"{prefix}/daily/activity/aggregated/keys", + key_page, + methods=["GET"], + name=resolver.operation_names["key_page"], + tags=list(resolver.tags), + dependencies=(Depends(user_api_key_auth),), + response_model=DailyActivityKeyPageResponse, + include_in_schema=prefix != "/end_user", + ) + + +def _register_search_route(router: APIRouter, resolver: EntityScopeResolver, prefix: str) -> None: + @management_endpoint_wrapper + async def search( + entity_query: Annotated[EntityQuery, Depends(resolver.query)], + search: Annotated[str, Query(min_length=1)], + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + repository: Annotated[DailyActivityRepository, Depends(get_daily_activity_repository)], + prisma_client: Annotated[PrismaClient, Depends(get_daily_activity_prisma_client)], + limit: Annotated[int, Query(ge=1, le=constants.USAGE_KEY_SEARCH_MAX)] = constants.USAGE_KEY_SEARCH_DEFAULT, + ) -> DailyActivityKeySearchResponse: + try: + resolved: ResolvedScope = await _resolved_scope( + resolver, + entity_query, + user_api_key_dict, + prisma_client, + user_aggregated=False, + ) + keys: tuple[str, ...] = await repository.search_keys( + resolved.scope, + search=search, + limit=limit, + ) + if not keys: + return DailyActivityKeySearchResponse(api_keys=[]) + search_response: SpendAnalyticsPaginatedResponse = await get_daily_activity_aggregated( + repository, + replace(resolved.scope, api_keys=keys), + include_entity_breakdown=False, + ) + return DailyActivityKeySearchResponse(api_keys=_search_rows(keys, search_response)) + except HTTPException: + raise + except Exception as exc: + verbose_proxy_logger.exception("Daily activity key search failed: %s", exc) + raise HTTPException(status_code=500, detail={"error": f"Failed to fetch analytics: {exc}"}) + + router.add_api_route( + f"{prefix}/daily/activity/aggregated/search", + search, + methods=["GET"], + name=resolver.operation_names["search"], + tags=list(resolver.tags), + dependencies=(Depends(user_api_key_auth),), + response_model=DailyActivityKeySearchResponse, + include_in_schema=prefix != "/end_user", + ) + + +def _register_model_top_keys_route(router: APIRouter, resolver: EntityScopeResolver, prefix: str) -> None: + @management_endpoint_wrapper + async def model_top_keys( + entity_query: Annotated[EntityQuery, Depends(resolver.query)], + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + repository: Annotated[DailyActivityRepository, Depends(get_daily_activity_repository)], + prisma_client: Annotated[PrismaClient, Depends(get_daily_activity_prisma_client)], + model_group: Annotated[str, Query(min_length=1)], + by_model_group: Annotated[bool, Query()] = True, + limit: Annotated[int, Query(ge=1, le=constants.USAGE_MODEL_TOP_KEYS_MAX)] = ( + constants.USAGE_MODEL_TOP_KEYS_DEFAULT + ), + ) -> ModelTopKeysResponse: + try: + resolved: ResolvedScope = await _resolved_scope( + resolver, + entity_query, + user_api_key_dict, + prisma_client, + user_aggregated=False, + ) + rows: tuple[KeySpendRow, ...] = await repository.model_top_keys( + resolved.scope, + model_group=model_group, + by_model_group=by_model_group, + limit=limit, + ) + return ModelTopKeysResponse( + model=model_group, + by_model_group=by_model_group, + api_keys=await _key_activity_rows(repository, rows, resolved), + ) + except HTTPException: + raise + except Exception as exc: + verbose_proxy_logger.exception("Daily activity model top keys failed: %s", exc) + raise HTTPException(status_code=500, detail={"error": f"Failed to fetch analytics: {exc}"}) + + router.add_api_route( + f"{prefix}/daily/activity/aggregated/model_top_keys", + model_top_keys, + methods=["GET"], + name=resolver.operation_names["model_top_keys"], + tags=list(resolver.tags), + dependencies=(Depends(user_api_key_auth),), + response_model=ModelTopKeysResponse, + include_in_schema=prefix != "/end_user", + ) + + +def _register_export_route(router: APIRouter, resolver: EntityScopeResolver, prefix: str) -> None: + @management_endpoint_wrapper + async def export( + entity_query: Annotated[EntityQuery, Depends(resolver.query)], + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + repository: Annotated[DailyActivityRepository, Depends(get_daily_activity_repository)], + prisma_client: Annotated[PrismaClient, Depends(get_daily_activity_prisma_client)], + export_type: Annotated[ExportType, Query()] = ExportType.DAILY, + file_format: Annotated[Literal["csv", "json"], Query(alias="format")] = "csv", + ) -> StreamingResponse: + try: + resolved: ResolvedScope = await _resolved_scope( + resolver, + entity_query, + user_api_key_dict, + prisma_client, + user_aggregated=False, + ) + rows: Final = repository.export_rows(resolved.scope, export_type=export_type) + first_row: Final = await anext(rows, None) + return StreamingResponse( + _stream_export_rows(first_row, rows, file_format), + media_type="text/csv" if file_format == "csv" else "application/json", + headers={ + "Cache-Control": "no-store", + "Content-Disposition": _content_disposition( + resolver.entity, + date.fromisoformat(resolved.scope.start_date), + date.fromisoformat(resolved.scope.end_date), + export_type, + file_format, + ), + }, + ) + except HTTPException: + raise + except Exception as exc: + verbose_proxy_logger.exception("Daily activity export failed: %s", exc) + raise HTTPException(status_code=500, detail={"error": f"Failed to fetch analytics: {exc}"}) + + router.add_api_route( + f"{prefix}/daily/activity/export", + export, + methods=["GET"], + name=resolver.operation_names["export"], + tags=list(resolver.tags), + dependencies=(Depends(user_api_key_auth),), + response_class=StreamingResponse, + responses={ + 200: { + "description": "Streamed daily activity export", + "content": { + "text/csv": {"schema": {"type": "string"}}, + "application/json": {"schema": {"type": "array", "items": {"type": "object"}}}, + }, + } + }, + include_in_schema=prefix != "/end_user", + ) + + +def _register_cache_leakage_route(router: APIRouter, resolver: EntityScopeResolver, prefix: str) -> None: + @management_endpoint_wrapper + async def cache_leakage_keys( + entity_query: Annotated[EntityQuery, Depends(resolver.query)], + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + repository: Annotated[DailyActivityRepository, Depends(get_daily_activity_repository)], + prisma_client: Annotated[PrismaClient, Depends(get_daily_activity_prisma_client)], + limit: Annotated[ + int, Query(ge=1, le=constants.USAGE_CACHE_LEAKAGE_KEYS_MAX) + ] = constants.USAGE_CACHE_LEAKAGE_KEYS_DEFAULT, + ) -> CacheLeakageKeysResponse: + try: + resolved: ResolvedScope = await _resolved_scope( + resolver, + entity_query, + user_api_key_dict, + prisma_client, + user_aggregated=False, + ) + rows: tuple[KeySpendRow, ...] = await repository.cache_leakage_keys( + resolved.scope, + limit=limit, + ) + return CacheLeakageKeysResponse( + api_keys=await _key_activity_rows(repository, rows, resolved), + ) + except HTTPException: + raise + except Exception as exc: + verbose_proxy_logger.exception("Daily activity cache leakage keys failed: %s", exc) + raise HTTPException(status_code=500, detail={"error": f"Failed to fetch analytics: {exc}"}) + + router.add_api_route( + f"{prefix}/daily/activity/aggregated/cache_leakage_keys", + cache_leakage_keys, + methods=["GET"], + name=resolver.operation_names["cache_leakage_keys"], + tags=list(resolver.tags), + dependencies=(Depends(user_api_key_auth),), + response_model=CacheLeakageKeysResponse, + include_in_schema=prefix != "/end_user", + ) + + +def register_daily_activity_routes(router: APIRouter, resolver: EntityScopeResolver) -> None: + for prefix in resolver.route_prefixes: + _register_aggregated_route(router, resolver, prefix) + _register_key_page_route(router, resolver, prefix) + _register_search_route(router, resolver, prefix) + _register_model_top_keys_route(router, resolver, prefix) + _register_export_route(router, resolver, prefix) + if resolver.entity == "user": + _register_cache_leakage_route(router, resolver, prefix) + + +for _resolver in ( + USER_RESOLVER, + TEAM_RESOLVER, + TAG_RESOLVER, + ORGANIZATION_RESOLVER, + CUSTOMER_RESOLVER, + AGENT_RESOLVER, +): + register_daily_activity_routes(router, _resolver) diff --git a/litellm/proxy/management_endpoints/daily_activity_scopes.py b/litellm/proxy/management_endpoints/daily_activity_scopes.py new file mode 100644 index 00000000000..4ed3c680f15 --- /dev/null +++ b/litellm/proxy/management_endpoints/daily_activity_scopes.py @@ -0,0 +1,472 @@ +from collections.abc import Awaitable, Callable, Mapping, Sequence +from dataclasses import dataclass +from types import MappingProxyType +from typing import Final, Literal + +from fastapi import Query + +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.agent_endpoints.endpoints import resolve_agent_daily_activity_scope +from litellm.proxy.management_endpoints.common_daily_activity import ScopeDenied +from litellm.proxy.management_endpoints.customer_endpoints import ( + customer_daily_activity_is_admin, + resolve_customer_daily_activity_scope, +) +from litellm.proxy.management_endpoints.internal_user_endpoints import resolve_user_daily_activity_entity_ids +from litellm.proxy.management_endpoints.organization_endpoints import resolve_organization_daily_activity_scope +from litellm.proxy.management_endpoints.tag_management_endpoints import get_tag_daily_activity_api_key_filter +from litellm.proxy.management_endpoints.team_endpoints import resolve_team_daily_activity_scope +from litellm.proxy.utils import PrismaClient +from litellm.types.repositories.daily_activity import DailyActivityScope, DailyActivityTable + + +@dataclass(frozen=True, slots=True) +class EntityQuery: + entity_ids: tuple[str, ...] | None + exclude_entity_ids: tuple[str, ...] + api_key: str | None + start_date: str | None + end_date: str | None + model: str | None + timezone_offset_minutes: int | None + include_current_utc_day: bool + + +@dataclass(frozen=True, slots=True) +class ResolvedScope: + scope: DailyActivityScope + entity_metadata: Mapping[str, dict[str, object]] | None + + +Entity = Literal["user", "team", "tag", "organization", "customer", "agent"] +EntityScopeResolution = ResolvedScope | ScopeDenied +EntityScopeQuery = Callable[..., EntityQuery] +EntityScopeResolve = Callable[[UserAPIKeyAuth, EntityQuery, PrismaClient], Awaitable[EntityScopeResolution]] +OperationNames = Mapping[str, str] + + +@dataclass(frozen=True, slots=True) +class EntityScopeResolver: + entity: Entity + table: DailyActivityTable + entity_id_field: str + route_prefixes: tuple[str, ...] + tags: tuple[str, ...] + query: EntityScopeQuery + resolve: EntityScopeResolve + include_entity_breakdown: bool + operation_names: OperationNames + + +def _query_ids(value: str | None) -> tuple[str, ...] | None: + return tuple(value.split(",")) if value else None + + +def _query_excluded_ids(value: str | None) -> tuple[str, ...]: + return tuple(value.split(",")) if value else () + + +def _build_scope( + resolver: EntityScopeResolver, + query: EntityQuery, + entity_ids: Sequence[str] | None, + exclude_entity_ids: Sequence[str], + api_key_filter: str | Sequence[str] | None, + entity_metadata: Mapping[str, dict[str, object]] | None, +) -> ResolvedScope: + start_date: Final[str] = query.start_date or "" + end_date: Final[str] = query.end_date or "" + api_keys: Final[tuple[str, ...] | None] = ( + None + if api_key_filter is None or api_key_filter == "" + else (api_key_filter,) + if isinstance(api_key_filter, str) + else tuple(api_key_filter) + ) + return ResolvedScope( + scope=DailyActivityScope( + table=resolver.table, + entity_id_field=resolver.entity_id_field, + entity_ids=None if entity_ids is None else tuple(entity_ids), + exclude_entity_ids=tuple(exclude_entity_ids), + api_keys=api_keys, + start_date=start_date, + end_date=end_date, + model=query.model, + timezone_offset_minutes=query.timezone_offset_minutes, + include_current_utc_day=query.include_current_utc_day, + ), + entity_metadata=entity_metadata, + ) + + +async def _resolve_user( + user_api_key_dict: UserAPIKeyAuth, query: EntityQuery, prisma_client: PrismaClient +) -> EntityScopeResolution: + entity_ids: Final = resolve_user_daily_activity_entity_ids( + user_id=query.entity_ids[0] if query.entity_ids is not None else None, + user_api_key_dict=user_api_key_dict, + ) + if isinstance(entity_ids, ScopeDenied): + return entity_ids + return _build_scope( + USER_RESOLVER, + query, + entity_ids, + query.exclude_entity_ids, + query.api_key, + None, + ) + + +async def _resolve_team( + user_api_key_dict: UserAPIKeyAuth, query: EntityQuery, prisma_client: PrismaClient +) -> EntityScopeResolution: + from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache + + team_scope: Final = await resolve_team_daily_activity_scope( + team_ids=",".join(query.entity_ids) if query.entity_ids is not None else None, + exclude_team_ids=",".join(query.exclude_entity_ids) if query.exclude_entity_ids else None, + api_key=query.api_key, + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + return _build_scope( + TEAM_RESOLVER, + query, + team_scope.team_ids, + team_scope.exclude_team_ids or (), + team_scope.api_key_filter, + team_scope.team_alias_metadata, + ) + + +async def _resolve_tag( + user_api_key_dict: UserAPIKeyAuth, query: EntityQuery, prisma_client: PrismaClient +) -> EntityScopeResolution: + api_key_filter: Final = await get_tag_daily_activity_api_key_filter( + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + requested_api_key=query.api_key, + ) + return _build_scope( + TAG_RESOLVER, + query, + query.entity_ids, + query.exclude_entity_ids, + api_key_filter, + None, + ) + + +async def _resolve_organization( + user_api_key_dict: UserAPIKeyAuth, query: EntityQuery, prisma_client: PrismaClient +) -> EntityScopeResolution: + org_scope: Final = await resolve_organization_daily_activity_scope( + organization_ids=query.entity_ids, + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + ) + return _build_scope( + ORGANIZATION_RESOLVER, + query, + org_scope.organization_ids, + query.exclude_entity_ids, + query.api_key, + org_scope.organization_metadata, + ) + + +async def _resolve_customer( + user_api_key_dict: UserAPIKeyAuth, query: EntityQuery, prisma_client: PrismaClient +) -> EntityScopeResolution: + if not customer_daily_activity_is_admin(user_api_key_dict): + return ScopeDenied(403, f"Admin-only endpoint. Your user role={user_api_key_dict.user_role}") + customer_scope: Final = await resolve_customer_daily_activity_scope( + end_user_ids=query.entity_ids, + prisma_client=prisma_client, + ) + return _build_scope( + CUSTOMER_RESOLVER, + query, + customer_scope.end_user_ids, + query.exclude_entity_ids, + query.api_key, + customer_scope.end_user_metadata, + ) + + +async def _resolve_agent( + user_api_key_dict: UserAPIKeyAuth, query: EntityQuery, prisma_client: PrismaClient +) -> EntityScopeResolution: + agent_scope: Final = await resolve_agent_daily_activity_scope( + agent_ids=query.entity_ids, + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + ) + return _build_scope( + AGENT_RESOLVER, + query, + agent_scope.agent_ids, + query.exclude_entity_ids, + query.api_key, + agent_scope.agent_metadata, + ) + + +def _user_query( + start_date: str | None = Query(default=None, description="Start date in YYYY-MM-DD format"), + end_date: str | None = Query(default=None, description="End date in YYYY-MM-DD format"), + model: str | None = Query(default=None, description="Filter by specific model"), + api_key: str | None = Query(default=None, description="Filter by specific API key"), + user_id: str | None = Query( + default=None, + description="Filter by specific user ID. Admins can filter by any user or omit for global view. " + "Non-admins must provide their own user_id.", + ), + timezone: int | None = Query( + default=None, + description="Timezone offset in minutes from UTC (e.g., 480 for PST). " + "Matches JavaScript's Date.getTimezoneOffset() convention.", + ), + include_current_utc_day: bool = Query( + default=False, + description="When the range ends on the caller's current local day, extend it to " + "today's UTC bucket so spend written after the caller's local midnight (in UTC " + "terms) is included. Requires the timezone parameter. Historical ranges are " + "never extended.", + ), +) -> EntityQuery: + return EntityQuery( + entity_ids=(user_id,) if user_id is not None else None, + exclude_entity_ids=(), + api_key=api_key, + start_date=start_date, + end_date=end_date, + model=model, + timezone_offset_minutes=timezone, + include_current_utc_day=include_current_utc_day, + ) + + +def _team_query( + team_ids: str | None = None, + start_date: str | None = None, + end_date: str | None = None, + model: str | None = None, + api_key: str | None = None, + exclude_team_ids: str | None = None, + timezone: int | None = None, +) -> EntityQuery: + return EntityQuery( + entity_ids=_query_ids(team_ids), + exclude_entity_ids=_query_excluded_ids(exclude_team_ids), + api_key=api_key, + start_date=start_date, + end_date=end_date, + model=model, + timezone_offset_minutes=timezone, + include_current_utc_day=False, + ) + + +def _tag_query( + start_date: str | None = None, + end_date: str | None = None, + model: str | None = None, + api_key: str | None = None, + tags: str | None = None, + timezone: int | None = None, +) -> EntityQuery: + return EntityQuery( + entity_ids=_query_ids(tags), + exclude_entity_ids=(), + api_key=api_key, + start_date=start_date, + end_date=end_date, + model=model, + timezone_offset_minutes=timezone, + include_current_utc_day=False, + ) + + +def _organization_query( + organization_ids: str | None = None, + start_date: str | None = None, + end_date: str | None = None, + model: str | None = None, + api_key: str | None = None, + exclude_organization_ids: str | None = None, + timezone: int | None = None, +) -> EntityQuery: + return EntityQuery( + entity_ids=_query_ids(organization_ids), + exclude_entity_ids=_query_excluded_ids(exclude_organization_ids), + api_key=api_key, + start_date=start_date, + end_date=end_date, + model=model, + timezone_offset_minutes=timezone, + include_current_utc_day=False, + ) + + +def _customer_query( + end_user_ids: str | None = None, + start_date: str | None = None, + end_date: str | None = None, + model: str | None = None, + api_key: str | None = None, + exclude_end_user_ids: str | None = None, + timezone: int | None = None, +) -> EntityQuery: + return EntityQuery( + entity_ids=_query_ids(end_user_ids), + exclude_entity_ids=_query_excluded_ids(exclude_end_user_ids), + api_key=api_key, + start_date=start_date, + end_date=end_date, + model=model, + timezone_offset_minutes=timezone, + include_current_utc_day=False, + ) + + +def _agent_query( + agent_ids: str | None = None, + start_date: str | None = None, + end_date: str | None = None, + model: str | None = None, + api_key: str | None = None, + exclude_agent_ids: str | None = None, + timezone: int | None = None, +) -> EntityQuery: + return EntityQuery( + entity_ids=_query_ids(agent_ids), + exclude_entity_ids=_query_excluded_ids(exclude_agent_ids), + api_key=api_key, + start_date=start_date, + end_date=end_date, + model=model, + timezone_offset_minutes=timezone, + include_current_utc_day=False, + ) + + +USER_RESOLVER = EntityScopeResolver( + entity="user", + table=DailyActivityTable.USER, + entity_id_field="user_id", + route_prefixes=("/user",), + tags=("Budget & Spend Tracking", "Internal User management"), + query=_user_query, + resolve=_resolve_user, + include_entity_breakdown=False, + operation_names=MappingProxyType( + { + "aggregated": "get_user_daily_activity_aggregated", + "search": "get_user_daily_activity_aggregated_search", + "key_page": "get_user_daily_activity_aggregated_keys", + "model_top_keys": "get_user_daily_activity_model_top_keys", + "export": "get_user_daily_activity_export", + "cache_leakage_keys": "get_user_daily_activity_cache_leakage_keys", + } + ), +) +TEAM_RESOLVER = EntityScopeResolver( + entity="team", + table=DailyActivityTable.TEAM, + entity_id_field="team_id", + route_prefixes=("/team",), + tags=("team management",), + query=_team_query, + resolve=_resolve_team, + include_entity_breakdown=True, + operation_names=MappingProxyType( + { + "aggregated": "get_team_daily_activity_aggregated", + "search": "get_team_daily_activity_aggregated_search", + "key_page": "get_team_daily_activity_aggregated_keys", + "model_top_keys": "get_team_daily_activity_model_top_keys", + "export": "get_team_daily_activity_export", + } + ), +) +TAG_RESOLVER = EntityScopeResolver( + entity="tag", + table=DailyActivityTable.TAG, + entity_id_field="tag", + route_prefixes=("/tag",), + tags=("tag management",), + query=_tag_query, + resolve=_resolve_tag, + include_entity_breakdown=True, + operation_names=MappingProxyType( + { + "aggregated": "get_tag_daily_activity_aggregated", + "search": "get_tag_daily_activity_aggregated_search", + "key_page": "get_tag_daily_activity_aggregated_keys", + "model_top_keys": "get_tag_daily_activity_model_top_keys", + "export": "get_tag_daily_activity_export", + } + ), +) +ORGANIZATION_RESOLVER = EntityScopeResolver( + entity="organization", + table=DailyActivityTable.ORGANIZATION, + entity_id_field="organization_id", + route_prefixes=("/organization",), + tags=("organization management",), + query=_organization_query, + resolve=_resolve_organization, + include_entity_breakdown=True, + operation_names=MappingProxyType( + { + "aggregated": "get_organization_daily_activity_aggregated", + "search": "get_organization_daily_activity_aggregated_search", + "key_page": "get_organization_daily_activity_aggregated_keys", + "model_top_keys": "get_organization_daily_activity_model_top_keys", + "export": "get_organization_daily_activity_export", + } + ), +) +CUSTOMER_RESOLVER = EntityScopeResolver( + entity="customer", + table=DailyActivityTable.CUSTOMER, + entity_id_field="end_user_id", + route_prefixes=("/customer", "/end_user"), + tags=("Customer Management",), + query=_customer_query, + resolve=_resolve_customer, + include_entity_breakdown=True, + operation_names=MappingProxyType( + { + "aggregated": "get_customer_daily_activity_aggregated", + "search": "get_customer_daily_activity_aggregated_search", + "key_page": "get_customer_daily_activity_aggregated_keys", + "model_top_keys": "get_customer_daily_activity_model_top_keys", + "export": "get_customer_daily_activity_export", + } + ), +) +AGENT_RESOLVER = EntityScopeResolver( + entity="agent", + table=DailyActivityTable.AGENT, + entity_id_field="agent_id", + route_prefixes=("/agent",), + tags=("Agent Management",), + query=_agent_query, + resolve=_resolve_agent, + include_entity_breakdown=True, + operation_names=MappingProxyType( + { + "aggregated": "get_agent_daily_activity_aggregated", + "search": "get_agent_daily_activity_aggregated_search", + "key_page": "get_agent_daily_activity_aggregated_keys", + "model_top_keys": "get_agent_daily_activity_model_top_keys", + "export": "get_agent_daily_activity_export", + } + ), +) diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index f6fe58e25b9..04b8ec56ae2 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -54,8 +54,9 @@ from litellm.proxy.hooks.user_management_event_hooks import UserManagementEventH from litellm.proxy.management.teams.access import is_team_admin from litellm.proxy.management_endpoints.common_daily_activity import ( DailySpendRecord, + ScopeDenied, get_daily_activity, - get_daily_activity_aggregated, + raise_public, ) from litellm.proxy.management_endpoints.common_utils import ( _user_has_admin_view, @@ -2862,6 +2863,18 @@ async def ui_view_users( # Using shared metric helper implementations from common_daily_activity +def resolve_user_daily_activity_entity_ids( + *, user_id: str | None, user_api_key_dict: UserAPIKeyAuth +) -> tuple[str, ...] | None | ScopeDenied: + if _user_has_admin_view(user_api_key_dict): + return (user_id,) if user_id is not None else None + + caller_user_id: Final = require_caller_user_id_for_non_admin(user_api_key_dict) + if user_id is not None and user_id != caller_user_id: + return ScopeDenied(403, "Non-admin users can only view their own spend data.") + return (caller_user_id,) + + async def _resolve_user_email_metadata( prisma_client: "PrismaClient", records: Sequence[DailySpendRecord] ) -> dict[str, dict]: @@ -2956,20 +2969,13 @@ async def get_user_daily_activity( ) try: - is_admin: Final = _user_has_admin_view(user_api_key_dict) - - if is_admin: - entity_id = user_id # None means global view, otherwise filter by user - else: - caller_user_id: Final = require_caller_user_id_for_non_admin(user_api_key_dict) - if user_id is None: - user_id = caller_user_id - if user_id != caller_user_id: - raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail={"error": "Non-admin users can only view their own spend data."}, - ) - entity_id = user_id + resolved_entity_ids: Final = resolve_user_daily_activity_entity_ids( + user_id=user_id, + user_api_key_dict=user_api_key_dict, + ) + if isinstance(resolved_entity_ids, ScopeDenied): + raise_public(resolved_entity_ids) + entity_id: Final[str | None] = resolved_entity_ids[0] if resolved_entity_ids is not None else None return await get_daily_activity( prisma_client=prisma_client, @@ -2996,108 +3002,3 @@ async def get_user_daily_activity( status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail={"error": f"Failed to fetch analytics: {e}"}, ) - - -@router.get( - "/user/daily/activity/aggregated", - tags=["Budget & Spend Tracking", "Internal User management"], - dependencies=[Depends(user_api_key_auth)], - response_model=SpendAnalyticsPaginatedResponse, -) -@management_endpoint_wrapper -async def get_user_daily_activity_aggregated( - start_date: str | None = fastapi.Query( - default=None, - description="Start date in YYYY-MM-DD format", - ), - end_date: str | None = fastapi.Query( - default=None, - description="End date in YYYY-MM-DD format", - ), - model: str | None = fastapi.Query( - default=None, - description="Filter by specific model", - ), - api_key: str | None = fastapi.Query( - default=None, - description="Filter by specific API key", - ), - user_id: str | None = fastapi.Query( - default=None, - description="Filter by specific user ID. Admins can filter by any user or omit for global view. Non-admins must provide their own user_id.", - ), - timezone: int | None = fastapi.Query( - default=None, - description="Timezone offset in minutes from UTC (e.g., 480 for PST). " - "Matches JavaScript's Date.getTimezoneOffset() convention.", - ), - include_current_utc_day: bool = fastapi.Query( - default=False, - description="When the range ends on the caller's current local day, extend it to " - "today's UTC bucket so spend written after the caller's local midnight (in UTC " - "terms) is included. Requires the timezone parameter. Historical ranges are " - "never extended.", - ), - user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), -) -> SpendAnalyticsPaginatedResponse: - """ - Aggregated analytics for a user's daily activity without pagination. - Returns the same response shape as the paginated endpoint with page metadata set to single-page. - - Reads daily spend records that only ever accumulate and are never affected by budget - resets. Their total can legitimately exceed the `spend` field returned by - `/v2/user/info`, which is a running budget counter that every budget reset sets back - to zero (or to the overage above `max_budget` when `budget_rollover` is enabled). - """ - from litellm.proxy.proxy_server import prisma_client - - if prisma_client is None: - raise HTTPException( - status_code=500, - detail={"error": CommonProxyErrors.db_not_connected_error.value}, - ) - - if start_date is None or end_date is None: - raise HTTPException( - status_code=status.HTTP_400_BAD_REQUEST, - detail={"error": "Please provide start_date and end_date"}, - ) - - try: - is_admin: Final = _user_has_admin_view(user_api_key_dict) - - if is_admin: - entity_id = user_id # None means global view, otherwise filter by user - else: - caller_user_id: Final = require_caller_user_id_for_non_admin(user_api_key_dict) - if user_id is None: - user_id = caller_user_id - if user_id != caller_user_id: - raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail={"error": "Non-admin users can only view their own spend data."}, - ) - entity_id = user_id - - return await get_daily_activity_aggregated( - prisma_client=prisma_client, - table_name="litellm_dailyuserspend", - entity_id_field="user_id", - entity_id=entity_id, - entity_metadata_field=None, - start_date=start_date, - end_date=end_date, - model=model, - api_key=api_key, - timezone_offset_minutes=timezone, - include_current_utc_day=include_current_utc_day, - ) - - except HTTPException: - raise - except Exception as e: - verbose_proxy_logger.exception("/user/daily/activity/aggregated: Exception occured - %s", e) - raise HTTPException( - status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, - detail={"error": f"Failed to fetch analytics: {e}"}, - ) diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 0974a2a952d..8ed8ec8752e 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -184,6 +184,7 @@ if MCP_AVAILABLE: from litellm.proxy._experimental.mcp_server.ui_session_utils import ( admitted_user_context, build_effective_auth_contexts, + granted_toolset_ids, is_ui_session_credential, ) from litellm.proxy._types import ( @@ -3332,18 +3333,15 @@ if MCP_AVAILABLE: ): """Return toolsets the calling key is allowed to access.""" prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") - is_admin: Final = _user_has_admin_view(user_api_key_dict) - op: Final = user_api_key_dict.object_permission - # mcp_toolsets=None or [] both mean "not restricted by toolsets". - # For admins: either value → no restriction → return all. - # For non-admins: either value → no toolsets explicitly granted → return nothing. - # (An admin whose DB row has mcp_toolsets=[] should still see all toolsets.) - raw_toolsets: Final = getattr(op, "mcp_toolsets", None) if op else None - if not raw_toolsets: - if is_admin: + if _user_has_admin_view(user_api_key_dict): + op: Final = user_api_key_dict.object_permission + if op is None or not op.mcp_toolsets: return await list_mcp_toolsets(prisma_client) + return await list_mcp_toolsets(prisma_client, toolset_ids=op.mcp_toolsets) + granted: Final = await granted_toolset_ids(user_api_key_dict) + if not granted: return [] - return await list_mcp_toolsets(prisma_client, toolset_ids=raw_toolsets) + return await list_mcp_toolsets(prisma_client, toolset_ids=sorted(granted)) @router.get( "/toolset/{toolset_id}", @@ -3355,15 +3353,13 @@ if MCP_AVAILABLE: user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): prisma_client: Final = get_prisma_client_or_throw("Database not connected. Connect a database to your proxy") - # Non-admin keys may only fetch toolsets they've been explicitly granted. - if not _user_has_admin_view(user_api_key_dict): - op: Final = user_api_key_dict.object_permission - granted: Final = getattr(op, "mcp_toolsets", None) if op else None - if granted is None or toolset_id not in granted: - raise HTTPException( - status_code=status.HTTP_403_FORBIDDEN, - detail={"error": "API key does not have access to this toolset."}, - ) + if not _user_has_admin_view(user_api_key_dict) and toolset_id not in await granted_toolset_ids( + user_api_key_dict + ): + raise HTTPException( + status_code=status.HTTP_403_FORBIDDEN, + detail={"error": "API key does not have access to this toolset."}, + ) toolset: Final = await get_mcp_toolset(prisma_client, toolset_id) if toolset is None: raise HTTPException( diff --git a/litellm/proxy/management_endpoints/organization_endpoints.py b/litellm/proxy/management_endpoints/organization_endpoints.py index 24bbd2b4b1f..c8ae7af41db 100644 --- a/litellm/proxy/management_endpoints/organization_endpoints.py +++ b/litellm/proxy/management_endpoints/organization_endpoints.py @@ -14,10 +14,12 @@ Endpoints for /organization operations #### ORGANIZATION MANAGEMENT #### from collections.abc import Mapping, Sequence +from types import MappingProxyType from typing import ( TYPE_CHECKING, Annotated, Final, + NamedTuple, Protocol, cast, # noqa: TID251 # prisma types Json columns as fields.Json but reads back plain python values overload, @@ -62,6 +64,7 @@ from litellm.proxy.management_helpers.utils import ( ) from litellm.proxy.utils import PrismaClient, ProxyLogging from litellm.repositories.budget_repository import BudgetRepository +from litellm.repositories.chunked_in import find_many_in from litellm.repositories.object_permission_repository import ObjectPermissionRepository from litellm.repositories.organization_repository import OrganizationRepository from litellm.repositories.table_repositories import OrganizationMembershipRepository @@ -583,43 +586,23 @@ async def get_organization_daily_activity( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - # Parse comma-separated ids - org_ids_list = organization_ids.split(",") if organization_ids else None + org_ids: Final = tuple(organization_ids.split(",")) if organization_ids else None exclude_org_ids_list: list[str] | None = None if exclude_organization_ids: exclude_org_ids_list = exclude_organization_ids.split(",") if exclude_organization_ids else None - # Restrict non-proxy-admins to only organizations where they are org_admin - if not _user_has_admin_view(user_api_key_dict): - memberships: Final = await _table(OrganizationMembershipRepository(prisma_client)).find_many( - where={"user_id": user_api_key_dict.user_id} - ) - admin_org_ids = [m.organization_id for m in memberships if m.user_role == LitellmUserRoles.ORG_ADMIN.value] - if org_ids_list is None: - # Default to orgs where user is org_admin - org_ids_list = admin_org_ids - else: - # Ensure user is org_admin for all requested orgs - for org_id in org_ids_list: - if org_id not in admin_org_ids: - raise HTTPException( - status_code=403, - detail={"error": f"User is not org_admin for Organization= {org_id}."}, - ) + org_scope: Final = await resolve_organization_daily_activity_scope( + organization_ids=org_ids, + prisma_client=prisma_client, + user_api_key_dict=user_api_key_dict, + ) - # Fetch organization aliases for metadata - where_condition: Final = _STR_OBJECT_DICT_ADAPTER.validate_python({}) - if org_ids_list is not None: - where_condition["organization_id"] = {"in": list(org_ids_list)} - org_aliases: Final = await _table(OrganizationRepository(prisma_client)).find_many(where=where_condition) - - # Query daily activity for organizations return await get_daily_activity( prisma_client=prisma_client, table_name="litellm_dailyorganizationspend", entity_id_field="organization_id", - entity_id=org_ids_list, - entity_metadata_field={o.organization_id: {"organization_alias": o.organization_alias} for o in org_aliases}, + entity_id=None if org_scope.organization_ids is None else list(org_scope.organization_ids), + entity_metadata_field=org_scope.organization_metadata, exclude_entity_ids=exclude_org_ids_list, start_date=start_date, end_date=end_date, @@ -630,6 +613,56 @@ async def get_organization_daily_activity( ) +class _OrganizationDailyActivityScope(NamedTuple): + organization_ids: tuple[str, ...] | None + organization_metadata: Mapping[str, dict[str, object]] + + +async def resolve_organization_daily_activity_scope( + *, + organization_ids: tuple[str, ...] | None, + prisma_client: PrismaClient, + user_api_key_dict: UserAPIKeyAuth, +) -> _OrganizationDailyActivityScope: + is_admin: Final = _user_has_admin_view(user_api_key_dict) + memberships: Final = ( + await _table(OrganizationMembershipRepository(prisma_client)).find_many( + where={"user_id": user_api_key_dict.user_id} + ) + if not is_admin + else () + ) + admin_organization_ids: Final = tuple( + membership.organization_id + for membership in memberships + if membership.user_role == LitellmUserRoles.ORG_ADMIN.value + ) + if not is_admin and organization_ids is not None: + for organization_id in organization_ids: + if organization_id not in admin_organization_ids: + raise HTTPException( + status_code=403, + detail={"error": f"User is not org_admin for Organization= {organization_id}."}, + ) + resolved_organization_ids: Final[tuple[str, ...] | None] = ( + organization_ids if is_admin or organization_ids is not None else admin_organization_ids + ) + + organization_table: Final = _table(OrganizationRepository(prisma_client)) + organization_rows: Final = ( + await find_many_in(organization_table, "organization_id", resolved_organization_ids) + if resolved_organization_ids is not None + else await organization_table.find_many(where={}) + ) + metadata: Final = MappingProxyType( + { + organization.organization_id: {"organization_alias": organization.organization_alias} + for organization in organization_rows + } + ) + return _OrganizationDailyActivityScope(resolved_organization_ids, metadata) + + async def _set_object_permission( data: NewOrganizationRequest, prisma_client: PrismaClient | None, diff --git a/litellm/proxy/management_endpoints/tag_management_endpoints.py b/litellm/proxy/management_endpoints/tag_management_endpoints.py index ab33d4bd766..5bf16379d05 100644 --- a/litellm/proxy/management_endpoints/tag_management_endpoints.py +++ b/litellm/proxy/management_endpoints/tag_management_endpoints.py @@ -190,7 +190,7 @@ async def _get_tag_list_scope( return {"api_key": {"in": scoped_api_keys}} -async def _get_tag_daily_activity_api_key_filter( +async def get_tag_daily_activity_api_key_filter( prisma_client: "PrismaClient", user_api_key_dict: UserAPIKeyAuth, requested_api_key: str | None, @@ -757,7 +757,7 @@ async def get_tag_daily_activity( # Convert comma-separated tags string to list if provided tag_list: Final = tags.split(",") if tags else None - scoped_api_key_filter: Final = await _get_tag_daily_activity_api_key_filter( + scoped_api_key_filter: Final = await get_tag_daily_activity_api_key_filter( prisma_client=prisma_client, user_api_key_dict=user_api_key_dict, requested_api_key=api_key, diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 86b37f31090..8943f5a7416 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -125,7 +125,8 @@ from litellm.proxy.hooks.model_max_budget_limiter import ( from litellm.proxy.management.teams.access import TEAM_OR_ORG_ADMIN, TeamRole, is_team_admin, team_access_denied from litellm.proxy.management.teams.dependencies import get_team_access from litellm.proxy.management_endpoints.common_daily_activity import ( - get_daily_activity_aggregated, + InvalidDateRange, + parse_canonical_date_range, ) from litellm.proxy.management_endpoints.common_utils import ( _check_disable_global_guardrails_caller_permission, @@ -3466,6 +3467,12 @@ async def team_member_add( litellm_proxy_admin_name=litellm_proxy_admin_name, ) + await delete_cache_team_object( + team_id=data.team_id, + team_alias=complete_team_data.team_alias, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) await evict_and_broadcast( cache_keys=tuple(sorted(user.user_id for user in updated_users)), user_api_key_cache=user_api_key_cache, @@ -6515,7 +6522,7 @@ class _TeamDailyActivityScope(NamedTuple): api_key_filter: str | list[str] | None # mutable-ok: downstream daily-activity signatures take str | list unions -async def _resolve_team_daily_activity_scope( +async def resolve_team_daily_activity_scope( *, team_ids: str | None, exclude_team_ids: str | None, @@ -6598,11 +6605,14 @@ async def _resolve_team_daily_activity_scope( user_api_keys = [key.token for key in user_keys if key.token] # If user has no API keys, return empty result if not user_api_keys: - user_api_keys = [""] # Use empty string to ensure no matches + user_api_keys = [] - # If api_key parameter is provided, use it; otherwise use user_api_keys if set - final_api_key_filter: str | list[str] | None = api_key - if final_api_key_filter is None and user_api_keys is not None: + final_api_key_filter: str | list[str] | None + if user_api_keys is None: + final_api_key_filter = api_key + elif api_key: + final_api_key_filter = api_key if api_key in user_api_keys else [] + else: final_api_key_filter = user_api_keys return _TeamDailyActivityScope( @@ -6653,7 +6663,7 @@ async def get_team_daily_activity( if prisma_client is None: raise _daily_activity_error(status_code=500, message=CommonProxyErrors.db_not_connected_error.value) - scope: Final = await _resolve_team_daily_activity_scope( + scope: Final = await resolve_team_daily_activity_scope( team_ids=team_ids, exclude_team_ids=exclude_team_ids, api_key=api_key, @@ -6682,95 +6692,19 @@ async def get_team_daily_activity( _MAX_AGGREGATED_RANGE_DAYS: Final = 400 -def _aggregated_date_range_error(start_date: str | None, end_date: str | None) -> str | None: +def aggregated_date_range_error(start_date: str | None, end_date: str | None) -> str | None: """The aggregated endpoint has no pagination to bound its work, so malformed dates and ranges wider than the UI ever requests are rejected before querying.""" - if start_date is None or end_date is None: - return "Please provide start_date and end_date" - try: - parsed_start: Final = datetime.strptime(start_date, "%Y-%m-%d").replace(tzinfo=timezone.utc) - parsed_end: Final = datetime.strptime(end_date, "%Y-%m-%d").replace(tzinfo=timezone.utc) - except ValueError: - return "start_date and end_date must be valid YYYY-MM-DD dates" - if parsed_end < parsed_start: + date_range: Final = parse_canonical_date_range(start_date, end_date) + if isinstance(date_range, InvalidDateRange): + return date_range.reason + if date_range.end < date_range.start: return "end_date must be on or after start_date" - if (parsed_end - parsed_start).days > _MAX_AGGREGATED_RANGE_DAYS: + if (date_range.end - date_range.start).days > _MAX_AGGREGATED_RANGE_DAYS: return f"Date range must be at most {_MAX_AGGREGATED_RANGE_DAYS} days" return None -@router.get( - "/team/daily/activity/aggregated", - response_model=SpendAnalyticsPaginatedResponse, - tags=["team management"], -) -async def get_team_daily_activity_aggregated( - user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], - team_ids: str | None = None, - start_date: str | None = None, - end_date: str | None = None, - model: str | None = None, - api_key: str | None = None, - exclude_team_ids: str | None = None, - timezone: int | None = None, -): - """ - Aggregated daily activity for teams without pagination, including per-team breakdown. - - One SQL GROUPING SETS pass returns every day in the range regardless of row - volume, so callers never reassemble pages. Same response shape as the - paginated endpoint with page metadata pinned to a single page. - - Args: - team_ids (Optional[str]): Comma-separated list of team IDs to filter by. If not provided, returns data for all teams. - start_date (Optional[str]): Start date for the activity period (YYYY-MM-DD). - end_date (Optional[str]): End date for the activity period (YYYY-MM-DD). - model (Optional[str]): Filter by model name. - api_key (Optional[str]): Filter by API key. - exclude_team_ids (Optional[str]): Comma-separated list of team IDs to exclude. - timezone (Optional[int]): Timezone offset in minutes from UTC, matching JavaScript's Date.getTimezoneOffset() convention. - Returns: - SpendAnalyticsPaginatedResponse: Response containing all daily activity data for the range. - """ - from litellm.proxy.proxy_server import ( - prisma_client, - proxy_logging_obj, - user_api_key_cache, - ) - - if prisma_client is None: - raise _daily_activity_error(status_code=500, message=CommonProxyErrors.db_not_connected_error.value) - - range_error: Final = _aggregated_date_range_error(start_date, end_date) - if range_error is not None: - raise _daily_activity_error(status_code=400, message=range_error) - - scope: Final = await _resolve_team_daily_activity_scope( - team_ids=team_ids, - exclude_team_ids=exclude_team_ids, - api_key=api_key, - user_api_key_dict=user_api_key_dict, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - proxy_logging_obj=proxy_logging_obj, - ) - - return await get_daily_activity_aggregated( - prisma_client=prisma_client, - table_name="litellm_dailyteamspend", - entity_id_field="team_id", - entity_id=scope.team_ids, - entity_metadata_field=scope.team_alias_metadata, - start_date=start_date, - end_date=end_date, - model=model, - api_key=scope.api_key_filter, - exclude_entity_ids=scope.exclude_team_ids, - timezone_offset_minutes=timezone, - include_entity_breakdown=True, - ) - - def _team_user_spend_sql(*, team_count: int, restrict_to_user: bool) -> str: team_placeholders: Final = ", ".join(f"${i}" for i in range(3, 3 + team_count)) user_clause: Final = f' AND sl."user" = ${3 + team_count}' if restrict_to_user else "" @@ -6838,14 +6772,14 @@ async def get_team_spend_by_user( if prisma_client is None: raise _daily_activity_error(status_code=500, message=CommonProxyErrors.db_not_connected_error.value) - range_error: Final = _aggregated_date_range_error(start_date, end_date) + range_error: Final = aggregated_date_range_error(start_date, end_date) if range_error is not None or start_date is None or end_date is None: raise _daily_activity_error(status_code=400, message=range_error or "Please provide start_date and end_date") if not team_ids: raise _daily_activity_error(status_code=400, message="Please provide team_ids") - scope: Final = await _resolve_team_daily_activity_scope( + scope: Final = await resolve_team_daily_activity_scope( team_ids=team_ids, exclude_team_ids=None, api_key=None, diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 19444dfe33b..0d7094f66dc 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -1853,10 +1853,36 @@ def _should_use_role_from_sso_response(sso_role: str | None) -> bool: return True +class _SsoUserNames(Protocol): + id: str | None + display_name: str | None + first_name: str | None + last_name: str | None + + +def _get_sso_user_alias(result: _SsoUserNames | Mapping[str, object] | None) -> str | None: + """Display name the IdP sent for the user, falling back to the joined first/last name.""" + if result is None: + return None + if isinstance(result, Mapping): + raw_names: tuple[object, ...] = tuple( + result.get(key) for key in ("id", "display_name", "first_name", "last_name") + ) + else: + raw_names = (result.id, result.display_name, result.first_name, result.last_name) + user_id, display_name, first_name, last_name = ( + name.strip() or None if isinstance(name, str) else None for name in raw_names + ) + if display_name and display_name != user_id: + return display_name + return " ".join(part for part in (first_name, last_name) if part) or None + + def _build_sso_user_update_data( - result: Union["CustomOpenID", OpenID, dict] | None, + result: Union["CustomOpenID", OpenID, Mapping[str, object]] | None, user_email: str | None, user_id: str | None, + existing_user_alias: str | None = None, ) -> dict[str, object]: """ Build the update data dictionary for SSO user upsert. @@ -1865,14 +1891,19 @@ def _build_sso_user_update_data( result: The SSO response containing user information user_email: The user's email from SSO user_id: The user's ID for logging purposes + existing_user_alias: The user's current alias in the DB; only an empty alias is filled from SSO Returns: - dict: Update data containing user_email and optionally user_role if valid + dict: Update data containing user_email, user_alias when newly available, and user_role if valid """ - update_data: Final[dict[str, object]] = {"user_email": normalize_email(user_email)} + sso_user_alias: Final = None if existing_user_alias else _get_sso_user_alias(result) + update_data: Final[dict[str, object]] = { + "user_email": normalize_email(user_email), + **({"user_alias": sso_user_alias} if sso_user_alias is not None else {}), + } # Get SSO role from result and include if valid - sso_role: Final = getattr(result, "user_role", None) + sso_role: Final = result.user_role if isinstance(result, CustomOpenID) else None if sso_role is not None: # Convert enum to string if needed sso_role_str: Final = sso_role.value if isinstance(sso_role, LitellmUserRoles) else sso_role @@ -2616,6 +2647,7 @@ async def insert_sso_user( new_user_request: Final = NewUserRequest( user_id=user_defined_values["user_id"], user_email=normalize_email(user_defined_values["user_email"]), + user_alias=_get_sso_user_alias(result_openid), user_role=user_defined_values["user_role"], max_budget=user_defined_values["max_budget"], budget_duration=user_defined_values["budget_duration"], @@ -3249,6 +3281,7 @@ class SSOAuthenticationHandler: result=result, user_email=user_email, user_id=user_id, + existing_user_alias=user_info.user_alias if isinstance(user_info, LiteLLM_UserTable) else None, ) await _user_meta_db(UserRepository(prisma_client)).update_many( @@ -3280,7 +3313,7 @@ class SSOAuthenticationHandler: if user_info is None: verbose_proxy_logger.debug("User not found in LiteLLM DB, skipping team member addition") return - sso_teams: Final = getattr(result, "team_ids", []) + sso_teams: Final = result.team_ids if isinstance(result, CustomOpenID) else [] await add_missing_team_member(user_info=user_info, sso_teams=sso_teams) @staticmethod diff --git a/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py b/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py index 1265da99d89..3c2300f14fd 100644 --- a/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py +++ b/litellm/proxy/management_endpoints/usage_endpoints/ai_usage_chat.py @@ -8,11 +8,13 @@ from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Callable, from datetime import date from typing import Final, Literal, NamedTuple, Protocol, cast, overload +from fastapi import HTTPException from typing_extensions import ReadOnly, TypedDict import litellm from litellm._logging import verbose_proxy_logger from litellm.constants import DEFAULT_COMPETITOR_DISCOVERY_MODEL +from litellm.proxy._types import CommonProxyErrors from litellm.types.proxy.management_endpoints.common_daily_activity import ( SpendAnalyticsPaginatedResponse, ) @@ -259,22 +261,34 @@ async def _query_activity( ) -> SpendAnalyticsPaginatedResponse: """Shared helper that calls the daily activity query layer.""" from litellm.proxy.management_endpoints.common_daily_activity import ( + daily_activity_repository, + daily_activity_scope, get_daily_activity, get_daily_activity_aggregated, ) from litellm.proxy.proxy_server import prisma_client if use_aggregated: + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": CommonProxyErrors.db_not_connected_error.value}, + ) + repository: Final = daily_activity_repository(prisma_client) + scope: Final = daily_activity_scope( + table_name, + entity_id_field, + entity_id, + None, + None, + start_date, + end_date, + None, + None, + ) return await get_daily_activity_aggregated( - prisma_client=prisma_client, - table_name=table_name, - entity_id_field=entity_id_field, - entity_id=entity_id, - entity_metadata_field=None, - start_date=start_date, - end_date=end_date, - model=None, - api_key=None, + repository, + scope, ) return await get_daily_activity( prisma_client=prisma_client, diff --git a/litellm/proxy/mcp_registry.json b/litellm/proxy/mcp_registry.json index b117f35600d..70e81f127c9 100644 --- a/litellm/proxy/mcp_registry.json +++ b/litellm/proxy/mcp_registry.json @@ -217,6 +217,17 @@ {"name": "GOOGLE_CLIENT_SECRET", "description": "Google OAuth Client Secret", "secret": true} ] }, + { + "name": "microsoft_365", + "title": "Microsoft 365 (Graph)", + "description": "Outlook mail and calendar, OneDrive and SharePoint files, and Teams through Microsoft Graph, with each user's own Entra ID sign-in. Self-hosted: run ms-365-mcp-server next to the proxy and point the URL at it", + "icon_url": "/ui/assets/logos/microsoft_365.svg", + "category": "Productivity", + "registry_url": "https://registry.modelcontextprotocol.io/servers/io.github.Softeria%2Fms-365-mcp-server", + "transport": "http", + "url": "http://localhost:3000/mcp", + "env_vars": [] + }, { "name": "obsidian", "title": "Obsidian", diff --git a/litellm/proxy/prisma_migration.py b/litellm/proxy/prisma_migration.py index 7e3aff75cef..3561f190808 100644 --- a/litellm/proxy/prisma_migration.py +++ b/litellm/proxy/prisma_migration.py @@ -1,9 +1,9 @@ """Standalone entrypoint for applying database migrations and generating the Prisma client. -Migration failures fail the entrypoint by default; set ENFORCE_PRISMA_MIGRATION_CHECK=false -for log-only behavior. A failed 'prisma generate' is always log-only: every shipped image -bakes the client at build time, and refreshing it writes into site-packages, which an -arbitrary non-root uid or a read-only root filesystem cannot do. +A failed migration fails the entrypoint, the same way it fails proxy startup. A failed +'prisma generate' is log-only: every shipped image bakes the client at build time, and +refreshing it writes into site-packages, which an arbitrary non-root uid or a read-only +root filesystem cannot do. """ import os @@ -18,17 +18,10 @@ from litellm_proxy_extras.prisma_toolchain import resolve_prisma_argv from litellm._logging import verbose_proxy_logger from litellm.proxy.proxy_cli import run_server -from litellm.secret_managers.main import str_to_bool def main() -> int: - enforce_prisma_migration_check: Final = str_to_bool(os.getenv("ENFORCE_PRISMA_MIGRATION_CHECK")) is not False - run_server_args: Final = ( - ("--skip_server_startup", "--enforce_prisma_migration_check") - if enforce_prisma_migration_check - else ("--skip_server_startup",) - ) - run_server(run_server_args, standalone_mode=False) + run_server(("--skip_server_startup",), standalone_mode=False) verbose_proxy_logger.info("Running 'prisma generate'...") result: Final = subprocess.run(resolve_prisma_argv(("prisma", "generate")), capture_output=True, text=True) diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 27c03d2d5d7..cb65cae79a4 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -966,8 +966,11 @@ class ProxyInitializationHelpers: "--enforce_prisma_migration_check", is_flag=True, default=False, - help="Exit with error if database migration fails on startup.", - envvar="ENFORCE_PRISMA_MIGRATION_CHECK", + hidden=True, + help=( + "Deprecated and ignored: the proxy always exits when database setup fails at " + "startup. It is still accepted so existing commands keep working." + ), ) @click.option( "--use_v2_migration_resolver", @@ -1098,6 +1101,12 @@ def run_server( if validate_config is True: ProxyInitializationHelpers._run_config_validation(config) return + if enforce_prisma_migration_check: + print( + "\033[1;33mLiteLLM Proxy: --enforce_prisma_migration_check is " + "deprecated and has no effect, because the proxy always exits " + "when database setup fails at startup. You can safely remove it.\033[0m" + ) if model and "ollama" in model and api_base is None: ProxyInitializationHelpers._run_ollama_serve() if health is True: @@ -1410,10 +1419,15 @@ def run_server( "LiteLLM versions contend for the same DB.\033[0m" ) try: - setup_ok: Final = PrismaManager.setup_database( + migrated: Final = PrismaManager.setup_database( use_migrate=not use_prisma_db_push, use_v2_resolver=use_v2_resolver, ) + setup_ok: Final = migrated and ( + not skip_server_startup or PrismaManager.build_request_log_indexes() + ) + if migrated and not skip_server_startup: + PrismaManager.start_request_log_index_build() except RuntimeError as e: # Raised on unrecoverable migration errors: the v2 # resolver's non-idempotent failures and permission @@ -1426,17 +1440,11 @@ def run_server( ) sys.exit(2) if not setup_ok: - if enforce_prisma_migration_check: - print( - "\033[1;31mLiteLLM Proxy: Database setup failed after multiple retries. " - "The proxy cannot start safely. Please check your database connection and migration status.\033[0m" - ) - sys.exit(1) - else: - print( - "\033[1;33mLiteLLM Proxy: Database migration failed but continuing startup. " - "Set --enforce_prisma_migration_check or ENFORCE_PRISMA_MIGRATION_CHECK=true to exit on failure.\033[0m" - ) + print( + "\033[1;31mLiteLLM Proxy: Database setup failed after multiple retries. " + "The proxy cannot start safely. Please check your database connection and migration status.\033[0m" + ) + sys.exit(1) else: print( "Unable to connect to DB. DATABASE_URL found in environment, but the prisma CLI is neither on " diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 9abf949ec1e..9282f216a52 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -604,6 +604,7 @@ from litellm.proxy.management_endpoints.cost_tracking_settings import ( from litellm.proxy.management_endpoints.customer_endpoints import ( router as customer_router, ) +from litellm.proxy.management_endpoints.daily_activity_routes import router as daily_activity_router from litellm.proxy.management_endpoints.fallback_management_endpoints import ( router as fallback_management_router, ) @@ -19951,6 +19952,7 @@ app.include_router(pass_through_router) app.include_router(health_router) app.include_router(key_management_router) app.include_router(internal_user_router) +app.include_router(daily_activity_router) app.include_router(password_management_router) app.include_router(session_management_router) app.include_router(team_router) diff --git a/litellm/proxy/tracing_endpoints.py b/litellm/proxy/tracing_endpoints.py index 46d29c50c1b..e4d46560f08 100644 --- a/litellm/proxy/tracing_endpoints.py +++ b/litellm/proxy/tracing_endpoints.py @@ -15,12 +15,15 @@ from types import MappingProxyType from typing import Annotated, Final from fastapi import APIRouter, Depends, HTTPException, Query, Request, Response +from pydantic import BaseModel, ConfigDict +from litellm._logging import verbose_proxy_logger from litellm.constants import OTLP_RETRY_AFTER_SECONDS from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.http_parsing_utils import is_otlp_trace_request from litellm.proxy.tracing_runtime import provide_receiver, require_receiver +from litellm.rust_bridge.traces import ClickHouseStorage, QueryScope from litellm.tracing import ( Tenant, TraceReceiver, @@ -117,7 +120,7 @@ async def ingest_otlp_traces( return Response(content=body, media_type=media_type) -@router.get("/v1/traces", response_model=None) +@router.get("/v1/traces", response_model=TracePage) async def list_agent_traces( context: Annotated[TraceAccessContext, Depends(provide_trace_access)], start_ms: Annotated[int | None, Query(description="Window start, unix ms. Default: 24h ago")] = None, @@ -137,7 +140,78 @@ async def list_agent_traces( raise HTTPException(status_code=400, detail=str(error)) from error -@router.get("/v1/traces/{trace_id}", response_model=None) +class TraceQueryRequest(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid") + sql: str + + +@dataclass(frozen=True, slots=True) +class TraceQueryAccess: + storage: ClickHouseStorage + scope: QueryScope + secret: str + + +def provide_trace_query_secret() -> str: + from litellm.proxy.proxy_server import master_key + + if not master_key: + raise HTTPException(status_code=503, detail="Trace SQL queries require a configured proxy master key") + return master_key + + +def trace_query_scope(auth: UserAPIKeyAuth) -> QueryScope: + if auth.user_role in (LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY): + return {"kind": "admin"} + if auth.project_id and auth.token: + return {"kind": "key", "team_id": auth.team_id or "", "api_key_hash": auth.token} + if auth.project_id: + raise HTTPException(status_code=403, detail="Project trace SQL queries require a project key") + if auth.team_id: + return {"kind": "team", "team_id": auth.team_id} + if auth.token: + return {"kind": "key", "team_id": "", "api_key_hash": auth.token} + raise HTTPException(status_code=403, detail="Trace SQL queries require an authenticated trace scope") + + +async def provide_trace_query_access( + auth: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + tracing: Annotated[TraceReceiver | None, Depends(provide_receiver)], + secret: Annotated[str, Depends(provide_trace_query_secret)], +) -> TraceQueryAccess: + return TraceQueryAccess(require_receiver(tracing).store.storage, trace_query_scope(auth), secret) + + +@router.post("/v1/traces/query") +async def query_agent_traces( + body: TraceQueryRequest, + access: Annotated[TraceQueryAccess, Depends(provide_trace_query_access)], +) -> Response: + try: + return Response( + content=await access.storage.query_sql(body.sql, access.scope, access.secret), media_type="application/json" + ) + except ValueError as error: + raise HTTPException(status_code=400, detail=str(error)) from error + except RuntimeError as error: + verbose_proxy_logger.warning("Trace SQL query unavailable: %s", error) + raise HTTPException(status_code=503, detail="Trace SQL query failed or exceeded reader limits") from error + + +@router.get("/v1/traces/query/help") +async def help_agent_trace_queries( + access: Annotated[TraceQueryAccess, Depends(provide_trace_query_access)], +) -> Response: + try: + return Response( + content=await access.storage.query_help(access.scope, access.secret), media_type="application/json" + ) + except RuntimeError as error: + verbose_proxy_logger.warning("Trace query help unavailable: %s", error) + raise HTTPException(status_code=503, detail="Trace query help is temporarily unavailable") from error + + +@router.get("/v1/traces/{trace_id}", response_model=Trace) async def get_agent_trace( trace_id: str, context: Annotated[TraceAccessContext, Depends(provide_trace_access)], @@ -150,7 +224,7 @@ async def get_agent_trace( return trace -@router.get("/v1/traces/{trace_id}/spans/{span_id}", response_model=None) +@router.get("/v1/traces/{trace_id}/spans/{span_id}", response_model=SpanDetail) async def get_agent_trace_span( trace_id: str, span_id: str, diff --git a/litellm/repositories/daily_activity_repository.py b/litellm/repositories/daily_activity_repository.py new file mode 100644 index 00000000000..2e34582091a --- /dev/null +++ b/litellm/repositories/daily_activity_repository.py @@ -0,0 +1,311 @@ +import asyncio +from collections.abc import AsyncIterator, Mapping, Sequence +from datetime import datetime +from itertools import groupby +from types import MappingProxyType +from typing import Final, Protocol + +from pydantic import StrictStr, TypeAdapter, ValidationError +from typing_extensions import assert_never + +from litellm import constants +from litellm._logging import verbose_proxy_logger +from litellm.repositories.chunked_in import find_many_in +from litellm.repositories.daily_activity_sql import ( + ExportCursor, + SqlQuery, + adjust_dates_for_timezone, + build_aggregated_sql, + build_cache_leakage_keys_sql, + build_entity_rollup_sql, + build_export_sql, + build_key_page_sql, + build_key_search_sql, + build_model_top_keys_sql, +) +from litellm.repositories.prisma_protocols import TableActions +from litellm.types.repositories.daily_activity import ( + AggregatedRows, + DailyActivityProxyReads, + DailyActivityRow, + DailyActivityScope, + DailyActivityTable, + DailyRowsPage, + EntityRollupRow, + ExportRow, + ExportType, + GroupingSetsRow, + KeyMetadataRow, + KeyPage, + KeySpendRow, + SpendLogsWindow, +) + + +class _VerificationTokenRow(Protocol): + token: str + key_alias: str | None + team_id: str | None + user_id: str | None + metadata: Mapping[str, object] | None + + +class _DeletedVerificationTokenRow(_VerificationTokenRow, Protocol): + deleted_at: datetime + + +class _QueryRaw(Protocol): + async def __call__(self, query: str, *values: object) -> Sequence[Mapping[str, object]] | None: ... + + +class _DailyActivityDatabase(Protocol): + query_raw: _QueryRaw + litellm_verificationtoken: TableActions[_VerificationTokenRow] + litellm_deletedverificationtoken: TableActions[_DeletedVerificationTokenRow] + + @property + def litellm_dailyuserspend(self) -> TableActions[DailyActivityRow]: ... + + @property + def litellm_dailyteamspend(self) -> TableActions[DailyActivityRow]: ... + + @property + def litellm_dailytagspend(self) -> TableActions[DailyActivityRow]: ... + + @property + def litellm_dailyorganizationspend(self) -> TableActions[DailyActivityRow]: ... + + @property + def litellm_dailyenduserspend(self) -> TableActions[DailyActivityRow]: ... + + @property + def litellm_dailyagentspend(self) -> TableActions[DailyActivityRow]: ... + + +class DailyActivityDatabase(Protocol): + @property + def db(self) -> _DailyActivityDatabase: ... + + +_GROUPING_ADAPTER: Final = TypeAdapter(tuple[GroupingSetsRow, ...]) +_ENTITY_ADAPTER: Final = TypeAdapter(tuple[EntityRollupRow, ...]) +_KEY_SPEND_ADAPTER: Final = TypeAdapter(tuple[KeySpendRow, ...]) +_KEY_PAGE_TOTAL_ADAPTER: Final[TypeAdapter[int]] = TypeAdapter(int) +_EXPORT_ADAPTER: Final = TypeAdapter(tuple[ExportRow, ...]) +_METADATA_TAGS_ADAPTER: Final = TypeAdapter(list[StrictStr]) + + +def _metadata_tags(value: object) -> tuple[str, ...]: + stable_value: Final = value + if not isinstance(value, list): + return () + try: + return tuple(_METADATA_TAGS_ADAPTER.validate_python(stable_value)) + except ValidationError: + return () + + +def _daily_rows_table( + prisma_client: DailyActivityDatabase, table: DailyActivityTable +) -> TableActions[DailyActivityRow]: + if table is DailyActivityTable.USER: + return prisma_client.db.litellm_dailyuserspend + if table is DailyActivityTable.TEAM: + return prisma_client.db.litellm_dailyteamspend + if table is DailyActivityTable.TAG: + return prisma_client.db.litellm_dailytagspend + if table is DailyActivityTable.ORGANIZATION: + return prisma_client.db.litellm_dailyorganizationspend + if table is DailyActivityTable.CUSTOMER: + return prisma_client.db.litellm_dailyenduserspend + if table is DailyActivityTable.AGENT: + return prisma_client.db.litellm_dailyagentspend + assert_never(table) + raise AssertionError("unreachable") + + +def _next_export_cursor(batch: tuple[ExportRow, ...], export_type: ExportType) -> ExportCursor: + last: Final = batch[-1] + cursor_key: Final = ( + last.api_key + if export_type is ExportType.DAILY_WITH_KEYS + else last.model + if export_type is ExportType.DAILY_WITH_MODELS + else last.user_id + if export_type is ExportType.DAILY_WITH_USERS + else "" + ) + return ExportCursor(date=last.date, entity_id=last.entity_id, group_key=cursor_key or "") + + +class DailyActivityRepository: + def __init__(self, prisma_client: DailyActivityDatabase, *, proxy_reads: DailyActivityProxyReads) -> None: + self._prisma_client = prisma_client + self._proxy_reads = proxy_reads + + async def _query(self, query: SqlQuery) -> tuple[Mapping[str, object], ...]: + first_line: Final = query.sql.lstrip().splitlines()[0].lstrip("(").strip() + verbose_proxy_logger.debug("DailyActivityRepository query: %s", first_line) + result: Sequence[Mapping[str, object]] | None = await self._prisma_client.db.query_raw(query.sql, *query.params) + if result is None: + return () + return tuple(result) + + async def aggregated( + self, scope: DailyActivityScope, *, include_entity_breakdown: bool, api_key_limit: int + ) -> AggregatedRows: + grouping_query: Final = build_aggregated_sql(scope, api_key_limit=api_key_limit) + entity_query: Final = ( + build_entity_rollup_sql(scope, api_key_limit=api_key_limit) if include_entity_breakdown else None + ) + grouping_result, entity_result = await asyncio.gather( + self._query(grouping_query), + self._query(entity_query) if entity_query is not None else asyncio.sleep(0, result=None), + ) + grouping_rows: Final = _GROUPING_ADAPTER.validate_python(grouping_result) + entity_rows: Final = None if entity_result is None else _ENTITY_ADAPTER.validate_python(entity_result) + distinct_api_keys: Final = next( + (row.distinct_api_keys for row in grouping_rows if row.distinct_api_keys is not None), 0 + ) + return AggregatedRows( + grouping_rows=grouping_rows, + entity_rows=entity_rows, + distinct_api_keys=distinct_api_keys, + ) + + async def search_keys(self, scope: DailyActivityScope, *, search: str, limit: int) -> tuple[str, ...]: + if not 1 <= limit <= constants.USAGE_KEY_SEARCH_MAX: + raise ValueError(f"limit must be between 1 and {constants.USAGE_KEY_SEARCH_MAX}") + query: Final = build_key_search_sql(scope, search=search, limit=limit) + rows: Final = _KEY_SPEND_ADAPTER.validate_python(await self._query(query)) + return tuple(row.api_key for row in rows) + + async def key_page(self, scope: DailyActivityScope, *, offset: int, limit: int) -> KeyPage: + query: Final = build_key_page_sql(scope, offset=offset, limit=limit) + result: Final = await self._query(query) + total_api_keys_value: Final = result[0].get("total_api_keys") if result else 0 + total_api_keys: Final = ( + _KEY_PAGE_TOTAL_ADAPTER.validate_python(total_api_keys_value) if total_api_keys_value is not None else 0 + ) + rows: Final = _KEY_SPEND_ADAPTER.validate_python(tuple(row for row in result if row.get("api_key") is not None)) + return KeyPage(rows=rows, total_api_keys=total_api_keys) + + async def model_top_keys( + self, scope: DailyActivityScope, *, model_group: str, by_model_group: bool, limit: int + ) -> tuple[KeySpendRow, ...]: + if not 1 <= limit <= constants.USAGE_MODEL_TOP_KEYS_MAX: + raise ValueError(f"limit must be between 1 and {constants.USAGE_MODEL_TOP_KEYS_MAX}") + query: Final = build_model_top_keys_sql( + scope, + model_group=model_group, + by_model_group=by_model_group, + limit=limit, + ) + return _KEY_SPEND_ADAPTER.validate_python(await self._query(query)) + + async def cache_leakage_keys(self, scope: DailyActivityScope, *, limit: int) -> tuple[KeySpendRow, ...]: + if not 1 <= limit <= constants.USAGE_CACHE_LEAKAGE_KEYS_MAX: + raise ValueError(f"limit must be between 1 and {constants.USAGE_CACHE_LEAKAGE_KEYS_MAX}") + query: Final = build_cache_leakage_keys_sql(scope, limit=limit) + return _KEY_SPEND_ADAPTER.validate_python(await self._query(query)) + + async def export_rows(self, scope: DailyActivityScope, *, export_type: ExportType) -> AsyncIterator[ExportRow]: + batch_size: Final = constants.USAGE_EXPORT_BATCH_SIZE + cursor: ExportCursor | None = None # rebind-ok: each page advances the export keyset cursor + while True: + batch: tuple[ExportRow, ...] = _EXPORT_ADAPTER.validate_python( + await self._query(build_export_sql(scope, export_type=export_type, after=cursor, batch_size=batch_size)) + ) + for row in batch: + yield row + if len(batch) < batch_size: + return + cursor = _next_export_cursor(batch, export_type) + + async def _active_token_rows(self, values: tuple[str, ...]) -> tuple[_VerificationTokenRow, ...]: + return await find_many_in(self._prisma_client.db.litellm_verificationtoken, "token", values) + + async def _deleted_token_rows(self, values: tuple[str, ...]) -> tuple[_DeletedVerificationTokenRow, ...]: + try: + return await find_many_in(self._prisma_client.db.litellm_deletedverificationtoken, "token", values) + except Exception as exc: + verbose_proxy_logger.warning("Could not read deleted verification token metadata: %s", exc) + return () + + async def key_metadata( + self, api_keys: frozenset[str], window: SpendLogsWindow | None + ) -> Mapping[str, KeyMetadataRow]: + if not api_keys: + return {} + values: Final = tuple(api_keys) + active_rows: Final = await self._active_token_rows(values) + active: Final = MappingProxyType({row.token: self._metadata_row(row, key_exists=True) for row in active_rows}) + missing: Final = tuple(key for key in values if key not in active) + deleted_rows: Final = await self._deleted_token_rows(missing) + deleted_by_token: Final = MappingProxyType( + { + token: max(rows, key=lambda row: row.deleted_at) + for token, rows in groupby( + sorted(deleted_rows, key=lambda row: row.token), + key=lambda row: row.token, + ) + } + ) + deleted: Final = MappingProxyType( + { + key: self._metadata_row(deleted_by_token[key], key_exists=False) + for key in missing + if key in deleted_by_token + } + ) + resolved: Final = MappingProxyType({**deleted, **active}) + return await self._proxy_reads.recover_key_metadata(resolved, api_keys, window) + + @staticmethod + def _metadata_row(row: _VerificationTokenRow, *, key_exists: bool) -> KeyMetadataRow: + tags: Final = _metadata_tags(row.metadata.get("tags") if row.metadata is not None else None) + return KeyMetadataRow( + api_key=row.token, + key_alias=row.key_alias, + team_id=row.team_id, + user_id=row.user_id, + user_email=None, + key_exists=key_exists, + tags=tags, + ) + + async def daily_rows(self, scope: DailyActivityScope, *, page: int, page_size: int) -> DailyRowsPage: + table: Final = _daily_rows_table(self._prisma_client, scope.table) + adjusted_start, adjusted_end = adjust_dates_for_timezone( + scope.start_date, + scope.end_date, + scope.timezone_offset_minutes, + include_current_utc_day=scope.include_current_utc_day, + ) + exclusion_filter: Final = ( + { + "OR": [ + {scope.entity_id_field: None}, + {scope.entity_id_field: {"not": {"in": list(scope.exclude_entity_ids)}}}, + ] + } + if scope.exclude_entity_ids + else {} + ) + conditions: Final = { + "date": {"gte": adjusted_start, "lte": adjusted_end}, + **({scope.entity_id_field: {"in": list(scope.entity_ids)}} if scope.entity_ids is not None else {}), + **exclusion_filter, + **({"model": scope.model} if scope.model else {}), + **({"api_key": {"in": list(scope.api_keys)}} if scope.api_keys is not None else {}), + } + count, rows = await asyncio.gather( + table.count(where=conditions), + table.find_many( + where=conditions, + skip=(page - 1) * page_size, + take=page_size, + order=({"date": "desc"}, {"id": "asc"}), + ), + ) + return DailyRowsPage(total_count=count, rows=tuple(rows)) diff --git a/litellm/repositories/daily_activity_sql.py b/litellm/repositories/daily_activity_sql.py new file mode 100644 index 00000000000..f12dc1be5ae --- /dev/null +++ b/litellm/repositories/daily_activity_sql.py @@ -0,0 +1,526 @@ +from collections.abc import Mapping +from dataclasses import dataclass +from datetime import datetime, timedelta, timezone +from itertools import count, islice +from types import MappingProxyType +from typing import Final + +from typing_extensions import assert_never + +from litellm import constants +from litellm.constants import PTU_SENTINEL_API_KEY +from litellm.types.repositories.daily_activity import DailyActivityScope, DailyActivityTable, ExportType + +_API_KEY_ROLLED_UP_BIT: Final = 32 +_MODEL_GROUP_EXPR: Final = "COALESCE(NULLIF(model_group, ''), model)" + + +@dataclass(frozen=True, slots=True) +class SqlQuery: + sql: str + params: tuple[object, ...] + + +@dataclass(frozen=True, slots=True) +class ExportCursor: + date: str + entity_id: str + group_key: str + + +PRISMA_TO_PG_TABLE: Final[Mapping[DailyActivityTable, str]] = MappingProxyType( + { + DailyActivityTable.USER: "LiteLLM_DailyUserSpend", + DailyActivityTable.TEAM: "LiteLLM_DailyTeamSpend", + DailyActivityTable.TAG: "LiteLLM_DailyTagSpend", + DailyActivityTable.ORGANIZATION: "LiteLLM_DailyOrganizationSpend", + DailyActivityTable.CUSTOMER: "LiteLLM_DailyEndUserSpend", + DailyActivityTable.AGENT: "LiteLLM_DailyAgentSpend", + } +) + + +def adjust_dates_for_timezone( + start_date: str, + end_date: str, + timezone_offset_minutes: int | None, + include_current_utc_day: bool = False, + utc_now: datetime | None = None, +) -> tuple[str, str]: + if not include_current_utc_day or timezone_offset_minutes is None: + return start_date, end_date + now: Final = utc_now if utc_now is not None else datetime.now(timezone.utc) + caller_local_today: Final = (now - timedelta(minutes=timezone_offset_minutes)).date().isoformat() + if end_date < caller_local_today: + return start_date, end_date + return start_date, max(end_date, now.date().isoformat()) + + +def build_where_clause(scope: DailyActivityScope, *, start_index: int = 1) -> tuple[str, tuple[object, ...]]: + adjusted_start, adjusted_end = adjust_dates_for_timezone( + scope.start_date, + scope.end_date, + scope.timezone_offset_minutes, + scope.include_current_utc_day, + ) + entity_index: Final = start_index + 2 + has_entity_array: Final = scope.entity_ids is not None and bool(scope.entity_ids) + exclusion_index: Final = entity_index + int(has_entity_array) + model_index: Final = exclusion_index + int(bool(scope.exclude_entity_ids)) + api_keys_index: Final = model_index + int(bool(scope.model)) + conditions: Final = ( + f"date >= ${start_index}", + f"date <= ${start_index + 1}", + *( + ("FALSE",) + if scope.entity_ids == () + else (f'"{scope.entity_id_field}" = ANY(${entity_index}::text[])',) + if has_entity_array + else () + ), + *( + ( + f'("{scope.entity_id_field}" IS NULL ' + f'OR NOT ("{scope.entity_id_field}" = ANY(${exclusion_index}::text[])))', + ) + if scope.exclude_entity_ids + else () + ), + *((f"model = ${model_index}",) if scope.model else ()), + *( + ("FALSE",) + if scope.api_keys == () + else (f"api_key = ANY(${api_keys_index}::text[])",) + if scope.api_keys + else () + ), + ) + params: Final = ( + adjusted_start, + adjusted_end, + *((list(scope.entity_ids or ()),) if has_entity_array else ()), + *((list(scope.exclude_entity_ids),) if scope.exclude_entity_ids else ()), + *((scope.model,) if scope.model else ()), + *((list(scope.api_keys),) if scope.api_keys else ()), + ) + return " AND ".join(conditions), params + + +def _ptu_flat_cost_select(table: DailyActivityTable, *, aggregate: bool = True) -> str: + if table is DailyActivityTable.TEAM: + return "SUM(ptu_flat_cost)::float AS ptu_flat_cost" if aggregate else "SUM(scoped.ptu_flat_cost)::float" + return "0::float AS ptu_flat_cost" if aggregate else "0::float" + + +def _rollup_metric_select(table: DailyActivityTable) -> str: + return f""" + SUM(spend)::float AS spend, + {_ptu_flat_cost_select(table)}, + SUM(prompt_tokens)::bigint AS prompt_tokens, + SUM(completion_tokens)::bigint AS completion_tokens, + SUM(cache_read_input_tokens)::bigint AS cache_read_input_tokens, + SUM(cache_creation_input_tokens)::bigint AS cache_creation_input_tokens, + SUM(compression_saved_tokens)::bigint AS compression_saved_tokens, + SUM(compression_savings_spend)::float AS compression_savings_spend, + SUM(prompt_caching_savings_spend)::float AS prompt_caching_savings_spend, + SUM(gateway_injected_caching_savings_spend)::float AS gateway_injected_caching_savings_spend, + SUM(autorouter_savings_spend)::float AS autorouter_savings_spend, + SUM(api_requests)::bigint AS api_requests, + SUM(successful_requests)::bigint AS successful_requests, + SUM(failed_requests)::bigint AS failed_requests, + SUM(total_response_time_ms)::bigint AS total_response_time_ms, + SUM(timed_requests)::bigint AS timed_requests""" + + +def _validate_api_key_limit(api_key_limit: int) -> None: + if not 1 <= api_key_limit <= constants.USAGE_TOP_API_KEYS_MAX: + raise ValueError(f"api_key_limit must be between 1 and {constants.USAGE_TOP_API_KEYS_MAX}") + + +def _top_api_keys_sql(pg_table: str, where_clause: str, *, sentinel_param: int, limit_param: int) -> str: + return f""" + SELECT api_key, COUNT(*) OVER () AS distinct_api_keys + FROM "{pg_table}" + WHERE {where_clause} AND api_key <> ${sentinel_param} + GROUP BY api_key + ORDER BY SUM(spend::numeric) DESC, api_key + LIMIT ${limit_param} + """ + + +def build_aggregated_sql(scope: DailyActivityScope, *, api_key_limit: int) -> SqlQuery: + pg_table: Final = PRISMA_TO_PG_TABLE[scope.table] + where_clause, where_params = build_where_clause(scope) + _validate_api_key_limit(api_key_limit) + sentinel_param: Final = len(where_params) + 1 + top_keys_limit_param: Final = len(where_params) + 2 + top_api_keys: Final = _top_api_keys_sql( + pg_table, where_clause, sentinel_param=sentinel_param, limit_param=top_keys_limit_param + ) + metric_select: Final = _rollup_metric_select(scope.table) + sql: Final = f""" + (SELECT + date, + NULL::text AS api_key, + model, + {_MODEL_GROUP_EXPR} AS model_group, + custom_llm_provider, + mcp_namespaced_tool_name, + endpoint, + (GROUPING(date) << 6) | {_API_KEY_ROLLED_UP_BIT} + | GROUPING(model, {_MODEL_GROUP_EXPR}, + custom_llm_provider, mcp_namespaced_tool_name, + endpoint) AS group_level, + NULL::bigint AS distinct_api_keys,{metric_select} + FROM "{pg_table}" + WHERE {where_clause} + GROUP BY GROUPING SETS ( + (date), + (date, model), + (date, {_MODEL_GROUP_EXPR}), + (date, custom_llm_provider), + (date, mcp_namespaced_tool_name), + (date, endpoint), + () + )) + UNION ALL + (WITH top_api_keys AS ( + {top_api_keys} + ) + SELECT + date, + api_key, + model, + {_MODEL_GROUP_EXPR} AS model_group, + custom_llm_provider, + mcp_namespaced_tool_name, + endpoint, + GROUPING(date, api_key, model, {_MODEL_GROUP_EXPR}, + custom_llm_provider, mcp_namespaced_tool_name, + endpoint) AS group_level, + MAX(top_api_keys.distinct_api_keys) AS distinct_api_keys,{metric_select} + FROM "{pg_table}" JOIN top_api_keys USING (api_key) + WHERE {where_clause} + GROUP BY GROUPING SETS ( + (date, api_key), + (date, model, api_key), + (date, {_MODEL_GROUP_EXPR}, api_key), + (date, custom_llm_provider, api_key), + (date, mcp_namespaced_tool_name, api_key), + (date, endpoint, api_key) + )) + """ + return SqlQuery( + sql=sql, + params=(*where_params, PTU_SENTINEL_API_KEY, api_key_limit), + ) + + +def build_entity_rollup_sql(scope: DailyActivityScope, *, api_key_limit: int) -> SqlQuery: + pg_table: Final = PRISMA_TO_PG_TABLE[scope.table] + where_clause, where_params = build_where_clause(scope) + _validate_api_key_limit(api_key_limit) + sentinel_param: Final = len(where_params) + 1 + top_keys_limit_param: Final = len(where_params) + 2 + top_api_keys: Final = _top_api_keys_sql( + pg_table, where_clause, sentinel_param=sentinel_param, limit_param=top_keys_limit_param + ) + metric_select: Final = _rollup_metric_select(scope.table) + sql: Final = f""" + WITH top_api_keys AS ( + {top_api_keys} + ), + entity_api_keys AS ( + SELECT COALESCE("{scope.entity_id_field}", '') AS entity_id, + COUNT(DISTINCT api_key)::bigint AS distinct_api_keys + FROM "{pg_table}" + WHERE {where_clause} AND api_key <> ${sentinel_param} + GROUP BY COALESCE("{scope.entity_id_field}", '') + ) + (SELECT e.*, COALESCE(k.distinct_api_keys, 0)::bigint AS distinct_api_keys + FROM ( + SELECT COALESCE("{scope.entity_id_field}", '') AS entity_id, + date, + NULL::text AS api_key, + 1 AS api_key_rolled,{metric_select} + FROM "{pg_table}" + WHERE {where_clause} + GROUP BY date, COALESCE("{scope.entity_id_field}", '') + ) e + LEFT JOIN entity_api_keys k ON k.entity_id = e.entity_id) + UNION ALL + (SELECT COALESCE("{scope.entity_id_field}", '') AS entity_id, + date, + api_key, + 0 AS api_key_rolled,{metric_select}, + NULL::bigint AS distinct_api_keys + FROM "{pg_table}" JOIN top_api_keys USING (api_key) + WHERE {where_clause} + GROUP BY date, COALESCE("{scope.entity_id_field}", ''), api_key) + """ + return SqlQuery(sql=sql, params=(*where_params, PTU_SENTINEL_API_KEY, api_key_limit)) + + +def _key_spend_select() -> str: + return """ + COALESCE(SUM(spend), 0)::float AS spend, + COALESCE(SUM(prompt_tokens), 0)::bigint AS prompt_tokens, + COALESCE(SUM(completion_tokens), 0)::bigint AS completion_tokens, + (COALESCE(SUM(prompt_tokens), 0) + COALESCE(SUM(completion_tokens), 0))::bigint AS total_tokens, + COALESCE(SUM(api_requests), 0)::bigint AS api_requests, + COALESCE(SUM(successful_requests), 0)::bigint AS successful_requests, + COALESCE(SUM(failed_requests), 0)::bigint AS failed_requests, + COALESCE(SUM(cache_read_input_tokens), 0)::bigint AS cache_read_input_tokens, + COALESCE(SUM(cache_creation_input_tokens), 0)::bigint AS cache_creation_input_tokens""" + + +def build_key_page_sql(scope: DailyActivityScope, *, offset: int, limit: int) -> SqlQuery: + if not 1 <= limit <= constants.USAGE_KEY_PAGE_MAX: + raise ValueError(f"limit must be between 1 and {constants.USAGE_KEY_PAGE_MAX}") + if offset < 0: + raise ValueError("offset must be non-negative") + where_clause, where_params = build_where_clause(scope) + sentinel_param: Final = len(where_params) + 1 + limit_param: Final = sentinel_param + 1 + offset_param: Final = limit_param + 1 + sql: Final = f""" + WITH ranked AS ( + SELECT api_key,{_key_spend_select()}, SUM(spend::numeric) AS rank_spend + FROM "{PRISMA_TO_PG_TABLE[scope.table]}" + WHERE {where_clause} AND api_key <> ${sentinel_param} + GROUP BY api_key + ) + SELECT (SELECT COUNT(*) FROM ranked)::bigint AS total_api_keys, page.* + FROM (SELECT 1) AS one + LEFT JOIN LATERAL ( + SELECT * FROM ranked + ORDER BY rank_spend DESC, api_key + LIMIT ${limit_param} OFFSET ${offset_param} + ) AS page ON TRUE + """ + return SqlQuery(sql=sql, params=(*where_params, PTU_SENTINEL_API_KEY, limit, offset)) + + +def _bounded_limit(limit: int, *, minimum: int = 1) -> None: + if limit < minimum: + raise ValueError(f"limit must be at least {minimum}") + + +def build_key_search_sql(scope: DailyActivityScope, *, search: str, limit: int) -> SqlQuery: + _bounded_limit(limit) + where_clause, where_params = build_where_clause(scope) + search_param: Final = len(where_params) + 1 + sentinel_param: Final = search_param + 1 + limit_param: Final = sentinel_param + 1 + escaped: Final = search.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") + sql: Final = f""" + SELECT api_key,{_key_spend_select()} + FROM "{PRISMA_TO_PG_TABLE[scope.table]}" + WHERE {where_clause} + AND api_key <> ${sentinel_param} + AND ( + api_key ILIKE ${search_param} ESCAPE '\\' + OR api_key IN ( + SELECT v.token FROM "LiteLLM_VerificationToken" v + LEFT JOIN "LiteLLM_UserTable" u ON u.user_id = v.user_id + WHERE v.key_alias ILIKE ${search_param} ESCAPE '\\' + OR v.user_id ILIKE ${search_param} ESCAPE '\\' + OR u.user_email ILIKE ${search_param} ESCAPE '\\' + UNION + SELECT d.token FROM "LiteLLM_DeletedVerificationToken" d + LEFT JOIN "LiteLLM_UserTable" u ON u.user_id = d.user_id + WHERE d.key_alias ILIKE ${search_param} ESCAPE '\\' + OR d.user_id ILIKE ${search_param} ESCAPE '\\' + OR u.user_email ILIKE ${search_param} ESCAPE '\\' + ) + ) + GROUP BY api_key + ORDER BY SUM(spend::numeric) DESC, api_key + LIMIT ${limit_param} + """ + return SqlQuery(sql=sql, params=(*where_params, f"%{escaped}%", PTU_SENTINEL_API_KEY, limit)) + + +def build_model_top_keys_sql( + scope: DailyActivityScope, *, model_group: str, by_model_group: bool, limit: int +) -> SqlQuery: + _bounded_limit(limit) + where_clause, where_params = build_where_clause(scope) + model_param: Final = len(where_params) + 1 + sentinel_param: Final = model_param + 1 + limit_param: Final = sentinel_param + 1 + model_clause: Final = ( + f"COALESCE(NULLIF(model_group, ''), model) = ${model_param}" if by_model_group else f"model = ${model_param}" + ) + sql: Final = f""" + SELECT api_key,{_key_spend_select()} + FROM "{PRISMA_TO_PG_TABLE[scope.table]}" + WHERE {where_clause} AND {model_clause} AND api_key <> ${sentinel_param} + GROUP BY api_key + ORDER BY SUM(spend::numeric) DESC, api_key + LIMIT ${limit_param} + """ + return SqlQuery(sql=sql, params=(*where_params, model_group, PTU_SENTINEL_API_KEY, limit)) + + +def build_cache_leakage_keys_sql(scope: DailyActivityScope, *, limit: int) -> SqlQuery: + _bounded_limit(limit) + where_clause, where_params = build_where_clause(scope) + sentinel_param: Final = len(where_params) + 1 + limit_param: Final = sentinel_param + 1 + sql: Final = f""" + SELECT api_key,{_key_spend_select()} + FROM "{PRISMA_TO_PG_TABLE[scope.table]}" + WHERE {where_clause} AND api_key <> ${sentinel_param} + GROUP BY api_key + HAVING SUM(prompt_tokens) - SUM(cache_read_input_tokens) > 0 + ORDER BY SUM(prompt_tokens) - SUM(cache_read_input_tokens) DESC, api_key + LIMIT ${limit_param} + """ + return SqlQuery(sql=sql, params=(*where_params, PTU_SENTINEL_API_KEY, limit)) + + +def build_export_sql( + scope: DailyActivityScope, *, export_type: ExportType, after: ExportCursor | None, batch_size: int +) -> SqlQuery: + _bounded_limit(batch_size) + where_clause, where_params = build_where_clause(scope) + group_key, output_key, user_fields, type_joins = _export_grouping(export_type) + grouping_keys: Final = ( + f"scoped.date, COALESCE(scoped.\"{scope.entity_id_field}\", '')", + *((group_key,) if export_type is not ExportType.DAILY else ()), + ) + entity_joins: Final = ( + ('LEFT JOIN "LiteLLM_TeamTable" tt ON tt.team_id = scoped.team_id',) + if scope.table is DailyActivityTable.TEAM + else ('LEFT JOIN "LiteLLM_OrganizationTable" ot ON ot.organization_id = scoped.organization_id',) + if scope.table is DailyActivityTable.ORGANIZATION + else () + ) + joins: Final = (*type_joins, *entity_joins) + alias_expression: Final = ( + "MAX(tt.team_alias)" + if scope.table is DailyActivityTable.TEAM + else "MAX(ot.organization_alias)" + if scope.table is DailyActivityTable.ORGANIZATION + else "NULL::text" + ) + parameter_indexes: Final = count(len(where_params) + 1) + sentinel_param: Final = next(parameter_indexes) if export_type is not ExportType.DAILY else None + cursor_indexes: Final = tuple(islice(parameter_indexes, 3)) if after is not None else () + limit_param: Final = next(parameter_indexes) + cursor_clause, cursor_params = _export_cursor_clause( + scope, after=after, cursor_indexes=cursor_indexes, group_key=group_key + ) + sentinel_clause: Final = f" AND api_key <> ${sentinel_param}" if sentinel_param is not None else "" + table: Final = PRISMA_TO_PG_TABLE[scope.table] + flat_cost: Final = _ptu_flat_cost_select(scope.table, aggregate=False) + sql: Final = f""" + WITH scoped AS ( + SELECT * FROM "{table}" + WHERE {where_clause}{sentinel_clause} + ) + SELECT + scoped.date, + COALESCE(scoped."{scope.entity_id_field}", '') AS entity_id, + {alias_expression} AS entity_alias, + {output_key} AS api_key, + {user_fields}, + {"NULLIF(COALESCE(scoped.model, ''), '')" if export_type is ExportType.DAILY_WITH_MODELS else "NULL::text"} AS model, + COALESCE(SUM(scoped.spend), 0)::float AS spend, + {flat_cost} AS flat_cost, + COALESCE(SUM(scoped.prompt_tokens), 0)::bigint AS prompt_tokens, + COALESCE(SUM(scoped.completion_tokens), 0)::bigint AS completion_tokens, + COALESCE(SUM(scoped.api_requests), 0)::bigint AS api_requests, + COALESCE(SUM(scoped.successful_requests), 0)::bigint AS successful_requests, + COALESCE(SUM(scoped.failed_requests), 0)::bigint AS failed_requests, + COALESCE(SUM(scoped.cache_read_input_tokens), 0)::bigint AS cache_read_input_tokens, + COALESCE(SUM(scoped.cache_creation_input_tokens), 0)::bigint AS cache_creation_input_tokens + FROM scoped + {" ".join(joins)} + WHERE TRUE{cursor_clause} + GROUP BY {", ".join(grouping_keys)} + ORDER BY {", ".join(grouping_keys)} + LIMIT ${limit_param} + """ + return SqlQuery( + sql=sql, + params=( + *where_params, + *((PTU_SENTINEL_API_KEY,) if export_type is not ExportType.DAILY else ()), + *cursor_params, + batch_size, + ), + ) + + +def _export_grouping(export_type: ExportType) -> tuple[str, str, str, tuple[str, ...]]: + if export_type is ExportType.DAILY: + return ( + "''", + "NULL::text", + "NULL::text AS key_alias, NULL::text AS user_id, NULL::text AS user_email", + (), + ) + if export_type is ExportType.DAILY_WITH_KEYS: + return ( + "scoped.api_key", + "NULLIF(scoped.api_key, '')", + "MAX(COALESCE(vt.key_alias, dvt.key_alias)) AS key_alias, " + "MAX(COALESCE(vt.user_id, dvt.user_id)) AS user_id, MAX(u.user_email) AS user_email", + ( + 'LEFT JOIN "LiteLLM_VerificationToken" vt ON vt.token = scoped.api_key', + """LEFT JOIN LATERAL ( + SELECT key_alias, user_id + FROM "LiteLLM_DeletedVerificationToken" + WHERE token = scoped.api_key + ORDER BY deleted_at DESC + LIMIT 1 + ) dvt ON vt.token IS NULL""", + 'LEFT JOIN "LiteLLM_UserTable" u ON u.user_id = COALESCE(vt.user_id, dvt.user_id)', + ), + ) + if export_type is ExportType.DAILY_WITH_MODELS: + return ( + "COALESCE(scoped.model, '')", + "NULL::text", + "NULL::text AS key_alias, NULL::text AS user_id, NULL::text AS user_email", + (), + ) + if export_type is ExportType.DAILY_WITH_USERS: + return ( + "COALESCE(vt.user_id, dvt.user_id, '')", + "NULL::text", + "NULL::text AS key_alias, MAX(COALESCE(vt.user_id, dvt.user_id)) AS user_id, " + "MAX(u.user_email) AS user_email", + ( + 'LEFT JOIN "LiteLLM_VerificationToken" vt ON vt.token = scoped.api_key', + """LEFT JOIN LATERAL ( + SELECT key_alias, user_id + FROM "LiteLLM_DeletedVerificationToken" + WHERE token = scoped.api_key + ORDER BY deleted_at DESC + LIMIT 1 + ) dvt ON vt.token IS NULL""", + 'LEFT JOIN "LiteLLM_UserTable" u ON u.user_id = COALESCE(vt.user_id, dvt.user_id)', + ), + ) + assert_never(export_type) + raise AssertionError("unreachable") + + +def _export_cursor_clause( + scope: DailyActivityScope, + *, + after: ExportCursor | None, + cursor_indexes: tuple[int, ...], + group_key: str, +) -> tuple[str, tuple[object, ...]]: + if after is None: + return "", () + first_cursor_index: Final = cursor_indexes[0] + clause: Final = ( + f""" AND (scoped.date, COALESCE(scoped."{scope.entity_id_field}", ''), {group_key}) """ + f"> (${first_cursor_index}, ${cursor_indexes[1]}, ${cursor_indexes[2]})" + ) + return clause, (after.date, after.entity_id, after.group_key) diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 71f61079154..c9e0861935d 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -40,6 +40,7 @@ if TYPE_CHECKING: from mcp.types import CallToolResult from mcp.types import Tool as MCPTool + from litellm.proxy._experimental.mcp_server.ui_session_utils import GrantedToolsetIds from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.utils import ProxyLogging else: @@ -223,7 +224,9 @@ class LiteLLM_Proxy_MCP_Handler: mcp_servers=all_server_ids, mcp_tool_permissions=tool_permissions, ) - return user_api_key_auth.model_copy(update={"object_permission": updated_op}) + return user_api_key_auth.model_copy( + update={"object_permission": updated_op, "mcp_explicit_grants_only": True} + ) except Exception as _e: verbose_logger.debug("Could not apply toolset permissions: %s", _e) return user_api_key_auth @@ -237,6 +240,7 @@ class LiteLLM_Proxy_MCP_Handler: mcp_server_auth_headers: dict[str, dict[str, str]] | None = None, request_tags: list[str] | None = None, raw_headers: dict[str, str] | None = None, + granted_toolsets: "GrantedToolsetIds | None" = None, ) -> tuple[list[MCPTool], list[str]]: """ Get available tools from the MCP server manager. @@ -279,23 +283,19 @@ class LiteLLM_Proxy_MCP_Handler: if prisma_client is not None: toolset = await global_mcp_server_manager.get_toolset_by_name_cached(prisma_client, name) if toolset is not None: - # Access control: only allow if the key explicitly grants this toolset. if user_api_key_auth is not None: + from litellm.proxy._experimental.mcp_server.ui_session_utils import ( + granted_toolset_ids, + ) from litellm.proxy.management_endpoints.common_utils import ( _user_has_admin_view, ) - is_admin = _user_has_admin_view(user_api_key_auth) - if not is_admin: - op = user_api_key_auth.object_permission - granted = getattr(op, "mcp_toolsets", None) if op else None - # None means no grants configured → deny (consistent with - # fetch_mcp_toolsets which returns [] for unconfigured keys) - if granted is None or toolset.toolset_id not in granted: - verbose_logger.debug( - "Key does not have access to toolset '%s', skipping.", name - ) - continue + if not _user_has_admin_view(user_api_key_auth) and toolset.toolset_id not in ( + await (granted_toolsets or granted_toolset_ids)(user_api_key_auth) + ): + verbose_logger.debug("Key does not have access to toolset '%s', skipping.", name) + continue resolved_toolset_ids.append(toolset.toolset_id) # Don't add to resolved_mcp_servers — toolset scope # restricts via object_permission, not server name filter. diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index a3d8ba0e582..5d23b048e86 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -11,7 +11,7 @@ from litellm.rust_bridge.embeddings.entrypoints import LiteLLMEmbeddingRequest from litellm.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest from litellm.rust_bridge.ocr.entrypoints import LiteLLMOcrRequest from litellm.rust_bridge.responses.entrypoints import LiteLLMResponsesRequest -from litellm.rust_bridge.traces import DecodedSpan +from litellm.rust_bridge.traces import DecodedSpan, QueryScope from litellm.types.llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.utils import EmbeddingResponse, ModelResponse @@ -23,12 +23,15 @@ class ProcessReservedForForking(RuntimeError): ... def trace_decode_otlp(body: bytes, content_type: str | None) -> list[DecodedSpan]: ... def trace_encode_error(message: str) -> bytes: ... +def trace_normalized_field_definitions() -> list[dict[str, str]]: ... @final class NativeTraceStorage: def __new__(cls, database: str, url: str, reader_url: str | None = None) -> NativeTraceStorage: ... def ensure_schema(self, trace_retention_days: int, spend_log_retention_days: int) -> Future[None]: ... def insert_rows(self, table: str, rows: Sequence[Mapping[str, object]]) -> Future[None]: ... + def query_sql(self, sql: str, scope: QueryScope, secret: str) -> Future[str]: ... + def query_help(self, scope: QueryScope, secret: str) -> Future[str]: ... def lens_query(self, name: str, parameters: Mapping[str, str | int | Sequence[str]]) -> Future[str]: ... def query(self, query: str, parameters: Mapping[str, str | int | Sequence[str]]) -> Future[str]: ... @@ -353,6 +356,7 @@ __all__ = [ "responses", "trace_decode_otlp", "trace_encode_error", + "trace_normalized_field_definitions", "transcription", ] diff --git a/litellm/rust_bridge/traces.py b/litellm/rust_bridge/traces.py index 6724db41ad3..afc1620a480 100644 --- a/litellm/rust_bridge/traces.py +++ b/litellm/rust_bridge/traces.py @@ -2,7 +2,7 @@ from collections.abc import Awaitable, Mapping, Sequence from types import MappingProxyType from typing import Final, Literal, Protocol, TypedDict, cast -from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter +from pydantic import BaseModel, ConfigDict, Field, JsonValue, TypeAdapter from typing_extensions import ReadOnly from litellm.rust_bridge.loader import get_native_bridge @@ -13,6 +13,28 @@ class DecodedEvent(TypedDict): attributes: ReadOnly[dict[str, str]] +class NormalizedSpan(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid", strict=True) + + observation_type: Literal["agent", "llm", "tool", "chain", "framework"] + agent_name: str + litellm_request_id: str + model: str + input_tokens: int = Field(ge=0, le=2**32 - 1) + output_tokens: int = Field(ge=0, le=2**32 - 1) + input: str + output: str + + +class NormalizedFieldDefinition(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid", strict=True) + + name: str + clickhouse_column: str + clickhouse_type: str + meaning: str + + class DecodedSpan(TypedDict): trace_id: ReadOnly[str] span_id: ReadOnly[str] @@ -29,11 +51,31 @@ class DecodedSpan(TypedDict): status_code: ReadOnly[str] status_message: ReadOnly[str] events: ReadOnly[list[DecodedEvent]] + normalized: ReadOnly[NormalizedSpan] + consumed_attributes: ReadOnly[tuple[str, str]] ReadQueryName = Literal["list_traces", "trace_spans", "span_detail", "span_error", "spend_by_response_ids"] +class AdminQueryScope(TypedDict): + kind: ReadOnly[Literal["admin"]] + + +class TeamQueryScope(TypedDict): + kind: ReadOnly[Literal["team"]] + team_id: ReadOnly[str] + + +class KeyQueryScope(TypedDict): + kind: ReadOnly[Literal["key"]] + team_id: ReadOnly[str] + api_key_hash: ReadOnly[str] + + +QueryScope = AdminQueryScope | TeamQueryScope | KeyQueryScope + + class NativeStore(Protocol): def __init__(self, database: str, url: str, reader_url: str | None = None) -> None: ... @@ -41,6 +83,10 @@ class NativeStore(Protocol): def insert_rows(self, table: str, rows: Sequence[Mapping[str, object]]) -> Awaitable[None]: ... + def query_sql(self, sql: str, scope: QueryScope, secret: str) -> Awaitable[str]: ... + + def query_help(self, scope: QueryScope, secret: str) -> Awaitable[str]: ... + def lens_query(self, name: str, parameters: Mapping[str, str | int | Sequence[str]]) -> Awaitable[str]: ... def query(self, name: ReadQueryName, parameters: Mapping[str, str | int | Sequence[str]]) -> Awaitable[str]: ... @@ -57,6 +103,8 @@ class NativeTraces(Protocol): def trace_encode_error(self, message: str) -> bytes: ... + def trace_normalized_field_definitions(self) -> list[dict[str, str]]: ... + class QueryResponse(BaseModel): model_config = ConfigDict(frozen=True) @@ -64,6 +112,7 @@ class QueryResponse(BaseModel): QUERY_PARAMETERS: Final = TypeAdapter(dict[str, str | int | list[str]]) +_FIELD_DEFINITIONS_ADAPTER: Final = TypeAdapter(tuple[NormalizedFieldDefinition, ...]) def _native() -> NativeTraces: @@ -74,7 +123,17 @@ def _native() -> NativeTraces: def decode_otlp(body: bytes, content_type: str | None) -> list[DecodedSpan]: - return _native().trace_decode_otlp(body, content_type) + return [ + {**span, "normalized": NormalizedSpan.model_validate(span["normalized"])} + for span in _native().trace_decode_otlp(body, content_type) + ] + + +def normalized_field_definitions() -> tuple[NormalizedFieldDefinition, ...]: + fields: Final = _FIELD_DEFINITIONS_ADAPTER.validate_python(_native().trace_normalized_field_definitions()) + if frozenset(field.name for field in fields) != frozenset(NormalizedSpan.model_fields): + raise ValueError("Rust and Python normalized trace fields disagree") + return fields def encode_error(message: str) -> bytes: @@ -101,6 +160,12 @@ class ClickHouseStorage: ) return QueryResponse.model_validate_json(result).data + async def query_sql(self, sql: str, scope: QueryScope, secret: str) -> str: + return await self._native.query_sql(sql, scope, secret) + + async def query_help(self, scope: QueryScope, secret: str) -> str: + return await self._native.query_help(scope, secret) + async def _lens_query(self, name: str, parameters: Mapping[str, object]) -> list[dict[str, JsonValue]]: result: Final = await self._native.lens_query(name, QUERY_PARAMETERS.validate_python(parameters)) return QueryResponse.model_validate_json(result).data @@ -108,6 +173,12 @@ class ClickHouseStorage: async def lens_sample(self, parameters: Mapping[str, object]) -> list[dict[str, JsonValue]]: return await self._lens_query("sample", parameters) + async def lens_availability(self, parameters: Mapping[str, object]) -> list[dict[str, JsonValue]]: + return await self._lens_query("availability", parameters) + + async def lens_agents(self, parameters: Mapping[str, object]) -> list[dict[str, JsonValue]]: + return await self._lens_query("agents", parameters) + async def lens_content(self, parameters: Mapping[str, object]) -> list[dict[str, JsonValue]]: return await self._lens_query("content", parameters) diff --git a/litellm/sandbox/__init__.py b/litellm/sandbox/__init__.py index e69de29bb2d..1974859fb1e 100644 --- a/litellm/sandbox/__init__.py +++ b/litellm/sandbox/__init__.py @@ -0,0 +1,27 @@ +"""litellm.sandbox: code-interpreter providers (see main.py) plus harness sandboxes. + +`sandbox.local(path)` and `sandbox.docker(image, ...)` re-export litellm.harness.sandbox. +They resolve lazily so `import litellm` does not pull in litellm.harness. +""" + +import importlib +from typing import Final + +_HARNESS_SANDBOX_MODULE: Final = "litellm.harness.sandbox" +_HARNESS_EXPORTS: Final = frozenset( + { + "local", + "docker", + "LocalSandbox", + "DockerSandbox", + "Sandbox", + "Process", + "CompletedRun", + } +) + + +def __getattr__(name: str) -> object: + if name in _HARNESS_EXPORTS: + return getattr(importlib.import_module(_HARNESS_SANDBOX_MODULE), name) + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") diff --git a/litellm/tracing/decode.py b/litellm/tracing/decode.py index d8b5f70de68..e76c9aad97d 100644 --- a/litellm/tracing/decode.py +++ b/litellm/tracing/decode.py @@ -1,50 +1,23 @@ -""" -OTLP/HTTP trace export -> `SpanRow`s. - -Pure functions, no I/O. Two steps: -1. `decode_otlp()` protobuf / JSON / gzip `ExportTraceServiceRequest` -> flat spans -2. `normalize()` framework conventions -> LiteLLM columns (type, agent, input/output, - LiteLLM request id). Supported: LangSmith (LangChain, LangGraph, - Deep Agents), OTEL GenAI semconv, OpenInference. -""" - import gzip import json import zlib from collections.abc import Mapping -from dataclasses import dataclass from io import BytesIO from itertools import accumulate from types import MappingProxyType from typing import Final from pydantic import JsonValue, TypeAdapter, ValidationError -from typing_extensions import NotRequired, ReadOnly, TypedDict +from typing_extensions import ReadOnly, TypedDict from litellm.constants import OTLP_MAX_ATTRIBUTE_VALUE_BYTES, OTLP_MAX_BODY_BYTES from litellm.rust_bridge.traces import DecodedSpan from litellm.rust_bridge.traces import decode_otlp as native_decode_otlp from litellm.rust_bridge.traces import encode_error as native_encode_error -from litellm.tracing.normalizers.messages import content_text -from litellm.tracing.types import SpanRow, SpanType +from litellm.tracing.types import SpanRow -_FRAMEWORK_SUFFIXES: Final = ( - ".wrap_model_call", - ".wrap_tool_call", - ".before_agent", - ".after_agent", - ".before_model", - ".after_model", -) -_LLM_OPERATIONS: Final = frozenset({"chat", "text_completion", "generate_content"}) -_LC_ROLES: Final = MappingProxyType({"human": "user", "ai": "assistant", "system": "system", "tool": "tool"}) -_OPENINFERENCE_TYPES: Final[Mapping[str, SpanType]] = MappingProxyType({"AGENT": "agent", "LLM": "llm", "TOOL": "tool"}) - - -_JSON: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) _MESSAGE_LIST: Final = TypeAdapter(tuple[dict[str, JsonValue], ...]) _MAX_JSON_ESCAPE_BYTES: Final = 6 -_MAX_TOKENS: Final = (1 << 32) - 1 class InvalidOTLPPayloadError(ValueError): @@ -55,33 +28,10 @@ class OTLPPayloadTooLargeError(OverflowError): pass -class MessageExtras(TypedDict): - tool_calls: ReadOnly[NotRequired[JsonValue]] - name: ReadOnly[NotRequired[str]] - - -class NormalizedMessage(MessageExtras): - role: ReadOnly[str] - content: ReadOnly[str] - - class OTLPError(TypedDict): message: ReadOnly[str] -@dataclass(frozen=True, slots=True) -class NormalizedSpan: - kind: SpanType - agent: str = "" - model: str = "" - request_id: str = "" - input: str = "" - output: str = "" - input_tokens: int = 0 - output_tokens: int = 0 - consumed: frozenset[str] = frozenset() - - def _truncate(value: str) -> str: encoded: Final = value.encode("utf-8") if len(encoded) <= OTLP_MAX_ATTRIBUTE_VALUE_BYTES: @@ -206,7 +156,7 @@ def _exception_message(span: DecodedSpan) -> str: def _span_row(span: DecodedSpan) -> SpanRow: attributes: Final = span["attributes"] - normalized: Final = normalize(span) + normalized: Final = span["normalized"] return SpanRow( Timestamp=span["start_ns"], TraceId=span["trace_id"], @@ -220,17 +170,17 @@ def _span_row(span: DecodedSpan) -> SpanRow: ScopeName=span["scope_name"], ScopeVersion=span["scope_version"], SpanAttributes=MappingProxyType( - {key: _truncate(value) for key, value in attributes.items() if key not in normalized.consumed} + {key: _truncate(value) for key, value in attributes.items() if key not in span["consumed_attributes"]} ), Duration=span["end_ns"] - span["start_ns"], StatusCode=span["status_code"], StatusMessage=span["status_message"] or _exception_message(span), TeamId="", ApiKeyHash="", - ObservationType=normalized.kind, - AgentName=normalized.agent, + ObservationType=normalized.observation_type, + AgentName=normalized.agent_name, Model=normalized.model, - LiteLLMRequestId=attributes.get("gen_ai.response.id") or normalized.request_id, + LiteLLMRequestId=attributes.get("gen_ai.response.id") or normalized.litellm_request_id, InputTokens=normalized.input_tokens, OutputTokens=normalized.output_tokens, Input=_truncate_payload(normalized.input), @@ -238,173 +188,6 @@ def _span_row(span: DecodedSpan) -> SpanRow: ) -def _loads(value: str) -> JsonValue: - if len(value.encode("utf-8")) > OTLP_MAX_BODY_BYTES: - return None - try: - return _JSON.validate_json(value) - except ValidationError: - return None - - -def _text(value: JsonValue) -> str: - return value if isinstance(value, str) else "" - - -def _message(value: JsonValue) -> NormalizedMessage | None: - if not isinstance(value, dict): - return None - kwargs: Final = value.get("kwargs", value) - if not isinstance(kwargs, dict): - return None - kind: Final = _text(kwargs.get("type")) or _text(kwargs.get("role")) - if not kind: - return None - calls: Final = kwargs.get("tool_calls") - if calls is not None and (not isinstance(calls, list) or not all(isinstance(call, dict) for call in calls)): - return None - role: Final = _LC_ROLES.get(kind, kind) - content: Final = kwargs.get("content", "") - name: Final = kwargs.get("name") - tool_calls: Final = MessageExtras(tool_calls=calls) if calls else MessageExtras() - tool_name: Final = MessageExtras(name=name) if role == "tool" and isinstance(name, str) else MessageExtras() - message: Final[NormalizedMessage] = { - "role": role, - "content": content_text(content), - **tool_calls, - **tool_name, - } - return message - - -def _messages(value: JsonValue, raw: str) -> str: - if not isinstance(value, list): - return raw - messages: Final = tuple(_message(item) for item in value) - return json.dumps(messages) if all(message is not None for message in messages) else raw - - -def _langsmith_type(span: DecodedSpan) -> SpanType: - attributes: Final = span["attributes"] - kind: Final = attributes.get("langsmith.span.kind", "chain") - if kind in ("llm", "tool"): - return "llm" if kind == "llm" else "tool" - if not span["parent_span_id"] or span["name"] == attributes.get("langsmith.metadata.lc_agent_name"): - return "agent" - return "framework" if span["name"].endswith(_FRAMEWORK_SUFFIXES) else "chain" - - -def _langsmith_io(kind: SpanType, attributes: Mapping[str, str]) -> tuple[str, str, str]: - raw_prompt: Final = attributes.get("gen_ai.prompt", "") - raw_completion: Final = attributes.get("gen_ai.completion", "") - prompt: Final = _loads(raw_prompt) - completion: Final = _loads(raw_completion) - messages: Final = prompt.get("messages") if isinstance(prompt, dict) else None - if kind == "llm": - batch: Final = ( - messages[0] if isinstance(messages, list) and messages and isinstance(messages[0], list) else messages - ) - generations: Final = completion.get("generations") if isinstance(completion, dict) else None - first: Final = generations[0] if isinstance(generations, list) and generations else None - item: Final = first[0] if isinstance(first, list) and first else first - message: Final = item.get("message") if isinstance(item, dict) else None - parsed: Final = _message(message) - kwargs: Final = message.get("kwargs", message) if isinstance(message, dict) else None - metadata: Final = kwargs.get("response_metadata") if isinstance(kwargs, dict) else None - request_id: Final = _text(metadata.get("id")) if isinstance(metadata, dict) else "" - return _messages(batch, raw_prompt), json.dumps(parsed) if parsed is not None else raw_completion, request_id - if kind == "tool": - output: Final = completion.get("output", completion) if isinstance(completion, dict) else completion - update: Final = output.get("update") if isinstance(output, dict) else None - updates: Final = update.get("messages") if isinstance(update, dict) else None - final: Final = updates[-1] if isinstance(updates, list) and updates else output - content: Final = final.get("content", final) if isinstance(final, dict) else final - return ( - raw_prompt, - (content if isinstance(content, str) else json.dumps(content)) if content is not None else raw_completion, - "", - ) - if kind == "agent": - outputs: Final = completion.get("messages") if isinstance(completion, dict) else None - last: Final = _message(outputs[-1]) if isinstance(outputs, list) and outputs else None - return _messages(messages, raw_prompt), json.dumps(last) if last is not None else raw_completion, "" - return raw_prompt, raw_completion, "" - - -def _to_int(value: str | None) -> int: - try: - number: Final = int(value) if value else 0 - except ValueError: - return 0 - if not 0 <= number <= _MAX_TOKENS: - raise InvalidOTLPPayloadError("OTLP token count is outside the storage range") - return number - - -def normalize(span: DecodedSpan) -> NormalizedSpan: - attributes: Final = span["attributes"] - fallback: Final[SpanType] = "agent" if not span["parent_span_id"] else "chain" - input_tokens: Final = _to_int(attributes.get("gen_ai.usage.input_tokens")) - output_tokens: Final = _to_int(attributes.get("gen_ai.usage.output_tokens")) - if span["scope_name"] == "langsmith" or "langsmith.span.kind" in attributes: - kind: Final = _langsmith_type(span) - prompt, completion, request_id = _langsmith_io(kind, attributes) - return NormalizedSpan( - kind, - attributes.get("langsmith.metadata.lc_agent_name", ""), - attributes.get("gen_ai.request.model", ""), - request_id, - prompt, - completion, - input_tokens, - output_tokens, - frozenset({"gen_ai.prompt", "gen_ai.completion"}), - ) - if "openinference.span.kind" in attributes: - return NormalizedSpan( - _OPENINFERENCE_TYPES.get(attributes["openinference.span.kind"].upper(), fallback), - attributes.get("agent.name", ""), - attributes.get("llm.model_name", ""), - "", - attributes.get("input.value", ""), - attributes.get("output.value", ""), - _to_int(attributes.get("llm.token_count.prompt")) - if "llm.token_count.prompt" in attributes - else input_tokens, - _to_int(attributes.get("llm.token_count.completion")) - if "llm.token_count.completion" in attributes - else output_tokens, - frozenset({"input.value", "output.value"}), - ) - operation: Final = attributes.get("gen_ai.operation.name", "") - genai_kind: Final[SpanType] = ( - "llm" - if operation in _LLM_OPERATIONS - else "tool" - if operation == "execute_tool" - else "agent" - if operation == "invoke_agent" - else fallback - ) - input_key: Final = ( - "gen_ai.input.messages" if attributes.get("gen_ai.input.messages") else "gen_ai.tool.call.arguments" - ) - output_key: Final = ( - "gen_ai.output.messages" if attributes.get("gen_ai.output.messages") else "gen_ai.tool.call.result" - ) - return NormalizedSpan( - genai_kind, - attributes.get("gen_ai.agent.name", ""), - attributes.get("gen_ai.request.model") or attributes.get("gen_ai.response.model", ""), - "", - attributes.get(input_key, ""), - attributes.get(output_key, ""), - input_tokens, - output_tokens, - frozenset({input_key, output_key}), - ) - - def encode_otlp_response(content_type: str | None, error: str | None = None) -> tuple[bytes, str]: media_type: Final = (content_type or "application/x-protobuf").split(";", 1)[0].strip().lower() if media_type == "application/json": diff --git a/litellm/tracing/normalizers/messages.py b/litellm/tracing/messages.py similarity index 61% rename from litellm/tracing/normalizers/messages.py rename to litellm/tracing/messages.py index 8a9aa914dfd..09f57e25308 100644 --- a/litellm/tracing/normalizers/messages.py +++ b/litellm/tracing/messages.py @@ -1,7 +1,7 @@ import json from collections.abc import Mapping from types import MappingProxyType -from typing import Any, Final, Literal, TypeAlias +from typing import Final, Literal, TypeAlias from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError @@ -37,20 +37,3 @@ def content_text(content: object) -> str: if not all(block.text is not None or block.type in _NON_TEXT_BLOCKS for block in blocks): return json.dumps(content) return "\n\n".join(block.text for block in blocks if block.text is not None) - - -def lc_message(message: Mapping[str, Any]) -> dict[str, Any]: - """LangChain serialized message (or plain {role, content}) -> {role, content, tool_calls?}.""" - kwargs: Final = message.get("kwargs", message) - role: Final = MESSAGE_ROLES.get( - kwargs.get("type") or kwargs.get("role"), kwargs.get("role") or kwargs.get("type") or "" - ) - out: Final[dict[str, Any]] = { # mutable-ok: the framework message is built for JSON serialization - "role": role, - "content": content_text(kwargs.get("content", "")), - } - if kwargs.get("tool_calls"): - out["tool_calls"] = tuple({"name": t.get("name"), "args": t.get("args")} for t in kwargs["tool_calls"]) - if role == "tool" and kwargs.get("name"): - out["name"] = kwargs["name"] - return out diff --git a/litellm/tracing/normalizers/__init__.py b/litellm/tracing/normalizers/__init__.py deleted file mode 100644 index 2f861330a36..00000000000 --- a/litellm/tracing/normalizers/__init__.py +++ /dev/null @@ -1,32 +0,0 @@ -"""Per-convention span normalizers, tried in order: the first whose `matches()` is true wins.""" - -from collections.abc import Mapping, Sequence -from typing import Final - -from litellm.tracing.normalizers.base import SpanNormalizer -from litellm.tracing.normalizers.genai import GenAISemconvNormalizer -from litellm.tracing.normalizers.langsmith import LangSmithNormalizer -from litellm.tracing.normalizers.openinference import OpenInferenceNormalizer - -NORMALIZERS: Final[tuple[SpanNormalizer, ...]] = ( - LangSmithNormalizer(), - OpenInferenceNormalizer(), - GenAISemconvNormalizer(), -) -_FALLBACK: Final[SpanNormalizer] = GenAISemconvNormalizer() - - -def select_normalizer( - scope_name: str, attributes: Mapping[str, str], registry: Sequence[SpanNormalizer] = NORMALIZERS -) -> SpanNormalizer: - return next((n for n in registry if n.matches(scope_name, attributes)), _FALLBACK) - - -__all__ = ( - "NORMALIZERS", - "GenAISemconvNormalizer", - "LangSmithNormalizer", - "OpenInferenceNormalizer", - "SpanNormalizer", - "select_normalizer", -) diff --git a/litellm/tracing/normalizers/base.py b/litellm/tracing/normalizers/base.py deleted file mode 100644 index 37735113ce2..00000000000 --- a/litellm/tracing/normalizers/base.py +++ /dev/null @@ -1,22 +0,0 @@ -from collections.abc import Mapping -from typing import Protocol - -from litellm.tracing.types import SpanRow - - -class SpanNormalizer(Protocol): - """Maps one tracing convention's span attributes onto the LiteLLM `SpanRow` columns.""" - - @property - def name(self) -> str: ... - - def matches(self, scope_name: str, attributes: Mapping[str, str]) -> bool: ... - - def normalize(self, row: SpanRow, attributes: Mapping[str, str]) -> None: ... - - -def to_int(value: str | None) -> int: - try: - return int(value) if value else 0 - except ValueError: - return 0 diff --git a/litellm/tracing/normalizers/genai.py b/litellm/tracing/normalizers/genai.py deleted file mode 100644 index 16986607396..00000000000 --- a/litellm/tracing/normalizers/genai.py +++ /dev/null @@ -1,31 +0,0 @@ -from collections.abc import Mapping -from dataclasses import dataclass -from typing import Final - -from litellm.tracing.types import SpanRow - -_LLM_OPERATIONS: Final = frozenset({"chat", "text_completion", "generate_content"}) - - -@dataclass(frozen=True, slots=True) -class GenAISemconvNormalizer: - """OTEL `gen_ai.*` semantic conventions. Matches every span, so it belongs last as the fallback.""" - - name: str = "genai" - - def matches(self, scope_name: str, attributes: Mapping[str, str]) -> bool: - return True - - def normalize(self, row: SpanRow, attributes: Mapping[str, str]) -> None: - operation: Final = attributes.get("gen_ai.operation.name", "") - if operation == "invoke_agent" or not row["ParentSpanId"]: - row["ObservationType"] = "agent" - elif operation in _LLM_OPERATIONS: - row["ObservationType"] = "llm" - elif operation == "execute_tool": - row["ObservationType"] = "tool" - row["AgentName"] = attributes.get("gen_ai.agent.name", "") - row["Model"] = attributes.get("gen_ai.request.model") or attributes.get("gen_ai.response.model", "") - row["LiteLLMRequestId"] = attributes.get("gen_ai.response.id", "") - row["Input"] = attributes.get("gen_ai.input.messages") or attributes.get("gen_ai.tool.call.arguments", "") - row["Output"] = attributes.get("gen_ai.output.messages") or attributes.get("gen_ai.tool.call.result", "") diff --git a/litellm/tracing/normalizers/langsmith.py b/litellm/tracing/normalizers/langsmith.py deleted file mode 100644 index daca932d57d..00000000000 --- a/litellm/tracing/normalizers/langsmith.py +++ /dev/null @@ -1,115 +0,0 @@ -"""LangSmith OTEL mode, which LangChain, LangGraph and Deep Agents export through.""" - -import json -from collections.abc import Mapping -from dataclasses import dataclass -from types import MappingProxyType -from typing import Final - -from litellm.tracing.normalizers.messages import lc_message -from litellm.tracing.types import SpanRow, SpanType - -# LangChain / Deep Agents middleware wrappers: real spans, but noise in the UI -_FRAMEWORK_SUFFIXES: Final = ( - ".wrap_model_call", - ".wrap_tool_call", - ".before_agent", - ".after_agent", - ".before_model", - ".after_model", -) - - -def _loads(value: str) -> object: - try: - return json.loads(value) - except (ValueError, TypeError): - return None - - -def _span_type(row: SpanRow, attributes: Mapping[str, str]) -> SpanType: - kind: Final = attributes.get("langsmith.span.kind", "chain") - name: Final = row["SpanName"] - if not row["ParentSpanId"] or name == attributes.get("langsmith.metadata.lc_agent_name"): - return "agent" - if kind in ("llm", "tool"): - return kind - if name.endswith(_FRAMEWORK_SUFFIXES): - return "framework" - return "chain" - - -def _tool_output(completion: object) -> object: - raw: Final = completion.get("output", completion) if isinstance(completion, dict) else completion - update: Final = raw.get("update") if isinstance(raw, dict) else None - update_messages: Final = update.get("messages") or () if isinstance(update, dict) else () - is_command: Final = isinstance(raw, dict) and "update" in raw - # LangGraph Command (e.g. the Deep Agents `task` tool): the result is the last update message - output: Final = update_messages[-1] if is_command and update_messages else raw - return output.get("content", output) if isinstance(output, dict) else output - - -def _set_agent_io(row: SpanRow, attributes: Mapping[str, str], prompt: object, completion: object) -> None: - input_messages: Final = prompt.get("messages") if isinstance(prompt, dict) else None - output_messages: Final = completion.get("messages") if isinstance(completion, dict) else None - # agents built with @traceable take arbitrary args, not a message list: keep the raw payload then - row["Input"] = ( - json.dumps(tuple(lc_message(m) for m in input_messages if isinstance(m, dict))) - if input_messages - else attributes.get("gen_ai.prompt", "") - ) - row["Output"] = ( - json.dumps(lc_message(output_messages[-1])) - if output_messages and isinstance(output_messages[-1], dict) - else attributes.get("gen_ai.completion", "") - ) - - -def _set_io(row: SpanRow, attributes: Mapping[str, str]) -> None: - prompt: Final = _loads(attributes.get("gen_ai.prompt", "")) - completion: Final = _loads(attributes.get("gen_ai.completion", "")) - if row["ObservationType"] == "llm" and isinstance(completion, dict): - prompt_payload: Final = prompt if isinstance(prompt, dict) else MappingProxyType({}) - messages: Final = prompt_payload.get("messages") or ((),) - batch: Final = messages[0] if messages and isinstance(messages[0], list) else messages - row["Input"] = ( - json.dumps(tuple(lc_message(m) for m in batch if isinstance(m, dict))) - if isinstance(batch, (list, tuple)) - else "" - ) - generations: Final = completion.get("generations") - first: Final = generations[0] if isinstance(generations, list) and generations else None - item: Final = first[0] if isinstance(first, list) and first else None - message: Final = item.get("message") if isinstance(item, dict) else None - generation: Final = message.get("kwargs") if isinstance(message, dict) else None - if isinstance(generation, dict): - row["Output"] = json.dumps(lc_message(generation)) - metadata: Final = generation.get("response_metadata") - row["LiteLLMRequestId"] = metadata.get("id", "") if isinstance(metadata, dict) else "" - return - row["Output"] = attributes.get("gen_ai.completion", "") - return - if row["ObservationType"] == "tool": - output: Final = _tool_output(completion) - row["Input"] = attributes.get("gen_ai.prompt", "") - row["Output"] = output if isinstance(output, str) else json.dumps(output) - return - if row["ObservationType"] == "agent": - _set_agent_io(row, attributes, prompt, completion) - return - row["Input"] = attributes.get("gen_ai.prompt", "") - row["Output"] = attributes.get("gen_ai.completion", "") - - -@dataclass(frozen=True, slots=True) -class LangSmithNormalizer: - name: str = "langsmith" - - def matches(self, scope_name: str, attributes: Mapping[str, str]) -> bool: - return scope_name == "langsmith" or "langsmith.span.kind" in attributes - - def normalize(self, row: SpanRow, attributes: Mapping[str, str]) -> None: - row["ObservationType"] = _span_type(row, attributes) - row["AgentName"] = attributes.get("langsmith.metadata.lc_agent_name", "") - row["Model"] = attributes.get("gen_ai.request.model", "") - _set_io(row, attributes) diff --git a/litellm/tracing/normalizers/openinference.py b/litellm/tracing/normalizers/openinference.py deleted file mode 100644 index f9e1295148c..00000000000 --- a/litellm/tracing/normalizers/openinference.py +++ /dev/null @@ -1,27 +0,0 @@ -from collections.abc import Mapping -from dataclasses import dataclass -from types import MappingProxyType -from typing import Final - -from litellm.tracing.normalizers.base import to_int -from litellm.tracing.types import SpanRow, SpanType - -_OPENINFERENCE_TYPES: Final[Mapping[str, SpanType]] = MappingProxyType({"AGENT": "agent", "LLM": "llm", "TOOL": "tool"}) - - -@dataclass(frozen=True, slots=True) -class OpenInferenceNormalizer: - name: str = "openinference" - - def matches(self, scope_name: str, attributes: Mapping[str, str]) -> bool: - return "openinference.span.kind" in attributes - - def normalize(self, row: SpanRow, attributes: Mapping[str, str]) -> None: - kind: Final = attributes.get("openinference.span.kind", "").upper() - row["ObservationType"] = _OPENINFERENCE_TYPES.get(kind, "agent" if not row["ParentSpanId"] else "chain") - row["AgentName"] = attributes.get("agent.name", "") - row["Model"] = attributes.get("llm.model_name", "") - row["Input"] = attributes.get("input.value", "") - row["Output"] = attributes.get("output.value", "") - row["InputTokens"] = to_int(attributes.get("llm.token_count.prompt")) - row["OutputTokens"] = to_int(attributes.get("llm.token_count.completion")) diff --git a/litellm/tracing/receiver.py b/litellm/tracing/receiver.py index 6cef84ec6d0..1fae2361572 100644 --- a/litellm/tracing/receiver.py +++ b/litellm/tracing/receiver.py @@ -108,7 +108,7 @@ class TraceReceiver: ClickHouseStorage( database=os.getenv("CLICKHOUSE_DATABASE", "litellm"), url=os.environ["CLICKHOUSE_URL"], - reader_url=os.environ["CLICKHOUSE_READER_URL"], + reader_url=os.getenv("CLICKHOUSE_READER_URL", os.environ["CLICKHOUSE_URL"]), ) ) ) diff --git a/litellm/tracing/ui_format.py b/litellm/tracing/ui_format.py index d7ecf48078f..51ec2876fbb 100644 --- a/litellm/tracing/ui_format.py +++ b/litellm/tracing/ui_format.py @@ -7,7 +7,7 @@ from typing import Final, Literal, TypeAlias from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter, ValidationError from typing_extensions import NotRequired, ReadOnly, TypedDict -from litellm.tracing.normalizers.messages import MESSAGE_ROLES, ChatRole, content_text +from litellm.tracing.messages import MESSAGE_ROLES, ChatRole, content_text class UIToolCall(TypedDict): diff --git a/litellm/types/llms/anthropic.py b/litellm/types/llms/anthropic.py index 6cd0e55c517..ee357cd6581 100644 --- a/litellm/types/llms/anthropic.py +++ b/litellm/types/llms/anthropic.py @@ -777,6 +777,7 @@ ANTHROPIC_EFFORT_BETA_HEADER: Final = "effort-2025-11-24" ANTHROPIC_MID_CONVERSATION_OUTPUT_CONFIG_BETA_HEADER: Final = "mid-conversation-output-config-2026-07-01" ANTHROPIC_THINKING_DISPLAY_UPDATES_BETA_HEADER: Final = "thinking-display-updates-2026-08-18" +ANTHROPIC_MID_CONVERSATION_TOOL_CHANGES_BETA_HEADER: Final = "mid-conversation-tool-changes-2026-07-01" ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER: Final = "fine-grained-tool-streaming-2025-05-14" diff --git a/litellm/types/llms/custom_http.py b/litellm/types/llms/custom_http.py index 858123b5232..47f80c52845 100644 --- a/litellm/types/llms/custom_http.py +++ b/litellm/types/llms/custom_http.py @@ -36,6 +36,7 @@ class httpxSpecialProvider(str, Enum): ModelCostMap = "model_cost_map" PasswordBreachCheck = "password_breach_check" ASGI = "asgi" + AgentHarness = "agent_harness" VerifyTypes = str | bool | ssl.SSLContext diff --git a/litellm/types/proxy/management_endpoints/common_daily_activity.py b/litellm/types/proxy/management_endpoints/common_daily_activity.py index 28488ba7de6..37804032569 100644 --- a/litellm/types/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/types/proxy/management_endpoints/common_daily_activity.py @@ -101,6 +101,22 @@ class DailySpendMetadata(BaseModel): page: int = Field(default=1) total_pages: int = Field(default=1) has_more: bool = Field(default=False) + api_key_limit: int | None = Field( + default=None, + description="When set, api_keys and every api_key_breakdown list at most this many keys, " + "ranked by spend. Totals and the model, provider, mcp and endpoint rollups still cover every key.", + ) + total_api_keys: int | None = Field( + default=None, + description="Distinct API keys matching the filters. When this exceeds api_key_limit, the per-key " + "lists are truncated to the highest-spend keys.", + ) + entity_total_api_keys: dict[str, int] | None = Field( + default=None, + description="Distinct API keys per entity over the requested range, set when the entity breakdown is " + "included. When an entity's count exceeds api_key_limit, its api_key_breakdown lists only its keys " + "among the top api_key_limit keys overall.", + ) class SpendAnalyticsPaginatedResponse(BaseModel): @@ -108,6 +124,51 @@ class SpendAnalyticsPaginatedResponse(BaseModel): metadata: DailySpendMetadata = Field(default_factory=DailySpendMetadata) +class KeyActivityRow(BaseModel): + api_key: str + metrics: SpendMetrics + metadata: KeyMetadata + + +class KeySpendMetrics(BaseModel): + spend: float = 0.0 + prompt_tokens: int = 0 + completion_tokens: int = 0 + total_tokens: int = 0 + api_requests: int = 0 + successful_requests: int = 0 + failed_requests: int = 0 + cache_read_input_tokens: int = 0 + cache_creation_input_tokens: int = 0 + + +class KeySpendActivityRow(BaseModel): + api_key: str + metrics: KeySpendMetrics + metadata: KeyMetadata + + +class DailyActivityKeySearchResponse(BaseModel): + api_keys: list[KeyActivityRow] + + +class DailyActivityKeyPageResponse(BaseModel): + api_keys: list[KeySpendActivityRow] + total_api_keys: int + offset: int + limit: int + + +class ModelTopKeysResponse(BaseModel): + model: str + by_model_group: bool + api_keys: list[KeySpendActivityRow] + + +class CacheLeakageKeysResponse(BaseModel): + api_keys: list[KeySpendActivityRow] + + class LiteLLM_DailyUserSpend(BaseModel): id: str user_id: str diff --git a/litellm/types/repositories/__init__.py b/litellm/types/repositories/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/types/repositories/daily_activity.py b/litellm/types/repositories/daily_activity.py new file mode 100644 index 00000000000..df302234398 --- /dev/null +++ b/litellm/types/repositories/daily_activity.py @@ -0,0 +1,192 @@ +from collections.abc import Mapping +from dataclasses import dataclass, field +from datetime import datetime +from enum import Enum +from types import MappingProxyType +from typing import Protocol, TypeAlias + + +class DailyActivityTable(str, Enum): + USER = "litellm_dailyuserspend" + TEAM = "litellm_dailyteamspend" + TAG = "litellm_dailytagspend" + ORGANIZATION = "litellm_dailyorganizationspend" + CUSTOMER = "litellm_dailyenduserspend" + AGENT = "litellm_dailyagentspend" + + +_ENTITY_FIELDS: Mapping[DailyActivityTable, frozenset[str]] = MappingProxyType( + { + DailyActivityTable.USER: frozenset(("user_id",)), + DailyActivityTable.TEAM: frozenset(("team_id",)), + DailyActivityTable.TAG: frozenset(("tag",)), + DailyActivityTable.ORGANIZATION: frozenset(("organization_id",)), + DailyActivityTable.CUSTOMER: frozenset(("end_user_id",)), + DailyActivityTable.AGENT: frozenset(("agent_id",)), + } +) + + +@dataclass(frozen=True, slots=True) +class DailyActivityScope: + table: DailyActivityTable + entity_id_field: str + entity_ids: tuple[str, ...] | None + exclude_entity_ids: tuple[str, ...] + api_keys: tuple[str, ...] | None + start_date: str + end_date: str + model: str | None + timezone_offset_minutes: int | None + include_current_utc_day: bool = False + + def __post_init__(self) -> None: + if self.entity_id_field not in _ENTITY_FIELDS[self.table]: + raise ValueError(f"Invalid entity_id_field {self.entity_id_field!r} for {self.table.value}") + + +@dataclass(frozen=True, slots=True) +class KeySpendRow: + api_key: str + spend: float + prompt_tokens: int + completion_tokens: int + total_tokens: int + api_requests: int + successful_requests: int + failed_requests: int + cache_read_input_tokens: int + cache_creation_input_tokens: int + + +@dataclass(frozen=True, slots=True) +class KeyPage: + rows: tuple[KeySpendRow, ...] + total_api_keys: int + + +@dataclass(frozen=True, slots=True) +class KeyMetadataRow: + api_key: str + key_alias: str | None + team_id: str | None + user_id: str | None + user_email: str | None + key_exists: bool + tags: tuple[str, ...] + + +class ExportType(str, Enum): + DAILY = "daily" + DAILY_WITH_KEYS = "daily_with_keys" + DAILY_WITH_MODELS = "daily_with_models" + DAILY_WITH_USERS = "daily_with_users" + + +@dataclass(frozen=True, slots=True) +class ExportRow: + date: str + entity_id: str + entity_alias: str | None + api_key: str | None + key_alias: str | None + user_id: str | None + user_email: str | None + model: str | None + spend: float + flat_cost: float + prompt_tokens: int + completion_tokens: int + api_requests: int + successful_requests: int + failed_requests: int + cache_read_input_tokens: int + cache_creation_input_tokens: int + + +@dataclass(frozen=True, slots=True) +class RollupMetricsRow: + date: str | None + api_key: str | None + spend: float | None + ptu_flat_cost: float | None = field(default=None, kw_only=True) + prompt_tokens: int | None + completion_tokens: int | None + cache_read_input_tokens: int | None + cache_creation_input_tokens: int | None + compression_saved_tokens: int | None + compression_savings_spend: float | None + prompt_caching_savings_spend: float | None + gateway_injected_caching_savings_spend: float | None + autorouter_savings_spend: float | None + api_requests: int | None + successful_requests: int | None + failed_requests: int | None + total_response_time_ms: int | None + timed_requests: int | None + + +@dataclass(frozen=True, slots=True) +class GroupingSetsRow(RollupMetricsRow): + model: str | None + model_group: str | None + custom_llm_provider: str | None + mcp_namespaced_tool_name: str | None + endpoint: str | None + group_level: int + distinct_api_keys: int | None + + +@dataclass(frozen=True, slots=True) +class EntityRollupRow(RollupMetricsRow): + entity_id: str | None + api_key_rolled: int + distinct_api_keys: int | None + + +@dataclass(frozen=True, slots=True) +class AggregatedRows: + grouping_rows: tuple[GroupingSetsRow, ...] + entity_rows: tuple[EntityRollupRow, ...] | None + distinct_api_keys: int + + +SpendLogsWindow: TypeAlias = tuple[datetime, datetime] + + +class DailyActivityProxyReads(Protocol): + async def recover_key_metadata( + self, resolved: Mapping[str, KeyMetadataRow], api_keys: frozenset[str], window: SpendLogsWindow | None + ) -> Mapping[str, KeyMetadataRow]: ... + + +class DailyActivityRow(Protocol): + id: str + date: str + api_key: str + model: str | None + model_group: str | None + custom_llm_provider: str | None + mcp_namespaced_tool_name: str | None + endpoint: str | None + prompt_tokens: int + completion_tokens: int + cache_read_input_tokens: int + cache_creation_input_tokens: int + compression_saved_tokens: int + compression_savings_spend: float + prompt_caching_savings_spend: float + gateway_injected_caching_savings_spend: float + autorouter_savings_spend: float + spend: float + api_requests: int + successful_requests: int + failed_requests: int + total_response_time_ms: int + timed_requests: int + + +@dataclass(frozen=True, slots=True) +class DailyRowsPage: + total_count: int + rows: tuple[DailyActivityRow, ...] diff --git a/litellm/utils.py b/litellm/utils.py index 0eb2754ed0c..05d5986885c 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -308,6 +308,7 @@ if TYPE_CHECKING: CachingHandlerResponse, LLMCachingHandler, ) + from litellm.harness.types import Harness from litellm.integrations.custom_logger import CustomLogger # Type stubs for lazy-loaded functions and classes @@ -385,6 +386,7 @@ if TYPE_CHECKING: from litellm.llms.base_llm.google_genai.transformation import ( BaseGoogleGenAIGenerateContentConfig, ) + from litellm.llms.base_llm.harness.transformation import BaseHarnessConfig from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig from litellm.llms.base_llm.image_generation.transformation import ( BaseImageGenerationConfig, @@ -9883,6 +9885,37 @@ class ProviderConfigManager: return OpenSandboxSandboxConfig() return None + @staticmethod + def get_provider_harness_config(harness: Harness) -> BaseHarnessConfig | None: + """ + Get the agent-harness configuration (Claude Code, Codex, OpenCode, Deep Agents). + """ + from litellm.harness.types import Harness as _Harness + + if harness == _Harness.CLAUDE_CODE: + from litellm.llms.claude_code.harness.transformation import ( + ClaudeCodeHarnessConfig, + ) + + return ClaudeCodeHarnessConfig() + if harness == _Harness.CODEX: + from litellm.llms.codex.harness.transformation import CodexHarnessConfig + + return CodexHarnessConfig() + if harness == _Harness.OPENCODE: + from litellm.llms.opencode.harness.transformation import ( + OpenCodeHarnessConfig, + ) + + return OpenCodeHarnessConfig() + if harness == _Harness.DEEPAGENTS: + from litellm.llms.deepagents.harness.transformation import ( + DeepAgentsHarnessConfig, + ) + + return DeepAgentsHarnessConfig() + return None + @staticmethod def get_provider_text_to_speech_config( model: str, diff --git a/migrations/run.py b/migrations/run.py index 7ea80d48719..94e7cc7c0f0 100644 --- a/migrations/run.py +++ b/migrations/run.py @@ -2,7 +2,12 @@ Runs `prisma migrate deploy` against the LiteLLM writer database using the recovery logic in `litellm_proxy_extras.ProxyExtrasDBManager.setup_database` -(P3005 baseline + P3009/P3018 idempotent-error handling, retries, etc.). +(P3005 baseline + P3009/P3018 idempotent-error handling, retries, etc.), then +builds the request-log indexes the migrations leave out +(`litellm_proxy_extras.request_log_indexes`), waiting for them. The job exits +non-zero when an index could not be built so that it is rerun. A serving proxy +that runs the migrations itself builds the same indexes in the background once +it serves. Env vars: DATABASE_URL required unless it can be assembled at @@ -23,10 +28,11 @@ Env vars: import os import sys -from litellm.proxy.db.db_url_settings import DatabaseURLSettings from litellm_proxy_extras._logging import logger from litellm_proxy_extras.utils import ProxyExtrasDBManager, str_to_bool +from litellm.proxy.db.db_url_settings import DatabaseURLSettings + def main() -> int: # Assemble DATABASE_URL from the discrete DATABASE_* env vars, matching @@ -52,7 +58,7 @@ def main() -> int: not use_db_push, use_v2, ) - ok = ProxyExtrasDBManager.setup_database( + ok = ProxyExtrasDBManager.run_migration_job( use_migrate=not use_db_push, use_v2_resolver=use_v2, ) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 40fdf083cf8..02e658ca779 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -5516,7 +5516,7 @@ "supports_web_search": false }, "azure/gpt-4.1-nano": { - "deprecation_date": "2026-10-14", + "deprecation_date": "2027-04-14", "cache_read_input_token_cost": 2.5e-08, "input_cost_per_token": 1e-07, "input_cost_per_token_batches": 5e-08, @@ -5550,7 +5550,7 @@ "supports_vision": true }, "azure/gpt-4.1-nano-2025-04-14": { - "deprecation_date": "2026-10-14", + "deprecation_date": "2027-04-14", "cache_read_input_token_cost": 2.5e-08, "input_cost_per_token": 1e-07, "input_cost_per_token_batches": 5e-08, @@ -6133,7 +6133,7 @@ "cache_creation_input_audio_token_cost": 4e-07, "cache_read_input_audio_token_cost": 4e-07, "cache_read_input_token_cost": 4e-07, - "deprecation_date": "2027-06-25", + "deprecation_date": "2027-07-31", "input_cost_per_audio_token": 3.2e-05, "input_cost_per_image_token": 5e-06, "input_cost_per_token": 4e-06, @@ -6168,7 +6168,7 @@ "cache_creation_input_audio_token_cost": 3e-07, "cache_read_input_audio_token_cost": 3e-07, "cache_read_input_token_cost": 6e-08, - "deprecation_date": "2027-06-25", + "deprecation_date": "2027-07-31", "input_cost_per_audio_token": 1e-05, "input_cost_per_image_token": 8e-07, "input_cost_per_token": 6e-07, @@ -6346,7 +6346,7 @@ "supports_tool_choice": true }, "azure/gpt-4o-transcribe": { - "deprecation_date": "2026-10-15", + "deprecation_date": "2026-12-31", "input_cost_per_audio_token": 2.5e-06, "input_cost_per_token": 2.5e-06, "litellm_provider": "azure", @@ -6354,6 +6354,7 @@ "max_output_tokens": 2000, "mode": "audio_transcription", "output_cost_per_token": 1e-05, + "source": "https://management.azure.com/subscriptions/c873328e-b572-4770-8dff-aaeb6f1f0e79/providers/Microsoft.CognitiveServices/locations/eastus2/models?api-version=2024-10-01", "supported_endpoints": [ "/v1/audio/transcriptions" ] @@ -8571,6 +8572,7 @@ "cache_creation_input_token_cost_above_272k_tokens": 5e-06, "cache_read_input_token_cost": 1e-07, "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "deprecation_date": "2028-03-11", "input_cost_per_token": 2e-06, "input_cost_per_token_above_272k_tokens": 4e-06, "litellm_provider": "azure", @@ -8619,6 +8621,7 @@ "cache_creation_input_token_cost_above_272k_tokens": 5e-06, "cache_read_input_token_cost": 1e-07, "cache_read_input_token_cost_above_272k_tokens": 2e-07, + "deprecation_date": "2028-03-11", "input_cost_per_token": 2e-06, "input_cost_per_token_above_272k_tokens": 4e-06, "litellm_provider": "azure", @@ -10972,7 +10975,7 @@ "supports_web_search": false }, "azure/us/gpt-4.1-nano-2025-04-14": { - "deprecation_date": "2026-10-14", + "deprecation_date": "2027-04-14", "cache_read_input_token_cost": 2.8e-08, "input_cost_per_token": 1.1e-07, "input_cost_per_token_batches": 5.5e-08, @@ -11958,6 +11961,7 @@ "supports_vision": true }, "azure_ai/Meta-Llama-3-70B-Instruct": { + "deprecation_date": "2025-06-30", "input_cost_per_token": 1.1e-06, "litellm_provider": "azure_ai", "max_input_tokens": 8192, @@ -11965,9 +11969,11 @@ "max_tokens": 2048, "mode": "chat", "output_cost_per_token": 3.7e-07, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true }, "azure_ai/Meta-Llama-3.1-70B-Instruct": { + "deprecation_date": "2025-06-30", "input_cost_per_token": 2.68e-06, "litellm_provider": "azure_ai", "max_input_tokens": 128000, @@ -11975,10 +11981,11 @@ "max_tokens": 2048, "mode": "chat", "output_cost_per_token": 3.54e-06, - "source": "https://marketplace.microsoft.com/en-us/marketplace/apps/metagenai.meta-llama-3-1-70b-instruct-offer?tab=PlansAndPrice", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true }, "azure_ai/Phi-3-medium-128k-instruct": { + "deprecation_date": "2025-08-30", "input_cost_per_token": 1.7e-07, "litellm_provider": "azure_ai", "max_input_tokens": 128000, @@ -11986,11 +11993,12 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6.8e-07, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true, "supports_vision": false }, "azure_ai/Phi-3-medium-4k-instruct": { + "deprecation_date": "2025-08-30", "input_cost_per_token": 1.7e-07, "litellm_provider": "azure_ai", "max_input_tokens": 4096, @@ -11998,11 +12006,12 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6.8e-07, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true, "supports_vision": false }, "azure_ai/Phi-3-mini-128k-instruct": { + "deprecation_date": "2025-08-30", "input_cost_per_token": 1.3e-07, "litellm_provider": "azure_ai", "max_input_tokens": 128000, @@ -12010,11 +12019,12 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 5.2e-07, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true, "supports_vision": false }, "azure_ai/Phi-3-mini-4k-instruct": { + "deprecation_date": "2025-08-30", "input_cost_per_token": 1.3e-07, "litellm_provider": "azure_ai", "max_input_tokens": 4096, @@ -12022,11 +12032,12 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 5.2e-07, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true, "supports_vision": false }, "azure_ai/Phi-3-small-128k-instruct": { + "deprecation_date": "2025-08-30", "input_cost_per_token": 1.5e-07, "litellm_provider": "azure_ai", "max_input_tokens": 128000, @@ -12034,11 +12045,12 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-07, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true, "supports_vision": false }, "azure_ai/Phi-3-small-8k-instruct": { + "deprecation_date": "2025-08-30", "input_cost_per_token": 1.5e-07, "litellm_provider": "azure_ai", "max_input_tokens": 8192, @@ -12046,11 +12058,12 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-07, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true, "supports_vision": false }, "azure_ai/Phi-3.5-MoE-instruct": { + "deprecation_date": "2025-08-30", "input_cost_per_token": 1.6e-07, "litellm_provider": "azure_ai", "max_input_tokens": 128000, @@ -12058,11 +12071,12 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6.4e-07, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true, "supports_vision": false }, "azure_ai/Phi-3.5-mini-instruct": { + "deprecation_date": "2025-08-30", "input_cost_per_token": 1.3e-07, "litellm_provider": "azure_ai", "max_input_tokens": 128000, @@ -12070,11 +12084,12 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 5.2e-07, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true, "supports_vision": false }, "azure_ai/Phi-3.5-vision-instruct": { + "deprecation_date": "2025-08-30", "input_cost_per_token": 1.3e-07, "litellm_provider": "azure_ai", "max_input_tokens": 128000, @@ -12082,7 +12097,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 5.2e-07, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true, "supports_vision": true }, @@ -12217,6 +12232,7 @@ "source": "https://azure.microsoft.com/en-us/pricing/details/ai-document-intelligence/" }, "azure_ai/MAI-DS-R1": { + "deprecation_date": "2026-02-27", "input_cost_per_token": 1.35e-06, "litellm_provider": "azure_ai", "max_input_tokens": 128000, @@ -12224,11 +12240,12 @@ "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 5.4e-06, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_reasoning": true, "supports_tool_choice": true }, "azure_ai/cohere-rerank-v3-english": { + "deprecation_date": "2025-06-30", "input_cost_per_query": 0.002, "input_cost_per_token": 0.0, "litellm_provider": "azure_ai", @@ -12236,9 +12253,11 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "rerank", - "output_cost_per_token": 0.0 + "output_cost_per_token": 0.0, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models" }, "azure_ai/cohere-rerank-v3-multilingual": { + "deprecation_date": "2025-06-30", "input_cost_per_query": 0.002, "input_cost_per_token": 0.0, "litellm_provider": "azure_ai", @@ -12246,7 +12265,8 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "rerank", - "output_cost_per_token": 0.0 + "output_cost_per_token": 0.0, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models" }, "azure_ai/cohere-rerank-v4.0-pro": { "input_cost_per_query": 0.0025, @@ -12301,6 +12321,7 @@ "supports_tool_choice": true }, "azure_ai/deepseek-v3": { + "deprecation_date": "2025-08-31", "input_cost_per_token": 1.14e-06, "litellm_provider": "azure_ai", "max_input_tokens": 128000, @@ -12308,7 +12329,7 @@ "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 4.56e-06, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_tool_choice": true }, "azure_ai/deepseek-v4-pro": { @@ -12515,6 +12536,7 @@ "supports_web_search": true }, "azure_ai/jais-30b-chat": { + "deprecation_date": "2026-01-30", "input_cost_per_token": 0.0032, "litellm_provider": "azure_ai", "max_input_tokens": 8192, @@ -12522,7 +12544,7 @@ "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 0.00971, - "source": "https://ai.azure.com/catalog/models/jais-30b-chat" + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models" }, "azure_ai/jamba-instruct": { "input_cost_per_token": 5e-07, @@ -12588,6 +12610,7 @@ "supports_tool_choice": true }, "azure_ai/mistral-large": { + "deprecation_date": "2025-04-15", "input_cost_per_token": 4e-06, "litellm_provider": "azure_ai", "max_input_tokens": 32000, @@ -12595,10 +12618,12 @@ "max_tokens": 8191, "mode": "chat", "output_cost_per_token": 1.2e-05, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_function_calling": true, "supports_tool_choice": true }, "azure_ai/mistral-large-2407": { + "deprecation_date": "2025-05-13", "input_cost_per_token": 2e-06, "litellm_provider": "azure_ai", "max_input_tokens": 128000, @@ -12606,7 +12631,7 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-06, - "source": "https://marketplace.microsoft.com/en/marketplace/apps/000-000.mistral-ai-large-2407-offer?tab=Overview", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_function_calling": true, "supports_tool_choice": true }, @@ -12648,6 +12673,7 @@ "supports_tool_choice": true }, "azure_ai/mistral-nemo": { + "deprecation_date": "2026-01-30", "input_cost_per_token": 1.5e-07, "litellm_provider": "azure_ai", "max_input_tokens": 131072, @@ -12655,10 +12681,11 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 1.5e-07, - "source": "https://marketplace.microsoft.com/en/marketplace/apps/000-000.mistral-nemo-12b-2407?tab=PlansAndPrice", + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_function_calling": true }, "azure_ai/mistral-small": { + "deprecation_date": "2025-07-31", "input_cost_per_token": 1e-06, "litellm_provider": "azure_ai", "max_input_tokens": 32000, @@ -12666,6 +12693,7 @@ "max_tokens": 8191, "mode": "chat", "output_cost_per_token": 3e-06, + "source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models", "supports_function_calling": true, "supports_tool_choice": true }, @@ -42206,14 +42234,14 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro": { - "cache_read_input_token_cost": 6.525e-08, - "input_cost_per_token": 7.83e-07, + "cache_read_input_token_cost": 1.74e-08, + "input_cost_per_token": 2.088e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 1.566e-06, + "output_cost_per_token": 4.176e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42226,14 +42254,14 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4.1-flash": { - "cache_read_input_token_cost": 2.91e-09, - "input_cost_per_token": 1.98e-08, + "cache_read_input_token_cost": 1e-08, + "input_cost_per_token": 3e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 3.96e-07, + "output_cost_per_token": 5e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -42820,14 +42848,14 @@ "supports_web_search": false }, "openrouter/nvidia/nemotron-3.5-lightning": { - "cache_read_input_token_cost": 3e-08, - "input_cost_per_token": 6e-08, + "cache_read_input_token_cost": 2.975e-08, + "input_cost_per_token": 5.95e-08, "litellm_provider": "openrouter", "max_input_tokens": 262144, - "max_output_tokens": 32768, - "max_tokens": 32768, + "max_output_tokens": 131072, + "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 1.6e-07, + "output_cost_per_token": 1.7e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -43829,6 +43857,7 @@ "openrouter/z-ai/glm-4.7": { "cache_creation_input_token_cost": 0.0, "cache_read_input_token_cost": 1.1e-07, + "deprecation_date": "2026-12-31", "input_cost_per_token": 6e-07, "litellm_provider": "openrouter", "max_input_tokens": 204800, @@ -67382,13 +67411,13 @@ "supports_web_search": false }, "openrouter/z-ai/glm-5.3": { - "input_cost_per_token": 1.4e-06, - "output_cost_per_token": 4.4e-06, - "cache_read_input_token_cost": 2.6e-07, + "input_cost_per_token": 2.219e-07, + "output_cost_per_token": 3.39e-06, + "cache_read_input_token_cost": 1.775e-07, "litellm_provider": "openrouter", - "max_input_tokens": 1310720, - "max_output_tokens": 943717, - "max_tokens": 943717, + "max_input_tokens": 1048576, + "max_output_tokens": 943718, + "max_tokens": 943718, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, @@ -67519,8 +67548,8 @@ "supports_prompt_caching": true }, "openrouter/deepseek/deepseek-v4-flash-0731": { - "cache_read_input_token_cost": 8.9e-09, - "input_cost_per_token": 8.9e-09, + "cache_read_input_token_cost": 1.08e-08, + "input_cost_per_token": 1.08e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, @@ -67608,8 +67637,8 @@ "supports_web_search": false }, "openrouter/moonshotai/kimi-k3": { - "cache_read_input_token_cost": 2.7e-07, - "input_cost_per_token": 2.8e-07, + "cache_read_input_token_cost": 4.357e-07, + "input_cost_per_token": 4.357e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, @@ -67732,7 +67761,7 @@ }, "openrouter/z-ai/glm-5.2": { "cache_read_input_token_cost": 2.6e-07, - "input_cost_per_token": 3.249e-07, + "input_cost_per_token": 4.1e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, @@ -68094,14 +68123,14 @@ "supports_web_search": true }, "openrouter/deepseek/deepseek-v4-flash": { - "cache_read_input_token_cost": 1.5708e-08, - "input_cost_per_token": 7.854e-08, + "cache_read_input_token_cost": 8.372e-09, + "input_cost_per_token": 4.186e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, - "max_output_tokens": 384000, - "max_tokens": 384000, + "max_output_tokens": 131072, + "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 1.5708e-07, + "output_cost_per_token": 8.372e-08, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -68114,9 +68143,9 @@ "supports_web_search": false }, "openrouter/moonshotai/kimi-k2.6": { - "input_cost_per_token": 6.5e-07, - "output_cost_per_token": 3.41e-06, - "cache_read_input_token_cost": 1.5e-07, + "input_cost_per_token": 4.3415e-07, + "output_cost_per_token": 1.828e-06, + "cache_read_input_token_cost": 7.312e-08, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 235929, @@ -68563,7 +68592,7 @@ "openrouter/z-ai/glm-4.6v": { "input_cost_per_token": 3e-07, "output_cost_per_token": 9e-07, - "cache_read_input_token_cost": 5.5e-08, + "cache_read_input_token_cost": 5e-08, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 32768, @@ -69229,7 +69258,7 @@ "supports_web_search": false }, "openrouter/minimax/minimax-m1": { - "input_cost_per_token": 4e-07, + "input_cost_per_token": 5.5e-07, "output_cost_per_token": 2.2e-06, "litellm_provider": "openrouter", "max_input_tokens": 1000000, @@ -70653,7 +70682,7 @@ "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" }, "azure/eu/gpt-4.1-nano": { - "deprecation_date": "2026-10-14", + "deprecation_date": "2027-04-14", "cache_read_input_token_cost": 2.8e-08, "input_cost_per_token": 1.1e-07, "input_cost_per_token_batches": 5.5e-08, @@ -71102,7 +71131,7 @@ "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" }, "azure/us/gpt-4.1-nano": { - "deprecation_date": "2026-10-14", + "deprecation_date": "2027-04-14", "cache_read_input_token_cost": 2.8e-08, "input_cost_per_token": 1.1e-07, "input_cost_per_token_batches": 5.5e-08, diff --git a/osv-scanner.toml b/osv-scanner.toml index 24e6fa40c58..b3b6bb17d97 100644 --- a/osv-scanner.toml +++ b/osv-scanner.toml @@ -7,13 +7,3 @@ reason = "diskcache has no fixed release published; remove this entry once one e id = "GHSA-h7x2-h6g9-p789" ignoreUntil = 2026-10-14 reason = "mlflow has no fixed release published (3.16.0, 2026-09-04, and master still store gateway secret api_base unvalidated); remove this entry once one exists" - -[[IgnoredVulns]] -id = "GHSA-hj66-6f7g-4r5v" -ignoreUntil = 2026-10-02 -reason = "oauthlib 4.0.0 (the only fixed release, 2026-09-28) is inside the 3-day uv exclude-newer cooldown; bump oauthlib and remove this entry once it clears" - -[[IgnoredVulns]] -id = "GHSA-xpv3-w29h-x7cv" -ignoreUntil = 2026-10-02 -reason = "oauthlib 4.0.0 (the only fixed release, 2026-09-28) is inside the 3-day uv exclude-newer cooldown; bump oauthlib and remove this entry once it clears" diff --git a/pyproject.toml b/pyproject.toml index 2c5be546a65..ad1750fac0b 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -75,8 +75,8 @@ proxy = [ "mcp>=2.2.0,<3", "httpx2>=2.5.0,<3", "pydantic>=2.12.0,<3", - "litellm-proxy-extras==0.4.103", - "litellm-enterprise==0.1.72", + "litellm-proxy-extras==0.4.104", + "litellm-enterprise==0.1.73", "RestrictedPython>=8.5,<9.0", "rich>=13.9.4,<14.0", "InquirerPy>=0.3.4,<1.0", diff --git a/terraform/provider/tools/endpointaudit/coverage_allowlist.txt b/terraform/provider/tools/endpointaudit/coverage_allowlist.txt index 99ed82dcad8..fe89257e83b 100644 --- a/terraform/provider/tools/endpointaudit/coverage_allowlist.txt +++ b/terraform/provider/tools/endpointaudit/coverage_allowlist.txt @@ -12,14 +12,34 @@ # Read-only analytics and spend reporting; observability, not Terraform-managed state GET /agent/daily/activity +GET /agent/daily/activity/aggregated +GET /agent/daily/activity/aggregated/keys +GET /agent/daily/activity/aggregated/model_top_keys +GET /agent/daily/activity/aggregated/search +GET /agent/daily/activity/export GET /customer/daily/activity +GET /customer/daily/activity/aggregated +GET /customer/daily/activity/aggregated/keys +GET /customer/daily/activity/aggregated/model_top_keys +GET /customer/daily/activity/aggregated/search +GET /customer/daily/activity/export GET /guardrails/usage/detail/{guardrail_id} GET /guardrails/usage/logs GET /guardrails/usage/overview GET /key/spend/report GET /organization/daily/activity +GET /organization/daily/activity/aggregated +GET /organization/daily/activity/aggregated/keys +GET /organization/daily/activity/aggregated/model_top_keys +GET /organization/daily/activity/aggregated/search +GET /organization/daily/activity/export GET /organization/spend/report GET /tag/daily/activity +GET /tag/daily/activity/aggregated +GET /tag/daily/activity/aggregated/keys +GET /tag/daily/activity/aggregated/model_top_keys +GET /tag/daily/activity/aggregated/search +GET /tag/daily/activity/export GET /tag/dau GET /tag/distinct GET /tag/mau @@ -28,10 +48,19 @@ GET /tag/user-agent/per-user-analytics GET /tag/wau GET /team/daily/activity GET /team/daily/activity/aggregated +GET /team/daily/activity/aggregated/keys +GET /team/daily/activity/aggregated/model_top_keys +GET /team/daily/activity/aggregated/search +GET /team/daily/activity/export GET /team/spend/by_user GET /team/spend/report GET /user/daily/activity GET /user/daily/activity/aggregated +GET /user/daily/activity/aggregated/keys +GET /user/daily/activity/aggregated/cache_leakage_keys +GET /user/daily/activity/aggregated/model_top_keys +GET /user/daily/activity/aggregated/search +GET /user/daily/activity/export GET /user/spend/report # Admin UI helper endpoints; serve UI forms and caller-scoped views, not desired state diff --git a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py index d7da48ce933..5c694702c6e 100644 --- a/tests/code_coverage_tests/check_migrations_no_data_rewrites.py +++ b/tests/code_coverage_tests/check_migrations_no_data_rewrites.py @@ -7,15 +7,16 @@ anything whose cost scales with existing table size turns into downtime. A singl plus a doubled heap that plain autovacuum will not give back. What is banned is the row-rewriting DML behind that, not everything whose cost -scales that way. A non-concurrent `CREATE INDEX`, an `ALTER COLUMN ... TYPE` that is -not binary coercible, a volatile `DEFAULT` on a new column, a `CREATE TABLE ... AS -SELECT` or `SELECT ... INTO` filling a new table from an existing one, the rename -that pairs with one of those to swap a table out, and a `REFRESH MATERIALIZED VIEW` -all read the whole table and all pass. That is deliberate: a rule wide enough to -reach them fires on most ordinary migrations, and a marker everyone adds by reflex -stops carrying information. The outage this was written for was a backfill. +scales that way. A non-concurrent `CREATE INDEX` passes except on a request-log +table, where it blocks writes until the build finishes. An `ALTER COLUMN ... TYPE` +that is not binary coercible, a volatile `DEFAULT` on a new column, a `CREATE TABLE +... AS SELECT` or `SELECT ... INTO` filling a new table from an existing one, the +rename that pairs with one of those to swap a table out, and a `REFRESH MATERIALIZED +VIEW` all read the whole table and all pass. That is deliberate: a rule wide enough +to reach them fires on most ordinary migrations, and a marker everyone adds by +reflex stops carrying information. The outage this was written for was a backfill. -The one schema change banned outright is `ADD COLUMN ... DEFAULT` on a table in +One column change banned outright is `ADD COLUMN ... DEFAULT` on a table in `REQUEST_LOG_TABLES`, the tables that hold a row per request. Postgres 11 stores such a default as metadata and touches no rows, but Postgres 10, which is supported, rewrites the whole heap and rebuilds every index under an `ACCESS EXCLUSIVE` lock, @@ -23,6 +24,13 @@ which on a spend-log-sized table is the same outage as a backfill. Every other t is small enough that the rewrite is not worth a rule, and a column added to a log table without a default is still free on every version. +An index on a request-log table cannot ship as a migration at all. A plain `CREATE +INDEX` blocks writes to the table until the build finishes, and `CREATE INDEX +CONCURRENTLY` is refused by Postgres on a partitioned parent, which LiteLLM_SpendLogs +is wherever the operator ran db_scripts/partition_spend_logs.sql. The migration job builds +those indexes after `migrate deploy`, concurrently and per partition, from the list in +litellm_proxy_extras/request_log_indexes.py, so that list is where a new one goes. + Flagged, per statement, by its leading keyword: UPDATE rewrites every matching row, and `WHERE` does not bound the scan @@ -44,6 +52,8 @@ Flagged, per statement, by its leading keyword: actions adds a column with a `DEFAULT`. An `ALTER COLUMN ... SET DEFAULT` written after the column exists changes metadata alone, so it passes, as does an `ADD CONSTRAINT` + CREATE only a `CREATE [UNIQUE] INDEX` on a request-log table, concurrent or + not; the migration job builds those Referential actions (`ON DELETE CASCADE`, `ON UPDATE CASCADE`) are schema, never a statement's leading keyword, so they pass. @@ -85,7 +95,9 @@ below line up with the statements they exempt. Add a column and let the application populate it, or run the rewrite as an opt-in batched job outside boot. When a rewrite is genuinely bounded and must ship inside the migration, put `-- data-migration-ok: ` on the statement or on the line -above it, naming what bounds it. The reason is required. A marker sharing a line +above it, naming what bounds it. The reason is required. A marker never exempts a +`CREATE INDEX` on a request-log table, since no bound makes that statement safe: +the migration job is the only place such an index is built. A marker sharing a line with the statement it follows exempts that statement alone, so the next statement down is still checked rather than picking the marker up as its own. A marker on an `EXECUTE` or on the assignment feeding one covers the single-quoted SQL that @@ -108,6 +120,7 @@ import sys from collections.abc import Iterator, Mapping from dataclasses import dataclass from pathlib import Path +from typing import Final REPO_ROOT = Path(__file__).resolve().parents[2] MIGRATIONS_DIR = REPO_ROOT / "litellm-proxy-extras" / "litellm_proxy_extras" / "migrations" @@ -115,6 +128,9 @@ MIGRATIONS_DIR = REPO_ROOT / "litellm-proxy-extras" / "litellm_proxy_extras" / " GRANDFATHERED = frozenset( { "20250425182129_add_session_id", + "20250510142544_add_session_id_index_spend_logs", + "20260228100000_add_spend_logs_composite_index", + "20250326162113_baseline", "20260817000000_shadow_eval_multi_key", "20260818000000_add_spend_log_timestamps", "20260818224500_add_shadow_eval_stopped_by", @@ -139,9 +155,7 @@ WORD_OR_ASSIGN = re.compile(r"[A-Za-z_][A-Za-z0-9_]*|:=|(?!:=])=(?![=>])") PRECEDING_WORD = re.compile(r"([A-Za-z_][A-Za-z0-9_]*)[^A-Za-z0-9_]*$") QUALIFIER_GAP = re.compile(r"[\s.]*") EXPLAIN_OPTIONS = re.compile(r"\bEXPLAIN\b(?:\s+(?:ANALYZE|ANALYSE|VERBOSE)\b)+", re.IGNORECASE) -DEFINES_A_ROUTINE = re.compile( - r"\bCREATE\b(?:\s+OR\s+REPLACE)?\s+(?:FUNCTION|PROCEDURE)\b", re.IGNORECASE -) +DEFINES_A_ROUTINE = re.compile(r"\bCREATE\b(?:\s+OR\s+REPLACE)?\s+(?:FUNCTION|PROCEDURE)\b", re.IGNORECASE) QUALIFIED_NAME = r"(?:\"[^\"]*\"|[A-Za-z_][A-Za-z0-9_$]*)" ROUTINE_NAME = re.compile(rf"\s*(?:{QUALIFIED_NAME}\s*\.\s*)?({QUALIFIED_NAME})") TABLE_NAME = ROUTINE_NAME @@ -204,6 +218,12 @@ statement with the bound spelled out: -- data-migration-ok: UPDATE ... +An index on a request-log table is not a migration, and no marker exempts one. Declare it with `@@index` in +schema.prisma and add it to REQUEST_LOG_INDEXES in +litellm_proxy_extras/request_log_indexes.py under the name Prisma derives for it; the +migration job builds it after `migrate deploy`, concurrently on a plain table and per +partition on a partitioned one, which no single migration statement can do. + On Postgres 10 an `ADD COLUMN ... DEFAULT` on a request-log table rewrites the table too. Add the column nullable with no default, then set the default in a separate `ALTER COLUMN ... SET DEFAULT`, which never touches existing rows. @@ -215,10 +235,11 @@ class Violation: migration: str line: int keyword: str + consequence: str = "rewrites existing rows at boot" def render(self) -> str: location = f"{MIGRATIONS_DIR.relative_to(REPO_ROOT)}/{self.migration}/migration.sql" - return f"{location}:{self.line}: {self.keyword} rewrites existing rows at boot" + return f"{location}:{self.line}: {self.keyword} {self.consequence}" @dataclass(frozen=True, slots=True) @@ -578,6 +599,35 @@ def rewrites_a_log_table(clause: str, region: str, base: int) -> str | None: return f"ADD COLUMN ... DEFAULT on {named.group(1)}" +def builds_a_log_index(clause: str, region: str, base: int) -> str | None: + """The keyword to report when a `CREATE INDEX` targets a request-log table, concurrent or + not: a plain build blocks writes for its whole duration, and a concurrent one fails with + P3018 on a partitioned parent, so the migration job builds those instead.""" + created: Final[re.Match[str] | None] = re.match( + r"\s*CREATE\s+(?:UNIQUE\s+)?INDEX\b(?:\s+CONCURRENTLY\b)?", clause, re.IGNORECASE + ) + if created is None: + return None + on: Final[re.Match[str] | None] = re.search(r"\bON\b(?:\s+ONLY\b)?", clause[created.end() :], re.IGNORECASE) + if on is None: + return None + named: Final[re.Match[str] | None] = TABLE_NAME.match( + region, skip_comments(region, base + created.end() + on.end()) + ) + if named is None or named.group(1).strip('"') not in REQUEST_LOG_TABLES: + return None + return f"CREATE INDEX on {named.group(1)}" + + +def consequence_of(found: str) -> str: + if found.startswith("CREATE INDEX"): + return ( + "blocks writes until the build finishes, or fails on a partitioned table; " + "add it to REQUEST_LOG_INDEXES in litellm_proxy_extras/request_log_indexes.py instead" + ) + return "rewrites existing rows at boot" + + def skip_comments(sql: str, start: int) -> int: index = start while index < len(sql): @@ -746,8 +796,7 @@ def read_markers(sql: str) -> Markers: return Markers( sql, tuple( - Marker(match.start(), match.end(), alone_on_its_line(sql, match.start())) - for match in MARKER.finditer(sql) + Marker(match.start(), match.end(), alone_on_its_line(sql, match.start())) for match in MARKER.finditer(sql) ), ) @@ -760,9 +809,7 @@ def scan(sql: str, migration: str, markers: Markers) -> Iterator[Violation]: yield from scan_region(sql, sql, migration, markers, 0) -def scan_region( - document: str, region: str, migration: str, markers: Markers, offset: int -) -> Iterator[Violation]: +def scan_region(document: str, region: str, migration: str, markers: Markers, offset: int) -> Iterator[Violation]: """Violations in one region of `document`, whose text begins at `offset`. Positions are always counted against the whole document, so a statement nested in a dollar-quoted body reports its real file line and lines up with the markers read from that file. A single-quoted @@ -790,13 +837,18 @@ def scan_region( offset + start, ) - keyword = offending_keyword(clause) - if exempt: - continue - found = keyword or rewrites_a_log_table(clause, region, base) + index = builds_a_log_index(clause, region, base) + found = ( + index if exempt else offending_keyword(clause) or rewrites_a_log_table(clause, region, base) or index + ) if found is None: continue - yield Violation(migration, line_of(document, offset + keyword_start(clause, base)), found) + yield Violation( + migration, + line_of(document, offset + keyword_start(clause, base)), + found, + consequence_of(found), + ) for body in bodies: if not runs_when_applied(masked, region, bodies, runnable, identifiers, body): diff --git a/tests/code_coverage_tests/check_provider_folders_documented.py b/tests/code_coverage_tests/check_provider_folders_documented.py index 60afc55331f..08fbde3d979 100644 --- a/tests/code_coverage_tests/check_provider_folders_documented.py +++ b/tests/code_coverage_tests/check_provider_folders_documented.py @@ -28,6 +28,11 @@ EXCLUDED_FOLDERS = { "pass_through", "openai_like", # This is a generic handler, not a specific provider "aiohttp_openai", # Internal implementation detail for async HTTP + # Agent-harness configs for litellm.agent(), not LLM providers; documented under docs/harness + "claude_code", + "codex", + "opencode", + "deepagents", } diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index 659dc438f2d..863934d76f7 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -59,6 +59,8 @@ IGNORE_FUNCTIONS = [ "_iter_fallback_targets", # max depth set (2 * ROUTER_MAX_FALLBACKS); fails closed by raising ValueError at the cap. "_mergeable_branch", # max depth set (_MAX_SCHEMA_FLATTEN_DEPTH=32) plus a seen_refs cycle guard; passes the schema through untouched at the cap. "json_string_leaves", # max depth set (MAX_STRUCTURED_CONTENT_SCAN_DEPTH); fails closed by raising at the cap so nothing goes unscanned. + "strict_json_schema", # harness: max depth set (DEFAULT_MAX_RECURSE_DEPTH); fails closed by raising ValueError at the cap. + "toml_value", # harness/codex: max depth set (DEFAULT_MAX_RECURSE_DEPTH); fails closed by raising OptionsMismatch at the cap. "with_json_string_leaves", # transitively bounded: only runs on a tree json_string_leaves already walked under the cap. "json_unrewritable_labels", # max depth set (MAX_STRUCTURED_CONTENT_SCAN_DEPTH); returns the None sentinel at the cap so the caller blocks. "_flatten_form_field", # bounded by the nesting depth of the already-parsed request body (a finite JSON tree, no cycles possible). diff --git a/tests/code_coverage_tests/unbounded_in_baseline.txt b/tests/code_coverage_tests/unbounded_in_baseline.txt index 17e8d32dde0..c42a6b0ddf5 100644 --- a/tests/code_coverage_tests/unbounded_in_baseline.txt +++ b/tests/code_coverage_tests/unbounded_in_baseline.txt @@ -11,7 +11,6 @@ litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py backfill_null_oau litellm/proxy/_experimental/mcp_server/oauth2_flow_backfill.py backfill_null_oauth2_flows prisma server_id.in `server_ids` 0 litellm/proxy/_experimental/mcp_server/toolset_db.py list_mcp_toolsets prisma toolset_id.in `toolset_ids` 0 litellm/proxy/agent_endpoints/endpoints.py _attach_keys_to_agents prisma agent_id.in `agent_ids` 0 -litellm/proxy/agent_endpoints/endpoints.py get_agent_daily_activity prisma agent_id.in `list(agent_ids_list)` 0 litellm/proxy/agent_endpoints/endpoints.py get_agents prisma agent_id.in `agent_ids` 0 litellm/proxy/anthropic_endpoints/claude_code_endpoints/claude_code_skill_access.py SkillVisibility.where prisma name.in `sorted(self.granted)` 0 litellm/proxy/auth/auth_checks.py _fetch_uncached_model_access_group_budgets prisma access_group_name.in `list(uncached_groups)` 0 @@ -39,19 +38,10 @@ litellm/proxy/management_endpoints/auto_router_endpoints.py start_shadow_eval pr litellm/proxy/management_endpoints/auto_router_endpoints.py start_shadow_eval prisma token.in `list(data.api_key_ids)` 0 litellm/proxy/management_endpoints/auto_router_endpoints.py start_shadow_eval prisma user_id.in `list(data.user_ids)` 0 litellm/proxy/management_endpoints/budget_management_endpoints.py info_budget prisma budget_id.in `data.budgets` 0 -litellm/proxy/management_endpoints/common_daily_activity.py _build_aggregated_where_clause raw-sql api_key.IN `IN ({placeholders})` 0 -litellm/proxy/management_endpoints/common_daily_activity.py _build_aggregated_where_clause raw-sql {entity_id_field}.IN `IN ({placeholders})` 0 -litellm/proxy/management_endpoints/common_daily_activity.py _build_aggregated_where_clause raw-sql {entity_id_field}.IN `IN ({placeholders})` 1 -litellm/proxy/management_endpoints/common_daily_activity.py _build_where_conditions prisma [entity_id_field].in `entity_id` 0 -litellm/proxy/management_endpoints/common_daily_activity.py _build_where_conditions prisma api_key.in `api_key` 0 -litellm/proxy/management_endpoints/common_daily_activity.py _build_where_conditions prisma not.in `exclude_entity_ids` 0 -litellm/proxy/management_endpoints/common_daily_activity.py get_api_key_metadata prisma token.in `list(api_keys)` 0 -litellm/proxy/management_endpoints/common_daily_activity.py get_api_key_metadata prisma token.in `list(missing_keys)` 0 litellm/proxy/management_endpoints/common_utils.py _team_admin_can_invite_user prisma team_id.in `admin_user_obj.teams` 0 litellm/proxy/management_endpoints/common_utils.py _user_has_admin_privileges prisma team_id.in `user_obj.teams` 0 litellm/proxy/management_endpoints/customer_endpoints.py delete_end_user prisma user_id.in `data.user_ids` 0 litellm/proxy/management_endpoints/customer_endpoints.py delete_end_user prisma user_id.in `data.user_ids` 1 -litellm/proxy/management_endpoints/customer_endpoints.py get_customer_daily_activity prisma user_id.in `list(end_user_ids_list)` 0 litellm/proxy/management_endpoints/internal_user_endpoints.py _check_user_info_v2_access prisma team_id.in `caller_user.teams` 0 litellm/proxy/management_endpoints/internal_user_endpoints.py _resolve_user_email_metadata prisma user_id.in `list(user_ids)` 0 litellm/proxy/management_endpoints/internal_user_endpoints.py delete_user prisma created_by.in `data.user_ids` 0 @@ -81,7 +71,6 @@ litellm/proxy/management_endpoints/mcp_management_endpoints.py fetch_all_mcp_ser litellm/proxy/management_endpoints/model_access_group_management_endpoints.py update_deployments_with_access_group prisma model_name.in `model_names` 0 litellm/proxy/management_endpoints/model_management_endpoints.py delete_team_models prisma model_id.in `model_ids` 0 litellm/proxy/management_endpoints/organization_endpoints.py deprecated_info_organization prisma organization_id.in `data.organizations` 0 -litellm/proxy/management_endpoints/organization_endpoints.py get_organization_daily_activity prisma organization_id.in `list(org_ids_list)` 0 litellm/proxy/management_endpoints/organization_endpoints.py list_organization prisma organization_id.in `membership_org_ids` 0 litellm/proxy/management_endpoints/router_weights.py validate_router_settings_weights prisma model_id.in `list(deployment_ids)` 0 litellm/proxy/management_endpoints/session_endpoints.py revoke_ui_session_keys prisma token.in `revoked_tokens` 0 @@ -99,7 +88,7 @@ litellm/proxy/management_endpoints/team_endpoints.py _build_team_list_where_cond litellm/proxy/management_endpoints/team_endpoints.py _get_keys_count_by_team prisma team_id.in `page_team_ids` 0 litellm/proxy/management_endpoints/team_endpoints.py _hydrate_member_user_details prisma user_id.in `sorted(user_ids)` 0 litellm/proxy/management_endpoints/team_endpoints.py _resolve_existing_member_user_ids prisma user_id.in `sorted(requested_user_ids)` 0 -litellm/proxy/management_endpoints/team_endpoints.py _resolve_team_daily_activity_scope prisma team_id.in `list(team_ids_list)` 0 +litellm/proxy/management_endpoints/team_endpoints.py resolve_team_daily_activity_scope prisma team_id.in `list(team_ids_list)` 0 litellm/proxy/management_endpoints/team_endpoints.py _sweep_deleted_team_references prisma team_id.in `tuple(team_ids)` 0 litellm/proxy/management_endpoints/team_endpoints.py _sweep_deleted_team_references_tx prisma team_id.in `tuple(team_ids)` 0 litellm/proxy/management_endpoints/team_endpoints.py _team_member_delete prisma user_id.in `sorted(addressed_user_ids)` 0 @@ -149,4 +138,7 @@ litellm/proxy/utils.py PrismaClient.get_data prisma budget_id.in `budget_id_list litellm/proxy/utils.py PrismaClient.get_data prisma team_id.in `team_id_list` 0 litellm/proxy/utils.py PrismaClient.get_data prisma user_id.in `user_id_list` 0 litellm/proxy/utils.py prefetch_config_params prisma param_name.in `param_names` 0 +litellm/repositories/daily_activity_repository.py DailyActivityRepository.daily_rows prisma [scope.entity_id_field].in `list(scope.entity_ids)` 0 +litellm/repositories/daily_activity_repository.py DailyActivityRepository.daily_rows prisma api_key.in `list(scope.api_keys)` 0 +litellm/repositories/daily_activity_repository.py DailyActivityRepository.daily_rows prisma not.in `list(scope.exclude_entity_ids)` 0 litellm/router_utils/auto_router_model_naming.py raw-sql classifier_type.IN `IN ({_LLM_CLASSIFIER_TYPES_SQL})` 0 diff --git a/tests/e2e/router/reliability_support.py b/tests/e2e/router/reliability_support.py index 5984d5645d8..f600531e663 100644 --- a/tests/e2e/router/reliability_support.py +++ b/tests/e2e/router/reliability_support.py @@ -97,7 +97,9 @@ def create_never_benched_refusing_deployment(proxy: ProxyClient, name: str) -> s def create_timeout_deployment(proxy: ProxyClient, name: str) -> str: """Register a deployment with a 1ms deadline the real backend always exceeds.""" - return proxy.create_model(name, LiteLLMParamsBody(model=REAL_MODEL, api_key=REAL_KEY, timeout=0.001)) + return proxy.create_model( + name, LiteLLMParamsBody(model=REAL_MODEL, api_key=REAL_KEY, timeout=0.001), provider_live=True + ) def create_small_context_deployment(proxy: ProxyClient, name: str) -> str: @@ -149,7 +151,7 @@ def create_caching_deployment(proxy: ProxyClient, name: str) -> str: def _register_benched_on_first_failure( - proxy: ProxyClient, name: str, litellm_params: LiteLLMParamsBody, allowed_fails: str + proxy: ProxyClient, name: str, litellm_params: LiteLLMParamsBody, allowed_fails: str, *, provider_live: bool = False ) -> str: """The always-picked half of a failing pair: all of the group's shuffle weight, and a cooldown policy that benches it on its first failure of the given class, @@ -159,7 +161,8 @@ def _register_benched_on_first_failure( model_name=name, litellm_params=litellm_params, model_info=ModelInfoBody(allowed_fails_policy={allowed_fails: 0}), - ) + ), + provider_live=provider_live, ) @@ -170,6 +173,7 @@ def create_always_timing_out_deployment(proxy: ProxyClient, name: str, cooldown_ name, LiteLLMParamsBody(model=REAL_MODEL, api_key=REAL_KEY, timeout=0.001, weight=1, cooldown_time=cooldown_time), "TimeoutErrorAllowedFails", + provider_live=True, ) diff --git a/tests/harness_e2e/__init__.py b/tests/harness_e2e/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/harness_e2e/conftest.py b/tests/harness_e2e/conftest.py new file mode 100644 index 00000000000..ba679d6d5bf --- /dev/null +++ b/tests/harness_e2e/conftest.py @@ -0,0 +1,74 @@ +"""Fixtures for the litellm.agent() end-to-end tests. + +These run the real harness runtimes (claude, codex, opencode, deepagents) against a real +LiteLLM AI Gateway, routed with the `litellm_proxy/` model prefix. They skip unless +LITELLM_PROXY_API_BASE and LITELLM_PROXY_API_KEY are set. Model groups can be overridden +per harness with HARNESS_E2E_MODEL_. +""" + +import importlib.util +import os +import shutil +from collections.abc import Iterator +from pathlib import Path + +import pytest + +from litellm import Harness + +GATEWAY_BASE = os.environ.get("LITELLM_PROXY_API_BASE", "").strip() +GATEWAY_KEY = os.environ.get("LITELLM_PROXY_API_KEY", "").strip() + +DEFAULT_MODEL_GROUPS = { + Harness.CLAUDE_CODE: "claude-haiku-4-5-20251001", + Harness.CODEX: "bedrock_mantle/openai.gpt-5.4", + Harness.OPENCODE: "claude-haiku-4-5-20251001", + Harness.DEEPAGENTS: "claude-haiku-4-5-20251001", +} + +BINARIES = { + Harness.CLAUDE_CODE: "claude", + Harness.CODEX: "codex", + Harness.OPENCODE: "opencode", +} + +requires_gateway = pytest.mark.skipif( + not (GATEWAY_BASE and GATEWAY_KEY), + reason="LITELLM_PROXY_API_BASE / LITELLM_PROXY_API_KEY not set", +) + + +def model_for(harness: Harness) -> str: + """`litellm_proxy/`: every model call goes through the gateway.""" + override = os.environ.get(f"HARNESS_E2E_MODEL_{harness.name}", "").strip() + return f"litellm_proxy/{override or DEFAULT_MODEL_GROUPS[harness]}" + + +def harness_available(harness: Harness) -> bool: + if harness is Harness.DEEPAGENTS: + return all( + importlib.util.find_spec(m) is not None + for m in ("deepagents", "langchain_litellm") + ) + return shutil.which(BINARIES[harness]) is not None + + +def harness_params() -> list: + return [ + pytest.param( + h, + id=h.value, + marks=pytest.mark.skipif( + not harness_available(h), reason=f"{h.value} runtime not installed" + ), + ) + for h in Harness + ] + + +@pytest.fixture +def workspace(tmp_path: Path) -> Iterator[Path]: + repo = tmp_path / "repo" + repo.mkdir() + (repo / "README.md").write_text("# demo\n") + yield repo diff --git a/tests/harness_e2e/test_harness_e2e.py b/tests/harness_e2e/test_harness_e2e.py new file mode 100644 index 00000000000..8892085f4bc --- /dev/null +++ b/tests/harness_e2e/test_harness_e2e.py @@ -0,0 +1,150 @@ +"""End-to-end: every harness, real runtime, real LiteLLM AI Gateway via litellm_proxy/.""" + +from pathlib import Path + +import pytest +from pydantic import BaseModel + +import litellm +from litellm import Harness, sandbox +from litellm.harness import ( + CapabilityUnsupported, + Done, + FileChange, + State, + Text, + ToolCall, +) + +from .conftest import harness_params, model_for, requires_gateway + +pytestmark = [requires_gateway] + +TURN_TIMEOUT = 300 + + +class Answer(BaseModel): + city: str + country: str + + +@pytest.mark.parametrize("harness", harness_params()) +def test_agent_creates_file_and_reports_cost(harness: Harness, workspace: Path) -> None: + result = litellm.agent( + harness, + "Create a file named hello.txt whose entire content is the single word: hi", + sandbox=sandbox.local(workspace), + model=model_for(harness), + timeout=TURN_TIMEOUT, + ) + + assert result.stop_reason == "done", result.text + assert (workspace / "hello.txt").read_text().strip().lower() == "hi" + assert [f.path for f in result.files if f.kind == "created"] == ["hello.txt"] + assert result.usage.calls >= 1 + assert result.usage.input_tokens > 0 + assert result.cost >= 0 + + +@pytest.mark.parametrize("harness", harness_params()) +def test_agent_stream_event_order(harness: Harness, workspace: Path) -> None: + (workspace / "secret.txt").write_text("The secret word is ZEBRA.\n") + events = list( + litellm.agent( + harness, + "Read secret.txt and reply with just the secret word in it.", + sandbox=sandbox.local(workspace), + model=model_for(harness), + timeout=TURN_TIMEOUT, + stream=True, + ) + ) + + assert isinstance(events[-1], Done) + assert sum(isinstance(e, Done) for e in events) == 1 + assert any(isinstance(e, Text) for e in events) + assert any(isinstance(e, ToolCall) for e in events) + assert "zebra" in events[-1].result.text.lower() + + +@pytest.mark.parametrize("harness", harness_params()) +def test_agent_structured_output(harness: Harness, workspace: Path) -> None: + result = litellm.agent( + harness, + "What is the capital of France? Do not use any tools.", + sandbox=sandbox.local(workspace), + model=model_for(harness), + output=Answer, + permissions="read-only", + timeout=TURN_TIMEOUT, + ) + + assert isinstance(result.output, Answer) + assert result.output.city.lower() == "paris" + + +@pytest.mark.parametrize("harness", harness_params()) +def test_agent_session_remembers_previous_turn( + harness: Harness, workspace: Path +) -> None: + with litellm.agent_session( + harness, + sandbox=sandbox.local(workspace), + model=model_for(harness), + timeout=TURN_TIMEOUT, + ) as s: + s.run("Remember this code word: PELICAN. Reply with just OK.") + second = s.run( + "What code word did I ask you to remember? Reply with just the word." + ) + assert "pelican" in second.text.lower() + assert s.cost >= second.cost + + +@pytest.mark.parametrize("harness", harness_params()) +def test_agent_detach_and_resume(harness: Harness, workspace: Path) -> None: + box = sandbox.local(workspace) + s = litellm.agent_session( + harness, sandbox=box, model=model_for(harness), timeout=TURN_TIMEOUT + ) + s.run("Remember this number: 4817. Reply with just OK.") + raw = s.detach().dumps() + + with litellm.agent_resume( + State.loads(raw), sandbox=box, model=model_for(harness) + ) as resumed: + r = resumed.run( + "What number did I ask you to remember? Reply with just the number." + ) + assert "4817" in r.text + + +@pytest.mark.parametrize("harness", harness_params()) +def test_agent_read_only_blocks_writes(harness: Harness, workspace: Path) -> None: + result = litellm.agent( + harness, + "Create a file named blocked.txt containing x. If you cannot, just say you cannot.", + sandbox=sandbox.local(workspace), + model=model_for(harness), + permissions="read-only", + timeout=TURN_TIMEOUT, + ) + + assert not (workspace / "blocked.txt").exists() + assert not [f for f in result.files if isinstance(f, FileChange)] + + +def test_string_harness_rejected(workspace: Path) -> None: + with pytest.raises(TypeError, match=r"Harness\.CODEX"): + litellm.agent("codex", "hi", sandbox=sandbox.local(workspace)) # type: ignore[arg-type] + + +def test_capability_checked_before_start(workspace: Path) -> None: + with pytest.raises(CapabilityUnsupported): + litellm.agent( + Harness.CODEX, + "hi", + sandbox=sandbox.local(workspace), + model=model_for(Harness.CODEX), + disable_tools=["bash"], + ) diff --git a/tests/integration/_support/mcp_grants.py b/tests/integration/_support/mcp_grants.py index 5fa9eeaa0b6..82e79b2c3bd 100644 --- a/tests/integration/_support/mcp_grants.py +++ b/tests/integration/_support/mcp_grants.py @@ -51,12 +51,12 @@ def delete_toolset(gateway: Gateway, identity: str) -> None: assert response.status_code in (200, 202, 204), response.text -def create_toolset(scenario: Scenario, tools: tuple[tuple[str, str], ...]) -> str: +def create_toolset(scenario: Scenario, tools: tuple[tuple[str, str], ...], toolset_name: str | None = None) -> str: response: Final = scenario.gateway.request( "POST", "/v1/mcp/toolset", { - "toolset_name": f"integration-{uuid.uuid4().hex[:10]}", + "toolset_name": toolset_name or f"integration-{uuid.uuid4().hex[:10]}", "tools": [{"server_id": server_id, "tool_name": tool} for server_id, tool in tools], }, ) diff --git a/tests/integration/_support/process.py b/tests/integration/_support/process.py index b0941672be1..891874bdfa6 100644 --- a/tests/integration/_support/process.py +++ b/tests/integration/_support/process.py @@ -16,6 +16,10 @@ import httpx import psutil from integration._support.client import GATEWAY_LIMITS, Gateway +DB_PUSH: Final = ("--use_prisma_db_push",) +MIGRATE_DEPLOY: Final = () +LEGACY_MIGRATE_DEPLOY: Final = ("--use_legacy_migration_resolver",) + def proxy_database_environment() -> Mapping[str, str]: writer: Final = os.environ.get("INTEGRATION_PROXY_DATABASE_URL", "") @@ -73,9 +77,16 @@ def owned_proxy( config: Path | None = None, remove_environment: tuple[str, ...] = (), workers: int = 1, + database_setup: tuple[str, ...] = DB_PUSH, ) -> Iterator[Gateway]: with owned_proxy_process( - gateway, directory, overrides, config=config, remove_environment=remove_environment, workers=workers + gateway, + directory, + overrides, + config=config, + remove_environment=remove_environment, + workers=workers, + database_setup=database_setup, ) as owned: yield owned.gateway @@ -133,7 +144,7 @@ def _lost_port_race(launch: _Launch) -> bool: def _wait_until_ready(launch: _Launch) -> None: with httpx.Client(base_url=f"http://127.0.0.1:{launch.port}", timeout=15, trust_env=False) as client: - deadline: Final = time.monotonic() + 70 + deadline: Final = time.monotonic() + float(os.environ.get("INTEGRATION_PROXY_READY_SECONDS", "70")) while launch.process.poll() is None: try: if client.get("/health/readiness", timeout=2).status_code == 200: @@ -171,6 +182,7 @@ def owned_proxy_process( config: Path | None = None, remove_environment: tuple[str, ...] = (), workers: int = 1, + database_setup: tuple[str, ...] = DB_PUSH, ) -> Iterator[OwnedProxy]: root: Final = Path(os.environ.get("INTEGRATION_PROXY_ROOT") or Path(__file__).resolve().parents[3]) environment: Final = { @@ -196,8 +208,7 @@ def owned_proxy_process( "127.0.0.1", "--num_workers", str(workers), - "--use_prisma_db_push", - "--enforce_prisma_migration_check", + *database_setup, ) launch: Final = _launch_until_bound(command, root, environment, output, _PORT_ATTEMPTS) process: Final = launch.process diff --git a/tests/integration/configuration/test_lazy_routes_flag.py b/tests/integration/configuration/test_lazy_routes_flag.py new file mode 100644 index 00000000000..6135deca1c8 --- /dev/null +++ b/tests/integration/configuration/test_lazy_routes_flag.py @@ -0,0 +1,605 @@ +"""Route table contract for the LITELLM_DISABLE_LAZY_ROUTES startup flag. + +By default optional feature routers (``LAZY_FEATURES``) are registered on the first +request to their path prefix, so an operator inspecting the route table right after +boot cannot see or gate them. With the flag set every feature is registered at worker +startup, so ``GET /routes`` lists them before any feature request is served and the +first feature request changes nothing. +""" + +import asyncio +import json +import os +import re +import uuid +from collections.abc import Iterator, Mapping +from concurrent.futures import ThreadPoolExecutor +from contextlib import contextmanager +from dataclasses import dataclass +from functools import partial +from pathlib import Path +from typing import Final + +import anthropic +import httpx +import openai +import psutil +import pytest +import yaml +from pydantic import JsonValue, TypeAdapter + +from litellm.proxy._lazy_features import LAZY_FEATURES, LazyFeature +from tests.integration._support.client import Gateway, eventually, object_value, string_value +from tests.integration._support.mcp import McpPeer, call_tool, echo_tool, scripted_peer, tool_calls, tool_names +from tests.integration._support.process import OwnedProxy, owned_proxy_process +from tests.integration._support.wire import Reply, Request, Wire, wire_server + +TICKET_FEATURES: Final = ("mcp_management", "mcp_byok_oauth") +FLAG: Final = "LITELLM_DISABLE_LAZY_ROUTES" +WARMUP_ROUTE: Final = "/lazy/warm/{name}" +MCP_WARM_PATH: Final = "/mcp/enabled" +MARKER: Final = re.compile(rb"lazyroutes-[0-9a-f]{32}") +FAILED_FEATURE: Final = re.compile(r"Failed to lazy-load optional feature '([a-z_]+)'") +JSON: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) +HOOK_MODULE: Final = "lazy_routes_route_filter_hook" +HOOK_SOURCE: Final = """from litellm.proxy.proxy_server import app + + +def drop_mcp_routes() -> None: + app.router.routes[:] = [ + route for route in app.router.routes if not getattr(route, "path", "").startswith(("/mcp", "/v1/mcp")) + ] +""" + + +def _paths(candidate: Gateway) -> tuple[str, ...]: + routes: Final = candidate.get("/routes")["routes"] + assert isinstance(routes, list), routes + return tuple(string_value(object_value(route)["path"]) for route in routes) + + +def _routed_features(candidate: Gateway) -> Mapping[str, tuple[str, ...]]: + paths: Final = _paths(candidate) + return {feature.name: tuple(path for path in paths if feature.matches(path)) for feature in LAZY_FEATURES} + + +def _mcp_paths(candidate: Gateway) -> tuple[str, ...]: + return tuple(path for path in _paths(candidate) if path.startswith(("/mcp", "/v1/mcp"))) + + +def _route_filter_hook(directory: Path) -> Mapping[str, str]: + (directory / f"{HOOK_MODULE}.py").write_text(HOOK_SOURCE) + search_path: Final = (str(directory), os.environ.get("PYTHONPATH", "")) + return { + "PYTHONPATH": os.pathsep.join(entry for entry in search_path if entry), + "LITELLM_WORKER_STARTUP_HOOKS": f"{HOOK_MODULE}:drop_mcp_routes", + } + + +def test_lazy_routes_are_absent_from_the_route_table_until_first_request(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy_process(gateway, tmp_path, {}, remove_environment=(FLAG,)) as owned: + at_boot: Final = _routed_features(owned.gateway) + assert {name: at_boot[name] for name in TICKET_FEATURES} == {name: () for name in TICKET_FEATURES}, at_boot + listing: Final = owned.gateway.request("GET", "/v1/mcp/server") + assert listing.status_code == 200, listing.text + after_first_request: Final = _routed_features(owned.gateway) + assert after_first_request["mcp_management"] != (), "first request did not register the router" + assert after_first_request["mcp_byok_oauth"] == (), "only the requested feature is mounted" + + +@pytest.mark.parametrize("workers", (1, 4)) +def test_disable_lazy_routes_flag_registers_every_feature_at_startup( + gateway: Gateway, tmp_path: Path, workers: int +) -> None: + with owned_proxy_process(gateway, tmp_path, {FLAG: "true"}, workers=workers) as owned: + at_boot: Final = tuple(_routed_features(owned.gateway) for _ in range(2 * workers)) + unregistered: Final = sorted(name for name, paths in at_boot[0].items() if not paths) + assert unregistered == [], f"features still missing from /routes at startup: {unregistered}" + assert all(table == at_boot[0] for table in at_boot), "workers disagree on the route table" + listing: Final = owned.gateway.request("GET", "/v1/mcp/server") + assert listing.status_code == 200, listing.text + assert _routed_features(owned.gateway) == at_boot[0], "first feature request changed the route table" + + +def test_startup_hook_cannot_remove_lazy_routes_that_register_after_it_ran(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy_process(gateway, tmp_path, _route_filter_hook(tmp_path), remove_environment=(FLAG,)) as owned: + assert _mcp_paths(owned.gateway) == (), "hook should have removed the routes registered before it ran" + listing: Final = owned.gateway.request("GET", "/v1/mcp/server") + assert listing.status_code == 200, listing.text + assert _mcp_paths(owned.gateway) != (), "first request should have registered the routes the hook never saw" + + +def test_disable_lazy_routes_flag_lets_a_startup_hook_remove_optional_routes_for_good( + gateway: Gateway, tmp_path: Path +) -> None: + overrides: Final = {**_route_filter_hook(tmp_path), FLAG: "true"} + with owned_proxy_process(gateway, tmp_path, overrides) as owned: + assert _mcp_paths(owned.gateway) == (), "hook should have seen and removed every MCP route" + listing: Final = owned.gateway.request("GET", "/v1/mcp/server") + assert listing.status_code == 404, listing.text + mounted: Final = owned.gateway.request("POST", "/mcp", {"jsonrpc": "2.0", "id": 1, "method": "tools/list"}) + assert mounted.status_code == 404, mounted.text + guardrails: Final = owned.gateway.request("GET", "/guardrails/list") + assert guardrails.status_code == 200, guardrails.text + assert _mcp_paths(owned.gateway) == (), "a feature request re-registered routes the hook removed" + + +def _openapi_paths(candidate: Gateway) -> Mapping[str, tuple[str, ...]]: + paths: Final = object_value(candidate.get("/openapi.json")["paths"]) + return {path: tuple(sorted(object_value(operations))) for path, operations in paths.items()} + + +def _published(feature: LazyFeature, paths: Mapping[str, tuple[str, ...]]) -> bool: + return any(feature.matches(path) for path in paths) + + +def _warm_every_feature(candidate: Gateway) -> None: + warmed: Final = tuple( + (feature.name, candidate.request("POST", f"/lazy/warm/{feature.name}")) for feature in LAZY_FEATURES + ) + cold: Final = [ + (name, response.status_code, response.text) for name, response in warmed if response.status_code != 200 + ] + assert cold == [], cold + enabled: Final = candidate.request("GET", MCP_WARM_PATH) + assert enabled.status_code == 200, enabled.text + unregistered: Final = sorted(name for name, paths in _routed_features(candidate).items() if not paths) + assert unregistered == [], f"features still missing after warming every one of them: {unregistered}" + + +def _shadowed_dependency(directory: Path) -> Mapping[str, str]: + package: Final = directory / "shadow" / "RestrictedPython" + package.mkdir(parents=True) + (package / "__init__.py").write_text('raise ImportError("shadowed by the lazy routes audit")\n') + search_path: Final = (str(package.parent), os.environ.get("PYTHONPATH", "")) + return {"PYTHONPATH": os.pathsep.join(entry for entry in search_path if entry)} + + +def _failed_features(owned: OwnedProxy) -> frozenset[str]: + return frozenset(FAILED_FEATURE.findall(owned.log.read_text())) + + +def _marker() -> str: + return "lazyroutes-" + uuid.uuid4().hex + + +def _chat_reply(identity: str, stream: bool) -> Reply: + if not stream: + return Reply( + body=json.dumps( + { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "lazy ok"}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 7, "completion_tokens": 2, "total_tokens": 9}, + } + ).encode() + ) + chunk: Final[dict[str, JsonValue]] = { + "id": identity, + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o-mini", + } + deltas: Final[tuple[dict[str, JsonValue], ...]] = ( + {**chunk, "choices": [{"index": 0, "delta": {"role": "assistant", "content": "lazy"}}]}, + {**chunk, "choices": [{"index": 0, "delta": {"content": " ok"}, "finish_reason": "stop"}]}, + {**chunk, "choices": [], "usage": {"prompt_tokens": 7, "completion_tokens": 2, "total_tokens": 9}}, + ) + return Reply( + content_type="text/event-stream", + chunks=(*(b"data: " + json.dumps(delta).encode() + b"\n\n" for delta in deltas), b"data: [DONE]\n\n"), + ) + + +def _responses_reply(identity: str, stream: bool) -> Reply: + response: Final[dict[str, JsonValue]] = { + "id": identity, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "id": "msg_" + identity, + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "lazy ok", "annotations": []}], + } + ], + "usage": {"input_tokens": 7, "output_tokens": 2, "total_tokens": 9}, + } + if not stream: + return Reply(body=json.dumps(response).encode()) + events: Final[tuple[dict[str, JsonValue], ...]] = ( + {"type": "response.created", "sequence_number": 0, "response": {**response, "status": "in_progress"}}, + { + "type": "response.output_text.delta", + "sequence_number": 1, + "item_id": "msg_" + identity, + "output_index": 0, + "content_index": 0, + "delta": "lazy ok", + }, + {"type": "response.completed", "sequence_number": 2, "response": response}, + ) + return Reply( + content_type="text/event-stream", + chunks=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events), + ) + + +def _upstream(request: Request) -> Reply: + found: Final = MARKER.search(request.body) + if found is None: + return Reply(status=404, body=b'{"error":"no marker"}') + marker: Final = found.group(0).decode() + stream: Final = object_value(JSON.validate_json(request.body)).get("stream") is True + if request.target.endswith("/responses"): + return _responses_reply(f"resp_{marker}", stream) + return _chat_reply(f"chatcmpl-{marker}", stream) + + +@pytest.fixture(scope="module") +def provider() -> Iterator[Wire]: + with wire_server(_upstream) as wire: + yield wire + + +async def _stream_chat(base_url: str, key: str, model: str, marker: str) -> tuple[frozenset[str], str]: + client: Final = openai.AsyncOpenAI(base_url=base_url + "/v1", api_key=key, max_retries=0) + stream: Final = await client.chat.completions.create( + model=model, messages=[{"role": "user", "content": marker}], stream=True + ) + chunks: Final = [chunk async for chunk in stream] + text: Final = "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) + return frozenset(chunk.id for chunk in chunks), text + + +async def _stream_message(base_url: str, key: str, model: str, marker: str) -> str: + client: Final = anthropic.AsyncAnthropic(base_url=base_url, api_key=key, max_retries=0) + async with client.messages.stream( + model=model, max_tokens=16, messages=[{"role": "user", "content": marker}] + ) as stream: + return "".join([text async for text in stream.text_stream]) + + +def _status(candidate: Gateway, path: str) -> int: + return candidate.request("GET", path).status_code + + +def _workers(owned: OwnedProxy) -> tuple[psutil.Process, ...]: + return tuple(child for child in psutil.Process(owned.process.pid).children() if _is_worker(child)) + + +def _is_worker(child: psutil.Process) -> bool: + try: + return "spawn_main" in " ".join(child.cmdline()) and child.status() != psutil.STATUS_ZOMBIE + except psutil.Error: + return False + + +@pytest.mark.parametrize("spelling", ("1", "Yes", "ON")) +def test_every_truthy_spelling_of_the_flag_registers_the_ticket_features_at_startup( + gateway: Gateway, tmp_path: Path, spelling: str +) -> None: + with owned_proxy_process(gateway, tmp_path, {FLAG: spelling}) as owned: + at_boot: Final = _routed_features(owned.gateway) + assert all(at_boot[name] for name in TICKET_FEATURES), {name: at_boot[name] for name in TICKET_FEATURES} + + +@pytest.mark.parametrize( + "spelling", ("", "0", "off", "maybe", "x" * 5000), ids=("empty", "zero", "off", "unknown-word", "five-kilobytes") +) +def test_a_falsey_or_unknown_flag_value_keeps_the_default_lazy_registration( + gateway: Gateway, tmp_path: Path, spelling: str +) -> None: + with owned_proxy_process(gateway, tmp_path, {FLAG: spelling}) as owned: + at_boot: Final = _routed_features(owned.gateway) + assert {name: at_boot[name] for name in TICKET_FEATURES} == {name: () for name in TICKET_FEATURES}, at_boot + listing: Final = owned.gateway.request("GET", "/v1/mcp/server") + assert listing.status_code == 200, listing.text + assert _routed_features(owned.gateway)["mcp_management"] != (), "first request did not register the router" + + +def test_disable_lazy_routes_flag_publishes_the_live_route_table_in_openapi_before_any_feature_request( + gateway: Gateway, tmp_path: Path +) -> None: + with owned_proxy_process(gateway, tmp_path, {FLAG: "true"}) as owned: + at_boot: Final = _openapi_paths(owned.gateway) + unpublished: Final = sorted( + feature.name + for feature in LAZY_FEATURES + if feature.name in TICKET_FEATURES and not _published(feature, at_boot) + ) + assert unpublished == [], f"ticket features missing from /openapi.json at startup: {unpublished}" + assert "get" in at_boot["/v1/mcp/server"], at_boot["/v1/mcp/server"] + assert WARMUP_ROUTE not in at_boot + stranger: Final = owned.gateway.request("GET", "/v1/mcp/server", key="sk-not-a-real-key") + assert stranger.status_code == 401, stranger.text + listing: Final = owned.gateway.request("GET", "/v1/mcp/server") + assert listing.status_code == 200, listing.text + assert _openapi_paths(owned.gateway) == at_boot, "first feature request changed /openapi.json" + + +def test_disable_lazy_routes_flag_removes_the_warmup_route(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy_process(gateway, tmp_path, {FLAG: "true"}) as owned: + assert WARMUP_ROUTE not in _paths(owned.gateway) + warmed: Final = owned.gateway.request("POST", "/lazy/warm/mcp_management") + assert warmed.status_code == 404, warmed.text + listing: Final = owned.gateway.request("GET", "/v1/mcp/server") + assert listing.status_code == 200, listing.text + + +def test_the_warmup_route_registers_a_feature_on_demand_by_default(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy_process(gateway, tmp_path, {}, remove_environment=(FLAG,)) as owned: + assert WARMUP_ROUTE in _paths(owned.gateway) + warmed: Final = owned.gateway.request("POST", "/lazy/warm/mcp_management") + assert warmed.status_code == 200, warmed.text + assert "/v1/mcp/server" in object_value(object_value(JSON.validate_json(warmed.content))["paths"]) + assert _routed_features(owned.gateway)["mcp_management"] != (), "warmup did not register the router" + + +def test_disable_lazy_routes_flag_keeps_the_fixed_mcp_proxy_route_ahead_of_the_mcp_mount( + gateway: Gateway, tmp_path: Path +) -> None: + with owned_proxy_process(gateway, tmp_path, {}, remove_environment=(FLAG,)) as lazy: + control: Final = lazy.gateway.request("POST", "/mcp/proxy", {}) + assert control.status_code == 400, control.text + warmed: Final = _paths(lazy.gateway) + assert warmed.count("/mcp") == 2, "expected the fixed /mcp route and the /mcp mount" + with owned_proxy_process(gateway, tmp_path, {FLAG: "true"}) as eager: + at_boot: Final = _paths(eager.gateway) + assert at_boot.count("/mcp") == 2, "expected the fixed /mcp route and the /mcp mount at startup" + mount: Final = max(index for index, path in enumerate(at_boot) if path == "/mcp") + assert at_boot.index("/mcp/proxy") < mount, "the /mcp mount shadows /mcp/proxy" + proxied: Final = eager.gateway.request("POST", "/mcp/proxy", {}) + assert (proxied.status_code, proxied.text) == (control.status_code, control.text) + + +def test_disable_lazy_routes_flag_matches_the_fully_warmed_lazy_route_table_and_openapi( + gateway: Gateway, tmp_path: Path +) -> None: + with owned_proxy_process(gateway, tmp_path, {}, remove_environment=(FLAG,)) as lazy: + _warm_every_feature(lazy.gateway) + warmed_paths: Final = tuple(path for path in _paths(lazy.gateway) if path != WARMUP_ROUTE) + warmed_openapi: Final = _openapi_paths(lazy.gateway) + with owned_proxy_process(gateway, tmp_path, {FLAG: "true"}) as eager: + assert _paths(eager.gateway) == warmed_paths + assert _openapi_paths(eager.gateway) == warmed_openapi + + +def test_disable_lazy_routes_flag_keeps_registering_after_an_optional_dependency_fails_to_import( + gateway: Gateway, tmp_path: Path +) -> None: + with owned_proxy_process(gateway, tmp_path, {**_shadowed_dependency(tmp_path), FLAG: "true"}) as owned: + at_boot: Final = _routed_features(owned.gateway) + unregistered: Final = frozenset(name for name, paths in at_boot.items() if not paths) + failed: Final = _failed_features(owned) + assert "guardrails" in failed, owned.log.read_text() + assert unregistered == failed, (sorted(unregistered), sorted(failed)) + guardrails: Final = owned.gateway.request("GET", "/guardrails/list") + assert guardrails.status_code == 404, guardrails.text + listing: Final = owned.gateway.request("GET", "/v1/mcp/server") + assert listing.status_code == 200, listing.text + stores: Final = owned.gateway.request("GET", "/vector_store/list") + assert stores.status_code == 200, stores.text + assert _routed_features(owned.gateway) == at_boot, "feature requests changed the route table" + + +def test_a_broken_optional_dependency_only_404s_its_own_feature_by_default(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy_process(gateway, tmp_path, _shadowed_dependency(tmp_path), remove_environment=(FLAG,)) as owned: + guardrails: Final = owned.gateway.request("GET", "/guardrails/list") + assert guardrails.status_code == 404, guardrails.text + assert "guardrails" in _failed_features(owned), owned.log.read_text() + listing: Final = owned.gateway.request("GET", "/v1/mcp/server") + assert listing.status_code == 200, listing.text + assert _routed_features(owned.gateway)["mcp_management"] != () + + +def test_disable_lazy_routes_flag_leaves_the_completion_endpoints_serving_every_client( + gateway: Gateway, tmp_path: Path, provider: Wire +) -> None: + with owned_proxy_process(gateway, tmp_path, {FLAG: "true"}) as owned, owned.gateway.scenario() as scenario: + model: Final = scenario.model(api_base=provider.url + "/v1") + base_url: Final = str(owned.gateway.client.base_url) + key: Final = owned.gateway.key + markers: Final = tuple(_marker() for _ in range(6)) + + completion: Final = openai.OpenAI( + base_url=base_url + "/v1", api_key=key, max_retries=0 + ).chat.completions.create(model=model, messages=[{"role": "user", "content": markers[0]}]) + assert (completion.id, completion.choices[0].message.content) == (f"chatcmpl-{markers[0]}", "lazy ok") + + assert asyncio.run(_stream_chat(base_url, key, model, markers[1])) == ( + frozenset({f"chatcmpl-{markers[1]}"}), + "lazy ok", + ) + + message: Final = anthropic.Anthropic(base_url=base_url, api_key=key, max_retries=0).messages.create( + model=model, max_tokens=16, messages=[{"role": "user", "content": markers[2]}] + ) + assert [block.text for block in message.content if block.type == "text"] == ["lazy ok"] + + assert asyncio.run(_stream_message(base_url, key, model, markers[3])) == "lazy ok" + + responded: Final = owned.gateway.request("POST", "/v1/responses", {"model": model, "input": markers[4]}) + assert responded.status_code == 200, responded.text + response: Final = object_value(JSON.validate_json(responded.content)) + assert (response["status"], response["object"]) == ("completed", "response"), responded.text + assert "lazy ok" in responded.text, responded.text + + streamed: Final = owned.gateway.request( + "POST", "/v1/responses", {"model": model, "input": markers[5], "stream": True} + ) + assert streamed.status_code == 200, streamed.text + assert "response.completed" in streamed.text and "lazy ok" in streamed.text, streamed.text + + reached: Final = tuple(request.target for request in provider.drain() if MARKER.search(request.body)) + assert len(reached) == 6, reached + + +def test_disable_lazy_routes_flag_route_table_survives_a_boot_burst_and_a_killed_worker( + gateway: Gateway, tmp_path: Path +) -> None: + probes: Final = ("/routes", "/v1/mcp/server", "/openapi.json", "/guardrails/list", "/vector_store/list") * 8 + with owned_proxy_process(gateway, tmp_path, {FLAG: "true"}, workers=2) as owned: + at_boot: Final = _routed_features(owned.gateway) + unregistered: Final = sorted(name for name, paths in at_boot.items() if not paths) + assert unregistered == [], f"features still missing from /routes at startup: {unregistered}" + with ThreadPoolExecutor(max_workers=8) as pool: + statuses: Final = tuple(pool.map(partial(_status, owned.gateway), probes)) + assert statuses == (200,) * len(probes), statuses + assert _routed_features(owned.gateway) == at_boot, "the boot burst changed the route table" + + victim: Final = eventually(lambda: _workers(owned), lambda workers: len(workers) == 2)[0] + victim.kill() + with httpx.Client(base_url=owned.gateway.client.base_url, timeout=15, trust_env=False) as fresh: + survivor: Final = Gateway(fresh, owned.gateway.key, owned.gateway.upstream_url) + during: Final = tuple(survivor.request("GET", "/v1/mcp/server").status_code for _ in range(10)) + assert during == (200,) * 10, during + respawned: Final = eventually( + lambda: frozenset(worker.pid for worker in _workers(owned)), + lambda pids: len(pids) == 2 and victim.pid not in pids, + seconds=30, + ) + assert f"Child process [{victim.pid}] died" in owned.log.read_text(), respawned + tables: Final = tuple(_routed_features(owned.gateway) for _ in range(4)) + assert all(table == at_boot for table in tables), "the respawned worker disagrees on the route table" + + +def test_disable_lazy_routes_flag_yields_the_same_route_table_after_a_restart(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy_process(gateway, tmp_path, {FLAG: "true"}) as first: + table: Final = _paths(first.gateway) + with owned_proxy_process(gateway, tmp_path, {FLAG: "true"}) as second: + assert _paths(second.gateway) == table + assert all(_routed_features(second.gateway).values()), "a feature is missing after restart" + + +SELF_HOSTED_LANGFUSE: Final = "/self-hosted-langfuse" + + +@dataclass(frozen=True, slots=True) +class _ConfiguredFeatures: + alias: str + config: Path + policy: Wire + langfuse: Wire + peer: McpPeer + + +def _allow(request: Request) -> Reply: + return Reply(body=json.dumps({"action": "NONE"}).encode()) + + +def _langfuse_health(request: Request) -> Reply: + return Reply(body=json.dumps({"status": "OK"}).encode()) + + +def _config_declaring(directory: Path, alias: str, policy: Wire, langfuse: Wire, peer: McpPeer) -> Path: + base: Final = object_value( + JSON.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + ) + config: Final = { + **base, + "guardrails": [ + { + "guardrail_name": alias, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": True, + "api_base": policy.url, + "api_key": "synthetic-guardrail-key", + }, + } + ], + "mcp_servers": {alias: peer.registration()}, + "general_settings": { + **object_value(base["general_settings"]), + "pass_through_endpoints": [ + { + "path": "/langfuse", + "target": langfuse.url + SELF_HOSTED_LANGFUSE, + "include_subpath": True, + "auth": True, + } + ], + }, + } + path: Final = directory / "configured-features.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +@contextmanager +def _configured_features(directory: Path) -> Iterator[_ConfiguredFeatures]: + alias: Final = "lazyroutes" + uuid.uuid4().hex[:8] + with ( + wire_server(_allow) as policy, + wire_server(_langfuse_health) as langfuse, + scripted_peer(echo_tool("add")) as peer, + ): + config: Final = _config_declaring(directory, alias, policy, langfuse, peer) + yield _ConfiguredFeatures(alias, config, policy, langfuse, peer) + + +def _config_server_id(candidate: Gateway, alias: str) -> str: + servers: Final = JSON.validate_json(candidate.request("GET", "/v1/mcp/server").content) + assert isinstance(servers, list), servers + return next( + string_value(object_value(server)["server_id"]) + for server in servers + if object_value(server)["server_name"] == alias + ) + + +def _assert_config_declared_features_serve(owned: OwnedProxy, features: _ConfiguredFeatures, provider: Wire) -> None: + marker: Final = _marker() + with owned.gateway.scenario() as scenario: + model: Final = scenario.model(api_base=provider.url + "/v1") + completion: Final = owned.gateway.request( + "POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": marker}]} + ) + assert completion.status_code == 200, completion.text + screened: Final = [request for request in features.policy.drain() if marker.encode() in request.body] + assert len(screened) == 1, "the config-declared guardrail did not screen the completion" + assert len([request for request in provider.drain() if marker.encode() in request.body]) == 1 + + identity: Final = _config_server_id(owned.gateway, features.alias) + key: Final = scenario.key(object_permission={"mcp_servers": [identity]}) + tool: Final = tool_names(owned.gateway, key, identity)["add"] + features.peer.drain() + called: Final = call_tool(owned.gateway, key, identity, tool, {"marker": marker}) + assert called.status_code == 200, called.text + reached_peer: Final = [ + object_value(object_value(call["body"])["params"]) for call in tool_calls(features.peer.drain()) + ] + assert [(params["name"], params["arguments"]) for params in reached_peer] == [("add", {"marker": marker})] + + forwarded: Final = owned.gateway.request("GET", "/langfuse/api/public/health") + assert forwarded.status_code == 200, forwarded.text + reached_langfuse: Final = tuple(request.target for request in features.langfuse.drain()) + assert reached_langfuse == (SELF_HOSTED_LANGFUSE + "/api/public/health",), ( + f"the config pass-through for /langfuse lost to the built-in Langfuse route: {reached_langfuse}" + ) + + +def test_disable_lazy_routes_flag_serves_config_declared_features_like_the_warmed_lazy_proxy( + gateway: Gateway, tmp_path: Path, provider: Wire +) -> None: + with _configured_features(tmp_path) as features: + with owned_proxy_process(gateway, tmp_path, {}, config=features.config, remove_environment=(FLAG,)) as lazy: + _assert_config_declared_features_serve(lazy, features, provider) + _warm_every_feature(lazy.gateway) + warmed_paths: Final = tuple(path for path in _paths(lazy.gateway) if path != WARMUP_ROUTE) + warmed_openapi: Final = _openapi_paths(lazy.gateway) + with owned_proxy_process(gateway, tmp_path, {FLAG: "true"}, config=features.config) as eager: + _assert_config_declared_features_serve(eager, features, provider) + assert _paths(eager.gateway) == warmed_paths + assert _openapi_paths(eager.gateway) == warmed_openapi diff --git a/tests/integration/database/test_request_log_indexes_at_boot.py b/tests/integration/database/test_request_log_indexes_at_boot.py new file mode 100644 index 00000000000..fe35ec2f1f9 --- /dev/null +++ b/tests/integration/database/test_request_log_indexes_at_boot.py @@ -0,0 +1,363 @@ +import os +import shutil +import subprocess +import sys +from collections.abc import Mapping +from dataclasses import dataclass +from itertools import product +from pathlib import Path +from types import MappingProxyType +from typing import Final + +import psycopg +import pytest +from integration._support.client import Gateway, eventually +from integration._support.database import scratch_database +from integration._support.process import LEGACY_MIGRATE_DEPLOY, MIGRATE_DEPLOY, owned_proxy_process +from psycopg import sql +from psycopg.rows import class_row + +REPO_ROOT: Final = Path(__file__).resolve().parents[3] +PRISMA_DIR: Final = REPO_ROOT / "litellm-proxy-extras" / "litellm_proxy_extras" +PARTITION_SCRIPT: Final = REPO_ROOT / "db_scripts" / "partition_spend_logs.sql" +SHIPPED_MIGRATIONS: Final = tuple(sorted(path.name for path in (PRISMA_DIR / "migrations").iterdir() if path.is_dir())) +API_KEY_INDEX_MIGRATION: Final = "20260823000000_add_spend_logs_api_key_starttime_index" +CALL_ID_INDEX_MIGRATION: Final = "20260831120001_spend_logs_litellm_call_id_index" +ORIGINAL_MIGRATION_SQL: Final = MappingProxyType( + { + API_KEY_INDEX_MIGRATION: ( + "-- CreateIndex\n" + 'CREATE INDEX IF NOT EXISTS "LiteLLM_SpendLogs_api_key_startTime_idx" ' + 'ON "LiteLLM_SpendLogs"("api_key", "startTime");\n' + ), + CALL_ID_INDEX_MIGRATION: ( + "-- CreateIndex\n" + 'CREATE INDEX CONCURRENTLY IF NOT EXISTS "LiteLLM_SpendLogs_litellm_call_id_idx" ' + 'ON "LiteLLM_SpendLogs"("litellm_call_id");\n' + ), + } +) +MIGRATION_JOB_SECONDS: Final = 300 +INDEXES_IN_PLACE: Final = "Request-log indexes are all in place" +INDEX_BUILD_LINES: Final = ("Building index", "Attached index") +BUILD_SECONDS: Final = 60 +SPEND_LOGS_INDEXES: Final = ("LiteLLM_SpendLogs_api_key_startTime_idx", "LiteLLM_SpendLogs_litellm_call_id_idx") +POPULATED_PARTITIONS: Final = MappingProxyType( + { + "LiteLLM_SpendLogs_p2026_08": ("2026-08-01", "2026-09-01"), + "LiteLLM_SpendLogs_p2026_09": ("2026-09-01", "2026-10-01"), + } +) +DEFAULT_PARTITION: Final = "LiteLLM_SpendLogs_pdefault" +ROWS_PER_PARTITION: Final = 500 +PARTITIONED_PARENT_ERROR: Final = 'cannot create index on partitioned table "LiteLLM_SpendLogs" concurrently' + + +@dataclass(frozen=True, slots=True) +class Resolver: + """One migration resolver as the serving proxy selects it (CLI flags) and as the + migration job selects it (environment).""" + + proxy_flags: tuple[str, ...] + job_environment: Mapping[str, str] + + +V2: Final = Resolver(MIGRATE_DEPLOY, MappingProxyType({"USE_V2_MIGRATION_RESOLVER": "true"})) +LEGACY: Final = Resolver(LEGACY_MIGRATE_DEPLOY, MappingProxyType({"USE_V2_MIGRATION_RESOLVER": "false"})) +RESOLVERS: Final = pytest.mark.parametrize("resolver", (V2, LEGACY), ids=("v2", "legacy")) + + +def release_layout(directory: Path, migrations: tuple[str, ...]) -> Path: + """The Prisma layout of the release that shipped `migrations`: the two index migrations + carry the SQL they shipped with, not the inert files of this build.""" + (directory / "migrations").mkdir(parents=True) + shutil.copy(PRISMA_DIR / "schema.prisma", directory / "schema.prisma") + shutil.copy(PRISMA_DIR / "migrations" / "migration_lock.toml", directory / "migrations" / "migration_lock.toml") + for name in migrations: + shutil.copytree(PRISMA_DIR / "migrations" / name, directory / "migrations" / name) + for name, original in ORIGINAL_MIGRATION_SQL.items(): + if name in migrations: + (directory / "migrations" / name / "migration.sql").write_text(original) + return directory / "schema.prisma" + + +def migrate_deploy(database_url: str, schema: Path) -> subprocess.CompletedProcess[str]: + return subprocess.run( + [sys.executable, "-I", "-m", "prisma", "migrate", "deploy", "--schema", str(schema)], + capture_output=True, + text=True, + timeout=300, + env={**os.environ, "DATABASE_URL": database_url}, + ) + + +def migration_job(database_url: str, resolver: Resolver) -> subprocess.CompletedProcess[str]: + """The migrations image entrypoint, as the helm migration Job runs it.""" + return subprocess.run( + [sys.executable, "-I", str(REPO_ROOT / "migrations" / "run.py")], + capture_output=True, + text=True, + timeout=MIGRATION_JOB_SECONDS, + cwd=REPO_ROOT, + env={**os.environ, "DATABASE_URL": database_url, **resolver.job_environment}, + ) + + +def migration_cli(database_url: str, gateway: Gateway, resolver: Resolver) -> subprocess.CompletedProcess[str]: + """The proxy CLI as a migration job: `--skip_server_startup` migrates, builds the + indexes and exits by their result.""" + return subprocess.run( + [ + sys.executable, + "-P", + "-m", + "integration._support.proxy", + "--config", + "tests/integration/proxy_config.yaml", + *resolver.proxy_flags, + "--skip_server_startup", + ], + capture_output=True, + text=True, + timeout=MIGRATION_JOB_SECONDS, + cwd=REPO_ROOT, + env={**os.environ, "DATABASE_URL": database_url, "LITELLM_MASTER_KEY": gateway.key}, + ) + + +def deploy_schema_before_the_index_migrations(database_url: str, directory: Path) -> None: + older: Final = tuple(name for name in SHIPPED_MIGRATIONS if name < API_KEY_INDEX_MIGRATION) + deployed: Final = migrate_deploy(database_url, release_layout(directory / "older-release", older)) + assert deployed.returncode == 0, deployed.stdout + deployed.stderr + + +def deploy_the_original_index_migrations(database_url: str, directory: Path) -> subprocess.CompletedProcess[str]: + """Boot the v1.103.0 layout once: on a partitioned table its CONCURRENTLY call_id index + fails with P3018 and leaves the ledger row unfinished, on a plain table both apply.""" + return migrate_deploy(database_url, release_layout(directory / "v1.103.0", SHIPPED_MIGRATIONS)) + + +def fail_the_call_id_index_migration_like_the_shipped_release(database_url: str, directory: Path) -> None: + deployed: Final = deploy_the_original_index_migrations(database_url, directory) + assert deployed.returncode != 0, deployed.stdout + assert "P3018" in deployed.stderr and PARTITIONED_PARENT_ERROR in deployed.stderr, deployed.stderr + assert ledger(database_url)[CALL_ID_INDEX_MIGRATION] is False + + +def partition_spend_logs(database_url: str) -> None: + with psycopg.connect(database_url, autocommit=True) as connection: + connection.execute(PARTITION_SCRIPT.read_bytes()) + for partition, (start, stop) in POPULATED_PARTITIONS.items(): + add_partition(connection, partition, start, stop) + connection.execute( + 'INSERT INTO "LiteLLM_SpendLogs" ("request_id", "call_type", "api_key", "startTime", "endTime") ' + "SELECT %s || n, 'acompletion', 'sk-' || (n %% 7), %s::timestamp + (n * interval '1 minute'), " + "%s::timestamp + (n * interval '1 minute') + interval '1 second' FROM generate_series(1, %s) AS n", + (partition, start, start, ROWS_PER_PARTITION), + ) + connection.execute( + 'INSERT INTO "LiteLLM_SpendLogs" ("request_id", "call_type", "startTime", "endTime") ' + "SELECT 'default-' || n, 'acompletion', '2026-07-01'::timestamp + (n * interval '1 minute'), " + "'2026-07-01'::timestamp + (n * interval '1 minute') FROM generate_series(1, %s) AS n", + (ROWS_PER_PARTITION,), + ) + + +@dataclass(frozen=True, slots=True) +class _LedgerRow: + name: str + finished: bool + + +@dataclass(frozen=True, slots=True) +class _IndexRow: + index: str + valid: bool + + +@dataclass(frozen=True, slots=True) +class _OidRow: + index: str + oid: int + + +@dataclass(frozen=True, slots=True) +class _AttachedRow: + partition: str + parent_index: str + valid: bool + + +def ledger(database_url: str) -> Mapping[str, bool]: + """Every migration in the ledger that was not rolled back, mapped to whether it finished.""" + with psycopg.connect(database_url) as connection, connection.cursor(row_factory=class_row(_LedgerRow)) as cursor: + rows: Final = cursor.execute( + 'SELECT migration_name AS name, finished_at IS NOT NULL AS finished FROM "_prisma_migrations" ' + "WHERE rolled_back_at IS NULL ORDER BY migration_name" + ).fetchall() + return MappingProxyType({row.name: row.finished for row in rows}) + + +def parent_index_validity(database_url: str) -> Mapping[str, bool]: + with psycopg.connect(database_url) as connection, connection.cursor(row_factory=class_row(_IndexRow)) as cursor: + rows: Final = cursor.execute( + "SELECT c.relname AS index, x.indisvalid AS valid FROM pg_index x JOIN pg_class c ON c.oid = x.indexrelid " + "WHERE x.indrelid = '\"LiteLLM_SpendLogs\"'::regclass AND c.relname = ANY(%s)", + (list(SPEND_LOGS_INDEXES),), + ).fetchall() + return MappingProxyType({row.index: row.valid for row in rows}) + + +def index_oids(database_url: str) -> Mapping[str, int]: + """index name -> oid for the SpendLogs indexes on the parent or any partition; a rebuild changes the oid.""" + with psycopg.connect(database_url) as connection, connection.cursor(row_factory=class_row(_OidRow)) as cursor: + rows: Final = cursor.execute( + "SELECT c.relname AS index, c.oid::int AS oid FROM pg_index x JOIN pg_class c ON c.oid = x.indexrelid " + "WHERE x.indrelid = '\"LiteLLM_SpendLogs\"'::regclass OR x.indrelid IN " + "(SELECT inhrelid FROM pg_inherits WHERE inhparent = '\"LiteLLM_SpendLogs\"'::regclass)" + ).fetchall() + return MappingProxyType({row.index: row.oid for row in rows}) + + +def attached_partition_indexes(database_url: str) -> frozenset[tuple[str, str, bool]]: + """(partition, parent index, child is valid) for every child index attached under a SpendLogs parent index.""" + with psycopg.connect(database_url) as connection, connection.cursor(row_factory=class_row(_AttachedRow)) as cursor: + rows: Final = cursor.execute( + "SELECT part.relname AS partition, parent_index.relname AS parent_index, child.indisvalid AS valid " + "FROM pg_inherits attached " + "JOIN pg_class parent_index ON parent_index.oid = attached.inhparent " + "JOIN pg_index child ON child.indexrelid = attached.inhrelid " + "JOIN pg_class part ON part.oid = child.indrelid " + "WHERE parent_index.relname = ANY(%s)", + (list(SPEND_LOGS_INDEXES),), + ).fetchall() + return frozenset((row.partition, row.parent_index, row.valid) for row in rows) + + +def expected_attachments(partitions: tuple[str, ...]) -> frozenset[tuple[str, str, bool]]: + return frozenset((partition, index, True) for partition, index in product(partitions, SPEND_LOGS_INDEXES)) + + +def add_partition(connection: psycopg.Connection[tuple[object, ...]], partition: str, start: str, stop: str) -> None: + connection.execute( + sql.SQL('CREATE TABLE {} PARTITION OF "LiteLLM_SpendLogs" FOR VALUES FROM ({}) TO ({})').format( + sql.Identifier(partition), sql.Literal(start), sql.Literal(stop) + ) + ) + + +def assert_ready(booted_gateway: Gateway) -> None: + readiness: Final = booted_gateway.request("GET", "/health/readiness") + assert readiness.status_code == 200, readiness.text + assert readiness.json()["db"] == "connected", readiness.text + + +def assert_both_indexes_cover_every_partition(database_url: str) -> None: + """Every populated partition's index is attached and valid, both parents are valid, and + a partition created afterwards inherits both indexes.""" + populated: Final = (*POPULATED_PARTITIONS, DEFAULT_PARTITION) + assert ledger(database_url) == {name: True for name in SHIPPED_MIGRATIONS} + assert attached_partition_indexes(database_url) == expected_attachments(populated) + assert parent_index_validity(database_url) == {index: True for index in SPEND_LOGS_INDEXES} + with psycopg.connect(database_url, autocommit=True) as connection: + add_partition(connection, "LiteLLM_SpendLogs_p2026_10", "2026-10-01", "2026-11-01") + assert attached_partition_indexes(database_url) == expected_attachments((*populated, "LiteLLM_SpendLogs_p2026_10")) + + +def assert_the_serving_proxy_boots_and_finds_the_indexes_in_place( + gateway: Gateway, directory: Path, database_url: str, resolver: Resolver +) -> None: + """The serving proxy applies the inert files, reports ready, and its background build + finds every index already there, so it builds nothing and the catalog is untouched.""" + oids: Final = index_oids(database_url) + with owned_proxy_process( + gateway, directory, {"DATABASE_URL": database_url}, database_setup=resolver.proxy_flags + ) as booted: + assert_ready(booted.gateway) + log: Final = eventually( + lambda: booted.log.read_text(errors="replace"), lambda text: INDEXES_IN_PLACE in text, seconds=BUILD_SECONDS + ) + assert ledger(database_url) == {name: True for name in SHIPPED_MIGRATIONS} + assert not any(line in log for line in INDEX_BUILD_LINES), log[-4000:] + assert index_oids(database_url) == oids + + +@RESOLVERS +def test_the_migration_job_gives_a_partitioned_table_at_the_pre_index_schema_both_indexes_per_partition( + gateway: Gateway, tmp_path: Path, resolver: Resolver +) -> None: + with scratch_database() as database_url: + deploy_schema_before_the_index_migrations(database_url, tmp_path) + partition_spend_logs(database_url) + job: Final = migration_job(database_url, resolver) + assert job.returncode == 0, job.stdout + job.stderr + assert INDEXES_IN_PLACE in job.stderr + job.stdout, job.stdout + job.stderr + assert_both_indexes_cover_every_partition(database_url) + assert_the_serving_proxy_boots_and_finds_the_indexes_in_place(gateway, tmp_path, database_url, resolver) + + +@RESOLVERS +def test_the_migration_job_heals_a_partitioned_table_left_with_the_failed_call_id_ledger_row( + gateway: Gateway, tmp_path: Path, resolver: Resolver +) -> None: + with scratch_database() as database_url: + deploy_schema_before_the_index_migrations(database_url, tmp_path) + partition_spend_logs(database_url) + fail_the_call_id_index_migration_like_the_shipped_release(database_url, tmp_path) + job: Final = migration_job(database_url, resolver) + assert job.returncode == 0, job.stdout + job.stderr + assert_both_indexes_cover_every_partition(database_url) + assert_the_serving_proxy_boots_and_finds_the_indexes_in_place(gateway, tmp_path, database_url, resolver) + + +@RESOLVERS +def test_the_migration_job_leaves_a_plain_table_that_applied_the_original_index_migrations_alone( + gateway: Gateway, tmp_path: Path, resolver: Resolver +) -> None: + with scratch_database() as database_url: + deploy_schema_before_the_index_migrations(database_url, tmp_path) + deployed: Final = deploy_the_original_index_migrations(database_url, tmp_path) + assert deployed.returncode == 0, deployed.stdout + deployed.stderr + before: Final = index_oids(database_url) + assert set(SPEND_LOGS_INDEXES) <= set(before), before + job: Final = migration_job(database_url, resolver) + assert job.returncode == 0, job.stdout + job.stderr + assert INDEXES_IN_PLACE in job.stderr + job.stdout, job.stdout + job.stderr + assert "Building index" not in job.stderr + job.stdout, job.stdout + job.stderr + assert ledger(database_url) == {name: True for name in SHIPPED_MIGRATIONS} + assert index_oids(database_url) == before + assert_the_serving_proxy_boots_and_finds_the_indexes_in_place(gateway, tmp_path, database_url, resolver) + + +@RESOLVERS +def test_a_serving_proxy_that_runs_the_migrations_itself_builds_both_indexes_after_it_is_ready( + gateway: Gateway, tmp_path: Path, resolver: Resolver +) -> None: + """A deployment that runs migrate deploy from the serving proxy and never runs the + migration job answers readiness with the inert files applied, then its background build + puts both indexes on every partition.""" + with scratch_database() as database_url: + deploy_schema_before_the_index_migrations(database_url, tmp_path) + partition_spend_logs(database_url) + assert "LiteLLM_SpendLogs_litellm_call_id_idx" not in parent_index_validity(database_url) + with owned_proxy_process( + gateway, tmp_path, {"DATABASE_URL": database_url}, database_setup=resolver.proxy_flags + ) as booted: + assert_ready(booted.gateway) + assert ledger(database_url) == {name: True for name in SHIPPED_MIGRATIONS} + log: Final = eventually( + lambda: booted.log.read_text(errors="replace"), + lambda text: INDEXES_IN_PLACE in text, + seconds=BUILD_SECONDS, + ) + assert "Building index" in log and "Attached index" in log, log[-4000:] + assert_both_indexes_cover_every_partition(database_url) + + +def test_the_cli_run_as_a_migration_job_builds_both_indexes_before_it_exits(gateway: Gateway, tmp_path: Path) -> None: + with scratch_database() as database_url: + deploy_schema_before_the_index_migrations(database_url, tmp_path) + partition_spend_logs(database_url) + job: Final = migration_cli(database_url, gateway, V2) + assert job.returncode == 0, job.stdout + job.stderr + assert_both_indexes_cover_every_partition(database_url) diff --git a/tests/integration/mcp/test_mcp_llm_endpoints.py b/tests/integration/mcp/test_mcp_llm_endpoints.py index 26f5ced4de7..01bf006a03e 100644 --- a/tests/integration/mcp/test_mcp_llm_endpoints.py +++ b/tests/integration/mcp/test_mcp_llm_endpoints.py @@ -9,6 +9,7 @@ import httpx import pytest from integration._support.client import Gateway, Scenario from integration._support.mcp import McpPeer, mcp_peer, register_mcp, tool_calls +from integration._support.mcp_grants import create_toolset from integration._support.wire import Reply, Request, Wire, wire_server Surface = Literal["chat", "responses", "messages", "messages_bridge"] @@ -194,9 +195,7 @@ class Rig: ) def upstream_tools(self) -> tuple[tuple[str, ...], ...]: - return tuple( - _tool_names(json.loads(request.body)) for request in self.wire.drain() if request.method == "POST" - ) + return tuple(_tool_names(json.loads(request.body)) for request in self.wire.drain() if request.method == "POST") def final_text(self, body: Mapping[str, object]) -> str: if self.surface == "chat": @@ -314,6 +313,25 @@ def test_allowed_tools_narrows_the_tool_list_handed_to_the_model(gateway: Gatewa assert [call["body"]["params"]["name"] for call in _peer_add_calls(rig.peer)] == ["add"] +@pytest.mark.parametrize("surface", ("chat", "responses", "messages")) +def test_toolset_gateway_url_serves_a_team_granted_toolset_to_a_key_without_its_own_grant( + gateway: Gateway, surface: Surface +) -> None: + with _rig(gateway, surface) as rig: + register_mcp(rig.scenario, rig.peer, "open" + uuid.uuid4().hex[:8], allow_all_keys=True) + toolset_name: Final = "ts" + uuid.uuid4().hex[:8] + toolset_id: Final = create_toolset(rig.scenario, ((rig.server_id, "add"),), toolset_name=toolset_name) + sibling_id: Final = create_toolset(rig.scenario, ((rig.server_id, "multiply"),)) + team_id: Final = rig.scenario.team(object_permission={"mcp_toolsets": [toolset_id, sibling_id]}) + key: Final = rig.scenario.key(team_id=team_id) + response: Final = rig.send(key, [{**AUTO, "server_url": f"litellm_proxy/mcp/{toolset_name}"}]) + assert response.status_code == 200, response.text + requests: Final = rig.upstream_tools() + assert requests, "model was never called" + assert all(names == (rig.tool,) for names in requests), requests + assert [call["body"]["params"]["name"] for call in _peer_add_calls(rig.peer)] == ["add"] + + @pytest.mark.parametrize("surface", ("chat", "responses", "messages")) def test_server_scoped_gateway_url_exposes_only_that_servers_tools(gateway: Gateway, surface: Surface) -> None: with _rig(gateway, surface) as rig, mcp_peer() as other_peer: @@ -347,3 +365,41 @@ def test_streaming_chat_executes_the_tool_once_and_streams_the_follow_up(gateway assert text == ANSWER, response.text assert [call["body"]["params"]["name"] for call in _peer_add_calls(rig.peer)] == ["add"] assert len(rig.upstream_tools()) == 2 + + +def test_streaming_chat_through_a_toolset_gateway_url_serves_a_team_key_without_its_own_grant( + gateway: Gateway, +) -> None: + with _rig(gateway, "chat") as rig: + register_mcp(rig.scenario, rig.peer, "open" + uuid.uuid4().hex[:8], allow_all_keys=True) + toolset_name: Final = "ts" + uuid.uuid4().hex[:8] + toolset_id: Final = create_toolset(rig.scenario, ((rig.server_id, "add"),), toolset_name=toolset_name) + key: Final = rig.scenario.key(team_id=rig.scenario.team(object_permission={"mcp_toolsets": [toolset_id]})) + response: Final = rig.send(key, [{**AUTO, "server_url": f"litellm_proxy/mcp/{toolset_name}"}], stream=True) + assert response.status_code == 200, response.text + chunks: Final = tuple( + json.loads(line.removeprefix("data: ")) + for line in response.text.splitlines() + if line.startswith("data: ") and line != "data: [DONE]" + ) + text: Final = "".join( + str(chunk["choices"][0]["delta"].get("content") or "") for chunk in chunks if chunk.get("choices") + ) + assert text == ANSWER, response.text + assert [call["body"]["params"]["name"] for call in _peer_add_calls(rig.peer)] == ["add"] + requests: Final = rig.upstream_tools() + assert len(requests) == 2 and all(names == (rig.tool,) for names in requests), requests + + +@pytest.mark.parametrize("surface", ("chat", "responses", "messages")) +def test_toolset_gateway_url_gives_a_key_of_an_ungranted_team_no_tools_and_never_reaches_the_peer( + gateway: Gateway, surface: Surface +) -> None: + with _rig(gateway, surface) as rig: + toolset_name: Final = "ts" + uuid.uuid4().hex[:8] + create_toolset(rig.scenario, ((rig.server_id, "add"),), toolset_name=toolset_name) + key: Final = rig.scenario.key(team_id=rig.scenario.team()) + response: Final = rig.send(key, [{**AUTO, "server_url": f"litellm_proxy/mcp/{toolset_name}"}]) + assert _peer_add_calls(rig.peer) == (), "denied caller reached the peer" + assert all(rig.tool not in names for names in rig.upstream_tools()), rig.upstream_tools() + assert response.status_code in (200, 400, 401, 403), response.text diff --git a/tests/integration/mcp/test_mcp_oauth_flows.py b/tests/integration/mcp/test_mcp_oauth_flows.py index 7a60c8ede30..60563e7aacd 100644 --- a/tests/integration/mcp/test_mcp_oauth_flows.py +++ b/tests/integration/mcp/test_mcp_oauth_flows.py @@ -1,25 +1,31 @@ import base64 import hashlib import secrets +import time import uuid from dataclasses import dataclass from typing import Final from urllib.parse import parse_qs, urlsplit import httpx +import jwt import pytest from integration._support.client import Gateway, eventually from integration._support.database import read_rows from integration._support.mcp import ( ENTRY_POINTS, + INITIALIZE, EntryPoint, McpCaller, McpPeer, + Outcome, + _outcome_from_rpc, call_tool, mcp_peer, register_mcp, tool_calls, ) +from integration._support.mcp_grants import create_toolset from integration._support.oauth_server import AuthorizationServer, oauth_server ADD: Final = {"a": 2, "b": 3} @@ -383,3 +389,117 @@ def test_dcr_bridge_relays_client_registration_and_advertises_gateway_endpoints( assert issuer.json()["authorization_endpoint"] == f"{_base(gateway)}/{alias}/authorize" assert issuer.json()["token_endpoint"] == f"{_base(gateway)}/{alias}/token" assert "S256" in issuer.json()["code_challenge_methods_supported"] + + +def _ui_session_cookie(gateway: Gateway, user_id: str) -> dict[str, str]: + claims: Final = {"user_id": user_id, "login_method": "username_password", "exp": int(time.time()) + 600} + return {"token": jwt.encode(claims, gateway.key, algorithm="HS256")} + + +def _gateway_session_bearer(gateway: Gateway, user_id: str, resource: str | None = None) -> str: + registered: Final = gateway.client.post( + "/register", json={"redirect_uris": [CLIENT_REDIRECT], "client_name": "integration"} + ) + assert registered.status_code in (200, 201), registered.text + client_id: Final = registered.json()["client_id"] + pkce: Final = _Pkce(secrets.token_urlsafe(48)) + cookies: Final = _ui_session_cookie(gateway, user_id) + started: Final = gateway.client.get( + "/authorize", + params={ + "client_id": client_id, + "redirect_uri": CLIENT_REDIRECT, + "response_type": "code", + "state": "lit6029", + "code_challenge": pkce.challenge, + "code_challenge_method": "S256", + **({} if resource is None else {"resource": resource}), + }, + cookies=cookies, + ) + assert started.status_code == 303, started.text + handle: Final = parse_qs(urlsplit(started.headers["location"]).query)["connect_flow"][0] + completed: Final = gateway.client.post( + "/authorize/complete", data={"flow": handle}, cookies={**cookies, **dict(started.cookies)} + ) + assert completed.status_code == 303, completed.text + callback: Final = parse_qs(urlsplit(completed.headers["location"]).query) + assert "code" in callback, completed.headers["location"] + issued: Final = gateway.client.post( + "/token", + data={ + "grant_type": "authorization_code", + "code": callback["code"][0], + "redirect_uri": CLIENT_REDIRECT, + "client_id": client_id, + "code_verifier": pkce.verifier, + }, + ) + assert issued.status_code == 200, issued.text + return _issued_token(issued.json()) + + +def _toolset_rpc(gateway: Gateway, bearer: str, name: str, method: str, params: dict[str, object]) -> Outcome: + def post(rpc_method: str, rpc_params: dict[str, object]) -> httpx.Response: + return gateway.client.post( + f"/toolset/{name}/mcp", + json={"jsonrpc": "2.0", "id": 1, "method": rpc_method, "params": rpc_params}, + headers={"Authorization": f"Bearer {bearer}", "Accept": "application/json, text/event-stream"}, + ) + + initialized: Final = _outcome_from_rpc(post("initialize", dict(INITIALIZE))) + if not initialized.ok: + return initialized + return _outcome_from_rpc(post(method, params)) + + +def test_gateway_session_bearer_of_a_team_member_is_served_the_team_toolset_on_its_route(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029sess" + uuid.uuid4().hex[:6] + server_id: Final = register_mcp(scenario, peer, alias) + granted_name: Final = "lit6029g" + uuid.uuid4().hex[:8] + withheld_name: Final = "lit6029w" + uuid.uuid4().hex[:8] + granted_id: Final = create_toolset(scenario, ((server_id, "add"),), toolset_name=granted_name) + create_toolset(scenario, ((server_id, "multiply"),), toolset_name=withheld_name) + member: Final = scenario.member(scenario.team(object_permission={"mcp_toolsets": [granted_id]})) + bearer: Final = _gateway_session_bearer(gateway, member) + assert bearer.startswith("llm_session_"), bearer[:16] + listed: Final = _toolset_rpc(gateway, bearer, granted_name, "tools/list", {}) + assert listed.tools == (f"{alias}-add",), listed.raw + peer.drain() + called: Final = _toolset_rpc( + gateway, bearer, granted_name, "tools/call", {"name": f"{alias}-add", "arguments": {"a": 4, "b": 5}} + ) + assert called.ok and called.text == "9", called.raw + assert len(tool_calls(peer.drain())) == 1 + denied: Final = _toolset_rpc(gateway, bearer, withheld_name, "tools/list", {}) + assert denied.status == 403, denied.raw + assert tool_calls(peer.drain()) == () + + +def test_resource_scoped_session_bearer_opens_a_team_toolset_inside_its_server_and_none_outside( + gateway: Gateway, +) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + inside: Final = "lit6029in" + uuid.uuid4().hex[:6] + outside: Final = "lit6029out" + uuid.uuid4().hex[:6] + inside_server: Final = register_mcp(scenario, peer, inside) + outside_server: Final = register_mcp(scenario, peer, outside) + inside_name: Final = "lit6029i" + uuid.uuid4().hex[:8] + outside_name: Final = "lit6029o" + uuid.uuid4().hex[:8] + inside_id: Final = create_toolset(scenario, ((inside_server, "add"),), toolset_name=inside_name) + outside_id: Final = create_toolset(scenario, ((outside_server, "add"),), toolset_name=outside_name) + member: Final = scenario.member(scenario.team(object_permission={"mcp_toolsets": [inside_id, outside_id]})) + bearer: Final = _gateway_session_bearer(gateway, member, resource=f"{_base(gateway)}/{inside}/mcp") + assert bearer.startswith("llm_session_"), bearer[:16] + listed: Final = _toolset_rpc(gateway, bearer, inside_name, "tools/list", {}) + assert listed.tools == (f"{inside}-add",), listed.raw + peer.drain() + called: Final = _toolset_rpc( + gateway, bearer, inside_name, "tools/call", {"name": f"{inside}-add", "arguments": {"a": 4, "b": 5}} + ) + assert called.ok and called.text == "9", called.raw + assert len(tool_calls(peer.drain())) == 1 + refused: Final = _toolset_rpc(gateway, bearer, outside_name, "tools/list", {}) + assert refused.status == 403, refused.raw + assert tool_calls(peer.drain()) == () diff --git a/tests/integration/mcp/test_mcp_toolsets.py b/tests/integration/mcp/test_mcp_toolsets.py new file mode 100644 index 00000000000..3dd309db665 --- /dev/null +++ b/tests/integration/mcp/test_mcp_toolsets.py @@ -0,0 +1,509 @@ +import secrets +import uuid +from datetime import datetime, timedelta, timezone +from pathlib import Path +from typing import Final + +import httpx +import pytest +import yaml +from integration._support.client import Gateway, Scenario, object_value +from integration._support.mcp import ( + INITIALIZE, + Outcome, + _outcome_from_rest, + _outcome_from_rpc, + mcp_peer, + register_mcp, + tool_calls, +) +from integration._support.mcp_grants import create_toolset +from integration._support.process import owned_proxy + +from litellm.models.user import LiteLLM_UserTable +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.auth.auth_checks import LITELLM_SESSION_TOKEN_PREFIX, ExperimentalUIJWTToken +from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_bearer_token + +ADD: Final = {"a": 4, "b": 5} + + +def _dashboard_ui_session_token(user_id: str) -> str: + user: Final = LiteLLM_UserTable(user_id=user_id, user_role="internal_user", models=[]) + return ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(user) + + +def _toolset(scenario: Scenario, server_id: str, tool: str) -> tuple[str, str]: + name: Final = "lit6029_" + uuid.uuid4().hex[:10] + return create_toolset(scenario, ((server_id, tool),), toolset_name=name), name + + +def _toolset_rpc( + gateway: Gateway, headers: dict[str, str], name: str, method: str, params: dict[str, object] +) -> Outcome: + def post(rpc_method: str, rpc_params: dict[str, object]) -> httpx.Response: + return gateway.client.post( + f"/toolset/{name}/mcp", + json={"jsonrpc": "2.0", "id": 1, "method": rpc_method, "params": rpc_params}, + headers={**headers, "Accept": "application/json, text/event-stream"}, + ) + + initialized: Final = _outcome_from_rpc(post("initialize", dict(INITIALIZE))) + if not initialized.ok: + return initialized + return _outcome_from_rpc(post(method, params)) + + +def _listed_toolset_ids(gateway: Gateway, headers: dict[str, str]) -> tuple[str, ...]: + response: Final = gateway.client.get("/v1/mcp/toolset", headers=headers) + assert response.status_code == 200, response.text + return tuple(toolset["toolset_id"] for toolset in response.json()) + + +def _assert_team_grants_only(gateway: Gateway, team_id: str, key: str, toolset_id: str) -> None: + team: Final = object_value(gateway.get("/team/info", {"team_id": team_id})["team_info"]) + assert object_value(team["object_permission"])["mcp_toolsets"] == [toolset_id], team + key_info: Final = object_value(gateway.get("/key/info", {"key": key})["info"]) + assert key_info.get("object_permission") is None, f"key must carry no grant of its own: {key_info}" + + +def test_team_granted_toolset_is_listed_and_served_to_a_team_key(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029_" + uuid.uuid4().hex[:8] + server_id: Final = register_mcp(scenario, peer, alias) + granted_id, granted_name = _toolset(scenario, server_id, "add") + withheld_id, withheld_name = _toolset(scenario, server_id, "multiply") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id]}) + key: Final = scenario.key(team_id=team_id) + _assert_team_grants_only(gateway, team_id, key, granted_id) + headers: Final = {"Authorization": f"Bearer {key}"} + + assert _listed_toolset_ids(gateway, headers) == (granted_id,) + detail: Final = gateway.client.get(f"/v1/mcp/toolset/{granted_id}", headers=headers) + assert detail.status_code == 200, detail.text + assert detail.json()["toolset_name"] == granted_name, detail.text + withheld_detail: Final = gateway.client.get(f"/v1/mcp/toolset/{withheld_id}", headers=headers) + assert withheld_detail.status_code == 403, withheld_detail.text + + listed: Final = _toolset_rpc(gateway, headers, granted_name, "tools/list", {}) + assert listed.ok, listed.raw + assert listed.tools == (f"{alias}-add",), listed.raw + peer.drain() + called: Final = _toolset_rpc( + gateway, headers, granted_name, "tools/call", {"name": f"{alias}-add", "arguments": ADD} + ) + assert called.ok and called.text == "9", called.raw + assert len(tool_calls(peer.drain())) == 1 + denied: Final = _toolset_rpc(gateway, headers, withheld_name, "tools/list", {}) + assert denied.status == 403, denied.raw + + +def test_dashboard_session_of_a_team_member_lists_the_team_granted_toolset( + gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-integration-salt") + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029_" + uuid.uuid4().hex[:8] + server_id: Final = register_mcp(scenario, peer, alias) + granted_id, granted_name = _toolset(scenario, server_id, "add") + _toolset(scenario, server_id, "multiply") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id]}) + user_id: Final = scenario.user(user_role="internal_user", teams=[team_id]) + user: Final = object_value(gateway.get("/user/info", {"user_id": user_id})["user_info"]) + assert user["teams"] == [team_id], user + headers: Final = {"Authorization": f"Bearer {_dashboard_ui_session_token(user_id)}"} + + assert _listed_toolset_ids(gateway, headers) == (granted_id,) + detail: Final = gateway.client.get(f"/v1/mcp/toolset/{granted_id}", headers=headers) + assert detail.status_code == 200, detail.text + assert detail.json()["toolset_name"] == granted_name, detail.text + + +def test_direct_grants_no_grants_and_admin_listing_are_unchanged_by_team_resolution(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029_" + uuid.uuid4().hex[:8] + server_id: Final = register_mcp(scenario, peer, alias) + granted_id, granted_name = _toolset(scenario, server_id, "add") + withheld_id, withheld_name = _toolset(scenario, server_id, "multiply") + direct: Final = {"Authorization": f"Bearer {scenario.key(object_permission={'mcp_toolsets': [granted_id]})}"} + ungranted_team: Final = scenario.team() + no_grant: Final = {"Authorization": f"Bearer {scenario.key(team_id=ungranted_team)}"} + admin: Final = {"Authorization": f"Bearer {gateway.key}"} + + assert _listed_toolset_ids(gateway, direct) == (granted_id,) + assert _toolset_rpc(gateway, direct, granted_name, "tools/list", {}).tools == (f"{alias}-add",) + assert _toolset_rpc(gateway, direct, withheld_name, "tools/list", {}).status == 403 + assert gateway.client.get(f"/v1/mcp/toolset/{withheld_id}", headers=direct).status_code == 403 + + assert _listed_toolset_ids(gateway, no_grant) == () + assert gateway.client.get(f"/v1/mcp/toolset/{granted_id}", headers=no_grant).status_code == 403 + assert _toolset_rpc(gateway, no_grant, granted_name, "tools/list", {}).status == 403 + + assert {granted_id, withheld_id} <= set(_listed_toolset_ids(gateway, admin)) + assert gateway.client.get(f"/v1/mcp/toolset/{withheld_id}", headers=admin).status_code == 200 + assert _toolset_rpc(gateway, admin, withheld_name, "tools/list", {}).tools == (f"{alias}-multiply",) + + +def test_a_key_with_its_own_toolset_grant_does_not_inherit_the_team_toolset(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029_" + uuid.uuid4().hex[:8] + server_id: Final = register_mcp(scenario, peer, alias) + own_id, own_name = _toolset(scenario, server_id, "add") + team_only_id, team_only_name = _toolset(scenario, server_id, "multiply") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [own_id, team_only_id]}) + key: Final = scenario.key(team_id=team_id, object_permission={"mcp_toolsets": [own_id]}) + headers: Final = {"Authorization": f"Bearer {key}"} + + assert _listed_toolset_ids(gateway, headers) == (own_id,) + assert gateway.client.get(f"/v1/mcp/toolset/{team_only_id}", headers=headers).status_code == 403 + assert _toolset_rpc(gateway, headers, team_only_name, "tools/list", {}).status == 403 + assert _toolset_rpc(gateway, headers, own_name, "tools/list", {}).tools == (f"{alias}-add",) + + +def _team_member_with_own_grant(scenario: Scenario, team_id: str, own_server_id: str) -> str: + user_id: Final = scenario.user(user_role="internal_user", object_permission={"mcp_servers": [own_server_id]}) + scenario.gateway.post("/team/member_add", {"team_id": team_id, "member": {"role": "user", "user_id": user_id}}) + return user_id + + +def test_dashboard_session_serves_the_team_toolset_despite_a_disjoint_grant_on_the_user_row( + gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + """The member's own row grants a different server outright. That grant must not cap the team's toolset + to nothing, and the team's sibling toolset must not leak onto the granted toolset's route.""" + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-integration-salt") + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029_" + uuid.uuid4().hex[:8] + server_id: Final = register_mcp(scenario, peer, alias) + own_server_id: Final = register_mcp(scenario, peer, "lit6029_own_" + uuid.uuid4().hex[:8]) + granted_id, granted_name = _toolset(scenario, server_id, "add") + sibling_id, sibling_name = _toolset(scenario, server_id, "multiply") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id, sibling_id]}) + user_id: Final = _team_member_with_own_grant(scenario, team_id, own_server_id) + headers: Final = {"Authorization": f"Bearer {_dashboard_ui_session_token(user_id)}"} + + assert set(_listed_toolset_ids(gateway, headers)) == {granted_id, sibling_id} + listed: Final = _toolset_rpc(gateway, headers, granted_name, "tools/list", {}) + assert listed.ok, listed.raw + assert listed.tools == (f"{alias}-add",), listed.raw + assert _toolset_rpc(gateway, headers, sibling_name, "tools/list", {}).tools == (f"{alias}-multiply",) + peer.drain() + called: Final = _toolset_rpc( + gateway, headers, granted_name, "tools/call", {"name": f"{alias}-add", "arguments": ADD} + ) + assert called.ok and called.text == "9", called.raw + assert len(tool_calls(peer.drain())) == 1 + stranger: Final = { + "Authorization": f"Bearer {_dashboard_ui_session_token(scenario.user(user_role='internal_user'))}" + } + assert _toolset_rpc(gateway, stranger, granted_name, "tools/list", {}).status == 403 + + +def test_a_member_removed_from_the_team_loses_its_toolset_on_the_dashboard_session( + gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-integration-salt") + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029_" + uuid.uuid4().hex[:8] + server_id: Final = register_mcp(scenario, peer, alias) + granted_id, granted_name = _toolset(scenario, server_id, "add") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id]}) + user_id: Final = scenario.member(team_id) + headers: Final = {"Authorization": f"Bearer {_dashboard_ui_session_token(user_id)}"} + assert _listed_toolset_ids(gateway, headers) == (granted_id,) + assert _toolset_rpc(gateway, headers, granted_name, "tools/list", {}).tools == (f"{alias}-add",) + + gateway.post("/team/member_delete", {"team_id": team_id, "user_id": user_id}) + + assert _listed_toolset_ids(gateway, headers) == () + assert gateway.client.get(f"/v1/mcp/toolset/{granted_id}", headers=headers).status_code == 403 + assert _toolset_rpc(gateway, headers, granted_name, "tools/list", {}).status == 403 + + +def _bearer(token: str) -> dict[str, str]: + return {"Authorization": f"Bearer {token}"} + + +def _expired_dashboard_token(user_id: str) -> str: + expired: Final = (datetime.now(timezone.utc) - timedelta(minutes=5)).strftime("%Y-%m-%dT%H:%M:%S.%f")[:-3] + stale: Final = UserAPIKeyAuth( + token="ui-token", + key_name="ui-token", + key_alias="ui-token", + expires=expired + "+00:00", + user_id=user_id, + team_id="litellm-dashboard", + models=[], + user_role=LitellmUserRoles.INTERNAL_USER, + ) + return encrypt_bearer_token(stale.model_dump_json(exclude_none=True), prefix=LITELLM_SESSION_TOKEN_PREFIX) + + +def _rest_list(gateway: Gateway, headers: dict[str, str], params: object) -> httpx.Response: + return gateway.client.get("/mcp-rest/tools/list", headers=headers, params=params) + + +def _rest_call(gateway: Gateway, headers: dict[str, str], name: str, server_id: str) -> Outcome: + return _outcome_from_rest( + gateway.client.post( + "/mcp-rest/tools/call", + headers=headers, + json={"name": name, "arguments": dict(ADD), "server_id": server_id}, + ) + ) + + +def _route_tools(gateway: Gateway, headers: dict[str, str], name: str) -> Outcome: + return _toolset_rpc(gateway, headers, name, "tools/list", {}) + + +def _route_call(gateway: Gateway, headers: dict[str, str], name: str, tool: str) -> Outcome: + return _toolset_rpc(gateway, headers, name, "tools/call", {"name": tool, "arguments": dict(ADD)}) + + +def _strict_config(directory: Path) -> Path: + base: Final = yaml.safe_load((Path(__file__).resolve().parents[1] / "proxy_config.yaml").read_text()) + strict: Final = { + **base, + "general_settings": {**base.get("general_settings", {}), "require_key_mcp_access_defined": True}, + } + path: Final = directory / "require_key_mcp_access.yaml" + path.write_text(yaml.safe_dump(strict)) + return path + + +def test_mcp_rest_toolset_name_narrows_the_list_and_serves_the_call_for_a_team_key_and_a_dashboard_member( + gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-integration-salt") + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029rest" + uuid.uuid4().hex[:6] + server_id: Final = register_mcp(scenario, peer, alias) + granted_id, granted_name = _toolset(scenario, server_id, "add") + _, withheld_name = _toolset(scenario, server_id, "multiply") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id]}) + key: Final = scenario.key(team_id=team_id) + member: Final = scenario.member(team_id) + for headers in (_bearer(key), _bearer(_dashboard_ui_session_token(member))): + listed: Final = _rest_list(gateway, headers, {"toolset_name": granted_name}) + assert listed.status_code == 200, listed.text + listed_names: Final = tuple(tool["name"] for tool in listed.json()["tools"]) + assert len(listed_names) == 1 and listed_names[0].endswith("add"), listed.text + denied: Final = _rest_list(gateway, headers, {"toolset_name": withheld_name}) + assert denied.status_code == 200 and denied.json()["tools"] == [], denied.text + assert "does not have access to toolset" in denied.json()["message"], denied.text + peer.drain() + called: Final = _rest_call(gateway, headers, listed_names[0], server_id) + assert called.ok and called.text == "9", called.raw + assert len(tool_calls(peer.drain())) == 1 + + +def test_a_key_restricted_to_its_own_servers_does_not_inherit_the_team_toolset(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029own" + uuid.uuid4().hex[:6] + server_id: Final = register_mcp(scenario, peer, alias) + other_id: Final = register_mcp(scenario, peer, "lit6029other" + uuid.uuid4().hex[:6]) + granted_id, granted_name = _toolset(scenario, server_id, "add") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id], "mcp_servers": [other_id]}) + key: Final = scenario.key(team_id=team_id, object_permission={"mcp_servers": [other_id]}) + headers: Final = _bearer(key) + assert _listed_toolset_ids(gateway, headers) == () + assert gateway.client.get(f"/v1/mcp/toolset/{granted_id}", headers=headers).status_code == 403 + assert _route_tools(gateway, headers, granted_name).status == 403 + assert tool_calls(peer.drain()) == () + + +def test_a_team_with_an_empty_or_absent_toolset_grant_gives_its_keys_nothing(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + server_id: Final = register_mcp(scenario, peer, "lit6029empty" + uuid.uuid4().hex[:6]) + granted_id, granted_name = _toolset(scenario, server_id, "add") + teams: Final = ( + scenario.team(object_permission={"mcp_toolsets": []}), + scenario.team(object_permission={"mcp_toolsets": None}), + scenario.team(), + ) + for team_id in teams: + headers: Final = _bearer(scenario.key(team_id=team_id)) + assert _listed_toolset_ids(gateway, headers) == (), team_id + assert gateway.client.get(f"/v1/mcp/toolset/{granted_id}", headers=headers).status_code == 403 + assert _route_tools(gateway, headers, granted_name).status == 403, team_id + assert tool_calls(peer.drain()) == () + + +def test_an_unknown_or_malformed_toolset_name_is_refused_without_peer_traffic_and_the_route_keeps_serving( + gateway: Gateway, +) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029bad" + uuid.uuid4().hex[:6] + server_id: Final = register_mcp(scenario, peer, alias) + granted_id, granted_name = _toolset(scenario, server_id, "add") + _, withheld_name = _toolset(scenario, server_id, "multiply") + headers: Final = _bearer(scenario.key(team_id=scenario.team(object_permission={"mcp_toolsets": [granted_id]}))) + unknown: Final = "missing" + uuid.uuid4().hex[:8] + assert _route_tools(gateway, headers, unknown).status == 404 + assert _rest_list(gateway, headers, {"toolset_name": unknown}).status_code == 404 + malformed: Final = ( + [("toolset_name", granted_name), ("toolset_name", withheld_name)], + {"toolset_name": "x" * 5000}, + {"toolset_name": ""}, + {"toolset_name": granted_name + "\x00"}, + ) + statuses: Final = tuple(_rest_list(gateway, headers, params).status_code for params in malformed) + assert all(status < 500 for status in statuses), statuses + assert tool_calls(peer.drain()) == () + assert gateway.client.get("/health/liveliness").status_code == 200 + served: Final = _route_tools(gateway, headers, granted_name) + assert served.tools == (f"{alias}-add",), served.raw + + +def test_garbage_expired_and_tampered_credentials_are_refused_on_every_toolset_surface( + gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-integration-salt") + with mcp_peer() as peer, gateway.scenario() as scenario: + server_id: Final = register_mcp(scenario, peer, "lit6029cred" + uuid.uuid4().hex[:6]) + granted_id, granted_name = _toolset(scenario, server_id, "add") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id]}) + member: Final = scenario.member(team_id) + forged: Final = ( + "sk-" + secrets.token_urlsafe(24), + _expired_dashboard_token(member), + "llm_session_" + secrets.token_urlsafe(32), + ) + for bearer in forged: + headers: Final = _bearer(bearer) + listed: Final = gateway.client.get("/v1/mcp/toolset", headers=headers) + assert listed.status_code == 401, (bearer[:12], listed.text) + detail: Final = gateway.client.get(f"/v1/mcp/toolset/{granted_id}", headers=headers) + assert detail.status_code == 401, (bearer[:12], detail.text) + routed: Final = _route_tools(gateway, headers, granted_name) + assert routed.status == 401, (bearer[:12], routed.raw) + rest: Final = _rest_list(gateway, headers, {"toolset_name": granted_name}) + assert rest.status_code == 401, (bearer[:12], rest.text) + assert peer.drain() == () + + +def test_a_dashboard_member_of_a_deleted_team_loses_the_toolset_while_a_direct_grant_survives( + gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-integration-salt") + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029gone" + uuid.uuid4().hex[:6] + server_id: Final = register_mcp(scenario, peer, alias) + team_id_, team_name = _toolset(scenario, server_id, "add") + own_id, own_name = _toolset(scenario, server_id, "multiply") + doomed: Final = scenario.gateway.post( + "/team/new", + {"team_alias": f"integration-{uuid.uuid4().hex}", "object_permission": {"mcp_toolsets": [team_id_]}}, + ) + doomed_team: Final = str(doomed["team_id"]) + member: Final = scenario.user(user_role="internal_user", teams=[doomed_team]) + granted: Final = scenario.user( + user_role="internal_user", teams=[doomed_team], object_permission={"mcp_toolsets": [own_id]} + ) + member_headers: Final = _bearer(_dashboard_ui_session_token(member)) + granted_headers: Final = _bearer(_dashboard_ui_session_token(granted)) + assert _listed_toolset_ids(gateway, member_headers) == (team_id_,) + assert set(_listed_toolset_ids(gateway, granted_headers)) == {team_id_, own_id} + scenario.delete_team(doomed_team) + assert _listed_toolset_ids(gateway, member_headers) == () + assert gateway.client.get(f"/v1/mcp/toolset/{team_id_}", headers=member_headers).status_code == 403 + assert _route_tools(gateway, member_headers, team_name).status == 403 + assert _listed_toolset_ids(gateway, granted_headers) == (own_id,) + assert _route_tools(gateway, granted_headers, team_name).status == 403 + assert _route_tools(gateway, granted_headers, own_name).tools == (f"{alias}-multiply",) + peer.drain() + kept: Final = _route_call(gateway, granted_headers, own_name, f"{alias}-multiply") + assert kept.ok and kept.text == "20", kept.raw + assert [call["body"]["params"]["name"] for call in tool_calls(peer.drain())] == ["multiply"] + + +def test_require_key_mcp_access_defined_stops_key_inheritance_but_not_the_dashboard_member( + gateway: Gateway, tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-integration-salt") + with ( + owned_proxy(gateway, tmp_path, {}, config=_strict_config(tmp_path), workers=2) as strict, + mcp_peer() as peer, + strict.scenario() as scenario, + ): + alias: Final = "lit6029strict" + uuid.uuid4().hex[:6] + server_id: Final = register_mcp(scenario, peer, alias) + granted_id, granted_name = _toolset(scenario, server_id, "add") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id]}) + inheriting: Final = _bearer(scenario.key(team_id=team_id)) + own: Final = _bearer(scenario.key(team_id=team_id, object_permission={"mcp_toolsets": [granted_id]})) + member: Final = _bearer(_dashboard_ui_session_token(scenario.member(team_id))) + assert _listed_toolset_ids(strict, inheriting) == () + assert _route_tools(strict, inheriting, granted_name).status == 403 + assert _listed_toolset_ids(strict, own) == (granted_id,) + assert _route_tools(strict, own, granted_name).tools == (f"{alias}-add",) + assert _listed_toolset_ids(strict, member) == (granted_id,) + assert _route_tools(strict, member, granted_name).tools == (f"{alias}-add",) + peer.drain() + called: Final = _route_call(strict, member, granted_name, f"{alias}-add") + assert called.ok and called.text == "9", called.raw + assert len(tool_calls(peer.drain())) == 1 + + +def test_a_key_with_only_a_vector_store_grant_still_inherits_the_team_toolset(gateway: Gateway) -> None: + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029vs" + uuid.uuid4().hex[:6] + server_id: Final = register_mcp(scenario, peer, alias) + granted_id, granted_name = _toolset(scenario, server_id, "add") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id]}) + headers: Final = _bearer( + scenario.key(team_id=team_id, object_permission={"vector_stores": ["vs-" + uuid.uuid4().hex[:8]]}) + ) + assert _listed_toolset_ids(gateway, headers) == (granted_id,) + assert _route_tools(gateway, headers, granted_name).tools == (f"{alias}-add",) + peer.drain() + called: Final = _route_call(gateway, headers, granted_name, f"{alias}-add") + assert called.ok and called.text == "9", called.raw + assert len(tool_calls(peer.drain())) == 1 + + +def test_a_member_added_after_the_team_was_cached_sees_the_toolset_on_every_following_request( + gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-integration-salt") + with mcp_peer() as peer, gateway.scenario() as scenario: + server_id: Final = register_mcp(scenario, peer, "lit6029cache" + uuid.uuid4().hex[:6]) + granted_id, _ = _toolset(scenario, server_id, "add") + team_id: Final = scenario.team(object_permission={"mcp_toolsets": [granted_id]}) + user_id: Final = scenario.user(user_role="internal_user") + headers: Final = _bearer(_dashboard_ui_session_token(user_id)) + warm: Final = _bearer(scenario.key(team_id=team_id)) + assert tuple(_listed_toolset_ids(gateway, warm) for _ in range(4)) == ((granted_id,),) * 4 + assert tuple(_listed_toolset_ids(gateway, headers) for _ in range(4)) == ((),) * 4 + added: Final = gateway.request( + "POST", "/team/member_add", {"team_id": team_id, "member": {"user_id": user_id, "role": "user"}} + ) + assert added.status_code == 200, added.text + listings: Final = tuple(_listed_toolset_ids(gateway, headers) for _ in range(8)) + assert listings == ((granted_id,),) * 8, listings + + +def test_a_member_of_two_teams_sees_the_union_and_each_route_stays_narrowed_to_its_own_toolset( + gateway: Gateway, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-integration-salt") + with mcp_peer() as peer, gateway.scenario() as scenario: + alias: Final = "lit6029two" + uuid.uuid4().hex[:6] + server_id: Final = register_mcp(scenario, peer, alias) + first_id, first_name = _toolset(scenario, server_id, "add") + second_id, second_name = _toolset(scenario, server_id, "multiply") + teams: Final = ( + scenario.team(object_permission={"mcp_toolsets": [first_id]}), + scenario.team(object_permission={"mcp_toolsets": [second_id]}), + ) + headers: Final = _bearer( + _dashboard_ui_session_token(scenario.user(user_role="internal_user", teams=list(teams))) + ) + assert set(_listed_toolset_ids(gateway, headers)) == {first_id, second_id} + assert _route_tools(gateway, headers, first_name).tools == (f"{alias}-add",) + assert _route_tools(gateway, headers, second_name).tools == (f"{alias}-multiply",) + peer.drain() + crossed: Final = _route_call(gateway, headers, first_name, f"{alias}-multiply") + assert not crossed.ok, crossed.raw + assert tool_calls(peer.drain()) == () diff --git a/tests/integration/providers/test_bedrock_converse_missing_content_wire.py b/tests/integration/providers/test_bedrock_converse_missing_content_wire.py new file mode 100644 index 00000000000..057596941a2 --- /dev/null +++ b/tests/integration/providers/test_bedrock_converse_missing_content_wire.py @@ -0,0 +1,966 @@ +import asyncio +import base64 +import json +import os +import re +import signal +import threading +import uuid +from collections.abc import Iterable, Iterator, Mapping, Sequence +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final, Literal +from urllib.parse import unquote, urlsplit + +import anthropic +import httpx +import openai +import psutil +import pytest +import yaml +from integration._support.client import ( + Gateway, + Scenario, + eventually, + gateway_from_environment, + object_value, + string_value, +) +from integration._support.database import read_rows +from integration._support.process import owned_proxy_process +from integration._support.upstream import _aws_event_frame +from integration._support.wire import Reply, Request, Wire, wire_server +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_if_encrypted_with +from pydantic import JsonValue, TypeAdapter + +_MODEL_ID: Final = "anthropic.claude-3-haiku-20240307-v1:0" +_CONVERSE_MODEL: Final = f"bedrock/converse/{_MODEL_ID}" +_INVOKE_MODEL: Final = f"bedrock/invoke/{_MODEL_ID}" +_CONVERSE_TARGET: Final = f"/model/{_MODEL_ID}/converse" +_STREAM_TARGET: Final = f"/model/{_MODEL_ID}/converse-stream" +_INVOKE_TARGET: Final = f"/model/{_MODEL_ID}/invoke" +_ANSWER: Final = "bedrock missing content control" +_RESPONSE: Final = json.dumps( + { + "output": {"message": {"role": "assistant", "content": [{"text": _ANSWER}]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15}, + "metrics": {"latencyMs": 1}, + } +).encode() +_EVENT_STREAM: Final = "application/vnd.amazon.eventstream" +_STREAM_EVENTS: Final[tuple[tuple[str, dict[str, JsonValue]], ...]] = ( + ("messageStart", {"role": "assistant"}), + ("contentBlockDelta", {"delta": {"text": _ANSWER}, "contentBlockIndex": 0}), + ("messageStop", {"stopReason": "end_turn"}), + ("metadata", {"usage": {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15}}), +) +_STREAM_BYTES: Final = b"".join(_aws_event_frame(kind, payload, "sc", "u") for kind, payload in _STREAM_EVENTS) +_DEFAULT_CONTINUE: Final = "Please continue." +_DEPLOYMENT_CONTINUE: Final = "Deployment says continue." +_DEPLOYMENT_CONTINUE_MESSAGE: Final[dict[str, JsonValue]] = {"role": "user", "content": _DEPLOYMENT_CONTINUE} +_NO_NON_SYSTEM_MESSAGE: Final = "bedrock requires at least one non-system message" +_JSON: Final = TypeAdapter(dict[str, JsonValue]) +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_CALL_INDEX: Final = re.compile(r"call-[0-9a-f]{32}-(\d+)") +_QUESTION: Final = "What is the capital of France?" +_ANSWERED: Final = "Paris." +_FOLLOW_UP: Final = "And the capital of Spain?" +_QUESTION_TURN: Final[dict[str, JsonValue]] = {"role": "user", "content": _QUESTION} +_ANSWERED_TURN: Final[dict[str, JsonValue]] = {"role": "assistant", "content": _ANSWERED} +_FOLLOW_UP_TURN: Final[dict[str, JsonValue]] = {"role": "user", "content": _FOLLOW_UP} +_NO_CONTENT_USER: Final[dict[str, JsonValue]] = {"role": "user"} +_NULL_CONTENT_USER: Final[dict[str, JsonValue]] = {"role": "user", "content": None} +_EMPTY_CONTENT_USER: Final[dict[str, JsonValue]] = {"role": "user", "content": ""} +_NO_CONTENT_SYSTEM: Final[dict[str, JsonValue]] = {"role": "system"} +_NULL_CONTENT_SYSTEM: Final[dict[str, JsonValue]] = {"role": "system", "content": None} +_NO_CONTENT_ASSISTANT: Final[dict[str, JsonValue]] = {"role": "assistant"} +_TOOL_CALL_TURN: Final[dict[str, JsonValue]] = { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": "call_1", "type": "function", "function": {"name": "get_weather", "arguments": '{"city": "Boston"}'}} + ], +} +_NO_CONTENT_TOOL: Final[dict[str, JsonValue]] = {"role": "tool", "tool_call_id": "call_1"} +_TOOLS: Final[tuple[dict[str, JsonValue], ...]] = ( + { + "type": "function", + "function": { + "name": "get_weather", + "parameters": {"type": "object", "properties": {"city": {"type": "string"}}}, + }, + }, +) +_CONVERSE_QUESTION: Final[dict[str, JsonValue]] = {"role": "user", "content": [{"text": _QUESTION}]} +_CONVERSE_ANSWERED: Final[dict[str, JsonValue]] = {"role": "assistant", "content": [{"text": _ANSWERED}]} +_CONVERSE_FOLLOW_UP: Final[dict[str, JsonValue]] = {"role": "user", "content": [{"text": _FOLLOW_UP}]} +_CONVERSE_TOOL_USE: Final[dict[str, JsonValue]] = { + "role": "assistant", + "content": [{"toolUse": {"toolUseId": "call_1", "name": "get_weather", "input": {"city": "Boston"}}}], +} +_CONVERSE_EMPTY_TOOL_RESULT: Final[dict[str, JsonValue]] = { + "role": "user", + "content": [{"toolResult": {"toolUseId": "call_1", "content": []}}], +} +_NEUTRALIZED_TOOL_CALL: Final[dict[str, JsonValue]] = { + "role": "assistant", + "content": [{"text": '[tool call call_1: get_weather({"city": "Boston"})]'}], +} +_NEUTRALIZED_TOOL_RESULT: Final[dict[str, JsonValue]] = { + "role": "user", + "content": [{"text": "[tool result for call_1: ]"}], +} +_EXTRA: Final[dict[str, JsonValue]] = {"num_retries": 0, "cache": {"no-cache": True}} +_SIGNING_KEY: Final = os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt") +_AWS: Final[dict[str, JsonValue]] = { + "aws_access_key_id": "AKIASCRIPTEDPROVIDER", + "aws_secret_access_key": "scripted-secret", + "aws_region_name": "us-east-1", +} +_PLAIN: Final = "bedrock-missing-content-plain" +_CONTINUE: Final = "bedrock-missing-content-continue" + +Endpoint = Literal["chat", "responses"] + + +@dataclass(frozen=True, slots=True) +class _Call: + endpoint: Endpoint + stream: bool + user: str + index: int + + +@dataclass(frozen=True, slots=True) +class _Served: + call: _Call + status: int + text: str + + +@dataclass(frozen=True, slots=True) +class _ModifyParamsCell: + row: str + model: str + messages: tuple[dict[str, JsonValue], ...] + expected: tuple[dict[str, JsonValue], ...] + extra: Mapping[str, JsonValue] = MappingProxyType({}) + + +def _continue_turn(text: str) -> dict[str, JsonValue]: + return {"role": "user", "content": [{"text": text}]} + + +_MODIFY_PARAMS_CELLS: Final = ( + _ModifyParamsCell("r11", _PLAIN, (_NO_CONTENT_USER,), (_continue_turn(_DEFAULT_CONTINUE),)), + _ModifyParamsCell( + "r12", + _PLAIN, + (_QUESTION_TURN, _ANSWERED_TURN, _NULL_CONTENT_USER), + (_CONVERSE_QUESTION, _CONVERSE_ANSWERED, _continue_turn(_DEFAULT_CONTINUE)), + ), + _ModifyParamsCell( + "r13", + _PLAIN, + (_QUESTION_TURN, _TOOL_CALL_TURN, _NO_CONTENT_TOOL), + (_CONVERSE_QUESTION, _CONVERSE_TOOL_USE, _CONVERSE_EMPTY_TOOL_RESULT), + MappingProxyType({"tools": list(_TOOLS)}), + ), + _ModifyParamsCell("r14", _PLAIN, (_NO_CONTENT_SYSTEM, _QUESTION_TURN), (_CONVERSE_QUESTION,)), + _ModifyParamsCell("r15", _CONTINUE, (_NO_CONTENT_USER,), (_continue_turn(_DEPLOYMENT_CONTINUE),)), +) + + +def _converse_peer(request: Request) -> Reply: + if unquote(request.target) == _STREAM_TARGET: + return Reply(body=_STREAM_BYTES, content_type=_EVENT_STREAM) + return Reply(body=_RESPONSE) + + +def _scripted_error(status: int, message: str) -> Reply: + return Reply(status=status, body=json.dumps({"message": message}).encode()) + + +def _converse_deployment(scenario: Scenario, wire: Wire, **extra: JsonValue) -> str: + return scenario.model(model=_CONVERSE_MODEL, api_base=wire.url, **_AWS, **extra) + + +def _auth(gateway: Gateway) -> dict[str, str]: + return {"Authorization": f"Bearer {gateway.key}"} + + +def _proxy_url(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + + +def _chat( + model: str, + messages: Sequence[Mapping[str, JsonValue]], + *, + stream: bool = False, + cached: bool = False, + **extra: JsonValue, +) -> dict[str, JsonValue]: + return { + "model": model, + "messages": [dict(message) for message in messages], + "max_tokens": 16, + "stream": stream, + "num_retries": 0, + **({} if cached else {"cache": {"no-cache": True}}), + **extra, + } + + +def _post_chat(gateway: Gateway, body: Mapping[str, JsonValue]) -> httpx.Response: + return gateway.request("POST", "/v1/chat/completions", body) + + +def _chat_answer(response: httpx.Response) -> str: + assert response.status_code == 200, response.text + body: Final = response.json() + assert body["choices"][0]["message"]["content"] == _ANSWER, response.text + return body["id"] + + +def _only_received(wire: Wire) -> tuple[str, dict[str, JsonValue]]: + (request,) = wire.drain() + return unquote(request.target), json.loads(request.body) + + +def _sse_payloads(lines: Iterable[str]) -> tuple[dict[str, JsonValue], ...]: + return tuple(json.loads(line[6:]) for line in lines if line.startswith("data: ") and line != "data: [DONE]") + + +def _stream_lines(gateway: Gateway, path: str, body: Mapping[str, JsonValue]) -> tuple[str, ...]: + with gateway.client.stream("POST", path, json=body, headers=_auth(gateway)) as response: + lines: Final = tuple(line for line in response.iter_lines() if line) + status_code: Final = response.status_code + assert status_code == 200, "\n".join(lines) + return lines + + +def _chat_stream_text(chunks: Iterable[dict[str, JsonValue]]) -> str: + return "".join(chunk["choices"][0]["delta"].get("content") or "" for chunk in chunks if chunk["choices"]) + + +def _spend_row(request_id: str) -> dict[str, JsonValue]: + (row,) = eventually( + lambda: read_rows( + 'SELECT request_id, status, call_type, end_user FROM "LiteLLM_SpendLogs" WHERE request_id=%s', + (request_id,), + ), + lambda found: len(found) >= 1, + seconds=70, + ) + return row + + +def _success_rows(prefix: str, expected: int) -> tuple[dict[str, JsonValue], ...]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id, call_type, end_user FROM "LiteLLM_SpendLogs" WHERE end_user LIKE %s AND status=%s', + (f"{prefix}%", "success"), + ), + lambda found: len(found) >= expected, + seconds=70, + ) + assert len(rows) == expected, rows + assert len({row["request_id"] for row in rows}) == expected, rows + return tuple(rows) + + +def _converse_cell(gateway: Gateway, wire: Wire, body: Mapping[str, JsonValue]) -> tuple[str, dict[str, JsonValue]]: + identity: Final = _chat_answer(_post_chat(gateway, body)) + target, received = _only_received(wire) + assert target == _CONVERSE_TARGET, target + assert _spend_row(identity)["status"] == "success" + return identity, received + + +def _owned_config(wire: Wire, directory: Path, *, modify_params: bool) -> Path: + base: Final = _JSON.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + deployment: Final[dict[str, JsonValue]] = { + "model": _CONVERSE_MODEL, + "api_base": wire.url, + "api_key": "integration-provider-key", + **_AWS, + } + config: Final[dict[str, JsonValue]] = { + **base, + "model_list": [ + {"model_name": _PLAIN, "litellm_params": deployment}, + { + "model_name": _CONTINUE, + "litellm_params": {**deployment, "user_continue_message": _DEPLOYMENT_CONTINUE_MESSAGE}, + }, + ], + "litellm_settings": {**_JSON.validate_python(base["litellm_settings"]), "modify_params": modify_params}, + "router_settings": {**_JSON.validate_python(base["router_settings"]), "num_retries": 0}, + } + path: Final = directory / f"bedrock-missing-content-{'modify-params' if modify_params else 'plain'}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +@pytest.fixture(scope="module") +def modify_params_proxy(tmp_path_factory: pytest.TempPathFactory) -> Iterator[tuple[Gateway, Wire]]: + directory: Final = tmp_path_factory.mktemp("bedrock-modify-params") + with gateway_from_environment() as gateway, wire_server(_converse_peer) as wire: + config: Final = _owned_config(wire, directory, modify_params=True) + with owned_proxy_process(gateway, directory, {}, config=config, workers=2) as owned: + yield owned.gateway, wire + + +def test_r01_openai_sync_lone_user_without_content_sends_no_converse_block(gateway: Gateway) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + with openai.OpenAI(base_url=f"{_proxy_url(gateway)}/v1", api_key=gateway.key, max_retries=0) as client: + completion: Final = client.chat.completions.create( + model=model, messages=[_NO_CONTENT_USER], max_tokens=16, extra_body=_EXTRA + ) + assert completion.choices[0].message.content == _ANSWER, completion + target, received = _only_received(wire) + assert target == _CONVERSE_TARGET, target + assert received["messages"] == [] and "system" not in received, received + assert _spend_row(completion.id)["status"] == "success" + + +async def test_r02_openai_async_user_with_null_content_after_an_assistant_turn_keeps_the_earlier_turns( + gateway: Gateway, +) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + async with openai.AsyncOpenAI( + base_url=f"{_proxy_url(gateway)}/v1", api_key=gateway.key, max_retries=0 + ) as client: + completion: Final = await client.chat.completions.create( + model=model, + messages=[_QUESTION_TURN, _ANSWERED_TURN, _NULL_CONTENT_USER], + max_tokens=16, + extra_body=_EXTRA, + ) + assert completion.choices[0].message.content == _ANSWER, completion + target, received = _only_received(wire) + assert target == _CONVERSE_TARGET, target + assert received["messages"] == [_CONVERSE_QUESTION, _CONVERSE_ANSWERED], received + assert _spend_row(completion.id)["status"] == "success" + + +def test_r03_httpx_raw_sse_user_without_content_after_an_assistant_turn_streams_to_done(gateway: Gateway) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + lines: Final = _stream_lines( + gateway, + "/v1/chat/completions", + _chat(model, (_QUESTION_TURN, _ANSWERED_TURN, _NO_CONTENT_USER), stream=True), + ) + assert lines[-1] == "data: [DONE]", lines + chunks: Final = _sse_payloads(lines) + assert _chat_stream_text(chunks) == _ANSWER, lines + (identity,) = {chunk["id"] for chunk in chunks} + target, received = _only_received(wire) + assert target == _STREAM_TARGET, target + assert received["messages"] == [_CONVERSE_QUESTION, _CONVERSE_ANSWERED], received + assert _spend_row(identity)["status"] == "success" + + +async def test_r04_openai_async_stream_tool_turn_without_content_sends_an_empty_tool_result(gateway: Gateway) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + async with openai.AsyncOpenAI( + base_url=f"{_proxy_url(gateway)}/v1", api_key=gateway.key, max_retries=0 + ) as client: + stream: Final = await client.chat.completions.create( + model=model, + messages=[_QUESTION_TURN, _TOOL_CALL_TURN, _NO_CONTENT_TOOL], + tools=list(_TOOLS), + max_tokens=16, + stream=True, + extra_body=_EXTRA, + ) + chunks: Final = [chunk async for chunk in stream] + assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) == _ANSWER, chunks + (identity,) = {chunk.id for chunk in chunks} + target, received = _only_received(wire) + assert target == _STREAM_TARGET, target + assert received["messages"] == [_CONVERSE_QUESTION, _CONVERSE_TOOL_USE, _CONVERSE_EMPTY_TOOL_RESULT], received + assert received["toolConfig"]["tools"][0]["toolSpec"]["name"] == "get_weather", received + assert _spend_row(identity)["status"] == "success" + + +@pytest.mark.parametrize("system_turn", (_NO_CONTENT_SYSTEM, _NULL_CONTENT_SYSTEM), ids=("r05", "r06")) +def test_r05_r06_leading_system_without_content_is_dropped(gateway: Gateway, system_turn: dict[str, JsonValue]) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + _, received = _converse_cell(gateway, wire, _chat(model, (system_turn, _QUESTION_TURN))) + assert "system" not in received, received + assert received["messages"] == [_CONVERSE_QUESTION], received + + +def test_r07_mid_conversation_system_without_content_is_dropped(gateway: Gateway) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + _, received = _converse_cell( + gateway, wire, _chat(model, (_QUESTION_TURN, _ANSWERED_TURN, _NO_CONTENT_SYSTEM, _FOLLOW_UP_TURN)) + ) + assert "system" not in received, received + assert received["messages"] == [_CONVERSE_QUESTION, _CONVERSE_ANSWERED, _CONVERSE_FOLLOW_UP], received + + +def test_r08_assistant_without_content_between_two_user_turns_merges_them(gateway: Gateway) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + _, received = _converse_cell( + gateway, wire, _chat(model, (_QUESTION_TURN, _NO_CONTENT_ASSISTANT, _FOLLOW_UP_TURN)) + ) + assert received["messages"] == [{"role": "user", "content": [{"text": _QUESTION}, {"text": _FOLLOW_UP}]}], ( + received + ) + + +def test_r09_lone_user_with_empty_string_content_sends_no_converse_block(gateway: Gateway) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + _, received = _converse_cell(gateway, wire, _chat(model, (_EMPTY_CONTENT_USER,))) + assert received["messages"] == [], received + + +def test_r10_empty_null_and_missing_content_produce_byte_identical_converse_bodies(gateway: Gateway) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + identities: Final = tuple( + _chat_answer(_post_chat(gateway, _chat(model, (turn,)))) + for turn in (_EMPTY_CONTENT_USER, _NULL_CONTENT_USER, _NO_CONTENT_USER) + ) + assert len(set(identities)) == 3, identities + received: Final = wire.drain() + assert [unquote(request.target) for request in received] == [_CONVERSE_TARGET] * 3, received + assert len({request.body for request in received}) == 1, received + assert json.loads(received[0].body)["messages"] == [], received + for identity in identities: + assert _spend_row(identity)["status"] == "success" + + +@pytest.mark.parametrize("cell", _MODIFY_PARAMS_CELLS, ids=lambda cell: cell.row) +def test_r11_to_r15_modify_params_fills_the_missing_user_content( + cell: _ModifyParamsCell, modify_params_proxy: tuple[Gateway, Wire] +) -> None: + gateway, wire = modify_params_proxy + _, received = _converse_cell(gateway, wire, _chat(cell.model, cell.messages, **cell.extra)) + assert received["messages"] == list(cell.expected), received + assert "system" not in received, received + + +def test_r16_deployment_user_continue_message_fills_the_missing_content_without_modify_params(gateway: Gateway) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire, user_continue_message=_DEPLOYMENT_CONTINUE_MESSAGE) + _, received = _converse_cell(gateway, wire, _chat(model, (_NO_CONTENT_USER,))) + assert received["messages"] == [_continue_turn(_DEPLOYMENT_CONTINUE)], received + + +def test_r17_anthropic_sync_lone_user_without_content_is_rejected_before_any_peer_call(gateway: Gateway) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + with anthropic.Anthropic(base_url=_proxy_url(gateway), api_key=gateway.key, max_retries=0) as client: + with pytest.raises(anthropic.BadRequestError, match=_NO_NON_SYSTEM_MESSAGE): + client.messages.create(model=model, max_tokens=16, messages=[_NO_CONTENT_USER], extra_body=_EXTRA) + assert wire.drain() == () + + +async def test_r18_anthropic_async_stream_user_without_content_after_an_assistant_turn_keeps_the_earlier_turns( + gateway: Gateway, +) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + async with anthropic.AsyncAnthropic(base_url=_proxy_url(gateway), api_key=gateway.key, max_retries=0) as client: + async with client.messages.stream( + model=model, + max_tokens=16, + messages=[_QUESTION_TURN, _ANSWERED_TURN, _NO_CONTENT_USER], + extra_body=_EXTRA, + ) as stream: + events: Final = [event async for event in stream] + final: Final = await stream.get_final_message() + (started,) = tuple(event for event in events if event.type == "message_start") + assert final.content[0].text == _ANSWER, final + assert final.id == started.message.id, (final.id, started.message.id) + target, received = _only_received(wire) + assert target == _STREAM_TARGET, target + assert received["messages"] == [_CONVERSE_QUESTION, _CONVERSE_ANSWERED], received + assert _spend_row(final.id)["call_type"] == "anthropic_messages" + + +def test_r19_anthropic_native_invoke_forwards_the_turn_verbatim_and_relays_the_scripted_400(gateway: Gateway) -> None: + with ( + wire_server(lambda _: _scripted_error(400, "scripted invoke validation")) as wire, + gateway.scenario() as scenario, + ): + model: Final = scenario.model( + model=_INVOKE_MODEL, api_key=None, aws_bedrock_runtime_endpoint=wire.url, api_base=wire.url, **_AWS + ) + with anthropic.Anthropic(base_url=_proxy_url(gateway), api_key=gateway.key, max_retries=0) as client: + with pytest.raises(anthropic.BadRequestError, match="scripted invoke validation"): + client.messages.create(model=model, max_tokens=16, messages=[_NO_CONTENT_USER], extra_body=_EXTRA) + target, received = _only_received(wire) + assert target == _INVOKE_TARGET, target + assert received["messages"] == [_NO_CONTENT_USER], received + + +def test_r20_responses_lone_input_item_without_content_is_rejected_before_any_peer_call(gateway: Gateway) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + response: Final = gateway.request( + "POST", "/v1/responses", {"model": model, "input": [_NO_CONTENT_USER], **_EXTRA} + ) + assert response.status_code == 400, response.text + assert _NO_NON_SYSTEM_MESSAGE in response.text, response.text + assert wire.drain() == () + + +def test_r21_responses_stream_input_item_without_content_after_an_assistant_item_keeps_the_earlier_turns( + gateway: Gateway, +) -> None: + marker: Final = f"call-{uuid.uuid4().hex}" + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + lines: Final = _stream_lines( + gateway, + "/v1/responses", + { + "model": model, + "input": [_QUESTION_TURN, _ANSWERED_TURN, _NO_CONTENT_USER], + "stream": True, + "user": marker, + **_EXTRA, + }, + ) + events: Final = _sse_payloads(lines) + (completed,) = tuple(event for event in events if event["type"] == "response.completed") + assert completed["response"]["output"][0]["content"][0]["text"] == _ANSWER, lines + target, received = _only_received(wire) + assert target == _STREAM_TARGET, target + assert received["messages"] == [_CONVERSE_QUESTION, _CONVERSE_ANSWERED], received + (row,) = _success_rows(marker, 1) + assert row["call_type"] == "aresponses" and row["end_user"] == marker, row + assert row["request_id"] == _inner_response_id(str(completed["response"]["id"])), (row, completed) + + +def test_r22_passthrough_converse_forwards_a_message_without_content_verbatim(gateway: Gateway) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + deployment: Final = scenario.model( + model=f"bedrock/{_MODEL_ID}", api_base=wire.url, aws_bedrock_runtime_endpoint=wire.url, **_AWS + ) + response: Final = gateway.request( + "POST", f"/bedrock/model/{deployment}/converse", {"messages": [_NO_CONTENT_USER]} + ) + assert response.status_code == 200, response.text + assert response.content == _RESPONSE, response.text + (request,) = wire.drain() + assert unquote(request.target) == _CONVERSE_TARGET, request.target + assert json.loads(request.body)["messages"] == [_NO_CONTENT_USER], request.body + + +def test_r23_tool_turn_without_content_and_without_tools_is_neutralized_to_text(gateway: Gateway) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + _, received = _converse_cell(gateway, wire, _chat(model, (_QUESTION_TURN, _TOOL_CALL_TURN, _NO_CONTENT_TOOL))) + assert received["messages"] == [_CONVERSE_QUESTION, _NEUTRALIZED_TOOL_CALL, _NEUTRALIZED_TOOL_RESULT], received + assert "toolConfig" not in received, received + + +def test_s01_lone_user_with_integer_content_errors_before_any_peer_call(gateway: Gateway) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + response: Final = _post_chat(gateway, _chat(model, ({"role": "user", "content": 42},))) + assert response.status_code >= 400, response.text + assert "error" in response.json(), response.text + assert wire.drain() == () + + +def test_s02_lone_user_with_list_content_sends_the_text_block(gateway: Gateway) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + _, received = _converse_cell( + gateway, wire, _chat(model, ({"role": "user", "content": [{"type": "text", "text": _QUESTION}]},)) + ) + assert received["messages"] == [_CONVERSE_QUESTION], received + + +def test_s03_lone_user_with_a_five_kilobyte_string_reaches_the_peer_whole(gateway: Gateway) -> None: + text: Final = "k" * 5120 + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + _, received = _converse_cell(gateway, wire, _chat(model, ({"role": "user", "content": text},))) + assert received["messages"] == [{"role": "user", "content": [{"text": text}]}], received + + +def test_s04_duplicate_content_keys_in_the_raw_body_let_the_last_value_win(gateway: Gateway) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + raw: Final = ( + f'{{"model": "{model}", "messages": [{{"role": "user", "content": "first", "content": "second"}}],' + ' "max_tokens": 16, "num_retries": 0, "cache": {"no-cache": true}}' + ) + response: Final = gateway.client.post( + "/v1/chat/completions", + content=raw.encode(), + headers={**_auth(gateway), "content-type": "application/json"}, + ) + identity: Final = _chat_answer(response) + target, received = _only_received(wire) + assert target == _CONVERSE_TARGET, target + assert received["messages"] == [{"role": "user", "content": [{"text": "second"}]}], received + assert _spend_row(identity)["status"] == "success" + + +def test_s05_unauthenticated_content_less_request_never_reaches_the_peer(gateway: Gateway) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + response: Final = gateway.request( + "POST", "/v1/chat/completions", _chat(model, (_NO_CONTENT_USER,)), key="sk-integration-bogus" + ) + assert response.status_code == 401, response.text + assert wire.drain() == (), response.text + + +@pytest.mark.parametrize( + ("peer_status", "message", "expected"), + ( + (400, "ValidationException: scripted validation", 400), + (429, "ThrottlingException: scripted throttle", 429), + (500, "scripted outage", 503), + ), + ids=("s06", "s07", "s08"), +) +def test_s06_to_s08_peer_errors_on_a_content_less_turn_reach_the_caller_after_one_attempt( + gateway: Gateway, peer_status: int, message: str, expected: int +) -> None: + with wire_server(lambda _: _scripted_error(peer_status, message)) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + response: Final = _post_chat(gateway, _chat(model, (_NO_CONTENT_USER,))) + assert response.status_code == expected, response.text + assert message in response.text, response.text + assert len(wire.drain()) == 1, response.text + + +def test_s09_unknown_model_with_a_content_less_turn_never_reaches_the_peer(gateway: Gateway) -> None: + with wire_server(_converse_peer) as wire: + response: Final = _post_chat(gateway, _chat(f"integration-missing-{uuid.uuid4().hex}", (_NO_CONTENT_USER,))) + assert response.status_code in (400, 404), response.text + assert wire.drain() == (), response.text + + +def test_s10_a_deployment_continue_message_without_content_adds_no_converse_block(gateway: Gateway) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire, user_continue_message={"role": "user"}) + _, received = _converse_cell(gateway, wire, _chat(model, (_NO_CONTENT_USER,))) + assert received["messages"] == [], received + + +def test_s11_a_deployment_continue_message_given_as_a_string_errors_in_the_body_and_leaves_the_proxy_serving( + gateway: Gateway, +) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + broken: Final = _converse_deployment(scenario, wire, user_continue_message=_DEPLOYMENT_CONTINUE) + healthy: Final = _converse_deployment(scenario, wire) + response: Final = _post_chat(gateway, _chat(broken, (_NO_CONTENT_USER,))) + assert response.status_code >= 400, response.text + assert "error" in response.json(), response.text + assert wire.drain() == () + _, received = _converse_cell(gateway, wire, _chat(healthy, (_QUESTION_TURN,))) + assert received["messages"] == [_CONVERSE_QUESTION], received + + +def test_e01_the_same_content_less_request_twice_with_no_cache_hits_the_peer_twice(gateway: Gateway) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + first: Final = _chat_answer(_post_chat(gateway, _chat(model, (_NO_CONTENT_USER,)))) + second: Final = _chat_answer(_post_chat(gateway, _chat(model, (_NO_CONTENT_USER,)))) + assert first != second + assert len(wire.drain()) == 2 + assert _spend_row(first)["status"] == "success" + assert _spend_row(second)["status"] == "success" + + +def test_e02_the_same_content_less_request_twice_is_served_from_the_response_cache(gateway: Gateway) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + body: Final = _chat(model, (_NO_CONTENT_USER,), cached=True) + first: Final = _chat_answer(_post_chat(gateway, body)) + second: Final = _chat_answer(_post_chat(gateway, body)) + assert first == second + assert len(wire.drain()) == 1 + assert _spend_row(first)["status"] == "success" + cache_hits: Final = eventually( + lambda: read_rows( + 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id LIKE %s', (f"{first}_cache_hit%",) + ), + lambda rows: len(rows) >= 1, + seconds=70, + ) + assert len(cache_hits) == 1, cache_hits + + +def _model_id(gateway: Gateway, name: str) -> str: + entries: Final = gateway.get("/model/info")["data"] + assert isinstance(entries, list), entries + (identity,) = ( + string_value(object_value(object_value(entry)["model_info"])["id"]) + for entry in entries + if object_value(entry)["model_name"] == name + ) + return identity + + +def _content_less_bodies(gateway: Gateway, wire: Wire, model: str, count: int) -> tuple[JsonValue, ...]: + identities: Final = tuple( + _chat_answer(_post_chat(gateway, _chat(model, (_NO_CONTENT_USER,)))) for _ in range(count) + ) + received: Final = wire.drain() + assert len(received) == len(identities), (identities, received) + bodies: Final = tuple(_JSON.validate_python(json.loads(request.body))["messages"] for request in received) + assert all(body in ([], [_continue_turn(_DEPLOYMENT_CONTINUE)]) for body in bodies), bodies + return bodies + + +@pytest.mark.timeout(180) +def test_e03_updating_the_deployment_continue_message_under_traffic_never_breaks_a_content_less_turn( + gateway: Gateway, +) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + assert _content_less_bodies(gateway, wire, model, 4) == ([],) * 4 + gateway.post( + "/model/update", + { + "model_info": {"id": _model_id(gateway, model)}, + "litellm_params": {"user_continue_message": _DEPLOYMENT_CONTINUE_MESSAGE}, + }, + ) + settled: Final = eventually( + lambda: _content_less_bodies(gateway, wire, model, 8), + lambda bodies: all(body == [_continue_turn(_DEPLOYMENT_CONTINUE)] for body in bodies), + seconds=90, + ) + assert len(settled) == 8, settled + + +@pytest.mark.parametrize( + ("continue_message", "expected"), + ((None, []), ({}, [_continue_turn(_DEFAULT_CONTINUE)])), + ids=("e04", "e05"), +) +def test_e04_e05_a_null_continue_message_means_absent_and_an_empty_one_means_the_default( + gateway: Gateway, continue_message: JsonValue, expected: list[JsonValue] +) -> None: + with wire_server(_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire, user_continue_message=continue_message) + _, received = _converse_cell(gateway, wire, _chat(model, (_NO_CONTENT_USER,))) + assert received["messages"] == expected, received + + +def _path(endpoint: Endpoint) -> str: + match endpoint: + case "chat": + return "/v1/chat/completions" + case "responses": + return "/v1/responses" + + +def _calls(marker: str, endpoint: Endpoint, stream: bool, indexes: range) -> tuple[_Call, ...]: + return tuple(_Call(endpoint, stream, f"{marker}-{index}", index) for index in indexes) + + +def _burst_body(model: str, call: _Call) -> dict[str, JsonValue]: + turns: Final = ({"role": "user", "content": f"{_QUESTION} {call.user}"}, _ANSWERED_TURN, _NO_CONTENT_USER) + if call.endpoint == "chat": + return {**_chat(model, turns, stream=call.stream), "user": call.user} + return {"model": model, "input": [dict(turn) for turn in turns], "stream": call.stream, "user": call.user, **_EXTRA} + + +def _call_index(request: Request) -> int: + found: Final = _CALL_INDEX.search(request.body.decode()) + assert found is not None, request.body + return int(found.group(1)) + + +async def _send(client: httpx.AsyncClient, key: str, model: str, call: _Call) -> _Served: + async with client.stream( + "POST", _path(call.endpoint), json=_burst_body(model, call), headers={"Authorization": f"Bearer {key}"} + ) as response: + raw: Final = await response.aread() + return _Served(call=call, status=response.status_code, text=raw.decode()) + + +async def _burst( + base_url: str, key: str, model: str, calls: tuple[_Call, ...], *, tolerate_transport_errors: bool = False +) -> tuple[_Served, ...]: + async with httpx.AsyncClient(base_url=base_url, timeout=60, trust_env=False) as client: + results: Final = await asyncio.gather( + *(_send(client, key, model, call) for call in calls), return_exceptions=tolerate_transport_errors + ) + for result in results: + assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result) + return tuple(result for result in results if isinstance(result, _Served)) + + +def _inner_response_id(identity: str) -> str: + managed: Final = decrypt_if_encrypted_with(identity.removeprefix("resp_"), _SIGNING_KEY) + assert managed is not None, identity + issued: Final = managed.split(";", 1)[0].rsplit("response_id:", 1)[1] + decoded: Final = base64.b64decode(issued.removeprefix("resp_")).decode() + return decoded.rsplit("response_id:", 1)[1] + + +def _served_id(served: _Served) -> str: + if served.call.stream: + (identity,) = {chunk["id"] for chunk in _sse_payloads(served.text.splitlines())} + return identity + return json.loads(served.text)["id"] + + +def _spend_row_id(served: _Served) -> str: + identity: Final = _served_id(served) + return identity if served.call.endpoint == "chat" else _inner_response_id(identity) + + +def _answered(served: Iterable[_Served]) -> frozenset[str]: + ids: Final = tuple(_spend_row_id(item) for item in served) + assert len(set(ids)) == len(ids), ids + return frozenset(ids) + + +def _open_peer_connections(pid: int, peer_url: str) -> int: + port: Final = urlsplit(peer_url).port + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port + ) + + +async def test_c01_a_mixed_burst_of_content_less_calls_survives_a_peer_outage_window(gateway: Gateway) -> None: + marker: Final = f"call-{uuid.uuid4().hex}" + calls: Final = ( + *_calls(marker, "chat", False, range(0, 10)), + *_calls(marker, "chat", True, range(10, 20)), + *_calls(marker, "responses", False, range(20, 30)), + ) + + def respond(request: Request) -> Reply: + if _call_index(request) % 3 == 1: + return _scripted_error(500, "scripted outage") + return _converse_peer(request) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = _converse_deployment(scenario, wire) + served: Final = await _burst(_proxy_url(gateway), gateway.key, model, calls) + assert len(served) == 30 + failed: Final = tuple(item for item in served if item.call.index % 3 == 1) + answered: Final = tuple(item for item in served if item.call.index % 3 != 1) + assert len(failed) == 10 and len(answered) == 20 + for item in failed: + assert item.status == 503 and "scripted outage" in item.text, (item.call, item.status, item.text) + for item in answered: + assert item.status == 200 and _ANSWER in item.text, (item.call, item.status, item.text) + identities: Final = _answered(answered) + assert len(identities) == 20 + assert {row["request_id"] for row in _success_rows(marker, 20)} == identities + assert len(wire.drain()) == 30 + + +@pytest.mark.timeout(180) +async def test_c02_worker_sigkill_mid_burst_leaves_the_sibling_serving_content_less_turns( + gateway: Gateway, tmp_path: Path +) -> None: + marker: Final = f"call-{uuid.uuid4().hex}" + again: Final = f"call-{uuid.uuid4().hex}" + calls: Final = _calls(marker, "chat", False, range(20)) + release: Final = threading.Event() + held_indexes: Final[SimpleQueue[int]] = SimpleQueue() + + def held(request: Request) -> Reply: + held_indexes.put(_call_index(request)) + assert release.wait(timeout=60), "The burst was never released" + return _converse_peer(request) + + with wire_server(held) as wire: + config: Final = _owned_config(wire, tmp_path, modify_params=False) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + candidate: Final = owned.gateway + workers: Final = eventually( + lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=30, + ) + burst: Final = asyncio.create_task( + _burst(_proxy_url(candidate), candidate.key, _PLAIN, calls, tolerate_transport_errors=True) + ) + await asyncio.to_thread(eventually, held_indexes.qsize, lambda size: size == 20, 60) + held_by: Final = MappingProxyType({pid: _open_peer_connections(pid, wire.url) for pid in workers}) + assert sum(held_by.values()) == 20, held_by + victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__) + psutil.Process(victim_pid).send_signal(signal.SIGKILL) + release.set() + served: Final = await burst + assert held_by[survivor_pid] >= 10, held_by + assert len(served) == held_by[survivor_pid], (held_by, len(served)) + for item in served: + assert item.status == 200 and _ANSWER in item.text, (item.call, item.status, item.text) + eventually( + lambda: len(_STARTED_WORKER.findall(owned.log.read_text())), lambda count: count == 3, seconds=60 + ) + follow_up: Final = await _burst( + _proxy_url(candidate), candidate.key, _PLAIN, _calls(again, "chat", False, range(6)) + ) + assert len(follow_up) == 6 + for item in follow_up: + assert item.status == 200 and _ANSWER in item.text, (item.call, item.status, item.text) + assert len(wire.drain()) == 26 + assert {row["request_id"] for row in _success_rows(marker, len(served))} == _answered(served) + assert {row["request_id"] for row in _success_rows(again, 6)} == _answered(follow_up) + + +@pytest.mark.timeout(180) +async def test_c03_proxy_terminated_mid_burst_lands_every_answered_content_less_call_at_most_once( + gateway: Gateway, tmp_path: Path +) -> None: + marker: Final = f"call-{uuid.uuid4().hex}" + calls: Final = _calls(marker, "chat", False, range(12)) + release: Final = threading.Event() + held_indexes: Final[SimpleQueue[int]] = SimpleQueue() + + def held(request: Request) -> Reply: + held_indexes.put(_call_index(request)) + assert release.wait(timeout=60), "The burst was never released" + return _converse_peer(request) + + with wire_server(held) as wire: + config: Final = _owned_config(wire, tmp_path, modify_params=False) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=1) as owned: + candidate: Final = owned.gateway + burst: Final = asyncio.create_task( + _burst(_proxy_url(candidate), candidate.key, _PLAIN, calls, tolerate_transport_errors=True) + ) + await asyncio.to_thread(eventually, held_indexes.qsize, lambda size: size == 12, 60) + owned.process.terminate() + release.set() + served: Final = await burst + eventually(owned.process.poll, lambda code: code is not None, seconds=60) + answered: Final = _answered(item for item in served if item.status == 200) + assert len(served) <= 12 + landed: Final = tuple( + row["request_id"] + for row in read_rows( + 'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE end_user LIKE %s AND status=%s', + (f"{marker}%", "success"), + ) + ) + assert len(landed) == len(set(landed)), landed + stray: Final = set(landed) - answered + assert len(stray) <= 12 - len(answered), (landed, answered) + assert len(wire.drain()) == 12 diff --git a/tests/integration/spend/_daily_activity_fixtures.py b/tests/integration/spend/_daily_activity_fixtures.py new file mode 100644 index 00000000000..9cde3ed0fb9 --- /dev/null +++ b/tests/integration/spend/_daily_activity_fixtures.py @@ -0,0 +1,341 @@ +from itertools import product +from typing import Final + +import psycopg +from psycopg import sql + +_TABLE_NAMES: Final = ( + "LiteLLM_DailyUserSpend", + "LiteLLM_DailyTeamSpend", + "LiteLLM_VerificationToken", + "LiteLLM_DeletedVerificationToken", + "LiteLLM_UserTable", + "LiteLLM_TeamTable", +) + +_TAG_KEY_MEMBERSHIPS: Final = ( + ("tag-a", "entity-key-0"), + ("tag-a", "entity-key-1"), + ("tag-a", "entity-key-2"), + ("tag-a", "entity-key-3"), + ("tag-a", "entity-key-4"), + ("tag-b", "entity-key-0"), + ("tag-b", "entity-key-1"), + ("tag-b", "entity-key-5"), + ("tag-c", "entity-key-2"), + ("tag-c", "entity-key-3"), + ("tag-c", "entity-key-6"), + ("tag-c", "entity-key-7"), + ("tag-d", "entity-key-4"), + ("tag-d", "entity-key-5"), + ("tag-d", "entity-key-6"), + ("tag-d", "entity-key-7"), +) +_TAG_ACTIVITY_DATES: Final = ("2026-06-01", "2026-06-02") + + +def seed_daily_activity_fixture(connection: psycopg.Connection, *, schema: str, ptu_sentinel_api_key: str) -> None: + daily_user_table: Final = sql.Identifier(schema, "LiteLLM_DailyUserSpend") + daily_team_table: Final = sql.Identifier(schema, "LiteLLM_DailyTeamSpend") + verification_token_table: Final = sql.Identifier(schema, "LiteLLM_VerificationToken") + deleted_token_table: Final = sql.Identifier(schema, "LiteLLM_DeletedVerificationToken") + user_table: Final = sql.Identifier(schema, "LiteLLM_UserTable") + team_table: Final = sql.Identifier(schema, "LiteLLM_TeamTable") + keys: Final = ( + ("key-a", "model-popular", 100.0, 2, 1), + ("key-b", "model-popular", 90.0, 3, 1), + ("key-c", "model-popular", 80.0, 4, 1), + ("key-target", "model-target", 1.0, 5, 2), + ("key-cache", "model-cache", 2.0, 1000, 1), + ) + user_rows: Final = tuple( + ( + f"user-row-{index}", + "user-1", + "2026-06-01", + api_key, + model, + "", + "provider-a", + None, + "/v1/chat/completions", + prompt_tokens, + 2, + cache_read_tokens, + 0, + spend, + 1, + 1, + 0, + "2026-06-01 12:00:00", + ) + for index, (api_key, model, spend, prompt_tokens, cache_read_tokens) in enumerate(keys) + ) + team_rows: Final = tuple( + ( + f"team-row-{index}", + "team-1", + "2026-06-01", + api_key, + model, + "", + "provider-a", + None, + "/v1/chat/completions", + prompt_tokens, + 2, + cache_read_tokens, + 0, + spend, + 1, + 1, + 0, + 0.0, + "2026-06-01 12:00:00", + ) + for index, (api_key, model, spend, prompt_tokens, cache_read_tokens) in enumerate(keys) + ) + sentinel_user_row: Final = ( + "user-row-ptu", + "user-1", + "2026-06-01", + ptu_sentinel_api_key, + "model-ptu", + "", + "provider-a", + None, + "/v1/chat/completions", + 0, + 0, + 0, + 0, + 1000.0, + 0, + 0, + 0, + "2026-06-01 12:00:00", + ) + sentinel_team_row: Final = ( + "team-row-ptu", + "team-1", + "2026-06-01", + ptu_sentinel_api_key, + "model-ptu", + "", + "provider-a", + None, + "/v1/chat/completions", + 0, + 0, + 0, + 0, + 1000.0, + 0, + 0, + 0, + 42.0, + "2026-06-01 12:00:00", + ) + with connection.cursor() as cursor: + for table_name in _TABLE_NAMES: + cursor.execute( + sql.SQL("CREATE TABLE {} (LIKE {} INCLUDING DEFAULTS INCLUDING CONSTRAINTS)").format( + sql.Identifier(schema, table_name), + sql.Identifier(table_name), + ) + ) + cursor.executemany( + sql.SQL(""" + INSERT INTO {} + (id, user_id, date, api_key, model, model_group, custom_llm_provider, + mcp_namespaced_tool_name, endpoint, prompt_tokens, completion_tokens, + cache_read_input_tokens, cache_creation_input_tokens, spend, api_requests, + successful_requests, failed_requests, updated_at) + VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s) + """).format(daily_user_table), + (*user_rows, sentinel_user_row), + ) + cursor.executemany( + sql.SQL(""" + INSERT INTO {} + (id, team_id, date, api_key, model, model_group, custom_llm_provider, + mcp_namespaced_tool_name, endpoint, prompt_tokens, completion_tokens, + cache_read_input_tokens, cache_creation_input_tokens, spend, api_requests, + successful_requests, failed_requests, ptu_flat_cost, updated_at) + VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s) + """).format(daily_team_table), + (*team_rows, sentinel_team_row), + ) + cursor.executemany( + sql.SQL( + "INSERT INTO {} (token, key_alias, team_id, user_id, metadata, models) VALUES (%s, %s, %s, %s, %s, %s)" + ).format(verification_token_table), + ( + ("key-a", "alias-a", "team-1", "user-1", '{"tags": ["blue", "gold"]}', []), + ("key-b", "alias-b", "team-1", "user-1", '{"tags": []}', []), + ("key-c", "alias-c", "team-1", "user-1", '{"tags": []}', []), + ("key-cache", "alias-cache", "team-1", "user-1", '{"tags": []}', []), + ), + ) + cursor.executemany( + sql.SQL(""" + INSERT INTO {} + (id, token, key_alias, team_id, user_id, metadata, models, deleted_at) + VALUES (%s, %s, %s, %s, %s, %s, %s, %s) + """).format(deleted_token_table), + ( + ("deleted-old", "key-target", "older-target", "team-1", "user-1", '{"tags": []}', [], "2026-06-01"), + ( + "deleted-new", + "key-target", + "deleted-target", + "team-1", + "user-1", + '{"tags": ["archived"]}', + [], + "2026-06-02", + ), + ), + ) + cursor.execute( + sql.SQL("INSERT INTO {} (user_id, user_email, models) VALUES (%s, %s, %s)").format(user_table), + ("user-1", "user@example.com", []), + ) + cursor.execute( + sql.SQL("INSERT INTO {} (team_id, team_alias, admins, members, models) VALUES (%s, %s, %s, %s, %s)").format( + team_table + ), + ("team-1", "Usage Team", [], [], []), + ) + connection.commit() + + +def seed_daily_tag_activity_fixture(connection: psycopg.Connection, *, schema: str) -> None: + tag_table: Final = sql.Identifier(schema, "LiteLLM_DailyTagSpend") + rows: Final = tuple( + ( + f"tag-rollup-{row_index}", + tag, + date, + api_key, + "entity-rollup-model", + "", + "provider-a", + None, + "/v1/chat/completions", + row_index + 1, + float(row_index + 1), + row_index % 5 + 1, + f"{date} 12:00:00", + ) + for row_index, (date, (tag, api_key)) in enumerate(product(_TAG_ACTIVITY_DATES, _TAG_KEY_MEMBERSHIPS)) + ) + with connection.cursor() as cursor: + cursor.execute( + sql.SQL("CREATE TABLE {} (LIKE {} INCLUDING DEFAULTS INCLUDING CONSTRAINTS)").format( + tag_table, + sql.Identifier("LiteLLM_DailyTagSpend"), + ) + ) + cursor.executemany( + sql.SQL(""" + INSERT INTO {} + (id, tag, date, api_key, model, model_group, custom_llm_provider, + mcp_namespaced_tool_name, endpoint, prompt_tokens, spend, api_requests, updated_at) + VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s) + """).format(tag_table), + rows, + ) + connection.commit() + + +def seed_daily_tag_float_tie_fixture(connection: psycopg.Connection, *, schema: str) -> None: + tag_table: Final = sql.Identifier(schema, "LiteLLM_DailyTagSpend") + key_spends: Final = ( + ("key-z", 0.1), + ("key-z", 0.2), + ("key-z", 0.3), + ("key-a", 0.3), + ("key-a", 0.2), + ("key-a", 0.1), + ) + rows: Final = ( + ( + f"float-tie-{row_index}", + "tag-float-tie", + "2026-06-01", + api_key, + "float-tie-model", + "", + "provider-a", + None, + "/v1/chat/completions", + 1, + spend, + 1, + "2026-06-01 12:00:00", + ) + for row_index, (api_key, spend) in enumerate(key_spends, start=1) + ) + with connection.cursor() as cursor: + cursor.execute( + sql.SQL("CREATE TABLE {} (LIKE {} INCLUDING DEFAULTS INCLUDING CONSTRAINTS)").format( + tag_table, + sql.Identifier("LiteLLM_DailyTagSpend"), + ) + ) + cursor.executemany( + sql.SQL(""" + INSERT INTO {} + (id, tag, date, api_key, model, model_group, custom_llm_provider, + mcp_namespaced_tool_name, endpoint, prompt_tokens, spend, api_requests, updated_at) + VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s) + """).format(tag_table), + rows, + ) + connection.commit() + + +def seed_daily_team_unassigned_fixture( + connection: psycopg.Connection, *, schema: str, ptu_sentinel_api_key: str +) -> None: + team_table: Final = sql.Identifier(schema, "LiteLLM_DailyTeamSpend") + rows: Final = ( + ("unassigned-null", None, "key-unassigned-null", 3.0, 0.0), + ("unassigned-empty", "", "key-unassigned-empty", 7.0, 0.0), + ("unassigned-ptu", None, ptu_sentinel_api_key, 13.0, 13.0), + ) + with connection.cursor() as cursor: + cursor.executemany( + sql.SQL(""" + INSERT INTO {} + (id, team_id, date, api_key, model, model_group, custom_llm_provider, + mcp_namespaced_tool_name, endpoint, prompt_tokens, spend, api_requests, ptu_flat_cost, updated_at) + VALUES (%s, %s, '2026-06-03', %s, 'model-a', '', 'provider-a', NULL, '/v1/chat/completions', + 1, %s, 1, %s, '2026-06-03 12:00:00') + """).format(team_table), + rows, + ) + connection.commit() + + +def seed_daily_team_exclusion_fixture(connection: psycopg.Connection, *, schema: str) -> None: + team_table: Final = sql.Identifier(schema, "LiteLLM_DailyTeamSpend") + rows: Final = ( + ("exclusion-null", None, "key-excluded-null", 3.0), + ("exclusion-empty", "", "key-excluded-empty", 7.0), + ("exclusion-dashboard", "litellm-dashboard", "key-excluded-dashboard", 11.0), + ("exclusion-normal", "team-normal", "key-excluded-normal", 13.0), + ) + with connection.cursor() as cursor: + cursor.executemany( + sql.SQL(""" + INSERT INTO {} + (id, team_id, date, api_key, model, model_group, custom_llm_provider, + mcp_namespaced_tool_name, endpoint, prompt_tokens, spend, api_requests, updated_at) + VALUES (%s, %s, '2026-06-04', %s, 'model-a', '', 'provider-a', NULL, '/v1/chat/completions', + 1, %s, 1, '2026-06-04 12:00:00') + """).format(team_table), + rows, + ) + connection.commit() diff --git a/tests/integration/spend/fixtures/daily_activity_team.json b/tests/integration/spend/fixtures/daily_activity_team.json new file mode 100644 index 00000000000..d833a695ede --- /dev/null +++ b/tests/integration/spend/fixtures/daily_activity_team.json @@ -0,0 +1 @@ +{"results":[{"date":"2026-06-01","metrics":{"spend":1273.0,"flat_cost":0.0,"prompt_tokens":1014,"completion_tokens":10,"cache_read_input_tokens":6,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1024,"successful_requests":5,"failed_requests":0,"api_requests":5,"total_response_time_ms":0,"timed_requests":0},"breakdown":{"mcp_servers":{},"models":{"model-ptu":{"metrics":{"spend":1000.0,"flat_cost":0.0,"prompt_tokens":0,"completion_tokens":0,"cache_read_input_tokens":0,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":0,"successful_requests":0,"failed_requests":0,"api_requests":0,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{}},"model-cache":{"metrics":{"spend":2.0,"flat_cost":0.0,"prompt_tokens":1000,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1002,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{"key-cache":{"metrics":{"spend":2.0,"flat_cost":0.0,"prompt_tokens":1000,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1002,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-cache","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}}}},"model-target":{"metrics":{"spend":1.0,"flat_cost":0.0,"prompt_tokens":5,"completion_tokens":2,"cache_read_input_tokens":2,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":7,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{"key-target":{"metrics":{"spend":1.0,"flat_cost":0.0,"prompt_tokens":5,"completion_tokens":2,"cache_read_input_tokens":2,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":7,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"deleted-target","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":false}}}},"model-popular":{"metrics":{"spend":270.0,"flat_cost":0.0,"prompt_tokens":9,"completion_tokens":6,"cache_read_input_tokens":3,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":15,"successful_requests":3,"failed_requests":0,"api_requests":3,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{"key-a":{"metrics":{"spend":100.0,"flat_cost":0.0,"prompt_tokens":2,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":4,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-a","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-b":{"metrics":{"spend":90.0,"flat_cost":0.0,"prompt_tokens":3,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":5,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-b","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-c":{"metrics":{"spend":80.0,"flat_cost":0.0,"prompt_tokens":4,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":6,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-c","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}}}}},"model_groups":{"model-ptu":{"metrics":{"spend":1000.0,"flat_cost":0.0,"prompt_tokens":0,"completion_tokens":0,"cache_read_input_tokens":0,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":0,"successful_requests":0,"failed_requests":0,"api_requests":0,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{}},"model-cache":{"metrics":{"spend":2.0,"flat_cost":0.0,"prompt_tokens":1000,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1002,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{"key-cache":{"metrics":{"spend":2.0,"flat_cost":0.0,"prompt_tokens":1000,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1002,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-cache","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}}}},"model-target":{"metrics":{"spend":1.0,"flat_cost":0.0,"prompt_tokens":5,"completion_tokens":2,"cache_read_input_tokens":2,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":7,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{"key-target":{"metrics":{"spend":1.0,"flat_cost":0.0,"prompt_tokens":5,"completion_tokens":2,"cache_read_input_tokens":2,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":7,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"deleted-target","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":false}}}},"model-popular":{"metrics":{"spend":270.0,"flat_cost":0.0,"prompt_tokens":9,"completion_tokens":6,"cache_read_input_tokens":3,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":15,"successful_requests":3,"failed_requests":0,"api_requests":3,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{"key-c":{"metrics":{"spend":80.0,"flat_cost":0.0,"prompt_tokens":4,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":6,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-c","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-a":{"metrics":{"spend":100.0,"flat_cost":0.0,"prompt_tokens":2,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":4,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-a","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-b":{"metrics":{"spend":90.0,"flat_cost":0.0,"prompt_tokens":3,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":5,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-b","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}}}}},"providers":{"provider-a":{"metrics":{"spend":1273.0,"flat_cost":0.0,"prompt_tokens":1014,"completion_tokens":10,"cache_read_input_tokens":6,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1024,"successful_requests":5,"failed_requests":0,"api_requests":5,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{"key-target":{"metrics":{"spend":1.0,"flat_cost":0.0,"prompt_tokens":5,"completion_tokens":2,"cache_read_input_tokens":2,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":7,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"deleted-target","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":false}},"key-cache":{"metrics":{"spend":2.0,"flat_cost":0.0,"prompt_tokens":1000,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1002,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-cache","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-b":{"metrics":{"spend":90.0,"flat_cost":0.0,"prompt_tokens":3,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":5,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-b","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-c":{"metrics":{"spend":80.0,"flat_cost":0.0,"prompt_tokens":4,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":6,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-c","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-a":{"metrics":{"spend":100.0,"flat_cost":0.0,"prompt_tokens":2,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":4,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-a","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}}}}},"endpoints":{"/v1/chat/completions":{"metrics":{"spend":1273.0,"flat_cost":0.0,"prompt_tokens":1014,"completion_tokens":10,"cache_read_input_tokens":6,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1024,"successful_requests":5,"failed_requests":0,"api_requests":5,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{"key-c":{"metrics":{"spend":80.0,"flat_cost":0.0,"prompt_tokens":4,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":6,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-c","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-target":{"metrics":{"spend":1.0,"flat_cost":0.0,"prompt_tokens":5,"completion_tokens":2,"cache_read_input_tokens":2,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":7,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"deleted-target","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":false}},"key-b":{"metrics":{"spend":90.0,"flat_cost":0.0,"prompt_tokens":3,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":5,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-b","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-a":{"metrics":{"spend":100.0,"flat_cost":0.0,"prompt_tokens":2,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":4,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-a","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-cache":{"metrics":{"spend":2.0,"flat_cost":0.0,"prompt_tokens":1000,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1002,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-cache","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}}}}},"api_keys":{"key-a":{"metrics":{"spend":100.0,"flat_cost":0.0,"prompt_tokens":2,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":4,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-a","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-b":{"metrics":{"spend":90.0,"flat_cost":0.0,"prompt_tokens":3,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":5,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-b","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-c":{"metrics":{"spend":80.0,"flat_cost":0.0,"prompt_tokens":4,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":6,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-c","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-cache":{"metrics":{"spend":2.0,"flat_cost":0.0,"prompt_tokens":1000,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1002,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-cache","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-target":{"metrics":{"spend":1.0,"flat_cost":0.0,"prompt_tokens":5,"completion_tokens":2,"cache_read_input_tokens":2,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":7,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"deleted-target","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":false}}},"entities":{"team-1":{"metrics":{"spend":1273.0,"flat_cost":0.0,"prompt_tokens":1014,"completion_tokens":10,"cache_read_input_tokens":6,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1024,"successful_requests":5,"failed_requests":0,"api_requests":5,"total_response_time_ms":0,"timed_requests":0},"metadata":{"team_alias":"Usage Team"},"api_key_breakdown":{"key-a":{"metrics":{"spend":100.0,"flat_cost":0.0,"prompt_tokens":2,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":4,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-a","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-b":{"metrics":{"spend":90.0,"flat_cost":0.0,"prompt_tokens":3,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":5,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-b","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-c":{"metrics":{"spend":80.0,"flat_cost":0.0,"prompt_tokens":4,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":6,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-c","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-cache":{"metrics":{"spend":2.0,"flat_cost":0.0,"prompt_tokens":1000,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1002,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-cache","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-target":{"metrics":{"spend":1.0,"flat_cost":0.0,"prompt_tokens":5,"completion_tokens":2,"cache_read_input_tokens":2,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":7,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"deleted-target","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":false}}}}}}}],"metadata":{"total_spend":1273.0,"total_flat_cost":0.0,"total_prompt_tokens":1014,"total_completion_tokens":10,"total_tokens":1024,"total_api_requests":5,"total_successful_requests":5,"total_failed_requests":0,"total_cache_read_input_tokens":6,"total_cache_creation_input_tokens":0,"total_compression_saved_tokens":0,"total_compression_savings_spend":0.0,"total_prompt_caching_savings_spend":0.0,"total_gateway_injected_caching_savings_spend":0.0,"total_autorouter_savings_spend":0.0,"total_response_time_ms":0,"total_timed_requests":0,"page":1,"total_pages":1,"has_more":false,"api_key_limit":100,"total_api_keys":5,"entity_total_api_keys":{"team-1":5}}} diff --git a/tests/integration/spend/fixtures/daily_activity_user.json b/tests/integration/spend/fixtures/daily_activity_user.json new file mode 100644 index 00000000000..36701c20d9b --- /dev/null +++ b/tests/integration/spend/fixtures/daily_activity_user.json @@ -0,0 +1 @@ +{"results":[{"date":"2026-06-01","metrics":{"spend":1273.0,"flat_cost":0.0,"prompt_tokens":1014,"completion_tokens":10,"cache_read_input_tokens":6,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1024,"successful_requests":5,"failed_requests":0,"api_requests":5,"total_response_time_ms":0,"timed_requests":0},"breakdown":{"mcp_servers":{},"models":{"model-ptu":{"metrics":{"spend":1000.0,"flat_cost":0.0,"prompt_tokens":0,"completion_tokens":0,"cache_read_input_tokens":0,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":0,"successful_requests":0,"failed_requests":0,"api_requests":0,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{}},"model-cache":{"metrics":{"spend":2.0,"flat_cost":0.0,"prompt_tokens":1000,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1002,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{"key-cache":{"metrics":{"spend":2.0,"flat_cost":0.0,"prompt_tokens":1000,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1002,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-cache","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}}}},"model-target":{"metrics":{"spend":1.0,"flat_cost":0.0,"prompt_tokens":5,"completion_tokens":2,"cache_read_input_tokens":2,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":7,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{"key-target":{"metrics":{"spend":1.0,"flat_cost":0.0,"prompt_tokens":5,"completion_tokens":2,"cache_read_input_tokens":2,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":7,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"deleted-target","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":false}}}},"model-popular":{"metrics":{"spend":270.0,"flat_cost":0.0,"prompt_tokens":9,"completion_tokens":6,"cache_read_input_tokens":3,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":15,"successful_requests":3,"failed_requests":0,"api_requests":3,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{"key-a":{"metrics":{"spend":100.0,"flat_cost":0.0,"prompt_tokens":2,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":4,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-a","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-b":{"metrics":{"spend":90.0,"flat_cost":0.0,"prompt_tokens":3,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":5,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-b","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-c":{"metrics":{"spend":80.0,"flat_cost":0.0,"prompt_tokens":4,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":6,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-c","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}}}}},"model_groups":{"model-ptu":{"metrics":{"spend":1000.0,"flat_cost":0.0,"prompt_tokens":0,"completion_tokens":0,"cache_read_input_tokens":0,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":0,"successful_requests":0,"failed_requests":0,"api_requests":0,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{}},"model-cache":{"metrics":{"spend":2.0,"flat_cost":0.0,"prompt_tokens":1000,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1002,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{"key-cache":{"metrics":{"spend":2.0,"flat_cost":0.0,"prompt_tokens":1000,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1002,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-cache","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}}}},"model-target":{"metrics":{"spend":1.0,"flat_cost":0.0,"prompt_tokens":5,"completion_tokens":2,"cache_read_input_tokens":2,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":7,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{"key-target":{"metrics":{"spend":1.0,"flat_cost":0.0,"prompt_tokens":5,"completion_tokens":2,"cache_read_input_tokens":2,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":7,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"deleted-target","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":false}}}},"model-popular":{"metrics":{"spend":270.0,"flat_cost":0.0,"prompt_tokens":9,"completion_tokens":6,"cache_read_input_tokens":3,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":15,"successful_requests":3,"failed_requests":0,"api_requests":3,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{"key-c":{"metrics":{"spend":80.0,"flat_cost":0.0,"prompt_tokens":4,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":6,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-c","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-a":{"metrics":{"spend":100.0,"flat_cost":0.0,"prompt_tokens":2,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":4,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-a","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-b":{"metrics":{"spend":90.0,"flat_cost":0.0,"prompt_tokens":3,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":5,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-b","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}}}}},"providers":{"provider-a":{"metrics":{"spend":1273.0,"flat_cost":0.0,"prompt_tokens":1014,"completion_tokens":10,"cache_read_input_tokens":6,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1024,"successful_requests":5,"failed_requests":0,"api_requests":5,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{"key-target":{"metrics":{"spend":1.0,"flat_cost":0.0,"prompt_tokens":5,"completion_tokens":2,"cache_read_input_tokens":2,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":7,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"deleted-target","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":false}},"key-cache":{"metrics":{"spend":2.0,"flat_cost":0.0,"prompt_tokens":1000,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1002,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-cache","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-b":{"metrics":{"spend":90.0,"flat_cost":0.0,"prompt_tokens":3,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":5,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-b","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-c":{"metrics":{"spend":80.0,"flat_cost":0.0,"prompt_tokens":4,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":6,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-c","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-a":{"metrics":{"spend":100.0,"flat_cost":0.0,"prompt_tokens":2,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":4,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-a","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}}}}},"endpoints":{"/v1/chat/completions":{"metrics":{"spend":1273.0,"flat_cost":0.0,"prompt_tokens":1014,"completion_tokens":10,"cache_read_input_tokens":6,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1024,"successful_requests":5,"failed_requests":0,"api_requests":5,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{"key-c":{"metrics":{"spend":80.0,"flat_cost":0.0,"prompt_tokens":4,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":6,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-c","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-target":{"metrics":{"spend":1.0,"flat_cost":0.0,"prompt_tokens":5,"completion_tokens":2,"cache_read_input_tokens":2,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":7,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"deleted-target","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":false}},"key-b":{"metrics":{"spend":90.0,"flat_cost":0.0,"prompt_tokens":3,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":5,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-b","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-a":{"metrics":{"spend":100.0,"flat_cost":0.0,"prompt_tokens":2,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":4,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-a","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-cache":{"metrics":{"spend":2.0,"flat_cost":0.0,"prompt_tokens":1000,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1002,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-cache","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}}}}},"api_keys":{"key-a":{"metrics":{"spend":100.0,"flat_cost":0.0,"prompt_tokens":2,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":4,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-a","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-b":{"metrics":{"spend":90.0,"flat_cost":0.0,"prompt_tokens":3,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":5,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-b","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-c":{"metrics":{"spend":80.0,"flat_cost":0.0,"prompt_tokens":4,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":6,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-c","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-cache":{"metrics":{"spend":2.0,"flat_cost":0.0,"prompt_tokens":1000,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1002,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-cache","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-target":{"metrics":{"spend":1.0,"flat_cost":0.0,"prompt_tokens":5,"completion_tokens":2,"cache_read_input_tokens":2,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":7,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"deleted-target","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":false}}},"entities":{"user-1":{"metrics":{"spend":1273.0,"flat_cost":0.0,"prompt_tokens":1014,"completion_tokens":10,"cache_read_input_tokens":6,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1024,"successful_requests":5,"failed_requests":0,"api_requests":5,"total_response_time_ms":0,"timed_requests":0},"metadata":{},"api_key_breakdown":{"key-a":{"metrics":{"spend":100.0,"flat_cost":0.0,"prompt_tokens":2,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":4,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-a","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-b":{"metrics":{"spend":90.0,"flat_cost":0.0,"prompt_tokens":3,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":5,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-b","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-c":{"metrics":{"spend":80.0,"flat_cost":0.0,"prompt_tokens":4,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":6,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-c","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-cache":{"metrics":{"spend":2.0,"flat_cost":0.0,"prompt_tokens":1000,"completion_tokens":2,"cache_read_input_tokens":1,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":1002,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"alias-cache","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":true}},"key-target":{"metrics":{"spend":1.0,"flat_cost":0.0,"prompt_tokens":5,"completion_tokens":2,"cache_read_input_tokens":2,"cache_creation_input_tokens":0,"compression_saved_tokens":0,"compression_savings_spend":0.0,"prompt_caching_savings_spend":0.0,"gateway_injected_caching_savings_spend":0.0,"autorouter_savings_spend":0.0,"total_tokens":7,"successful_requests":1,"failed_requests":0,"api_requests":1,"total_response_time_ms":0,"timed_requests":0},"metadata":{"key_alias":"deleted-target","team_id":"team-1","user_id":"user-1","user_email":"user@example.com","key_exists":false}}}}}}}],"metadata":{"total_spend":1273.0,"total_flat_cost":0.0,"total_prompt_tokens":1014,"total_completion_tokens":10,"total_tokens":1024,"total_api_requests":5,"total_successful_requests":5,"total_failed_requests":0,"total_cache_read_input_tokens":6,"total_cache_creation_input_tokens":0,"total_compression_saved_tokens":0,"total_compression_savings_spend":0.0,"total_prompt_caching_savings_spend":0.0,"total_gateway_injected_caching_savings_spend":0.0,"total_autorouter_savings_spend":0.0,"total_response_time_ms":0,"total_timed_requests":0,"page":1,"total_pages":1,"has_more":false,"api_key_limit":100,"total_api_keys":5,"entity_total_api_keys":{"user-1":5}}} diff --git a/tests/integration/spend/golden/daily_activity_team_aggregated.json b/tests/integration/spend/golden/daily_activity_team_aggregated.json new file mode 100644 index 00000000000..f434e767700 --- /dev/null +++ b/tests/integration/spend/golden/daily_activity_team_aggregated.json @@ -0,0 +1,1169 @@ +{ + "metadata": {"entity_total_api_keys":{"team-1":5}, + "api_key_limit": 100, + "has_more": false, + "page": 1, + "total_api_keys": 5, + "total_api_requests": 5, + "total_autorouter_savings_spend": 0.0, + "total_cache_creation_input_tokens": 0, + "total_cache_read_input_tokens": 6, + "total_completion_tokens": 10, + "total_compression_saved_tokens": 0, + "total_compression_savings_spend": 0.0, + "total_failed_requests": 0, + "total_flat_cost": 0.0, + "total_gateway_injected_caching_savings_spend": 0.0, + "total_pages": 1, + "total_prompt_caching_savings_spend": 0.0, + "total_prompt_tokens": 1014, + "total_response_time_ms": 0, + "total_spend": 1273.0, + "total_successful_requests": 5, + "total_timed_requests": 0, + "total_tokens": 1024 + }, + "results": [ + { + "breakdown": { + "api_keys": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + }, + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "endpoints": { + "/v1/chat/completions": { + "api_key_breakdown": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + }, + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 5, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 6, + "completion_tokens": 10, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1014, + "spend": 1273.0, + "successful_requests": 5, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1024 + } + } + }, + "entities": { + "team-1": { + "api_key_breakdown": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + }, + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "metadata": { + "team_alias": "Usage Team" + }, + "metrics": { + "api_requests": 5, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 6, + "completion_tokens": 10, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1014, + "spend": 1273.0, + "successful_requests": 5, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1024 + } + } + }, + "mcp_servers": {}, + "model_groups": { + "model-cache": { + "api_key_breakdown": { + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "model-popular": { + "api_key_breakdown": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 3, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 3, + "completion_tokens": 6, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 9, + "spend": 270.0, + "successful_requests": 3, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 15 + } + }, + "model-ptu": { + "api_key_breakdown": {}, + "metadata": {}, + "metrics": { + "api_requests": 0, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "completion_tokens": 0, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 0, + "spend": 1000.0, + "successful_requests": 0, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 0 + } + }, + "model-target": { + "api_key_breakdown": { + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "models": { + "model-cache": { + "api_key_breakdown": { + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "model-popular": { + "api_key_breakdown": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 3, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 3, + "completion_tokens": 6, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 9, + "spend": 270.0, + "successful_requests": 3, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 15 + } + }, + "model-ptu": { + "api_key_breakdown": {}, + "metadata": {}, + "metrics": { + "api_requests": 0, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "completion_tokens": 0, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 0, + "spend": 1000.0, + "successful_requests": 0, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 0 + } + }, + "model-target": { + "api_key_breakdown": { + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "providers": { + "provider-a": { + "api_key_breakdown": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + }, + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 5, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 6, + "completion_tokens": 10, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1014, + "spend": 1273.0, + "successful_requests": 5, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1024 + } + } + } + }, + "date": "2026-06-01", + "metrics": { + "api_requests": 5, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 6, + "completion_tokens": 10, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1014, + "spend": 1273.0, + "successful_requests": 5, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1024 + } + } + ] +} diff --git a/tests/integration/spend/golden/daily_activity_team_paginated.json b/tests/integration/spend/golden/daily_activity_team_paginated.json new file mode 100644 index 00000000000..e56ebbffe55 --- /dev/null +++ b/tests/integration/spend/golden/daily_activity_team_paginated.json @@ -0,0 +1,1197 @@ +{ + "metadata": {"entity_total_api_keys":null, + "api_key_limit": null, + "has_more": false, + "page": 1, + "total_api_keys": null, + "total_api_requests": 5, + "total_autorouter_savings_spend": 0.0, + "total_cache_creation_input_tokens": 0, + "total_cache_read_input_tokens": 6, + "total_completion_tokens": 10, + "total_compression_saved_tokens": 0, + "total_compression_savings_spend": 0.0, + "total_failed_requests": 0, + "total_flat_cost": 0.0, + "total_gateway_injected_caching_savings_spend": 0.0, + "total_pages": 1, + "total_prompt_caching_savings_spend": 0.0, + "total_prompt_tokens": 1014, + "total_response_time_ms": 0, + "total_spend": 1273.0, + "total_successful_requests": 5, + "total_timed_requests": 0, + "total_tokens": 1024 + }, + "results": [ + { + "breakdown": { + "api_keys": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + }, + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "endpoints": { + "/v1/chat/completions": { + "api_key_breakdown": { + "__ptu_flat_cost__": { + "metadata": { + "key_alias": null, + "key_exists": false, + "team_id": null, + "user_email": null, + "user_id": null + }, + "metrics": { + "api_requests": 0, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "completion_tokens": 0, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 0, + "spend": 1000.0, + "successful_requests": 0, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 0 + } + }, + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + }, + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 5, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 6, + "completion_tokens": 10, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1014, + "spend": 1273.0, + "successful_requests": 5, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1024 + } + } + }, + "entities": { + "team-1": { + "api_key_breakdown": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + }, + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "metadata": { + "team_alias": "Usage Team" + }, + "metrics": { + "api_requests": 5, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 6, + "completion_tokens": 10, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1014, + "spend": 1273.0, + "successful_requests": 5, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1024 + } + } + }, + "mcp_servers": {}, + "model_groups": { + "model-cache": { + "api_key_breakdown": { + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "model-popular": { + "api_key_breakdown": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 3, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 3, + "completion_tokens": 6, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 9, + "spend": 270.0, + "successful_requests": 3, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 15 + } + }, + "model-ptu": { + "api_key_breakdown": {}, + "metadata": {}, + "metrics": { + "api_requests": 0, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "completion_tokens": 0, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 0, + "spend": 1000.0, + "successful_requests": 0, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 0 + } + }, + "model-target": { + "api_key_breakdown": { + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "models": { + "model-cache": { + "api_key_breakdown": { + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "model-popular": { + "api_key_breakdown": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 3, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 3, + "completion_tokens": 6, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 9, + "spend": 270.0, + "successful_requests": 3, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 15 + } + }, + "model-ptu": { + "api_key_breakdown": {}, + "metadata": {}, + "metrics": { + "api_requests": 0, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "completion_tokens": 0, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 0, + "spend": 1000.0, + "successful_requests": 0, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 0 + } + }, + "model-target": { + "api_key_breakdown": { + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "providers": { + "provider-a": { + "api_key_breakdown": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + }, + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 5, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 6, + "completion_tokens": 10, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1014, + "spend": 273.0, + "successful_requests": 5, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1024 + } + } + } + }, + "date": "2026-06-01", + "metrics": { + "api_requests": 5, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 6, + "completion_tokens": 10, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1014, + "spend": 1273.0, + "successful_requests": 5, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1024 + } + } + ] +} diff --git a/tests/integration/spend/golden/daily_activity_user_aggregated.json b/tests/integration/spend/golden/daily_activity_user_aggregated.json new file mode 100644 index 00000000000..de1bbd63977 --- /dev/null +++ b/tests/integration/spend/golden/daily_activity_user_aggregated.json @@ -0,0 +1,1002 @@ +{ + "metadata": {"entity_total_api_keys":null, + "api_key_limit": 100, + "has_more": false, + "page": 1, + "total_api_keys": 5, + "total_api_requests": 5, + "total_autorouter_savings_spend": 0.0, + "total_cache_creation_input_tokens": 0, + "total_cache_read_input_tokens": 6, + "total_completion_tokens": 10, + "total_compression_saved_tokens": 0, + "total_compression_savings_spend": 0.0, + "total_failed_requests": 0, + "total_flat_cost": 0.0, + "total_gateway_injected_caching_savings_spend": 0.0, + "total_pages": 1, + "total_prompt_caching_savings_spend": 0.0, + "total_prompt_tokens": 1014, + "total_response_time_ms": 0, + "total_spend": 1273.0, + "total_successful_requests": 5, + "total_timed_requests": 0, + "total_tokens": 1024 + }, + "results": [ + { + "breakdown": { + "api_keys": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + }, + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "endpoints": { + "/v1/chat/completions": { + "api_key_breakdown": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + }, + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 5, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 6, + "completion_tokens": 10, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1014, + "spend": 1273.0, + "successful_requests": 5, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1024 + } + } + }, + "entities": {}, + "mcp_servers": {}, + "model_groups": { + "model-cache": { + "api_key_breakdown": { + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "model-popular": { + "api_key_breakdown": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 3, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 3, + "completion_tokens": 6, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 9, + "spend": 270.0, + "successful_requests": 3, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 15 + } + }, + "model-ptu": { + "api_key_breakdown": {}, + "metadata": {}, + "metrics": { + "api_requests": 0, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "completion_tokens": 0, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 0, + "spend": 1000.0, + "successful_requests": 0, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 0 + } + }, + "model-target": { + "api_key_breakdown": { + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "models": { + "model-cache": { + "api_key_breakdown": { + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "model-popular": { + "api_key_breakdown": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 3, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 3, + "completion_tokens": 6, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 9, + "spend": 270.0, + "successful_requests": 3, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 15 + } + }, + "model-ptu": { + "api_key_breakdown": {}, + "metadata": {}, + "metrics": { + "api_requests": 0, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "completion_tokens": 0, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 0, + "spend": 1000.0, + "successful_requests": 0, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 0 + } + }, + "model-target": { + "api_key_breakdown": { + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "providers": { + "provider-a": { + "api_key_breakdown": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + }, + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 5, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 6, + "completion_tokens": 10, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1014, + "spend": 1273.0, + "successful_requests": 5, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1024 + } + } + } + }, + "date": "2026-06-01", + "metrics": { + "api_requests": 5, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 6, + "completion_tokens": 10, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1014, + "spend": 1273.0, + "successful_requests": 5, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1024 + } + } + ] +} diff --git a/tests/integration/spend/golden/daily_activity_user_paginated.json b/tests/integration/spend/golden/daily_activity_user_paginated.json new file mode 100644 index 00000000000..b27bc405ebd --- /dev/null +++ b/tests/integration/spend/golden/daily_activity_user_paginated.json @@ -0,0 +1,1198 @@ +{ + "metadata": {"entity_total_api_keys":null, + "api_key_limit": null, + "has_more": false, + "page": 1, + "total_api_keys": null, + "total_api_requests": 5, + "total_autorouter_savings_spend": 0.0, + "total_cache_creation_input_tokens": 0, + "total_cache_read_input_tokens": 6, + "total_completion_tokens": 10, + "total_compression_saved_tokens": 0, + "total_compression_savings_spend": 0.0, + "total_failed_requests": 0, + "total_flat_cost": 0.0, + "total_gateway_injected_caching_savings_spend": 0.0, + "total_pages": 1, + "total_prompt_caching_savings_spend": 0.0, + "total_prompt_tokens": 1014, + "total_response_time_ms": 0, + "total_spend": 1273.0, + "total_successful_requests": 5, + "total_timed_requests": 0, + "total_tokens": 1024 + }, + "results": [ + { + "breakdown": { + "api_keys": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + }, + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "endpoints": { + "/v1/chat/completions": { + "api_key_breakdown": { + "__ptu_flat_cost__": { + "metadata": { + "key_alias": null, + "key_exists": false, + "team_id": null, + "user_email": null, + "user_id": null + }, + "metrics": { + "api_requests": 0, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "completion_tokens": 0, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 0, + "spend": 1000.0, + "successful_requests": 0, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 0 + } + }, + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + }, + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 5, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 6, + "completion_tokens": 10, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1014, + "spend": 1273.0, + "successful_requests": 5, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1024 + } + } + }, + "entities": { + "user-1": { + "api_key_breakdown": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + }, + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "metadata": { + "user_alias": null, + "user_email": "user@example.com" + }, + "metrics": { + "api_requests": 5, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 6, + "completion_tokens": 10, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1014, + "spend": 1273.0, + "successful_requests": 5, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1024 + } + } + }, + "mcp_servers": {}, + "model_groups": { + "model-cache": { + "api_key_breakdown": { + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "model-popular": { + "api_key_breakdown": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 3, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 3, + "completion_tokens": 6, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 9, + "spend": 270.0, + "successful_requests": 3, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 15 + } + }, + "model-ptu": { + "api_key_breakdown": {}, + "metadata": {}, + "metrics": { + "api_requests": 0, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "completion_tokens": 0, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 0, + "spend": 1000.0, + "successful_requests": 0, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 0 + } + }, + "model-target": { + "api_key_breakdown": { + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "models": { + "model-cache": { + "api_key_breakdown": { + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "model-popular": { + "api_key_breakdown": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 3, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 3, + "completion_tokens": 6, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 9, + "spend": 270.0, + "successful_requests": 3, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 15 + } + }, + "model-ptu": { + "api_key_breakdown": {}, + "metadata": {}, + "metrics": { + "api_requests": 0, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "completion_tokens": 0, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 0, + "spend": 1000.0, + "successful_requests": 0, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 0 + } + }, + "model-target": { + "api_key_breakdown": { + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "providers": { + "provider-a": { + "api_key_breakdown": { + "key-a": { + "metadata": { + "key_alias": "alias-a", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 2, + "spend": 100.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 4 + } + }, + "key-b": { + "metadata": { + "key_alias": "alias-b", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 3, + "spend": 90.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 5 + } + }, + "key-c": { + "metadata": { + "key_alias": "alias-c", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 4, + "spend": 80.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 6 + } + }, + "key-cache": { + "metadata": { + "key_alias": "alias-cache", + "key_exists": true, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 1, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1000, + "spend": 2.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1002 + } + }, + "key-target": { + "metadata": { + "key_alias": "deleted-target", + "key_exists": false, + "team_id": "team-1", + "user_email": "user@example.com", + "user_id": "user-1" + }, + "metrics": { + "api_requests": 1, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 2, + "completion_tokens": 2, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 5, + "spend": 1.0, + "successful_requests": 1, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 7 + } + } + }, + "metadata": {}, + "metrics": { + "api_requests": 5, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 6, + "completion_tokens": 10, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1014, + "spend": 273.0, + "successful_requests": 5, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1024 + } + } + } + }, + "date": "2026-06-01", + "metrics": { + "api_requests": 5, + "autorouter_savings_spend": 0.0, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 6, + "completion_tokens": 10, + "compression_saved_tokens": 0, + "compression_savings_spend": 0.0, + "failed_requests": 0, + "flat_cost": 0.0, + "gateway_injected_caching_savings_spend": 0.0, + "prompt_caching_savings_spend": 0.0, + "prompt_tokens": 1014, + "spend": 1273.0, + "successful_requests": 5, + "timed_requests": 0, + "total_response_time_ms": 0, + "total_tokens": 1024 + } + } + ] +} diff --git a/tests/integration/spend/test_daily_activity_aggregated_breakdowns.py b/tests/integration/spend/test_daily_activity_aggregated_breakdowns.py index 56c936a716f..33baf9d0e2e 100644 --- a/tests/integration/spend/test_daily_activity_aggregated_breakdowns.py +++ b/tests/integration/spend/test_daily_activity_aggregated_breakdowns.py @@ -6,7 +6,7 @@ from typing import Final import pytest from pydantic import JsonValue, TypeAdapter -from litellm.constants import PTU_SENTINEL_API_KEY +from litellm.constants import PTU_SENTINEL_API_KEY, USAGE_TOP_API_KEYS_DEFAULT from tests.integration._support.client import Gateway, object_value from tests.integration._support.database import write_rows @@ -43,57 +43,86 @@ def _row_id() -> str: return f"agg-{uuid.uuid4().hex}" +def _ranked_key_rows(day: str, count: int) -> list[tuple[object, ...]]: + return [ + ( + _row_id(), + f"user-{i:03d}", + day, + f"key-{i:03d}", + "gpt-5", + "", + "openai", + None, + "/v1/chat/completions", + 10, + 6.0 if i == 4 else float(i + 1), + 1, + 1, + ) + for i in range(count) + ] + + @pytest.mark.asyncio -async def test_get_daily_activity_aggregated_returns_every_api_key(gateway: Gateway) -> None: +async def test_get_daily_activity_aggregated_bounds_api_key_rollups(gateway: Gateway) -> None: + """key-004 and key-005 tie on spend exactly at the default api_key_limit cutoff; the api_key + tiebreaker keeps key-004 and drops key-005. The PTU sentinel outspends every key but takes no + slot. Dropped keys and the sentinel still count toward the totals and the model rollup.""" + key_count: Final = USAGE_TOP_API_KEYS_DEFAULT + 5 day: Final = _unique_day() _seed( day, [ - *[ - ( - _row_id(), - f"user-{i:03d}", - day, - f"key-{i:03d}", - "gpt-5", - "", - "openai", - None, - "/v1/chat/completions", - 10, - 6.0 if i == 4 else float(i + 1), - 1, - 1, - ) - for i in range(105) - ], + *_ranked_key_rows(day, key_count), (_row_id(), None, day, PTU_SENTINEL_API_KEY, "gpt-5", "", "azure", None, None, 0, 1000.0, 0, 0), ], ) + key_spend: Final = sum(6.0 if i == 4 else float(i + 1) for i in range(key_count)) try: body: Final = _activity(gateway, day) metadata: Final = object_value(body["metadata"]) - assert metadata["total_spend"] == pytest.approx(6566.0) - assert metadata["total_api_requests"] == 105 + assert metadata["total_spend"] == pytest.approx(key_spend + 1000.0) + assert metadata["total_api_requests"] == key_count + assert metadata["total_api_keys"] == key_count + assert metadata["api_key_limit"] == USAGE_TOP_API_KEYS_DEFAULT results: Final = _RESULTS.validate_python(body["results"]) assert len(results) == 1 result_day: Final = object_value(results[0]) - assert object_value(result_day["metrics"])["spend"] == pytest.approx(6566.0) + assert object_value(result_day["metrics"])["spend"] == pytest.approx(key_spend + 1000.0) breakdown: Final = object_value(result_day["breakdown"]) - expected_api_keys: Final = {f"key-{i:03d}" for i in range(105)} + expected_top: Final = {f"key-{i:03d}" for i in range(6, key_count)} | {"key-004"} api_keys: Final = object_value(breakdown["api_keys"]) - assert set(api_keys) == expected_api_keys + assert set(api_keys) == expected_top + assert object_value(object_value(api_keys["key-004"])["metrics"])["spend"] == 6.0 assert PTU_SENTINEL_API_KEY not in api_keys models: Final = object_value(breakdown["models"]) gpt5: Final = object_value(models["gpt-5"]) - assert object_value(gpt5["metrics"])["spend"] == pytest.approx(6566.0) - assert set(object_value(gpt5["api_key_breakdown"])) == expected_api_keys + assert object_value(gpt5["metrics"])["spend"] == pytest.approx(key_spend + 1000.0) + assert set(object_value(gpt5["api_key_breakdown"])) == expected_top providers: Final = object_value(breakdown["providers"]) openai: Final = object_value(providers["openai"]) - assert object_value(openai["metrics"])["spend"] == pytest.approx(5566.0) - assert set(object_value(openai["api_key_breakdown"])) == expected_api_keys + assert object_value(openai["metrics"])["spend"] == pytest.approx(key_spend) + assert set(object_value(openai["api_key_breakdown"])) == expected_top endpoints: Final = object_value(breakdown["endpoints"]) - assert object_value(object_value(endpoints["/v1/chat/completions"])["metrics"])["api_requests"] == 105 + assert object_value(object_value(endpoints["/v1/chat/completions"])["metrics"])["api_requests"] == key_count + finally: + _clean(day) + + +@pytest.mark.asyncio +async def test_get_daily_activity_aggregated_reports_exact_limit_key_count_as_complete(gateway: Gateway) -> None: + """With exactly USAGE_TOP_API_KEYS_DEFAULT keys nothing is dropped and total_api_keys equals the limit.""" + day: Final = _unique_day() + _seed(day, _ranked_key_rows(day, USAGE_TOP_API_KEYS_DEFAULT)) + try: + body: Final = _activity(gateway, day) + metadata: Final = object_value(body["metadata"]) + assert metadata["total_api_keys"] == USAGE_TOP_API_KEYS_DEFAULT + assert metadata["api_key_limit"] == USAGE_TOP_API_KEYS_DEFAULT + results: Final = _RESULTS.validate_python(body["results"]) + api_keys: Final = object_value(object_value(object_value(results[0])["breakdown"])["api_keys"]) + assert set(api_keys) == {f"key-{i:03d}" for i in range(USAGE_TOP_API_KEYS_DEFAULT)} finally: _clean(day) @@ -126,7 +155,9 @@ async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_resu ) try: body: Final = _activity(gateway, day, api_key="key-1") - assert object_value(body["metadata"])["total_spend"] == 2.0 + metadata: Final = object_value(body["metadata"]) + assert metadata["total_spend"] == 2.0 + assert metadata["total_api_keys"] == 1 results: Final = _RESULTS.validate_python(body["results"]) assert len(results) == 1 breakdown: Final = object_value(object_value(results[0])["breakdown"]) diff --git a/tests/integration/spend/test_daily_activity_repository.py b/tests/integration/spend/test_daily_activity_repository.py new file mode 100644 index 00000000000..c6ffeef2eba --- /dev/null +++ b/tests/integration/spend/test_daily_activity_repository.py @@ -0,0 +1,663 @@ +import os +import uuid +from collections.abc import AsyncIterator, Mapping +from contextlib import asynccontextmanager +from dataclasses import dataclass +from math import isclose +from pathlib import Path +from types import MappingProxyType +from typing import Final, cast +from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit + +import psycopg +import pytest +from integration.spend._daily_activity_fixtures import ( + seed_daily_activity_fixture, + seed_daily_tag_activity_fixture, + seed_daily_tag_float_tie_fixture, + seed_daily_team_exclusion_fixture, + seed_daily_team_unassigned_fixture, +) +from prisma import Prisma +from psycopg import sql +from pydantic import TypeAdapter + +from litellm import constants +from litellm.proxy.management_endpoints.common_daily_activity import get_daily_activity_aggregated +from litellm.repositories.chunked_in import find_many_in +from litellm.repositories.daily_activity_repository import DailyActivityDatabase, DailyActivityRepository +from litellm.types.repositories.daily_activity import ( + DailyActivityProxyReads, + DailyActivityScope, + DailyActivityTable, + ExportType, + KeyMetadataRow, + SpendLogsWindow, +) + + +@dataclass(frozen=True, slots=True) +class _TagRollupMetrics: + tag: str | None + date: str + spend: float + api_requests: int + prompt_tokens: int + + +@dataclass(frozen=True, slots=True) +class _TagApiKeyCount: + tag: str | None + distinct_api_keys: int + + +@dataclass(frozen=True, slots=True) +class _TagKeyMembershipCount: + api_key: str + tag_count: int + + +@dataclass(frozen=True, slots=True) +class _TagFloatSpend: + api_key: str + spend: float + + +@dataclass(frozen=True, slots=True) +class _TagRankedKey: + api_key: str + + +@dataclass(frozen=True, slots=True) +class _TagDistinctKeyCount: + total_api_keys: int + + +_TAG_ROLLUP_METRICS_ADAPTER: Final = TypeAdapter(tuple[_TagRollupMetrics, ...]) +_TAG_API_KEY_COUNT_ADAPTER: Final = TypeAdapter(tuple[_TagApiKeyCount, ...]) +_TAG_KEY_MEMBERSHIP_COUNT_ADAPTER: Final = TypeAdapter(tuple[_TagKeyMembershipCount, ...]) +_TAG_FLOAT_SPEND_ADAPTER: Final = TypeAdapter(tuple[_TagFloatSpend, ...]) +_TAG_RANKED_KEY_ADAPTER: Final = TypeAdapter(tuple[_TagRankedKey, ...]) +_TAG_DISTINCT_KEY_COUNT_ADAPTER: Final = TypeAdapter(tuple[_TagDistinctKeyCount, ...]) + + +def _scoped_url(url: str, schema: str) -> str: + parsed: Final = urlsplit(url) + return urlunsplit(parsed._replace(query=urlencode({**dict(parse_qsl(parsed.query)), "schema": schema}))) + + +@asynccontextmanager +async def _daily_activity_database( + *, + include_tag_activity: bool = False, + include_tag_float_tie_activity: bool = False, + include_team_unassigned_activity: bool = False, + include_team_exclusion_activity: bool = False, +) -> AsyncIterator[Prisma]: + schema: Final = f"integration_{uuid.uuid4().hex}" + url: Final = os.environ["DATABASE_URL"] + with psycopg.connect(url, autocommit=True) as setup: + setup.execute(sql.SQL("CREATE SCHEMA {}").format(sql.Identifier(schema))) + try: + with psycopg.connect(url) as connection: + seed_daily_activity_fixture( + connection, + schema=schema, + ptu_sentinel_api_key=constants.PTU_SENTINEL_API_KEY, + ) + if include_tag_activity: + seed_daily_tag_activity_fixture(connection, schema=schema) + if include_tag_float_tie_activity: + seed_daily_tag_float_tie_fixture(connection, schema=schema) + if include_team_unassigned_activity: + seed_daily_team_unassigned_fixture( + connection, schema=schema, ptu_sentinel_api_key=constants.PTU_SENTINEL_API_KEY + ) + if include_team_exclusion_activity: + seed_daily_team_exclusion_fixture(connection, schema=schema) + database: Final = Prisma(datasource={"url": _scoped_url(url, schema)}) + await database.connect() + try: + yield database + finally: + await database.disconnect() + finally: + setup.execute(sql.SQL("DROP SCHEMA {} CASCADE").format(sql.Identifier(schema))) + + +@dataclass(frozen=True, slots=True) +class _PrismaDatabase: + db: Prisma + + +@dataclass(frozen=True, slots=True) +class _ProxyReads(DailyActivityProxyReads): + database: Prisma + + async def recover_key_metadata( + self, resolved: Mapping[str, KeyMetadataRow], api_keys: frozenset[str], window: SpendLogsWindow | None + ) -> Mapping[str, KeyMetadataRow]: + user_ids: Final = frozenset(row.user_id for row in resolved.values() if row.user_id) + user_rows: Final = await find_many_in(self.database.litellm_usertable, "user_id", user_ids) if user_ids else () + user_emails: Final = MappingProxyType({row.user_id: row.user_email for row in user_rows if row.user_email}) + return MappingProxyType( + { + key: KeyMetadataRow( + api_key=row.api_key, + key_alias=row.key_alias, + team_id=row.team_id, + user_id=row.user_id, + user_email=row.user_email or user_emails.get(row.user_id), + key_exists=row.key_exists, + tags=row.tags, + ) + for key, row in resolved.items() + } + ) + + +def _repository(database: Prisma) -> DailyActivityRepository: + client: Final = cast(DailyActivityDatabase, _PrismaDatabase(database)) + return DailyActivityRepository(client, proxy_reads=_ProxyReads(database)) + + +def _scope( + table: DailyActivityTable, + entity_id_field: str, + entity_id: str, + api_keys: tuple[str, ...] | None = None, +) -> DailyActivityScope: + return DailyActivityScope( + table=table, + entity_id_field=entity_id_field, + entity_ids=(entity_id,), + exclude_entity_ids=(), + api_keys=api_keys, + start_date="2026-06-01", + end_date="2026-06-01", + model=None, + timezone_offset_minutes=None, + ) + + +@pytest.mark.asyncio +async def test_repository_queries_and_exports_seeded_daily_activity(monkeypatch: pytest.MonkeyPatch) -> None: + async with _daily_activity_database() as database: + repository: Final = _repository(database) + team_scope: Final = _scope(DailyActivityTable.TEAM, "team_id", "team-1") + monkeypatch.setattr(constants, "USAGE_EXPORT_BATCH_SIZE", 2) + aggregate: Final = await repository.aggregated(team_scope, include_entity_breakdown=True, api_key_limit=3) + totals: Final = tuple(row for row in aggregate.grouping_rows if row.group_level == 127) + assert len(totals) == 1 + assert totals[0].spend == 1273.0 + assert totals[0].ptu_flat_cost == 42.0 + assert aggregate.distinct_api_keys == 5 + grouped_keys: Final = frozenset( + row.api_key for row in aggregate.grouping_rows if row.group_level == 31 and row.api_key + ) + assert grouped_keys == frozenset(("key-a", "key-b", "key-c")) + entity_totals: Final = tuple(row for row in aggregate.entity_rows or () if row.api_key_rolled) + assert len(entity_totals) == 1 + assert entity_totals[0].spend == 1273.0 + assert entity_totals[0].ptu_flat_cost == 42.0 + + targeted_model_keys: Final = await repository.model_top_keys( + team_scope, model_group="model-target", by_model_group=False, limit=3 + ) + assert tuple(row.api_key for row in targeted_model_keys) == ("key-target",) + popular_model_keys: Final = await repository.model_top_keys( + team_scope, model_group="model-popular", by_model_group=False, limit=3 + ) + assert tuple(row.api_key for row in popular_model_keys) == ("key-a", "key-b", "key-c") + assert await repository.search_keys(team_scope, search="target", limit=10) == ("key-target",) + assert await repository.search_keys(team_scope, search="deleted-target", limit=20) == ("key-target",) + leakage_keys: Final = await repository.cache_leakage_keys(team_scope, limit=2) + assert tuple(row.api_key for row in leakage_keys) == ("key-cache", "key-c") + + exports: Final = tuple( + [row async for row in repository.export_rows(team_scope, export_type=ExportType.DAILY_WITH_KEYS)] + ) + assert tuple(row.api_key for row in exports) == ( + "key-a", + "key-b", + "key-c", + "key-cache", + "key-target", + ) + assert sum(row.spend for row in exports) == 273.0 + assert sum(row.flat_cost for row in exports) == 0.0 + deleted_key_export: Final = next(row for row in exports if row.api_key == "key-target") + assert (deleted_key_export.key_alias, deleted_key_export.user_id, deleted_key_export.user_email) == ( + "deleted-target", + "user-1", + "user@example.com", + ) + user_exports: Final = tuple( + [row async for row in repository.export_rows(team_scope, export_type=ExportType.DAILY_WITH_USERS)] + ) + assert len(user_exports) == 1 + assert (user_exports[0].user_id, user_exports[0].user_email, user_exports[0].spend) == ( + "user-1", + "user@example.com", + 273.0, + ) + daily_export: Final = tuple( + [row async for row in repository.export_rows(team_scope, export_type=ExportType.DAILY)] + ) + assert len(daily_export) == 1 + assert daily_export[0].spend == 1273.0 + assert daily_export[0].flat_cost == 42.0 + + metadata: Final = await repository.key_metadata(frozenset(("key-a", "key-target")), None) + assert metadata["key-a"].key_exists is True + assert metadata["key-a"].key_alias == "alias-a" + assert metadata["key-a"].tags == ("blue", "gold") + assert metadata["key-a"].user_email == "user@example.com" + assert metadata["key-target"].key_exists is False + assert metadata["key-target"].key_alias == "deleted-target" + assert metadata["key-target"].tags == ("archived",) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("search", "api_keys", "expected_keys"), + ( + ("needle-alias", None, ("needle-key",)), + ("needle-user", None, ("needle-key",)), + ("needle@example.com", None, ("needle-key",)), + ("needle-alias", ("key-a",), ()), + ), +) +async def test_search_keys_matches_token_metadata_outside_top_n_and_respects_scope( + search: str, + api_keys: tuple[str, ...] | None, + expected_keys: tuple[str, ...], +) -> None: + async with _daily_activity_database() as database: + await database.execute_raw( + """ + INSERT INTO "LiteLLM_DailyTeamSpend" ( + id, team_id, date, api_key, model, model_group, custom_llm_provider, + mcp_namespaced_tool_name, endpoint, prompt_tokens, completion_tokens, + cache_read_input_tokens, cache_creation_input_tokens, spend, api_requests, + successful_requests, failed_requests, ptu_flat_cost, updated_at + ) VALUES ( + $1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15, $16, $17, $18, $19::timestamp + ) + """, + "needle-row", + "team-1", + "2026-06-01", + "needle-key", + "model-needle", + "", + "provider-a", + None, + "/v1/chat/completions", + 1, + 1, + 0, + 0, + 0.5, + 1, + 1, + 0, + 0.0, + "2026-06-01 12:00:00", + ) + await database.execute_raw( + """ + INSERT INTO "LiteLLM_VerificationToken" (token, key_alias, team_id, user_id, metadata, models) + VALUES ($1, $2, $3, $4, $5::jsonb, $6::text[]) + """, + "needle-key", + "needle-alias", + "team-1", + "needle-user", + '{"tags": []}', + [], + ) + await database.execute_raw( + """ + INSERT INTO "LiteLLM_UserTable" (user_id, user_email, models) + VALUES ($1, $2, $3::text[]) + """, + "needle-user", + "needle@example.com", + [], + ) + + repository: Final = _repository(database) + team_scope: Final = _scope(DailyActivityTable.TEAM, "team_id", "team-1", api_keys=api_keys) + aggregate: Final = await repository.aggregated(team_scope, include_entity_breakdown=False, api_key_limit=1) + top_keys: Final = frozenset( + row.api_key for row in aggregate.grouping_rows if row.group_level == 31 and row.api_key + ) + assert "needle-key" not in top_keys + assert await repository.search_keys(team_scope, search=search, limit=10) == expected_keys + + +@pytest.mark.asyncio +async def test_aggregated_returns_totals_with_a_one_key_limit() -> None: + async with _daily_activity_database() as database: + aggregate: Final = await _repository(database).aggregated( + _scope(DailyActivityTable.TEAM, "team_id", "team-1"), + include_entity_breakdown=False, + api_key_limit=1, + ) + + totals: Final = tuple(row for row in aggregate.grouping_rows if row.group_level == 127) + per_key_rows: Final = tuple(row for row in aggregate.grouping_rows if row.api_key is not None) + per_key_names: Final = frozenset(row.api_key for row in per_key_rows) + assert len(totals) == 1 + assert totals[0].spend == 1273.0 + assert len(per_key_rows) == 6 + assert len(per_key_names) == 1 + + +@pytest.mark.asyncio +async def test_tag_entity_rollups_bound_keys_and_preserve_full_scope_totals() -> None: + async with _daily_activity_database(include_tag_activity=True) as database: + scope: Final = DailyActivityScope( + table=DailyActivityTable.TAG, + entity_id_field="tag", + entity_ids=None, + exclude_entity_ids=(), + api_keys=None, + start_date="2026-06-01", + end_date="2026-06-02", + model=None, + timezone_offset_minutes=None, + ) + repository: Final = _repository(database) + aggregate: Final = await repository.aggregated(scope, include_entity_breakdown=True, api_key_limit=3) + independent_metrics: Final = _TAG_ROLLUP_METRICS_ADAPTER.validate_python( + await database.query_raw( + """ + SELECT tag, date, SUM(spend)::float AS spend, + SUM(api_requests)::bigint AS api_requests, + SUM(prompt_tokens)::bigint AS prompt_tokens + FROM "LiteLLM_DailyTagSpend" + WHERE date >= $1 AND date <= $2 + GROUP BY tag, date + """, + "2026-06-01", + "2026-06-02", + ) + ) + independent_key_counts: Final = _TAG_API_KEY_COUNT_ADAPTER.validate_python( + await database.query_raw( + """ + SELECT tag, COUNT(DISTINCT api_key)::bigint AS distinct_api_keys + FROM "LiteLLM_DailyTagSpend" + WHERE date >= $1 AND date <= $2 AND api_key <> $3 + GROUP BY tag + """, + "2026-06-01", + "2026-06-02", + constants.PTU_SENTINEL_API_KEY, + ) + ) + independent_key_memberships: Final = _TAG_KEY_MEMBERSHIP_COUNT_ADAPTER.validate_python( + await database.query_raw( + """ + SELECT api_key, COUNT(DISTINCT tag)::bigint AS tag_count + FROM "LiteLLM_DailyTagSpend" + WHERE date >= $1 AND date <= $2 + GROUP BY api_key + """, + "2026-06-01", + "2026-06-02", + ) + ) + expected_metrics: Final = MappingProxyType({(row.date, row.tag): row for row in independent_metrics}) + expected_key_counts: Final = MappingProxyType( + {row.tag: row.distinct_api_keys for row in independent_key_counts} + ) + assert {row.date for row in independent_metrics} == {"2026-06-01", "2026-06-02"} + assert len(expected_key_counts) == 4 + assert len(independent_key_memberships) == 8 + assert all(row.tag_count == 2 for row in independent_key_memberships) + assert max(expected_key_counts.values()) > 3 + + entity_rows: Final = aggregate.entity_rows or () + rolled_rows: Final = tuple(row for row in entity_rows if row.api_key_rolled) + keyed_rows: Final = tuple(row for row in entity_rows if not row.api_key_rolled and row.api_key) + top_level_keys: Final = frozenset( + row.api_key for row in aggregate.grouping_rows if row.group_level == 31 and row.api_key + ) + entity_day_keys: Final = MappingProxyType( + { + key: frozenset(row.api_key for row in keyed_rows if (row.date, row.entity_id) == key and row.api_key) + for key in frozenset((row.date, row.entity_id) for row in keyed_rows) + } + ) + assert entity_day_keys + assert max(len(keys) for keys in entity_day_keys.values()) <= 3 + assert frozenset(row.api_key for row in keyed_rows) <= top_level_keys + + rolled_by_entity_day: Final = MappingProxyType( + {(row.date, row.entity_id): row for row in rolled_rows if row.date is not None} + ) + assert set(rolled_by_entity_day) == set(expected_metrics) + for key, row in rolled_by_entity_day.items(): + assert row.spend is not None + assert isclose(row.spend, expected_metrics[key].spend, rel_tol=1e-9, abs_tol=1e-9) + assert row.api_requests == expected_metrics[key].api_requests + assert row.prompt_tokens == expected_metrics[key].prompt_tokens + assert row.distinct_api_keys == expected_key_counts[row.entity_id] + + response: Final = await get_daily_activity_aggregated( + repository, + scope, + include_entity_breakdown=True, + api_key_limit=3, + ) + assert response.metadata.entity_total_api_keys == { + tag: count for tag, count in expected_key_counts.items() if tag is not None + } + assert all( + all(len(entity.api_key_breakdown) <= 3 for entity in day.breakdown.entities.values()) + for day in response.results + ) + + +@pytest.mark.asyncio +async def test_key_pages_match_full_tag_ranking_and_aggregate_top_keys() -> None: + async with _daily_activity_database(include_tag_activity=True) as database: + scope: Final = DailyActivityScope( + table=DailyActivityTable.TAG, + entity_id_field="tag", + entity_ids=None, + exclude_entity_ids=(), + api_keys=None, + start_date="2026-06-01", + end_date="2026-06-02", + model=None, + timezone_offset_minutes=None, + ) + repository: Final = _repository(database) + first_page: Final = await repository.key_page(scope, offset=0, limit=3) + remaining_pages: Final = tuple( + [ + await repository.key_page(scope, offset=offset, limit=3) + for offset in range(3, first_page.total_api_keys, 3) + ] + ) + pages: Final = (first_page, *remaining_pages) + actual_keys: Final = tuple(row.api_key for page in pages for row in page.rows) + expected_rows: Final = _TAG_RANKED_KEY_ADAPTER.validate_python( + await database.query_raw( + """ + SELECT api_key + FROM "LiteLLM_DailyTagSpend" + WHERE date >= $1 AND date <= $2 AND api_key <> $3 + GROUP BY api_key + ORDER BY SUM(spend::numeric) DESC, api_key + """, + "2026-06-01", + "2026-06-02", + constants.PTU_SENTINEL_API_KEY, + ) + ) + expected_keys: Final = tuple(row.api_key for row in expected_rows) + independent_count: Final = _TAG_DISTINCT_KEY_COUNT_ADAPTER.validate_python( + await database.query_raw( + """ + SELECT COUNT(DISTINCT api_key)::bigint AS total_api_keys + FROM "LiteLLM_DailyTagSpend" + WHERE date >= $1 AND date <= $2 AND api_key <> $3 + """, + "2026-06-01", + "2026-06-02", + constants.PTU_SENTINEL_API_KEY, + ) + )[0].total_api_keys + aggregate: Final = await repository.aggregated(scope, include_entity_breakdown=False, api_key_limit=3) + aggregate_top_keys: Final = frozenset( + row.api_key for row in aggregate.grouping_rows if row.group_level == 31 and row.api_key is not None + ) + empty_page: Final = await repository.key_page(scope, offset=independent_count + 3, limit=3) + + assert actual_keys == expected_keys + assert len(actual_keys) == len(frozenset(actual_keys)) + assert all(page.total_api_keys == independent_count for page in pages) + assert first_page.total_api_keys == independent_count + assert frozenset(row.api_key for row in first_page.rows) == aggregate_top_keys + assert empty_page.rows == () + assert empty_page.total_api_keys == independent_count + + +@pytest.mark.asyncio +async def test_top_api_key_rank_is_order_independent_for_float_ties() -> None: + async with _daily_activity_database(include_tag_float_tie_activity=True) as database: + float_totals: Final = _TAG_FLOAT_SPEND_ADAPTER.validate_python( + await database.query_raw( + """ + SELECT api_key, SUM(spend)::float AS spend + FROM "LiteLLM_DailyTagSpend" + WHERE tag = $1 AND date = $2 + GROUP BY api_key + """, + "tag-float-tie", + "2026-06-01", + ) + ) + float_spends: Final = MappingProxyType({row.api_key: row.spend for row in float_totals}) + assert float_spends["key-z"] > float_spends["key-a"] + + scope: Final = DailyActivityScope( + table=DailyActivityTable.TAG, + entity_id_field="tag", + entity_ids=None, + exclude_entity_ids=(), + api_keys=None, + start_date="2026-06-01", + end_date="2026-06-01", + model=None, + timezone_offset_minutes=None, + ) + aggregate: Final = await _repository(database).aggregated(scope, include_entity_breakdown=True, api_key_limit=1) + key_page: Final = await _repository(database).key_page(scope, offset=0, limit=1) + top_level_keys: Final = frozenset( + row.api_key for row in aggregate.grouping_rows if row.group_level == 31 and row.api_key + ) + entity_keyed_keys: Final = frozenset( + row.api_key for row in aggregate.entity_rows or () if not row.api_key_rolled and row.api_key is not None + ) + expected_keys: Final = frozenset(("key-a",)) + assert (top_level_keys, entity_keyed_keys) == ( + expected_keys, + expected_keys, + ), f"plain float SUM totals: {float_spends}" + assert frozenset(row.api_key for row in key_page.rows) == expected_keys + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("table", "entity_field", "entity_id", "fixture_name"), + ( + (DailyActivityTable.USER, "user_id", "user-1", "daily_activity_user.json"), + (DailyActivityTable.TEAM, "team_id", "team-1", "daily_activity_team.json"), + ), +) +async def test_aggregated_response_matches_base_golden( + table: DailyActivityTable, entity_field: str, entity_id: str, fixture_name: str +) -> None: + async with _daily_activity_database() as database: + result: Final = await get_daily_activity_aggregated( + _repository(database), + _scope(table, entity_field, entity_id), + entity_metadata_field=MappingProxyType({"team-1": {"team_alias": "Usage Team"}}), + include_entity_breakdown=True, + ) + golden_path: Final = Path(__file__).with_name("fixtures") / fixture_name + assert result.model_dump_json() + "\n" == golden_path.read_text() + + +@pytest.mark.asyncio +async def test_team_entity_rollups_merge_null_and_empty_entity_ids() -> None: + async with _daily_activity_database(include_team_unassigned_activity=True) as database: + scope: Final = DailyActivityScope( + table=DailyActivityTable.TEAM, + entity_id_field="team_id", + entity_ids=None, + exclude_entity_ids=(), + api_keys=None, + start_date="2026-06-03", + end_date="2026-06-03", + model=None, + timezone_offset_minutes=None, + ) + aggregate: Final = await _repository(database).aggregated(scope, include_entity_breakdown=True, api_key_limit=3) + + totals: Final = tuple(row for row in aggregate.grouping_rows if row.group_level == 127) + assert len(totals) == 1 + assert totals[0].spend == 23.0 + assert aggregate.distinct_api_keys == 2 + + rolled_rows: Final = tuple(row for row in aggregate.entity_rows or () if row.api_key_rolled) + assert len(rolled_rows) == 1 + assert rolled_rows[0].entity_id == "" + assert rolled_rows[0].spend == 23.0 + assert rolled_rows[0].ptu_flat_cost == 13.0 + assert rolled_rows[0].distinct_api_keys == 2 + + keyed_rows: Final = tuple(row for row in aggregate.entity_rows or () if not row.api_key_rolled) + assert {row.entity_id for row in keyed_rows} == {""} + assert {row.api_key for row in keyed_rows} == {"key-unassigned-null", "key-unassigned-empty"} + + +@pytest.mark.asyncio +async def test_team_exclusion_keeps_null_and_empty_entity_rows() -> None: + async with _daily_activity_database(include_team_exclusion_activity=True) as database: + repository: Final = _repository(database) + scope: Final = DailyActivityScope( + table=DailyActivityTable.TEAM, + entity_id_field="team_id", + entity_ids=None, + exclude_entity_ids=("litellm-dashboard",), + api_keys=None, + start_date="2026-06-04", + end_date="2026-06-04", + model=None, + timezone_offset_minutes=None, + ) + aggregate: Final = await repository.aggregated(scope, include_entity_breakdown=True, api_key_limit=10) + + totals: Final = tuple(row for row in aggregate.grouping_rows if row.group_level == 127) + assert len(totals) == 1 + assert totals[0].spend == 23.0 + assert aggregate.distinct_api_keys == 3 + + keyed_rows: Final = tuple(row for row in aggregate.entity_rows or () if not row.api_key_rolled) + assert {row.api_key for row in keyed_rows} == {"key-excluded-null", "key-excluded-empty", "key-excluded-normal"} + assert {row.entity_id for row in keyed_rows} == {"", "team-normal"} + + page: Final = await repository.key_page(scope, offset=0, limit=10) + assert page.total_api_keys == 3 + assert {row.api_key for row in page.rows} == {"key-excluded-null", "key-excluded-empty", "key-excluded-normal"} + + daily: Final = await repository.daily_rows(scope, page=1, page_size=10) + assert daily.total_count == 3 + assert {row.api_key for row in daily.rows} == {"key-excluded-null", "key-excluded-empty", "key-excluded-normal"} diff --git a/tests/integration/spend/test_daily_activity_routes.py b/tests/integration/spend/test_daily_activity_routes.py new file mode 100644 index 00000000000..9e4c16e573c --- /dev/null +++ b/tests/integration/spend/test_daily_activity_routes.py @@ -0,0 +1,548 @@ +import csv +import hashlib +import io +import uuid +from datetime import datetime, timedelta, timezone +from itertools import chain +from pathlib import Path +from typing import Final + +import httpx +import pytest +from fastapi import FastAPI +from integration._support.client import JSON_OBJECT, Gateway, eventually, object_value, string_value +from integration._support.database import read_rows +from integration._support.process import owned_proxy +from integration._support.upstream import JsonResponse, delete_scenario, register_scenario +from integration.spend.test_daily_activity_repository import _daily_activity_database, _PrismaDatabase, _repository + +from litellm import constants +from litellm.proxy import proxy_server +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.management_endpoints.daily_activity_routes import ( + get_daily_activity_prisma_client, + get_daily_activity_repository, +) +from litellm.proxy.management_endpoints.daily_activity_routes import ( + router as daily_activity_router, +) +from litellm.proxy.management_endpoints.internal_user_endpoints import router as internal_user_router +from litellm.proxy.management_endpoints.team_endpoints import router as team_router +from litellm.types.proxy.management_endpoints.common_daily_activity import DailyActivityKeyPageResponse + + +def _delete_organization(gateway: Gateway, organization_id: str) -> None: + response: Final = gateway.request("DELETE", "/organization/delete", {"organization_ids": [organization_id]}) + assert response.status_code == 200, response.text + + +def _delete_tag(gateway: Gateway, tag: str) -> None: + response: Final = gateway.request("POST", "/tag/delete", {"name": tag}) + assert response.status_code == 200, response.text + + +def _delete_end_user(gateway: Gateway, end_user_id: str) -> None: + response: Final = gateway.request("POST", "/end_user/delete", {"user_ids": [end_user_id]}) + assert response.status_code == 200, response.text + + +def _delete_agent(gateway: Gateway, agent_id: str) -> None: + response: Final = gateway.request("DELETE", f"/v1/agents/{agent_id}") + assert response.status_code == 200, response.text + + +def _daily_activity_request( + gateway: Gateway, + *, + model: str, + key: str, + end_user_id: str, + tag: str, + request_number: int, +) -> None: + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + JSON_OBJECT.validate_python( + { + "model": model, + "messages": [{"role": "user", "content": f"daily activity {request_number}"}], + "metadata": {"tags": [tag]}, + "user": end_user_id, + } + ), + key=key, + ) + assert response.status_code == 200, response.text + + +def _aggregate_result_api_keys(result: object) -> tuple[str, ...]: + result_body: Final = object_value(result) + breakdown: Final = object_value(result_body["breakdown"]) + api_keys: Final = object_value(breakdown["api_keys"]) + return tuple(api_keys) + + +def _aggregate_top_keys(results: object) -> frozenset[str]: + assert isinstance(results, list) + api_keys_by_result: Final = tuple(_aggregate_result_api_keys(result) for result in results) + return frozenset(chain.from_iterable(api_keys_by_result)) + + +def _assert_entity_activity_routes( + gateway: Gateway, + *, + prefix: str, + entity_param: str, + entity_id: str, + table: str, + entity_column: str, + date_params: dict[str, str], + target_digest: str, + model: str, +) -> None: + params: Final = {**date_params, entity_param: entity_id} + persisted_rows: Final = eventually( + lambda: read_rows( + f'SELECT api_key FROM "{table}" WHERE "{entity_column}"=%s AND date BETWEEN %s AND %s', + (entity_id, date_params["start_date"], date_params["end_date"]), + ), + lambda rows: len(rows) == 6, + seconds=70, + ) + aggregated: Final = gateway.request( + "GET", + f"{prefix}/daily/activity/aggregated", + params={**params, "api_key_limit": "3"}, + ) + assert aggregated.status_code == 200, aggregated.text + aggregate_body: Final = object_value(aggregated.json()) + metadata: Final = object_value(aggregate_body["metadata"]) + total_api_keys: Final = metadata["total_api_keys"] + api_key_limit: Final = metadata["api_key_limit"] + assert isinstance(total_api_keys, int) and total_api_keys == 6, aggregated.text + assert isinstance(api_key_limit, int) and api_key_limit == 3, aggregated.text + assert total_api_keys > api_key_limit, aggregated.text + assert metadata["total_api_requests"] == 8, aggregated.text + top_api_keys: Final = _aggregate_top_keys(aggregate_body["results"]) + assert target_digest not in top_api_keys, aggregated.text + ranked_rows: Final = read_rows( + f'SELECT api_key FROM "{table}" WHERE "{entity_column}"=%s AND date BETWEEN %s AND %s ' + "AND api_key <> %s GROUP BY api_key ORDER BY SUM(spend::numeric) DESC, api_key", + ( + entity_id, + date_params["start_date"], + date_params["end_date"], + constants.PTU_SENTINEL_API_KEY, + ), + ) + ranked_keys: Final = tuple(string_value(row["api_key"]) for row in ranked_rows) + page_responses: Final = tuple( + gateway.request( + "GET", + f"{prefix}/daily/activity/aggregated/keys", + params={**params, "offset": str(offset), "limit": "2"}, + ) + for offset in range(0, len(ranked_keys), 2) + ) + assert all(response.status_code == 200 for response in page_responses), tuple( + response.text for response in page_responses + ) + page_bodies: Final = tuple( + DailyActivityKeyPageResponse.model_validate_json(response.content) for response in page_responses + ) + page_api_keys: Final = tuple(tuple(row.api_key for row in body.api_keys) for body in page_bodies) + paged_keys: Final = tuple(chain.from_iterable(page_api_keys)) + assert tuple(body.total_api_keys for body in page_bodies) == (6,) * len(page_bodies) + assert paged_keys == ranked_keys + assert len(paged_keys) == len(frozenset(paged_keys)) + assert frozenset(paged_keys[:3]) == top_api_keys, aggregated.text + + key_details: Final = gateway.request( + "GET", + f"{prefix}/daily/activity/aggregated", + params={**params, "api_key": target_digest}, + ) + assert key_details.status_code == 200, key_details.text + key_details_body: Final = JSON_OBJECT.validate_json(key_details.content) + assert object_value(key_details_body["metadata"])["total_api_keys"] == 1, key_details.text + assert _aggregate_top_keys(key_details_body["results"]) == frozenset((target_digest,)), key_details.text + + searched: Final = gateway.request( + "GET", + f"{prefix}/daily/activity/aggregated/search", + params={**params, "search": target_digest}, + ) + assert searched.status_code == 200, searched.text + search_body: Final = object_value(searched.json()) + search_rows: Final = search_body["api_keys"] + assert isinstance(search_rows, list) and len(search_rows) == 1, searched.text + assert object_value(search_rows[0])["api_key"] == target_digest, searched.text + + top_keys: Final = gateway.request( + "GET", + f"{prefix}/daily/activity/aggregated/model_top_keys", + params={**params, "model_group": model}, + ) + assert top_keys.status_code == 200, top_keys.text + top_body: Final = object_value(top_keys.json()) + top_rows: Final = top_body["api_keys"] + assert isinstance(top_rows, list) and len(top_rows) == 5, top_keys.text + top_spends: Final = tuple(object_value(object_value(row)["metrics"])["spend"] for row in top_rows[:2]) + assert top_spends == ( + pytest.approx(0.12), + pytest.approx(0.12), + ), top_keys.text + + exported: Final = gateway.request( + "GET", + f"{prefix}/daily/activity/export", + params={**params, "export_type": "daily_with_keys"}, + ) + assert exported.status_code == 200, exported.text + export_rows: Final = tuple(csv.reader(io.StringIO(exported.text))) + assert len(export_rows) == len(persisted_rows) + 1, exported.text + + +@pytest.mark.timeout(90) +def test_daily_activity_routes_cover_all_entities_and_bounded_key_search(gateway: Gateway, tmp_path: Path) -> None: + with owned_proxy(gateway, tmp_path, {}) as proxy: + _assert_daily_activity_routes(proxy) + + +def _assert_daily_activity_routes(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002) + target_scenario_id: Final = f"usage-cache-{uuid.uuid4().hex}" + target_response: Final = JsonResponse( + content_type="application/json", + body=JSON_OBJECT.validate_python( + { + "id": "$UNIQUE_ID", + "object": "chat.completion", + "created": 1_700_000_000, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "cached response"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 40, + "completion_tokens": 20, + "total_tokens": 60, + "prompt_tokens_details": {"cached_tokens": 20}, + }, + } + ), + ) + target_upstream: Final = register_scenario(target_scenario_id, target_response) + scenario.cleanups.callback(delete_scenario, target_upstream) + cache_model: Final = scenario.model( + api_base=target_upstream.api_base(), + api_key=target_scenario_id, + input_cost_per_token=0.0, + output_cost_per_token=0.0, + ) + organization: Final = gateway.post( + "/organization/new", + {"organization_alias": f"integration-{uuid.uuid4().hex}"}, + ) + organization_id: Final = string_value(organization["organization_id"]) + scenario.cleanups.callback(_delete_organization, gateway, organization_id) + team_id: Final = scenario.team(organization_id=organization_id) + user_id: Final = scenario.user() + tag: Final = f"integration-{uuid.uuid4().hex}" + gateway.post("/tag/new", {"name": tag}) + scenario.cleanups.callback(_delete_tag, gateway, tag) + end_user_id: Final = f"integration-{uuid.uuid4().hex}" + gateway.post("/end_user/new", {"user_id": end_user_id}) + scenario.cleanups.callback(_delete_end_user, gateway, end_user_id) + agent_response: Final = gateway.request( + "POST", + "/v1/agents", + { + "agent_name": f"integration-{uuid.uuid4().hex}", + "agent_card_params": { + "protocolVersion": "0.3", + "name": "integration", + "description": "integration agent", + "url": "http://127.0.0.1:1/agent", + "version": "1", + "capabilities": {}, + "defaultInputModes": ["text"], + "defaultOutputModes": ["text"], + "skills": [], + }, + }, + ) + assert agent_response.status_code == 200, agent_response.text + agent_id: Final = string_value(object_value(agent_response.json())["agent_id"]) + scenario.cleanups.callback(_delete_agent, gateway, agent_id) + keys: Final = tuple( + scenario.key( + models=[model, cache_model], + team_id=team_id, + user_id=user_id, + organization_id=organization_id, + agent_id=agent_id, + ) + for _ in range(5) + ) + target_key: Final = scenario.key( + models=[model, cache_model], + team_id=team_id, + user_id=user_id, + organization_id=organization_id, + agent_id=agent_id, + ) + for request_number, key in enumerate(keys): + _daily_activity_request( + gateway, + model=model, + key=key, + end_user_id=end_user_id, + tag=tag, + request_number=request_number, + ) + for request_number, key in enumerate(keys[:2]): + _daily_activity_request( + gateway, + model=model, + key=key, + end_user_id=end_user_id, + tag=tag, + request_number=100 + request_number, + ) + target_request: Final = gateway.request( + "POST", + "/v1/chat/completions", + JSON_OBJECT.validate_python( + { + "model": cache_model, + "messages": [{"role": "user", "content": "cached activity response"}], + "metadata": {"tags": [tag]}, + "user": end_user_id, + } + ), + key=target_key, + ) + assert target_request.status_code == 200, target_request.text + + today: Final = datetime.now(timezone.utc).date() + start_date: Final = (today - timedelta(days=1)).isoformat() + end_date: Final = (today + timedelta(days=1)).isoformat() + date_params: Final = {"start_date": start_date, "end_date": end_date, "timezone": "0"} + route_cases: Final = ( + ("/user", "user_id", user_id, "LiteLLM_DailyUserSpend", "user_id"), + ("/team", "team_ids", team_id, "LiteLLM_DailyTeamSpend", "team_id"), + ("/tag", "tags", tag, "LiteLLM_DailyTagSpend", "tag"), + ( + "/organization", + "organization_ids", + organization_id, + "LiteLLM_DailyOrganizationSpend", + "organization_id", + ), + ("/customer", "end_user_ids", end_user_id, "LiteLLM_DailyEndUserSpend", "end_user_id"), + ("/agent", "agent_ids", agent_id, "LiteLLM_DailyAgentSpend", "agent_id"), + ) + target_digest: Final = hashlib.sha256(target_key.encode()).hexdigest() + for prefix, entity_param, entity_id, table, entity_column in route_cases: + _assert_entity_activity_routes( + gateway, + prefix=prefix, + entity_param=entity_param, + entity_id=entity_id, + table=table, + entity_column=entity_column, + date_params=date_params, + target_digest=target_digest, + model=model, + ) + + user_cache_keys: Final = gateway.request( + "GET", + "/user/daily/activity/aggregated/cache_leakage_keys", + params={**date_params, "user_id": user_id}, + ) + assert user_cache_keys.status_code == 200, user_cache_keys.text + cache_rows: Final = object_value(user_cache_keys.json())["api_keys"] + assert isinstance(cache_rows, list) and cache_rows, user_cache_keys.text + cache_api_keys: Final = tuple(string_value(object_value(row)["api_key"]) for row in cache_rows) + assert target_digest in cache_api_keys, user_cache_keys.text + + +async def _assert_route_matches_golden(client: httpx.AsyncClient, route: str, golden_name: str) -> None: + response: Final = await client.get( + route, + params={"start_date": "2026-06-01", "end_date": "2026-06-01"}, + ) + assert response.status_code == 200, response.text + golden: Final = (Path(__file__).parent / "golden" / golden_name).read_text() + expected: Final = JSON_OBJECT.validate_json(golden) + actual: Final = object_value(response.json()) + assert actual == expected, route + + +@pytest.mark.asyncio +async def test_existing_activity_routes_match_base_branch_goldens(monkeypatch: pytest.MonkeyPatch) -> None: + async with _daily_activity_database() as database: + repository: Final = _repository(database) + app: Final = FastAPI() + app.include_router(internal_user_router) + app.include_router(team_router) + app.include_router(daily_activity_router) + monkeypatch.setattr(proxy_server, "prisma_client", _PrismaDatabase(database)) + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="integration-admin", user_role=LitellmUserRoles.PROXY_ADMIN + ) + app.dependency_overrides[get_daily_activity_prisma_client] = lambda: _PrismaDatabase(database) + app.dependency_overrides[get_daily_activity_repository] = lambda: repository + + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://testserver") as client: + route_goldens: Final = ( + ("/user/daily/activity", "daily_activity_user_paginated.json"), + ("/user/daily/activity/aggregated", "daily_activity_user_aggregated.json"), + ("/team/daily/activity", "daily_activity_team_paginated.json"), + ("/team/daily/activity/aggregated", "daily_activity_team_aggregated.json"), + ) + for route, golden_name in route_goldens: + await _assert_route_matches_golden(client, route, golden_name) + + +@pytest.mark.asyncio +async def test_user_key_pages_and_details_respect_caller_scope() -> None: + async with _daily_activity_database() as database: + await database.query_raw( + 'INSERT INTO "LiteLLM_UserTable" (user_id, user_email, models) VALUES ($1, $2, $3)', + "user-2", + "other@example.test", + [], + ) + await database.query_raw( + """ + INSERT INTO "LiteLLM_VerificationToken" + (token, key_alias, team_id, user_id, metadata, models) + VALUES ($1, $2, $3, $4, $5::jsonb, $6) + """, + "key-other-user", + "Other user key", + None, + "user-2", + "{}", + [], + ) + await database.query_raw( + """ + INSERT INTO "LiteLLM_DailyUserSpend" + (id, user_id, date, api_key, model, model_group, custom_llm_provider, + mcp_namespaced_tool_name, endpoint, prompt_tokens, completion_tokens, + cache_read_input_tokens, cache_creation_input_tokens, spend, api_requests, + successful_requests, failed_requests, updated_at) + VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11, $12, $13, $14, $15, $16, $17, $18::timestamp) + """, + "other-user-row", + "user-2", + "2026-06-01", + "key-other-user", + "model", + "", + "provider-a", + None, + "/v1/chat/completions", + 1, + 1, + 0, + 0, + 50.0, + 1, + 1, + 0, + "2026-06-01 12:00:00", + ) + repository: Final = _repository(database) + app: Final = FastAPI() + app.include_router(daily_activity_router) + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="user-1", + user_role=LitellmUserRoles.INTERNAL_USER, + ) + app.dependency_overrides[get_daily_activity_prisma_client] = lambda: _PrismaDatabase(database) + app.dependency_overrides[get_daily_activity_repository] = lambda: repository + + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), + base_url="http://testserver", + ) as client: + params: Final = {"start_date": "2026-06-01", "end_date": "2026-06-01"} + page: Final = await client.get( + "/user/daily/activity/aggregated/keys", + params={**params, "user_id": "user-1", "limit": 100}, + ) + assert page.status_code == 200, page.text + page_body: Final = DailyActivityKeyPageResponse.model_validate_json(page.content) + page_keys: Final = frozenset(row.api_key for row in page_body.api_keys) + assert page_body.total_api_keys == len(page_keys) == 5 + assert "key-other-user" not in page_keys + + denied: Final = await client.get( + "/user/daily/activity/aggregated/keys", + params={**params, "user_id": "user-2"}, + ) + assert denied.status_code == 403, denied.text + + own_details: Final = await client.get( + "/user/daily/activity/aggregated", + params={**params, "user_id": "user-1", "api_key": "key-a"}, + ) + assert own_details.status_code == 200, own_details.text + own_body: Final = JSON_OBJECT.validate_json(own_details.content) + assert object_value(own_body["metadata"])["total_api_keys"] == 1 + assert _aggregate_top_keys(own_body["results"]) == frozenset(("key-a",)) + + other_details: Final = await client.get( + "/user/daily/activity/aggregated", + params={**params, "user_id": "user-1", "api_key": "key-other-user"}, + ) + assert other_details.status_code == 200, other_details.text + other_body: Final = JSON_OBJECT.validate_json(other_details.content) + assert object_value(other_body["metadata"])["total_api_keys"] == 0 + assert _aggregate_top_keys(other_body["results"]) == frozenset() + + +@pytest.mark.asyncio +async def test_team_routes_exclusion_keeps_unassigned_keys() -> None: + async with _daily_activity_database(include_team_exclusion_activity=True) as database: + repository: Final = _repository(database) + app: Final = FastAPI() + app.include_router(daily_activity_router) + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="integration-admin", user_role=LitellmUserRoles.PROXY_ADMIN + ) + app.dependency_overrides[get_daily_activity_prisma_client] = lambda: _PrismaDatabase(database) + app.dependency_overrides[get_daily_activity_repository] = lambda: repository + + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://testserver") as client: + params: Final = { + "start_date": "2026-06-04", + "end_date": "2026-06-04", + "exclude_team_ids": "litellm-dashboard", + } + surviving_keys: Final = frozenset(("key-excluded-null", "key-excluded-empty", "key-excluded-normal")) + + aggregated: Final = await client.get("/team/daily/activity/aggregated", params=params) + assert aggregated.status_code == 200, aggregated.text + aggregated_body: Final = JSON_OBJECT.validate_json(aggregated.content) + assert object_value(aggregated_body["metadata"])["total_spend"] == 23.0 + assert object_value(aggregated_body["metadata"])["total_api_keys"] == 3 + assert _aggregate_top_keys(aggregated_body["results"]) == surviving_keys + + page: Final = await client.get("/team/daily/activity/aggregated/keys", params={**params, "limit": 10}) + assert page.status_code == 200, page.text + page_body: Final = DailyActivityKeyPageResponse.model_validate_json(page.content) + assert page_body.total_api_keys == 3 + assert frozenset(row.api_key for row in page_body.api_keys) == surviving_keys diff --git a/tests/integration/spend/test_key_metadata_recovery_probe_bounds.py b/tests/integration/spend/test_key_metadata_recovery_probe_bounds.py index 81be9e56f30..fd67721cc6b 100644 --- a/tests/integration/spend/test_key_metadata_recovery_probe_bounds.py +++ b/tests/integration/spend/test_key_metadata_recovery_probe_bounds.py @@ -1,12 +1,11 @@ from collections.abc import Mapping, Sequence from dataclasses import dataclass from datetime import datetime, timedelta -from pathlib import Path from typing import Final -import litellm_proxy_extras import psycopg import pytest +from litellm_proxy_extras.request_log_indexes import REQUEST_LOG_INDEXES from psycopg.types.json import Jsonb from pydantic import JsonValue @@ -29,11 +28,8 @@ _SPEND_LOGS_DDL: Final = """ ) """ -_API_KEY_START_TIME_INDEX_MIGRATION: Final = ( - Path(litellm_proxy_extras.__file__).parent - / "migrations" - / "20260823000000_add_spend_logs_api_key_starttime_index" - / "migration.sql" +_API_KEY_START_TIME_INDEX: Final = next( + index for index in REQUEST_LOG_INDEXES if index.name == "LiteLLM_SpendLogs_api_key_startTime_idx" ) _STATS_SQL: Final = """ @@ -56,7 +52,12 @@ class _Settle: def _create_spend_logs_table(database_url: str) -> None: write_rows(_SPEND_LOGS_DDL, (), database_url=database_url) - write_rows(_API_KEY_START_TIME_INDEX_MIGRATION.read_text(), (), database_url=database_url) + write_rows( + f'CREATE INDEX "{_API_KEY_START_TIME_INDEX.name}" ON "{_API_KEY_START_TIME_INDEX.table}" ' # pyright: ignore[reportArgumentType] # DDL from the migration job index list + f"{_API_KEY_START_TIME_INDEX.definition}", + (), + database_url=database_url, + ) def _spend_log_stats(database_url: str) -> dict[str, int]: diff --git a/tests/integration/spend/test_lens_billing.py b/tests/integration/spend/test_lens_billing.py index bedcf6c5380..d8eded62b39 100644 --- a/tests/integration/spend/test_lens_billing.py +++ b/tests/integration/spend/test_lens_billing.py @@ -124,6 +124,9 @@ def test_lens_bills_selected_key_and_rechecks_its_permissions(gateway: Gateway, seconds=70, ) assert second_rows[0]["spend"] == pytest.approx(expected) + active_revoke: Final = gateway.request("DELETE", f"/lens/workers/{worker_id}") + assert active_revoke.status_code == 409, active_revoke.text + gateway.post(f"/lens/{lens_id}/cancel", {}) revoked: Final = gateway.request("DELETE", f"/lens/workers/{worker_id}") assert revoked.status_code == 200, revoked.text denied_worker: Final = gateway.request( @@ -134,7 +137,6 @@ def test_lens_bills_selected_key_and_rechecks_its_permissions(gateway: Gateway, "PUT", f"/lens/workers/{worker_id}/billing-key", {"analysis_key_id": replacement_id} ) assert forbidden_change.status_code == 409, forbidden_change.text - gateway.post(f"/lens/{lens_id}/cancel", {}) @pytest.mark.parametrize("cancel_on_disconnect", (False, True)) diff --git a/tests/proxy_migration_tests/test_db_schema_migration.py b/tests/proxy_migration_tests/test_db_schema_migration.py index b0d44cd3e1c..70498a3bcbd 100644 --- a/tests/proxy_migration_tests/test_db_schema_migration.py +++ b/tests/proxy_migration_tests/test_db_schema_migration.py @@ -5,6 +5,7 @@ import tempfile from pathlib import Path import pytest +from litellm_proxy_extras.request_log_indexes import filter_request_log_index_diff @pytest.mark.skipif( @@ -16,7 +17,9 @@ def test_schema_migration_in_sync(): Applies every committed migration to an empty database, then diffs the result against schema.prisma. A non-empty diff means the schema was changed without a - matching migration being generated. + matching migration being generated. The request-log indexes the migration job + builds are declared in the schema and deliberately absent from the migrations, + so those statements are filtered out before the diff is judged. """ db_url = os.environ["DATABASE_URL"] source_migrations_dir = Path( @@ -60,11 +63,14 @@ def test_schema_migration_in_sync(): ) if diff.returncode == 2: - pytest.fail( - "Schema changes detected that no migration captures. Run " - "`python litellm/ci_cd/run_migration.py `.\n\n" - + diff.stdout - ) - assert diff.returncode == 0, f"prisma migrate diff errored: {diff.stderr}" + drift = filter_request_log_index_diff(diff.stdout) + if drift.strip(): + pytest.fail( + "Schema changes detected that no migration captures. Run " + "`python litellm/ci_cd/run_migration.py `.\n\n" + + drift + ) + else: + assert diff.returncode == 0, f"prisma migrate diff errored: {diff.stderr}" finally: shutil.rmtree(temp_base, ignore_errors=True) diff --git a/tests/proxy_migration_tests/test_invalid_index_repair.py b/tests/proxy_migration_tests/test_invalid_index_repair.py index 741fa7386df..0c971b5d073 100644 --- a/tests/proxy_migration_tests/test_invalid_index_repair.py +++ b/tests/proxy_migration_tests/test_invalid_index_repair.py @@ -1,12 +1,14 @@ import os import threading +import time import uuid from collections.abc import Iterator, Mapping from types import MappingProxyType from typing import Final import pytest -from litellm_proxy_extras.utils import INDEX_REPAIR_ADVISORY_LOCK_KEY, ProxyExtrasDBManager +from litellm_proxy_extras.migration_lock import MIGRATION_LOCK_KEY +from litellm_proxy_extras.utils import INDEX_REPAIR_ADVISORY_LOCK_KEY, ProxyExtrasDBManager, _InvalidIndex psycopg = pytest.importorskip("psycopg") @@ -20,6 +22,7 @@ requires_db: Final = pytest.mark.skipif( HEALTH_TABLE: Final = "LiteLLM_HealthCheckTable" HEALTH_INDEX: Final = "LiteLLM_HealthCheckTable_model_id_model_name_checked_at_idx" HEALTH_INDEX_COLUMNS: Final = '"model_id", "model_name", "checked_at" DESC' +SECOND_HEALTH_INDEX: Final = "LiteLLM_HealthCheckTable_model_name_idx" LOOKALIKE_TABLE: Final = "LiteLLMLookalikeTable" LOOKALIKE_INDEX: Final = "LiteLLMLookalikeTable_id_idx" PARTITIONED_TABLE: Final = "LiteLLM_PartitionedTable" @@ -167,7 +170,74 @@ def test_repair_yields_to_the_replica_holding_the_repair_lock(scratch_schema: st @requires_db -def test_repair_gives_up_on_a_blocked_rebuild_and_finishes_it_on_the_next_startup(scratch_schema: str) -> None: +def test_repair_yields_to_the_migration_job_building_indexes_under_the_migration_lock(scratch_schema: str) -> None: + """A migration job's index build holds the migration lock while its CREATE INDEX CONCURRENTLY + is cataloged as invalid; the repair must not rebuild that in-flight index.""" + _leave_invalid_index(scratch_schema, HEALTH_TABLE, HEALTH_INDEX, HEALTH_INDEX_COLUMNS) + + with psycopg.connect(_base_url(), autocommit=True) as index_builder: + index_builder.execute("SELECT pg_advisory_lock(%s)", (MIGRATION_LOCK_KEY,)) + assert ProxyExtrasDBManager.repair_invalid_indexes() is False + assert _index_validity(scratch_schema) == {HEALTH_INDEX: False} + + assert ProxyExtrasDBManager.repair_invalid_indexes() is True + assert _index_validity(scratch_schema) == {HEALTH_INDEX: True} + + +def _hold_migration_lock_once_free(release: threading.Event) -> None: + with psycopg.connect(_base_url(), autocommit=True) as resolver: + resolver.execute("SELECT pg_advisory_lock(%s)", (MIGRATION_LOCK_KEY,)) + release.wait(timeout=60) + + +def _wait_until_a_session_queues_for_the_migration_lock() -> None: + with psycopg.connect(_base_url(), autocommit=True) as conn: + for _ in range(200): + queued: Final = conn.execute( + "SELECT count(*) FROM pg_locks WHERE locktype = 'advisory' AND NOT granted " + "AND classid = %s AND objid = %s", + (MIGRATION_LOCK_KEY >> 32, MIGRATION_LOCK_KEY & 0xFFFFFFFF), + ).fetchone() + if queued is not None and queued[0]: + return + time.sleep(0.05) + pytest.fail("no session queued for the migration lock") + + +@requires_db +def test_repair_releases_the_migration_lock_between_indexes_so_a_booting_resolver_gets_in( + scratch_schema: str, +) -> None: + """A v2 resolver on another replica waits for the migration lock; with two invalid + indexes to rebuild it must get the lock after the first REINDEX, not after both.""" + _leave_invalid_index(scratch_schema, HEALTH_TABLE, HEALTH_INDEX, HEALTH_INDEX_COLUMNS) + _leave_invalid_index(scratch_schema, HEALTH_TABLE, SECOND_HEALTH_INDEX, '"model_name"') + release: Final = threading.Event() + resolver: Final = threading.Thread(target=_hold_migration_lock_once_free, args=(release,)) + + def repair_then_let_a_resolver_queue_for_the_lock( + conn: "psycopg.Connection[tuple[str, str, str]]", index: _InvalidIndex + ) -> None: + ProxyExtrasDBManager._repair_index(conn, index) + if not resolver.is_alive(): + resolver.start() + _wait_until_a_session_queues_for_the_migration_lock() + + try: + assert ( + ProxyExtrasDBManager.repair_invalid_indexes(repair=repair_then_let_a_resolver_queue_for_the_lock) is False + ) + assert sorted(_index_validity(scratch_schema).values()) == [False, True] + finally: + release.set() + resolver.join() + + assert ProxyExtrasDBManager.repair_invalid_indexes() is True + assert _index_validity(scratch_schema) == {HEALTH_INDEX: True, SECOND_HEALTH_INDEX: True} + + +@requires_db +def test_repair_gives_up_on_a_blocked_rebuild_and_finishes_it_on_the_next_boot(scratch_schema: str) -> None: _leave_invalid_index(scratch_schema, HEALTH_TABLE, HEALTH_INDEX, HEALTH_INDEX_COLUMNS) with psycopg.connect(_base_url()) as pin: diff --git a/tests/proxy_migration_tests/test_request_log_indexes.py b/tests/proxy_migration_tests/test_request_log_indexes.py new file mode 100644 index 00000000000..23e4adce477 --- /dev/null +++ b/tests/proxy_migration_tests/test_request_log_indexes.py @@ -0,0 +1,912 @@ +import os +import queue +import shutil +import subprocess +import sys +import threading +import time +import uuid +from collections.abc import Callable, Iterator, Mapping +from dataclasses import dataclass +from pathlib import Path +from types import MappingProxyType +from typing import Final + +import psycopg +import pytest +from litellm_proxy_extras import request_log_indexes +from litellm_proxy_extras.migration_lock import MIGRATION_LOCK_KEY, migration_lock +from litellm_proxy_extras.migration_recovery import roll_back_failed_inert_migration +from litellm_proxy_extras.request_log_indexes import ( + REQUEST_LOG_INDEXES, + RequestLogIndex, + build_index_on_partitioned_table, + ensure_request_log_indexes, +) +from litellm_proxy_extras.utils import ProxyExtrasDBManager +from psycopg import sql +from psycopg.abc import Params, QueryNoTemplate +from psycopg.rows import class_row + +pytestmark = pytest.mark.timeout(900) + +requires_db: Final = pytest.mark.skipif( + "DATABASE_URL" not in os.environ, + reason="requires a postgres database (DATABASE_URL)", +) + +REPO: Final = Path(__file__).resolve().parents[2] +PACKAGE: Final = REPO / "litellm-proxy-extras" / "litellm_proxy_extras" +PARTITION_SCRIPT: Final = REPO / "db_scripts" / "partition_spend_logs.sql" +API_KEY_INDEX_MIGRATION: Final = "20260823000000_add_spend_logs_api_key_starttime_index" +CALL_ID_INDEX_MIGRATION: Final = "20260831120001_spend_logs_litellm_call_id_index" +API_KEY_INDEX: Final = "LiteLLM_SpendLogs_api_key_startTime_idx" +CALL_ID_INDEX: Final = "LiteLLM_SpendLogs_litellm_call_id_idx" +PARTITIONED_PARENT_ERROR: Final = 'cannot create index on partitioned table "LiteLLM_SpendLogs" concurrently' +ORIGINAL_MIGRATION_SQL: Final = MappingProxyType( + { + API_KEY_INDEX_MIGRATION: ( + "-- CreateIndex\n" + 'CREATE INDEX IF NOT EXISTS "LiteLLM_SpendLogs_api_key_startTime_idx" ' + 'ON "LiteLLM_SpendLogs"("api_key", "startTime");\n' + ), + CALL_ID_INDEX_MIGRATION: ( + "-- CreateIndex\n" + 'CREATE INDEX CONCURRENTLY IF NOT EXISTS "LiteLLM_SpendLogs_litellm_call_id_idx" ' + 'ON "LiteLLM_SpendLogs"("litellm_call_id");\n' + ), + } +) +CALL_ID_INDEX_DEFINITION: Final = next(index for index in REQUEST_LOG_INDEXES if index.name == CALL_ID_INDEX) +RELEASES: Final = pytest.mark.parametrize( + "release", (API_KEY_INDEX_MIGRATION, CALL_ID_INDEX_MIGRATION), ids=("v1.102.1", "v1.103.0") +) +PARTITIONS: Final = MappingProxyType( + { + "LiteLLM_SpendLogs_p2026_08": ("2026-08-01", "2026-09-01"), + "LiteLLM_SpendLogs_p2026_09": ("2026-09-01", "2026-10-01"), + } +) +DEFAULT_PARTITION: Final = "LiteLLM_SpendLogs_pdefault" +ROWS_PER_PARTITION: Final = 200 +RESOLVERS: Final = pytest.mark.parametrize("use_v2_resolver", (True, False), ids=("v2", "v1")) + + +def _base_url() -> str: + return os.environ["DATABASE_URL"].split("?")[0] + + +def _migrate_deploy(database_url: str, schema: Path) -> "subprocess.CompletedProcess[str]": + return subprocess.run( + [sys.executable, "-I", "-m", "prisma", "migrate", "deploy", "--schema", str(schema)], + capture_output=True, + text=True, + env={**os.environ, "DATABASE_URL": database_url}, + ) + + +def _release_layout(prisma_dir: Path, before: str) -> Path: + """The shipped migrations older than `before`, with the two index migrations written + the way the releases that shipped them did: the Prisma layout of a proxy on that release.""" + (prisma_dir / "migrations").mkdir(parents=True) + shutil.copy(PACKAGE / "schema.prisma", prisma_dir / "schema.prisma") + for migration in sorted((PACKAGE / "migrations").iterdir()): + if migration.is_dir() and migration.name < before: + shutil.copytree(migration, prisma_dir / "migrations" / migration.name) + for name, original in ORIGINAL_MIGRATION_SQL.items(): + if (prisma_dir / "migrations" / name).is_dir(): + (prisma_dir / "migrations" / name / "migration.sql").write_text(original) + return prisma_dir / "schema.prisma" + + +def _deploy_release(database_url: str, prisma_dir: Path, before: str) -> None: + deployed: Final = _migrate_deploy(database_url, _release_layout(prisma_dir, before)) + assert deployed.returncode == 0, deployed.stderr + + +def _insert_spend_log( + conn: "psycopg.Connection[tuple[object, ...]]", request_id: str, day: str, table: str = "LiteLLM_SpendLogs" +) -> None: + conn.execute( + sql.SQL( + 'INSERT INTO {} ("request_id", "call_type", "startTime", "endTime", "api_key") VALUES (%s, %s, %s, %s, %s)' + ).format(sql.Identifier(table)), + (request_id, "acompletion", day, day, f"key-{request_id[-1]}"), + ) + + +def _partition_spend_logs(database_url: str) -> None: + with psycopg.connect(database_url, autocommit=True) as conn: + conn.execute(PARTITION_SCRIPT.read_bytes()) + for partition, (start, stop) in PARTITIONS.items(): + conn.execute( + sql.SQL('CREATE TABLE {} PARTITION OF "LiteLLM_SpendLogs" FOR VALUES FROM ({}) TO ({})').format( + sql.Identifier(partition), sql.Literal(start), sql.Literal(stop) + ) + ) + for row in range(ROWS_PER_PARTITION): + _insert_spend_log(conn, f"{partition}-{row}", start) + for row in range(ROWS_PER_PARTITION): + _insert_spend_log(conn, f"default-{row}", "2020-01-01") + + +@pytest.fixture +def release() -> str: + """The first migration a database has not applied yet; the v1.103.0 shape unless a test parametrizes it.""" + return CALL_ID_INDEX_MIGRATION + + +@pytest.fixture +def scratch_database(release: str, monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> Iterator[str]: + """A deployment stopped before `release`, with DATABASE_URL pointed at it so + ProxyExtrasDBManager upgrades it like a booting proxy.""" + admin_url: Final = _base_url() + name: Final = f"spend_logs_index_{uuid.uuid4().hex[:8]}" + with psycopg.connect(admin_url, autocommit=True) as conn: + conn.execute(sql.SQL("CREATE DATABASE {}").format(sql.Identifier(name))) + database_url: Final = f"{admin_url.rsplit('/', 1)[0]}/{name}" + try: + _deploy_release(database_url, tmp_path / "prisma", release) + monkeypatch.delenv("DIRECT_URL", raising=False) + monkeypatch.setenv("DATABASE_URL", database_url) + yield database_url + finally: + with psycopg.connect(admin_url, autocommit=True) as conn: + conn.execute(sql.SQL("DROP DATABASE {} WITH (FORCE)").format(sql.Identifier(name))) + + +@pytest.fixture +def partitioned_database(scratch_database: str) -> str: + _partition_spend_logs(scratch_database) + return scratch_database + + +def _fail_the_call_id_migration_like_the_shipped_release(database_url: str, tmp_path: Path) -> None: + """Boot the original v1.103.0 layout once: its CONCURRENTLY statement fails on the + partitioned parent and leaves the call_id ledger row unfinished.""" + failed: Final = _migrate_deploy(database_url, _release_layout(tmp_path / "v1.103.0", "99999999999999")) + assert failed.returncode != 0 and PARTITIONED_PARENT_ERROR in failed.stderr, failed.stderr + assert _ledger(database_url)[CALL_ID_INDEX_MIGRATION] == (False, False) + + +@dataclass(frozen=True, slots=True) +class _IndexRow: + name: str + valid: bool + + +@dataclass(frozen=True, slots=True) +class _AttachedRow: + table: str + index: str + + +@dataclass(frozen=True, slots=True) +class _LedgerRow: + name: str + finished: bool + rolled_back: bool + + +@dataclass(frozen=True, slots=True) +class _OidRow: + name: str + oid: int + + +def _index_validity(database_url: str, suffix: str) -> Mapping[str, bool]: + """index name -> indisvalid for every index ending in `suffix` on the SpendLogs parent or one of its partitions.""" + with psycopg.connect(database_url) as conn, conn.cursor(row_factory=class_row(_IndexRow)) as cursor: + rows: Final = cursor.execute( + "SELECT c.relname AS name, i.indisvalid AS valid FROM pg_index i JOIN pg_class c ON c.oid = i.indexrelid " + "WHERE c.relname LIKE %s AND (i.indrelid = to_regclass('\"LiteLLM_SpendLogs\"') OR i.indrelid IN " + "(SELECT inhrelid FROM pg_inherits WHERE inhparent = to_regclass('\"LiteLLM_SpendLogs\"'))) " + "ORDER BY c.relname", + (f"%{suffix}",), + ).fetchall() + return MappingProxyType({row.name: row.valid for row in rows}) + + +def _attached_children(database_url: str, parent_index: str) -> frozenset[tuple[str, str]]: + """(partition, child index) pairs attached under the parent index.""" + with psycopg.connect(database_url) as conn, conn.cursor(row_factory=class_row(_AttachedRow)) as cursor: + rows: Final = cursor.execute( + 'SELECT t.relname AS "table", c.relname AS index FROM pg_inherits i ' + "JOIN pg_class c ON c.oid = i.inhrelid JOIN pg_index x ON x.indexrelid = c.oid " + "JOIN pg_class t ON t.oid = x.indrelid " + "WHERE i.inhparent = to_regclass(%s)", + (f'"{parent_index}"',), + ).fetchall() + return frozenset((row.table, row.index) for row in rows) + + +@dataclass(frozen=True, slots=True) +class _TableRow: + name: str + + +def _indexed_table(database_url: str, index: str) -> "str | None": + with psycopg.connect(database_url) as conn, conn.cursor(row_factory=class_row(_TableRow)) as cursor: + row: Final = cursor.execute( + "SELECT t.relname AS name FROM pg_index x JOIN pg_class t ON t.oid = x.indrelid " + "WHERE x.indexrelid = to_regclass(%s)", + (f'"{index}"',), + ).fetchone() + return None if row is None else row.name + + +def _ledger(database_url: str) -> Mapping[str, tuple[bool, bool]]: + """migration name -> (finished, rolled back) for the newest ledger row of each migration.""" + with psycopg.connect(database_url) as conn, conn.cursor(row_factory=class_row(_LedgerRow)) as cursor: + rows: Final = cursor.execute( + "SELECT DISTINCT ON (migration_name) migration_name AS name, finished_at IS NOT NULL AS finished, " + "rolled_back_at IS NOT NULL AS rolled_back FROM _prisma_migrations ORDER BY migration_name, started_at DESC" + ).fetchall() + return MappingProxyType({row.name: (row.finished, row.rolled_back) for row in rows}) + + +def _index_oids(database_url: str) -> Mapping[str, int]: + """index name -> oid for every index on the SpendLogs parent or one of its partitions; a rebuild changes the oid.""" + with psycopg.connect(database_url) as conn, conn.cursor(row_factory=class_row(_OidRow)) as cursor: + rows: Final = cursor.execute( + "SELECT c.relname AS name, c.oid::int AS oid FROM pg_index i JOIN pg_class c ON c.oid = i.indexrelid " + "WHERE i.indrelid = to_regclass('\"LiteLLM_SpendLogs\"') OR i.indrelid IN " + "(SELECT inhrelid FROM pg_inherits WHERE inhparent = to_regclass('\"LiteLLM_SpendLogs\"'))" + ).fetchall() + return MappingProxyType({row.name: row.oid for row in rows}) + + +def _migration_job(use_v2_resolver: bool) -> bool: + return ProxyExtrasDBManager.run_migration_job(use_migrate=True, use_v2_resolver=use_v2_resolver) + + +def _assert_no_pending_migrations(database_url: str) -> None: + status: Final = _migrate_deploy(database_url, PACKAGE / "schema.prisma") + assert status.returncode == 0 and "No pending migrations" in status.stdout, status.stdout + status.stderr + + +def _assert_every_ledger_row_is_finished(database_url: str) -> Mapping[str, tuple[bool, bool]]: + ledger: Final = _ledger(database_url) + assert ledger[API_KEY_INDEX_MIGRATION] == (True, False) and ledger[CALL_ID_INDEX_MIGRATION] == (True, False) + assert all(finished and not rolled_back for finished, rolled_back in ledger.values()), ledger + return ledger + + +def _expected_children(partitions: tuple[str, ...], suffix: str) -> frozenset[tuple[str, str]]: + return frozenset((partition, f"{partition}_{suffix}") for partition in partitions) + + +def _assert_index_covers_every_partition(database_url: str, parent_index: str, suffix: str) -> None: + partitions: Final = (*PARTITIONS, DEFAULT_PARTITION) + assert _index_validity(database_url, suffix) == {parent_index: True} | {f"{p}_{suffix}": True for p in partitions} + assert _attached_children(database_url, parent_index) == _expected_children(partitions, suffix) + + +@requires_db +@RESOLVERS +@RELEASES +def test_a_partitioned_spend_logs_upgrade_builds_both_indexes_per_partition_and_a_rerun_is_idempotent( + partitioned_database: str, use_v2_resolver: bool +) -> None: + assert _migration_job(use_v2_resolver) is True + + _assert_index_covers_every_partition(partitioned_database, API_KEY_INDEX, "api_key_startTime_idx") + _assert_index_covers_every_partition(partitioned_database, CALL_ID_INDEX, "litellm_call_id_idx") + ledger: Final = _assert_every_ledger_row_is_finished(partitioned_database) + _assert_no_pending_migrations(partitioned_database) + oids: Final = _index_oids(partitioned_database) + + with psycopg.connect(partitioned_database, autocommit=True) as conn: + conn.execute( + 'CREATE TABLE "LiteLLM_SpendLogs_p2026_10" PARTITION OF "LiteLLM_SpendLogs" ' + "FOR VALUES FROM ('2026-10-01') TO ('2026-11-01')" + ) + inherited: Final = frozenset( + ("LiteLLM_SpendLogs_p2026_10", f"LiteLLM_SpendLogs_p2026_10_{suffix}") + for suffix in ("api_key_startTime_idx", "litellm_call_id_idx") + ) + attached: Final = _attached_children(partitioned_database, API_KEY_INDEX) | _attached_children( + partitioned_database, CALL_ID_INDEX + ) + assert inherited <= attached, attached + + assert _migration_job(use_v2_resolver) is True + assert _ledger(partitioned_database) == ledger + assert {name: oid for name, oid in _index_oids(partitioned_database).items() if name in oids} == oids + + +@requires_db +@RESOLVERS +@RELEASES +def test_a_plain_spend_logs_upgrade_builds_both_indexes_and_a_second_job_run_rebuilds_nothing( + scratch_database: str, use_v2_resolver: bool +) -> None: + with psycopg.connect(scratch_database, autocommit=True) as conn: + for row in range(ROWS_PER_PARTITION): + _insert_spend_log(conn, f"flat-{row}", "2026-09-01") + + assert _migration_job(use_v2_resolver) is True + + assert _index_validity(scratch_database, "api_key_startTime_idx") == {API_KEY_INDEX: True} + assert _index_validity(scratch_database, "litellm_call_id_idx") == {CALL_ID_INDEX: True} + _assert_every_ledger_row_is_finished(scratch_database) + _assert_no_pending_migrations(scratch_database) + oids: Final = _index_oids(scratch_database) + + assert _migration_job(use_v2_resolver) is True + assert _index_oids(scratch_database) == oids + + +@requires_db +@RESOLVERS +def test_a_database_that_applied_the_original_migration_files_sees_no_pending_migrations_and_no_rebuild( + scratch_database: str, use_v2_resolver: bool, tmp_path: Path +) -> None: + """A plain table upgraded on v1.103.0 applied both original files. The inert files in + this build must neither re-run nor fail those rows, and the migration job must keep the + indexes the migrations built.""" + deployed: Final = _migrate_deploy(scratch_database, _release_layout(tmp_path / "v1.103.0", "99999999999999")) + assert deployed.returncode == 0, deployed.stderr + before: Final = _ledger(scratch_database) + assert before[API_KEY_INDEX_MIGRATION] == (True, False) and before[CALL_ID_INDEX_MIGRATION] == (True, False) + oids: Final = _index_oids(scratch_database) + assert {API_KEY_INDEX, CALL_ID_INDEX} <= set(oids) + + _assert_no_pending_migrations(scratch_database) + assert _migration_job(use_v2_resolver) is True + + assert _ledger(scratch_database) == before + assert _index_oids(scratch_database) == oids + + +@requires_db +@RESOLVERS +def test_a_failed_call_id_ledger_row_from_a_v1_103_boot_is_rolled_back_and_the_inert_file_applied( + partitioned_database: str, use_v2_resolver: bool, tmp_path: Path +) -> None: + _fail_the_call_id_migration_like_the_shipped_release(partitioned_database, tmp_path) + + assert _migration_job(use_v2_resolver) is True + + _assert_every_ledger_row_is_finished(partitioned_database) + _assert_index_covers_every_partition(partitioned_database, API_KEY_INDEX, "api_key_startTime_idx") + _assert_index_covers_every_partition(partitioned_database, CALL_ID_INDEX, "litellm_call_id_idx") + _assert_no_pending_migrations(partitioned_database) + with psycopg.connect(partitioned_database) as conn: + rows: Final = conn.execute( + "SELECT finished_at IS NOT NULL, rolled_back_at IS NOT NULL FROM _prisma_migrations " + "WHERE migration_name = %s ORDER BY started_at", + (CALL_ID_INDEX_MIGRATION,), + ).fetchall() + assert rows == [(False, True), (True, False)], rows + + +@requires_db +def test_a_failed_row_whose_migration_still_runs_sql_in_this_build_is_left_for_the_operator( + partitioned_database: str, tmp_path: Path +) -> None: + _fail_the_call_id_migration_like_the_shipped_release(partitioned_database, tmp_path) + still_building: Final = tmp_path / "edited" / CALL_ID_INDEX_MIGRATION / "migration.sql" + still_building.parent.mkdir(parents=True) + still_building.write_text(ORIGINAL_MIGRATION_SQL[CALL_ID_INDEX_MIGRATION]) + + with migration_lock(partitioned_database) as coordinator: + assert roll_back_failed_inert_migration(coordinator, "public", still_building) is False + + assert _ledger(partitioned_database)[CALL_ID_INDEX_MIGRATION] == (False, False) + + +@requires_db +def test_a_migration_without_a_failed_row_is_not_touched(partitioned_database: str) -> None: + inert: Final = PACKAGE / "migrations" / CALL_ID_INDEX_MIGRATION / "migration.sql" + before: Final = _ledger(partitioned_database) + + with migration_lock(partitioned_database) as coordinator: + assert roll_back_failed_inert_migration(coordinator, "public", inert) is False + + assert _ledger(partitioned_database) == before + + +def _pin_a_snapshot_on(database_url: str, table: str) -> "psycopg.Connection[tuple[object, ...]]": + pin: Final = psycopg.connect(database_url) + pin.isolation_level = psycopg.IsolationLevel.REPEATABLE_READ + pin.execute(sql.SQL("SELECT count(*) FROM {}").format(sql.Identifier(table))) + return pin + + +def _leave_an_invalid_index(database_url: str, name: str, table: str, column: str) -> None: + with _pin_a_snapshot_on(database_url, table): + with psycopg.connect(database_url, autocommit=True) as builder: + builder.execute("SET statement_timeout = '1s'") + with pytest.raises(psycopg.errors.QueryCanceled): + builder.execute( + sql.SQL("CREATE INDEX CONCURRENTLY {} ON {} ({})").format( + sql.Identifier(name), sql.Identifier(table), sql.Identifier(column) + ) + ) + + +@requires_db +def test_an_invalid_index_of_the_managed_name_on_a_plain_table_is_rebuilt(scratch_database: str) -> None: + _leave_an_invalid_index(scratch_database, CALL_ID_INDEX, "LiteLLM_SpendLogs", "litellm_call_id") + assert _index_validity(scratch_database, "litellm_call_id_idx") == {CALL_ID_INDEX: False} + + assert ensure_request_log_indexes(scratch_database, "public") is True + + assert _index_validity(scratch_database, "litellm_call_id_idx") == {CALL_ID_INDEX: True} + + +def _rebuild_as_another_replica(database_url: str, name: str, table: str, column: str) -> int: + """Drop and rebuild the index from a second connection, as a replica that won the + race would, and return the oid of the index it built.""" + with psycopg.connect(database_url, autocommit=True) as other_replica: + other_replica.execute(sql.SQL("DROP INDEX {}").format(sql.Identifier(name))) + other_replica.execute( + sql.SQL("CREATE INDEX {} ON {} ({})").format( + sql.Identifier(name), sql.Identifier(table), sql.Identifier(column) + ) + ) + return _index_oids(database_url)[name] + + +def _connecting_with_another_replica_acting_first( + statement: str, other_replica: Callable[[QueryNoTemplate], None] +) -> Callable[[str], "psycopg.Connection[tuple[object, ...]]"]: + """A connect function whose cursors let `other_replica` act, once, right before the + first statement containing `statement` runs: the interleaving two replicas booting + together can produce, made deterministic.""" + raced: Final = threading.Event() + + class _RacedCursor(psycopg.Cursor[tuple[object, ...]]): + def execute( # pyright: ignore[reportIncompatibleMethodOverride] # the builder never runs a Template query + self, + query: QueryNoTemplate, + params: "Params | None" = None, + *, + prepare: "bool | None" = None, + binary: "bool | None" = None, + ) -> "_RacedCursor": + text: Final = query.as_string(self.connection) if isinstance(query, sql.Composable) else query + if isinstance(text, str) and statement in text and not raced.is_set(): + raced.set() + other_replica(query) + return super().execute(query, params, prepare=prepare, binary=binary) + + def connect(database_url: str) -> "psycopg.Connection[tuple[object, ...]]": + return psycopg.connect(database_url, autocommit=True, cursor_factory=_RacedCursor) + + return connect + + +@requires_db +def test_an_index_another_replica_made_valid_before_the_lock_was_taken_is_kept(scratch_database: str) -> None: + """Two replicas boot against the same invalid index. The one that takes the lock + second must read the catalog again under it, or it drops the valid index the first + one just finished and starts the whole build over.""" + _leave_an_invalid_index(scratch_database, CALL_ID_INDEX, "LiteLLM_SpendLogs", "litellm_call_id") + theirs: Final[queue.SimpleQueue[int]] = queue.SimpleQueue() + connect: Final = _connecting_with_another_replica_acting_first( + "pg_try_advisory_lock", + lambda _: theirs.put( + _rebuild_as_another_replica(scratch_database, CALL_ID_INDEX, "LiteLLM_SpendLogs", "litellm_call_id") + ), + ) + + assert ensure_request_log_indexes(scratch_database, "public", (CALL_ID_INDEX_DEFINITION,), connect) is True + + assert _index_oids(scratch_database)[CALL_ID_INDEX] == theirs.get_nowait() + assert _index_validity(scratch_database, "litellm_call_id_idx") == {CALL_ID_INDEX: True} + + +@requires_db +def test_a_child_index_another_replica_attached_first_is_not_attached_twice(partitioned_database: str) -> None: + """A replica that reaches the attach step after another one attached the same child + relies on ATTACH PARTITION being a no-op for an index already under that parent + (PostgreSQL 14 ALTER INDEX, ATExecAttachPartitionIdx, checked 2026-10-01); this test + is where that would surface if a future version or a code change made it an error.""" + + def attach_as_another_replica(statement: QueryNoTemplate) -> None: + with psycopg.connect(partitioned_database, autocommit=True) as other_replica: + other_replica.execute(statement) + + connect: Final = _connecting_with_another_replica_acting_first("ATTACH PARTITION", attach_as_another_replica) + + assert ensure_request_log_indexes(partitioned_database, "public", (CALL_ID_INDEX_DEFINITION,), connect) is True + + _assert_index_covers_every_partition(partitioned_database, CALL_ID_INDEX, "litellm_call_id_idx") + + +@requires_db +def test_an_invalid_child_index_left_by_an_interrupted_build_is_rebuilt_and_attached( + partitioned_database: str, +) -> None: + partition: Final = "LiteLLM_SpendLogs_p2026_08" + child: Final = f"{partition}_litellm_call_id_idx" + _leave_an_invalid_index(partitioned_database, child, partition, "litellm_call_id") + assert _index_validity(partitioned_database, "litellm_call_id_idx") == {child: False} + + assert ensure_request_log_indexes(partitioned_database, "public") is True + + _assert_index_covers_every_partition(partitioned_database, CALL_ID_INDEX, "litellm_call_id_idx") + + +@requires_db +def test_an_index_of_that_name_on_another_table_is_left_alone_and_reported(partitioned_database: str) -> None: + with psycopg.connect(partitioned_database, autocommit=True) as conn: + conn.execute(f'CREATE INDEX "{CALL_ID_INDEX}" ON "LiteLLM_ErrorLogs" ("request_id")') + + assert ensure_request_log_indexes(partitioned_database, "public") is False + + assert _index_validity(partitioned_database, "litellm_call_id_idx") == {} + assert _indexed_table(partitioned_database, CALL_ID_INDEX) == "LiteLLM_ErrorLogs" + _assert_index_covers_every_partition(partitioned_database, API_KEY_INDEX, "api_key_startTime_idx") + + +@requires_db +def test_an_invalid_index_of_a_child_name_on_another_table_is_not_dropped(partitioned_database: str) -> None: + child: Final = "LiteLLM_SpendLogs_p2026_08_litellm_call_id_idx" + _leave_an_invalid_index(partitioned_database, child, "LiteLLM_ErrorLogs", "request_id") + + assert ensure_request_log_indexes(partitioned_database, "public") is False + + assert _indexed_table(partitioned_database, child) == "LiteLLM_ErrorLogs" + assert _attached_children(partitioned_database, CALL_ID_INDEX) == frozenset() + + +@requires_db +def test_a_process_holding_the_migration_lock_makes_the_build_wait_for_the_next_job_run(scratch_database: str) -> None: + with psycopg.connect(scratch_database, autocommit=True) as other_replica: + other_replica.execute("SELECT pg_advisory_lock(%s)", (MIGRATION_LOCK_KEY,)) + assert ensure_request_log_indexes(scratch_database, "public") is False + assert _index_validity(scratch_database, "litellm_call_id_idx") == {} + + assert ensure_request_log_indexes(scratch_database, "public") is True + assert _index_validity(scratch_database, "litellm_call_id_idx") == {CALL_ID_INDEX: True} + + +@requires_db +@RESOLVERS +def test_a_migration_job_that_could_not_build_the_indexes_reports_failure_and_succeeds_when_rerun( + scratch_database: str, use_v2_resolver: bool +) -> None: + """The migration job waits for the build and exits by run_migration_job's result; a job + that exits 0 with the indexes missing would leave the table unindexed until the next + deploy or until a serving proxy's background build gets to them.""" + with psycopg.connect(scratch_database, autocommit=True) as other_replica: + other_replica.execute("SELECT pg_advisory_lock(%s)", (MIGRATION_LOCK_KEY,)) + assert _migration_job(use_v2_resolver) is False + _assert_every_ledger_row_is_finished(scratch_database) + assert _index_validity(scratch_database, "litellm_call_id_idx") == {} + + assert _migration_job(use_v2_resolver) is True + assert _index_validity(scratch_database, "litellm_call_id_idx") == {CALL_ID_INDEX: True} + assert _index_validity(scratch_database, "api_key_startTime_idx") == {API_KEY_INDEX: True} + + +@requires_db +@RESOLVERS +def test_the_serving_proxy_setup_applies_the_inert_migrations_and_builds_no_index( + partitioned_database: str, use_v2_resolver: bool +) -> None: + """setup_database alone applies the inert files and builds nothing, so a serving proxy's + readiness is never held up by an index build; the build it starts afterwards, or the + migration job, is what puts the indexes in place.""" + api_key_index_before: Final = _index_validity(partitioned_database, "api_key_startTime_idx") + assert ProxyExtrasDBManager.setup_database(use_migrate=True, use_v2_resolver=use_v2_resolver) is True + + _assert_every_ledger_row_is_finished(partitioned_database) + _assert_no_pending_migrations(partitioned_database) + assert _index_validity(partitioned_database, "litellm_call_id_idx") == {} + assert _index_validity(partitioned_database, "api_key_startTime_idx") == api_key_index_before + + assert _migration_job(use_v2_resolver) is True + _assert_index_covers_every_partition(partitioned_database, API_KEY_INDEX, "api_key_startTime_idx") + _assert_index_covers_every_partition(partitioned_database, CALL_ID_INDEX, "litellm_call_id_idx") + + +@requires_db +def test_a_role_that_may_not_create_indexes_is_logged_and_left_for_the_next_job_run( + scratch_database: str, caplog: pytest.LogCaptureFixture +) -> None: + with psycopg.connect(scratch_database, autocommit=True) as conn: + conn.execute("REVOKE CREATE ON SCHEMA public FROM PUBLIC") + conn.execute("CREATE ROLE spend_logs_reader LOGIN PASSWORD 'reader'") + conn.execute("GRANT USAGE ON SCHEMA public TO spend_logs_reader") + conn.execute('GRANT SELECT ON "LiteLLM_SpendLogs" TO spend_logs_reader') + reader_url: Final = scratch_database.replace("postgres:postgres@", "spend_logs_reader:reader@", 1) + try: + with caplog.at_level("WARNING", logger="litellm_proxy_extras"): + assert ensure_request_log_indexes(reader_url, "public") is False + finally: + with psycopg.connect(scratch_database, autocommit=True) as conn: + conn.execute("DROP OWNED BY spend_logs_reader") + conn.execute("DROP ROLE spend_logs_reader") + assert "leaving them for the next index build" in caplog.text + assert _index_validity(scratch_database, "litellm_call_id_idx") == {} + + +@requires_db +def test_inserts_keep_flowing_while_the_partition_indexes_build(partitioned_database: str) -> None: + """With a write open on one partition, the parent index goes on ONLY the parent and + the CONCURRENTLY child build waits for that write without blocking new INSERTs. A + plain CREATE INDEX on the parent would wait for the same write while holding SHARE + on the parent, queueing every new INSERT behind it.""" + outcome: Final[list[bool]] = [] # mutable-ok: the builder thread hands its result back through it + with psycopg.connect(partitioned_database) as writer: + _insert_spend_log(writer, "LiteLLM_SpendLogs_p2026_08-open", "2026-08-15", table="LiteLLM_SpendLogs_p2026_08") + builder_thread: Final = threading.Thread( + target=lambda: outcome.append(_build_in_its_own_session(partitioned_database, CALL_ID_INDEX_DEFINITION)) + ) + builder_thread.start() + try: + _wait_until_the_build_is_waiting(partitioned_database) + with psycopg.connect(partitioned_database, autocommit=True) as late_writer: + late_writer.execute("SET lock_timeout = '1s'") + _insert_spend_log(late_writer, "LiteLLM_SpendLogs_p2026_08-late", "2026-08-16") + finally: + writer.commit() + builder_thread.join() + assert outcome == [True] + _assert_index_covers_every_partition(partitioned_database, CALL_ID_INDEX, "litellm_call_id_idx") + + +def _insert_for(database_url: str, seconds: float) -> None: + with psycopg.connect(database_url, autocommit=True) as conn: + conn.execute("SET lock_timeout = '1s'") + deadline: Final = time.monotonic() + seconds + while time.monotonic() < deadline: + _insert_spend_log(conn, f"lock-test-{uuid.uuid4().hex}", "2026-08-16") + time.sleep(0.05) + + +def _wait_for_blocked_ddl(database_url: str, query_pattern: str) -> bool: + with psycopg.connect(database_url, autocommit=True) as conn: + deadline: Final = time.monotonic() + 10 + while time.monotonic() < deadline: + if conn.execute( + "SELECT 1 FROM pg_stat_activity WHERE wait_event_type = 'Lock' AND query ILIKE %s", + (query_pattern,), + ).fetchone(): + return True + time.sleep(0.01) + return False + + +@requires_db +def test_inserts_are_never_held_back_while_the_parent_index_waits_for_an_open_write( + partitioned_database: str, +) -> None: + outcome: Final[list[bool]] = [] # mutable-ok: the builder thread hands its result back through it + with psycopg.connect(partitioned_database) as writer: + _insert_spend_log(writer, "parent-index-lock-owner", "2026-08-15") + builder_thread: Final = threading.Thread( + target=lambda: outcome.append(_build_in_its_own_session(partitioned_database, CALL_ID_INDEX_DEFINITION)) + ) + builder_thread.start() + try: + assert _wait_for_blocked_ddl(partitioned_database, "%CREATE INDEX%ON ONLY%") + _insert_for(partitioned_database, 3) + finally: + try: + writer.commit() + finally: + builder_thread.join() + assert outcome == [True] + _assert_index_covers_every_partition(partitioned_database, CALL_ID_INDEX, "litellm_call_id_idx") + + +@requires_db +def test_inserts_are_never_held_back_while_attach_partition_waits_for_a_reader_of_the_child_index( + partitioned_database: str, +) -> None: + partition: Final = "LiteLLM_SpendLogs_p2026_08" + child_index: Final = CALL_ID_INDEX_DEFINITION.partition_index_name(partition) + assert child_index == "LiteLLM_SpendLogs_p2026_08_litellm_call_id_idx" + with psycopg.connect(partitioned_database, autocommit=True) as conn: + conn.execute( + 'CREATE INDEX "LiteLLM_SpendLogs_litellm_call_id_idx" ON ONLY "LiteLLM_SpendLogs" ("litellm_call_id")' + ) + conn.execute( + sql.SQL('CREATE INDEX {} ON {} ("litellm_call_id")').format( + sql.Identifier(CALL_ID_INDEX_DEFINITION.partition_index_name(partition)), sql.Identifier(partition) + ) + ) + + outcome: Final[list[bool]] = [] # mutable-ok: the builder thread hands its result back through it + with psycopg.connect(partitioned_database) as reader: + reader.execute("SET enable_seqscan = off") + reader.execute( + sql.SQL('SELECT count(*) FROM {} WHERE "litellm_call_id" IS NULL').format(sql.Identifier(partition)) + ).fetchone() + reader_pid: Final = reader.execute("SELECT pg_backend_pid()").fetchone()[0] + with psycopg.connect(partitioned_database, autocommit=True) as inspector: + child_lock: Final = inspector.execute( + "SELECT 1 FROM pg_locks WHERE pid = %s AND relation = to_regclass(%s) " + "AND mode = 'AccessShareLock' AND granted", + (reader_pid, f'"{child_index}"'), + ).fetchone() + assert child_lock is not None + builder_thread: Final = threading.Thread( + target=lambda: outcome.append(_build_in_its_own_session(partitioned_database, CALL_ID_INDEX_DEFINITION)) + ) + builder_thread.start() + try: + assert _wait_for_blocked_ddl(partitioned_database, "%ATTACH PARTITION%") + _insert_for(partitioned_database, 3) + finally: + try: + reader.commit() + finally: + builder_thread.join() + assert outcome == [True] + _assert_index_covers_every_partition(partitioned_database, CALL_ID_INDEX, "litellm_call_id_idx") + + +@requires_db +def test_a_parent_index_that_never_gets_its_lock_is_left_for_the_next_index_build( + partitioned_database: str, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture +) -> None: + monkeypatch.setattr(request_log_indexes, "_DDL_LOCK_ATTEMPTS", 2) + with psycopg.connect(partitioned_database) as writer: + _insert_spend_log(writer, "parent-index-lock-owner", "2026-08-15") + with caplog.at_level("WARNING", logger="litellm_proxy_extras"): + assert _build_in_its_own_session(partitioned_database, CALL_ID_INDEX_DEFINITION) is False + assert "leaving it for the next index build" in caplog.text + assert _indexed_table(partitioned_database, CALL_ID_INDEX) is None + writer.commit() + assert _build_in_its_own_session(partitioned_database, CALL_ID_INDEX_DEFINITION) is True + _assert_index_covers_every_partition(partitioned_database, CALL_ID_INDEX, "litellm_call_id_idx") + + +@requires_db +def test_an_attach_that_never_gets_its_lock_is_left_for_the_next_index_build( + partitioned_database: str, monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture +) -> None: + partition: Final = "LiteLLM_SpendLogs_p2026_08" + child_index: Final = CALL_ID_INDEX_DEFINITION.partition_index_name(partition) + assert child_index == "LiteLLM_SpendLogs_p2026_08_litellm_call_id_idx" + with psycopg.connect(partitioned_database, autocommit=True) as conn: + conn.execute( + 'CREATE INDEX "LiteLLM_SpendLogs_litellm_call_id_idx" ON ONLY "LiteLLM_SpendLogs" ("litellm_call_id")' + ) + conn.execute( + sql.SQL('CREATE INDEX {} ON {} ("litellm_call_id")').format( + sql.Identifier(child_index), sql.Identifier(partition) + ) + ) + + monkeypatch.setattr(request_log_indexes, "_DDL_LOCK_ATTEMPTS", 2) + with psycopg.connect(partitioned_database) as reader: + reader.execute("SET enable_seqscan = off") + reader.execute( + sql.SQL('SELECT count(*) FROM {} WHERE "litellm_call_id" IS NULL').format(sql.Identifier(partition)) + ).fetchone() + reader_pid: Final = reader.execute("SELECT pg_backend_pid()").fetchone()[0] + with psycopg.connect(partitioned_database, autocommit=True) as inspector: + child_lock: Final = inspector.execute( + "SELECT 1 FROM pg_locks WHERE pid = %s AND relation = to_regclass(%s) " + "AND mode = 'AccessShareLock' AND granted", + (reader_pid, f'"{child_index}"'), + ).fetchone() + assert child_lock is not None + with caplog.at_level("WARNING", logger="litellm_proxy_extras"): + assert _build_in_its_own_session(partitioned_database, CALL_ID_INDEX_DEFINITION) is False + assert "Could not get the lock for attaching" in caplog.text + reader.commit() + + assert _build_in_its_own_session(partitioned_database, CALL_ID_INDEX_DEFINITION) is True + _assert_index_covers_every_partition(partitioned_database, CALL_ID_INDEX, "litellm_call_id_idx") + + +def _build_in_its_own_session(database_url: str, index: RequestLogIndex) -> bool: + with psycopg.connect(database_url, autocommit=True) as builder: + return build_index_on_partitioned_table(builder, "public", index) + + +def _wait_until_the_build_is_waiting(database_url: str) -> None: + deadline: Final = time.monotonic() + 30 + with psycopg.connect(database_url, autocommit=True) as conn: + while time.monotonic() < deadline: + waiting = conn.execute( + "SELECT 1 FROM pg_stat_activity WHERE query LIKE 'CREATE INDEX%' AND wait_event_type IS NOT NULL" + ).fetchone() + if waiting is not None: + return + time.sleep(0.05) + pytest.fail("the partition index build never started waiting on the open write") + + +def _create_index(database_url: str, name: str, table: str, columns: str) -> int: + """Create a plain index by hand, the way an operator's workaround would, and return its oid.""" + with psycopg.connect(database_url, autocommit=True) as conn: + conn.execute( + sql.SQL("CREATE INDEX {} ON {} {}").format(sql.Identifier(name), sql.Identifier(table), sql.SQL(columns)) + ) + return _index_oids(database_url)[name] + + +@requires_db +def test_a_valid_index_of_the_same_definition_under_another_name_is_renamed_instead_of_rebuilt( + scratch_database: str, +) -> None: + hand_built: Final = _create_index(scratch_database, "call_id_by_hand", "LiteLLM_SpendLogs", '("litellm_call_id")') + + assert ensure_request_log_indexes(scratch_database, "public") is True + + oids: Final = _index_oids(scratch_database) + assert "call_id_by_hand" not in oids and oids[CALL_ID_INDEX] == hand_built + assert _index_validity(scratch_database, "litellm_call_id_idx") == {CALL_ID_INDEX: True} + + +@requires_db +def test_a_hand_built_child_index_under_another_name_is_renamed_and_attached(partitioned_database: str) -> None: + partition: Final = "LiteLLM_SpendLogs_p2026_08" + hand_built: Final = _create_index( + partitioned_database, "p2026_08_call_id_by_hand", partition, '("litellm_call_id")' + ) + + assert ensure_request_log_indexes(partitioned_database, "public") is True + + _assert_index_covers_every_partition(partitioned_database, CALL_ID_INDEX, "litellm_call_id_idx") + oids: Final = _index_oids(partitioned_database) + assert "p2026_08_call_id_by_hand" not in oids and oids[f"{partition}_litellm_call_id_idx"] == hand_built + + +@requires_db +def test_an_index_with_another_definition_is_not_taken_for_the_managed_one(scratch_database: str) -> None: + with psycopg.connect(scratch_database, autocommit=True) as conn: + conn.execute(sql.SQL("DROP INDEX {}").format(sql.Identifier(API_KEY_INDEX))) + others: Final = { + "time_then_key": _create_index( + scratch_database, "time_then_key", "LiteLLM_SpendLogs", '("startTime", "api_key")' + ), + "call_id_desc": _create_index( + scratch_database, "call_id_desc", "LiteLLM_SpendLogs", '("litellm_call_id" DESC)' + ), + "call_id_then_key": _create_index( + scratch_database, "call_id_then_key", "LiteLLM_SpendLogs", '("litellm_call_id", "api_key")' + ), + "call_id_pattern": _create_index( + scratch_database, "call_id_pattern", "LiteLLM_SpendLogs", '("litellm_call_id" text_pattern_ops)' + ), + } + + assert ensure_request_log_indexes(scratch_database, "public") is True + + oids: Final = _index_oids(scratch_database) + assert {name: oids[name] for name in others} == others + assert _index_validity(scratch_database, "litellm_call_id_idx") == {CALL_ID_INDEX: True} + assert _index_validity(scratch_database, "api_key_startTime_idx") == {API_KEY_INDEX: True} + + +@requires_db +def test_a_valid_partitioned_parent_index_under_another_name_is_renamed_with_its_children_kept( + partitioned_database: str, caplog: pytest.LogCaptureFixture +) -> None: + hand_built: Final = _create_index( + partitioned_database, "call_id_parent_by_hand", "LiteLLM_SpendLogs", '("litellm_call_id")' + ) + children_before: Final = _attached_children(partitioned_database, "call_id_parent_by_hand") + + with caplog.at_level("INFO", logger="litellm_proxy_extras"): + assert ensure_request_log_indexes(partitioned_database, "public") is True + + assert "Building index" not in caplog.text + oids: Final = _index_oids(partitioned_database) + assert "call_id_parent_by_hand" not in oids and oids[CALL_ID_INDEX] == hand_built + assert _attached_children(partitioned_database, CALL_ID_INDEX) == children_before + assert _index_validity(partitioned_database, "litellm_call_id_idx")[CALL_ID_INDEX] is True + + +@requires_db +def test_a_second_copy_of_a_managed_index_is_reported_with_its_drop_statement_and_left_in_place( + scratch_database: str, caplog: pytest.LogCaptureFixture +) -> None: + assert ensure_request_log_indexes(scratch_database, "public") is True + copy: Final = _create_index(scratch_database, "call_id_copy", "LiteLLM_SpendLogs", '("litellm_call_id")') + + with caplog.at_level("WARNING", logger="litellm_proxy_extras"): + assert ensure_request_log_indexes(scratch_database, "public") is True + + assert 'remove it with: DROP INDEX CONCURRENTLY "public"."call_id_copy"' in caplog.text + assert _index_oids(scratch_database)["call_id_copy"] == copy diff --git a/tests/test_litellm/tracing/normalizers/test_registry.py b/tests/test_litellm/tracing/normalizers/test_registry.py deleted file mode 100644 index 4c2fd051d6c..00000000000 --- a/tests/test_litellm/tracing/normalizers/test_registry.py +++ /dev/null @@ -1,68 +0,0 @@ -from collections.abc import Mapping -from dataclasses import dataclass -from types import MappingProxyType -from typing import Final - -from litellm.tracing.normalizers import ( - NORMALIZERS, - GenAISemconvNormalizer, - LangSmithNormalizer, - OpenInferenceNormalizer, - select_normalizer, -) -from litellm.tracing.types import SpanRow - -_NO_ATTRIBUTES: Final[Mapping[str, str]] = MappingProxyType({}) - - -def test_langsmith_scope_selects_langsmith_without_any_attributes(): - assert isinstance(select_normalizer("langsmith", _NO_ATTRIBUTES), LangSmithNormalizer) - - -def test_langsmith_kind_attribute_selects_langsmith_under_any_scope(): - assert isinstance(select_normalizer("other", MappingProxyType({"langsmith.span.kind": "llm"})), LangSmithNormalizer) - - -def test_langsmith_wins_over_openinference_when_both_markers_present(): - attributes: Final = MappingProxyType({"langsmith.span.kind": "llm", "openinference.span.kind": "LLM"}) - assert isinstance(select_normalizer("other", attributes), LangSmithNormalizer) - - -def test_openinference_kind_attribute_selects_openinference(): - assert isinstance( - select_normalizer("other", MappingProxyType({"openinference.span.kind": "LLM"})), OpenInferenceNormalizer - ) - - -def test_unmarked_span_falls_back_to_genai(): - assert isinstance( - select_normalizer("other", MappingProxyType({"gen_ai.operation.name": "chat"})), GenAISemconvNormalizer - ) - - -def test_empty_registry_falls_back_to_genai(): - assert isinstance(select_normalizer("langsmith", _NO_ATTRIBUTES, registry=()), GenAISemconvNormalizer) - - -def test_registry_names_are_unique(): - names: Final = tuple(n.name for n in NORMALIZERS) - assert len(names) == len(frozenset(names)) - - -@dataclass(frozen=True, slots=True) -class _CustomNormalizer: - name: str = "custom" - - def matches(self, scope_name: str, attributes: Mapping[str, str]) -> bool: - return scope_name == "custom-sdk" - - def normalize(self, row: SpanRow, attributes: Mapping[str, str]) -> None: - return None - - -def test_normalizer_inserted_ahead_in_custom_registry_wins_only_where_it_matches(): - registry: Final = (_CustomNormalizer(), *NORMALIZERS) - assert isinstance( - select_normalizer("custom-sdk", MappingProxyType({"langsmith.span.kind": "llm"}), registry), _CustomNormalizer - ) - assert isinstance(select_normalizer("langsmith", _NO_ATTRIBUTES, registry), LangSmithNormalizer) diff --git a/tests/test_litellm_rust/test_traces.py b/tests/test_litellm_rust/test_traces.py index e6492c9bca6..5a09578661e 100644 --- a/tests/test_litellm_rust/test_traces.py +++ b/tests/test_litellm_rust/test_traces.py @@ -9,7 +9,7 @@ from urllib.parse import parse_qs, urlsplit import pytest from litellm.rust_bridge._native import NativeTraceStorage, trace_decode_otlp -from litellm.rust_bridge.traces import ClickHouseStorage +from litellm.rust_bridge.traces import ClickHouseStorage, NormalizedSpan, normalized_field_definitions from litellm.tracing import Tenant, TraceReceiver, TracingPayloadTooLargeError from litellm.tracing.decode import decode_otlp from litellm.tracing.store import TraceStore @@ -157,6 +157,15 @@ def test_decode_and_tenant_stamping_share_resources_without_crossing_groups() -> assert rows[0]["ResourceAttributes"] == {"shared": "x" * 128, "litellm.team_id": "spoofed"} +def test_normalized_field_contract_matches_decoded_rust_span() -> None: + body: Final = _resource_export(8, 1) + spans: Final = trace_decode_otlp(body, "application/json") + fields: Final = normalized_field_definitions() + assert len(spans) == 1 + assert {field.name for field in fields} == set(spans[0]["normalized"]) == set(NormalizedSpan.model_fields) + assert len({field.clickhouse_column for field in fields}) == len(fields) + + @pytest.mark.asyncio async def test_resource_fanout_reaches_insert_with_identical_values(recording_server: RecordingServer) -> None: body: Final = _resource_export(16 * 1024, 1024) @@ -198,3 +207,124 @@ async def test_insert_validates_values_without_pydantic_copy(recording_server: R assert stored["Timestamp"] == "1970-01-01T00:00:00.000000001Z" assert stored["ResourceAttributes"] == attributes assert stored["SpanAttributes"] == attributes + + +@pytest.mark.parametrize("role", ["proxy_admin", "proxy_admin_viewer", "internal_user"]) +def test_trace_sql_endpoint_executes_for_admin_and_preserves_clickhouse_envelope( + recording_server: RecordingServer, role: str +) -> None: + from fastapi import FastAPI + from fastapi.testclient import TestClient + + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + from litellm.proxy.tracing_endpoints import provide_receiver, provide_trace_query_secret, router + + envelope: Final = {"meta": [{"name": "answer", "type": "UInt8"}], "data": [{"answer": 42}], "rows": 1} + recording_server.expected_requests = 12 + for _ in range(11): + recording_server.enqueue(ResponseSpec(body="")) + recording_server.enqueue(ResponseSpec(body=envelope)) + storage: Final = ClickHouseStorage("trace_test", recording_server.base_url, recording_server.base_url) + app: Final = FastAPI() + app.include_router(router) + app.dependency_overrides[provide_trace_query_secret] = lambda: "test-master-secret" + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=role, token="test") + app.dependency_overrides[provide_receiver] = lambda: TraceReceiver(TraceStore(storage)) + with TestClient(app) as client: + result: Final = client.post("/v1/traces/query", json={"sql": "SELECT 42 AS answer"}) + assert result.status_code == 200, result.text + assert result.json() == envelope + assert recording_server.requests[-1].raw_body == b"SELECT 42 AS answer" + assert client.post("/v1/traces/query", json={"sql": " "}).status_code == 400 + assert client.post("/v1/traces/query", json={}).status_code == 422 + + +def test_trace_help_endpoint_runs_native_schema_and_metadata_discovery(recording_server: RecordingServer) -> None: + from fastapi import FastAPI + from fastapi.testclient import TestClient + + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + from litellm.proxy.tracing_endpoints import provide_receiver, provide_trace_query_secret, router + + recording_server.expected_requests = 17 + for _ in range(11): + recording_server.enqueue(ResponseSpec(body="")) + for response in ( + {"data": [{"name": "Model", "type": "String"}]}, + {"data": []}, + {"data": []}, + {"data": [{"metadata": '{"custom": {"label": "hello"}}'}]}, + {"data": [{"key": "custom.span"}]}, + {"data": [{"key": "custom.resource"}]}, + ): + recording_server.enqueue(ResponseSpec(body=response)) + storage: Final = ClickHouseStorage("trace_test", recording_server.base_url, recording_server.base_url) + app: Final = FastAPI() + app.include_router(router) + app.dependency_overrides[provide_trace_query_secret] = lambda: "test-master-secret" + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role="proxy_admin", token="test") + app.dependency_overrides[provide_receiver] = lambda: TraceReceiver(TraceStore(storage)) + with TestClient(app) as client: + result: Final = client.get("/v1/traces/query/help") + assert result.status_code == 200, result.text + body: Final = result.json() + assert body["guide"].startswith("Trace SQL query guide") + assert "JSONExtractRaw(metadata, 'custom', 'label')" in body["guide"] + assert body["tables"][0]["columns"] == [{"name": "Model", "type": "String"}] + assert body["metadata"]["fields"][1] == { + "path": ["custom", "label"], + "types": ["string"], + "expression": "JSONExtractRaw(metadata, 'custom', 'label')", + } + assert body["attributes"][0]["fields"][0]["expression"] == "SpanAttributes['custom.span']" + assert body["attributes"][1]["fields"][0]["expression"] == "ResourceAttributes['custom.resource']" + + +@pytest.mark.parametrize("clickhouse_status, expected_status", [(400, 400), (404, 400), (500, 503), (503, 503)]) +def test_trace_sql_endpoint_distinguishes_query_errors_from_reader_failures( + recording_server: RecordingServer, clickhouse_status: int, expected_status: int +) -> None: + from fastapi import FastAPI + from fastapi.testclient import TestClient + + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + from litellm.proxy.tracing_endpoints import provide_receiver, provide_trace_query_secret, router + + recording_server.expected_requests = 13 + for _ in range(11): + recording_server.enqueue(ResponseSpec(body="")) + recording_server.enqueue(ResponseSpec(status=clickhouse_status, body=b"ClickHouse rejected the query")) + envelope: Final = {"meta": [{"name": "answer", "type": "UInt8"}], "data": [{"answer": 42}], "rows": 1} + recording_server.enqueue(ResponseSpec(body=envelope)) + storage: Final = ClickHouseStorage("trace_test", recording_server.base_url, recording_server.base_url) + app: Final = FastAPI() + app.include_router(router) + app.dependency_overrides[provide_trace_query_secret] = lambda: "test-master-secret" + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role="proxy_admin", token="test") + app.dependency_overrides[provide_receiver] = lambda: TraceReceiver(TraceStore(storage)) + with TestClient(app) as client: + failed: Final = client.post("/v1/traces/query", json={"sql": "SELEC 42"}) + assert failed.status_code == expected_status, failed.text + recovered: Final = client.post("/v1/traces/query", json={"sql": "SELECT 42 AS answer"}) + assert recovered.status_code == 200, recovered.text + assert recovered.json() == envelope + assert recording_server.requests[-2].raw_body == b"SELEC 42" + + +@pytest.mark.asyncio +async def test_trace_receiver_reads_with_only_one_clickhouse_url( + recording_server: RecordingServer, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("CLICKHOUSE_URL", recording_server.base_url) + monkeypatch.setenv("CLICKHOUSE_DATABASE", "trace_test") + monkeypatch.delenv("CLICKHOUSE_READER_URL", raising=False) + recording_server.enqueue(ResponseSpec(body={"data": [{"trace_id": "trace-1"}]})) + receiver: Final = TraceReceiver.from_env() + rows: Final = await receiver.store.storage.query("trace_spans", {"trace_id": "trace-1"}) + assert rows == [{"trace_id": "trace-1"}] + parameters: Final = parse_qs(urlsplit(recording_server.requests[0].path).query) + assert parameters["database"] == ["trace_test"] + assert parameters["readonly"] == ["1"] diff --git a/tests/unit/harness/__init__.py b/tests/unit/harness/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/harness/core_fakes.py b/tests/unit/harness/core_fakes.py new file mode 100644 index 00000000000..c86cf39739b --- /dev/null +++ b/tests/unit/harness/core_fakes.py @@ -0,0 +1,246 @@ +"""Fake handler/config, sandbox and endpoint shared by the core runtime tests.""" + +from __future__ import annotations + +import asyncio +import os +from collections.abc import AsyncIterator, Callable +from dataclasses import dataclass +from typing import Any, ClassVar + +import pytest + +from litellm.harness import runtime +from litellm.harness.context import SessionContext +from litellm.harness.handlers.base import BaseHarnessHandler +from litellm.llms.base_llm.harness.transformation import BaseHarnessConfig +from litellm.harness.options import ClaudeCodeOptions +from litellm.harness.sandbox.base import CompletedRun +from litellm.harness.sandbox.snapshot import snapshot_local +from litellm.harness.types import ( + Approval, + Capabilities, + Event, + Harness, + Text, + ToolCall, + ToolResult, +) + +ALL_MODES = frozenset({"read-only", "ask", "edit", "full"}) +FULL_CAPS = Capabilities( + structured_output=True, + tool_approval=True, + tool_filtering=True, + history=True, + custom_tools=True, + skills=True, + resume=True, + permission_modes=ALL_MODES, +) +NARROW_CAPS = Capabilities( + structured_output=False, + tool_approval=False, + tool_filtering=False, + history=False, + custom_tools=False, + skills=False, + resume=False, + permission_modes=frozenset({"read-only", "full"}), +) + + +class FakeConfig(BaseHarnessConfig): + """Declares the fake harness; per-test subclasses override capabilities.""" + + harness: ClassVar[Harness] = Harness.CLAUDE_CODE + options_type: ClassVar[type] = ClaudeCodeOptions + capabilities: ClassVar[Capabilities] = FULL_CAPS + uses_model_endpoint: ClassVar[bool] = True + + +Script = Callable[["FakeAdapter", SessionContext, str], AsyncIterator[Event]] + + +class FakeSandbox: + """A LocalSandbox-like object over a temp dir; no subprocesses.""" + + def __init__(self, workdir: str) -> None: + self.workdir = workdir + self.closed = False + + def _path(self, path: str) -> str: + return path if os.path.isabs(path) else os.path.join(self.workdir, path) + + async def exec(self, cmd: list[str], *, env: Any = None, cwd: Any = None) -> Any: + raise NotImplementedError + + async def run( + self, cmd: list[str], *, env: Any = None, cwd: Any = None, timeout: Any = None + ) -> CompletedRun: + return CompletedRun(stdout="", stderr="", exit_code=0) + + async def read(self, path: str) -> bytes: + with open(self._path(path), "rb") as fh: + return fh.read() + + async def write(self, path: str, data: bytes) -> None: + full = self._path(path) + os.makedirs(os.path.dirname(full), exist_ok=True) + with open(full, "wb") as fh: + fh.write(data) + + def host_url(self, port: int) -> str: + return f"http://127.0.0.1:{port}" + + async def which(self, binary: str) -> str | None: + return None + + async def snapshot(self) -> dict[str, str]: + return await snapshot_local(self.workdir) + + async def close(self) -> None: + self.closed = True + + +@dataclass +class FakeUsage: + input_tokens: int = 0 + output_tokens: int = 0 + cost: float = 0.0 + calls: int = 0 + + def add(self, input_tokens: int, output_tokens: int, cost: float) -> None: + self.input_tokens += input_tokens + self.output_tokens += output_tokens + self.cost += cost + self.calls += 1 + + +class FakeEndpoint: + """Stands in for ModelEndpoint; records every instance.""" + + instances: ClassVar[list[FakeEndpoint]] = [] + + def __init__(self, harness: Harness, model: Any, gateway: Any, **kwargs: Any): + self.harness = harness + self.model = model + self.gateway = gateway + self.kwargs = kwargs + self.usage = FakeUsage() + self.url = "http://127.0.0.1:1" + self.token = "tok" + self.entered = False + self.exited = False + FakeEndpoint.instances.append(self) + + async def __aenter__(self) -> FakeEndpoint: + self.entered = True + return self + + async def __aexit__(self, *exc_info: object) -> None: + self.exited = True + + +async def script_hello( + adapter: FakeAdapter, ctx: SessionContext, prompt: str +) -> AsyncIterator[Event]: + yield Text("hello ") + yield ToolCall(id="t1", name="bash", native_name="Bash", input={"cmd": "ls"}) + yield ToolResult(id="t1", output="a.txt") + yield Text("world") + if ctx.endpoint is not None: + ctx.endpoint.usage.add(10, 5, 0.25) + else: + ctx.input_tokens += 10 + ctx.output_tokens += 5 + ctx.cost += 0.25 + ctx.calls += 1 + + +class FakeAdapter(BaseHarnessHandler): + """Configurable adapter; subclass per test and set `script` / `caps`.""" + + harness: ClassVar[Harness] = Harness.CLAUDE_CODE + options_type: ClassVar[type] = ClaudeCodeOptions + capabilities: ClassVar[Capabilities] = FULL_CAPS + uses_endpoint: ClassVar[bool] = True + script: ClassVar[Script] = script_hello + instances: ClassVar[list[FakeAdapter]] = [] + + def __init__(self, config: BaseHarnessConfig | None = None) -> None: + self.config = config if config is not None else FakeConfig() + self.calls: list[str] = [] + self.prompts: list[str] = [] + self.resumed_with: str | None = None + self.approvals: list[tuple[bool, str]] = [] + type(self).instances.append(self) + + async def start(self, ctx: SessionContext) -> None: + self.calls.append("start") + + async def turn(self, ctx: SessionContext, prompt: str) -> AsyncIterator[Event]: + self.calls.append("turn") + self.prompts.append(prompt) + async for event in type(self).script(self, ctx, prompt): + yield event + + async def stop(self, ctx: SessionContext) -> None: + self.calls.append("stop") + + def native_session_id(self) -> str | None: + return "native-123" + + async def resume(self, ctx: SessionContext, native_session_id: str) -> None: + self.calls.append("resume") + self.resumed_with = native_session_id + + async def history(self, ctx: SessionContext) -> list[dict[str, Any]]: + return [{"role": "user", "content": p} for p in self.prompts] + + +async def script_approval( + adapter: FakeAdapter, ctx: SessionContext, prompt: str +) -> AsyncIterator[Event]: + approval = Approval(tool="bash", input={"cmd": "rm"}) + yield approval + decision = await approval.wait() + adapter.approvals.append(decision) + yield Text("allowed" if decision[0] else "denied") + + +def install_adapter( + monkeypatch: pytest.MonkeyPatch, + script: Script = script_hello, + caps: Capabilities = FULL_CAPS, + uses_endpoint: bool = True, +) -> type[FakeAdapter]: + """Register a FakeAdapter subclass for every harness and fake the endpoint.""" + adapter_cls = type( + "TestAdapter", + (FakeAdapter,), + { + "script": staticmethod(script), + "capabilities": caps, + "uses_endpoint": uses_endpoint, + "instances": [], + }, + ) + config_cls = type( + "TestConfig", + (FakeConfig,), + {"capabilities": caps, "uses_model_endpoint": uses_endpoint}, + ) + monkeypatch.setattr(runtime, "get_harness_config", lambda harness: config_cls()) + monkeypatch.setattr( + runtime, "get_harness_handler", lambda config: adapter_cls(config) + ) + monkeypatch.setattr(runtime, "ModelEndpoint", FakeEndpoint) + monkeypatch.delenv("LITELLM_PROXY_API_BASE", raising=False) + monkeypatch.delenv("LITELLM_PROXY_API_KEY", raising=False) + FakeEndpoint.instances = [] + return adapter_cls + + +async def wait_forever() -> None: + await asyncio.Event().wait() diff --git a/tests/unit/harness/handlers/__init__.py b/tests/unit/harness/handlers/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/harness/handlers/test_deepagents_handler.py b/tests/unit/harness/handlers/test_deepagents_handler.py new file mode 100644 index 00000000000..6f0963cac46 --- /dev/null +++ b/tests/unit/harness/handlers/test_deepagents_handler.py @@ -0,0 +1,378 @@ +import asyncio +import builtins +import os +import sys +from pathlib import Path +from typing import Any + +import pytest +from pydantic import BaseModel + +from litellm.harness.context import GatewayTarget, SessionContext +from litellm.harness.errors import HarnessError, HarnessInstallFailed +from litellm.harness.handlers import deepagents_handler as dh +from litellm.harness.sandbox.local import LocalSandbox +from litellm.harness.types import Approval, Harness, Text, ToolCall, ToolResult +from litellm.llms.deepagents.harness.transformation import DeepAgentsHarnessConfig + +pytest.importorskip("deepagents") +pytest.importorskip("langchain_litellm") + +from langchain_core.language_models.fake_chat_models import ( # noqa: E402 + FakeMessagesListChatModel, +) +from langchain_core.messages import AIMessage # noqa: E402 + +from litellm.llms.deepagents.harness.sandbox_backend import ( # noqa: E402 + SandboxBackend, + message_cost, +) + +USAGE = {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15} + + +class FakeToolModel(FakeMessagesListChatModel): + """Canned responses; records the tool names bound on each call.""" + + bound: list = [] + + def bind_tools(self, tools: Any, **kwargs: Any) -> "FakeToolModel": + names = [getattr(t, "name", None) or t.get("name") for t in tools] + self.bound.append(sorted(n for n in names if n)) + return self + + +def tool_call(name: str, args: dict, call_id: str) -> AIMessage: + return AIMessage( + content="", + tool_calls=[{"name": name, "args": args, "id": call_id}], + usage_metadata=USAGE, + ) + + +def final(text: str) -> AIMessage: + return AIMessage(content=text, usage_metadata=USAGE) + + +@pytest.fixture +def fake_model(monkeypatch: pytest.MonkeyPatch): + def install(responses: list) -> FakeToolModel: + model = FakeToolModel(responses=responses, bound=[]) + monkeypatch.setattr(dh, "build_chat_model", lambda ctx, deps: model) + return model + + return install + + +def make_ctx(tmp_path: Path, **kwargs: Any) -> SessionContext: + base: dict[str, Any] = { + "harness": Harness.DEEPAGENTS, + "sandbox": LocalSandbox(tmp_path), + "session_id": f"s-{os.urandom(4).hex()}", + "model": "gpt-4o-mini", + } + return SessionContext(**{**base, **kwargs}) + + +def make_handler() -> dh.DeepAgentsHandler: + return dh.DeepAgentsHandler(DeepAgentsHarnessConfig()) + + +async def started(ctx: SessionContext) -> dh.DeepAgentsHandler: + handler = make_handler() + await handler.start(ctx) + return handler + + +async def run_turn( + handler: dh.DeepAgentsHandler, + ctx: SessionContext, + prompt: str, + approve: bool = True, +) -> list: + events = [] + async for event in handler.turn(ctx, prompt): + events.append(event) + if isinstance(event, Approval): + event.allow() if approve else event.deny("no") + return events + + +async def test_write_then_read_events_and_file(tmp_path: Path, fake_model) -> None: + fake_model( + [ + tool_call("write_file", {"file_path": "/hello.txt", "content": "hi"}, "c1"), + tool_call("read_file", {"file_path": "/hello.txt"}, "c2"), + final("done"), + ] + ) + ctx = make_ctx(tmp_path) + handler = await started(ctx) + events = await run_turn(handler, ctx, "write hello.txt with hi") + + calls = [e for e in events if isinstance(e, ToolCall)] + results = [e for e in events if isinstance(e, ToolResult)] + assert [(c.name, c.native_name, c.builtin) for c in calls] == [ + ("write", "write_file", True), + ("read", "read_file", True), + ] + assert [r.id for r in results] == ["c1", "c2"] + assert "hi" in results[1].output + assert not any(r.is_error for r in results) + assert "done" in "".join(e.delta for e in events if isinstance(e, Text)) + assert (tmp_path / "hello.txt").read_text() == "hi" + assert ctx.final_text == "done" + assert (ctx.input_tokens, ctx.output_tokens, ctx.calls) == (30, 15, 3) + assert ctx.cost > 0 + history = await handler.history(ctx) + assert history[0] == {"role": "user", "content": "write hello.txt with hi"} + assert history[-1]["content"] == "done" + + +async def test_read_only_hides_write_tools(tmp_path: Path, fake_model) -> None: + model = fake_model( + [ + tool_call("write_file", {"file_path": "/x.txt", "content": "no"}, "c1"), + final("ok"), + ] + ) + ctx = make_ctx(tmp_path, permissions="read-only") + handler = await started(ctx) + events = await run_turn(handler, ctx, "try to write") + + first = model.bound[0] + assert "read_file" in first and "ls" in first + assert not {"write_file", "edit_file", "execute", "delete"} & set(first) + result = next(e for e in events if isinstance(e, ToolResult)) + assert result.is_error + assert not (tmp_path / "x.txt").exists() + + +async def test_disable_tools_uses_normalized_names(tmp_path: Path, fake_model) -> None: + model = fake_model([final("ok")]) + ctx = make_ctx(tmp_path, disable_tools=["bash", "grep"]) + handler = await started(ctx) + await run_turn(handler, ctx, "hi") + assert "execute" not in model.bound[0] and "grep" not in model.bound[0] + assert "write_file" in model.bound[0] + + +class Answer(BaseModel): + city: str + + +async def test_structured_output(tmp_path: Path, fake_model) -> None: + fake_model([tool_call("Answer", {"city": "Paris"}, "c1")]) + ctx = make_ctx(tmp_path, output=Answer) + handler = await started(ctx) + events = await run_turn(handler, ctx, "capital of France?") + assert Answer.model_validate_json(ctx.output_json or "") == Answer(city="Paris") + assert not any(isinstance(e, ToolCall) for e in events) + + +async def test_custom_tool(tmp_path: Path, fake_model) -> None: + def add(a: int, b: int) -> int: + """Add two numbers.""" + return a + b + + fake_model([tool_call("add", {"a": 2, "b": 3}, "c1"), final("5")]) + ctx = make_ctx(tmp_path, tools=[add]) + handler = await started(ctx) + events = await run_turn(handler, ctx, "2+3") + call = next(e for e in events if isinstance(e, ToolCall)) + assert (call.name, call.builtin) == ("add", False) + assert next(e for e in events if isinstance(e, ToolResult)).output == "5" + + +@pytest.mark.parametrize("approve", [True, False]) +async def test_ask_permissions_emit_approval( + tmp_path: Path, fake_model, approve: bool +) -> None: + fake_model( + [ + tool_call("write_file", {"file_path": "/a.txt", "content": "x"}, "c1"), + final("end"), + ] + ) + ctx = make_ctx(tmp_path, permissions="ask") + handler = await started(ctx) + events = await run_turn(handler, ctx, "write a", approve=approve) + approval = next(e for e in events if isinstance(e, Approval)) + assert approval.tool == "write" + assert approval.input["file_path"] == "/a.txt" + assert (tmp_path / "a.txt").exists() is approve + assert ctx.final_text == "end" + + +async def test_edit_and_execute_through_sandbox(tmp_path: Path, fake_model) -> None: + (tmp_path / "f.txt").write_text("one two\n") + fake_model( + [ + tool_call( + "edit_file", + {"file_path": "/f.txt", "old_string": "two", "new_string": "three"}, + "c1", + ), + tool_call("execute", {"command": "cat f.txt"}, "c2"), + final("ok"), + ] + ) + ctx = make_ctx(tmp_path) + handler = await started(ctx) + events = await run_turn(handler, ctx, "edit") + calls = [e.name for e in events if isinstance(e, ToolCall)] + assert calls == ["edit", "bash"] + results = [e for e in events if isinstance(e, ToolResult)] + assert "one three" in results[1].output + assert (tmp_path / "f.txt").read_text() == "one three\n" + + +async def test_resume_keeps_thread(tmp_path: Path, fake_model) -> None: + fake_model([final("first"), final("second")]) + ctx = make_ctx(tmp_path) + handler = await started(ctx) + await run_turn(handler, ctx, "one") + native = handler.native_session_id() + assert native == ctx.session_id + + other = make_handler() + ctx2 = make_ctx(tmp_path) + await other.start(ctx2) + await other.resume(ctx2, native or "") + await run_turn(other, ctx2, "two") + history = await other.history(ctx2) + assert [m["content"] for m in history if m["role"] == "user"] == ["one", "two"] + + +async def test_skills_copied_and_loaded(tmp_path: Path, fake_model) -> None: + skill = tmp_path / "src-skills" / "greeter" + skill.mkdir(parents=True) + (skill / "SKILL.md").write_text( + "---\nname: greeter\ndescription: Says hi\n---\nSay hi.\n" + ) + work = tmp_path / "work" + work.mkdir() + fake_model([final("ok")]) + ctx = make_ctx(work, skills=[str(skill)]) + handler = await started(ctx) + await run_turn(handler, ctx, "hi") + assert (work / ".deepagents" / "skills" / "greeter" / "SKILL.md").exists() + + +async def test_turn_and_history_before_start_and_after_stop( + tmp_path: Path, fake_model +) -> None: + fake_model([final("ok")]) + ctx = make_ctx(tmp_path) + handler = make_handler() + with pytest.raises(HarnessError, match="not started"): + await run_turn(handler, ctx, "hi") + await handler.start(ctx) + await handler.stop(ctx) + with pytest.raises(HarnessError, match="not started"): + await handler.history(ctx) + + +async def test_start_validates_model(tmp_path: Path, fake_model) -> None: + fake_model([final("ok")]) + with pytest.raises(ValueError, match="needs model="): + await started(make_ctx(tmp_path, model=None)) + + +def test_build_chat_model_uses_chat_model_kwargs(tmp_path: Path) -> None: + deps = dh.load_deps() + gw = GatewayTarget(api_base="https://gw.example.com", api_key="sk-virtual") + model = dh.build_chat_model(make_ctx(tmp_path, gateway=gw), deps) + assert isinstance(model, deps.chat_litellm) + assert model.model == "litellm_proxy/gpt-4o-mini" + + +def test_shared_checkpointer_is_process_wide() -> None: + deps = dh.load_deps() + assert dh.shared_checkpointer(deps) is dh.shared_checkpointer(deps) + + +async def test_sandbox_backend_fs_ops(tmp_path: Path) -> None: + (tmp_path / "src").mkdir() + (tmp_path / "src" / "a.py").write_text("print('hello')\n") + (tmp_path / "b.txt").write_text("hello world\n") + backend = SandboxBackend(LocalSandbox(tmp_path), loop=asyncio.get_running_loop()) + + ls = await backend.als("/") + assert {e["path"] for e in ls.entries or []} == {"/b.txt", "/src/"} + assert (await backend.als("/missing")).error + globbed = await backend.aglob("*.py") + assert [m["path"] for m in globbed.matches or []] == ["/src/a.py"] + grep = await backend.agrep("hello", glob="*.txt") + assert [(m["path"], m["line"]) for m in grep.matches or []] == [("/b.txt", 1)] + read = await backend.aread("/b.txt") + assert read.file_data and read.file_data["content"] == "hello world\n" + assert (await backend.aread("/nope.txt")).error + assert (await backend.aread("/../etc/passwd")).error + edit = await backend.aedit("/b.txt", "hello", "bye") + assert edit.occurrences == 1 + assert (await backend.aedit("/b.txt", "zzz", "q")).error + assert (await backend.adelete("/src")).path == "/src" + assert not (tmp_path / "src").exists() + assert ( + backend.to_real(str(tmp_path / "b.txt")) + == str(LocalSandbox(tmp_path).workdir) + "/b.txt" + ) + sync_ls = await asyncio.to_thread(backend.ls, "/") + assert [e["path"] for e in sync_ls.entries or []] == ["/b.txt"] + + read_only = SandboxBackend( + LocalSandbox(tmp_path), + loop=asyncio.get_running_loop(), + writable=False, + allow_execute=False, + ) + assert (await read_only.awrite("/c.txt", "x")).error + assert (await read_only.aexecute("ls")).exit_code == 1 + assert not (tmp_path / "c.txt").exists() + + +def test_message_cost_prefers_reported_and_never_raises() -> None: + reported = AIMessage(content="", response_metadata={"response_cost": 0.5}) + assert message_cost(reported, "gpt-4o-mini", 1, 1) == 0.5 + assert message_cost(AIMessage(content=""), "not-a-real-model-xyz", 10, 10) == 0.0 + assert message_cost(AIMessage(content=""), "gpt-4o-mini", 1000, 1000) > 0 + + +def test_missing_deps_raise_install_hint(monkeypatch: pytest.MonkeyPatch) -> None: + real_import = builtins.__import__ + + def fake_import(name: str, *args: Any, **kwargs: Any) -> Any: + if name.startswith("deepagents"): + raise ImportError("No module named 'deepagents'") + return real_import(name, *args, **kwargs) + + for mod in [m for m in sys.modules if m.startswith("deepagents")]: + monkeypatch.delitem(sys.modules, mod) + monkeypatch.setattr(builtins, "__import__", fake_import) + with pytest.raises( + HarnessInstallFailed, match="pip install deepagents langchain-litellm" + ): + dh.load_deps() + + +LIVE_BASE = os.environ.get("LITELLM_PROXY_API_BASE", "") +LIVE_KEY = os.environ.get("LITELLM_PROXY_API_KEY", "") + + +@pytest.mark.skipif( + not (LIVE_BASE and LIVE_KEY), reason="LITELLM_PROXY_API_BASE / KEY not set" +) +async def test_live_gateway_write_file(tmp_path: Path) -> None: + model = os.environ.get("HARNESS_DEEPAGENTS_LIVE_MODEL", "claude-haiku-4-5-20251001") + ctx = make_ctx( + tmp_path, + model=model, + gateway=GatewayTarget(api_base=LIVE_BASE, api_key=LIVE_KEY), + max_turns=6, + ) + handler = await started(ctx) + events = await run_turn(handler, ctx, "write hello.txt with hi") + assert any(isinstance(e, ToolCall) and e.name == "write" for e in events) + assert (tmp_path / "hello.txt").read_text().strip() == "hi" + assert ctx.calls >= 1 and ctx.input_tokens > 0 diff --git a/tests/unit/harness/sandbox/__init__.py b/tests/unit/harness/sandbox/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/harness/sandbox/test_docker.py b/tests/unit/harness/sandbox/test_docker.py new file mode 100644 index 00000000000..10cb8a092f3 --- /dev/null +++ b/tests/unit/harness/sandbox/test_docker.py @@ -0,0 +1,241 @@ +import asyncio +import hashlib +import shutil +import subprocess +from typing import Optional + +import pytest + +from litellm import sandbox +from litellm.harness.errors import SandboxError +from litellm.harness.sandbox import DockerSandbox, Sandbox +from litellm.harness.sandbox.docker import parse_sha256sum + +DOCKER_IMAGE = "alpine:3.20" +CID = "cid123" + + +class FakeStdin: + def __init__(self) -> None: + self.data = b"" + self.closed = False + + def write(self, data: bytes) -> None: + self.data += data + + async def drain(self) -> None: + return None + + def close(self) -> None: + self.closed = True + + +class FakeHandle: + def __init__(self, stdout: bytes = b"", stderr: bytes = b"", code: int = 0): + self.stdin = FakeStdin() + self.stdout = asyncio.StreamReader() + self.stdout.feed_data(stdout) + self.stdout.feed_eof() + self.stderr = asyncio.StreamReader() + self.stderr.feed_data(stderr) + self.stderr.feed_eof() + self.returncode: Optional[int] = code + self._code = code + + async def wait(self) -> int: + return self._code + + async def kill(self) -> None: + return None + + +class Recorder: + """Stands in for DockerSandbox._spawn; scripted responses by docker subcommand.""" + + def __init__(self) -> None: + self.calls: list[list[str]] = [] + self.handles: list[FakeHandle] = [] + self.responses: dict[str, FakeHandle] = {} + + async def __call__(self, args: list[str]) -> FakeHandle: + self.calls.append(args) + handle = self.responses.pop(args[0], None) + if handle is None: + handle = FakeHandle(stdout=f"{CID}\n".encode() if args[0] == "run" else b"") + self.handles.append(handle) + return handle + + +@pytest.fixture +def fake(monkeypatch): + rec = Recorder() + monkeypatch.setattr(DockerSandbox, "_spawn", lambda self, args: rec(args)) + return rec + + +def test_run_args(): + box = sandbox.docker( + "img:1", + mounts={"/host/src": "/workspace"}, + env={"A": "1"}, + name="h1", + ) + assert isinstance(box, Sandbox) + assert box.run_args() == [ + "run", + "-d", + "--rm", + "--add-host=host.docker.internal:host-gateway", + "--name", + "h1", + "-v", + "/host/src:/workspace", + "-e", + "A=1", + "-w", + "/workspace", + "img:1", + "sleep", + "infinity", + ] + assert box.host_url(8080) == "http://host.docker.internal:8080" + + +def test_relative_workdir_rejected(): + with pytest.raises(SandboxError): + sandbox.docker("img", workdir="rel") + + +async def test_lazy_start_and_exec_args(fake): + box = sandbox.docker("img") + assert fake.calls == [] + await box.exec(["echo", "hi"], env={"K": "V"}, cwd="sub") + await box.exec(["true"]) + assert fake.calls[0][0] == "run" + assert [c for c in fake.calls if c[0] == "run"] == [fake.calls[0]] + assert fake.calls[1] == [ + "exec", + "-i", + "-w", + "/workspace/sub", + "-e", + "K=V", + CID, + "echo", + "hi", + ] + assert fake.calls[2] == ["exec", "-i", "-w", "/workspace", CID, "true"] + + +async def test_run_start_failure(fake): + fake.responses["run"] = FakeHandle(stderr=b"no such image", code=125) + box = sandbox.docker("img") + with pytest.raises(SandboxError, match="no such image"): + await box.run(["echo"]) + + +async def test_read_write_which_tempdir(fake): + box = sandbox.docker("img") + await box.start() + + fake.responses["exec"] = FakeHandle(stdout=b"content") + assert await box.read("a.txt") == b"content" + assert fake.calls[-1][-2:] == ["cat", "/workspace/a.txt"] + + await box.write("d/b.txt", b"payload") + assert fake.calls[-1][-5:-1] == ["sh", "-c", fake.calls[-1][-3], "sh"] + assert fake.calls[-1][-1] == "/workspace/d/b.txt" + assert fake.handles[-1].stdin.data == b"payload" + assert fake.handles[-1].stdin.closed + + fake.responses["exec"] = FakeHandle(stdout=b"/usr/bin/codex\n") + assert await box.which("codex") == "/usr/bin/codex" + assert fake.calls[-1][-5:] == ["sh", "-lc", 'command -v "$1"', "sh", "codex"] + + fake.responses["exec"] = FakeHandle(code=1) + assert await box.which("nope") is None + + fake.responses["exec"] = FakeHandle(stdout=b"/tmp/tmp.abc\n") + assert await box.tempdir() == "/tmp/tmp.abc" + + fake.responses["exec"] = FakeHandle(stderr=b"No such file", code=1) + with pytest.raises(SandboxError, match="No such file"): + await box.read("missing") + + +async def test_snapshot_parses_output(fake): + box = sandbox.docker("img") + digest = "a" * 64 + fake.responses["exec"] = FakeHandle( + stdout=f"{digest} ./x.txt\n{digest} ./dir/with space.txt\n".encode() + ) + snap = await box.snapshot() + assert snap == {"x.txt": digest, "dir/with space.txt": digest} + script = fake.calls[-1][-3] + assert "-name '.git'" in script and "-prune" in script + assert fake.calls[-1][-1] == "/workspace" + + +async def test_close_removes_container(fake): + box = sandbox.docker("img") + await box.start() + await box.close() + assert fake.calls[-1] == ["rm", "-f", CID] + with pytest.raises(SandboxError): + await box.start() + + +async def test_close_without_start_is_noop(fake): + await sandbox.docker("img").close() + assert fake.calls == [] + + +async def test_missing_docker_binary(monkeypatch): + monkeypatch.setattr(shutil, "which", lambda name, *a, **k: None) + with pytest.raises(SandboxError, match="docker"): + await sandbox.docker("img").start() + + +def test_parse_sha256sum_ignores_junk(): + assert parse_sha256sum("garbage\n\n") == {} + + +def _docker_usable() -> bool: + if shutil.which("docker") is None: + return False + try: + return ( + subprocess.run( + ["docker", "info"], capture_output=True, timeout=20 + ).returncode + == 0 + ) + except (OSError, subprocess.SubprocessError): + return False + + +@pytest.mark.skipif(not _docker_usable(), reason="docker daemon not available") +async def test_real_docker_roundtrip(): + box = sandbox.docker(DOCKER_IMAGE, workdir="/workspace") + try: + result = await box.run(["echo", "hello"], timeout=120) + assert result.stdout.strip() == "hello" + assert result.exit_code == 0 + + await box.write("seed.txt", b"seed") + await box.write("sub/out.txt", b"from host") + assert await box.read("sub/out.txt") == b"from host" + await box.write("node_modules/skip.js", b"x") + + assert await box.which("sh") is not None + assert await box.which("definitely-not-a-binary-xyz") is None + tmp = await box.tempdir() + assert tmp.startswith("/") + + snap = await box.snapshot() + assert snap == { + "seed.txt": hashlib.sha256(b"seed").hexdigest(), + "sub/out.txt": hashlib.sha256(b"from host").hexdigest(), + } + finally: + await box.close() diff --git a/tests/unit/harness/sandbox/test_local.py b/tests/unit/harness/sandbox/test_local.py new file mode 100644 index 00000000000..5d08dac13a1 --- /dev/null +++ b/tests/unit/harness/sandbox/test_local.py @@ -0,0 +1,180 @@ +import os +import sys + +import pytest + +from litellm import sandbox +from litellm.harness.errors import SandboxError +from litellm.harness.sandbox import LocalSandbox, Process, Sandbox +from litellm.harness.sandbox.local import filtered_environ, is_secret_env_name + +PY = sys.executable + + +@pytest.fixture +async def sbx(tmp_path): + box = sandbox.local(tmp_path) + yield box + await box.close() + + +def test_local_requires_existing_dir(tmp_path): + with pytest.raises(SandboxError): + sandbox.local(tmp_path / "missing") + + +def test_local_resolves_absolute_and_satisfies_protocol(tmp_path, monkeypatch): + monkeypatch.chdir(tmp_path) + (tmp_path / "ws").mkdir() + box = sandbox.local("ws") + assert isinstance(box, LocalSandbox) + assert isinstance(box, Sandbox) + assert box.workdir == os.path.realpath(tmp_path / "ws") + assert box.host_url(4321) == "http://127.0.0.1:4321" + + +async def test_run_collects_output(sbx): + result = await sbx.run( + [ + PY, + "-c", + "import os,sys;print(os.getcwd());print('err',file=sys.stderr);sys.exit(3)", + ] + ) + assert result.stdout.strip() == sbx.workdir + assert result.stderr.strip() == "err" + assert result.exit_code == 3 + + +async def test_run_cwd_inside_workdir(sbx): + os.mkdir(os.path.join(sbx.workdir, "sub")) + result = await sbx.run([PY, "-c", "import os;print(os.getcwd())"], cwd="sub") + assert result.stdout.strip() == os.path.join(sbx.workdir, "sub") + with pytest.raises(SandboxError): + await sbx.run([PY, "-c", "pass"], cwd="/") + + +async def test_exec_streams_stdin(sbx): + proc = await sbx.exec([PY, "-c", "import sys;print(sys.stdin.read().upper())"]) + assert isinstance(proc, Process) + assert proc.stdin is not None + proc.stdin.write(b"hello") + await proc.stdin.drain() + proc.stdin.close() + assert (await proc.stdout.read()).strip() == b"HELLO" + assert await proc.wait() == 0 + + +async def test_run_timeout_kills(sbx): + with pytest.raises(SandboxError, match="timed out"): + await sbx.run([PY, "-c", "import time;time.sleep(30)"], timeout=0.5) + + +async def test_missing_binary_raises(sbx): + with pytest.raises(SandboxError): + await sbx.run(["definitely-not-a-binary-xyz"]) + + +async def test_close_kills_live_processes(tmp_path): + box = sandbox.local(tmp_path) + proc = await box.exec([PY, "-c", "import time;time.sleep(30)"]) + await box.close() + assert proc.returncode is not None + with pytest.raises(SandboxError): + await box.run([PY, "-c", "pass"]) + + +async def test_read_write_roundtrip(sbx): + await sbx.write("a/b/c.txt", b"data") + assert await sbx.read("a/b/c.txt") == b"data" + abs_path = os.path.join(sbx.workdir, "a", "b", "c.txt") + assert await sbx.read(abs_path) == b"data" + + +@pytest.mark.parametrize("bad", ["../escape.txt", "a/../../escape.txt", "/etc/passwd"]) +async def test_path_escape_rejected(sbx, bad): + with pytest.raises(SandboxError, match="escapes"): + await sbx.read(bad) + with pytest.raises(SandboxError, match="escapes"): + await sbx.write(bad, b"x") + + +async def test_symlink_escape_rejected(sbx, tmp_path_factory): + outside = tmp_path_factory.mktemp("outside") + os.symlink(outside, os.path.join(sbx.workdir, "link")) + with pytest.raises(SandboxError, match="escapes"): + await sbx.write("link/x.txt", b"x") + + +async def test_tempdir_is_allowed_and_cleaned(tmp_path): + box = sandbox.local(tmp_path) + tmp = await box.tempdir() + assert os.path.isdir(tmp) + target = os.path.join(tmp, "config.toml") + await box.write(target, b"k = 1") + assert await box.read(target) == b"k = 1" + await box.close() + assert not os.path.exists(tmp) + + +async def test_env_filters_provider_secrets(sbx, monkeypatch): + monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant-fake") + monkeypatch.setenv("OPENAI_BASE_URL", "http://x") + monkeypatch.setenv("GITHUB_TOKEN", "ghp_fake") + monkeypatch.setenv("HARNESS_TEST_PLAIN", "visible") + script = ( + "import os;" + "print(os.environ.get('ANTHROPIC_API_KEY',''));" + "print(os.environ.get('OPENAI_BASE_URL',''));" + "print(os.environ.get('GITHUB_TOKEN',''));" + "print(os.environ.get('HARNESS_TEST_PLAIN',''));" + "print(os.environ.get('ANTHROPIC_BASE_URL',''))" + ) + result = await sbx.run( + [PY, "-c", script], env={"ANTHROPIC_BASE_URL": "http://127.0.0.1:1"} + ) + assert result.stdout.split() == [ + "", + "", + "", + "visible", + "http://127.0.0.1:1", + ] + + +@pytest.mark.parametrize( + "name,secret", + [ + ("ANTHROPIC_API_KEY", True), + ("AWS_REGION", True), + ("VERTEXAI_PROJECT", True), + ("GOOGLE_APPLICATION_CREDENTIALS", True), + ("MY_API_KEY", True), + ("SLACK_BOT_TOKEN", True), + ("CLIENT_SECRET", True), + ("PATH", False), + ("HOME", False), + ], +) +def test_is_secret_env_name(name, secret): + assert is_secret_env_name(name) is secret + + +def test_filtered_environ_overlay_wins(): + env = filtered_environ( + {"PATH": "/bin", "OPENAI_API_KEY": "x"}, {"PATH": "/usr/bin"} + ) + assert env == {"PATH": "/usr/bin"} + + +async def test_which_uses_filtered_path(sbx): + assert await sbx.which("sh") is not None + assert await sbx.which("definitely-not-a-binary-xyz") is None + + +async def test_snapshot_skips_dirs(sbx): + await sbx.write("keep.txt", b"k") + await sbx.write(".git/HEAD", b"ref") + await sbx.write("node_modules/x/index.js", b"x") + snap = await sbx.snapshot() + assert list(snap) == ["keep.txt"] diff --git a/tests/unit/harness/sandbox/test_snapshot.py b/tests/unit/harness/sandbox/test_snapshot.py new file mode 100644 index 00000000000..e3bfc8806db --- /dev/null +++ b/tests/unit/harness/sandbox/test_snapshot.py @@ -0,0 +1,129 @@ +import hashlib +import os + +import pytest + +from litellm import sandbox +from litellm.constants import HARNESS_MAX_DIFF_BYTES +from litellm.harness.sandbox.snapshot import ( + build_file_changes, + capture_text_contents, + diff_snapshots, + snapshot_local, + unified_diff, +) +from litellm.harness.types import FileChange + + +def test_diff_snapshots_kinds(): + before = {"a": "1", "b": "2", "c": "3"} + after = {"a": "1", "b": "9", "d": "4"} + assert diff_snapshots(before, after) == [ + ("b", "modified"), + ("c", "deleted"), + ("d", "created"), + ] + + +async def test_snapshot_local_hashes_and_skips(tmp_path): + (tmp_path / "x.txt").write_bytes(b"hello") + (tmp_path / "sub").mkdir() + (tmp_path / "sub" / "y.txt").write_bytes(b"y") + (tmp_path / "__pycache__").mkdir() + (tmp_path / "__pycache__" / "z.pyc").write_bytes(b"z") + os.symlink(tmp_path / "x.txt", tmp_path / "link.txt") + snap = await snapshot_local(str(tmp_path)) + assert snap == { + "x.txt": hashlib.sha256(b"hello").hexdigest(), + "sub/y.txt": hashlib.sha256(b"y").hexdigest(), + } + + +def test_unified_diff_created(): + diff = unified_diff("f.txt", None, "one\n") + assert diff.startswith("--- /dev/null\n+++ b/f.txt\n") + assert "+one\n" in diff + + +async def test_created_modified_deleted_end_to_end(tmp_path): + box = sandbox.local(tmp_path) + try: + await box.write("mod.txt", b"line1\nline2\n") + await box.write("gone.txt", b"bye\n") + await box.write("bin.dat", b"\x00\x01\x02") + before = await box.snapshot() + contents = await capture_text_contents(box, before) + assert set(contents) == {"mod.txt", "gone.txt"} + + await box.write("mod.txt", b"line1\nchanged\n") + await box.write("new.txt", b"fresh\n") + await box.write("bin.dat", b"\x00\x09") + os.remove(os.path.join(box.workdir, "gone.txt")) + after = await box.snapshot() + + changes = await build_file_changes(box, before, after, contents) + by_path = {c.path: c for c in changes} + assert all(isinstance(c, FileChange) for c in changes) + assert [(c.path, c.kind) for c in changes] == [ + ("bin.dat", "modified"), + ("gone.txt", "deleted"), + ("mod.txt", "modified"), + ("new.txt", "created"), + ] + assert by_path["bin.dat"].diff is None + assert "-line2\n" in by_path["mod.txt"].diff + assert "+changed\n" in by_path["mod.txt"].diff + assert "+fresh\n" in by_path["new.txt"].diff + assert "-bye\n" in by_path["gone.txt"].diff + finally: + await box.close() + + +async def test_modified_without_before_contents_has_no_diff(tmp_path): + box = sandbox.local(tmp_path) + try: + await box.write("f.txt", b"a\n") + before = await box.snapshot() + await box.write("f.txt", b"b\n") + after = await box.snapshot() + changes = await build_file_changes(box, before, after, None) + assert changes == [FileChange(path="f.txt", kind="modified", diff=None)] + finally: + await box.close() + + +async def test_large_text_file_has_no_diff(tmp_path): + box = sandbox.local(tmp_path) + try: + before = await box.snapshot() + await box.write("big.txt", b"a" * (HARNESS_MAX_DIFF_BYTES + 1)) + after = await box.snapshot() + changes = await build_file_changes(box, before, after, {}) + assert changes == [FileChange(path="big.txt", kind="created", diff=None)] + finally: + await box.close() + + +async def test_capture_respects_total_cap(tmp_path, monkeypatch): + monkeypatch.setattr( + "litellm.harness.sandbox.snapshot.HARNESS_SNAPSHOT_MAX_TOTAL_BYTES", 10 + ) + box = sandbox.local(tmp_path) + try: + await box.write("a.txt", b"x" * 8) + await box.write("b.txt", b"x" * 8) + await box.write("c.txt", b"x" * 8) + captured = await capture_text_contents(box, await box.snapshot()) + assert set(captured) == {"a.txt", "b.txt"} + finally: + await box.close() + + +@pytest.mark.parametrize("data", [b"\xff\xfe bad utf8", b"has\x00nul"]) +async def test_capture_skips_binary(tmp_path, data): + box = sandbox.local(tmp_path) + try: + await box.write("f", data) + assert await capture_text_contents(box, {"f": "h"}) == {} + finally: + await box.close() diff --git a/tests/unit/harness/test_endpoint.py b/tests/unit/harness/test_endpoint.py new file mode 100644 index 00000000000..8b56debdc29 --- /dev/null +++ b/tests/unit/harness/test_endpoint.py @@ -0,0 +1,359 @@ +import json +import sys +from collections.abc import AsyncIterator +from typing import Any + +import httpx +import pytest + +import litellm +from litellm.harness import endpoint as endpoint_module +from litellm.harness.endpoint import ( + ModelEndpoint, + SSEUsageParser, + UsageTracker, + compute_cost, + usage_from_body, +) +from litellm.harness.errors import HarnessInstallFailed +from litellm.harness.context import GatewayTarget +from litellm.harness.types import Harness, Usage +from litellm.types.utils import ModelResponse, ModelResponseStream + +GATEWAY = GatewayTarget(api_base="https://gw.example.com", api_key="sk-gateway-secret") + +ANTHROPIC_SSE = ( + b"event: message_start\n" + b'data: {"type":"message_start","message":{"usage":{"input_tokens":11,"output_tokens":1}}}\n\n' + b"event: content_block_delta\n" + b'data: {"type":"content_block_delta","delta":{"type":"text_delta","text":"hi"}}\n\n' + b"event: message_delta\n" + b'data: {"type":"message_delta","usage":{"output_tokens":7}}\n\n' + b"event: message_stop\n" + b'data: {"type":"message_stop"}\n\n' +) +CHAT_SSE = ( + b'data: {"choices":[{"delta":{"content":"hi"}}]}\n\n' + b'data: {"choices":[],"usage":{"prompt_tokens":5,"completion_tokens":3}}\n\n' + b"data: [DONE]\n\n" +) +RESPONSES_SSE = ( + b"event: response.output_text.delta\n" + b'data: {"type":"response.output_text.delta","delta":"hi"}\n\n' + b"event: response.completed\n" + b'data: {"type":"response.completed","response":{"usage":{"input_tokens":20,"output_tokens":4}}}\n\n' +) + + +class Recorder: + def __init__(self, response: httpx.Response) -> None: + self.response = response + self.requests: list[httpx.Request] = [] + + def __call__(self, request: httpx.Request) -> httpx.Response: + self.requests.append(request) + return self.response + + +def sse_response(body: bytes, headers: dict[str, str] | None = None) -> httpx.Response: + return httpx.Response( + 200, + content=body, + headers={"content-type": "text/event-stream", **(headers or {})}, + ) + + +def gateway_endpoint(recorder: Recorder, **kwargs: Any) -> ModelEndpoint: + return ModelEndpoint( + Harness.CLAUDE_CODE, + kwargs.pop("model", "claude-sonnet"), + GATEWAY, + client=httpx.AsyncClient(transport=httpx.MockTransport(recorder)), + **kwargs, + ) + + +def auth(ep: ModelEndpoint) -> dict[str, str]: + return {"authorization": f"Bearer {ep.token}"} + + +@pytest.fixture(autouse=True) +def no_real_cost(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr( + litellm, + "cost_per_token", + lambda model, prompt_tokens, completion_tokens: ( + prompt_tokens * 0.001, + completion_tokens * 0.002, + ), + ) + + +async def test_rejects_bad_token_and_accepts_both_header_styles() -> None: + recorder = Recorder(httpx.Response(200, json={"usage": {}})) + async with gateway_endpoint(recorder) as ep: + assert ep.url == f"http://127.0.0.1:{ep.port}" and ep.port > 0 + async with httpx.AsyncClient(base_url=ep.url) as client: + missing = await client.post("/v1/messages", json={}) + wrong = await client.post( + "/v1/messages", json={}, headers={"x-api-key": "nope"} + ) + bearer = await client.post("/v1/messages", json={}, headers=auth(ep)) + api_key = await client.post( + "/messages", json={}, headers={"x-api-key": ep.token} + ) + assert missing.status_code == 401 + assert wrong.status_code == 401 + assert "error" in wrong.json() + assert bearer.status_code == 200 + assert api_key.status_code == 200 + assert len(recorder.requests) == 2 + + +async def test_gateway_rewrites_headers_and_model() -> None: + recorder = Recorder( + httpx.Response( + 200, + json={"id": "m", "usage": {"input_tokens": 3, "output_tokens": 2}}, + ) + ) + async with gateway_endpoint(recorder, metadata={"run": "abc"}) as ep: + async with httpx.AsyncClient(base_url=ep.url) as client: + resp = await client.post( + "/v1/messages", + json={"model": "whatever", "max_tokens": 5}, + headers={ + "x-api-key": ep.token, + "anthropic-version": "2023-06-01", + "anthropic-beta": "tools-2024", + }, + ) + assert resp.status_code == 200 + sent = recorder.requests[0] + assert str(sent.url) == "https://gw.example.com/v1/messages" + assert sent.headers["authorization"] == "Bearer sk-gateway-secret" + assert "x-api-key" not in sent.headers + assert sent.headers["x-litellm-tags"] == "harness,claude_code" + assert json.loads(sent.headers["x-litellm-spend-logs-metadata"]) == {"run": "abc"} + assert sent.headers["anthropic-version"] == "2023-06-01" + assert sent.headers["anthropic-beta"] == "tools-2024" + assert json.loads(sent.content)["model"] == "claude-sonnet" + assert ep.usage.input_tokens == 3 and ep.usage.output_tokens == 2 + assert ep.usage.calls == 1 + + +@pytest.mark.parametrize( + "path,body,expected", + [ + ("/v1/messages", ANTHROPIC_SSE, (11, 7)), + ("/v1/chat/completions", CHAT_SSE, (5, 3)), + ("/responses", RESPONSES_SSE, (20, 4)), + ], +) +async def test_gateway_sse_passthrough_and_usage( + path: str, body: bytes, expected: tuple[int, int] +) -> None: + recorder = Recorder(sse_response(body)) + async with gateway_endpoint(recorder) as ep: + async with httpx.AsyncClient(base_url=ep.url) as client: + resp = await client.post(path, json={"stream": True}, headers=auth(ep)) + assert resp.status_code == 200 + assert resp.headers["content-type"].startswith("text/event-stream") + assert resp.content == body + assert (ep.usage.input_tokens, ep.usage.output_tokens) == expected + expected_cost = expected[0] * 0.001 + expected[1] * 0.002 + assert ep.usage.cost == pytest.approx(expected_cost) + + +async def test_cost_header_preferred_over_computed() -> None: + recorder = Recorder( + httpx.Response( + 200, + json={"usage": {"prompt_tokens": 100, "completion_tokens": 100}}, + headers={"x-litellm-response-cost": "0.42"}, + ) + ) + async with gateway_endpoint(recorder) as ep: + async with httpx.AsyncClient(base_url=ep.url) as client: + await client.post("/v1/chat/completions", json={}, headers=auth(ep)) + assert ep.usage.cost == pytest.approx(0.42) + assert ep.usage.snapshot() == Usage(input_tokens=100, output_tokens=100, calls=1) + + +async def test_gateway_error_status_preserved_and_not_counted() -> None: + recorder = Recorder(httpx.Response(429, json={"error": "rate limited"})) + async with gateway_endpoint(recorder) as ep: + async with httpx.AsyncClient(base_url=ep.url) as client: + resp = await client.post("/v1/chat/completions", json={}, headers=auth(ep)) + assert resp.status_code == 429 + assert ep.usage.calls == 0 + + +async def test_models_route() -> None: + recorder = Recorder(httpx.Response(200)) + async with gateway_endpoint(recorder) as ep: + async with httpx.AsyncClient(base_url=ep.url) as client: + with_model = await client.get("/v1/models", headers=auth(ep)) + unauth = await client.get("/models") + assert unauth.status_code == 401 + assert with_model.json()["object"] == "list" + assert [m["id"] for m in with_model.json()["data"]] == ["claude-sonnet"] + + async with ModelEndpoint(Harness.CODEX, None, None) as ep: + async with httpx.AsyncClient(base_url=ep.url) as client: + empty = await client.get("/models", headers=auth(ep)) + assert empty.json() == {"object": "list", "data": []} + + +async def test_sdk_chat_non_stream(monkeypatch: pytest.MonkeyPatch) -> None: + calls: list[dict[str, Any]] = [] + + async def fake_acompletion(**kwargs: Any) -> ModelResponse: + calls.append(kwargs) + response = ModelResponse( + model="gpt-x", + choices=[{"message": {"role": "assistant", "content": "hello"}}], + usage={"prompt_tokens": 9, "completion_tokens": 4, "total_tokens": 13}, + ) + response._hidden_params["response_cost"] = 0.5 + return response + + monkeypatch.setattr(litellm, "acompletion", fake_acompletion) + async with ModelEndpoint( + Harness.OPENCODE, "openai/gpt-x", None, api_key="sk-real", api_base="https://x" + ) as ep: + async with httpx.AsyncClient(base_url=ep.url) as client: + resp = await client.post( + "/v1/chat/completions", + json={ + "model": "ignored", + "messages": [{"role": "user", "content": "hi"}], + }, + headers=auth(ep), + ) + assert resp.status_code == 200 + assert resp.json()["choices"][0]["message"]["content"] == "hello" + assert calls[0]["model"] == "openai/gpt-x" + assert calls[0]["api_key"] == "sk-real" + assert calls[0]["api_base"] == "https://x" + assert (ep.usage.input_tokens, ep.usage.output_tokens) == (9, 4) + assert ep.usage.cost == pytest.approx(0.5) + + +async def fake_chat_stream() -> AsyncIterator[ModelResponseStream]: + yield ModelResponseStream(choices=[{"delta": {"content": "he"}}]) + yield ModelResponseStream(choices=[{"delta": {"content": "llo"}}]) + final = ModelResponseStream(choices=[]) + final.usage = litellm.Usage(prompt_tokens=6, completion_tokens=2, total_tokens=8) + yield final + + +async def test_sdk_chat_stream(monkeypatch: pytest.MonkeyPatch) -> None: + calls: list[dict[str, Any]] = [] + + async def fake_acompletion(**kwargs: Any) -> AsyncIterator[ModelResponseStream]: + calls.append(kwargs) + return fake_chat_stream() + + monkeypatch.setattr(litellm, "acompletion", fake_acompletion) + async with ModelEndpoint(Harness.OPENCODE, "openai/gpt-x", None) as ep: + async with httpx.AsyncClient(base_url=ep.url) as client: + resp = await client.post( + "/chat/completions", + json={"messages": [], "stream": True}, + headers=auth(ep), + ) + assert resp.headers["content-type"].startswith("text/event-stream") + lines = [line for line in resp.text.split("\n") if line.startswith("data: ")] + assert lines[-1] == "data: [DONE]" + assert json.loads(lines[0][6:])["choices"][0]["delta"]["content"] == "he" + assert calls[0]["stream_options"] == {"include_usage": True} + assert (ep.usage.input_tokens, ep.usage.output_tokens) == (6, 2) + assert ep.usage.cost == pytest.approx(6 * 0.001 + 2 * 0.002) + + +async def fake_anthropic_stream() -> AsyncIterator[Any]: + yield {"type": "message_start", "message": {"usage": {"input_tokens": 4}}} + yield b'event: message_delta\ndata: {"type":"message_delta","usage":{"output_tokens":9}}\n\n' + + +async def test_sdk_messages_stream_handles_dicts_and_bytes( + monkeypatch: pytest.MonkeyPatch, +) -> None: + async def fake_acreate(**kwargs: Any) -> AsyncIterator[Any]: + return fake_anthropic_stream() + + monkeypatch.setattr(litellm.anthropic.messages, "acreate", fake_acreate) + async with ModelEndpoint(Harness.CLAUDE_CODE, "anthropic/claude", None) as ep: + async with httpx.AsyncClient(base_url=ep.url) as client: + resp = await client.post( + "/v1/messages", json={"stream": True}, headers=auth(ep) + ) + assert "event: message_start" in resp.text + assert "event: message_delta" in resp.text + assert (ep.usage.input_tokens, ep.usage.output_tokens) == (4, 9) + + +async def test_sdk_error_is_sanitized(monkeypatch: pytest.MonkeyPatch) -> None: + async def failing(**kwargs: Any) -> Any: + raise litellm.RateLimitError( + message="too many requests for key sk-real", + llm_provider="openai", + model="gpt-x", + ) + + monkeypatch.setattr(litellm, "aresponses", failing) + async with ModelEndpoint(Harness.CODEX, "gpt-x", None, api_key="sk-real") as ep: + async with httpx.AsyncClient(base_url=ep.url) as client: + resp = await client.post("/v1/responses", json={}, headers=auth(ep)) + assert resp.status_code == 429 + assert "sk-real" not in resp.text + assert resp.json()["error"]["type"] == "RateLimitError" + + +async def test_missing_server_deps_raises_install_failed( + monkeypatch: pytest.MonkeyPatch, +) -> None: + def missing() -> Any: + raise HarnessInstallFailed(endpoint_module.MISSING_DEPS_MESSAGE) + + monkeypatch.setattr(endpoint_module, "_load_server_deps", missing) + with pytest.raises(HarnessInstallFailed, match="pip install starlette uvicorn"): + async with ModelEndpoint(Harness.CODEX, None, None): + pass + + +def test_load_server_deps_maps_import_error(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setitem(sys.modules, "uvicorn", None) + with pytest.raises(HarnessInstallFailed, match="starlette and uvicorn"): + endpoint_module._load_server_deps() + + +def test_usage_helpers() -> None: + assert usage_from_body({"usage": {"prompt_tokens": 1, "completion_tokens": 2}}) == ( + 1, + 2, + ) + assert usage_from_body({"response": {"usage": {"input_tokens": 3}}}) == (3, 0) + assert usage_from_body("nope") == (0, 0) + + parser = SSEUsageParser() + for i in range(0, len(ANTHROPIC_SSE), 7): # split across arbitrary chunk borders + parser.feed(ANTHROPIC_SSE[i : i + 7]) + parser.close() + assert (parser.input_tokens, parser.output_tokens) == (11, 7) + + tracker = UsageTracker() + tracker.add(1, 2, 0.1) + tracker.add(3, 4, 0.2) + assert tracker.snapshot() == Usage(input_tokens=4, output_tokens=6, calls=2) + assert tracker.cost == pytest.approx(0.3) + + +def test_compute_cost_never_raises(monkeypatch: pytest.MonkeyPatch) -> None: + def boom(**kwargs: Any) -> Any: + raise ValueError("unknown model") + + monkeypatch.setattr(litellm, "cost_per_token", boom) + assert compute_cost("mystery", 10, 10) == 0.0 + assert compute_cost(None, 10, 10) == 0.0 diff --git a/tests/unit/harness/test_init.py b/tests/unit/harness/test_init.py new file mode 100644 index 00000000000..effdc238249 --- /dev/null +++ b/tests/unit/harness/test_init.py @@ -0,0 +1,95 @@ +"""Tests for litellm/harness/__init__.py: the public API surface.""" + +from __future__ import annotations + + +from litellm import harness +from tests.test_litellm_rust.support.child_interpreter import run_child_interpreter +from litellm.utils import ProviderConfigManager + +PUBLIC_NAMES = [ + "Harness", + "agent", + "aagent", + "agent_session", + "aagent_session", + "agent_resume", + "aagent_resume", + "agent_capabilities", + "Result", + "Usage", + "State", + "Capabilities", + "Session", + "EventStream", + "Text", + "Reasoning", + "ToolCall", + "ToolResult", + "FileChange", + "Compaction", + "Approval", + "Done", + "Event", + "ClaudeCodeOptions", + "CodexOptions", + "OpenCodeOptions", + "DeepAgentsOptions", + "HarnessError", + "CapabilityUnsupported", + "OptionsMismatch", + "HarnessInstallFailed", + "SandboxError", + "SessionClosed", + "StateIncompatible", + "OutputInvalid", +] +ERROR_NAMES = [ + "CapabilityUnsupported", + "OptionsMismatch", + "HarnessInstallFailed", + "SandboxError", + "SessionClosed", + "StateIncompatible", + "OutputInvalid", +] +LAZY_IMPORT_CHECK = ( + "import sys, litellm\n" + "assert 'litellm.harness' not in sys.modules\n" + "h = litellm.harness\n" + "assert h.Harness.CODEX.value == 'codex'\n" + "assert 'starlette' not in sys.modules and 'uvicorn' not in sys.modules\n" + "print('ok')\n" +) + + +def test_public_api_names_exported(): + missing = [name for name in PUBLIC_NAMES if not hasattr(harness, name)] + assert missing == [] + assert set(PUBLIC_NAMES) <= set(harness.__all__) + + +def test_errors_share_base_class(): + for name in ERROR_NAMES: + assert issubclass(getattr(harness, name), harness.HarnessError) + + +def test_litellm_harness_attribute_is_lazy(): + out = run_child_interpreter(LAZY_IMPORT_CHECK, timeout=120) + assert out.returncode == 0, out.stderr + assert out.stdout.strip() == "ok" + + +def test_adapter_registry_paths_cover_every_harness(): + for member in harness.Harness: + config = ProviderConfigManager.get_provider_harness_config(member) + assert config is not None and config.harness is member + + +def test_litellm_agent_is_top_level_and_lazy(): + code = ( + "import sys, litellm; assert 'litellm.harness' not in sys.modules; " + "assert litellm.agent is litellm.harness.agent; assert litellm.Harness.CODEX.value == 'codex'" + ) + out = run_child_interpreter(code, timeout=120) + assert out.returncode == 0, out.stderr diff --git a/tests/unit/harness/test_runtime.py b/tests/unit/harness/test_runtime.py new file mode 100644 index 00000000000..e5f2235b794 --- /dev/null +++ b/tests/unit/harness/test_runtime.py @@ -0,0 +1,621 @@ +"""Tests for litellm/harness/runtime.py using a fake adapter, sandbox and endpoint.""" + +from __future__ import annotations + +import asyncio +import os +from collections.abc import AsyncIterator + +import pytest +from pydantic import BaseModel + +from litellm.harness import runtime +from litellm.harness.context import SessionContext +from litellm.harness.errors import ( + CapabilityUnsupported, + HarnessInstallFailed, + OptionsMismatch, + OutputInvalid, + SessionClosed, + StateIncompatible, +) +from litellm.harness.options import CodexOptions +from litellm.harness.types import ( + Approval, + Done, + Event, + FileChange, + Harness, + State, + Text, + ToolCall, +) +from tests.unit.harness.core_fakes import ( + NARROW_CAPS, + FakeAdapter, + FakeEndpoint, + FakeSandbox, + install_adapter, + script_approval, + wait_forever, +) + + +class Answer(BaseModel): + value: int + + +@pytest.fixture +def sandbox(tmp_path) -> FakeSandbox: + return FakeSandbox(str(tmp_path)) + + +async def _collect(stream) -> list[Event]: + return [event async for event in stream] + + +# -- validation --------------------------------------------------------------- + + +async def test_string_harness_raises_type_error_with_hint(monkeypatch, sandbox): + install_adapter(monkeypatch) + with pytest.raises(TypeError, match=r"Harness\.CODEX"): + await runtime.aagent("codex", "hi", sandbox=sandbox) # type: ignore[arg-type] + + +async def test_options_mismatch(monkeypatch, sandbox): + adapter_cls = install_adapter(monkeypatch) + with pytest.raises(OptionsMismatch, match="CodexOptions"): + await runtime.aagent( + Harness.CLAUDE_CODE, "hi", sandbox=sandbox, options=CodexOptions() + ) + assert adapter_cls.instances == [] + + +@pytest.mark.parametrize( + "kwargs", + [ + {"permissions": "edit"}, + {"output": Answer}, + {"tools": [print]}, + {"disable_tools": ["bash"]}, + {"permissions": "ask", "on_approval": lambda a: True}, + ], +) +async def test_capability_errors_before_start(monkeypatch, sandbox, kwargs): + adapter_cls = install_adapter(monkeypatch, caps=NARROW_CAPS) + with pytest.raises(CapabilityUnsupported): + await runtime.aagent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox, **kwargs) + assert all("start" not in a.calls for a in adapter_cls.instances) + assert FakeEndpoint.instances == [] + + +async def test_skills_capability_error_before_start(monkeypatch, sandbox, tmp_path): + skill = tmp_path / "skill" + skill.mkdir() + (skill / "SKILL.md").write_text("# s") + adapter_cls = install_adapter(monkeypatch, caps=NARROW_CAPS) + with pytest.raises(CapabilityUnsupported, match="skills"): + await runtime.aagent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox, skills=[skill]) + assert adapter_cls.instances == [] + + +async def test_skill_folder_without_skill_md_rejected(monkeypatch, sandbox, tmp_path): + install_adapter(monkeypatch) + with pytest.raises(ValueError, match=r"SKILL\.md"): + await runtime.aagent( + Harness.CLAUDE_CODE, "hi", sandbox=sandbox, skills=[tmp_path] + ) + + +async def test_ask_without_handler_only_allowed_for_stream(monkeypatch, sandbox): + install_adapter(monkeypatch) + with pytest.raises(ValueError, match="on_approval"): + await runtime.aagent( + Harness.CLAUDE_CODE, "hi", sandbox=sandbox, permissions="ask" + ) + stream = runtime.aagent( + Harness.CLAUDE_CODE, "hi", sandbox=sandbox, permissions="ask", stream=True + ) + events = await _collect(stream) + assert isinstance(events[-1], Done) + + +async def test_invalid_permissions_value(monkeypatch, sandbox): + install_adapter(monkeypatch) + with pytest.raises(ValueError, match="permissions"): + await runtime.aagent( + Harness.CLAUDE_CODE, "hi", sandbox=sandbox, permissions="yolo" # type: ignore[arg-type] + ) + + +# -- gateway routing (litellm_proxy/ prefix) --------------------------------- + + +def test_litellm_proxy_prefix_routes_through_gateway_env(monkeypatch): + monkeypatch.setenv("LITELLM_PROXY_API_BASE", "https://gw.example.com/") + monkeypatch.setenv("LITELLM_PROXY_API_KEY", "sk-test") + model, gateway = runtime.resolve_model_route("litellm_proxy/coder", None, None) + assert model == "coder" + assert gateway == runtime.GatewayTarget( + api_base="https://gw.example.com", api_key="sk-test" + ) + + +def test_litellm_proxy_call_args_win_over_env(monkeypatch): + monkeypatch.setenv("LITELLM_PROXY_API_BASE", "https://env.example.com") + monkeypatch.setenv("LITELLM_PROXY_API_KEY", "sk-env") + _, gateway = runtime.resolve_model_route( + "litellm_proxy/coder", "sk-arg", "https://arg.example.com" + ) + assert gateway == runtime.GatewayTarget( + api_base="https://arg.example.com", api_key="sk-arg" + ) + + +def test_litellm_proxy_without_base_raises(monkeypatch): + monkeypatch.delenv("LITELLM_PROXY_API_BASE", raising=False) + monkeypatch.setenv("LITELLM_PROXY_API_KEY", "sk-test") + with pytest.raises(ValueError, match="LITELLM_PROXY_API_BASE"): + runtime.resolve_model_route("litellm_proxy/coder", None, None) + + +def test_litellm_proxy_without_key_raises(monkeypatch): + monkeypatch.setenv("LITELLM_PROXY_API_BASE", "https://gw.example.com") + monkeypatch.setenv("LITELLM_PROXY_API_KEY", " ") + with pytest.raises(ValueError, match="LITELLM_PROXY_API_KEY"): + runtime.resolve_model_route("litellm_proxy/coder", None, None) + + +def test_plain_model_is_sdk_mode_even_with_gateway_env(monkeypatch): + monkeypatch.setenv("LITELLM_PROXY_API_BASE", "https://gw.example.com") + monkeypatch.setenv("LITELLM_PROXY_API_KEY", "sk-test") + assert runtime.resolve_model_route("anthropic/claude-sonnet-4-5", None, None) == ( + "anthropic/claude-sonnet-4-5", + None, + ) + + +def test_use_litellm_proxy_flag_routes_unprefixed_model(monkeypatch): + monkeypatch.setattr(runtime.litellm, "use_litellm_proxy", True) + monkeypatch.setenv("LITELLM_PROXY_API_BASE", "https://gw.example.com") + monkeypatch.setenv("LITELLM_PROXY_API_KEY", "sk-test") + model, gateway = runtime.resolve_model_route("coder", None, None) + assert model == "coder" and gateway is not None + + +async def test_gateway_passed_to_endpoint(monkeypatch, sandbox): + install_adapter(monkeypatch) + monkeypatch.setenv("LITELLM_PROXY_API_BASE", "https://gw.example.com") + monkeypatch.setenv("LITELLM_PROXY_API_KEY", "sk-test") + await runtime.aagent( + Harness.CLAUDE_CODE, "hi", sandbox=sandbox, model="litellm_proxy/m" + ) + endpoint = FakeEndpoint.instances[0] + assert endpoint.gateway.api_key == "sk-test" + assert endpoint.model == "m" + assert endpoint.entered and endpoint.exited + + +# -- event flow --------------------------------------------------------------- + + +async def test_text_and_tool_events_flow_and_done_last(monkeypatch, sandbox): + adapter_cls = install_adapter(monkeypatch) + events = await _collect( + runtime.aagent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox, stream=True) + ) + kinds = [type(e).__name__ for e in events] + assert kinds == ["Text", "ToolCall", "ToolResult", "Text", "Done"] + assert sum(isinstance(e, Done) for e in events) == 1 + result = events[-1].result + assert result.text == "hello world" + assert result.stop_reason == "done" + assert result.usage.input_tokens == 10 and result.usage.output_tokens == 5 + assert result.cost == pytest.approx(0.25) + assert adapter_cls.instances[0].calls == ["start", "turn", "stop"] + + +async def test_arun_returns_result(monkeypatch, sandbox): + install_adapter(monkeypatch) + result = await runtime.aagent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox) + assert result.text == "hello world" + assert len(result.events) == 4 + + +async def test_final_text_from_ctx_preferred(monkeypatch, sandbox): + async def script(adapter, ctx: SessionContext, prompt) -> AsyncIterator[Event]: + yield Text("partial") + ctx.final_text = "final answer" + + install_adapter(monkeypatch, script=script) + result = await runtime.aagent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox) + assert result.text == "final answer" + + +async def test_endpointless_adapter_usage(monkeypatch, sandbox): + install_adapter(monkeypatch, uses_endpoint=False) + result = await runtime.aagent(Harness.DEEPAGENTS, "hi", sandbox=sandbox) + assert FakeEndpoint.instances == [] + assert result.usage.calls == 1 + assert result.cost == pytest.approx(0.25) + + +async def test_stream_result_property(monkeypatch, sandbox): + install_adapter(monkeypatch) + stream = runtime.aagent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox, stream=True) + assert stream.result is None + await _collect(stream) + assert stream.result is not None and stream.result.text == "hello world" + + +# -- stop reasons ------------------------------------------------------------- + + +async def _tool_loop(adapter, ctx, prompt) -> AsyncIterator[Event]: + for i in range(10): + yield ToolCall(id=str(i), name="bash", native_name="Bash", input={}) + + +async def test_max_turns_stop(monkeypatch, sandbox): + adapter_cls = install_adapter(monkeypatch, script=_tool_loop) + result = await runtime.aagent( + Harness.CLAUDE_CODE, "hi", sandbox=sandbox, max_turns=3 + ) + assert result.stop_reason == "max_turns" + assert sum(isinstance(e, ToolCall) for e in result.events) == 3 + assert "stop" in adapter_cls.instances[0].calls + + +async def _slow(adapter, ctx, prompt) -> AsyncIterator[Event]: + yield Text("thinking") + await wait_forever() + yield Text("never") + + +async def test_timeout_stop(monkeypatch, sandbox): + install_adapter(monkeypatch, script=_slow) + result = await runtime.aagent( + Harness.CLAUDE_CODE, "hi", sandbox=sandbox, timeout=0.2 + ) + assert result.stop_reason == "timeout" + assert result.text == "thinking" + + +async def test_cancel_stop(monkeypatch, sandbox): + install_adapter(monkeypatch, script=_slow) + stream = runtime.aagent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox, stream=True) + events = [] + async for event in stream: + events.append(event) + if isinstance(event, Text): + stream.cancel() + assert isinstance(events[-1], Done) + assert events[-1].stop_reason == "cancelled" + + +async def _crash(adapter, ctx, prompt) -> AsyncIterator[Event]: + yield Text("partial") + raise RuntimeError("process exited with code 1") + + +async def test_runtime_error_stop_reason(monkeypatch, sandbox): + install_adapter(monkeypatch, script=_crash) + result = await runtime.aagent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox) + assert result.stop_reason == "runtime_error" + assert "process exited with code 1" in result.text + + +async def _missing_binary(adapter, ctx, prompt) -> AsyncIterator[Event]: + raise HarnessInstallFailed("claude not found on PATH") + yield Text("unreachable") # pragma: no cover + + +async def test_install_failed_propagates(monkeypatch, sandbox): + install_adapter(monkeypatch, script=_missing_binary) + with pytest.raises(HarnessInstallFailed, match="claude"): + await runtime.aagent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox) + assert FakeEndpoint.instances[0].exited + + +# -- approvals ---------------------------------------------------------------- + + +async def test_approval_on_approval_allow(monkeypatch, sandbox): + adapter_cls = install_adapter(monkeypatch, script=script_approval) + result = await runtime.aagent( + Harness.CLAUDE_CODE, + "hi", + sandbox=sandbox, + permissions="ask", + on_approval=lambda approval: approval.tool == "bash", + ) + assert adapter_cls.instances[0].approvals[0][0] is True + assert result.text == "allowed" + + +async def test_approval_async_handler_deny(monkeypatch, sandbox): + async def handler(approval: Approval) -> bool: + await asyncio.sleep(0) + return False + + adapter_cls = install_adapter(monkeypatch, script=script_approval) + result = await runtime.aagent( + Harness.CLAUDE_CODE, + "hi", + sandbox=sandbox, + permissions="ask", + on_approval=handler, + ) + assert adapter_cls.instances[0].approvals[0][0] is False + assert result.text == "denied" + + +async def test_approval_handler_raises_denies(monkeypatch, sandbox): + def handler(approval: Approval) -> bool: + raise RuntimeError("boom") + + adapter_cls = install_adapter(monkeypatch, script=script_approval) + result = await runtime.aagent( + Harness.CLAUDE_CODE, + "hi", + sandbox=sandbox, + permissions="ask", + on_approval=handler, + ) + allowed, reason = adapter_cls.instances[0].approvals[0] + assert allowed is False and "boom" in reason + assert result.stop_reason == "done" + + +async def test_stream_consumer_answers_approval(monkeypatch, sandbox): + adapter_cls = install_adapter(monkeypatch, script=script_approval) + stream = runtime.aagent( + Harness.CLAUDE_CODE, "hi", sandbox=sandbox, permissions="ask", stream=True + ) + async for event in stream: + if isinstance(event, Approval): + event.allow() + assert adapter_cls.instances[0].approvals[0][0] is True + + +async def test_unanswered_approval_denied(monkeypatch, sandbox): + adapter_cls = install_adapter(monkeypatch, script=script_approval) + events = await _collect( + runtime.aagent( + Harness.CLAUDE_CODE, "hi", sandbox=sandbox, permissions="ask", stream=True + ) + ) + allowed, reason = adapter_cls.instances[0].approvals[0] + assert allowed is False and "not answered" in reason + assert isinstance(events[-1], Done) + + +async def test_approval_without_ask_denied_in_run(monkeypatch, sandbox): + adapter_cls = install_adapter(monkeypatch, script=script_approval) + await runtime.aagent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox) + assert adapter_cls.instances[0].approvals[0][0] is False + + +# -- structured output -------------------------------------------------------- + + +def _answer_script(text: str, output_json: str | None = None): + async def script(adapter, ctx, prompt) -> AsyncIterator[Event]: + yield Text(text) + ctx.output_json = output_json + + return script + + +async def test_structured_output_from_text(monkeypatch, sandbox): + install_adapter( + monkeypatch, script=_answer_script('Sure. {"x": 1} then {"value": 42} done') + ) + result = await runtime.aagent( + Harness.CLAUDE_CODE, "hi", sandbox=sandbox, output=Answer + ) + assert result.output == Answer(value=42) + + +async def test_structured_output_from_ctx_output_json(monkeypatch, sandbox): + install_adapter(monkeypatch, script=_answer_script("ok", '{"value": 7}')) + result = await runtime.aagent( + Harness.CLAUDE_CODE, "hi", sandbox=sandbox, output=Answer + ) + assert result.output == Answer(value=7) + + +async def test_structured_output_invalid_carries_result(monkeypatch, sandbox): + install_adapter(monkeypatch, script=_answer_script('{"value": "nope"}')) + with pytest.raises(OutputInvalid) as info: + await runtime.aagent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox, output=Answer) + assert info.value.raw == '{"value": "nope"}' + assert info.value.result is not None + assert info.value.result.text == '{"value": "nope"}' + + +async def test_structured_output_missing_json(monkeypatch, sandbox): + install_adapter(monkeypatch, script=_answer_script("no json here")) + with pytest.raises(OutputInvalid, match="no JSON"): + await runtime.aagent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox, output=Answer) + + +async def test_stream_yields_done_before_output_invalid(monkeypatch, sandbox): + install_adapter(monkeypatch, script=_answer_script("nothing")) + stream = runtime.aagent( + Harness.CLAUDE_CODE, "hi", sandbox=sandbox, output=Answer, stream=True + ) + seen: list[object] = [] + with pytest.raises(OutputInvalid): + await _drain_into(stream, seen) + assert isinstance(seen[-1], Done) + + +async def _drain_into(stream, seen: list[object]) -> None: + async for event in stream: + seen.append(event) + + +def test_last_json_object(): + assert runtime.last_json_object('a {"a": {"b": 1}} b {"c": 2}') == '{"c": 2}' + assert runtime.last_json_object("{broken") is None + + +# -- files -------------------------------------------------------------------- + + +async def _edit_files(adapter, ctx, prompt) -> AsyncIterator[Event]: + root = ctx.sandbox.workdir + with open(os.path.join(root, "new.txt"), "w") as fh: + fh.write("new\n") + with open(os.path.join(root, "keep.txt"), "w") as fh: + fh.write("changed\n") + os.remove(os.path.join(root, "gone.txt")) + yield FileChange(path="new.txt", kind="created", diff=None) + yield Text("edited") + + +async def test_file_changes_emitted_once(monkeypatch, sandbox, tmp_path): + (tmp_path / "keep.txt").write_text("original\n") + (tmp_path / "gone.txt").write_text("bye\n") + install_adapter(monkeypatch, script=_edit_files) + events = await _collect( + runtime.aagent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox, stream=True) + ) + file_events = [e for e in events if isinstance(e, FileChange)] + assert sorted((e.path, e.kind) for e in file_events) == [ + ("gone.txt", "deleted"), + ("keep.txt", "modified"), + ("new.txt", "created"), + ] + result = events[-1].result + by_path = {f.path: f for f in result.files} + assert set(by_path) == {"gone.txt", "keep.txt", "new.txt"} + assert "+changed" in by_path["keep.txt"].diff + assert "-bye" in by_path["gone.txt"].diff + assert "+new" in by_path["new.txt"].diff + assert isinstance(events[-1], Done) + + +# -- sessions ----------------------------------------------------------------- + + +async def test_session_multi_turn_cost(monkeypatch, sandbox): + adapter_cls = install_adapter(monkeypatch) + async with runtime.aagent_session(Harness.CLAUDE_CODE, sandbox=sandbox) as session: + first = await session.arun("one") + second = await session.arun("two") + assert first.cost == pytest.approx(0.25) + assert second.cost == pytest.approx(0.25) + assert second.usage.input_tokens == 10 + assert session.cost == pytest.approx(0.5) + assert session.usage.calls == 2 + assert await session.history() == [ + {"role": "user", "content": "one"}, + {"role": "user", "content": "two"}, + ] + adapter = adapter_cls.instances[0] + assert adapter.calls == ["start", "turn", "turn", "stop"] + assert len(FakeEndpoint.instances) == 1 + with pytest.raises(SessionClosed): + await session.arun("three") + + +async def test_await_asession(monkeypatch, sandbox): + install_adapter(monkeypatch) + session = await runtime.aagent_session(Harness.CLAUDE_CODE, sandbox=sandbox) + result = await session.arun("hi") + await session.aclose() + assert result.text == "hello world" + + +async def test_session_restarts_after_timeout(monkeypatch, sandbox): + calls = {"n": 0} + + async def script(adapter, ctx, prompt) -> AsyncIterator[Event]: + calls["n"] += 1 + if calls["n"] == 1: + await wait_forever() + yield Text("ok") + + adapter_cls = install_adapter(monkeypatch, script=script) + async with runtime.aagent_session( + Harness.CLAUDE_CODE, sandbox=sandbox, timeout=0.2 + ) as session: + assert (await session.arun("one")).stop_reason == "timeout" + assert (await session.arun("two")).text == "ok" + adapter = adapter_cls.instances[0] + assert adapter.calls[:4] == ["start", "turn", "stop", "start"] + assert adapter.resumed_with == "native-123" + + +async def test_detach_state_round_trip_resume(monkeypatch, sandbox): + adapter_cls = install_adapter(monkeypatch) + async with runtime.aagent_session( + Harness.CLAUDE_CODE, sandbox=sandbox, model="m1" + ) as session: + await session.arun("one") + state = await session.adetach() + data = state.dumps() + assert b"sk-" not in data + restored = State.loads(data) + assert restored == state and restored.native_session_id == "native-123" + + async with runtime.aagent_resume(data, sandbox=sandbox) as resumed: + await resumed.arun("two") + new_adapter = adapter_cls.instances[-1] + assert new_adapter.calls[:2] == ["start", "resume"] + assert new_adapter.resumed_with == "native-123" + assert resumed.config.model == "m1" + + +async def test_resume_requires_capability(monkeypatch, sandbox): + install_adapter( + monkeypatch, + caps=NARROW_CAPS.__class__(**{**NARROW_CAPS.__dict__, "resume": False}), + ) + state = State(harness=Harness.CODEX, native_session_id="x", workdir="/tmp") + with pytest.raises(CapabilityUnsupported, match="resume"): + runtime.aagent_resume(state, sandbox=sandbox) + + +async def test_resume_state_without_native_id(monkeypatch, sandbox): + install_adapter(monkeypatch) + state = State(harness=Harness.CODEX, native_session_id=None, workdir="/tmp") + with pytest.raises(StateIncompatible): + runtime.aagent_resume(state, sandbox=sandbox) + + +async def test_history_requires_capability(monkeypatch, sandbox): + install_adapter(monkeypatch, caps=NARROW_CAPS) + async with runtime.aagent_session(Harness.CLAUDE_CODE, sandbox=sandbox) as session: + with pytest.raises(CapabilityUnsupported): + await session.history() + + +def test_capabilities_uses_registry(monkeypatch): + install_adapter(monkeypatch, caps=NARROW_CAPS) + assert runtime.agent_capabilities(Harness.CODEX) is NARROW_CAPS + with pytest.raises(TypeError): + runtime.agent_capabilities("codex") # type: ignore[arg-type] + + +def test_fake_adapter_is_a_harness_adapter(): + assert issubclass(FakeAdapter, runtime.BaseHarnessHandler) + + +async def test_turn_keeps_every_event_when_queue_overflows(monkeypatch, sandbox): + """A turn that emits more events than the queue holds must not drop any of them.""" + total = 40 + monkeypatch.setattr(runtime, "HARNESS_EVENT_QUEUE_MAX_SIZE", 4) + + async def burst(adapter, ctx, prompt): + for i in range(total): + yield Text(f"{i},") + + install_adapter(monkeypatch, script=burst) + result = await runtime.aagent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox) + texts = [e.delta for e in result.events if isinstance(e, Text)] + assert texts == [f"{i}," for i in range(total)] + assert result.stop_reason == "done" diff --git a/tests/unit/harness/test_sync.py b/tests/unit/harness/test_sync.py new file mode 100644 index 00000000000..a92f2957aee --- /dev/null +++ b/tests/unit/harness/test_sync.py @@ -0,0 +1,114 @@ +"""Tests for litellm/harness/sync.py: the sync bridge over the async runtime.""" + +from __future__ import annotations + +import asyncio +import threading + +import pytest + +from litellm.harness import sync +from litellm.harness.types import Approval, Done, Harness, State, Text +from tests.unit.harness.core_fakes import ( + FakeSandbox, + install_adapter, + script_approval, +) + + +@pytest.fixture +def sandbox(tmp_path) -> FakeSandbox: + return FakeSandbox(str(tmp_path)) + + +async def _call_run_in_loop(sandbox: FakeSandbox) -> None: + sync.agent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox) + + +def test_sync_run_from_plain_code(monkeypatch, sandbox): + install_adapter(monkeypatch) + result = sync.agent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox) + assert result.text == "hello world" + assert result.stop_reason == "done" + + +def test_sync_stream_from_plain_code(monkeypatch, sandbox): + install_adapter(monkeypatch) + stream = sync.agent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox, stream=True) + events = list(stream) + assert isinstance(events[-1], Done) + assert [e.delta for e in events if isinstance(e, Text)] == ["hello ", "world"] + assert stream.result is not None and stream.result.text == "hello world" + assert list(stream) == [] + + +def test_sync_stream_answers_approval(monkeypatch, sandbox): + adapter_cls = install_adapter(monkeypatch, script=script_approval) + for event in sync.agent( + Harness.CLAUDE_CODE, "hi", sandbox=sandbox, permissions="ask", stream=True + ): + if isinstance(event, Approval): + event.allow() + assert adapter_cls.instances[0].approvals[0][0] is True + + +def test_sync_stream_close_early(monkeypatch, sandbox): + adapter_cls = install_adapter(monkeypatch) + with sync.agent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox, stream=True) as stream: + next(stream) + assert adapter_cls.instances[0].calls[-1] == "stop" + + +def test_sync_validation_errors_raise_eagerly(monkeypatch, sandbox): + install_adapter(monkeypatch) + with pytest.raises(TypeError, match=r"Harness\.OPENCODE"): + sync.agent("opencode", "hi", sandbox=sandbox, stream=True) # type: ignore[arg-type] + + +def test_sync_session_multi_turn_and_detach(monkeypatch, sandbox): + adapter_cls = install_adapter(monkeypatch) + with sync.agent_session(Harness.CLAUDE_CODE, sandbox=sandbox) as session: + session.run("one") + events = list(session.stream("two")) + assert isinstance(events[-1], Done) + assert session.cost == pytest.approx(0.5) + assert len(session.history()) == 2 + state = session.detach() + assert isinstance(state, State) + with sync.agent_resume(state.dumps(), sandbox=sandbox) as resumed: + assert resumed.run("three").text == "hello world" + assert adapter_cls.instances[-1].resumed_with == "native-123" + + +def test_sync_session_stop_returns_state(monkeypatch, sandbox): + install_adapter(monkeypatch) + session = sync.agent_session(Harness.CLAUDE_CODE, sandbox=sandbox).start() + session.run("one") + state = session.stop() + assert state.native_session_id == "native-123" + + +def test_single_background_loop_thread(monkeypatch, sandbox): + install_adapter(monkeypatch) + sync.agent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox) + first = sync._LOOP.loop() + sync.agent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox) + assert sync._LOOP.loop() is first + names = [t.name for t in threading.enumerate()] + assert names.count("litellm-harness-loop") == 1 + + +async def test_run_inside_event_loop_raises(monkeypatch, sandbox): + install_adapter(monkeypatch) + with pytest.raises(RuntimeError, match="aagent"): + sync.agent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox) + with pytest.raises(RuntimeError, match="aagent"): + sync.agent(Harness.CLAUDE_CODE, "hi", sandbox=sandbox, stream=True) + with pytest.raises(RuntimeError, match="aagent_session"): + sync.agent_session(Harness.CLAUDE_CODE, sandbox=sandbox) + + +def test_run_inside_asyncio_run_raises(monkeypatch, sandbox): + install_adapter(monkeypatch) + with pytest.raises(RuntimeError, match=r"await litellm\.aagent"): + asyncio.run(_call_run_in_loop(sandbox)) diff --git a/tests/unit/harness/test_types.py b/tests/unit/harness/test_types.py new file mode 100644 index 00000000000..036d6dd625c --- /dev/null +++ b/tests/unit/harness/test_types.py @@ -0,0 +1,90 @@ +"""Tests for litellm/harness/types.py.""" + +from __future__ import annotations + +import asyncio + +import pytest + +from litellm.harness.errors import StateIncompatible +from litellm.harness.types import ( + Approval, + Done, + Harness, + Result, + State, + Usage, + require_harness, +) + + +def test_harness_is_plain_enum(): + assert Harness.CODEX.value == "codex" + assert not isinstance(Harness.CODEX, str) + + +@pytest.mark.parametrize( + "given,hint", + [ + ("codex", "Harness.CODEX"), + ("claude-code", "Harness.CLAUDE_CODE"), + ("OPENCODE", "Harness.OPENCODE"), + ], +) +def test_require_harness_hint(given, hint): + with pytest.raises(TypeError, match=hint): + require_harness(given) + + +def test_require_harness_no_hint_for_unknown(): + with pytest.raises(TypeError) as info: + require_harness(42) + assert "Did you mean" not in str(info.value) + assert require_harness(Harness.DEEPAGENTS) is Harness.DEEPAGENTS + + +def test_usage_total_tokens(): + assert Usage(input_tokens=3, output_tokens=4, calls=1).total_tokens == 7 + + +def test_done_exposes_result_fields(): + result = Result( + text="t", + output=None, + files=[], + events=[], + usage=Usage(1, 2, 1), + cost=0.5, + stop_reason="done", + session_id="s", + ) + done = Done(result) + assert done.usage.total_tokens == 3 + assert done.cost == 0.5 + assert done.stop_reason == "done" + + +def test_state_round_trip_and_errors(): + state = State(Harness.CODEX, "thread-1", "/work", model="gpt") + assert State.loads(state.dumps()) == state + with pytest.raises(StateIncompatible): + State.loads(b"not json") + with pytest.raises(StateIncompatible): + State.loads(b'{"harness": "nope", "version": 1, "workdir": "/"}') + with pytest.raises(StateIncompatible, match="version"): + State.loads(b'{"harness": "codex", "version": 99, "workdir": "/"}') + + +async def test_approval_allow_deny_once(): + approval = Approval(tool="bash", input={}) + assert not approval.answered + approval.allow() + approval.deny("late") + assert await approval.wait() == (True, "") + assert approval.answered + + +async def test_approval_resolved_from_other_thread(): + approval = Approval(tool="bash", input={}) + await asyncio.to_thread(approval.deny, "nope") + assert await approval.wait() == (False, "nope") diff --git a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 1e12a973cdb..8d2e6b9fd0c 100644 --- a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -28,6 +28,7 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( sanitize_messages_for_tool_calling, ) from litellm.types.llms.openai import ChatCompletionToolMessage +from litellm.utils import validate_and_fix_openai_messages def _get_gemini_function_response_inline_data_parts(result): @@ -4095,3 +4096,168 @@ def test_is_unsignable_thinking_block_treats_whitespace_only_as_empty(): } assert is_unsignable_thinking_block(whitespace_only_block) is True + + +_CONTENT_LESS_USER_MESSAGES: Final = ({"role": "user"}, {"role": "user", "content": None}) +_CONTENT_LESS_TOOL_MESSAGES: Final = ( + {"role": "tool", "tool_call_id": "call_1"}, + {"role": "tool", "tool_call_id": "call_1", "content": None}, +) +_BOSTON_WEATHER_TOOL_CALL_TURN: Final = ( + {"role": "user", "content": "What is the weather in Boston?"}, + { + "role": "assistant", + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": {"name": "get_weather", "arguments": '{"city": "Boston"}'}, + } + ], + }, +) + + +def _conversation_around( + content_less_user_message: dict[str, object], +) -> tuple[list[dict[str, object]], list[dict[str, object]]]: + with_message: Final = [ + {"role": "user", "content": "What is the capital of France?"}, + content_less_user_message, + {"role": "assistant", "content": "Paris."}, + {"role": "user", "content": "And of Spain?"}, + ] + without_message: Final = [message for message in with_message if message is not content_less_user_message] + return validate_and_fix_openai_messages(with_message), validate_and_fix_openai_messages(without_message) + + +@pytest.mark.parametrize("content_less_user_message", _CONTENT_LESS_USER_MESSAGES) +def test_bedrock_converse_messages_pt_user_message_without_content_adds_no_block( + content_less_user_message: dict[str, object], +): + with_message, without_message = _conversation_around(content_less_user_message) + + assert _bedrock_converse_messages_pt( + messages=with_message, model="anthropic.claude-haiku-4-5", llm_provider="bedrock" + ) == _bedrock_converse_messages_pt(messages=without_message, model="anthropic.claude-haiku-4-5", llm_provider="bedrock") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("content_less_user_message", _CONTENT_LESS_USER_MESSAGES) +async def test_bedrock_converse_messages_pt_async_user_message_without_content_adds_no_block( + content_less_user_message: dict[str, object], +): + with_message, without_message = _conversation_around(content_less_user_message) + + assert await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async( + messages=with_message, model="anthropic.claude-haiku-4-5", llm_provider="bedrock" + ) == await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async( + messages=without_message, model="anthropic.claude-haiku-4-5", llm_provider="bedrock" + ) + + +@pytest.mark.parametrize("content_less_tool_message", _CONTENT_LESS_TOOL_MESSAGES) +def test_bedrock_converse_messages_pt_tool_message_without_content_yields_empty_tool_result( + content_less_tool_message: dict[str, object], +): + result: Final = _bedrock_converse_messages_pt( + messages=validate_and_fix_openai_messages([*_BOSTON_WEATHER_TOOL_CALL_TURN, content_less_tool_message]), + model="anthropic.claude-haiku-4-5", + llm_provider="bedrock", + ) + + tool_result: Final = result[-1]["content"][0]["toolResult"] + assert result[-1]["role"] == "user" + assert tool_result["toolUseId"] == "call_1" + assert tool_result["content"] == [] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("content_less_tool_message", _CONTENT_LESS_TOOL_MESSAGES) +async def test_bedrock_converse_messages_pt_async_tool_message_without_content_yields_empty_tool_result( + content_less_tool_message: dict[str, object], +): + result: Final = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async( + messages=validate_and_fix_openai_messages([*_BOSTON_WEATHER_TOOL_CALL_TURN, content_less_tool_message]), + model="anthropic.claude-haiku-4-5", + llm_provider="bedrock", + ) + + tool_result: Final = result[-1]["content"][0]["toolResult"] + assert tool_result["toolUseId"] == "call_1" + assert tool_result["content"] == [] + + +def test_bedrock_converse_messages_pt_blank_user_text_sends_the_continue_message_text(): + continue_message: Final = {"role": "user", "content": "Please continue."} + blank_last_turn: Final = [ + {"role": "user", "content": "Hello"}, + {"role": "assistant", "content": "Hi."}, + {"role": "user", "content": " "}, + ] + explicit_last_turn: Final = [*blank_last_turn[:2], continue_message] + + assert _bedrock_converse_messages_pt( + messages=blank_last_turn, + model="anthropic.claude-haiku-4-5", + llm_provider="bedrock", + user_continue_message=continue_message, + ) == _bedrock_converse_messages_pt( + messages=explicit_last_turn, + model="anthropic.claude-haiku-4-5", + llm_provider="bedrock", + user_continue_message=continue_message, + ) + + +@pytest.mark.parametrize("content_less_user_message", _CONTENT_LESS_USER_MESSAGES) +def test_bedrock_converse_messages_pt_lone_content_less_user_turn_sends_the_continue_message( + content_less_user_message: dict[str, object], +): + continue_message: Final = {"role": "user", "content": "Please continue."} + + assert _bedrock_converse_messages_pt( + messages=validate_and_fix_openai_messages([content_less_user_message]), + model="anthropic.claude-haiku-4-5", + llm_provider="bedrock", + user_continue_message=continue_message, + ) == _bedrock_converse_messages_pt( + messages=[continue_message], + model="anthropic.claude-haiku-4-5", + llm_provider="bedrock", + user_continue_message=continue_message, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("content_less_user_message", _CONTENT_LESS_USER_MESSAGES) +async def test_bedrock_converse_messages_pt_async_lone_content_less_user_turn_continues_under_modify_params( + content_less_user_message: dict[str, object], monkeypatch: pytest.MonkeyPatch +): + monkeypatch.setattr(litellm, "modify_params", True) + + assert await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async( + messages=validate_and_fix_openai_messages([content_less_user_message]), + model="anthropic.claude-haiku-4-5", + llm_provider="bedrock", + ) == await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async( + messages=[{"role": "user", "content": ""}], + model="anthropic.claude-haiku-4-5", + llm_provider="bedrock", + ) + + +@pytest.mark.parametrize("content_less_user_message", _CONTENT_LESS_USER_MESSAGES) +def test_bedrock_converse_messages_pt_lone_content_less_user_turn_adds_no_block_without_a_continue_message( + content_less_user_message: dict[str, object], monkeypatch: pytest.MonkeyPatch +): + monkeypatch.setattr(litellm, "modify_params", False) + + assert ( + _bedrock_converse_messages_pt( + messages=validate_and_fix_openai_messages([content_less_user_message]), + model="anthropic.claude-haiku-4-5", + llm_provider="bedrock", + ) + == [] + ) diff --git a/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py b/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py index 755c7617701..5540cf54193 100644 --- a/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py +++ b/tests/unit/litellm_proxy_extras/test_litellm_proxy_extras_utils.py @@ -2,7 +2,9 @@ import glob import os import re import sys +import threading from pathlib import Path +from typing import Final import pytest @@ -1024,3 +1026,143 @@ class TestJWTKeyMappingCascade: f"{path} must declare onDelete: Cascade on the JWT key mapping " "relation (issue #33702)" ) + + + +class TestBuildRequestLogIndexes: + """The migration job hands the index build the direct database URL and the schema + the migrations target, waits for it, and reports its result.""" + + @pytest.fixture + def builds(self): + return [] + + @pytest.fixture + def build(self, builds): + def record(database_url: str, schema: str) -> bool: + builds.append((database_url, schema)) + return True + + return record + + def test_the_build_gets_the_direct_url_without_prisma_params_and_the_prisma_schema(self, monkeypatch, builds, build): + monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@pooler:6543/db?schema=tenant&pgbouncer=true") + monkeypatch.setenv("DIRECT_URL", "postgresql://u:p@primary:5432/db?connection_limit=1") + + assert ProxyExtrasDBManager.build_request_log_indexes(build=build) is True + + assert builds == [("postgresql://u:p@primary:5432/db", "tenant")] + + def test_the_build_defaults_to_the_database_url_and_the_public_schema(self, monkeypatch, builds, build): + monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@primary:5432/db") + monkeypatch.delenv("DIRECT_URL", raising=False) + + assert ProxyExtrasDBManager.build_request_log_indexes(build=build) is True + + assert builds == [("postgresql://u:p@primary:5432/db", "public")] + + def test_a_build_that_leaves_indexes_missing_is_reported_so_the_job_reruns(self, monkeypatch): + monkeypatch.setenv("DATABASE_URL", "postgresql://u:p@primary:5432/db") + + assert ProxyExtrasDBManager.build_request_log_indexes(build=lambda url, schema: False) is False + + def test_without_a_database_url_nothing_is_built(self, monkeypatch, builds, build): + monkeypatch.delenv("DATABASE_URL", raising=False) + + assert ProxyExtrasDBManager.build_request_log_indexes(build=build) is True + + assert builds == [] + + +class TestStartRequestLogIndexBuild: + """A serving proxy that ran the migrations starts the index build on a daemon thread + and goes on to serve while it runs.""" + + def test_the_build_runs_on_a_daemon_thread_that_does_not_hold_up_the_caller(self): + release: Final = threading.Event() + builds: Final[list[str]] = [] # mutable-ok: the builder thread hands back the thread it ran on + + def build() -> bool: + assert release.wait(5), "the caller never came back from start_request_log_index_build" + builds.append(threading.current_thread().name) + return True + + thread: Final = ProxyExtrasDBManager.start_request_log_index_build(build=build) + + assert builds == [], "the build ran before start_request_log_index_build returned" + assert thread.daemon is True + release.set() + thread.join(5) + assert builds == ["litellm-request-log-indexes"] + + +class TestRunMigrationJob: + """`run_migration_job` is `setup_database` followed by the index build, each step's + result deciding whether the job reports success.""" + + @pytest.fixture + def calls(self): + return [] + + @pytest.fixture + def setup(self, calls): + def record(result: bool): + def setup_database(use_migrate: bool, use_v2_resolver: bool) -> bool: + calls.append(("setup", use_migrate, use_v2_resolver)) + return result + + return setup_database + + return record + + @pytest.fixture + def build(self, calls): + def record(result: bool): + def build_request_log_indexes() -> bool: + calls.append(("build",)) + return result + + return build_request_log_indexes + + return record + + def test_the_job_builds_the_indexes_after_the_migrations_succeed(self, calls, setup, build): + assert ProxyExtrasDBManager.run_migration_job(True, False, setup=setup(True), build=build(True)) is True + + assert calls == [("setup", True, False), ("build",)] + + def test_the_job_fails_without_building_when_the_migrations_fail(self, calls, setup, build): + assert ProxyExtrasDBManager.run_migration_job(True, True, setup=setup(False), build=build(True)) is False + + assert calls == [("setup", True, True)] + + def test_the_job_fails_when_an_index_could_not_be_built(self, calls, setup, build): + assert ProxyExtrasDBManager.run_migration_job(True, True, setup=setup(True), build=build(False)) is False + + assert calls == [("setup", True, True), ("build",)] + + +class TestMigrationJobOwnedDrift: + JOB_INDEXES = ( + "-- CreateIndex\n" + 'CREATE INDEX "LiteLLM_SpendLogs_litellm_call_id_idx" ON "LiteLLM_SpendLogs"("litellm_call_id");\n' + "\n-- CreateIndex\n" + 'CREATE INDEX "LiteLLM_SpendLogs_api_key_startTime_idx" ON "LiteLLM_SpendLogs"("api_key", "startTime");\n' + ) + + def test_a_plain_spend_logs_table_only_loses_the_migration_job_indexes(self): + filtered = ProxyExtrasDBManager._filter_migration_job_owned_drift( + _PARTITIONED_DRIFT_SQL + self.JOB_INDEXES, partitioned=False + ) + assert "LiteLLM_SpendLogs_litellm_call_id_idx" not in filtered + assert "LiteLLM_SpendLogs_api_key_startTime_idx" not in filtered + assert 'PRIMARY KEY ("request_id")' in filtered + + def test_a_partitioned_spend_logs_table_also_loses_its_partitioning_artifacts(self): + filtered = ProxyExtrasDBManager._filter_migration_job_owned_drift( + _PARTITIONED_DRIFT_SQL + self.JOB_INDEXES, partitioned=True + ) + assert "LiteLLM_SpendLogs_litellm_call_id_idx" not in filtered + assert 'PRIMARY KEY ("request_id")' not in filtered + assert "LiteLLM_SpendLogs_legacy" not in filtered + assert 'ALTER TABLE "LiteLLM_BudgetTable" ADD COLUMN "updated_by" TEXT;' in filtered diff --git a/tests/unit/litellm_proxy_extras/test_request_log_indexes.py b/tests/unit/litellm_proxy_extras/test_request_log_indexes.py new file mode 100644 index 00000000000..5cf7593c5cd --- /dev/null +++ b/tests/unit/litellm_proxy_extras/test_request_log_indexes.py @@ -0,0 +1,149 @@ +import re +from pathlib import Path +from typing import Final + +import pytest +from litellm_proxy_extras.migration_recovery import is_inert_migration +from litellm_proxy_extras.request_log_indexes import ( + REQUEST_LOG_INDEXES, + RequestLogIndex, + filter_request_log_index_diff, +) + +PACKAGE: Final = Path(__file__).resolve().parents[3] / "litellm-proxy-extras" / "litellm_proxy_extras" +SCHEMA: Final = PACKAGE / "schema.prisma" +INERT_MIGRATIONS: Final = ( + "20260823000000_add_spend_logs_api_key_starttime_index", + "20260831120001_spend_logs_litellm_call_id_index", +) +CALL_ID_INDEX: Final = RequestLogIndex( + "LiteLLM_SpendLogs", "LiteLLM_SpendLogs_litellm_call_id_idx", '("litellm_call_id")' +) + + +def _prisma_indexes_of(schema: str, model: str) -> frozenset[str]: + """The index names Prisma derives for a model's @@index declarations: __idx.""" + body: Final = re.search(rf"model {model} \{{(.*?)\n\}}", schema, re.DOTALL) + assert body is not None, model + declarations: Final[tuple[str, ...]] = tuple( + match.group(1) for match in re.finditer(r"@@index\(\[([^\]]+)\]\)", body.group(1)) + ) + return frozenset( + f"{model}_{'_'.join(column.strip() for column in columns.split(','))}_idx" for columns in declarations + ) + + +class TestTheIndexList: + def test_every_migration_job_index_is_declared_in_the_prisma_schema_under_the_same_name(self): + schema: Final = SCHEMA.read_text() + for index in REQUEST_LOG_INDEXES: + assert index.name in _prisma_indexes_of(schema, index.table), index + + @pytest.mark.parametrize("name", INERT_MIGRATIONS) + def test_the_migrations_that_used_to_build_these_indexes_run_no_sql(self, name: str): + assert is_inert_migration((PACKAGE / "migrations" / name / "migration.sql").read_text()) + + +class TestIsInertMigration: + @pytest.mark.parametrize( + "script", + ( + "", + "-- only a comment\n", + "/* block */\n-- line\n", + "-- a semicolon; in a comment\n", + ";\n;", + "-- why\nSELECT 1;\n", + "select 1", + ), + ids=( + "empty", + "line-comment", + "both-comments", + "semicolon-in-comment", + "bare-separators", + "select-1", + "lowercase", + ), + ) + def test_comments_and_a_select_1_alone_are_inert(self, script: str): + assert is_inert_migration(script) is True + + @pytest.mark.parametrize( + "script", + ( + "SELECT 2;", + 'SELECT 1 FROM "LiteLLM_SpendLogs";', + '-- comment\nCREATE INDEX "ix" ON "t" ("a");', + "/* c */ ALTER TABLE t ADD COLUMN a TEXT", + 'SELECT 1; DROP INDEX "ix";', + ), + ids=("select-2", "select-from", "index-after-comment", "alter-after-block-comment", "drop-after-select-1"), + ) + def test_any_statement_is_not_inert(self, script: str): + assert is_inert_migration(script) is False + + +class TestPartitionIndexName: + def test_a_partition_gets_the_name_postgres_would_give_an_inherited_index(self): + assert CALL_ID_INDEX.partition_index_name("LiteLLM_SpendLogs_p2026_09") == ( + "LiteLLM_SpendLogs_p2026_09_litellm_call_id_idx" + ) + + def test_an_index_not_prefixed_by_its_table_keeps_its_whole_name(self): + index = RequestLogIndex("LiteLLM_SpendLogs", "call_id_lookup", '("litellm_call_id")') + assert index.partition_index_name("LiteLLM_SpendLogs_pdefault") == "LiteLLM_SpendLogs_pdefault_call_id_lookup" + + def test_a_long_name_is_cut_to_63_bytes_with_a_digest_that_keeps_partitions_apart(self): + first = CALL_ID_INDEX.partition_index_name("LiteLLM_SpendLogs_p" + "x" * 50 + "_2026_09") + second = CALL_ID_INDEX.partition_index_name("LiteLLM_SpendLogs_p" + "x" * 50 + "_2026_10") + assert len(first.encode()) == 63 and len(second.encode()) == 63 + assert first != second + assert first.startswith("LiteLLM_SpendLogs_p") and first[-9] == "_" + + def test_the_byte_limit_counts_multibyte_characters(self): + name = CALL_ID_INDEX.partition_index_name("é" * 40) + assert len(name.encode()) <= 63 and len(name) < 63 + + +class TestColumns: + def test_the_columns_are_the_quoted_names_of_the_definition_in_order(self): + index = RequestLogIndex( + "LiteLLM_SpendLogs", "LiteLLM_SpendLogs_api_key_startTime_idx", '("api_key", "startTime")' + ) + assert index.columns == ("api_key", "startTime") + + def test_every_migration_job_index_names_at_least_one_column(self): + assert all(index.columns for index in REQUEST_LOG_INDEXES) + + +DRIFT_WITH_BOTH_INDEXES: Final = ( + "-- CreateIndex\n" + 'CREATE INDEX "LiteLLM_SpendLogs_litellm_call_id_idx" ON "LiteLLM_SpendLogs"("litellm_call_id");\n' + "\n" + "-- CreateIndex\n" + 'CREATE INDEX "LiteLLM_SpendLogs_api_key_startTime_idx" ON "LiteLLM_SpendLogs"("api_key", "startTime");\n' +) + + +class TestFilterRequestLogIndexDiff: + def test_a_drift_script_that_only_creates_the_migration_job_indexes_becomes_empty(self): + assert filter_request_log_index_diff(DRIFT_WITH_BOTH_INDEXES) == "" + + def test_other_statements_survive_with_the_migration_job_indexes_removed(self): + other: Final = '-- AlterTable\nALTER TABLE "LiteLLM_BudgetTable" ADD COLUMN "updated_by" TEXT;\n' + filtered = filter_request_log_index_diff(other + DRIFT_WITH_BOTH_INDEXES) + assert 'ALTER TABLE "LiteLLM_BudgetTable" ADD COLUMN "updated_by" TEXT;' in filtered + assert "LiteLLM_SpendLogs_litellm_call_id_idx" not in filtered + assert "LiteLLM_SpendLogs_api_key_startTime_idx" not in filtered + + def test_an_index_of_another_name_on_spend_logs_is_kept(self): + sql: Final = 'CREATE INDEX "LiteLLM_SpendLogs_end_user_idx" ON "LiteLLM_SpendLogs"("end_user");\n' + assert filter_request_log_index_diff(sql) == sql + + def test_a_drop_of_a_migration_job_index_is_kept_for_the_operator_to_see(self): + sql: Final = 'DROP INDEX "LiteLLM_SpendLogs_litellm_call_id_idx";\n' + assert filter_request_log_index_diff(sql) == sql + + def test_an_empty_script_stays_empty(self): + assert filter_request_log_index_diff("") == "" diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_per_turn_control.py b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_per_turn_control.py index ef1fac9e120..05c12c6a285 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_per_turn_control.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_per_turn_control.py @@ -150,3 +150,33 @@ def test_native_messages_thinking_display_updates_beta(display: str | None, expl ) assert headers.get("anthropic-beta", "").split(",").count(beta) == int(display == "updates" or explicit_beta) + + +@pytest.mark.parametrize("action", (None, "tool_addition", "tool_removal")) +@pytest.mark.parametrize("explicit_beta", (False, True)) +def test_native_messages_tool_changes_beta(action: str | None, explicit_beta: bool) -> None: + from typing import Final + + from litellm.types.llms.anthropic import ( + ANTHROPIC_MID_CONVERSATION_TOOL_CHANGES_BETA_HEADER, + ANTHROPIC_THINKING_DISPLAY_UPDATES_BETA_HEADER, + ) + + beta: Final = ANTHROPIC_MID_CONVERSATION_TOOL_CHANGES_BETA_HEADER + content: Final = ( + [{"type": action, "tool": {"type": "tool_reference", "name": "mcp__test__ping"}}] + if action + else "Answer briefly" + ) + headers, _ = AnthropicMessagesConfig().validate_anthropic_messages_environment( + headers={"anthropic-beta": beta} if explicit_beta else {}, + model="claude-fable-5-1", + messages=["not a message dict", {"role": "user", "content": "Hello"}, {"role": "system", "content": content}], + optional_params={"thinking": {"type": "adaptive", "display": "updates"}}, + litellm_params={}, + api_key="sk-ant-test", + ) + + assert headers.get("anthropic-beta", "").split(",").count(beta) == int(action is not None or explicit_beta) + + assert ANTHROPIC_THINKING_DISPLAY_UPDATES_BETA_HEADER in headers.get("anthropic-beta", "").split(",") diff --git a/tests/unit/llms/anthropic/test_anthropic_common_utils.py b/tests/unit/llms/anthropic/test_anthropic_common_utils.py index 0904a20a16b..68e2e9650d5 100644 --- a/tests/unit/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/unit/llms/anthropic/test_anthropic_common_utils.py @@ -2450,3 +2450,57 @@ def test_shared_legacy_thinking_translation_preserves_supported_display( ) assert optional_params["thinking"] == expected_thinking + + +@pytest.mark.usefixtures("local_model_cost_map", "local_beta_headers_config") +@pytest.mark.parametrize("action", (None, "tool_addition", "tool_removal")) +@pytest.mark.parametrize("explicit_beta", (False, True)) +def test_validate_environment_adds_tool_changes_beta(action: str | None, explicit_beta: bool) -> None: + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + from litellm.types.llms.anthropic import ANTHROPIC_MID_CONVERSATION_TOOL_CHANGES_BETA_HEADER + + beta: Final = ANTHROPIC_MID_CONVERSATION_TOOL_CHANGES_BETA_HEADER + content: Final = ( + [{"type": action, "tool": {"type": "tool_reference", "name": "mcp__test__ping"}}] + if action + else "Answer briefly" + ) + headers: Final = AnthropicModelInfo().validate_environment( + headers={"anthropic-beta": beta} if explicit_beta else {}, + model="claude-fable-5-1", + messages=[{"role": "user", "content": "Hello"}, {"role": "system", "content": content}], + optional_params={}, + litellm_params={}, + api_key=FAKE_REGULAR_KEY, + ) + + assert headers.get("anthropic-beta", "").split(",").count(beta) == int(action is not None or explicit_beta) + assert headers["x-api-key"] == FAKE_REGULAR_KEY + + +@pytest.mark.parametrize( + ("role", "content"), + ( + ("user", [{"type": "tool_addition", "tool": {"type": "tool_reference", "name": "ping"}}]), + ("assistant", [{"type": "tool_addition", "tool": {"type": "tool_reference", "name": "ping"}}]), + ("system", "tool_addition"), + ("system", None), + ("system", ["tool_addition"]), + ("system", [{"type": "tool_reference", "name": "ping"}]), + ("system", [{"type": "tool_addition", "tool": {"type": "tool_definition", "definition": {"name": "ping"}}}]), + ), +) +def test_tool_changes_beta_requires_system_tool_reference(role: str, content: object) -> None: + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + from litellm.types.llms.anthropic import ANTHROPIC_MID_CONVERSATION_TOOL_CHANGES_BETA_HEADER + + headers: Final = AnthropicModelInfo().validate_environment( + headers={}, + model="claude-fable-5-1", + messages=[{"role": role, "content": content}], + optional_params={}, + litellm_params={}, + api_key=FAKE_REGULAR_KEY, + ) + + assert ANTHROPIC_MID_CONVERSATION_TOOL_CHANGES_BETA_HEADER not in headers.get("anthropic-beta", "").split(",") diff --git a/tests/unit/llms/base_llm/harness/__init__.py b/tests/unit/llms/base_llm/harness/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/bedrock/chat/test_converse_transformation.py b/tests/unit/llms/bedrock/chat/test_converse_transformation.py index 499096621c5..f6f98e3b9bd 100644 --- a/tests/unit/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/unit/llms/bedrock/chat/test_converse_transformation.py @@ -7769,6 +7769,17 @@ def test_mid_conversation_system_entry_without_text_is_dropped(empty_content): assert out_messages == [{"role": "user", "content": "hi"}, {"role": "user", "content": "done"}] +def test_system_entry_without_content_key_transforms_like_an_empty_one(): + config = AmazonConverseConfig() + leading_without_key = [{"role": "system"}, {"role": "user", "content": "hi"}] + leading_empty = [{"role": "system", "content": ""}, {"role": "user", "content": "hi"}] + assert config._transform_system_message(leading_without_key) == config._transform_system_message(leading_empty) + assert config._transform_system_message(leading_without_key) == ([{"role": "user", "content": "hi"}], []) + mid_without_key = [{"role": "user", "content": "hi"}, {"role": "system"}, {"role": "user", "content": "done"}] + mid_empty = [{"role": "user", "content": "hi"}, {"role": "system", "content": ""}, {"role": "user", "content": "done"}] + assert config._transform_system_message(mid_without_key) == config._transform_system_message(mid_empty) + + def _thinking_reply(text: str) -> dict: return { "role": "assistant", diff --git a/tests/unit/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py b/tests/unit/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py index d7f451dd6ee..a269d556262 100644 --- a/tests/unit/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py +++ b/tests/unit/llms/bedrock/messages/invoke_transformations/test_anthropic_claude3_transformation.py @@ -3527,3 +3527,54 @@ def test_bedrock_clear_thinking_preserves_display_updates() -> None: assert result.get("thinking") == {"type": "adaptive", "display": "updates"} assert ANTHROPIC_THINKING_DISPLAY_UPDATES_BETA_HEADER in result.get("anthropic_beta", []) + + +@pytest.mark.usefixtures("local_model_cost_map", "local_beta_headers_config") +@pytest.mark.parametrize("action", (None, "tool_addition", "tool_removal")) +@pytest.mark.parametrize("explicit_beta", (False, True)) +def test_bedrock_messages_tool_changes_beta(action: str | None, explicit_beta: bool) -> None: + from litellm.types.llms.anthropic import ANTHROPIC_MID_CONVERSATION_TOOL_CHANGES_BETA_HEADER + from litellm.types.router import GenericLiteLLMParams + + beta: Final = ANTHROPIC_MID_CONVERSATION_TOOL_CHANGES_BETA_HEADER + content: Final = ( + [{"type": action, "tool": {"type": "tool_reference", "name": "mcp__test__ping"}}] + if action + else "Answer briefly" + ) + messages: Final = [{"role": "user", "content": "Hello"}, {"role": "system", "content": content}] + result: Final = AmazonAnthropicClaudeMessagesConfig().transform_anthropic_messages_request( + model="global.anthropic.claude-fable-5-1", + messages=messages, + anthropic_messages_optional_request_params={"max_tokens": 512}, + litellm_params=GenericLiteLLMParams(), + headers={"anthropic-beta": beta} if explicit_beta else {}, + ) + + assert result.get("anthropic_beta", []).count(beta) == int(action is not None or explicit_beta) + assert result["messages"] == messages + + +@pytest.mark.usefixtures("local_model_cost_map", "local_beta_headers_config") +@pytest.mark.parametrize("explicit_beta", (False, True)) +def test_bedrock_removed_tool_change_does_not_add_beta(explicit_beta: bool) -> None: + from litellm.types.llms.anthropic import ANTHROPIC_MID_CONVERSATION_TOOL_CHANGES_BETA_HEADER + from litellm.types.router import GenericLiteLLMParams + + beta: Final = ANTHROPIC_MID_CONVERSATION_TOOL_CHANGES_BETA_HEADER + result: Final = AmazonAnthropicClaudeMessagesConfig().transform_anthropic_messages_request( + model="global.anthropic.claude-fable-5-1", + messages=[ + { + "role": "system", + "content": [{"type": "tool_addition", "tool": {"type": "tool_reference", "name": "ping"}}], + }, + {"role": "user", "content": "Reply with OK"}, + ], + anthropic_messages_optional_request_params={"max_tokens": 512}, + litellm_params=GenericLiteLLMParams(), + headers={"anthropic-beta": beta} if explicit_beta else {}, + ) + + assert result["messages"] == [{"role": "user", "content": "Reply with OK"}] + assert result.get("anthropic_beta", []).count(beta) == int(explicit_beta) diff --git a/tests/unit/llms/claude_code/__init__.py b/tests/unit/llms/claude_code/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/claude_code/harness/__init__.py b/tests/unit/llms/claude_code/harness/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/claude_code/harness/fixtures/__init__.py b/tests/unit/llms/claude_code/harness/fixtures/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/claude_code/harness/fixtures/api_error.jsonl b/tests/unit/llms/claude_code/harness/fixtures/api_error.jsonl new file mode 100644 index 00000000000..649b44345ce --- /dev/null +++ b/tests/unit/llms/claude_code/harness/fixtures/api_error.jsonl @@ -0,0 +1,3 @@ +{"type": "system", "subtype": "init", "cwd": "/workspace", "session_id": "53af83ee-c3e1-4b96-a70a-f15b6cb6c794", "tools": ["Task", "Bash", "CronCreate", "CronDelete", "CronList", "Edit", "EnterWorktree", "ExitWorktree", "ListAgents", "NotebookEdit", "Read", "ReportFindings", "ScheduleWakeup", "SendMessage", "Skill", "TaskStop", "WebFetch", "WebSearch", "Workflow", "Write"], "mcp_servers": [], "model": "does-not-exist-model-xyz", "permissionMode": "bypassPermissions", "slash_commands": ["deep-research", "dataviz", "update-config", "verify", "debug", "code-review", "simplify", "batch", "fewer-permission-prompts", "doctor", "loop", "claude-api", "workflow-authoring", "run", "run-skill-generator", "agents", "auto-mode-setup", "autocompact", "clear", "color", "compact", "config", "output-style", "context", "effort", "fast", "focus", "heapdump", "init", "mcp", "model", "__remote-workflow", "workflow-launch-exec", "reload-plugins", "reload-skills", "rename", "security-review", "usage", "insights", "recap", "goal", "list-agents", "team-onboarding"], "terminal_slash_commands": ["doctor", "color", "focus", "reload-plugins"], "apiKeySource": "none", "claude_code_version": "2.1.285", "output_style": "default", "agents": ["claude", "Explore", "general-purpose", "Plan", "statusline-setup"], "skills": ["deep-research", "dataviz", "update-config", "verify", "debug", "code-review", "simplify", "batch", "fewer-permission-prompts", "doctor", "loop", "claude-api", "workflow-authoring", "run", "run-skill-generator"], "plugins": [{"name": "cc-plugin-agents-md", "path": "builtin", "source": "cc-plugin-agents-md@builtin"}], "capabilities": ["interrupt_receipt_v1", "interrupt_cancel_queued_v1", "msg_lifecycle_v1", "mcp_read_resource_v1", "mcp_tool_ui_meta_v1"], "analytics_disabled": true, "product_feedback_disabled": true, "uuid": "a5308e12-af9c-41a0-97dc-91fc574b03dc", "fast_mode_state": "off", "fast_mode_disabled_reason": "sdk_opt_in_required", "per_turn_effort_active": false, "view_mode": "default"} +{"type": "assistant", "message": {"diagnostics": null, "id": "4a8ebe84-f673-472b-9f28-b38722e84b33", "container": null, "model": "", "role": "assistant", "stop_details": null, "stop_reason": "stop_sequence", "stop_sequence": "", "type": "message", "usage": {"output_tokens_details": null, "input_tokens": 0, "output_tokens": 0, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0, "server_tool_use": {"web_search_requests": 0, "web_fetch_requests": 0}, "service_tier": null, "cache_creation": {"ephemeral_1h_input_tokens": 0, "ephemeral_5m_input_tokens": 0}, "inference_geo": null, "iterations": null, "speed": null, "fallback_credit": null}, "content": [{"type": "text", "text": "API Error: 400 litellm.BadRequestError: You passed in model=does-not-exist-model-xyz. There are no healthy deployments for this model\n\nLiteLLM: model group 'does-not-exist-model-xyz' failed with the error above and no fallback model group was found for it, so the request was not retried on another model. Fallbacks are configured for: anthropic/*, anthropic/claude-opus-4-8, claude-mixed-router, anthropic/claude-fable-5, claude-opus-5, claude-sonnet-5, claude-fable-5, claude-fable-5-1, claude-haiku-4-5-20251001. Add a fallbacks entry for that model group (Router fallbacks or proxy router_settings.fallbacks) to retry on another model."}], "context_management": null}, "parent_tool_use_id": null, "session_id": "53af83ee-c3e1-4b96-a70a-f15b6cb6c794", "uuid": "f8c605c9-c1d6-43b0-89f8-aa64015d7895", "timestamp": "2026-09-30T17:00:14.185Z", "error": "unknown", "is_api_error_message": true} +{"duration_api_ms": 0, "stop_reason": "stop_sequence", "session_id": "53af83ee-c3e1-4b96-a70a-f15b6cb6c794", "total_cost_usd": 0, "usage": {"output_tokens_details": {"thinking_tokens": 0}, "input_tokens": 0, "cache_creation_input_tokens": 0, "cache_read_input_tokens": 0, "output_tokens": 0, "server_tool_use": {"web_search_requests": 0, "web_fetch_requests": 0}, "service_tier": "standard", "cache_creation": {"ephemeral_1h_input_tokens": 0, "ephemeral_5m_input_tokens": 0}, "inference_geo": "", "iterations": [], "speed": "standard", "fallback_credit": null}, "modelUsage": {}, "permission_denials": [], "terminal_reason": "api_error", "fast_mode_state": "off", "fast_mode_disabled_reason": "sdk_opt_in_required", "subagent_stats": {"spawned": 0, "requested": {"background": 0, "foreground": 0, "unset": 0}, "started_in_background": 0, "max_depth": 0, "spawned_by_subagents": 0, "completed": 0, "failed": 0, "killed": {"parent": 0, "user": 0, "system": 0}, "refused": {"depth_limit": 0, "concurrency_limit": 0, "budget": 0}, "by_type": {}}, "is_error": true, "num_turns": 1, "subtype": "success", "api_error_status": 400, "result": "API Error: 400 litellm.BadRequestError: You passed in model=does-not-exist-model-xyz. There are no healthy deployments for this model\n\nLiteLLM: model group 'does-not-exist-model-xyz' failed with the error above and no fallback model group was found for it, so the request was not retried on another model. Fallbacks are configured for: anthropic/*, anthropic/claude-opus-4-8, claude-mixed-router, anthropic/claude-fable-5, claude-opus-5, claude-sonnet-5, claude-fable-5, claude-fable-5-1, claude-haiku-4-5-20251001. Add a fallbacks entry for that model group (Router fallbacks or proxy router_settings.fallbacks) to retry on another model.", "type": "result", "duration_ms": 6781, "uuid": "c38e7f6d-8920-44f4-bab5-a599535509a0", "queued_turn_count": 0, "result_index": 0} diff --git a/tests/unit/llms/claude_code/harness/fixtures/max_turns.jsonl b/tests/unit/llms/claude_code/harness/fixtures/max_turns.jsonl new file mode 100644 index 00000000000..f40c7eae8a2 --- /dev/null +++ b/tests/unit/llms/claude_code/harness/fixtures/max_turns.jsonl @@ -0,0 +1,11 @@ +{"type": "system", "subtype": "init", "cwd": "/workspace", "session_id": "2270f364-6991-4263-8761-cd601a27cb8a", "tools": ["Task", "Bash", "CronCreate", "CronDelete", "CronList", "Edit", "EnterWorktree", "ExitWorktree", "ListAgents", "NotebookEdit", "Read", "ReportFindings", "ScheduleWakeup", "SendMessage", "Skill", "TaskCreate", "TaskGet", "TaskList", "TaskStop", "TaskUpdate", "WebFetch", "WebSearch", "Workflow", "Write"], "mcp_servers": [], "model": "claude-haiku-4-5-20251001", "permissionMode": "bypassPermissions", "slash_commands": ["deep-research", "dataviz", "update-config", "verify", "debug", "code-review", "simplify", "batch", "fewer-permission-prompts", "doctor", "loop", "claude-api", "workflow-authoring", "run", "run-skill-generator", "agents", "auto-mode-setup", "autocompact", "clear", "color", "compact", "config", "output-style", "context", "effort", "fast", "focus", "heapdump", "init", "mcp", "model", "__remote-workflow", "workflow-launch-exec", "reload-plugins", "reload-skills", "rename", "security-review", "usage", "insights", "recap", "goal", "list-agents", "team-onboarding"], "terminal_slash_commands": ["doctor", "color", "focus", "reload-plugins"], "apiKeySource": "none", "claude_code_version": "2.1.285", "output_style": "default", "agents": ["claude", "Explore", "general-purpose", "Plan", "statusline-setup"], "skills": ["deep-research", "dataviz", "update-config", "verify", "debug", "code-review", "simplify", "batch", "fewer-permission-prompts", "doctor", "loop", "claude-api", "workflow-authoring", "run", "run-skill-generator"], "plugins": [{"name": "cc-plugin-agents-md", "path": "builtin", "source": "cc-plugin-agents-md@builtin"}], "capabilities": ["interrupt_receipt_v1", "interrupt_cancel_queued_v1", "msg_lifecycle_v1", "mcp_read_resource_v1", "mcp_tool_ui_meta_v1"], "analytics_disabled": true, "product_feedback_disabled": true, "uuid": "a1474c31-391d-443a-8cb7-ff594f9bf72c", "fast_mode_state": "off", "fast_mode_disabled_reason": "sdk_opt_in_required", "per_turn_effort_active": false, "view_mode": "default"} +{"type": "system", "subtype": "thinking_tokens", "estimated_tokens": 7, "estimated_tokens_delta": 7, "session_id": "2270f364-6991-4263-8761-cd601a27cb8a", "uuid": "b263f250-510d-44b0-a0aa-79d269fe38cd"} +{"type": "system", "subtype": "thinking_tokens", "estimated_tokens": 11, "estimated_tokens_delta": 4, "session_id": "2270f364-6991-4263-8761-cd601a27cb8a", "uuid": "a84f3325-922e-4ca1-a0f2-acb8d5b60b91"} +{"type": "assistant", "message": {"model": "claude-haiku-4-5-20251001", "id": "msg_011CfZz9LnyzBujxEY6yNoWA", "type": "message", "role": "assistant", "content": [{"type": "thinking", "thinking": "The user wants me to create three files:\n1. a.txt with content \"1\"\n2. b.txt with content \"2\"\n3. c.txt with content \"3\"\n\nThey want one tool call for each, which means I should make three separate Write tool calls.\n\nSince these are independent operations (no dependencies between them), I should make all three calls in parallel in a single function_calls block.\n\nThe files should be created in the primary working directory: /workspace"}], "container": null, "stop_reason": null, "stop_sequence": null, "stop_details": null, "usage": {"input_tokens": 9, "cache_creation_input_tokens": 2730, "cache_read_input_tokens": 26136, "cache_creation": {"ephemeral_5m_input_tokens": 2730, "ephemeral_1h_input_tokens": 0}, "output_tokens": 6, "service_tier": "standard", "inference_geo": "not_available"}, "diagnostics": null, "context_management": null}, "parent_tool_use_id": null, "session_id": "2270f364-6991-4263-8761-cd601a27cb8a", "uuid": "1a577616-8652-49e8-8eba-195b5ea6d7d2", "timestamp": "2026-09-30T17:03:49.895Z"} +{"type": "assistant", "message": {"model": "claude-haiku-4-5-20251001", "id": "msg_011CfZz9LnyzBujxEY6yNoWA", "type": "message", "role": "assistant", "content": [{"type": "tool_use", "id": "toolu_01AAviEQjsijaCfLqz9HJexD", "name": "Write", "input": {"file_path": "/workspace/a.txt", "content": "1"}, "caller": {"type": "direct"}}], "container": null, "stop_reason": null, "stop_sequence": null, "stop_details": null, "usage": {"input_tokens": 9, "cache_creation_input_tokens": 2730, "cache_read_input_tokens": 26136, "cache_creation": {"ephemeral_5m_input_tokens": 2730, "ephemeral_1h_input_tokens": 0}, "output_tokens": 6, "service_tier": "standard", "inference_geo": "not_available"}, "diagnostics": null, "context_management": null}, "parent_tool_use_id": null, "session_id": "2270f364-6991-4263-8761-cd601a27cb8a", "uuid": "aec1ad61-fdb1-4e47-9238-68fad55c06b4", "timestamp": "2026-09-30T17:03:50.170Z"} +{"type": "user", "message": {"role": "user", "content": [{"tool_use_id": "toolu_01AAviEQjsijaCfLqz9HJexD", "type": "tool_result", "content": "File created successfully at: /workspace/a.txt (file state is current in your context — no need to Read it back)"}]}, "parent_tool_use_id": null, "session_id": "2270f364-6991-4263-8761-cd601a27cb8a", "uuid": "cfa2046c-604b-45ea-bcc4-a605207e77b8", "timestamp": "2026-09-30T17:03:50.177Z", "tool_use_result": {"type": "create", "filePath": "/workspace/a.txt", "content": "1", "structuredPatch": [], "originalFile": null, "userModified": false}} +{"type": "assistant", "message": {"model": "claude-haiku-4-5-20251001", "id": "msg_011CfZz9LnyzBujxEY6yNoWA", "type": "message", "role": "assistant", "content": [{"type": "tool_use", "id": "toolu_01E8ZcsQTKVnqC76NZzV3HCU", "name": "Write", "input": {"file_path": "/workspace/b.txt", "content": "2"}, "caller": {"type": "direct"}}], "container": null, "stop_reason": null, "stop_sequence": null, "stop_details": null, "usage": {"input_tokens": 9, "cache_creation_input_tokens": 2730, "cache_read_input_tokens": 26136, "cache_creation": {"ephemeral_5m_input_tokens": 2730, "ephemeral_1h_input_tokens": 0}, "output_tokens": 6, "service_tier": "standard", "inference_geo": "not_available"}, "diagnostics": null, "context_management": null}, "parent_tool_use_id": null, "session_id": "2270f364-6991-4263-8761-cd601a27cb8a", "uuid": "40e1dcdd-3c69-4318-9081-81c68ae13ae3", "timestamp": "2026-09-30T17:03:50.450Z"} +{"type": "user", "message": {"role": "user", "content": [{"tool_use_id": "toolu_01E8ZcsQTKVnqC76NZzV3HCU", "type": "tool_result", "content": "File created successfully at: /workspace/b.txt (file state is current in your context — no need to Read it back)"}]}, "parent_tool_use_id": null, "session_id": "2270f364-6991-4263-8761-cd601a27cb8a", "uuid": "a2b1dfa4-cde9-4bfb-8f2a-460a0eef2a18", "timestamp": "2026-09-30T17:03:50.456Z", "tool_use_result": {"type": "create", "filePath": "/workspace/b.txt", "content": "2", "structuredPatch": [], "originalFile": null, "userModified": false}} +{"type": "assistant", "message": {"model": "claude-haiku-4-5-20251001", "id": "msg_011CfZz9LnyzBujxEY6yNoWA", "type": "message", "role": "assistant", "content": [{"type": "tool_use", "id": "toolu_01N9f8PbyFhmyZ9wgiTqg3uG", "name": "Write", "input": {"file_path": "/workspace/c.txt", "content": "3"}, "caller": {"type": "direct"}}], "container": null, "stop_reason": null, "stop_sequence": null, "stop_details": null, "usage": {"input_tokens": 9, "cache_creation_input_tokens": 2730, "cache_read_input_tokens": 26136, "cache_creation": {"ephemeral_5m_input_tokens": 2730, "ephemeral_1h_input_tokens": 0}, "output_tokens": 6, "service_tier": "standard", "inference_geo": "not_available"}, "diagnostics": null, "context_management": null}, "parent_tool_use_id": null, "session_id": "2270f364-6991-4263-8761-cd601a27cb8a", "uuid": "8222d589-7c3e-4233-8bd2-4a73ffabeed2", "timestamp": "2026-09-30T17:03:50.725Z"} +{"type": "user", "message": {"role": "user", "content": [{"tool_use_id": "toolu_01N9f8PbyFhmyZ9wgiTqg3uG", "type": "tool_result", "content": "File created successfully at: /workspace/c.txt (file state is current in your context — no need to Read it back)"}]}, "parent_tool_use_id": null, "session_id": "2270f364-6991-4263-8761-cd601a27cb8a", "uuid": "98d09b2c-80b3-4b4f-8933-cfa7225ba8dc", "timestamp": "2026-09-30T17:03:50.737Z", "tool_use_result": {"type": "create", "filePath": "/workspace/c.txt", "content": "3", "structuredPatch": [], "originalFile": null, "userModified": false}} +{"duration_api_ms": 3347, "stop_reason": "tool_use", "session_id": "2270f364-6991-4263-8761-cd601a27cb8a", "total_cost_usd": 0.0077101, "usage": {"input_tokens": 9, "cache_creation_input_tokens": 2730, "cache_read_input_tokens": 26136, "output_tokens": 335, "output_tokens_details": {"thinking_tokens": 118}, "server_tool_use": {"web_search_requests": 0, "web_fetch_requests": 0}, "service_tier": "standard", "cache_creation": {"ephemeral_1h_input_tokens": 0, "ephemeral_5m_input_tokens": 2730}, "inference_geo": "not_available", "iterations": [], "speed": "standard", "fallback_credit": null}, "modelUsage": {"claude-haiku-4-5-20251001": {"inputTokens": 9, "outputTokens": 335, "cacheReadInputTokens": 26136, "cacheCreationInputTokens": 2730, "webSearchRequests": 0, "costUSD": 0.0077101, "contextWindow": 200000, "maxOutputTokens": 32000, "thinkingTokens": 118, "canonicalModel": "claude-haiku-4-5", "provider": "firstParty", "costBasis": "list"}}, "permission_denials": [], "terminal_reason": "max_turns", "fast_mode_state": "off", "fast_mode_disabled_reason": "sdk_opt_in_required", "subagent_stats": {"spawned": 0, "requested": {"background": 0, "foreground": 0, "unset": 0}, "started_in_background": 0, "max_depth": 0, "spawned_by_subagents": 0, "completed": 0, "failed": 0, "killed": {"parent": 0, "user": 0, "system": 0}, "refused": {"depth_limit": 0, "concurrency_limit": 0, "budget": 0}, "by_type": {}}, "is_error": true, "num_turns": 2, "subtype": "error_max_turns", "errors": ["Reached maximum number of turns (1)"], "type": "result", "duration_ms": 3381, "uuid": "db4afec5-5854-4439-89b4-f8d3539de3fd", "queued_turn_count": 0, "result_index": 0} diff --git a/tests/unit/llms/claude_code/harness/fixtures/resume_turn.jsonl b/tests/unit/llms/claude_code/harness/fixtures/resume_turn.jsonl new file mode 100644 index 00000000000..3b6994ff6a5 --- /dev/null +++ b/tests/unit/llms/claude_code/harness/fixtures/resume_turn.jsonl @@ -0,0 +1,6 @@ +{"type": "system", "subtype": "init", "cwd": "/workspace", "session_id": "5ef64ff1-d2af-4c38-a7ca-17b4a9d07d34", "tools": ["Task", "Bash", "CronCreate", "CronDelete", "CronList", "Edit", "EnterWorktree", "ExitWorktree", "ListAgents", "NotebookEdit", "Read", "ReportFindings", "ScheduleWakeup", "SendMessage", "Skill", "TaskCreate", "TaskGet", "TaskList", "TaskStop", "TaskUpdate", "WebFetch", "WebSearch", "Workflow", "Write"], "mcp_servers": [], "model": "claude-haiku-4-5-20251001", "permissionMode": "bypassPermissions", "slash_commands": ["deep-research", "dataviz", "update-config", "verify", "debug", "code-review", "simplify", "batch", "fewer-permission-prompts", "doctor", "loop", "claude-api", "workflow-authoring", "run", "run-skill-generator", "agents", "auto-mode-setup", "autocompact", "clear", "color", "compact", "config", "output-style", "context", "effort", "fast", "focus", "heapdump", "init", "mcp", "model", "__remote-workflow", "workflow-launch-exec", "reload-plugins", "reload-skills", "rename", "security-review", "usage", "insights", "recap", "goal", "list-agents", "team-onboarding"], "terminal_slash_commands": ["doctor", "color", "focus", "reload-plugins"], "apiKeySource": "none", "claude_code_version": "2.1.285", "output_style": "default", "agents": ["claude", "Explore", "general-purpose", "Plan", "statusline-setup"], "skills": ["deep-research", "dataviz", "update-config", "verify", "debug", "code-review", "simplify", "batch", "fewer-permission-prompts", "doctor", "loop", "claude-api", "workflow-authoring", "run", "run-skill-generator"], "plugins": [{"name": "cc-plugin-agents-md", "path": "builtin", "source": "cc-plugin-agents-md@builtin"}], "capabilities": ["interrupt_receipt_v1", "interrupt_cancel_queued_v1", "msg_lifecycle_v1", "mcp_read_resource_v1", "mcp_tool_ui_meta_v1"], "analytics_disabled": true, "product_feedback_disabled": true, "uuid": "753b841b-38d0-44e1-97a0-1cf3fc4875b0", "fast_mode_state": "off", "fast_mode_disabled_reason": "sdk_opt_in_required", "per_turn_effort_active": false, "view_mode": "default"} +{"type": "system", "subtype": "thinking_tokens", "estimated_tokens": 7, "estimated_tokens_delta": 7, "session_id": "5ef64ff1-d2af-4c38-a7ca-17b4a9d07d34", "uuid": "8c86350f-f3c9-44d8-98b1-b9bc64a5d89d"} +{"type": "system", "subtype": "thinking_tokens", "estimated_tokens": 10, "estimated_tokens_delta": 3, "session_id": "5ef64ff1-d2af-4c38-a7ca-17b4a9d07d34", "uuid": "533fe314-e617-4692-bc30-3985fd3cadc1"} +{"type": "assistant", "message": {"model": "claude-haiku-4-5-20251001", "id": "msg_011CfZyrEoMRLGNJffWtEf6r", "type": "message", "role": "assistant", "content": [{"type": "thinking", "thinking": "The user is asking what file I created, and they want a one-word answer. I created hello.txt, so the one-word answer would be \"hello.txt\" but that's technically two parts with a dot. If they want literally one word, it could be \"hello\" without the extension. But \"hello.txt\" is the filename. Let me think about what makes sense - they said \"one word\" so they probably want just \"hello\" or \"hello.txt\". The filename is \"hello.txt\" but if forced to one word, \"hello\" works.\n\nActually, looking at it more carefully, they asked \"What file\" - the file is \"hello.txt\". If they want one word, the answer should be \"hello.txt\" as a single unit (no spaces), or just \"hello\" if being very literal about \"one word\". I'll go with the filename as written."}], "container": null, "stop_reason": null, "stop_sequence": null, "stop_details": null, "usage": {"input_tokens": 9, "cache_creation_input_tokens": 104, "cache_read_input_tokens": 29266, "cache_creation": {"ephemeral_5m_input_tokens": 104, "ephemeral_1h_input_tokens": 0}, "output_tokens": 6, "service_tier": "standard", "inference_geo": "not_available"}, "diagnostics": null, "context_management": null}, "parent_tool_use_id": null, "session_id": "5ef64ff1-d2af-4c38-a7ca-17b4a9d07d34", "uuid": "dfb24f01-69d2-4a2c-8628-d78af662deb2", "timestamp": "2026-09-30T16:59:58.607Z"} +{"type": "assistant", "message": {"model": "claude-haiku-4-5-20251001", "id": "msg_011CfZyrEoMRLGNJffWtEf6r", "type": "message", "role": "assistant", "content": [{"type": "text", "text": "hello.txt"}], "container": null, "stop_reason": null, "stop_sequence": null, "stop_details": null, "usage": {"input_tokens": 9, "cache_creation_input_tokens": 104, "cache_read_input_tokens": 29266, "cache_creation": {"ephemeral_5m_input_tokens": 104, "ephemeral_1h_input_tokens": 0}, "output_tokens": 6, "service_tier": "standard", "inference_geo": "not_available"}, "diagnostics": null, "context_management": null}, "parent_tool_use_id": null, "session_id": "5ef64ff1-d2af-4c38-a7ca-17b4a9d07d34", "uuid": "09aae2d0-4408-4a2f-b6d1-238d15920f49", "timestamp": "2026-09-30T16:59:58.629Z"} +{"duration_api_ms": 27333, "stop_reason": "end_turn", "session_id": "5ef64ff1-d2af-4c38-a7ca-17b4a9d07d34", "total_cost_usd": 0.048190000000000004, "usage": {"input_tokens": 9, "cache_creation_input_tokens": 104, "cache_read_input_tokens": 29266, "output_tokens": 206, "output_tokens_details": {"thinking_tokens": 197}, "server_tool_use": {"web_search_requests": 0, "web_fetch_requests": 0}, "service_tier": "standard", "cache_creation": {"ephemeral_1h_input_tokens": 0, "ephemeral_5m_input_tokens": 104}, "inference_geo": "not_available", "iterations": [], "speed": "standard", "fallback_credit": null}, "modelUsage": {"claude-haiku-4-5-20251001": {"inputTokens": 34, "outputTokens": 544, "thinkingTokens": 354, "cacheReadInputTokens": 87235, "cacheCreationInputTokens": 29370, "webSearchRequests": 0, "costUSD": 0.048190000000000004, "contextWindow": 200000, "maxOutputTokens": 32000, "canonicalModel": "claude-haiku-4-5", "provider": "firstParty", "costBasis": "list"}}, "permission_denials": [], "terminal_reason": "completed", "fast_mode_state": "off", "fast_mode_disabled_reason": "sdk_opt_in_required", "subagent_stats": {"spawned": 0, "requested": {"background": 0, "foreground": 0, "unset": 0}, "started_in_background": 0, "max_depth": 0, "spawned_by_subagents": 0, "completed": 0, "failed": 0, "killed": {"parent": 0, "user": 0, "system": 0}, "refused": {"depth_limit": 0, "concurrency_limit": 0, "budget": 0}, "by_type": {}}, "is_error": false, "num_turns": 1, "subtype": "success", "api_error_status": null, "result": "hello.txt", "ttft_ms": 7545, "type": "result", "duration_ms": 7578, "uuid": "701ed1c8-48e5-4aad-b3fe-3982ad7aed00", "ttft_stream_ms": 5547, "time_to_request_ms": 23, "first_content_frame_ms": 5547, "queued_turn_count": 0, "result_index": 0} diff --git a/tests/unit/llms/claude_code/harness/fixtures/structured_output.jsonl b/tests/unit/llms/claude_code/harness/fixtures/structured_output.jsonl new file mode 100644 index 00000000000..29ba78c80e4 --- /dev/null +++ b/tests/unit/llms/claude_code/harness/fixtures/structured_output.jsonl @@ -0,0 +1,7 @@ +{"type": "system", "subtype": "init", "cwd": "/workspace", "session_id": "e0b4fb7e-b899-44ac-81fd-62841efa5380", "tools": ["Task", "Bash", "CronCreate", "CronDelete", "CronList", "Edit", "EnterWorktree", "ExitWorktree", "ListAgents", "NotebookEdit", "Read", "ReportFindings", "ScheduleWakeup", "SendMessage", "Skill", "StructuredOutput", "TaskCreate", "TaskGet", "TaskList", "TaskStop", "TaskUpdate", "WebFetch", "WebSearch", "Workflow", "Write"], "mcp_servers": [], "model": "claude-haiku-4-5-20251001", "permissionMode": "bypassPermissions", "slash_commands": ["deep-research", "dataviz", "update-config", "verify", "debug", "code-review", "simplify", "batch", "fewer-permission-prompts", "doctor", "loop", "claude-api", "workflow-authoring", "run", "run-skill-generator", "agents", "auto-mode-setup", "autocompact", "clear", "color", "compact", "config", "output-style", "context", "effort", "fast", "focus", "heapdump", "init", "mcp", "model", "__remote-workflow", "workflow-launch-exec", "reload-plugins", "reload-skills", "rename", "security-review", "usage", "insights", "recap", "goal", "list-agents", "team-onboarding"], "terminal_slash_commands": ["doctor", "color", "focus", "reload-plugins"], "apiKeySource": "none", "claude_code_version": "2.1.285", "output_style": "default", "agents": ["claude", "Explore", "general-purpose", "Plan", "statusline-setup"], "skills": ["deep-research", "dataviz", "update-config", "verify", "debug", "code-review", "simplify", "batch", "fewer-permission-prompts", "doctor", "loop", "claude-api", "workflow-authoring", "run", "run-skill-generator"], "plugins": [{"name": "cc-plugin-agents-md", "path": "builtin", "source": "cc-plugin-agents-md@builtin"}], "capabilities": ["interrupt_receipt_v1", "interrupt_cancel_queued_v1", "msg_lifecycle_v1", "mcp_read_resource_v1", "mcp_tool_ui_meta_v1"], "analytics_disabled": true, "product_feedback_disabled": true, "uuid": "a9fd450b-ce6f-4dc8-8d2c-1e2c1b654aa5", "fast_mode_state": "off", "fast_mode_disabled_reason": "sdk_opt_in_required", "per_turn_effort_active": false, "view_mode": "default"} +{"type": "system", "subtype": "thinking_tokens", "estimated_tokens": 5, "estimated_tokens_delta": 5, "session_id": "e0b4fb7e-b899-44ac-81fd-62841efa5380", "uuid": "c52043cf-0474-42ce-ae1b-1275fda7ee57"} +{"type": "system", "subtype": "thinking_tokens", "estimated_tokens": 9, "estimated_tokens_delta": 4, "session_id": "e0b4fb7e-b899-44ac-81fd-62841efa5380", "uuid": "dc90c9e0-4aef-4fe0-8dc9-25738677206c"} +{"type": "assistant", "message": {"model": "claude-haiku-4-5-20251001", "id": "msg_011CfZzH4cNfTXjBxTRhX5Vm", "type": "message", "role": "assistant", "content": [{"type": "thinking", "thinking": "The user is asking me to compute 2+3 and provide the answer and a short word. This is a straightforward math question.\n\n2 + 3 = 5\n\nThey want me to use the StructuredOutput tool to provide the answer in a structured format with an integer answer and a word.\n\nLooking at the tool definition:\n```\n{\n \"answer\": {\"type\": \"integer\"},\n \"word\": {\"type\": \"string\"}\n}\n```\n\nSo I need to return:\n- answer: 5\n- word: some short word\n\nI'll choose a simple word like \"sum\" or \"total\" or just a generic short word."}], "container": null, "stop_reason": null, "stop_sequence": null, "stop_details": null, "usage": {"input_tokens": 9, "cache_creation_input_tokens": 28946, "cache_read_input_tokens": 0, "cache_creation": {"ephemeral_5m_input_tokens": 28946, "ephemeral_1h_input_tokens": 0}, "output_tokens": 4, "service_tier": "standard", "inference_geo": "not_available"}, "diagnostics": null, "context_management": null}, "parent_tool_use_id": null, "session_id": "e0b4fb7e-b899-44ac-81fd-62841efa5380", "uuid": "27e9734a-846e-4ff2-b525-b249b1ff7b7f", "timestamp": "2026-09-30T17:05:35.124Z"} +{"type": "assistant", "message": {"model": "claude-haiku-4-5-20251001", "id": "msg_011CfZzH4cNfTXjBxTRhX5Vm", "type": "message", "role": "assistant", "content": [{"type": "tool_use", "id": "toolu_01NoK5Gm9S9NTEPEjjrjdfCf", "name": "StructuredOutput", "input": {"answer": 5, "word": "sum"}, "caller": {"type": "direct"}}], "container": null, "stop_reason": null, "stop_sequence": null, "stop_details": null, "usage": {"input_tokens": 9, "cache_creation_input_tokens": 28946, "cache_read_input_tokens": 0, "cache_creation": {"ephemeral_5m_input_tokens": 28946, "ephemeral_1h_input_tokens": 0}, "output_tokens": 4, "service_tier": "standard", "inference_geo": "not_available"}, "diagnostics": null, "context_management": null}, "parent_tool_use_id": null, "session_id": "e0b4fb7e-b899-44ac-81fd-62841efa5380", "uuid": "9e627359-3d66-4ce4-b893-330ebc0546f3", "timestamp": "2026-09-30T17:05:35.433Z"} +{"type": "user", "message": {"role": "user", "content": [{"tool_use_id": "toolu_01NoK5Gm9S9NTEPEjjrjdfCf", "type": "tool_result", "content": "Structured output provided successfully"}]}, "parent_tool_use_id": null, "session_id": "e0b4fb7e-b899-44ac-81fd-62841efa5380", "uuid": "939bc606-79ae-492f-a142-fbca6e400489", "timestamp": "2026-09-30T17:05:35.436Z", "tool_use_result": "Structured output provided successfully"} +{"duration_api_ms": 3171, "stop_reason": "tool_use", "session_id": "e0b4fb7e-b899-44ac-81fd-62841efa5380", "total_cost_usd": 0.0373165, "usage": {"input_tokens": 9, "cache_creation_input_tokens": 28946, "cache_read_input_tokens": 0, "output_tokens": 225, "output_tokens_details": {"thinking_tokens": 151}, "server_tool_use": {"web_search_requests": 0, "web_fetch_requests": 0}, "service_tier": "standard", "cache_creation": {"ephemeral_1h_input_tokens": 0, "ephemeral_5m_input_tokens": 28946}, "inference_geo": "not_available", "iterations": [], "speed": "standard", "fallback_credit": null}, "modelUsage": {"claude-haiku-4-5-20251001": {"inputTokens": 9, "outputTokens": 225, "cacheReadInputTokens": 0, "cacheCreationInputTokens": 28946, "webSearchRequests": 0, "costUSD": 0.0373165, "contextWindow": 200000, "maxOutputTokens": 32000, "thinkingTokens": 151, "canonicalModel": "claude-haiku-4-5", "provider": "firstParty", "costBasis": "list"}}, "permission_denials": [], "terminal_reason": "completed", "fast_mode_state": "off", "fast_mode_disabled_reason": "sdk_opt_in_required", "subagent_stats": {"spawned": 0, "requested": {"background": 0, "foreground": 0, "unset": 0}, "started_in_background": 0, "max_depth": 0, "spawned_by_subagents": 0, "completed": 0, "failed": 0, "killed": {"parent": 0, "user": 0, "system": 0}, "refused": {"depth_limit": 0, "concurrency_limit": 0, "budget": 0}, "by_type": {}}, "is_error": false, "num_turns": 2, "subtype": "success", "api_error_status": null, "result": "{\"answer\":5,\"word\":\"sum\"}", "structured_output": {"answer": 5, "word": "sum"}, "ttft_ms": 2886, "type": "result", "duration_ms": 3202, "uuid": "dca4570c-3d4e-4442-9544-5d47d9ce4268", "ttft_stream_ms": 1288, "time_to_request_ms": 31, "first_content_frame_ms": 1288, "queued_turn_count": 0, "result_index": 0} diff --git a/tests/unit/llms/claude_code/harness/fixtures/success_tools.jsonl b/tests/unit/llms/claude_code/harness/fixtures/success_tools.jsonl new file mode 100644 index 00000000000..b84d31447be --- /dev/null +++ b/tests/unit/llms/claude_code/harness/fixtures/success_tools.jsonl @@ -0,0 +1,12 @@ +{"type": "system", "subtype": "init", "cwd": "/workspace", "session_id": "5ef64ff1-d2af-4c38-a7ca-17b4a9d07d34", "tools": ["Task", "Bash", "CronCreate", "CronDelete", "CronList", "Edit", "EnterWorktree", "ExitWorktree", "ListAgents", "NotebookEdit", "Read", "ReportFindings", "ScheduleWakeup", "SendMessage", "Skill", "TaskCreate", "TaskGet", "TaskList", "TaskStop", "TaskUpdate", "WebFetch", "WebSearch", "Workflow", "Write"], "mcp_servers": [], "model": "claude-haiku-4-5-20251001", "permissionMode": "bypassPermissions", "slash_commands": ["deep-research", "dataviz", "update-config", "verify", "debug", "code-review", "simplify", "batch", "fewer-permission-prompts", "doctor", "loop", "claude-api", "workflow-authoring", "run", "run-skill-generator", "agents", "auto-mode-setup", "autocompact", "clear", "color", "compact", "config", "output-style", "context", "effort", "fast", "focus", "heapdump", "init", "mcp", "model", "__remote-workflow", "workflow-launch-exec", "reload-plugins", "reload-skills", "rename", "security-review", "usage", "insights", "recap", "goal", "list-agents", "team-onboarding"], "terminal_slash_commands": ["doctor", "color", "focus", "reload-plugins"], "apiKeySource": "none", "claude_code_version": "2.1.285", "output_style": "default", "agents": ["claude", "Explore", "general-purpose", "Plan", "statusline-setup"], "skills": ["deep-research", "dataviz", "update-config", "verify", "debug", "code-review", "simplify", "batch", "fewer-permission-prompts", "doctor", "loop", "claude-api", "workflow-authoring", "run", "run-skill-generator"], "plugins": [{"name": "cc-plugin-agents-md", "path": "builtin", "source": "cc-plugin-agents-md@builtin"}], "capabilities": ["interrupt_receipt_v1", "interrupt_cancel_queued_v1", "msg_lifecycle_v1", "mcp_read_resource_v1", "mcp_tool_ui_meta_v1"], "analytics_disabled": true, "product_feedback_disabled": true, "uuid": "71b745d7-51cf-49c7-9e0d-9fa08a069e09", "fast_mode_state": "off", "fast_mode_disabled_reason": "sdk_opt_in_required", "per_turn_effort_active": false, "view_mode": "default"} +{"type": "system", "subtype": "thinking_tokens", "estimated_tokens": 1, "estimated_tokens_delta": 1, "session_id": "5ef64ff1-d2af-4c38-a7ca-17b4a9d07d34", "uuid": "a2ed41b6-6db0-4667-aad5-f20ffcf19ce4"} +{"type": "system", "subtype": "thinking_tokens", "estimated_tokens": 4, "estimated_tokens_delta": 3, "session_id": "5ef64ff1-d2af-4c38-a7ca-17b4a9d07d34", "uuid": "f89d8c8f-2c16-4e1e-9d0f-16af3d94bb2c"} +{"type": "assistant", "message": {"model": "claude-haiku-4-5-20251001", "id": "msg_011CfZynHUDvD1eucc6XZnzm", "type": "message", "role": "assistant", "content": [{"type": "thinking", "thinking": "The user wants me to:\n1. Create a file hello.txt with the content \"hi\"\n2. Read it back\n3. Be brief\n\nI should use the Write tool to create the file, then the Read tool to read it back. Since these operations depend on each other (I need to write first), I should do them sequentially.\n\nLet me create the file first with Write, then read it back with Read."}], "container": null, "stop_reason": null, "stop_sequence": null, "stop_details": null, "usage": {"input_tokens": 9, "cache_creation_input_tokens": 28852, "cache_read_input_tokens": 0, "cache_creation": {"ephemeral_5m_input_tokens": 28852, "ephemeral_1h_input_tokens": 0}, "output_tokens": 1, "service_tier": "standard", "inference_geo": "not_available"}, "diagnostics": null, "context_management": null}, "parent_tool_use_id": null, "session_id": "5ef64ff1-d2af-4c38-a7ca-17b4a9d07d34", "uuid": "d2fab3ed-aee9-4122-99d3-ec5279016428", "timestamp": "2026-09-30T16:59:04.048Z"} +{"type": "assistant", "message": {"model": "claude-haiku-4-5-20251001", "id": "msg_011CfZynHUDvD1eucc6XZnzm", "type": "message", "role": "assistant", "content": [{"type": "tool_use", "id": "toolu_01DFhmKzT5x1NzxuestG2Hkj", "name": "Write", "input": {"file_path": "/workspace/hello.txt", "content": "hi"}, "caller": {"type": "direct"}}], "container": null, "stop_reason": null, "stop_sequence": null, "stop_details": null, "usage": {"input_tokens": 9, "cache_creation_input_tokens": 28852, "cache_read_input_tokens": 0, "cache_creation": {"ephemeral_5m_input_tokens": 28852, "ephemeral_1h_input_tokens": 0}, "output_tokens": 1, "service_tier": "standard", "inference_geo": "not_available"}, "diagnostics": null, "context_management": null}, "parent_tool_use_id": null, "session_id": "5ef64ff1-d2af-4c38-a7ca-17b4a9d07d34", "uuid": "8eca3aef-aa68-48c3-82cc-9073fb1fb653", "timestamp": "2026-09-30T16:59:04.408Z"} +{"type": "user", "message": {"role": "user", "content": [{"tool_use_id": "toolu_01DFhmKzT5x1NzxuestG2Hkj", "type": "tool_result", "content": "File created successfully at: /workspace/hello.txt (file state is current in your context — no need to Read it back)"}]}, "parent_tool_use_id": null, "session_id": "5ef64ff1-d2af-4c38-a7ca-17b4a9d07d34", "uuid": "5c07cb84-d183-4d41-abbf-aa81e6bccf24", "timestamp": "2026-09-30T16:59:04.421Z", "tool_use_result": {"type": "create", "filePath": "/workspace/hello.txt", "content": "hi", "structuredPatch": [], "originalFile": null, "userModified": false}} +{"type": "assistant", "message": {"model": "claude-haiku-4-5-20251001", "id": "msg_011CfZynchNSG2pNMhNYitM5", "type": "message", "role": "assistant", "content": [{"type": "thinking", "thinking": "The Write tool already created the file successfully. Now I need to read it back using the Read tool to confirm the contents."}], "container": null, "stop_reason": null, "stop_sequence": null, "stop_details": null, "usage": {"input_tokens": 8, "cache_creation_input_tokens": 265, "cache_read_input_tokens": 28852, "cache_creation": {"ephemeral_5m_input_tokens": 265, "ephemeral_1h_input_tokens": 0}, "output_tokens": 2, "service_tier": "standard", "inference_geo": "not_available"}, "diagnostics": null, "context_management": null}, "parent_tool_use_id": null, "session_id": "5ef64ff1-d2af-4c38-a7ca-17b4a9d07d34", "uuid": "4db547e1-d93c-47b7-bcc2-782e17f7e445", "timestamp": "2026-09-30T16:59:07.726Z"} +{"type": "assistant", "message": {"model": "claude-haiku-4-5-20251001", "id": "msg_011CfZynchNSG2pNMhNYitM5", "type": "message", "role": "assistant", "content": [{"type": "tool_use", "id": "toolu_017jG7ow2c8KCBUTq4Z9cDG8", "name": "Read", "input": {"file_path": "/workspace/hello.txt"}, "caller": {"type": "direct"}}], "container": null, "stop_reason": null, "stop_sequence": null, "stop_details": null, "usage": {"input_tokens": 8, "cache_creation_input_tokens": 265, "cache_read_input_tokens": 28852, "cache_creation": {"ephemeral_5m_input_tokens": 265, "ephemeral_1h_input_tokens": 0}, "output_tokens": 2, "service_tier": "standard", "inference_geo": "not_available"}, "diagnostics": null, "context_management": null}, "parent_tool_use_id": null, "session_id": "5ef64ff1-d2af-4c38-a7ca-17b4a9d07d34", "uuid": "fd6c4e76-87b1-46a9-80dd-98c1e0ee2244", "timestamp": "2026-09-30T16:59:07.881Z"} +{"type": "user", "message": {"role": "user", "content": [{"tool_use_id": "toolu_017jG7ow2c8KCBUTq4Z9cDG8", "type": "tool_result", "content": "1\thi"}]}, "parent_tool_use_id": null, "session_id": "5ef64ff1-d2af-4c38-a7ca-17b4a9d07d34", "uuid": "ac34772d-30f2-4fed-a0ec-5c9e50e2c864", "timestamp": "2026-09-30T16:59:07.895Z", "tool_use_result": {"type": "text", "file": {"filePath": "/workspace/hello.txt", "content": "hi", "numLines": 1, "startLine": 1, "totalLines": 1}}} +{"type": "assistant", "message": {"model": "claude-haiku-4-5-20251001", "id": "msg_011CfZyoDMCMiseBH9a2J4bZ", "type": "message", "role": "assistant", "content": [{"type": "thinking", "thinking": "Perfect! I've created the file and read it back. The file contains \"hi\" as expected. The task is complete."}], "container": null, "stop_reason": null, "stop_sequence": null, "stop_details": null, "usage": {"input_tokens": 8, "cache_creation_input_tokens": 149, "cache_read_input_tokens": 29117, "cache_creation": {"ephemeral_5m_input_tokens": 149, "ephemeral_1h_input_tokens": 0}, "output_tokens": 3, "service_tier": "standard", "inference_geo": "not_available"}, "diagnostics": null, "context_management": null}, "parent_tool_use_id": null, "session_id": "5ef64ff1-d2af-4c38-a7ca-17b4a9d07d34", "uuid": "ec814481-9a98-419f-8a2f-96e57be0f3d6", "timestamp": "2026-09-30T16:59:16.170Z"} +{"type": "assistant", "message": {"model": "claude-haiku-4-5-20251001", "id": "msg_011CfZyoDMCMiseBH9a2J4bZ", "type": "message", "role": "assistant", "content": [{"type": "text", "text": "Done. Created `hello.txt` with content \"hi\" and confirmed it reads back correctly."}], "container": null, "stop_reason": null, "stop_sequence": null, "stop_details": null, "usage": {"input_tokens": 8, "cache_creation_input_tokens": 149, "cache_read_input_tokens": 29117, "cache_creation": {"ephemeral_5m_input_tokens": 149, "ephemeral_1h_input_tokens": 0}, "output_tokens": 3, "service_tier": "standard", "inference_geo": "not_available"}, "diagnostics": null, "context_management": null}, "parent_tool_use_id": null, "session_id": "5ef64ff1-d2af-4c38-a7ca-17b4a9d07d34", "uuid": "563dcf5b-5a7d-4d10-9dbb-9924c2f0b09f", "timestamp": "2026-09-30T16:59:16.434Z"} +{"duration_api_ms": 19780, "stop_reason": "end_turn", "session_id": "5ef64ff1-d2af-4c38-a7ca-17b4a9d07d34", "total_cost_usd": 0.044094400000000006, "usage": {"input_tokens": 25, "cache_creation_input_tokens": 29266, "cache_read_input_tokens": 57969, "output_tokens": 338, "output_tokens_details": {"thinking_tokens": 157}, "server_tool_use": {"web_search_requests": 0, "web_fetch_requests": 0}, "service_tier": "standard", "cache_creation": {"ephemeral_1h_input_tokens": 0, "ephemeral_5m_input_tokens": 29266}, "inference_geo": "not_available", "iterations": [], "speed": "standard", "fallback_credit": null}, "modelUsage": {"claude-haiku-4-5-20251001": {"inputTokens": 25, "outputTokens": 338, "cacheReadInputTokens": 57969, "cacheCreationInputTokens": 29266, "webSearchRequests": 0, "costUSD": 0.044094400000000006, "contextWindow": 200000, "maxOutputTokens": 32000, "thinkingTokens": 157, "canonicalModel": "claude-haiku-4-5", "provider": "firstParty", "costBasis": "list"}}, "permission_denials": [], "terminal_reason": "completed", "fast_mode_state": "off", "fast_mode_disabled_reason": "sdk_opt_in_required", "subagent_stats": {"spawned": 0, "requested": {"background": 0, "foreground": 0, "unset": 0}, "started_in_background": 0, "max_depth": 0, "spawned_by_subagents": 0, "completed": 0, "failed": 0, "killed": {"parent": 0, "user": 0, "system": 0}, "refused": {"depth_limit": 0, "concurrency_limit": 0, "budget": 0}, "by_type": {}}, "is_error": false, "num_turns": 3, "subtype": "success", "api_error_status": null, "result": "Done. Created `hello.txt` with content \"hi\" and confirmed it reads back correctly.", "ttft_ms": 7217, "type": "result", "duration_ms": 19835, "uuid": "702bf2fe-16b0-41e8-afb9-efe37f17abe3", "ttft_stream_ms": 6540, "time_to_request_ms": 27, "first_content_frame_ms": 6541, "queued_turn_count": 0, "result_index": 0} diff --git a/tests/unit/llms/claude_code/harness/test_transformation.py b/tests/unit/llms/claude_code/harness/test_transformation.py new file mode 100644 index 00000000000..6348b4f6a3e --- /dev/null +++ b/tests/unit/llms/claude_code/harness/test_transformation.py @@ -0,0 +1,708 @@ +"""Unit tests for the Claude Code harness config. No network, no real CLI. + +Fixtures under fixtures/ are sanitized stream-json recorded from Claude Code +2.1.285 through a LiteLLM gateway. +""" + +from __future__ import annotations + +import asyncio +import json +import os +from collections.abc import Mapping +from pathlib import Path +from typing import Any + +import pytest +from pydantic import BaseModel + +from litellm.harness.context import SessionContext +from litellm.harness.errors import ( + HarnessError, + HarnessInstallFailed, + OptionsMismatch, +) +from litellm.harness.handlers.cli_handler import CLIHarnessHandler +from litellm.harness.options import ClaudeCodeOptions, CodexOptions +from litellm.harness.sandbox.base import CompletedRun +from litellm.harness.types import ( + Compaction, + Harness, + Reasoning, + Text, + ToolCall, + ToolResult, +) +from litellm.llms.base_llm.harness.transformation import ( + HarnessTurnError, + HarnessTurnRequest, +) +from litellm.llms.base_llm.harness.utils import ( + decode_json_line, + last_json_object, + native_tool_names, +) +from litellm.llms.claude_code.harness.transformation import ( + MANAGED_CONFIG_KEYS, + MANAGED_ENV_KEYS, + NORMALIZED_TO_NATIVE, + PERMISSION_MODES, + ClaudeCodeHarnessConfig, + ClaudeCodeStreamState, + build_system_prompt, + stringify_tool_output, + turn_error_message, +) + +FIXTURES = Path(__file__).parent / "fixtures" +SESSION_ID = "5ef64ff1-d2af-4c38-a7ca-17b4a9d07d34" +TOKEN = "per-session-token-abc" +PORT = 53211 +PRIV = "/priv" + + +def fixture_lines(name: str) -> list[str]: + return (FIXTURES / name).read_text().splitlines() + + +def parse_line(line: str, state: ClaudeCodeStreamState) -> list[Any]: + decoded = decode_json_line(line) + if decoded is None: + return [] + return ClaudeCodeHarnessConfig().transform_stream_line(decoded, state) + + +def parse_fixture(name: str) -> tuple[list[Any], ClaudeCodeStreamState]: + state = ClaudeCodeHarnessConfig().create_stream_state() + events: list[Any] = [] + for line in fixture_lines(name): + events.extend(parse_line(line, state)) + return events, state + + +class FakeEndpoint: + port = PORT + token = TOKEN + + +class FakeProcess: + def __init__(self, stdout: bytes, stderr: bytes, exit_code: int) -> None: + self.stdin_data = bytearray() + self.stdin_closed = False + self.killed = False + self._exit_code = exit_code + self.stdout = asyncio.StreamReader() + self.stdout.feed_data(stdout) + self.stdout.feed_eof() + self.stderr = asyncio.StreamReader() + self.stderr.feed_data(stderr) + self.stderr.feed_eof() + self.stdin = FakeStdin(self) + + async def wait(self) -> int: + return self._exit_code + + async def kill(self) -> None: + self.killed = True + + +class FakeStdin: + def __init__(self, proc: FakeProcess) -> None: + self._proc = proc + + def write(self, data: bytes) -> None: + self._proc.stdin_data.extend(data) + + async def drain(self) -> None: + return None + + def close(self) -> None: + self._proc.stdin_closed = True + + +class FakeSandbox: + def __init__( + self, + workdir: str, + outputs: list[tuple[str, bytes, int]], + binary: str | None = "/usr/bin/claude", + tempdir: str | None = None, + ) -> None: + self.workdir = workdir + self.binary = binary + self.outputs = list(outputs) + self.calls: list[dict[str, Any]] = [] + self.runs: list[list[str]] = [] + self.procs: list[FakeProcess] = [] + self.written: dict[str, bytes] = {} + self._tempdir = tempdir or os.path.join(workdir, "_cfg") + + async def exec( + self, + cmd: list[str], + *, + env: Mapping[str, str] | None = None, + cwd: str | None = None, + ) -> FakeProcess: + self.calls.append({"cmd": cmd, "env": dict(env or {}), "cwd": cwd}) + fixture, stderr, code = self.outputs.pop(0) + stdout = (FIXTURES / fixture).read_bytes() if fixture else b"" + proc = FakeProcess(stdout, stderr, code) + self.procs.append(proc) + return proc + + async def run(self, cmd: list[str], **kwargs: Any) -> CompletedRun: + self.runs.append(cmd) + return CompletedRun("", "", 0) + + async def read(self, path: str) -> bytes: + return self.written[path] + + async def write(self, path: str, data: bytes) -> None: + self.written[path] = data + + def host_url(self, port: int) -> str: + return f"http://host.docker.internal:{port}" + + async def which(self, binary: str) -> str | None: + return self.binary + + async def tempdir(self) -> str: + return self._tempdir + + async def snapshot(self) -> dict[str, str]: + return {} + + async def close(self) -> None: + return None + + +class Answer(BaseModel): + answer: int + word: str + + +def make_ctx(sandbox: FakeSandbox, **overrides: Any) -> SessionContext: + values: dict[str, Any] = { + "harness": Harness.CLAUDE_CODE, + "sandbox": sandbox, + "session_id": "hs_1", + "model": "claude-haiku-4-5-20251001", + "endpoint": FakeEndpoint(), + **overrides, + } + return SessionContext(**values) + + +def pure_ctx(tmp_path: Path, **overrides: Any) -> SessionContext: + return make_ctx(FakeSandbox(str(tmp_path), []), **overrides) + + +def make_handler() -> CLIHarnessHandler: + return CLIHarnessHandler(ClaudeCodeHarnessConfig()) + + +def request_for( + ctx: SessionContext, native_session_id: str | None = None, prompt: str = "hi" +) -> HarnessTurnRequest: + cfg = ClaudeCodeHarnessConfig() + setup = cfg.transform_session_setup(ctx, PRIV) + return cfg.transform_turn_request(ctx, setup, PRIV, prompt, native_session_id) + + +async def run_turn(handler: CLIHarnessHandler, ctx: SessionContext, prompt: str): + return [event async for event in handler.turn(ctx, prompt)] + + +# --------------------------------------------------------------------------- +# Parsing +# --------------------------------------------------------------------------- + + +def test_parse_success_fixture_events(): + events, state = parse_fixture("success_tools.jsonl") + kinds = [type(e).__name__ for e in events] + assert kinds == [ + "Reasoning", + "ToolCall", + "ToolResult", + "Reasoning", + "ToolCall", + "ToolResult", + "Reasoning", + "Text", + ] + write_call, read_call = events[1], events[4] + assert write_call == ToolCall( + id="toolu_01DFhmKzT5x1NzxuestG2Hkj", + name="write", + native_name="Write", + input={"file_path": "/workspace/hello.txt", "content": "hi"}, + builtin=True, + ) + assert read_call.name == "read" and read_call.native_name == "Read" + assert events[2].id == write_call.id and events[2].is_error is False + assert events[5].output == "1\thi" + assert state.session_id == SESSION_ID + assert ClaudeCodeHarnessConfig().get_native_session_id(state) == SESSION_ID + assert state.result_seen and not state.is_error + assert state.final_text.startswith("Done. Created `hello.txt`") + + +def test_parse_api_error_fixture_skips_synthetic_text(): + events, state = parse_fixture("api_error.jsonl") + assert events == [] + assert state.is_error + assert "no healthy deployments" in (state.result_text or "") + + +def test_parse_max_turns_fixture(): + events, state = parse_fixture("max_turns.jsonl") + assert [e.native_name for e in events if isinstance(e, ToolCall)] == [ + "Write", + "Write", + "Write", + ] + assert state.is_error and state.result_text is None + assert state.errors == ["Reached maximum number of turns (1)"] + + +def test_parse_structured_output_fixture(): + _, state = parse_fixture("structured_output.jsonl") + assert state.structured_output == {"answer": 5, "word": "sum"} + + +def test_parse_compaction_and_garbage(): + state = ClaudeCodeStreamState() + line = json.dumps( + { + "type": "system", + "subtype": "compact_boundary", + "compact_metadata": {"trigger": "auto", "pre_tokens": 1234}, + } + ) + assert parse_line(line, state) == [ + Compaction(tokens_before=1234, tokens_after=None) + ] + assert parse_line("not json", state) == [] + assert parse_line("", state) == [] + assert parse_line("[1,2]", state) == [] + cfg = ClaudeCodeHarnessConfig() + assert cfg.transform_stream_line({"type": "unknown"}, state) == [] + + +def test_parse_skips_subagent_messages_and_maps_errors(): + cfg = ClaudeCodeHarnessConfig() + state = ClaudeCodeStreamState() + sub = { + "type": "assistant", + "parent_tool_use_id": "toolu_parent", + "message": {"content": [{"type": "text", "text": "inner"}]}, + } + assert cfg.transform_stream_line(sub, state) == [] + err = { + "type": "user", + "message": { + "content": [ + { + "type": "tool_result", + "tool_use_id": "t1", + "is_error": True, + "content": [{"type": "text", "text": "boom"}], + } + ] + }, + } + assert cfg.transform_stream_line(err, state) == [ + ToolResult(id="t1", output="boom", is_error=True) + ] + + +def test_parse_thinking_and_mcp_tools(): + state = ClaudeCodeStreamState() + msg = { + "type": "assistant", + "message": { + "content": [ + {"type": "thinking", "thinking": "hmm"}, + {"type": "tool_use", "id": "t", "name": "mcp__x__y", "input": {}}, + {"type": "tool_use", "id": "u", "name": "MultiEdit", "input": {}}, + ] + }, + } + events = ClaudeCodeHarnessConfig().transform_stream_line(msg, state) + assert events[0] == Reasoning(delta="hmm") + assert events[1].name == "mcp__x__y" and events[1].builtin is False + assert events[2].name == "edit" + + +def test_stringify_tool_output_variants(): + assert stringify_tool_output(None) == "" + assert stringify_tool_output("x") == "x" + assert stringify_tool_output([{"type": "text", "text": "a"}, "b"]) == "a\nb" + assert stringify_tool_output({"k": 1}) == '{"k": 1}' + + +def test_extract_last_json_object(): + text = 'first {"a": 1} then {not json} and finally {"b": {"c": 2}}' + assert json.loads(last_json_object(text) or "") == {"b": {"c": 2}} + assert last_json_object("no json here") is None + + +# --------------------------------------------------------------------------- +# Session setup / turn request (argv + env) +# --------------------------------------------------------------------------- + + +def test_native_disallowed_tools_mapping(): + natives = native_tool_names(["edit", "bash", "Task", "edit"], NORMALIZED_TO_NATIVE) + assert natives == ["Edit", "MultiEdit", "Bash", "Task"] + + +@pytest.mark.parametrize( + "permissions,native", + [ + ("read-only", "plan"), + ("edit", "acceptEdits"), + ("full", "bypassPermissions"), + ], +) +def test_turn_request_permission_modes(tmp_path, permissions, native): + assert PERMISSION_MODES[permissions] == native + argv = list(request_for(pure_ctx(tmp_path, permissions=permissions)).argv) + assert argv[argv.index("--permission-mode") + 1] == native + assert "--resume" not in argv + assert argv[argv.index("--setting-sources") + 1] == "user" + + +def test_session_setup_and_turn_request_env_and_command(tmp_path): + ctx = pure_ctx( + tmp_path, + instructions="Be terse.", + disable_tools=["bash", "web_search"], + max_turns=7, + options=ClaudeCodeOptions(config={"cleanupPeriodDays": 1}, env={"X": "1"}), + ) + cfg = ClaudeCodeHarnessConfig() + setup = cfg.transform_session_setup(ctx, PRIV) + assert setup.persisted_dirs == [("projects", "claude_code/projects")] + assert setup.skills_dir == "skills" + request = cfg.transform_turn_request(ctx, setup, PRIV, "do the thing", None) + env, cmd = request.env, list(request.argv) + assert request.stdin == "do the thing" + assert env["ANTHROPIC_AUTH_TOKEN"] == TOKEN + assert env["ANTHROPIC_API_KEY"] == "" + assert env["ANTHROPIC_BASE_URL"] == f"http://host.docker.internal:{PORT}" + assert env["ANTHROPIC_MODEL"] == "claude-haiku-4-5-20251001" + assert env["ANTHROPIC_SMALL_FAST_MODEL"] == "claude-haiku-4-5-20251001" + assert env["CLAUDE_CONFIG_DIR"] == PRIV + assert env["DISABLE_TELEMETRY"] == "1" + assert env["CLAUDE_CODE_DISABLE_NONESSENTIAL_TRAFFIC"] == "1" + assert env["X"] == "1" + assert not any(TOKEN in a for a in cmd) + assert cmd[:7] == [ + "claude", + "-p", + "--output-format", + "stream-json", + "--verbose", + "--input-format", + "text", + ] + assert cmd[cmd.index("--model") + 1] == "claude-haiku-4-5-20251001" + assert cmd[cmd.index("--permission-mode") + 1] == "bypassPermissions" + assert cmd[cmd.index("--setting-sources") + 1] == "user" + assert cmd[cmd.index("--append-system-prompt") + 1] == "Be terse." + assert cmd[cmd.index("--max-turns") + 1] == "7" + assert json.loads(cmd[cmd.index("--settings") + 1]) == {"cleanupPeriodDays": 1} + assert cmd[cmd.index("--disallowedTools") + 1] == "Bash,WebSearch" + assert "--resume" not in cmd + + +def test_background_model_is_the_session_model(tmp_path): + env = ( + ClaudeCodeHarnessConfig().transform_session_setup(pure_ctx(tmp_path), PRIV).env + ) + assert env["ANTHROPIC_SMALL_FAST_MODEL"] == "claude-haiku-4-5-20251001" + + +def test_resume_argv(tmp_path): + argv = list(request_for(pure_ctx(tmp_path), "prior-session").argv) + assert argv[argv.index("--resume") + 1] == "prior-session" + + +def test_missing_endpoint_raises(tmp_path): + with pytest.raises(HarnessError, match="endpoint"): + ClaudeCodeHarnessConfig().transform_session_setup( + pure_ctx(tmp_path, endpoint=None), PRIV + ) + + +@pytest.mark.parametrize("key", sorted(MANAGED_ENV_KEYS)) +def test_options_env_cannot_override_managed_keys(tmp_path, key): + ctx = pure_ctx(tmp_path, options=ClaudeCodeOptions(env={key: "sk-real"})) + with pytest.raises(OptionsMismatch, match=key): + ClaudeCodeHarnessConfig().validate_environment(ctx) + + +def test_wrong_options_type_rejected(tmp_path): + with pytest.raises(OptionsMismatch): + ClaudeCodeHarnessConfig().validate_environment( + pure_ctx(tmp_path, options=CodexOptions()) + ) + + +def test_structured_output_system_prompt(tmp_path): + argv = list( + request_for(pure_ctx(tmp_path, output=Answer, instructions="Base.")).argv + ) + prompt = argv[argv.index("--append-system-prompt") + 1] + assert prompt.startswith("Base.\n\n") + assert json.dumps(Answer.model_json_schema()) in prompt + assert build_system_prompt(None, None) is None + + +# --------------------------------------------------------------------------- +# Turn response +# --------------------------------------------------------------------------- + + +def test_turn_response_api_error_includes_stderr(tmp_path): + _, state = parse_fixture("api_error.jsonl") + with pytest.raises(HarnessTurnError) as info: + ClaudeCodeHarnessConfig().transform_turn_response( + pure_ctx(tmp_path), state, 1, ["[claude-code:unrecognized_model] bad"] + ) + assert "no healthy deployments" in str(info.value) + assert "unrecognized_model" in str(info.value) + + +def test_turn_response_max_turns(tmp_path): + _, state = parse_fixture("max_turns.jsonl") + with pytest.raises(HarnessTurnError, match="maximum number of turns"): + ClaudeCodeHarnessConfig().transform_turn_response( + pure_ctx(tmp_path), state, 1, [] + ) + + +def test_turn_error_message_no_result(): + message = turn_error_message(ClaudeCodeStreamState(), 139, ["segfault", ""]) + assert message is not None + assert "code 139: no result event" in message and "segfault" in message + _, ok = parse_fixture("success_tools.jsonl") + assert turn_error_message(ok, 0, []) is None + + +def test_turn_response_structured_output_and_fallback(tmp_path): + cfg = ClaudeCodeHarnessConfig() + ctx = pure_ctx(tmp_path, output=Answer) + _, state = parse_fixture("structured_output.jsonl") + response = cfg.transform_turn_response(ctx, state, 0, []) + assert json.loads(response.output_json or "") == {"answer": 5, "word": "sum"} + + _, plain = parse_fixture("resume_turn.jsonl") + response = cfg.transform_turn_response(ctx, plain, 0, []) + assert response.final_text == "hello.txt" + assert response.output_json is None # "hello.txt" holds no JSON object + + text_json = ClaudeCodeStreamState( + result_seen=True, result_text='answer: {"answer": 1, "word": "x"}' + ) + response = cfg.transform_turn_response(ctx, text_json, 0, []) + assert json.loads(response.output_json or "") == {"answer": 1, "word": "x"} + + no_output = cfg.transform_turn_response(pure_ctx(tmp_path), state, 0, []) + assert no_output.output_json is None + + +# --------------------------------------------------------------------------- +# Through CLIHarnessHandler (start + turn) +# --------------------------------------------------------------------------- + + +async def test_start_and_turn_env_and_command(tmp_path): + sandbox = FakeSandbox(str(tmp_path), [("success_tools.jsonl", b"", 0)]) + ctx = make_ctx(sandbox, options=ClaudeCodeOptions(env={"X": "1"})) + handler = make_handler() + await handler.start(ctx) + assert len(sandbox.runs) == 1 + assert sandbox.runs[0][:2] == ["sh", "-c"] + assert sandbox.runs[0][-2:] == [ + f"{tmp_path / '_cfg'}/projects", + "claude_code/projects", + ] + events = await run_turn(handler, ctx, "do the thing") + + call = sandbox.calls[0] + env, cmd = call["env"], call["cmd"] + assert env["ANTHROPIC_AUTH_TOKEN"] == TOKEN + assert env["CLAUDE_CONFIG_DIR"] == str(tmp_path / "_cfg") + assert env["X"] == "1" + assert not any(TOKEN in a for a in cmd) + assert cmd[cmd.index("--setting-sources") + 1] == "user" + + proc = sandbox.procs[0] + assert bytes(proc.stdin_data) == b"do the thing" and proc.stdin_closed + assert any(isinstance(e, Text) for e in events) + assert ctx.final_text.startswith("Done.") + assert handler.native_session_id() == SESSION_ID + + +async def test_second_turn_resumes_session(tmp_path): + sandbox = FakeSandbox( + str(tmp_path), + [("success_tools.jsonl", b"", 0), ("resume_turn.jsonl", b"", 0)], + ) + ctx = make_ctx(sandbox) + handler = make_handler() + await handler.start(ctx) + await run_turn(handler, ctx, "one") + await run_turn(handler, ctx, "two") + cmd = sandbox.calls[1]["cmd"] + assert cmd[cmd.index("--resume") + 1] == SESSION_ID + assert ctx.final_text == "hello.txt" + + +async def test_resume_sets_native_session_id(tmp_path): + sandbox = FakeSandbox(str(tmp_path), [("resume_turn.jsonl", b"", 0)]) + ctx = make_ctx(sandbox) + handler = make_handler() + await handler.start(ctx) + await handler.resume(ctx, "prior-session") + assert handler.native_session_id() == "prior-session" + await run_turn(handler, ctx, "again") + cmd = sandbox.calls[0]["cmd"] + assert cmd[cmd.index("--resume") + 1] == "prior-session" + + +async def test_missing_binary_raises_install_failed(tmp_path): + sandbox = FakeSandbox(str(tmp_path), [], binary=None) + with pytest.raises(HarnessInstallFailed, match="claude"): + await make_handler().start(make_ctx(sandbox)) + + +async def test_start_missing_endpoint_raises(tmp_path): + sandbox = FakeSandbox(str(tmp_path), []) + with pytest.raises(HarnessError, match="endpoint"): + await make_handler().start(make_ctx(sandbox, endpoint=None)) + + +async def test_start_rejects_managed_env(tmp_path): + sandbox = FakeSandbox(str(tmp_path), []) + options = ClaudeCodeOptions(env={"ANTHROPIC_API_KEY": "sk-real"}) + with pytest.raises(OptionsMismatch, match="ANTHROPIC_API_KEY"): + await make_handler().start(make_ctx(sandbox, options=options)) + assert sandbox.runs == [] and sandbox.written == {} + + +async def test_api_error_raises_turn_error_with_stderr(tmp_path): + stderr = b"[claude-code:unrecognized_model] bad model\n" + sandbox = FakeSandbox(str(tmp_path), [("api_error.jsonl", stderr, 1)]) + ctx = make_ctx(sandbox) + handler = make_handler() + await handler.start(ctx) + with pytest.raises(HarnessTurnError) as info: + await run_turn(handler, ctx, "hi") + assert "no healthy deployments" in str(info.value) + assert "unrecognized_model" in str(info.value) + + +async def test_max_turns_raises_turn_error(tmp_path): + sandbox = FakeSandbox(str(tmp_path), [("max_turns.jsonl", b"", 1)]) + ctx = make_ctx(sandbox) + handler = make_handler() + await handler.start(ctx) + with pytest.raises(HarnessTurnError, match="maximum number of turns"): + await run_turn(handler, ctx, "hi") + + +async def test_nonzero_exit_without_result_raises(tmp_path): + sandbox = FakeSandbox(str(tmp_path), [("", b"segfault\n", 139)]) + ctx = make_ctx(sandbox) + handler = make_handler() + await handler.start(ctx) + with pytest.raises(HarnessTurnError, match=r"code 139.*no result event") as info: + await run_turn(handler, ctx, "hi") + assert "segfault" in str(info.value) + + +async def test_structured_output_prompt_and_json(tmp_path): + sandbox = FakeSandbox(str(tmp_path), [("structured_output.jsonl", b"", 0)]) + ctx = make_ctx(sandbox, output=Answer, instructions="Base.") + handler = make_handler() + await handler.start(ctx) + await run_turn(handler, ctx, "2+3?") + cmd = sandbox.calls[0]["cmd"] + prompt = cmd[cmd.index("--append-system-prompt") + 1] + assert prompt.startswith("Base.\n\n") + assert json.loads(ctx.output_json or "") == {"answer": 5, "word": "sum"} + + +async def test_structured_output_falls_back_to_final_text(tmp_path): + sandbox = FakeSandbox(str(tmp_path), [("resume_turn.jsonl", b"", 0)]) + ctx = make_ctx(sandbox, output=Answer) + handler = make_handler() + await handler.start(ctx) + await run_turn(handler, ctx, "hi") + assert ctx.output_json is None # "hello.txt" holds no JSON object + + +async def test_skills_copied_into_private_config(tmp_path): + skill = tmp_path / "skills_src" / "demo" + (skill / "scripts").mkdir(parents=True) + (skill / "SKILL.md").write_text("---\nname: demo\n---\nSay DEMO.\n") + (skill / "scripts" / "run.sh").write_text("echo hi\n") + workdir = tmp_path / "work" + workdir.mkdir() + sandbox = FakeSandbox(str(workdir), [], tempdir="/cfg") + await make_handler().start(make_ctx(sandbox, skills=[str(skill)])) + assert sandbox.written == { + "/cfg/skills/demo/SKILL.md": b"---\nname: demo\n---\nSay DEMO.\n", + "/cfg/skills/demo/scripts/run.sh": b"echo hi\n", + } + + +async def test_skill_without_manifest_rejected(tmp_path): + skill = tmp_path / "bad" + skill.mkdir() + sandbox = FakeSandbox(str(tmp_path), []) + with pytest.raises(ValueError, match=r"SKILL\.md"): + await make_handler().start(make_ctx(sandbox, skills=[str(skill)])) + + +async def test_stop_kills_live_process(tmp_path): + sandbox = FakeSandbox(str(tmp_path), [("success_tools.jsonl", b"", 0)]) + ctx = make_ctx(sandbox) + handler = make_handler() + await handler.start(ctx) + stream = handler.turn(ctx, "hi") + await stream.__anext__() + proc = sandbox.procs[0] + await handler.stop(ctx) + assert proc.killed + await stream.aclose() + await handler.stop(ctx) # safe twice + + +def test_capabilities_match_spec(): + cfg = ClaudeCodeHarnessConfig() + caps = cfg.capabilities + assert cfg.harness is Harness.CLAUDE_CODE + assert cfg.options_type is ClaudeCodeOptions + assert cfg.get_binary() == "claude" + assert "@anthropic-ai/claude-code" in cfg.get_install_hint() + assert caps.structured_output and caps.tool_filtering and caps.skills + assert caps.resume + assert not (caps.tool_approval or caps.custom_tools or caps.history) + assert caps.permission_modes == frozenset({"read-only", "edit", "full"}) + + +@pytest.mark.parametrize("key", sorted(MANAGED_CONFIG_KEYS)) +def test_options_config_cannot_set_managed_keys(tmp_path, key): + ctx = pure_ctx(tmp_path, options=ClaudeCodeOptions(config={key: "x"})) + with pytest.raises(OptionsMismatch, match=key): + ClaudeCodeHarnessConfig().validate_environment(ctx) + + +def test_no_settings_flag_without_config(tmp_path): + assert "--settings" not in list(request_for(pure_ctx(tmp_path), None).argv) diff --git a/tests/unit/llms/codex/__init__.py b/tests/unit/llms/codex/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/codex/harness/__init__.py b/tests/unit/llms/codex/harness/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/codex/harness/fixtures/__init__.py b/tests/unit/llms/codex/harness/fixtures/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/codex/harness/fixtures/reasoning.jsonl b/tests/unit/llms/codex/harness/fixtures/reasoning.jsonl new file mode 100644 index 00000000000..e38a212baed --- /dev/null +++ b/tests/unit/llms/codex/harness/fixtures/reasoning.jsonl @@ -0,0 +1,5 @@ +{"type":"thread.started","thread_id":"01a0f347-4945-7f40-ae19-d1a2724ddee1"} +{"type":"turn.started"} +{"type":"item.completed","item":{"id":"item_0","type":"reasoning","text":"**Calculating multiplication**\n\nAlright, I need to respond with just the number. I multiply 17 and 23 to get 391. Let me check that: 20 times 23 equals 460, and if I subtract 3 times 23, which is 69, from 460, I get 391. So, yes, 391 is correct! I’ll provide the final answer as \"391\" only, without any extra text."}} +{"type":"item.completed","item":{"id":"item_1","type":"agent_message","text":"391"}} +{"type":"turn.completed","usage":{"input_tokens":9098,"cached_input_tokens":0,"output_tokens":68,"reasoning_output_tokens":0}} diff --git a/tests/unit/llms/codex/harness/fixtures/structured_output.jsonl b/tests/unit/llms/codex/harness/fixtures/structured_output.jsonl new file mode 100644 index 00000000000..0ecb79d9f1e --- /dev/null +++ b/tests/unit/llms/codex/harness/fixtures/structured_output.jsonl @@ -0,0 +1,6 @@ +{"type":"thread.started","thread_id":"01a0f344-34bf-7b82-9025-bc6db70f867e"} +{"type":"turn.started"} +{"type":"item.started","item":{"id":"item_0","type":"command_execution","command":"/bin/zsh -lc \"rg --files -g 'AGENTS.md' -g 'hello.txt' . && printf '\\\\n---\\\\n' && cat hello.txt\"","aggregated_output":"","exit_code":null,"status":"in_progress"}} +{"type":"item.completed","item":{"id":"item_0","type":"command_execution","command":"/bin/zsh -lc \"rg --files -g 'AGENTS.md' -g 'hello.txt' . && printf '\\\\n---\\\\n' && cat hello.txt\"","aggregated_output":"./hello.txt\n\n---\nhello world\n","exit_code":0,"status":"completed"}} +{"type":"item.completed","item":{"id":"item_1","type":"agent_message","text":"{\"file\":\"hello.txt\",\"content\":\"Name: `hello.txt`\\nContent: `hello world`\"}"}} +{"type":"turn.completed","usage":{"input_tokens":24271,"cached_input_tokens":3758,"output_tokens":96,"reasoning_output_tokens":0}} diff --git a/tests/unit/llms/codex/harness/fixtures/turn1_bash.jsonl b/tests/unit/llms/codex/harness/fixtures/turn1_bash.jsonl new file mode 100644 index 00000000000..110ae70999d --- /dev/null +++ b/tests/unit/llms/codex/harness/fixtures/turn1_bash.jsonl @@ -0,0 +1,7 @@ +{"type":"thread.started","thread_id":"01a0f341-fe37-7072-93b3-055358e8147f"} +{"type":"turn.started"} +{"type":"item.completed","item":{"id":"item_0","type":"agent_message","text":"I’ll create the file, then print it back to confirm."}} +{"type":"item.started","item":{"id":"item_1","type":"command_execution","command":"/bin/zsh -lc \"printf 'hi\n' > hello.txt && cat hello.txt\"","aggregated_output":"","exit_code":null,"status":"in_progress"}} +{"type":"item.completed","item":{"id":"item_1","type":"command_execution","command":"/bin/zsh -lc \"printf 'hi\n' > hello.txt && cat hello.txt\"","aggregated_output":"hi\n","exit_code":0,"status":"completed"}} +{"type":"item.completed","item":{"id":"item_2","type":"agent_message","text":"Done — `hello.txt` now contains `hi`, and `cat hello.txt` prints:\n\n```text\nhi\n```"}} +{"type":"turn.completed","usage":{"input_tokens":24372,"cached_input_tokens":13998,"output_tokens":103,"reasoning_output_tokens":0}} diff --git a/tests/unit/llms/codex/harness/fixtures/turn2_resume_apply_patch.jsonl b/tests/unit/llms/codex/harness/fixtures/turn2_resume_apply_patch.jsonl new file mode 100644 index 00000000000..4cb5d924352 --- /dev/null +++ b/tests/unit/llms/codex/harness/fixtures/turn2_resume_apply_patch.jsonl @@ -0,0 +1,7 @@ +{"type":"thread.started","thread_id":"01a0f341-fe37-7072-93b3-055358e8147f"} +{"type":"turn.started"} +{"type":"item.completed","item":{"id":"item_0","type":"agent_message","text":"I’ll patch `hello.txt` directly, then I’ll reply exactly as requested."}} +{"type":"item.started","item":{"id":"item_1","type":"command_execution","command":"/bin/zsh -lc \"apply_patch '*** Begin Patch\n*** Delete File: hello.txt\n*** Add File: hello.txt\n+hello world\n*** End Patch'\"","aggregated_output":"","exit_code":null,"status":"in_progress"}} +{"type":"item.completed","item":{"id":"item_1","type":"command_execution","command":"/bin/zsh -lc \"apply_patch '*** Begin Patch\n*** Delete File: hello.txt\n*** Add File: hello.txt\n+hello world\n*** End Patch'\"","aggregated_output":"Success. Updated the following files:\nA hello.txt\nD hello.txt\n","exit_code":0,"status":"completed"}} +{"type":"item.completed","item":{"id":"item_2","type":"agent_message","text":"done"}} +{"type":"turn.completed","usage":{"input_tokens":49569,"cached_input_tokens":38236,"output_tokens":204,"reasoning_output_tokens":0}} diff --git a/tests/unit/llms/codex/harness/fixtures/turn_failed.jsonl b/tests/unit/llms/codex/harness/fixtures/turn_failed.jsonl new file mode 100644 index 00000000000..7e4404f6ff5 --- /dev/null +++ b/tests/unit/llms/codex/harness/fixtures/turn_failed.jsonl @@ -0,0 +1,4 @@ +{"type": "thread.started", "thread_id": "01a0f344-5b83-7d40-a964-7d56c9e4ec9a"} +{"type": "turn.started"} +{"type": "error", "message": "{\"error\":{\"message\":\"litellm.BadRequestError: You passed in model=no-such-model-xyz. There are no healthy deployments for this model\",\"type\":\"invalid_request_error\",\"param\":null,\"code\":\"400\"}}"} +{"type": "turn.failed", "error": {"message": "{\"error\":{\"message\":\"litellm.BadRequestError: You passed in model=no-such-model-xyz. There are no healthy deployments for this model\",\"type\":\"invalid_request_error\",\"param\":null,\"code\":\"400\"}}"}} diff --git a/tests/unit/llms/codex/harness/test_transformation.py b/tests/unit/llms/codex/harness/test_transformation.py new file mode 100644 index 00000000000..6371f74fba6 --- /dev/null +++ b/tests/unit/llms/codex/harness/test_transformation.py @@ -0,0 +1,670 @@ +"""Unit tests for the Codex harness config. No network, no real CLI. + +Fixtures under fixtures/ are sanitized `codex exec --json` output recorded from +codex-cli through a LiteLLM gateway. +""" + +import asyncio +import json +from dataclasses import dataclass, field +from pathlib import Path +from typing import Optional + +import pytest +from pydantic import BaseModel + +from litellm.harness.context import SessionContext +from litellm.harness.errors import HarnessError, HarnessInstallFailed, OptionsMismatch +from litellm.harness.handlers.cli_handler import CLIHarnessHandler +from litellm.harness.options import ClaudeCodeOptions, CodexOptions +from litellm.harness.sandbox.base import CompletedRun +from litellm.harness.sandbox.docker import DockerSandbox +from litellm.harness.types import Harness, Reasoning, Text, ToolCall, ToolResult +from litellm.llms.base_llm.harness.transformation import ( + HarnessTurnError, + HarnessTurnRequest, +) +from litellm.llms.base_llm.harness.utils import strict_json_schema +from litellm.llms.codex.harness.transformation import ( + CODEX_SCHEMA_FILENAME, + CODEX_TOKEN_ENV, + MANAGED_CONFIG_KEYS, + CodexHarnessConfig, + CodexStreamState, + config_overrides, + toml_key, + toml_value, +) + +FIXTURES = Path(__file__).parent / "fixtures" +TOKEN = "tok-secret-123" +HOME = "/tmp/codex-home" +THREAD_ID = "01a0f341-fe37-7072-93b3-055358e8147f" + + +def load_fixture(name: str) -> list[dict]: + return [ + json.loads(line) for line in (FIXTURES / name).read_text().splitlines() if line + ] + + +def parse_event(obj: dict, state: CodexStreamState) -> list: + return CodexHarnessConfig().transform_stream_line(obj, state) + + +def parse_all(name: str, state: Optional[CodexStreamState] = None): + state = state or CodexHarnessConfig().create_stream_state() + events = [] + for obj in load_fixture(name): + events.extend(parse_event(obj, state)) + return events, state + + +# --------------------------------------------------------------------------- fakes + + +class FakeProcess: + def __init__(self, stdout: bytes, stderr: bytes = b"", exit_code: int = 0): + self.stdin = FakeStdin() + self.stdout = asyncio.StreamReader() + self.stdout.feed_data(stdout) + self.stdout.feed_eof() + self.stderr = asyncio.StreamReader() + self.stderr.feed_data(stderr) + self.stderr.feed_eof() + self._exit_code = exit_code + self.killed = False + + async def wait(self) -> int: + return self._exit_code + + async def kill(self) -> None: + self.killed = True + + +class FakeStdin: + def __init__(self): + self.data = b"" + self.closed = False + + def write(self, data: bytes) -> None: + self.data += data + + async def drain(self) -> None: + return None + + def close(self) -> None: + self.closed = True + + +@dataclass +class FakeSandbox: + workdir: str = "/work" + has_codex: bool = True + outputs: list = field(default_factory=list) + files: dict = field(default_factory=dict) + execs: list = field(default_factory=list) + runs: list = field(default_factory=list) + processes: list = field(default_factory=list) + + async def exec(self, cmd, *, env=None, cwd=None): + self.execs.append({"cmd": cmd, "env": dict(env or {}), "cwd": cwd}) + proc = self.outputs.pop(0) + self.processes.append(proc) + return proc + + async def run(self, cmd, *, env=None, cwd=None, timeout=None): + self.runs.append(cmd) + return CompletedRun("", "", 0) + + async def read(self, path): + return self.files[path] + + async def write(self, path, data): + self.files[path] = data + + def host_url(self, port): + return f"http://127.0.0.1:{port}" + + async def which(self, binary): + return f"/usr/bin/{binary}" if self.has_codex else None + + async def tempdir(self): + return HOME + + async def snapshot(self): + return {} + + async def close(self): + return None + + +@dataclass +class FakeEndpoint: + port: int = 4555 + token: str = TOKEN + + +class Answer(BaseModel): + file: str + content: str + + +class Nested(BaseModel): + answer: Answer + tags: list[str] = [] + note: Optional[str] = None + + +def make_ctx(sandbox, **kwargs) -> SessionContext: + return SessionContext( + harness=Harness.CODEX, + sandbox=sandbox, + session_id="s1", + model=kwargs.pop("model", "gpt-5.4"), + endpoint=kwargs.pop("endpoint", FakeEndpoint()), + **kwargs, + ) + + +def request_for( + ctx: SessionContext, native_session_id: Optional[str] = None, prompt: str = "hi" +) -> HarnessTurnRequest: + cfg = CodexHarnessConfig() + setup = cfg.transform_session_setup(ctx, HOME) + return cfg.transform_turn_request(ctx, setup, HOME, prompt, native_session_id) + + +def argv_for(ctx: SessionContext, native_session_id: Optional[str] = None) -> list: + return list(request_for(ctx, native_session_id).argv) + + +def fixture_proc(name: str, **kwargs) -> FakeProcess: + return FakeProcess((FIXTURES / name).read_bytes(), **kwargs) + + +def config_values(argv: list[str]) -> list[str]: + return [argv[i + 1] for i, a in enumerate(argv) if a == "-c"] + + +def make_handler() -> CLIHarnessHandler: + return CLIHarnessHandler(CodexHarnessConfig()) + + +async def collect(handler, ctx, prompt): + return [e async for e in handler.turn(ctx, prompt)] + + +# --------------------------------------------------------------------------- parsing + + +def test_parse_bash_turn(): + events, state = parse_all("turn1_bash.jsonl") + assert state.thread_id == THREAD_ID + assert CodexHarnessConfig().get_native_session_id(state) == THREAD_ID + assert [type(e) for e in events] == [Text, ToolCall, ToolResult, Text] + call, result = events[1], events[2] + assert call.name == "bash" and call.native_name == "command_execution" + assert call.builtin is True + assert "hello.txt" in call.input["command"] + assert result.id == call.id == "item_1" + assert result.output == "hi\n" and result.is_error is False + assert state.final_text.startswith("Done") + assert not state.failed + + +def test_parse_reasoning(): + events, state = parse_all("reasoning.jsonl") + assert isinstance(events[0], Reasoning) and "391" in events[0].delta + assert events[1] == Text(delta="391") + assert state.final_text == "391" + + +def test_parse_turn_failed(): + events, state = parse_all("turn_failed.jsonl") + assert events == [] + assert state.failed + assert "no healthy deployments" in state.error + + +def test_parse_file_change_and_mcp_and_web_search(): + state = CodexStreamState() + change = { + "id": "i1", + "type": "file_change", + "changes": [{"path": "a.txt", "kind": "add"}], + "status": "completed", + } + events = parse_event({"type": "item.completed", "item": change}, state) + assert events[0] == ToolCall( + id="i1", + name="edit", + native_name="apply_patch", + input={"changes": [{"path": "a.txt", "kind": "add"}]}, + ) + assert events[1] == ToolResult(id="i1", output="add a.txt", is_error=False) + + mcp = { + "id": "i2", + "type": "mcp_tool_call", + "server": "docs", + "tool": "search", + "arguments": {"q": "x"}, + "status": "in_progress", + } + started = parse_event({"type": "item.started", "item": mcp}, state) + assert started == [ + ToolCall( + id="i2", + name="docs.search", + native_name="search", + input={"q": "x"}, + builtin=False, + ) + ] + done = {**mcp, "status": "failed", "error": {"message": "boom"}} + assert parse_event({"type": "item.completed", "item": done}, state) == [ + ToolResult(id="i2", output="boom", is_error=True) + ] + + web = {"id": "i3", "type": "web_search", "query": "litellm"} + events = parse_event({"type": "item.completed", "item": web}, state) + assert events[0].name == "web_search" and events[0].input == {"query": "litellm"} + + +def test_parse_failed_command_is_error_and_unknown_events_ignored(): + state = CodexStreamState() + item = { + "id": "c", + "type": "command_execution", + "command": "false", + "aggregated_output": "", + "exit_code": 1, + "status": "failed", + } + events = parse_event({"type": "item.completed", "item": item}, state) + assert events[1].is_error is True + usage = {"type": "turn.completed", "usage": {"input_tokens": 5}} + assert parse_event(usage, state) == [] + todo = {"type": "item.completed", "item": {"type": "todo_list"}} + assert parse_event(todo, state) == [] + assert parse_event({"type": "item.completed", "item": "nope"}, state) == [] + + +def test_parse_error_event_then_turn_failed(): + state = CodexStreamState() + assert parse_event({"type": "error", "message": "reconnecting"}, state) == [] + assert state.error == "reconnecting" and not state.failed + assert parse_event({"type": "turn.failed", "error": {"message": "x"}}, state) == [] + assert state.failed and state.error == "x" + + +# --------------------------------------------------------------------------- helpers + + +def test_strict_json_schema_recursive(): + schema = strict_json_schema(Nested.model_json_schema()) + assert schema["additionalProperties"] is False + assert schema["required"] == ["answer", "tags", "note"] + assert schema["properties"]["answer"] == {"$ref": "#/$defs/Answer"} + assert "default" not in schema["properties"]["tags"] + answer = schema["$defs"]["Answer"] + assert answer["additionalProperties"] is False + assert answer["required"] == ["file", "content"] + + +def test_config_overrides_rejects_managed_keys(): + for key in ( + "model_provider", + "model_providers.x.base_url", + "approval_policy", + "sandbox_mode", + "mcp_servers.a", + ): + with pytest.raises(OptionsMismatch): + config_overrides({key: "x"}) + for key in sorted(MANAGED_CONFIG_KEYS): + with pytest.raises(OptionsMismatch, match="managed by LiteLLM"): + config_overrides({key: "x"}) + for bad in ("", "a=b"): + with pytest.raises(OptionsMismatch, match="Invalid"): + config_overrides({bad: "x"}) + assert config_overrides( + { + "sandbox_workspace_write.network_access": True, + "notice": {"a b": 1}, + "x": ["y"], + } + ) == [ + "sandbox_workspace_write.network_access=true", + 'notice={"a b" = 1}', + 'x=["y"]', + ] + + +def test_toml_value_and_key(): + assert toml_value('say "hi"') == '"say \\"hi\\""' + assert toml_value(False) == "false" + assert toml_value(1.5) == "1.5" + assert toml_value(("a", 2)) == '["a", 2]' + assert toml_key("plain_key-1") == "plain_key-1" + assert toml_key("a b") == '"a b"' + with pytest.raises(OptionsMismatch): + toml_value(object()) + + +# --------------------------------------------------------------------------- session setup / turn request + + +def test_session_setup_env_and_schema(): + ctx = make_ctx(FakeSandbox(), output=Answer, options=CodexOptions(env={"X": "1"})) + setup = CodexHarnessConfig().transform_session_setup(ctx, HOME) + assert setup.env == {"X": "1", CODEX_TOKEN_ENV: TOKEN, "CODEX_HOME": HOME} + assert setup.persisted_dirs == [("sessions", "codex/sessions")] + assert setup.skills_dir == "skills" + schema = json.loads(setup.files[CODEX_SCHEMA_FILENAME]) + assert schema["additionalProperties"] is False + assert schema["required"] == ["file", "content"] + no_schema = CodexHarnessConfig().transform_session_setup( + make_ctx(FakeSandbox()), HOME + ) + assert no_schema.files == {} + + +def test_missing_endpoint_raises(): + ctx = make_ctx(FakeSandbox(), endpoint=None) + with pytest.raises(HarnessError, match="endpoint"): + CodexHarnessConfig().transform_session_setup(ctx, HOME) + + +def test_validate_environment_rejects_managed_config_and_wrong_options(): + cfg = CodexHarnessConfig() + with pytest.raises(OptionsMismatch): + cfg.validate_environment( + make_ctx( + FakeSandbox(), options=CodexOptions(config={"model_provider": "openai"}) + ) + ) + with pytest.raises(OptionsMismatch): + cfg.validate_environment(make_ctx(FakeSandbox(), options=ClaudeCodeOptions())) + + +def test_first_turn_argv_env(): + ctx = make_ctx( + FakeSandbox(), + instructions="Be terse.", + options=CodexOptions( + reasoning_effort="low", + config={"sandbox_workspace_write.network_access": True}, + ), + ) + request = request_for(ctx, prompt="create hello.txt") + argv, env = list(request.argv), request.env + assert request.stdin == "create hello.txt" + assert request.cwd == "/work" + assert argv[:4] == ["codex", "exec", "--json", "--skip-git-repo-check"] + assert argv[-1] == "-" and argv[argv.index("-C") + 1] == "/work" + assert argv[argv.index("-m") + 1] == "gpt-5.4" + assert argv[argv.index("--sandbox") + 1] == "workspace-write" + cfg = config_values(argv) + assert "model_provider=litellm" in cfg + assert 'model_providers.litellm.base_url="http://127.0.0.1:4555/v1"' in cfg + assert "model_providers.litellm.env_key=LITELLM_HARNESS_TOKEN" in cfg + assert "model_providers.litellm.wire_api=responses" in cfg + assert "approval_policy=never" in cfg + assert "model_reasoning_effort=low" in cfg + assert "model_reasoning_summary=auto" in cfg + assert "web_search=disabled" in cfg + assert 'developer_instructions="Be terse."' in cfg + assert "sandbox_workspace_write.network_access=true" in cfg + assert "--output-schema" not in argv + assert not any(TOKEN in a for a in argv) + assert env["LITELLM_HARNESS_TOKEN"] == TOKEN + assert env["CODEX_HOME"] == HOME + + +def test_resume_argv(): + argv = argv_for(make_ctx(FakeSandbox()), THREAD_ID) + assert argv[:4] == ["codex", "exec", "resume", THREAD_ID] + assert "--sandbox" not in argv and "-C" not in argv + assert 'sandbox_mode="workspace-write"' in config_values(argv) + assert argv[-1] == "-" + + +def test_permission_modes(): + ro = make_ctx(FakeSandbox(), permissions="read-only") + argv = argv_for(ro) + assert argv[argv.index("--sandbox") + 1] == "read-only" + assert 'sandbox_mode="read-only"' in config_values(argv_for(ro, "t")) + + # The container is the boundary: DockerSandbox opts out of codex's own sandbox. + assert DockerSandbox.is_container is True + container = FakeSandbox(workdir="/workspace") + container.is_container = True + argv = argv_for(make_ctx(container, permissions="full")) + assert "--dangerously-bypass-approvals-and-sandbox" in argv + assert "--sandbox" not in argv + assert argv[argv.index("-C") + 1] == "/workspace" + resumed = argv_for(make_ctx(container, permissions="full"), "t") + assert "--dangerously-bypass-approvals-and-sandbox" in resumed + assert not any(v.startswith("sandbox_mode=") for v in config_values(resumed)) + + # read-only wins even inside a container + argv = argv_for(make_ctx(container, permissions="read-only")) + assert argv[argv.index("--sandbox") + 1] == "read-only" + assert "--dangerously-bypass-approvals-and-sandbox" not in argv + + web = make_ctx(FakeSandbox(), options=CodexOptions(web_search=True)) + assert "web_search=live" in config_values(argv_for(web)) + + +def test_structured_output_argv(): + argv = argv_for(make_ctx(FakeSandbox(), output=Answer)) + assert argv[argv.index("--output-schema") + 1] == f"{HOME}/{CODEX_SCHEMA_FILENAME}" + + +def test_no_model_omits_flag(): + assert "-m" not in argv_for(make_ctx(FakeSandbox(), model=None)) + + +# --------------------------------------------------------------------------- turn response + + +def test_turn_response_failed_raises(): + _, state = parse_all("turn_failed.jsonl") + with pytest.raises(HarnessTurnError, match="no healthy deployments"): + CodexHarnessConfig().transform_turn_response( + make_ctx(FakeSandbox()), state, 1, [] + ) + + +def test_turn_response_nonzero_exit_uses_stderr_tail(): + with pytest.raises(HarnessTurnError, match=r"code 1: Error loading config\.toml"): + CodexHarnessConfig().transform_turn_response( + make_ctx(FakeSandbox()), + CodexStreamState(), + 1, + ["Error loading config.toml: bad", ""], + ) + with pytest.raises(HarnessTurnError, match="code 2: no output"): + CodexHarnessConfig().transform_turn_response( + make_ctx(FakeSandbox()), CodexStreamState(), 2, [] + ) + + +def test_turn_response_output_json_only_with_output(): + state = CodexStreamState(final_text='{"file": "a", "content": "b"}') + cfg = CodexHarnessConfig() + with_out = cfg.transform_turn_response( + make_ctx(FakeSandbox(), output=Answer), state, 0, [] + ) + assert with_out.output_json == state.final_text + without = cfg.transform_turn_response(make_ctx(FakeSandbox()), state, 0, []) + assert without.output_json is None and without.final_text == state.final_text + + +# --------------------------------------------------------------------------- handler + + +async def test_start_missing_binary(): + ctx = make_ctx(FakeSandbox(has_codex=False)) + with pytest.raises(HarnessInstallFailed, match="codex"): + await make_handler().start(ctx) + + +async def test_start_rejects_managed_config_and_wrong_options(): + with pytest.raises(OptionsMismatch): + await make_handler().start( + make_ctx( + FakeSandbox(), options=CodexOptions(config={"model_provider": "openai"}) + ) + ) + with pytest.raises(OptionsMismatch): + await make_handler().start(make_ctx(FakeSandbox(), options=ClaudeCodeOptions())) + + +async def test_start_writes_skills_and_schema(tmp_path): + skill = tmp_path / "my-skill" + (skill / "scripts").mkdir(parents=True) + (skill / "SKILL.md").write_text("---\nname: my-skill\n---\nbody") + (skill / "scripts" / "run.sh").write_text("echo hi") + sbx = FakeSandbox() + await make_handler().start(make_ctx(sbx, skills=[str(skill)], output=Answer)) + assert sbx.files[f"{HOME}/skills/my-skill/SKILL.md"].startswith(b"---") + assert sbx.files[f"{HOME}/skills/my-skill/scripts/run.sh"] == b"echo hi" + schema = json.loads(sbx.files[f"{HOME}/output_schema.json"]) + assert schema["additionalProperties"] is False + assert schema["required"] == ["file", "content"] + assert sbx.runs[0][:2] == ["sh", "-c"] + assert sbx.runs[0][-2:] == [f"{HOME}/sessions", "codex/sessions"] + + +async def test_first_turn_then_resume_argv_env(): + sbx = FakeSandbox( + outputs=[ + fixture_proc("turn1_bash.jsonl"), + fixture_proc("turn2_resume_apply_patch.jsonl"), + ] + ) + handler = make_handler() + ctx = make_ctx( + sbx, + instructions="Be terse.", + options=CodexOptions( + reasoning_effort="low", + config={"sandbox_workspace_write.network_access": True}, + ), + ) + await handler.start(ctx) + events = await collect(handler, ctx, "create hello.txt") + + first = sbx.execs[0] + argv, env = first["cmd"], first["env"] + assert argv[:4] == ["codex", "exec", "--json", "--skip-git-repo-check"] + assert argv[-1] == "-" and argv[argv.index("-C") + 1] == "/work" + assert first["cwd"] == "/work" + assert "model_provider=litellm" in config_values(argv) + assert not any(TOKEN in a for a in argv) + assert env["LITELLM_HARNESS_TOKEN"] == TOKEN + assert env["CODEX_HOME"] == HOME + assert sbx.processes[0].stdin.data == b"create hello.txt" + assert sbx.processes[0].stdin.closed + + assert isinstance(events[-1], Text) + assert ctx.final_text.startswith("Done") + assert handler.native_session_id() == THREAD_ID + + await collect(handler, ctx, "edit it") + argv2 = sbx.execs[1]["cmd"] + assert argv2[:4] == ["codex", "exec", "resume", THREAD_ID] + assert "--sandbox" not in argv2 and "-C" not in argv2 + assert 'sandbox_mode="workspace-write"' in config_values(argv2) + assert ctx.final_text == "done" + + +async def test_resume_sets_thread_id(): + sbx = FakeSandbox(outputs=[fixture_proc("reasoning.jsonl")]) + handler = make_handler() + ctx = make_ctx(sbx) + await handler.start(ctx) + await handler.resume(ctx, "thread-9") + assert handler.native_session_id() == "thread-9" + await collect(handler, ctx, "again") + assert sbx.execs[0]["cmd"][:4] == ["codex", "exec", "resume", "thread-9"] + + +async def test_structured_output_sets_output_json(): + sbx = FakeSandbox(outputs=[fixture_proc("structured_output.jsonl")]) + handler = make_handler() + ctx = make_ctx(sbx, output=Answer, permissions="read-only") + await handler.start(ctx) + await collect(handler, ctx, "read hello.txt") + argv = sbx.execs[0]["cmd"] + assert argv[argv.index("--output-schema") + 1] == f"{HOME}/output_schema.json" + assert argv[argv.index("--sandbox") + 1] == "read-only" + assert Answer.model_validate_json(ctx.output_json).file == "hello.txt" + + +async def test_turn_failed_raises(): + sbx = FakeSandbox(outputs=[fixture_proc("turn_failed.jsonl", exit_code=1)]) + handler = make_handler() + ctx = make_ctx(sbx) + await handler.start(ctx) + with pytest.raises(HarnessTurnError, match="no healthy deployments"): + await collect(handler, ctx, "hi") + + +async def test_nonzero_exit_raises_with_stderr_tail(): + sbx = FakeSandbox( + outputs=[ + FakeProcess(b"", stderr=b"Error loading config.toml: bad\n", exit_code=1) + ] + ) + handler = make_handler() + ctx = make_ctx(sbx) + await handler.start(ctx) + with pytest.raises(HarnessTurnError, match=r"code 1: Error loading config\.toml"): + await collect(handler, ctx, "hi") + + +async def test_early_close_kills_process_and_stop_is_idempotent(): + sbx = FakeSandbox(outputs=[fixture_proc("turn1_bash.jsonl")]) + handler = make_handler() + ctx = make_ctx(sbx) + await handler.start(ctx) + gen = handler.turn(ctx, "hi") + await gen.__anext__() + await gen.aclose() + assert sbx.processes[0].killed + await handler.stop(ctx) + await handler.stop(ctx) + + +async def test_long_jsonl_line_is_parsed(): + text = "x" * 200_000 + line = json.dumps( + { + "type": "item.completed", + "item": {"id": "a", "type": "agent_message", "text": text}, + } + ) + sbx = FakeSandbox(outputs=[FakeProcess(line.encode() + b"\n")]) + handler = make_handler() + ctx = make_ctx(sbx) + await handler.start(ctx) + events = await collect(handler, ctx, "hi") + assert events == [Text(delta=text)] + + +def test_capabilities(): + cfg = CodexHarnessConfig() + caps = cfg.capabilities + assert cfg.harness is Harness.CODEX + assert cfg.options_type is CodexOptions + assert cfg.get_binary() == "codex" + assert "@openai/codex" in cfg.get_install_hint() + assert caps.structured_output and caps.skills and caps.resume + assert not ( + caps.tool_approval or caps.tool_filtering or caps.custom_tools or caps.history + ) + assert caps.permission_modes == frozenset({"read-only", "full"}) diff --git a/tests/unit/llms/deepagents/__init__.py b/tests/unit/llms/deepagents/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/deepagents/harness/__init__.py b/tests/unit/llms/deepagents/harness/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/deepagents/harness/test_sandbox_backend_symlinks.py b/tests/unit/llms/deepagents/harness/test_sandbox_backend_symlinks.py new file mode 100644 index 00000000000..022771400e5 --- /dev/null +++ b/tests/unit/llms/deepagents/harness/test_sandbox_backend_symlinks.py @@ -0,0 +1,78 @@ +"""A repository must not be able to reach host files through symlinks, in any file tool.""" + +import asyncio +import os +from pathlib import Path + +import pytest + +from litellm.harness.sandbox.local import LocalSandbox + +backend = pytest.importorskip("litellm.llms.deepagents.harness.sandbox_backend") + +SECRET = "AWS_SECRET_ACCESS_KEY=leaked-from-host" + + +@pytest.fixture +def repo_with_escape_links(tmp_path: Path) -> Path: + host = tmp_path / "host_home" + host.mkdir() + (host / "credentials").write_text(SECRET + "\n") + repo = tmp_path / "repo" + repo.mkdir() + (repo / "README.md").write_text("hello\n") + os.symlink(host / "credentials", repo / "creds_link") + os.symlink(host, repo / "home_link") + return repo + + +async def _backend(repo: Path) -> object: + return backend.SandboxBackend(LocalSandbox(str(repo)), loop=asyncio.get_running_loop(), writable=False) + + +async def test_grep_whole_repo_skips_symlinks_out_of_workspace(repo_with_escape_links: Path) -> None: + b = await _backend(repo_with_escape_links) + result = await b.agrep("AWS_SECRET") + assert not result.matches, f"grep followed a symlink out of the repo: {result}" + + +async def test_grep_rooted_at_symlink_dir_is_refused(repo_with_escape_links: Path) -> None: + b = await _backend(repo_with_escape_links) + result = await b.agrep("AWS_SECRET", path="/home_link") + assert result.error and "outside the workspace" in result.error + assert not result.matches + + +async def test_read_through_symlink_is_refused(repo_with_escape_links: Path) -> None: + b = await _backend(repo_with_escape_links) + result = await b.aread("/creds_link") + assert result.error and "outside the workspace" in result.error + assert SECRET not in str(result.file_data) + + +async def test_glob_does_not_list_files_behind_symlinks(repo_with_escape_links: Path) -> None: + b = await _backend(repo_with_escape_links) + result = await b.aglob("**/*") + paths = [m["path"] for m in result.matches or []] + assert paths == ["/README.md"] + + +async def test_grep_still_finds_real_repo_files(repo_with_escape_links: Path) -> None: + b = await _backend(repo_with_escape_links) + result = await b.agrep("hello") + assert [(m["path"], m["line"]) for m in result.matches] == [("/README.md", 1)] + + +async def test_write_into_new_nested_directory_is_allowed(tmp_path: Path) -> None: + repo = tmp_path / "repo" + repo.mkdir() + b = backend.SandboxBackend(LocalSandbox(str(repo)), loop=asyncio.get_running_loop(), writable=True) + result = await b.awrite("/new_dir/sub/file.py", "print('hi')\n") + assert result.error is None, result.error + assert (repo / "new_dir" / "sub" / "file.py").read_text() == "print('hi')\n" + + +async def test_write_under_symlinked_dir_is_refused(repo_with_escape_links: Path) -> None: + b = backend.SandboxBackend(LocalSandbox(str(repo_with_escape_links)), loop=asyncio.get_running_loop(), writable=True) + result = await b.awrite("/home_link/new_dir/evil.txt", "x") + assert result.error and "outside the workspace" in result.error diff --git a/tests/unit/llms/deepagents/harness/test_transformation.py b/tests/unit/llms/deepagents/harness/test_transformation.py new file mode 100644 index 00000000000..4409cead299 --- /dev/null +++ b/tests/unit/llms/deepagents/harness/test_transformation.py @@ -0,0 +1,185 @@ +import os +from pathlib import Path +from types import SimpleNamespace +from typing import Any + +import pytest +from pydantic import BaseModel + +from litellm.harness.context import GatewayTarget, SessionContext +from litellm.harness.errors import OptionsMismatch +from litellm.harness.options import CodexOptions, DeepAgentsOptions +from litellm.harness.sandbox.local import LocalSandbox +from litellm.harness.types import Harness, Reasoning, Text, ToolCall, ToolResult +from litellm.llms.deepagents.harness import transformation as da + + +def make_ctx(tmp_path: Path, **kwargs: Any) -> SessionContext: + base: dict[str, Any] = { + "harness": Harness.DEEPAGENTS, + "sandbox": LocalSandbox(tmp_path), + "session_id": f"s-{os.urandom(4).hex()}", + "model": "gpt-4o-mini", + } + return SessionContext(**{**base, **kwargs}) + + +def msg(kind: str, **fields: Any) -> SimpleNamespace: + return SimpleNamespace(type=kind, **fields) + + +def test_blocked_tools_modes() -> None: + assert da.blocked_tools("full", []) == frozenset() + assert da.blocked_tools("edit", []) == frozenset({"execute"}) + assert {"write_file", "edit_file", "delete", "execute"} <= da.blocked_tools( + "read-only", [] + ) + assert da.blocked_tools("full", ["read", "ls"]) == frozenset({"read_file", "ls"}) + assert da.blocked_tools("full", ["bash", "grep"]) == frozenset({"execute", "grep"}) + + +def test_interrupt_config_only_for_ask() -> None: + assert da.interrupt_config("full", frozenset()) is None + config = da.interrupt_config("ask", frozenset({"execute"})) + assert set(config) == {"write_file", "edit_file", "delete"} + assert all( + v == {"allowed_decisions": ["approve", "reject"]} for v in config.values() + ) + + +def test_normalized_tool_name() -> None: + assert da.normalized_tool_name("write_file") == "write" + assert da.normalized_tool_name("read_file") == "read" + assert da.normalized_tool_name("edit_file") == "edit" + assert da.normalized_tool_name("execute") == "bash" + assert da.normalized_tool_name("add") == "add" + + +def test_chat_model_kwargs_gateway_and_sdk(tmp_path: Path) -> None: + gw = GatewayTarget(api_base="https://gw.example.com", api_key="sk-virtual") + ctx = make_ctx(tmp_path, gateway=gw, metadata={"team": "a"}) + kwargs = da.chat_model_kwargs(ctx) + assert kwargs["model"] == "litellm_proxy/gpt-4o-mini" + assert kwargs["api_base"] == "https://gw.example.com" + assert kwargs["api_key"] == "sk-virtual" + assert kwargs["extra_headers"]["x-litellm-tags"] == "harness,deepagents" + assert '"team": "a"' in kwargs["extra_headers"]["x-litellm-spend-logs-metadata"] + no_meta = da.chat_model_kwargs(make_ctx(tmp_path, gateway=gw)) + assert "x-litellm-spend-logs-metadata" not in no_meta["extra_headers"] + + sdk = da.chat_model_kwargs(make_ctx(tmp_path, api_key="k", api_base="http://b")) + assert sdk == {"model": "gpt-4o-mini", "api_key": "k", "api_base": "http://b"} + with pytest.raises(ValueError, match="needs model="): + da.chat_model_kwargs(make_ctx(tmp_path, model=None)) + + +def test_recursion_limit(tmp_path: Path) -> None: + assert ( + da.recursion_limit(make_ctx(tmp_path)) == da.DEEPAGENTS_DEFAULT_RECURSION_LIMIT + ) + assert da.recursion_limit(make_ctx(tmp_path, max_turns=2)) == ( + da.DEEPAGENTS_BASE_RECURSION_LIMIT + 2 * da.DEEPAGENTS_STEPS_PER_TURN + ) + opts = DeepAgentsOptions(recursion_limit=7) + assert da.recursion_limit(make_ctx(tmp_path, max_turns=2, options=opts)) == 7 + + +def test_stream_events_text_and_reasoning() -> None: + assert da.stream_events(msg("human", content="hi")) == [] + events = da.stream_events( + msg( + "AIMessageChunk", + content=[ + {"type": "thinking", "thinking": "hmm"}, + {"type": "text", "text": "a"}, + "b", + ], + additional_kwargs={}, + ) + ) + assert events == [Reasoning(delta="hmm"), Text(delta="ab")] + extra = da.stream_events( + msg("ai", content="x", additional_kwargs={"reasoning_content": "r"}) + ) + assert extra == [Reasoning(delta="r"), Text(delta="x")] + + +def test_update_events_tool_calls_results_and_skip() -> None: + ai = msg( + "ai", + tool_calls=[ + {"name": "write_file", "args": {"file_path": "/a"}, "id": "c1"}, + {"name": "Answer", "args": {"city": "Paris"}, "id": "c2"}, + {"name": "add", "args": None, "id": "c3"}, + ], + ) + tool = msg("tool", name="write_file", tool_call_id="c1", content="ok", status=None) + err = msg("tool", name="execute", tool_call_id="c4", content="x", status="error") + skipped = msg("tool", name="Answer", tool_call_id="c2", content="", status=None) + update = { + "model": {"messages": [ai]}, + "tools": {"messages": [tool, err, skipped]}, + "SomeMiddleware.after_model": {"messages": [ai]}, + } + events = da.update_events(update, frozenset({"Answer"})) + assert events == [ + ToolCall( + id="c1", + name="write", + native_name="write_file", + input={"file_path": "/a"}, + builtin=True, + ), + ToolCall( + id="c3", name="add", native_name="add", input={"args": None}, builtin=False + ), + ToolResult(id="c1", output="ok", is_error=False), + ToolResult(id="c4", output="x", is_error=True), + ] + assert da.update_events(None, frozenset()) == [] + assert da.update_events({"model": None}, frozenset()) == [] + + +def test_interrupts_and_approval_requests() -> None: + assert da.interrupts_in({"__interrupt__": ("i",)}) == ["i"] + assert da.interrupts_in({}) == [] and da.interrupts_in(None) == [] + value = {"action_requests": [{"name": "write_file", "args": {}}, "junk"]} + assert da.approval_requests(value) == [{"name": "write_file", "args": {}}] + assert da.approval_requests(None) == [] + assert da.approval_requests({"action_requests": "x"}) == [] + + +def test_decision() -> None: + assert da.decision(True, "") == {"type": "approve"} + assert da.decision(False, "no") == {"type": "reject", "message": "no"} + assert da.decision(False, "")["message"] + + +class Answer(BaseModel): + city: str + + +def test_final_ai_text_and_structured_json() -> None: + messages = [ + msg("ai", content="first"), + msg("tool", content="t"), + msg("ai", content=""), + ] + assert da.final_ai_text(messages) == "first" + assert da.final_ai_text([]) == "" + assert da.structured_json(None) is None + assert Answer.model_validate_json(da.structured_json(Answer(city="Paris"))) + assert da.structured_json({"city": "Paris"}) == '{"city": "Paris"}' + + +def test_config_capabilities_and_validation(tmp_path: Path) -> None: + config = da.DeepAgentsHarnessConfig() + assert config.uses_model_endpoint is False + assert config.capabilities.tool_approval and config.capabilities.history + assert "ask" in config.capabilities.permission_modes + config.validate_environment(make_ctx(tmp_path)) + with pytest.raises(ValueError, match="needs model="): + config.validate_environment(make_ctx(tmp_path, model=None)) + with pytest.raises(OptionsMismatch): + config.validate_environment(make_ctx(tmp_path, options=CodexOptions())) + assert "pip install deepagents langchain-litellm" in da.INSTALL_HINT diff --git a/tests/unit/llms/opencode/__init__.py b/tests/unit/llms/opencode/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/opencode/harness/__init__.py b/tests/unit/llms/opencode/harness/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/opencode/harness/fixtures/__init__.py b/tests/unit/llms/opencode/harness/fixtures/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/opencode/harness/fixtures/api_error.jsonl b/tests/unit/llms/opencode/harness/fixtures/api_error.jsonl new file mode 100644 index 00000000000..b2b3148285e --- /dev/null +++ b/tests/unit/llms/opencode/harness/fixtures/api_error.jsonl @@ -0,0 +1 @@ +{"type":"error","timestamp":1790788205744,"sessionID":"ses_f0cb48565ffeMWhVl1J584kSti","error":{"name":"APIError","data":{"message":"litellm.BadRequestError: You passed in model=no-such-model-xyz. There are no healthy deployments for this model","statusCode":400,"isRetryable":false}}} diff --git a/tests/unit/llms/opencode/harness/fixtures/endpoint_requests.jsonl b/tests/unit/llms/opencode/harness/fixtures/endpoint_requests.jsonl new file mode 100644 index 00000000000..0b4c3d22bf4 --- /dev/null +++ b/tests/unit/llms/opencode/harness/fixtures/endpoint_requests.jsonl @@ -0,0 +1,4 @@ +{"method":"POST","path":"/v1/chat/completions","headers":{"Content-Type":"application/json","User-Agent":"opencode/1.14.41 ai-sdk/provider-utils/4.0.23 runtime/bun/1.3.13","x-session-affinity":"ses_f0cb977cdffeoCMeplOiw1KY25","Connection":"keep-alive","Accept":"*/*"},"auth_prefix":"Bearer ","body_keys":["max_tokens","messages","model","stream","stream_options"],"model":"claude-haiku-4-5-20251001","stream":true,"stream_options":{"include_usage":true},"tools":[],"n_messages":3,"roles":["system","user","user"]} +{"method":"POST","path":"/v1/chat/completions","headers":{"Content-Type":"application/json","User-Agent":"opencode/1.14.41 ai-sdk/provider-utils/4.0.23 runtime/bun/1.3.13","x-session-affinity":"ses_f0cb977cdffeoCMeplOiw1KY25","Connection":"keep-alive","Accept":"*/*"},"auth_prefix":"Bearer ","body_keys":["max_tokens","messages","model","stream","stream_options","tool_choice","tools"],"model":"claude-haiku-4-5-20251001","stream":true,"stream_options":{"include_usage":true},"tools":["bash","read","glob","grep","edit","write","task","webfetch","todowrite","skill"],"n_messages":2,"roles":["system","user"]} +{"method":"POST","path":"/v1/chat/completions","headers":{"Content-Type":"application/json","User-Agent":"opencode/1.14.41 ai-sdk/provider-utils/4.0.23 runtime/bun/1.3.13","x-session-affinity":"ses_f0cb977cdffeoCMeplOiw1KY25","Connection":"keep-alive","Accept":"*/*"},"auth_prefix":"Bearer ","body_keys":["max_tokens","messages","model","stream","stream_options","tool_choice","tools"],"model":"claude-haiku-4-5-20251001","stream":true,"stream_options":{"include_usage":true},"tools":["bash","read","glob","grep","edit","write","task","webfetch","todowrite","skill"],"n_messages":4,"roles":["system","user","assistant","tool"]} +{"method":"POST","path":"/v1/chat/completions","headers":{"Content-Type":"application/json","User-Agent":"opencode/1.14.41 ai-sdk/provider-utils/4.0.23 runtime/bun/1.3.13","x-session-affinity":"ses_f0cb977cdffeoCMeplOiw1KY25","Connection":"keep-alive","Accept":"*/*"},"auth_prefix":"Bearer ","body_keys":["max_tokens","messages","model","stream","stream_options","tool_choice","tools"],"model":"claude-haiku-4-5-20251001","stream":true,"stream_options":{"include_usage":true},"tools":["bash","read","glob","grep","edit","write","task","webfetch","todowrite","skill"],"n_messages":6,"roles":["system","user","assistant","tool","assistant","tool"]} diff --git a/tests/unit/llms/opencode/harness/fixtures/readonly_denied_bash.jsonl b/tests/unit/llms/opencode/harness/fixtures/readonly_denied_bash.jsonl new file mode 100644 index 00000000000..e31cb95503c --- /dev/null +++ b/tests/unit/llms/opencode/harness/fixtures/readonly_denied_bash.jsonl @@ -0,0 +1,7 @@ +{"type":"step_start","timestamp":1790787993225,"sessionID":"ses_f0cb7c593ffesLm4De78bjOZjn","part":{"id":"prt_0f3483e81001WZXwo1sKkhgOIo","messageID":"msg_0f3483ac90012uFburOGJS91gV","sessionID":"ses_f0cb7c593ffesLm4De78bjOZjn","type":"step-start"}} +{"type":"tool_use","timestamp":1790787993571,"sessionID":"ses_f0cb7c593ffesLm4De78bjOZjn","part":{"type":"tool","tool":"invalid","callID":"toolu_0126GuE9NKXyF4HoXLY3rREx","state":{"status":"completed","input":{"tool":"bash","error":"Model tried to call unavailable tool 'bash'. Available tools: invalid, read, glob, grep, task, todowrite, skill."},"output":"The arguments provided to the tool are invalid: Model tried to call unavailable tool 'bash'. Available tools: invalid, read, glob, grep, task, todowrite, skill.","metadata":{"truncated":false},"title":"Invalid Tool","time":{"start":1790787993567,"end":1790787993570}},"id":"prt_0f3483f5e001tfzTsZE0IrrCxH","sessionID":"ses_f0cb7c593ffesLm4De78bjOZjn","messageID":"msg_0f3483ac90012uFburOGJS91gV"}} +{"type":"text","timestamp":1790787993573,"sessionID":"ses_f0cb7c593ffesLm4De78bjOZjn","part":{"id":"prt_0f3483e85001y7vW8x70lZJBwu","messageID":"msg_0f3483ac90012uFburOGJS91gV","sessionID":"ses_f0cb7c593ffesLm4De78bjOZjn","type":"text","text":"I'll run that shell command for you using bash.","time":{"start":1790787993221,"end":1790787993572}}} +{"type":"step_finish","timestamp":1790787993575,"sessionID":"ses_f0cb7c593ffesLm4De78bjOZjn","part":{"id":"prt_0f3483fe5001VEbV0Sy9d0ybb0","reason":"tool-calls","messageID":"msg_0f3483ac90012uFburOGJS91gV","sessionID":"ses_f0cb7c593ffesLm4De78bjOZjn","type":"step-finish","tokens":{"total":7484,"input":7416,"output":68,"reasoning":0,"cache":{"write":0,"read":0}},"cost":0}} +{"type":"step_start","timestamp":1790787994554,"sessionID":"ses_f0cb7c593ffesLm4De78bjOZjn","part":{"id":"prt_0f34843b6001QhtCAWotsESvRK","messageID":"msg_0f3483fea0021zWRKRkh32701g","sessionID":"ses_f0cb7c593ffesLm4De78bjOZjn","type":"step-start"}} +{"type":"text","timestamp":1790787994748,"sessionID":"ses_f0cb7c593ffesLm4De78bjOZjn","part":{"id":"prt_0f34843b900143p3b7263tu3pc","messageID":"msg_0f3483fea0021zWRKRkh32701g","sessionID":"ses_f0cb7c593ffesLm4De78bjOZjn","type":"text","text":"FAILED\n\nThe bash tool is not available in this environment, so I cannot execute that shell command.","time":{"start":1790787994553,"end":1790787994747}}} +{"type":"step_finish","timestamp":1790787994749,"sessionID":"ses_f0cb7c593ffesLm4De78bjOZjn","part":{"id":"prt_0f348447c001hnUGx5uag7BAYm","reason":"stop","messageID":"msg_0f3483fea0021zWRKRkh32701g","sessionID":"ses_f0cb7c593ffesLm4De78bjOZjn","type":"step-finish","tokens":{"total":7596,"input":158,"output":24,"reasoning":0,"cache":{"write":0,"read":7414}},"cost":0}} diff --git a/tests/unit/llms/opencode/harness/fixtures/turn1_write_read.jsonl b/tests/unit/llms/opencode/harness/fixtures/turn1_write_read.jsonl new file mode 100644 index 00000000000..83eab9e49a8 --- /dev/null +++ b/tests/unit/llms/opencode/harness/fixtures/turn1_write_read.jsonl @@ -0,0 +1,11 @@ +{"type":"step_start","timestamp":1790787882230,"sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","part":{"id":"prt_0f3468cf4001Lp09mwOdY2Nk0v","messageID":"msg_0f3468893001gnm1GTIsdViRXS","sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","type":"step-start"}} +{"type":"text","timestamp":1790787882762,"sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","part":{"id":"prt_0f3468cf5001EnUgA050AAqsKb","messageID":"msg_0f3468893001gnm1GTIsdViRXS","sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","type":"text","text":"I'll create a hello.txt file containing \"hi\" and then read it.","time":{"start":1790787882229,"end":1790787882762}}} +{"type":"tool_use","timestamp":1790787882770,"sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","part":{"type":"tool","tool":"write","callID":"toolu_015FFUwEf2dazoWfCrMbMCnm","state":{"status":"completed","input":{"filePath":"/workspace/hello.txt","content":"hi"},"output":"Wrote file successfully.","metadata":{"diagnostics":{},"filepath":"/workspace/hello.txt","exists":false,"truncated":false},"title":"private/workspace/hello.txt","time":{"start":1790787882760,"end":1790787882768}},"id":"prt_0f3468dcf001rxCR4QNHtzY272","sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","messageID":"msg_0f3468893001gnm1GTIsdViRXS"}} +{"type":"step_finish","timestamp":1790787882770,"sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","part":{"id":"prt_0f3468f11001woc6vGYjk4ErHI","reason":"tool-calls","messageID":"msg_0f3468893001gnm1GTIsdViRXS","sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","type":"step-finish","tokens":{"total":11018,"input":10918,"output":100,"reasoning":0,"cache":{"write":0,"read":0}},"cost":0}} +{"type":"step_start","timestamp":1790787905598,"sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","part":{"id":"prt_0f346e8390015Ikck6AJPLpFr0","messageID":"msg_0f3468f14001j9Au0eVNVNWev1","sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","type":"step-start"}} +{"type":"text","timestamp":1790787905940,"sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","part":{"id":"prt_0f346e83d001bMWqsgMxNP5FGG","messageID":"msg_0f3468f14001j9Au0eVNVNWev1","sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","type":"text","text":"Now let me read the file:","time":{"start":1790787905597,"end":1790787905940}}} +{"type":"tool_use","timestamp":1790787905949,"sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","part":{"type":"tool","tool":"read","callID":"toolu_017766NprB4499fkNFoqLx6k","state":{"status":"completed","input":{"filePath":"/workspace/hello.txt"},"output":"/workspace/hello.txt\nfile\n\n1: hi\n\n(End of file - total 1 lines)\n","metadata":{"preview":"hi","truncated":false,"loaded":[]},"title":"private/workspace/hello.txt","time":{"start":1790787905937,"end":1790787905947}},"id":"prt_0f346e8be001TkCzxNO3iCWyHh","sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","messageID":"msg_0f3468f14001j9Au0eVNVNWev1"}} +{"type":"step_finish","timestamp":1790787905949,"sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","part":{"id":"prt_0f346e99c0018zQ7l02y6UXvHZ","reason":"tool-calls","messageID":"msg_0f3468f14001j9Au0eVNVNWev1","sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","type":"step-finish","tokens":{"total":11107,"input":118,"output":73,"reasoning":0,"cache":{"write":0,"read":10916}},"cost":0}} +{"type":"step_start","timestamp":1790787907026,"sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","part":{"id":"prt_0f346edd1001UEsH2Hitcy599I","messageID":"msg_0f346e99f0010o3y8k52od8t3I","sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","type":"step-start"}} +{"type":"text","timestamp":1790787907596,"sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","part":{"id":"prt_0f346edd10036kWCbbi5l5qxja","messageID":"msg_0f346e99f0010o3y8k52od8t3I","sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","type":"text","text":"Done! I've created hello.txt with the content \"hi\" and read it back. The file contains exactly what was requested.","time":{"start":1790787907025,"end":1790787907595}}} +{"type":"step_finish","timestamp":1790787907597,"sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","part":{"id":"prt_0f346f00c002zjKza4p4eIg1Ox","reason":"stop","messageID":"msg_0f346e99f0010o3y8k52od8t3I","sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","type":"step-finish","tokens":{"total":11198,"input":139,"output":30,"reasoning":0,"cache":{"write":0,"read":11029}},"cost":0}} diff --git a/tests/unit/llms/opencode/harness/fixtures/turn2_session_skill.jsonl b/tests/unit/llms/opencode/harness/fixtures/turn2_session_skill.jsonl new file mode 100644 index 00000000000..76d40e08c99 --- /dev/null +++ b/tests/unit/llms/opencode/harness/fixtures/turn2_session_skill.jsonl @@ -0,0 +1,7 @@ +{"type":"step_start","timestamp":1790787953184,"sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","part":{"id":"prt_0f347a21e001EmD408wvQkLCan","messageID":"msg_0f3479cf2001JlcVj8lSmhqAum","sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","type":"step-start"}} +{"type":"text","timestamp":1790787953736,"sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","part":{"id":"prt_0f347a21f001Bu2RKXM9urS1v9","messageID":"msg_0f3479cf2001JlcVj8lSmhqAum","sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","type":"text","text":"I created `/workspace/hello.txt`.\n\nFor the secret greeting word, let me load the greeter skill:","time":{"start":1790787953183,"end":1790787953736}}} +{"type":"tool_use","timestamp":1790787953752,"sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","part":{"type":"tool","tool":"skill","callID":"toolu_01Q1ip1d6snmazhbYCaEqcsT","state":{"status":"completed","input":{"name":"greeter"},"output":"\n# Skill: greeter\n\nThe secret greeting word is PINEAPPLE.\n\nBase directory for this skill: file:///tmp/litellm-harness-x/skills/greeter\nRelative paths in this skill (e.g., scripts/, reference/) are relative to this base directory.\nNote: file list is sampled.\n\n\n\n\n","metadata":{"name":"greeter","dir":"/tmp/litellm-harness-x/skills/greeter","truncated":false},"title":"Loaded skill: greeter","time":{"start":1790787953733,"end":1790787953749}},"id":"prt_0f347a3ab001knlkXMZcS3jV9G","sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","messageID":"msg_0f3479cf2001JlcVj8lSmhqAum"}} +{"type":"step_finish","timestamp":1790787953752,"sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","part":{"id":"prt_0f347a457001JH8PaFwNLIRfA9","reason":"tool-calls","messageID":"msg_0f3479cf2001JlcVj8lSmhqAum","sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","type":"step-finish","tokens":{"total":11531,"input":11446,"output":85,"reasoning":0,"cache":{"write":0,"read":0}},"cost":0}} +{"type":"step_start","timestamp":1790787970574,"sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","part":{"id":"prt_0f347e60c001CUPeXPcH733B3G","messageID":"msg_0f347a45b001XQWc2pM6IPZrsO","sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","type":"step-start"}} +{"type":"text","timestamp":1790787970675,"sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","part":{"id":"prt_0f347e60d001pQ2aAgMwHkYRx3","messageID":"msg_0f347a45b001XQWc2pM6IPZrsO","sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","type":"text","text":"The secret greeting word is **PINEAPPLE**.","time":{"start":1790787970573,"end":1790787970674}}} +{"type":"step_finish","timestamp":1790787970676,"sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","part":{"id":"prt_0f347e673002k2mm651nyRtr07","reason":"stop","messageID":"msg_0f347a45b001XQWc2pM6IPZrsO","sessionID":"ses_f0cb977cdffeoCMeplOiw1KY25","type":"step-finish","tokens":{"total":11663,"input":204,"output":15,"reasoning":0,"cache":{"write":0,"read":11444}},"cost":0}} diff --git a/tests/unit/llms/opencode/harness/test_transformation.py b/tests/unit/llms/opencode/harness/test_transformation.py new file mode 100644 index 00000000000..2caf1ef80b8 --- /dev/null +++ b/tests/unit/llms/opencode/harness/test_transformation.py @@ -0,0 +1,726 @@ +import asyncio +import json +from dataclasses import dataclass, field +from pathlib import Path + +import pytest +from pydantic import BaseModel + +from litellm.harness.context import SessionContext +from litellm.harness.errors import ( + CapabilityUnsupported, + HarnessError, + HarnessInstallFailed, + OptionsMismatch, +) +from litellm.harness.handlers.cli_handler import PERSIST_DIR_SCRIPT, CLIHarnessHandler +from litellm.harness.options import CodexOptions, OpenCodeOptions +from litellm.harness.sandbox.base import CompletedRun +from litellm.harness.types import Harness, Reasoning, Text, ToolCall, ToolResult +from litellm.llms.base_llm.harness.transformation import ( + HarnessSessionSetup, + HarnessTurnError, +) +from litellm.llms.opencode.harness.transformation import ( + INSTRUCTIONS_FILENAME, + OPENCODE_ISOLATION_ENV, + OPENCODE_SESSION_TITLE, + TOKEN_FILENAME, + XDG_DIRNAME, + OpenCodeHarnessConfig, + OpenCodeStreamState, + build_instructions, + build_opencode_config, + permission_rules, + turn_prompt, + validate_user_config, +) + +FIXTURES = Path(__file__).parent / "fixtures" +TOKEN = "tok-secret-123" +SESSION = "ses_f0cb977cdffeoCMeplOiw1KY25" +PRIVATE = "/tmp/oc-1" +CONFIG = OpenCodeHarnessConfig() + + +def load_fixture(name: str) -> list[dict]: + return [ + json.loads(line) for line in (FIXTURES / name).read_text().splitlines() if line + ] + + +def parse(obj: dict, state: OpenCodeStreamState) -> list: + return CONFIG.transform_stream_line(obj, state) + + +def parse_all(name: str, state: OpenCodeStreamState | None = None): + state = state or CONFIG.create_stream_state() + events = [] + for obj in load_fixture(name): + events.extend(parse(obj, state)) + return events, state + + +# --------------------------------------------------------------------------- fakes + + +class FakeStdin: + def __init__(self): + self.data = b"" + self.closed = False + + def write(self, data: bytes) -> None: + self.data += data + + async def drain(self) -> None: + return None + + def close(self) -> None: + self.closed = True + + +class FakeProcess: + def __init__(self, stdout: bytes, stderr: bytes = b"", exit_code: int = 0): + self.stdin = FakeStdin() + self.stdout = asyncio.StreamReader() + self.stdout.feed_data(stdout) + self.stdout.feed_eof() + self.stderr = asyncio.StreamReader() + self.stderr.feed_data(stderr) + self.stderr.feed_eof() + self._exit_code = exit_code + self.killed = False + + async def wait(self) -> int: + return self._exit_code + + async def kill(self) -> None: + self.killed = True + + +@dataclass +class FakeSandbox: + workdir: str = "/work" + has_binary: bool = True + persist_ok: bool = True + outputs: list = field(default_factory=list) + files: dict = field(default_factory=dict) + execs: list = field(default_factory=list) + runs: list = field(default_factory=list) + tempdirs: int = 0 + + async def exec(self, cmd, *, env=None, cwd=None): + self.execs.append({"cmd": cmd, "env": dict(env or {}), "cwd": cwd}) + return self.outputs.pop(0) + + async def run(self, cmd, *, env=None, cwd=None, timeout=None): + self.runs.append(cmd) + if self.persist_ok: + return CompletedRun("", "", 0) + return CompletedRun("", "read-only fs", 1) + + async def read(self, path): + return self.files[path] + + async def write(self, path, data): + self.files[path] = data + + def host_url(self, port): + return f"http://host.docker.internal:{port}" + + async def which(self, binary): + return f"/usr/bin/{binary}" if self.has_binary else None + + async def tempdir(self): + self.tempdirs += 1 + return f"/tmp/oc-{self.tempdirs}" + + async def snapshot(self): + return {} + + async def close(self): + return None + + +@dataclass +class FakeEndpoint: + port: int = 4555 + token: str = TOKEN + model: str | None = None + + +class Answer(BaseModel): + file: str + content: str + + +def make_ctx(sandbox=None, **kwargs) -> SessionContext: + return SessionContext( + harness=Harness.OPENCODE, + sandbox=sandbox or FakeSandbox(), + session_id="s1", + model=kwargs.pop("model", "claude-haiku-4-5-20251001"), + endpoint=kwargs.pop("endpoint", FakeEndpoint()), + **kwargs, + ) + + +def setup_for(ctx: SessionContext) -> HarnessSessionSetup: + return CONFIG.transform_session_setup(ctx, PRIVATE) + + +def setup_config(setup: HarnessSessionSetup) -> dict: + return json.loads(setup.env["OPENCODE_CONFIG_CONTENT"]) + + +def fixture_proc(name: str, **kwargs) -> FakeProcess: + return FakeProcess((FIXTURES / name).read_bytes(), **kwargs) + + +async def collect(handler, ctx, prompt): + return [e async for e in handler.turn(ctx, prompt)] + + +async def started(sandbox=None, **kwargs): + sandbox = sandbox or FakeSandbox() + handler = CLIHarnessHandler(OpenCodeHarnessConfig()) + ctx = make_ctx(sandbox, **kwargs) + await handler.start(ctx) + return handler, ctx, sandbox + + +def exec_config(sandbox, index=0) -> dict: + return json.loads(sandbox.execs[index]["env"]["OPENCODE_CONFIG_CONTENT"]) + + +# --------------------------------------------------------------------------- parsing + + +def test_parse_write_read_turn(): + events, state = parse_all("turn1_write_read.jsonl") + assert CONFIG.get_native_session_id(state) == SESSION + assert [type(e) for e in events] == [ + Text, + ToolCall, + ToolResult, + Text, + ToolCall, + ToolResult, + Text, + ] + write, write_result = events[1], events[2] + assert write.name == "write" and write.native_name == "write" + assert write.builtin is True + assert write.input == {"filePath": "/workspace/hello.txt", "content": "hi"} + assert write_result.id == write.id == "toolu_015FFUwEf2dazoWfCrMbMCnm" + assert write_result.is_error is False + read, read_result = events[4], events[5] + assert read.name == "read" and "1: hi" in read_result.output + assert state.final_text.startswith("Done!") + assert state.error is None + + +def test_parse_skill_tool_on_continued_session(): + events, state = parse_all("turn2_session_skill.jsonl") + assert state.session_id == SESSION + skill = next(e for e in events if isinstance(e, ToolCall)) + assert skill.name == "skill" and skill.input == {"name": "greeter"} + assert skill.builtin is True + assert "PINEAPPLE" in state.final_text + + +def test_parse_denied_tool_is_error_result(): + events, state = parse_all("readonly_denied_bash.jsonl") + call = next(e for e in events if isinstance(e, ToolCall)) + result = next(e for e in events if isinstance(e, ToolResult)) + assert call.native_name == "invalid" and call.input["tool"] == "bash" + assert result.is_error is True + assert state.final_text.startswith("FAILED") + + +def test_parse_api_error_records_error(): + events, state = parse_all("api_error.jsonl") + assert events == [] + assert "no healthy deployments" in state.error + + +def test_parse_reasoning_and_tool_error_and_name_mapping(): + state = OpenCodeStreamState() + reasoning = parse( + {"type": "reasoning", "sessionID": "s", "part": {"text": "thinking hard"}}, + state, + ) + assert reasoning == [Reasoning(delta="thinking hard")] + failed = parse( + { + "type": "tool_use", + "part": { + "tool": "bash", + "callID": "c1", + "state": { + "status": "error", + "input": {"command": "x"}, + "error": "boom", + }, + }, + }, + state, + ) + assert failed[1] == ToolResult(id="c1", output="boom", is_error=True) + for native, normalized in [ + ("list", "ls"), + ("webfetch", "web_search"), + ("glob", "glob"), + ("grep", "grep"), + ("apply_patch", "edit"), + ]: + call = parse( + { + "type": "tool_use", + "part": { + "tool": native, + "callID": "x", + "state": {"status": "completed", "input": {}, "output": ""}, + }, + }, + state, + )[0] + assert call.name == normalized + mcp = parse( + { + "type": "tool_use", + "part": { + "tool": "github_search", + "callID": "m", + "state": {"status": "completed", "input": {}, "output": {"a": 1}}, + }, + }, + state, + ) + assert mcp[0].builtin is False and mcp[1].output == '{"a": 1}' + assert state.session_id == "s" + + +def test_final_text_is_last_step_text(): + state = OpenCodeStreamState() + parse({"type": "step_start"}, state) + parse({"type": "text", "part": {"text": "working"}}, state) + parse({"type": "step_start"}, state) + parse({"type": "text", "part": {"text": "a"}}, state) + parse({"type": "text", "part": {"text": "b"}}, state) + assert state.final_text == "a\n\nb" + + +def test_error_event_message_shapes(): + state = OpenCodeStreamState() + parse({"type": "error", "error": {"data": {"message": "m1"}}}, state) + parse({"type": "error", "error": {"name": "APIError"}}, state) + parse({"type": "error", "error": "raw"}, state) + assert state.error == "m1\nAPIError\nraw" + + +# --------------------------------------------------------------------------- config + + +def test_permission_mapping(): + assert permission_rules("full", ()) == {"*": "allow"} + assert permission_rules("read-only", ()) == { + "edit": "deny", + "bash": "deny", + "webfetch": "deny", + } + edit = permission_rules("edit", ()) + assert edit["edit"] == "allow" and edit["bash"] == "deny" + with pytest.raises(CapabilityUnsupported): + permission_rules("ask", ()) + + +def test_disable_tools_map_to_native_denies_after_wildcard(): + rules = permission_rules("full", ["bash", "web_search", "ls", "write"]) + assert list(rules)[0] == "*" + assert rules["bash"] == "deny" + assert rules["webfetch"] == rules["websearch"] == "deny" + assert rules["list"] == "deny" + assert rules["edit"] == "deny" + + +def test_build_config_merges_user_config_under_managed_keys(): + config = build_opencode_config( + model="m1", + base_url="http://h:1/v1", + token_path="/tmp/p/token", + permissions="full", + user_config={"instructions": ["RULES.md"], "compaction": {"auto": False}}, + instructions_path="/tmp/p/instructions.md", + skills_path="/tmp/p/skills", + ) + provider = config["provider"]["litellm"] + assert provider["npm"] == "@ai-sdk/openai-compatible" + assert provider["options"] == { + "baseURL": "http://h:1/v1", + "apiKey": "{file:/tmp/p/token}", + } + assert provider["models"] == {"m1": {}} + assert config["model"] == config["small_model"] == "litellm/m1" + assert config["enabled_providers"] == ["litellm"] + assert config["instructions"] == ["RULES.md", "/tmp/p/instructions.md"] + assert config["skills"] == {"paths": ["/tmp/p/skills"]} + assert config["compaction"] == {"auto": False} + + +@pytest.mark.parametrize( + "config", + [ + {"provider": {}}, + {"model": "openai/gpt-5"}, + {"permission": {"*": "allow"}}, + {"tools": {"bash": True}}, + {"agent": {"build": {"permission": {"bash": "allow"}}}}, + {"mode": {"x": {"model": "a/b"}}}, + {"agent": "not-a-mapping"}, + ], +) +def test_managed_keys_rejected(config): + with pytest.raises(OptionsMismatch): + validate_user_config(config) + + +def test_config_metadata(): + assert CONFIG.get_binary() == "opencode" + assert "opencode" in CONFIG.get_install_hint() + assert CONFIG.uses_model_endpoint is True + assert CONFIG.capabilities.permission_modes == {"read-only", "edit", "full"} + + +def test_validate_environment_rejects_wrong_options_and_managed_config(): + CONFIG.validate_environment(make_ctx()) + with pytest.raises(OptionsMismatch): + CONFIG.validate_environment(make_ctx(options=CodexOptions())) + with pytest.raises(OptionsMismatch): + CONFIG.validate_environment( + make_ctx(options=OpenCodeOptions(config={"model": "openai/x"})) + ) + + +def test_session_setup_token_only_in_private_file(): + setup = setup_for(make_ctx()) + assert setup.files == {TOKEN_FILENAME: TOKEN.encode()} + assert TOKEN not in json.dumps(dict(setup.env)) + config = setup_config(setup) + assert config["provider"]["litellm"]["options"] == { + "baseURL": "http://host.docker.internal:4555/v1", + "apiKey": "{file:/tmp/oc-1/token}", + } + assert config["permission"] == {"*": "allow"} + + +def test_session_setup_env_and_persisted_xdg(): + setup = setup_for(make_ctx(options=OpenCodeOptions(env={"FOO": "1"}))) + assert list(setup.persisted_dirs) == [(XDG_DIRNAME, "opencode")] + assert setup.skills_dir == "skills" + env = setup.env + for sub in ("config", "data", "state", "cache"): + assert env[f"XDG_{sub.upper()}_HOME"] == f"{PRIVATE}/xdg/{sub}" + for key, value in OPENCODE_ISOLATION_ENV.items(): + assert env[key] == value + assert env["OPENCODE_CONFIG"] == "" and env["OPENCODE_PERMISSION"] == "" + assert env["FOO"] == "1" + + +def test_session_setup_errors(): + with pytest.raises(HarnessError): + setup_for(make_ctx(endpoint=None)) + with pytest.raises(ValueError, match="needs model="): + setup_for(make_ctx(model=None, endpoint=FakeEndpoint(model=None))) + + +def test_session_setup_read_only_and_disable_tools(): + setup = setup_for(make_ctx(permissions="read-only", disable_tools=["grep"])) + assert setup_config(setup)["permission"] == { + "edit": "deny", + "bash": "deny", + "webfetch": "deny", + "grep": "deny", + } + + +def test_session_setup_instructions_and_skills(): + ctx = make_ctx(instructions="Be terse.", output=Answer, skills=["/s/greeter"]) + setup = setup_for(ctx) + written = setup.files[INSTRUCTIONS_FILENAME].decode() + assert written == build_instructions(ctx) + assert written.startswith("Be terse.") + assert '"file"' in written and "single JSON object" in written + config = setup_config(setup) + assert config["instructions"] == ["/tmp/oc-1/instructions.md"] + assert config["skills"] == {"paths": ["/tmp/oc-1/skills"]} + assert build_instructions(make_ctx()) is None + assert "skills" not in setup_config(setup_for(make_ctx())) + + +def test_turn_request_argv_and_session_continuation(): + ctx = make_ctx(options=OpenCodeOptions(agent="build")) + setup = setup_for(ctx) + first = CONFIG.transform_turn_request(ctx, setup, PRIVATE, "hello", None) + assert list(first.argv) == [ + "opencode", + "run", + "--pure", + "--format", + "json", + "--thinking", + "-m", + "litellm/claude-haiku-4-5-20251001", + "--agent", + "build", + "--title", + OPENCODE_SESSION_TITLE, + ] + assert first.cwd == "/work" + assert first.stdin == "hello" + assert first.env == setup.env + second = CONFIG.transform_turn_request(ctx, setup, PRIVATE, "again", SESSION) + argv = list(second.argv) + assert argv[argv.index("--session") + 1] == SESSION + assert "--title" not in argv + assert "again" not in " ".join(argv) + + +def test_turn_prompt_repeats_schema_when_output_set(): + assert turn_prompt(make_ctx(), "hi") == "hi" + prompt = turn_prompt(make_ctx(output=Answer), "hi") + assert prompt.startswith("hi\n\n") and '"file"' in prompt + ctx = make_ctx(output=Answer) + request = CONFIG.transform_turn_request(ctx, setup_for(ctx), PRIVATE, "hi", None) + assert request.stdin == prompt + + +def test_turn_response_paths(): + state = OpenCodeStreamState(final_text='Here: {"file": "a", "content": "hi"}') + ok = CONFIG.transform_turn_response(make_ctx(output=Answer), state, 0, []) + assert json.loads(ok.output_json) == {"file": "a", "content": "hi"} + plain = CONFIG.transform_turn_response(make_ctx(), state, 0, []) + assert plain.output_json is None and plain.final_text == state.final_text + with pytest.raises(HarnessTurnError, match="boom"): + CONFIG.transform_turn_response( + make_ctx(), OpenCodeStreamState(error="boom"), 0, [] + ) + with pytest.raises(HarnessTurnError, match="code 3: no output"): + CONFIG.transform_turn_response(make_ctx(), OpenCodeStreamState(), 3, []) + + +# --------------------------------------------------------------------------- handler + + +async def test_start_writes_token_only_in_private_file(): + handler, ctx, sandbox = await started() + assert sandbox.files["/tmp/oc-1/token"] == TOKEN.encode() + sandbox.outputs.append(fixture_proc("turn1_write_read.jsonl")) + await collect(handler, ctx, "create hello.txt containing hi then read it") + call = sandbox.execs[0] + assert TOKEN not in json.dumps(call["cmd"]) + assert TOKEN not in json.dumps(call["env"]) + config = exec_config(sandbox) + assert config["provider"]["litellm"]["options"] == { + "baseURL": "http://host.docker.internal:4555/v1", + "apiKey": "{file:/tmp/oc-1/token}", + } + assert config["permission"] == {"*": "allow"} + + +async def test_start_persists_xdg_dir(): + _, _, sandbox = await started() + assert sandbox.runs == [ + ["sh", "-c", PERSIST_DIR_SCRIPT, "sh", "/tmp/oc-1/xdg", "opencode"] + ] + + +async def test_persist_failure_still_uses_private_xdg(): + handler, ctx, sandbox = await started(FakeSandbox(persist_ok=False)) + sandbox.outputs.append(fixture_proc("turn1_write_read.jsonl")) + await collect(handler, ctx, "x") + assert sandbox.execs[0]["env"]["XDG_DATA_HOME"] == "/tmp/oc-1/xdg/data" + + +async def test_turn_argv_env_and_session_continuation(): + handler, ctx, sandbox = await started( + options=OpenCodeOptions(agent="build", env={"FOO": "1"}) + ) + sandbox.outputs.append(fixture_proc("turn1_write_read.jsonl")) + events = await collect(handler, ctx, "create hello.txt containing hi then read it") + first = sandbox.execs[0] + assert first["cmd"] == [ + "opencode", + "run", + "--pure", + "--format", + "json", + "--thinking", + "-m", + "litellm/claude-haiku-4-5-20251001", + "--agent", + "build", + "--title", + OPENCODE_SESSION_TITLE, + ] + assert first["cwd"] == "/work" + env = first["env"] + assert env["XDG_CONFIG_HOME"] == "/tmp/oc-1/xdg/config" + assert env["XDG_DATA_HOME"] == "/tmp/oc-1/xdg/data" + assert env["XDG_STATE_HOME"] == "/tmp/oc-1/xdg/state" + assert env["XDG_CACHE_HOME"] == "/tmp/oc-1/xdg/cache" + assert env["OPENCODE_DISABLE_AUTOUPDATE"] == "1" + assert env["OPENCODE_DISABLE_MODELS_FETCH"] == "1" + assert env["OPENCODE_CONFIG"] == "" and env["OPENCODE_PERMISSION"] == "" + assert env["FOO"] == "1" + assert any(isinstance(e, ToolCall) for e in events) + assert ctx.final_text.startswith("Done!") + assert handler.native_session_id() == SESSION + + sandbox.outputs.append(fixture_proc("turn2_session_skill.jsonl")) + await collect(handler, ctx, "what file did you create?") + second = sandbox.execs[1]["cmd"] + assert second[second.index("--session") + 1] == SESSION + assert "--title" not in second + assert "what file" not in " ".join(second) + + +async def test_prompt_is_sent_on_stdin_not_argv(): + handler, ctx, sandbox = await started() + proc = fixture_proc("turn1_write_read.jsonl") + sandbox.outputs.append(proc) + await collect(handler, ctx, "secret prompt text") + assert proc.stdin.data == b"secret prompt text" and proc.stdin.closed + assert "secret prompt text" not in sandbox.execs[0]["cmd"] + + +async def test_resume_sets_session(): + handler, ctx, sandbox = await started() + await handler.resume(ctx, "ses_prev") + sandbox.outputs.append(fixture_proc("turn2_session_skill.jsonl")) + await collect(handler, ctx, "hi") + cmd = sandbox.execs[0]["cmd"] + assert cmd[cmd.index("--session") + 1] == "ses_prev" + + +async def test_read_only_and_disable_tools_config(): + handler, ctx, sandbox = await started( + permissions="read-only", disable_tools=["grep"] + ) + sandbox.outputs.append(fixture_proc("readonly_denied_bash.jsonl")) + await collect(handler, ctx, "x") + assert exec_config(sandbox)["permission"] == { + "edit": "deny", + "bash": "deny", + "webfetch": "deny", + "grep": "deny", + } + + +async def test_instructions_and_structured_output(): + handler, ctx, sandbox = await started(instructions="Be terse.", output=Answer) + written = sandbox.files["/tmp/oc-1/instructions.md"].decode() + assert written.startswith("Be terse.") + assert '"file"' in written and "single JSON object" in written + lines = [ + {"type": "step_start", "sessionID": "s"}, + { + "type": "text", + "sessionID": "s", + "part": {"text": 'Here: {"file": "a", "content": "hi"}'}, + }, + ] + proc = FakeProcess("\n".join(json.dumps(line) for line in lines).encode()) + sandbox.outputs.append(proc) + await collect(handler, ctx, "x") + assert exec_config(sandbox)["instructions"] == ["/tmp/oc-1/instructions.md"] + assert '"file"' in proc.stdin.data.decode() + assert json.loads(ctx.output_json) == {"file": "a", "content": "hi"} + + +async def test_skills_copied_to_private_skills_path(tmp_path): + skill = tmp_path / "greeter" + (skill / "ref").mkdir(parents=True) + (skill / "SKILL.md").write_text("---\nname: greeter\ndescription: d\n---\nbody") + (skill / "ref" / "notes.txt").write_text("n") + handler, ctx, sandbox = await started(skills=[str(skill)]) + assert sandbox.files["/tmp/oc-1/skills/greeter/SKILL.md"].startswith(b"---") + assert sandbox.files["/tmp/oc-1/skills/greeter/ref/notes.txt"] == b"n" + sandbox.outputs.append(fixture_proc("turn2_session_skill.jsonl")) + await collect(handler, ctx, "x") + assert exec_config(sandbox)["skills"] == {"paths": ["/tmp/oc-1/skills"]} + + +async def test_missing_binary(): + with pytest.raises(HarnessInstallFailed, match="opencode"): + await started(FakeSandbox(has_binary=False)) + + +async def test_wrong_options_and_managed_config_rejected(): + with pytest.raises(OptionsMismatch): + await started(options=CodexOptions()) + with pytest.raises(OptionsMismatch): + await started(options=OpenCodeOptions(config={"model": "openai/x"})) + + +async def test_turn_before_start_raises(): + handler = CLIHarnessHandler(OpenCodeHarnessConfig()) + with pytest.raises(RuntimeError, match="before start"): + await collect(handler, make_ctx(), "x") + + +async def test_api_error_event_raises_even_on_exit_zero(): + handler, ctx, sandbox = await started() + sandbox.outputs.append(fixture_proc("api_error.jsonl")) + with pytest.raises(HarnessTurnError, match="no healthy deployments"): + await collect(handler, ctx, "x") + + +async def test_nonzero_exit_raises_with_stderr_tail(): + handler, ctx, sandbox = await started() + sandbox.outputs.append( + FakeProcess(b"", stderr=b"line1\nfatal: bad config\n", exit_code=2) + ) + with pytest.raises(HarnessTurnError, match="code 2: line1\nfatal: bad config"): + await collect(handler, ctx, "x") + + +async def test_early_close_kills_process_and_stop_is_idempotent(): + handler, ctx, sandbox = await started() + proc = fixture_proc("turn1_write_read.jsonl") + sandbox.outputs.append(proc) + gen = handler.turn(ctx, "x") + await gen.__anext__() + await gen.aclose() + assert proc.killed + await handler.stop(ctx) + await handler.stop(ctx) + + +async def test_model_falls_back_to_endpoint_model(): + handler, ctx, sandbox = await started( + model=None, endpoint=FakeEndpoint(model="gw-model") + ) + sandbox.outputs.append(fixture_proc("turn1_write_read.jsonl")) + await collect(handler, ctx, "x") + assert "litellm/gw-model" in sandbox.execs[0]["cmd"] + assert exec_config(sandbox)["provider"]["litellm"]["models"] == {"gw-model": {}} + + +def test_turn_request_never_loads_plugins(): + """A repo's .opencode/plugin/*.js would run as the host user at startup; --pure blocks it.""" + ctx = make_ctx() + argv = list(CONFIG.transform_turn_request(ctx, setup_for(ctx), PRIVATE, "hi", None).argv) + assert argv[:3] == ["opencode", "run", "--pure"] + + +def test_options_config_cannot_add_plugins(): + with pytest.raises(OptionsMismatch, match="plugin"): + validate_user_config({"plugin": ["./evil.js"]}) + + +def test_endpoint_request_fixture_documents_contract(): + requests = load_fixture("endpoint_requests.jsonl") + assert {r["path"] for r in requests} == {"/v1/chat/completions"} + assert all(r["stream"] is True for r in requests) + assert all(r["stream_options"] == {"include_usage": True} for r in requests) diff --git a/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 0d0c65e3650..02ac1540071 100644 --- a/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/unit/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -8201,7 +8201,7 @@ class TestGatewaySessionAdmission: assert not any(k.lower() == "authorization" for k in (raw_headers or {})) -def _make_team(team_id, mcp_servers, *, org_id=None, tool_perms=None, members=("sso-user",)): +def _make_team(team_id, mcp_servers, *, org_id=None, tool_perms=None, members=("sso-user",), toolsets=None): from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable, Member return LiteLLM_TeamTable( @@ -8210,7 +8210,10 @@ def _make_team(team_id, mcp_servers, *, org_id=None, tool_perms=None, members=(" members_with_roles=[Member(user_id=u, role="user") for u in members], access_group_ids=[], object_permission=LiteLLM_ObjectPermissionTable( - object_permission_id=f"op-{team_id}", mcp_servers=mcp_servers, mcp_tool_permissions=tool_perms + object_permission_id=f"op-{team_id}", + mcp_servers=mcp_servers, + mcp_tool_permissions=tool_perms, + mcp_toolsets=toolsets, ), ) @@ -8271,6 +8274,64 @@ class TestUserSubjectTeamUnion: result = await MCPRequestHandler.get_allowed_mcp_servers(auth) assert set(result) == {"srv1", "srv2", "srv3"} + async def test_toolsets_of_a_team_that_dropped_the_user_from_its_roster_are_not_granted(self): + """The user's cached team list still names team-revoked, but its live roster no longer lists + the user, so its toolset is withheld exactly as its servers are on the aggregate /mcp.""" + from litellm.proxy._experimental.mcp_server.ui_session_utils import granted_toolset_ids + + teams = { + "team-kept": _make_team("team-kept", [], toolsets=["ts-kept"]), + "team-revoked": _make_team("team-revoked", [], toolsets=["ts-revoked"], members=("someone-else",)), + } + auth = _make_admitted_subject("sso-user") + with self._patch(teams_by_id=teams, user_teams=["team-kept", "team-revoked"]): + granted = await granted_toolset_ids(auth) + assert granted == {"ts-kept"} + + async def test_a_pinned_toolset_narrows_every_source_to_the_toolset_servers_and_tools(self): + """On /toolset/{name}/mcp the admitted subject carries mcp_toolset_id; team-a's grant on srv1 and + srv2 with every tool collapses to the toolset's srv1 and its one tool, and team-b's srv3 drops.""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + + teams = {"team-a": _make_team("team-a", ["srv1", "srv2"]), "team-b": _make_team("team-b", ["srv3"])} + auth = _make_admitted_subject("sso-user") + pinned = auth.model_copy(update={"mcp_toolset_id": "ts-1"}) + resolve = AsyncMock(return_value={"srv1": ["add"]}) + with ( + self._patch(teams_by_id=teams, user_teams=["team-a", "team-b"]), + patch.object(global_mcp_server_manager, "resolve_toolset_tool_permissions", resolve), + ): + servers = await MCPRequestHandler.resolve_admitted_subject_servers(pinned) + tools = await MCPRequestHandler.resolve_admitted_subject_tools("srv1", pinned) + unpinned_servers = await MCPRequestHandler.resolve_admitted_subject_servers(auth) + unpinned_tools = await MCPRequestHandler.resolve_admitted_subject_tools("srv1", auth) + assert servers == ["srv1"] + assert tools == ["add"] + assert set(unpinned_servers) == {"srv1", "srv2", "srv3"} + assert unpinned_tools is None + assert {call.kwargs["toolset_ids"][0] for call in resolve.await_args_list} == {"ts-1"} + + async def test_a_fresh_policy_pinned_toolset_bypasses_the_toolset_permission_cache(self): + """A session admitted under requires_fresh_policy reads the pinned toolset from the writer, so a + tool revoked from the toolset is gone on the very next request (Devin Review 4150024092).""" + from litellm.proxy._experimental.mcp_server.mcp_server_manager import global_mcp_server_manager + + teams = {"team-a": _make_team("team-a", ["srv1", "srv2"])} + auth = _make_admitted_subject("sso-user") + auth.requires_fresh_policy = True + pinned = auth.model_copy(update={"mcp_toolset_id": "ts-1"}) + resolve = AsyncMock(return_value={"srv1": ["add"]}) + with ( + self._patch(teams_by_id=teams, user_teams=["team-a"]), + patch.object(global_mcp_server_manager, "resolve_toolset_tool_permissions", resolve), + ): + servers = await MCPRequestHandler.resolve_admitted_subject_servers(pinned) + tools = await MCPRequestHandler.resolve_admitted_subject_tools("srv1", pinned) + assert servers == ["srv1"] + assert tools == ["add"] + assert resolve.await_args_list + assert all(call.kwargs == {"toolset_ids": ["ts-1"], "requires_fresh_policy": True} for call in resolve.await_args_list) + async def test_key_based_caller_uses_single_team_only(self): """A key-based caller (api_key set) with a team_id sees ONLY that team, even though the same user belongs to other teams: key auth must be byte-identical to before.""" diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_discovery.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_discovery.py index 43cf35c152d..d35cb7234dc 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_discovery.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_discovery.py @@ -1,5 +1,6 @@ import json import os +from typing import Final import pytest @@ -95,10 +96,42 @@ class TestMCPRegistryFile: with open(registry_path, "r") as f: data = json.load(f) names = {s["name"] for s in data["servers"]} - expected = {"github", "slack", "postgresql", "snowflake", "atlassian"} + expected = {"github", "slack", "postgresql", "snowflake", "atlassian", "microsoft_365"} missing = expected - names assert not missing, f"Missing well-known servers: {missing}" + def test_microsoft_365_is_a_self_hosted_streamable_http_server(self, registry_path): + """The Graph server runs next to the proxy in org mode, so the entry must be streamable HTTP at /mcp.""" + with open(registry_path, "r") as f: + data = json.load(f) + entry: Final = next(s for s in data["servers"] if s["name"] == "microsoft_365") + assert entry["transport"] == "http" + assert entry["url"].endswith("/mcp") + assert entry["category"] == "Productivity" + assert "ms-365-mcp-server" in entry["registry_url"] + + def test_bundled_icons_exist(self, registry_path): + """An icon served from the proxy's own assets ships twice, as the built copy the wheel packages and as + the dashboard source copy every Docker image rebuilds from. Both must exist and match or a card goes blank.""" + with open(registry_path, "r") as f: + data = json.load(f) + proxy_dir: Final = os.path.dirname(registry_path) + built_logos_dir: Final = os.path.join(proxy_dir, "_experimental", "out", "assets", "logos") + source_logos_dir: Final = os.path.join( + proxy_dir, "..", "..", "ui", "litellm-dashboard", "public", "assets", "logos" + ) + bundled: Final = [s for s in data["servers"] if s.get("icon_url", "").startswith("/ui/assets/logos/")] + assert bundled, "at least one registry entry ships its own icon" + for server in bundled: + file_name: Final = os.path.basename(server["icon_url"]) + built: Final = os.path.join(built_logos_dir, file_name) + source: Final = os.path.join(source_logos_dir, file_name) + assert os.path.isfile(built), f"{server['name']}: {server['icon_url']} missing from the built dashboard" + assert os.path.isfile(source), f"{server['name']}: {server['icon_url']} missing from the dashboard source" + with open(built, "rb") as built_file, open(source, "rb") as source_file: + same_bytes: Final = built_file.read() == source_file.read() + assert same_bytes, f"{server['name']}: built and source copies of {file_name} differ" + def test_env_vars_structure(self, registry_path): with open(registry_path, "r") as f: data = json.load(f) diff --git a/tests/unit/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py b/tests/unit/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py index 1398884783e..95be2b8b12b 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_mcp_toolset_scope.py @@ -1,10 +1,12 @@ """Tests for MCP toolset scope enforcement.""" import asyncio +from collections.abc import Awaitable, Callable from typing import Dict, List, Optional from unittest.mock import AsyncMock, MagicMock, patch import pytest +from fastapi import HTTPException from litellm.proxy._types import ( LiteLLM_ObjectPermissionTable, @@ -30,6 +32,19 @@ def _make_auth( ) +def _granted_through_team(*team_toolset_ids: str) -> Callable[[UserAPIKeyAuth], Awaitable[frozenset[str]]]: + """The real grant resolver over a team that holds ``team_toolset_ids``, with no key access rule.""" + from litellm.proxy._experimental.mcp_server.ui_session_utils import granted_toolset_ids + + async def team_permission(context: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable: + return LiteLLM_ObjectPermissionTable(object_permission_id="team-op", mcp_toolsets=list(team_toolset_ids)) + + async def granted(context: UserAPIKeyAuth) -> frozenset[str]: + return await granted_toolset_ids(context, team_object_permission=team_permission, require_key_access=False) + + return granted + + class TestApplyToolsetScope: """Tests for _apply_toolset_scope helper.""" @@ -97,6 +112,122 @@ class TestApplyToolsetScope: assert op.mcp_servers == ["server-a"] assert op.mcp_tool_permissions == toolset_perms + @pytest.mark.asyncio + async def test_team_granted_toolset_is_served_to_a_key_without_its_own_grant(self): + """A team key whose own row carries no toolset grant is admitted to the toolset its team + holds (LIT-6029), scoped to that toolset's servers and tools.""" + from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + + toolset_perms = {"server-a": ["tool1"]} + auth = UserAPIKeyAuth(api_key="sk-test", team_id="team-a", object_permission=None) + with patch( + "litellm.proxy._experimental.mcp_server.server." + "global_mcp_server_manager.resolve_toolset_tool_permissions", + new=AsyncMock(return_value=toolset_perms), + ): + result = await _apply_toolset_scope(auth, "toolset-123", granted=_granted_through_team("toolset-123")) + + assert result.mcp_toolset_id == "toolset-123" + assert result.object_permission is not None + assert result.object_permission.mcp_servers == ["server-a"] + assert result.object_permission.mcp_tool_permissions == toolset_perms + + @pytest.mark.asyncio + async def test_a_non_admin_dashboard_session_is_pinned_as_its_admitted_user_instead_of_rewritten(self): + """The dashboard session acts as its admitted user, whose team grants resolve per source, so a + team-granted toolset is not capped by the user's own row: the row stays intact and the toolset + rides along as mcp_toolset_id (LIT-6029).""" + from litellm.constants import UI_SESSION_TOKEN_TEAM_ID + from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + + session = UserAPIKeyAuth(team_id=UI_SESSION_TOKEN_TEAM_ID, user_id="user-1") + own_row = LiteLLM_ObjectPermissionTable(object_permission_id="user-op", mcp_servers=["server-own"]) + admitted = UserAPIKeyAuth(user_id="user-1", object_permission=own_row) + admitted.mcp_admitted_user_subject = True + granted = AsyncMock(return_value=frozenset({"toolset-123"})) + resolve = AsyncMock(return_value={"server-team": ["tool1"]}) + with patch( + "litellm.proxy._experimental.mcp_server.server." + "global_mcp_server_manager.resolve_toolset_tool_permissions", + new=resolve, + ): + result = await _apply_toolset_scope( + session, "toolset-123", acting_user=AsyncMock(return_value=admitted), granted=granted + ) + + assert granted.await_args is not None and granted.await_args.args[0].mcp_admitted_user_subject is True + assert result.mcp_admitted_user_subject is True + assert result.mcp_toolset_id == "toolset-123" + assert result.object_permission == own_row + resolve.assert_not_awaited() + + @pytest.mark.asyncio + async def test_a_gateway_admitted_user_without_the_toolset_in_any_source_is_denied(self): + from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + + admitted = UserAPIKeyAuth(user_id="user-1", object_permission=None) + admitted.mcp_admitted_user_subject = True + granted = AsyncMock(return_value=frozenset({"toolset-other"})) + with pytest.raises(HTTPException) as exc_info: + await _apply_toolset_scope(admitted, "toolset-123", granted=granted) + + assert exc_info.value.status_code == 403 + granted.assert_awaited_once_with(admitted) + + @pytest.mark.asyncio + async def test_a_resource_scoped_admitted_user_is_denied_a_team_toolset_on_another_server(self): + """A gateway bearer scoped to server-own (RFC 8707 resource) cannot open a team toolset whose + servers lie outside that resource, even though the team grants it (Devin Review 4150024267).""" + from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + + admitted = UserAPIKeyAuth(user_id="user-1", object_permission=None) + admitted.mcp_admitted_user_subject = True + admitted.mcp_session_resource_server_id = "server-own" + admitted.requires_fresh_policy = True + granted = AsyncMock(return_value=frozenset({"toolset-123"})) + resolve = AsyncMock(return_value={"server-team": ["tool1"]}) + with patch( + "litellm.proxy._experimental.mcp_server.server." + "global_mcp_server_manager.resolve_toolset_tool_permissions", + new=resolve, + ): + with pytest.raises(HTTPException) as exc_info: + await _apply_toolset_scope(admitted, "toolset-123", granted=granted) + + assert exc_info.value.status_code == 403 + resolve.assert_awaited_once_with(toolset_ids=["toolset-123"], requires_fresh_policy=True) + + @pytest.mark.asyncio + async def test_a_resource_scoped_admitted_user_opens_a_toolset_inside_its_resource(self): + from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + + admitted = UserAPIKeyAuth(user_id="user-1", object_permission=None) + admitted.mcp_admitted_user_subject = True + admitted.mcp_session_resource_server_id = "server-team" + granted = AsyncMock(return_value=frozenset({"toolset-123"})) + resolve = AsyncMock(return_value={"server-team": ["tool1"], "server-other": ["tool2"]}) + with patch( + "litellm.proxy._experimental.mcp_server.server." + "global_mcp_server_manager.resolve_toolset_tool_permissions", + new=resolve, + ): + result = await _apply_toolset_scope(admitted, "toolset-123", granted=granted) + + assert result.mcp_toolset_id == "toolset-123" + assert result.mcp_session_resource_server_id == "server-team" + resolve.assert_awaited_once_with(toolset_ids=["toolset-123"], requires_fresh_policy=False) + + @pytest.mark.asyncio + async def test_team_grant_for_another_toolset_does_not_admit_a_key_to_this_one(self): + from litellm.proxy._experimental.mcp_server.server import _apply_toolset_scope + + auth = _make_auth(mcp_toolsets=[]) + auth.team_id = "team-a" + with pytest.raises(HTTPException) as exc_info: + await _apply_toolset_scope(auth, "toolset-123", granted=_granted_through_team("toolset-other")) + + assert exc_info.value.status_code == 403 + @pytest.mark.asyncio async def test_non_admin_no_object_permission_raises_403(self): """Non-admin key with object_permission=None is denied (no grants configured).""" @@ -250,6 +381,131 @@ class TestFetchMCPToolsetsAccess: assert len(result) == 2 mock_list.assert_called_once_with(mock_client, toolset_ids=["ts-1", "ts-2"]) + @pytest.mark.asyncio + async def test_team_granted_toolsets_are_listed_for_a_key_without_its_own_grant(self): + """GET /v1/mcp/toolset for a team key lists the team's toolsets (LIT-6029).""" + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + fetch_mcp_toolsets, + ) + + auth = UserAPIKeyAuth(api_key="sk-test", team_id="team-a", object_permission=None) + team_permission = LiteLLM_ObjectPermissionTable(object_permission_id="team-op", mcp_toolsets=["ts-team"]) + fake_toolsets = [MagicMock(toolset_id="ts-team")] + mock_client = MagicMock() + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=mock_client, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.list_mcp_toolsets", + new=AsyncMock(return_value=fake_toolsets), + ) as mock_list, + patch.object( + MCPRequestHandler, + "_get_team_object_permission", + new=AsyncMock(return_value=team_permission), + ), + ): + result = await fetch_mcp_toolsets(user_api_key_dict=auth) + + assert result == fake_toolsets + mock_list.assert_called_once_with(mock_client, toolset_ids=["ts-team"]) + + @pytest.mark.asyncio + async def test_admin_with_own_grants_is_not_narrowed_by_a_team_lookup(self): + """An admin's own grant list is the only filter; no team lookup runs for admins.""" + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + fetch_mcp_toolsets, + ) + + auth = _make_auth(mcp_toolsets=["ts-1"]) + auth.user_role = LitellmUserRoles.PROXY_ADMIN + mock_client = MagicMock() + own_toolsets = [{"toolset_id": "ts-1", "toolset_name": "own"}] + + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=mock_client, + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.list_mcp_toolsets", + new=AsyncMock(return_value=own_toolsets), + ) as mock_list, + patch.object( + MCPRequestHandler, "_get_team_object_permission", new=AsyncMock(return_value=None) + ) as team_lookup, + ): + result = await fetch_mcp_toolsets(user_api_key_dict=auth) + + assert result == own_toolsets + mock_list.assert_called_once_with(mock_client, toolset_ids=["ts-1"]) + team_lookup.assert_not_awaited() + + +class TestFetchMCPToolsetAccess: + """Tests for GET /v1/mcp/toolset/{toolset_id} access control.""" + + @staticmethod + async def _fetch(auth: UserAPIKeyAuth, toolset_id: str, team_toolsets: list[str] | None): + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + fetch_mcp_toolset, + ) + + team_permission = ( + LiteLLM_ObjectPermissionTable(object_permission_id="team-op", mcp_toolsets=team_toolsets) + if team_toolsets is not None + else None + ) + toolset = MagicMock(toolset_id=toolset_id) + with ( + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_prisma_client_or_throw", + return_value=MagicMock(), + ), + patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.get_mcp_toolset", + new=AsyncMock(return_value=toolset), + ), + patch.object( + MCPRequestHandler, + "_get_team_object_permission", + new=AsyncMock(return_value=team_permission), + ), + ): + return await fetch_mcp_toolset(toolset_id=toolset_id, user_api_key_dict=auth) + + @pytest.mark.asyncio + async def test_team_granted_toolset_detail_is_served_to_a_key_without_its_own_grant(self): + auth = UserAPIKeyAuth(api_key="sk-test", team_id="team-a", object_permission=None) + + toolset = await self._fetch(auth, "ts-team", team_toolsets=["ts-team"]) + + assert toolset.toolset_id == "ts-team" + + @pytest.mark.asyncio + async def test_toolset_detail_stays_forbidden_when_neither_key_nor_team_holds_it(self): + from fastapi import HTTPException + + auth = _make_auth(mcp_toolsets=["ts-own"]) + auth.team_id = "team-a" + + with pytest.raises(HTTPException) as exc_info: + await self._fetch(auth, "ts-withheld", team_toolsets=["ts-team"]) + + assert exc_info.value.status_code == 403 + class TestToolsetPrefixResolution: """Regression for LIT-3419. diff --git a/tests/unit/proxy/_experimental/mcp_server/test_ui_session_utils.py b/tests/unit/proxy/_experimental/mcp_server/test_ui_session_utils.py index 816ccc5e7e6..293d9443ced 100644 --- a/tests/unit/proxy/_experimental/mcp_server/test_ui_session_utils.py +++ b/tests/unit/proxy/_experimental/mcp_server/test_ui_session_utils.py @@ -6,13 +6,14 @@ import pytest from fastapi import HTTPException from litellm.constants import UI_SESSION_TOKEN_TEAM_ID -from litellm.proxy._types import UserAPIKeyAuth - from litellm.proxy._experimental.mcp_server.ui_session_utils import ( build_effective_auth_contexts, clone_user_api_key_auth_with_team, + granted_toolset_ids, + toolset_grant_contexts, resolve_ui_session_team_ids, ) +from litellm.proxy._types import LiteLLM_ObjectPermissionTable, UserAPIKeyAuth def test_clone_user_api_key_auth_with_team_creates_independent_copy(): @@ -258,3 +259,234 @@ async def test_admitted_user_context_carries_the_request_span(monkeypatch): assert (await acting_user_auth(user_auth)).parent_otel_span is parent_span assert (await build_effective_auth_contexts(user_auth))[-1].parent_otel_span is parent_span + + +def _toolset_permission(*toolset_ids: str) -> LiteLLM_ObjectPermissionTable: + return LiteLLM_ObjectPermissionTable( + object_permission_id=f"op-{'-'.join(toolset_ids)}", mcp_toolsets=list(toolset_ids) + ) + + +@pytest.mark.asyncio +async def test_granted_toolset_ids_unions_own_and_team_grants_over_every_effective_context(): + """A dashboard session of a user in two teams holds the toolsets of both teams plus the ones on + the user row itself, exactly the grant sources the aggregate /mcp listing expands.""" + session = UserAPIKeyAuth(team_id=UI_SESSION_TOKEN_TEAM_ID, user_id="user-1") + team_a = UserAPIKeyAuth(team_id="team-a", user_id="user-1") + team_b = UserAPIKeyAuth(team_id="team-b", user_id="user-1", object_permission=_toolset_permission()) + admitted = UserAPIKeyAuth(user_id="user-1", object_permission=_toolset_permission("ts-user")) + team_grants = {"team-a": _toolset_permission("ts-a", "ts-shared"), "team-b": _toolset_permission("ts-b")} + + async def effective_contexts(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]: + assert auth is session + return [team_a, team_b, admitted] + + async def team_permission(auth: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable | None: + return team_grants.get(auth.team_id or "") + + granted = await granted_toolset_ids(session, effective_contexts, team_permission) + + assert granted == frozenset({"ts-a", "ts-shared", "ts-b", "ts-user"}) + + +@pytest.mark.asyncio +async def test_granted_toolset_ids_is_empty_when_neither_key_nor_team_grants_a_toolset(): + key = UserAPIKeyAuth(api_key="sk-test", team_id="team-a", object_permission=_toolset_permission()) + + async def effective_contexts(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]: + return [auth] + + async def no_team_permission(auth: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable | None: + return None + + assert await granted_toolset_ids(key, effective_contexts, no_team_permission) == frozenset() + + +async def _same_context(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]: + return [auth] + + +async def _team_grants_ts_team(auth: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable | None: + return _toolset_permission("ts-team") + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "own", + [ + LiteLLM_ObjectPermissionTable(object_permission_id="op", mcp_toolsets=["ts-own"]), + LiteLLM_ObjectPermissionTable(object_permission_id="op", mcp_servers=["srv-own"]), + LiteLLM_ObjectPermissionTable(object_permission_id="op", mcp_tool_permissions={"srv-own": ["add"]}), + LiteLLM_ObjectPermissionTable(object_permission_id="op", mcp_access_groups=["group-own"]), + ], +) +async def test_a_key_declaring_its_own_mcp_grant_does_not_inherit_the_team_toolsets(own): + """The key/team rule of the aggregate listing: a key's own MCP grant is a ceiling the team cannot widen.""" + key = UserAPIKeyAuth(api_key="sk-test", team_id="team-a", object_permission=own) + + granted = await granted_toolset_ids(key, _same_context, _team_grants_ts_team, require_key_access=False) + + assert granted == frozenset(own.mcp_toolsets or ()) + + +@pytest.mark.asyncio +async def test_require_key_mcp_access_defined_stops_a_key_inheriting_team_toolsets_but_not_a_session(): + key = UserAPIKeyAuth(api_key="sk-test", team_id="team-a") + session = UserAPIKeyAuth(team_id=UI_SESSION_TOKEN_TEAM_ID, user_id="user-1") + + assert await granted_toolset_ids(key, _same_context, _team_grants_ts_team, require_key_access=False) == {"ts-team"} + assert await granted_toolset_ids(key, _same_context, _team_grants_ts_team, require_key_access=True) == frozenset() + assert await granted_toolset_ids(session, _same_context, _team_grants_ts_team, require_key_access=True) == { + "ts-team" + } + + +@pytest.mark.asyncio +async def test_toolset_grant_contexts_of_a_virtual_key_is_the_key_alone(): + key = UserAPIKeyAuth(api_key="sk-test", team_id="team-a") + + async def never(auth: UserAPIKeyAuth) -> None: + raise AssertionError("a virtual key has no admitted sources") + + assert await toolset_grant_contexts(key, admitted_context=never, admitted_sources=never) == (key,) + + +def _admitted(user_id: str, own: LiteLLM_ObjectPermissionTable | None = None) -> UserAPIKeyAuth: + subject = UserAPIKeyAuth(user_id=user_id, object_permission=own) + subject.mcp_admitted_user_subject = True + return subject + + +@pytest.mark.asyncio +async def test_toolset_grant_contexts_of_a_dashboard_session_are_its_admitted_users_grant_sources(): + """The dashboard session fans out through the same roster-checked source builder as the aggregate + /mcp resolution, applied to the admitted user it acts as, so a cached membership a team has since + revoked never reaches the toolset check.""" + session = UserAPIKeyAuth(team_id=UI_SESSION_TOKEN_TEAM_ID, user_id="user-1") + admitted = _admitted("user-1") + own_source = UserAPIKeyAuth(user_id="user-1") + team_source = UserAPIKeyAuth(user_id="user-1", team_id="team-a") + + async def admitted_context(auth: UserAPIKeyAuth) -> UserAPIKeyAuth: + assert auth is session + return admitted + + async def admitted_sources(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]: + assert auth is admitted + return [own_source, team_source] + + assert await toolset_grant_contexts(session, admitted_context, admitted_sources) == (own_source, team_source) + + +@pytest.mark.asyncio +async def test_toolset_grant_contexts_of_a_gateway_admitted_user_are_its_own_grant_sources(): + admitted = _admitted("user-1") + team_source = UserAPIKeyAuth(user_id="user-1", team_id="team-a") + + async def no_dashboard_context(auth: UserAPIKeyAuth) -> None: + return None + + async def admitted_sources(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]: + assert auth is admitted + return [team_source] + + assert await toolset_grant_contexts(admitted, no_dashboard_context, admitted_sources) == (team_source,) + + +@pytest.mark.asyncio +async def test_a_source_declaring_its_own_mcp_grant_never_reads_its_team(): + key = UserAPIKeyAuth(api_key="sk-test", team_id="team-a", object_permission=_toolset_permission("ts-own")) + team_reads: list[str | None] = [] # mutable-ok: records the lookups the code under test performs + + async def team_permission(auth: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable | None: + team_reads.append(auth.team_id) + return _toolset_permission("ts-team") + + granted = await granted_toolset_ids(key, _same_context, team_permission, require_key_access=False) + + assert granted == {"ts-own"} + assert team_reads == [] + + +async def _team_a_unreadable(auth: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable | None: + if auth.team_id == "team-a": + raise RuntimeError("team row unreadable") + return _toolset_permission("ts-b") + + +@pytest.mark.asyncio +async def test_an_unreadable_team_grants_nothing_while_the_direct_and_other_team_grants_still_count(): + """A dashboard user whose own row grants ts-user and who sits on team-a and team-b keeps ts-user and + ts-b when team-a cannot be read; team-a itself contributes nothing rather than failing the lookup.""" + admitted = _admitted("user-1", _toolset_permission("ts-user")) + team_a = UserAPIKeyAuth(user_id="user-1", team_id="team-a") + team_b = UserAPIKeyAuth(user_id="user-1", team_id="team-b") + + async def sources(auth: UserAPIKeyAuth) -> list[UserAPIKeyAuth]: + return [admitted, team_a, team_b] + + assert await granted_toolset_ids(admitted, sources, _team_a_unreadable) == {"ts-user", "ts-b"} + + +@pytest.mark.asyncio +async def test_a_key_whose_only_grant_source_is_an_unreadable_team_is_granted_nothing(): + key = UserAPIKeyAuth(api_key="sk-test", team_id="team-a") + + assert await granted_toolset_ids(key, _same_context, _team_a_unreadable, require_key_access=False) == frozenset() + + +async def _hydrates_op_key_to_srv_own(auth: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable | None: + if auth.object_permission is not None: + return auth.object_permission + if auth.object_permission_id == "op-key": + return LiteLLM_ObjectPermissionTable(object_permission_id="op-key", mcp_servers=["srv-own"]) + return None + + +@pytest.mark.asyncio +async def test_a_key_cached_with_its_own_grant_unhydrated_is_scoped_to_that_grant_not_its_team(): + """The main auth flow can cache a key with object_permission_id set and object_permission None. The + row it names is the key's ceiling, so it is loaded and read as the key's own grant instead of letting the + key inherit its team's toolsets.""" + key = UserAPIKeyAuth(api_key="sk-test", team_id="team-a", object_permission_id="op-key") + + granted = await granted_toolset_ids( + key, + _same_context, + _team_grants_ts_team, + require_key_access=False, + own_object_permission=_hydrates_op_key_to_srv_own, + ) + + assert granted == frozenset() + + +async def _own_row_unreadable(auth: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable | None: + raise RuntimeError("object permission row unreadable") + + +async def _own_row_gone(auth: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable | None: + return None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("load_own", [_own_row_unreadable, _own_row_gone]) +async def test_a_key_naming_an_own_grant_that_cannot_be_read_is_granted_nothing_rather_than_its_team(load_own): + key = UserAPIKeyAuth(api_key="sk-test", team_id="team-a", object_permission_id="op-key") + + granted = await granted_toolset_ids( + key, _same_context, _team_grants_ts_team, require_key_access=False, own_object_permission=load_own + ) + + assert granted == frozenset() + + +@pytest.mark.asyncio +async def test_a_key_naming_no_own_grant_is_not_hydrated_before_inheriting_its_team(): + key = UserAPIKeyAuth(api_key="sk-test", team_id="team-a") + + granted = await granted_toolset_ids( + key, _same_context, _team_grants_ts_team, require_key_access=False, own_object_permission=_own_row_unreadable + ) + + assert granted == {"ts-team"} diff --git a/tests/unit/proxy/auth/test_route_checks.py b/tests/unit/proxy/auth/test_route_checks.py index d8ee58a52ea..6503e5bf50e 100644 --- a/tests/unit/proxy/auth/test_route_checks.py +++ b/tests/unit/proxy/auth/test_route_checks.py @@ -16,6 +16,93 @@ from litellm.proxy._types import ( from litellm.proxy.auth.auth_checks_organization import _user_is_org_admin from litellm.proxy.auth.route_checks import RouteChecks +DAILY_ACTIVITY_ROUTE_PAIRS: Final[tuple[tuple[str, str], ...]] = ( + ("/user/daily/activity", "/user/daily/activity/aggregated"), + ("/user/daily/activity", "/user/daily/activity/aggregated/keys"), + ("/user/daily/activity", "/user/daily/activity/aggregated/search"), + ("/user/daily/activity", "/user/daily/activity/aggregated/model_top_keys"), + ("/user/daily/activity", "/user/daily/activity/export"), + ("/user/daily/activity", "/user/daily/activity/aggregated/cache_leakage_keys"), + ("/team/daily/activity", "/team/daily/activity/aggregated"), + ("/team/daily/activity", "/team/daily/activity/aggregated/keys"), + ("/team/daily/activity", "/team/daily/activity/aggregated/search"), + ("/team/daily/activity", "/team/daily/activity/aggregated/model_top_keys"), + ("/team/daily/activity", "/team/daily/activity/export"), + ("/tag/daily/activity", "/tag/daily/activity/aggregated"), + ("/tag/daily/activity", "/tag/daily/activity/aggregated/keys"), + ("/tag/daily/activity", "/tag/daily/activity/aggregated/search"), + ("/tag/daily/activity", "/tag/daily/activity/aggregated/model_top_keys"), + ("/tag/daily/activity", "/tag/daily/activity/export"), + ("/organization/daily/activity", "/organization/daily/activity/aggregated"), + ("/organization/daily/activity", "/organization/daily/activity/aggregated/keys"), + ("/organization/daily/activity", "/organization/daily/activity/aggregated/search"), + ("/organization/daily/activity", "/organization/daily/activity/aggregated/model_top_keys"), + ("/organization/daily/activity", "/organization/daily/activity/export"), + ("/customer/daily/activity", "/customer/daily/activity/aggregated"), + ("/customer/daily/activity", "/customer/daily/activity/aggregated/keys"), + ("/customer/daily/activity", "/customer/daily/activity/aggregated/search"), + ("/customer/daily/activity", "/customer/daily/activity/aggregated/model_top_keys"), + ("/customer/daily/activity", "/customer/daily/activity/export"), + ("/customer/daily/activity", "/end_user/daily/activity/aggregated"), + ("/customer/daily/activity", "/end_user/daily/activity/aggregated/keys"), + ("/customer/daily/activity", "/end_user/daily/activity/aggregated/search"), + ("/customer/daily/activity", "/end_user/daily/activity/aggregated/model_top_keys"), + ("/customer/daily/activity", "/end_user/daily/activity/export"), + ("/agent/daily/activity", "/agent/daily/activity/aggregated"), + ("/agent/daily/activity", "/agent/daily/activity/aggregated/keys"), + ("/agent/daily/activity", "/agent/daily/activity/aggregated/search"), + ("/agent/daily/activity", "/agent/daily/activity/aggregated/model_top_keys"), + ("/agent/daily/activity", "/agent/daily/activity/export"), +) + +DAILY_ACTIVITY_ROLES: Final[tuple[LitellmUserRoles, ...]] = ( + LitellmUserRoles.PROXY_ADMIN, + LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, + LitellmUserRoles.INTERNAL_USER, + LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, + LitellmUserRoles.ORG_ADMIN, + LitellmUserRoles.TEAM, + LitellmUserRoles.CUSTOMER, +) + + +def _daily_activity_route_outcome(route: str, user_role: LitellmUserRoles) -> str: + if user_role == LitellmUserRoles.PROXY_ADMIN: + return "allowed" + user_obj = LiteLLM_UserTable( + user_id="test_user", + user_email="test@example.com", + user_role=user_role.value, + ) + valid_token = UserAPIKeyAuth(user_id="test_user", user_role=user_role) + request = MagicMock(spec=Request) + request.method = "GET" + request.query_params = {} + try: + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=user_obj, + _user_role=user_role.value, + route=route, + request=request, + valid_token=valid_token, + request_data={}, + ) + except HTTPException as exc: + return f"denied:{exc.status_code}" + except Exception as exc: + return f"denied:{type(exc).__name__}" + return "allowed" + + +@pytest.mark.parametrize(("existing_path", "new_path"), DAILY_ACTIVITY_ROUTE_PAIRS) +@pytest.mark.parametrize("user_role", DAILY_ACTIVITY_ROLES) +def test_daily_activity_routes_preserve_route_access_outcomes( + existing_path: str, new_path: str, user_role: LitellmUserRoles +) -> None: + assert _daily_activity_route_outcome(new_path, user_role) == _daily_activity_route_outcome( + existing_path, user_role + ) + def test_non_admin_config_update_route_rejected(): """Test that non-admin users are rejected when trying to call /config/update""" @@ -2219,7 +2306,7 @@ def test_internal_user_can_access_logs_drawer_detail_route(user_role): request_data={}, ) except Exception as e: - pytest.fail(f"{user_role.value} should be able to access {route}. Got error: {str(e)}") + pytest.fail(f"{user_role.value} should be able to access {route}. Got error: {e!s}") @pytest.mark.parametrize( diff --git a/tests/unit/proxy/db/test_prisma_client.py b/tests/unit/proxy/db/test_prisma_client.py index 99e494fccd5..7b8d000a8d5 100644 --- a/tests/unit/proxy/db/test_prisma_client.py +++ b/tests/unit/proxy/db/test_prisma_client.py @@ -448,9 +448,25 @@ def test_db_push_without_the_prisma_runner_fails_the_migration_instead_of_crashi ): """ An ImportError out of setup_database escapes the caller's RuntimeError handler and - kills boot, bypassing the operator's enforce_prisma_migration_check choice. + kills boot with a traceback instead of the failed-setup message and exit code. """ monkeypatch.setitem(sys.modules, "litellm_proxy_extras.prisma_toolchain", None) assert PrismaManager.setup_database(use_migrate=False) is False assert fake_prisma_cli.calls == [] + + +@pytest.mark.parametrize( + ("run", "outcome"), + ( + (PrismaManager.build_request_log_indexes, False), + (PrismaManager.start_request_log_index_build, None), + ), + ids=("wait-for-the-build", "start-the-build"), +) +def test_without_proxy_extras_the_index_build_reports_failure_instead_of_raising(monkeypatch, run, outcome): + """The migration job exits non-zero and a serving proxy keeps booting when the extras + package that owns the index build is not installed.""" + monkeypatch.setitem(sys.modules, "litellm_proxy_extras.utils", None) + + assert run() is outcome diff --git a/tests/unit/proxy/lens/test_endpoints.py b/tests/unit/proxy/lens/test_endpoints.py index 97bb7759a02..ca19277da08 100644 --- a/tests/unit/proxy/lens/test_endpoints.py +++ b/tests/unit/proxy/lens/test_endpoints.py @@ -4,7 +4,22 @@ import pytest from fastapi import HTTPException from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth -from litellm.proxy.lens.endpoints import user_scope +from litellm.proxy.lens.endpoints import list_agents, user_scope + + +@pytest.mark.parametrize("role", (LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY)) +@pytest.mark.asyncio +async def test_agent_discovery_without_trace_storage_is_empty(role: LitellmUserRoles) -> None: + auth: Final = UserAPIKeyAuth(user_role=role) + assert await list_agents(auth, None) == () + + +@pytest.mark.asyncio +async def test_agent_discovery_without_trace_storage_still_requires_admin_access() -> None: + auth: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER) + with pytest.raises(HTTPException) as error: + await list_agents(auth, None) + assert error.value.status_code == 403 @pytest.mark.parametrize( diff --git a/tests/unit/proxy/lens/test_sources.py b/tests/unit/proxy/lens/test_sources.py index 5dc6e2652f0..f063f496314 100644 --- a/tests/unit/proxy/lens/test_sources.py +++ b/tests/unit/proxy/lens/test_sources.py @@ -61,3 +61,48 @@ async def test_sample_never_returns_authentication_attributes() -> None: assert sample.executions[0].metadata == (MetadataFilter(key="environment", value="production"),) assert "opaque-oauth-bearer" not in sample.model_dump_json() assert sample.eligible == 1 + + +@pytest.mark.asyncio +async def test_agents_use_the_same_team_and_key_scope_as_samples() -> None: + class AgentStorage: + async def lens_agents(self, parameters): + assert parameters["all_teams"] == 0 + assert parameters["team"] == "alpha" + assert parameters["key_hash"] == "key-hash" + return [{"agent_name": "research_agent"}, {"agent_name": "support_agent"}] + + names: Final = await SourceReader(AgentStorage()).agents(Scope(team_id="alpha", api_key_hash="key-hash")) + assert names == ("research_agent", "support_agent") + + +@pytest.mark.asyncio +async def test_request_only_storage_is_available_for_investigation() -> None: + class RequestStorage: + async def lens_availability(self, parameters): + assert parameters["team"] == "alpha" + return [{"traces": 0, "requests": 1}] + + available: Final = await SourceReader(RequestStorage()).availability(Scope(team_id="alpha")) + assert available.requests + assert not available.traces + + +@pytest.mark.asyncio +async def test_agent_filter_is_independent_of_service_and_metadata() -> None: + class SampleStorage: + async def lens_sample(self, parameters): + assert parameters["agent_name"] == "research_agent" + assert parameters["service"] == "shared-app" + assert parameters["filter_keys"] == ("enduser.id",) + assert parameters["filter_values"] == ("user-42",) + return [] + + settings: Final = lens().settings.model_copy( + update={ + "agent_name": "research_agent", + "service": "shared-app", + "filters": (MetadataFilter(key="enduser.id", value="user-42"),), + } + ) + assert not (await SourceReader(SampleStorage()).sample(Scope(all_teams=True), settings, 1, 2)).executions diff --git a/tests/unit/proxy/lens/test_state.py b/tests/unit/proxy/lens/test_state.py index ac70a22077e..0e01085fb04 100644 --- a/tests/unit/proxy/lens/test_state.py +++ b/tests/unit/proxy/lens/test_state.py @@ -86,7 +86,7 @@ def test_behavior_description_is_sufficient_without_separate_checks() -> None: @pytest.mark.parametrize( - "field,value", (("sample_percent", 0), ("sample_percent", 101), ("sample_size", 0), ("concurrency", 0)) + "field,value", (("sample_percent", 0), ("sample_percent", 101), ("sample_size", 0), ("concurrency", 0), ("lookback_hours", 0), ("lookback_hours", 8761)) ) def test_invalid_selection_and_parallelism_are_rejected(field: str, value: int) -> None: from pydantic import ValidationError @@ -145,11 +145,11 @@ def test_monthly_budget_renews_without_erasing_job_costs() -> None: assert renew_budget(spent, NOW) is spent -@pytest.mark.parametrize("hours", (24, 168, 720)) +@pytest.mark.parametrize("hours", (24, 168, 720, 4800, 8760)) def test_every_scan_uses_the_configured_lookback_window(hours: int) -> None: original: Final = lens() configured: Final = original.model_copy( - update={"settings": original.settings.model_copy(update={"lookback_hours": hours})} + update={"settings": LensSettings.model_validate({**original.settings.model_dump(), "lookback_hours": hours})} ) first: Final = queue_job(configured, NOW, "first") assert first.jobs[0].start == NOW - timedelta(hours=hours) diff --git a/tests/unit/proxy/lens/test_worker.py b/tests/unit/proxy/lens/test_worker.py index a0212e03319..dce3fc04d45 100644 --- a/tests/unit/proxy/lens/test_worker.py +++ b/tests/unit/proxy/lens/test_worker.py @@ -15,7 +15,7 @@ from litellm.proxy.lens.models import ( TracePart, ) from litellm.proxy.lens.state import queue_job -from litellm.proxy.lens.worker import LensWorker +from litellm.proxy.lens.worker import LensWorker, failure_message from tests.unit.proxy.lens.test_state import NOW, lens @@ -126,6 +126,34 @@ async def test_worker_reads_claimed_activity_and_reports_analysis_or_failure(mod assert result.coverage.screened == 1 assert result.coverage.unassessable == 0 elif model_status == 402: - assert result.error == "Monthly budget reached" + assert "HTTP 402" in result.error and "remaining budget" in result.error else: - assert result.error.startswith("Analysis interrupted.") + assert result.error.startswith("Model request failed (HTTP 503).") + + +@pytest.mark.parametrize("status", (400, 401, 402, 403, 404, 409, 429, 503)) +def test_failure_reports_action_and_status_without_private_response_content(status: int) -> None: + request: Final = httpx.Request( + "POST", "https://private-host.test/lens/worker/private-lens/private-run/model?token=secret" + ) + response: Final = httpx.Response(status, request=request, text="private trace content and key") + error: Final = httpx.HTTPStatusError("private exception details", request=request, response=response) + message: Final = failure_message(error) + assert message.startswith(f"Model request failed (HTTP {status}).") + assert "private" not in message and "secret" not in message + + +@pytest.mark.parametrize( + "route,action", (("sample", "Reading trace data"), ("content", "Reading trace data"), ("result", "Saving results")) +) +def test_failure_identifies_the_failing_worker_operation(route: str, action: str) -> None: + request: Final = httpx.Request("GET", f"https://proxy.test/lens/worker/lens/job/{route}") + response: Final = httpx.Response(503, request=request) + error: Final = httpx.HTTPStatusError("private body", request=request, response=response) + assert failure_message(error).startswith(f"{action} failed (HTTP 503).") + + +def test_connection_timeout_and_invalid_response_have_distinct_private_diagnostics() -> None: + assert "connect to the proxy" in failure_message(httpx.ConnectError("private hostname")) + assert "timed out" in failure_message(httpx.ReadTimeout("private prompt")) + assert "structured JSON" in failure_message(ValueError("private model response")) diff --git a/tests/unit/proxy/management_endpoints/test_activity_tenant_scoping.py b/tests/unit/proxy/management_endpoints/test_activity_tenant_scoping.py index 8c80429aa92..bd1436dcd10 100644 --- a/tests/unit/proxy/management_endpoints/test_activity_tenant_scoping.py +++ b/tests/unit/proxy/management_endpoints/test_activity_tenant_scoping.py @@ -387,3 +387,59 @@ async def test_agent_activity_non_admin_no_access_returns_empty_page(): assert result.results == [] fake_get_daily.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("owned_tokens", "requested_api_key"), + [ + ([], None), + (["alice-key-1"], "bob-key-1"), + ], +) +async def test_team_activity_member_without_matching_keys_queries_nothing( + owned_tokens: list[str], requested_api_key: str | None +) -> None: + """A member without full team view whose key list is empty, or who asks for + a key they do not own, must reach the repository with an empty key filter, + never with no filter at all.""" + from litellm.proxy.management_endpoints import common_daily_activity, team_endpoints + from litellm.repositories.daily_activity_sql import build_where_clause + from litellm.types.repositories.daily_activity import DailyRowsPage + + user = UserAPIKeyAuth(user_id="alice", user_role=LitellmUserRoles.INTERNAL_USER.value) + prisma = MagicMock() + prisma.db.litellm_teamtable.find_many = AsyncMock(return_value=[_make_team("team-B", admin_user_ids=["bob"])]) + prisma.db.litellm_verificationtoken.find_many = AsyncMock( + return_value=[MagicMock(token=token) for token in owned_tokens] + ) + user_info = MagicMock() + user_info.teams = ["team-B"] + repository = MagicMock() + repository.daily_rows = AsyncMock(return_value=DailyRowsPage(total_count=0, rows=())) + + with ( + patch.object(team_endpoints, "prisma_client", prisma, create=True), + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_user_object", + new=AsyncMock(return_value=user_info), + ), + patch.object(common_daily_activity, "daily_activity_repository", return_value=repository), + patch("litellm.proxy.proxy_server.prisma_client", prisma), + patch("litellm.proxy.proxy_server.user_api_key_cache", MagicMock()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + ): + response = await team_endpoints.get_team_daily_activity( + team_ids="team-B", + start_date="2026-01-01", + end_date="2026-01-02", + api_key=requested_api_key, + user_api_key_dict=user, + ) + + scope = repository.daily_rows.await_args.args[0] + assert scope.api_keys == () + sql, _params = build_where_clause(scope) + assert sql.endswith(" AND FALSE") + assert response.results == [] + assert response.metadata.total_spend == 0 diff --git a/tests/unit/proxy/management_endpoints/test_common_daily_activity.py b/tests/unit/proxy/management_endpoints/test_common_daily_activity.py index 7cc5100037e..8808b73f89d 100644 --- a/tests/unit/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/unit/proxy/management_endpoints/test_common_daily_activity.py @@ -1,28 +1,205 @@ -from collections.abc import Sequence -from datetime import datetime, timedelta, timezone +from collections.abc import Mapping, Sequence +from datetime import date, datetime from types import SimpleNamespace from typing import Final from unittest.mock import AsyncMock, MagicMock import pytest +from fastapi import HTTPException +import litellm.proxy.management_endpoints.common_daily_activity as common_daily_activity_module +from litellm.constants import USAGE_TOP_API_KEYS_DEFAULT from litellm.proxy.management_endpoints.common_daily_activity import ( - _adjust_dates_for_timezone, - _build_aggregated_sql_query, - _build_entity_rollup_sql_query, + CanonicalDateRange, + InvalidDateRange, _is_user_agent_tag, + _ProxyDailyActivityReads, _record_to_spend_metrics, + compute_tag_metadata_totals, + daily_activity_repository, + daily_activity_scope, get_api_key_metadata, get_daily_activity, - get_daily_activity_aggregated, + parse_canonical_date, + parse_canonical_date_range, + raise_public, update_metrics, ) +from litellm.proxy.management_endpoints.common_daily_activity import ( + get_daily_activity_aggregated as _get_daily_activity_aggregated, +) from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR -from litellm.proxy.utils import hash_token +from litellm.proxy.utils import PrismaClient, hash_token from litellm.types.proxy.management_endpoints.common_daily_activity import ( DailySpendMetadata, + SpendAnalyticsPaginatedResponse, SpendMetrics, ) +from litellm.types.repositories.daily_activity import GroupingSetsRow, KeyMetadataRow + + +async def _run_aggregated_daily_activity( + *, + prisma_client: PrismaClient, + table_name: str, + entity_id_field: str, + entity_id: str | list[str] | None, + entity_metadata_field: Mapping[str, dict[str, object]] | None = None, + start_date: str, + end_date: str, + model: str | None, + api_key: str | list[str] | None, + exclude_entity_ids: list[str] | None = None, + timezone_offset_minutes: int | None = None, + include_current_utc_day: bool = False, + include_entity_breakdown: bool = False, + api_key_limit: int = USAGE_TOP_API_KEYS_DEFAULT, +) -> SpendAnalyticsPaginatedResponse: + repository: Final = daily_activity_repository(prisma_client) + scope: Final = daily_activity_scope( + table_name, + entity_id_field, + entity_id, + exclude_entity_ids, + api_key, + start_date, + end_date, + model, + timezone_offset_minutes, + include_current_utc_day, + ) + return await _get_daily_activity_aggregated( + repository, + scope, + entity_metadata_field=entity_metadata_field, + include_entity_breakdown=include_entity_breakdown, + api_key_limit=api_key_limit, + ) + + +async def get_daily_activity_aggregated( + *, + prisma_client: PrismaClient, + table_name: str, + entity_id_field: str, + entity_id: str | list[str] | None, + entity_metadata_field: Mapping[str, dict[str, object]] | None = None, + start_date: str, + end_date: str, + model: str | None, + api_key: str | list[str] | None, + exclude_entity_ids: list[str] | None = None, + timezone_offset_minutes: int | None = None, + include_current_utc_day: bool = False, + include_entity_breakdown: bool = False, + api_key_limit: int = USAGE_TOP_API_KEYS_DEFAULT, +) -> SpendAnalyticsPaginatedResponse: + return await _run_aggregated_daily_activity( + prisma_client=prisma_client, + table_name=table_name, + entity_id_field=entity_id_field, + entity_id=entity_id, + entity_metadata_field=entity_metadata_field, + start_date=start_date, + end_date=end_date, + model=model, + api_key=api_key, + exclude_entity_ids=exclude_entity_ids, + timezone_offset_minutes=timezone_offset_minutes, + include_current_utc_day=include_current_utc_day, + include_entity_breakdown=include_entity_breakdown, + api_key_limit=api_key_limit, + ) + + +@pytest.mark.asyncio +async def test_get_daily_activity_requires_a_database(): + with pytest.raises(HTTPException) as error: + await get_daily_activity( + prisma_client=None, + table_name="litellm_dailyuserspend", + entity_id_field="user_id", + entity_id="user-1", + entity_metadata_field=None, + start_date="2026-06-16", + end_date="2026-06-16", + model=None, + api_key=None, + page=1, + page_size=10, + ) + + assert error.value.status_code == 500 + assert error.value.detail == {"error": common_daily_activity_module.CommonProxyErrors.db_not_connected_error.value} + + +@pytest.mark.asyncio +async def test_get_daily_activity_maps_repository_failures_to_http_errors(): + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_table = MagicMock() + mock_table.count = AsyncMock(return_value=0) + mock_table.find_many = AsyncMock(side_effect=RuntimeError("daily rows unavailable")) + mock_prisma.db.litellm_dailyuserspend = mock_table + + with pytest.raises(HTTPException) as error: + await get_daily_activity( + prisma_client=mock_prisma, + table_name="litellm_dailyuserspend", + entity_id_field="user_id", + entity_id="user-1", + entity_metadata_field=None, + start_date="2026-06-16", + end_date="2026-06-16", + model=None, + api_key=None, + page=1, + page_size=10, + ) + + assert error.value.status_code == 500 + assert error.value.detail == {"error": "Failed to fetch analytics: daily rows unavailable"} + + +@pytest.mark.asyncio +async def test_get_daily_activity_aggregated_maps_repository_failures_to_http_errors(): + repository = MagicMock() + repository.aggregated = AsyncMock(side_effect=RuntimeError("daily aggregate unavailable")) + scope = daily_activity_scope( + "litellm_dailyuserspend", + "user_id", + "user-1", + None, + None, + "2026-06-16", + "2026-06-16", + None, + None, + ) + + with pytest.raises(HTTPException) as error: + await _get_daily_activity_aggregated(repository, scope) + + assert error.value.status_code == 500 + assert error.value.detail == {"error": "Failed to fetch analytics: daily aggregate unavailable"} + + +def test_compute_tag_metadata_totals_deduplicates_and_ignores_user_agent_tags(): + smaller = _spend_record("key-1", spend=1.0) + smaller.request_id = "request-1" + smaller.tag = "environment: small" + larger = _spend_record("key-1", spend=4.0) + larger.request_id = "request-1" + larger.tag = "environment: large" + larger.api_requests = 1 + user_agent = _spend_record("key-2", spend=10.0) + user_agent.request_id = "request-2" + user_agent.tag = "User-Agent: test" + user_agent.api_requests = 3 + + totals = compute_tag_metadata_totals((smaller, larger, user_agent)) + + assert (totals.spend, totals.api_requests) == (4.0, 1) @pytest.mark.asyncio @@ -39,12 +216,12 @@ async def test_get_daily_activity_empty_entity_id_list(): mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) # Set the table name dynamically - mock_prisma.db.litellm_dailyspend = mock_table + mock_prisma.db.litellm_dailyteamspend = mock_table # Call the function with empty entity_id list - result = await get_daily_activity( + await get_daily_activity( prisma_client=mock_prisma, - table_name="litellm_dailyspend", + table_name="litellm_dailyteamspend", entity_id_field="team_id", entity_id=[], entity_metadata_field=None, @@ -87,11 +264,11 @@ async def test_get_daily_activity_order_has_id_tiebreaker(): mock_table.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_verificationtoken = MagicMock() mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) - mock_prisma.db.litellm_dailyspend = mock_table + mock_prisma.db.litellm_dailyteamspend = mock_table await get_daily_activity( prisma_client=mock_prisma, - table_name="litellm_dailyspend", + table_name="litellm_dailyteamspend", entity_id_field="team_id", entity_id="team-1", entity_metadata_field=None, @@ -105,7 +282,7 @@ async def test_get_daily_activity_order_has_id_tiebreaker(): mock_table.find_many.assert_called_once() order = mock_table.find_many.call_args[1]["order"] - assert order == [{"date": "desc"}, {"id": "asc"}], ( + assert order == ({"date": "desc"}, {"id": "asc"}), ( f"order must include the id tiebreaker after date for stable offset pagination (see #30164); got {order!r}" ) @@ -169,6 +346,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown(): "endpoint": "/v1/chat/completions", "api_key": None, "group_level": 62, + "distinct_api_keys": None, "spend": 15.0, "prompt_tokens": 150, "completion_tokens": 75, @@ -181,31 +359,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown(): "endpoint": "/v1/embeddings", "api_key": None, "group_level": 62, - "spend": 3.0, - "prompt_tokens": 30, - "completion_tokens": 0, - "api_requests": 1, - "successful_requests": 1, - }, - # (date, endpoint, api_key) — populates the per-key sub-bucket - { - **base, - "date": "2024-01-01", - "endpoint": "/v1/chat/completions", - "api_key": "key-1", - "group_level": 30, - "spend": 15.0, - "prompt_tokens": 150, - "completion_tokens": 75, - "api_requests": 2, - "successful_requests": 2, - }, - { - **base, - "date": "2024-01-01", - "endpoint": "/v1/embeddings", - "api_key": "key-2", - "group_level": 30, + "distinct_api_keys": None, "spend": 3.0, "prompt_tokens": 30, "completion_tokens": 0, @@ -219,6 +373,7 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown(): "endpoint": None, "api_key": None, "group_level": 63, + "distinct_api_keys": None, "spend": 18.0, "prompt_tokens": 180, "completion_tokens": 75, @@ -232,12 +387,40 @@ async def test_get_daily_activity_aggregated_with_endpoint_breakdown(): "endpoint": None, "api_key": None, "group_level": 127, + "distinct_api_keys": None, "spend": 18.0, "prompt_tokens": 180, "completion_tokens": 75, "api_requests": 3, "successful_requests": 3, }, + # (date, endpoint, api_key) — populates the per-key sub-bucket + { + **base, + "date": "2024-01-01", + "endpoint": "/v1/chat/completions", + "api_key": "key-1", + "group_level": 30, + "distinct_api_keys": 2, + "spend": 15.0, + "prompt_tokens": 150, + "completion_tokens": 75, + "api_requests": 2, + "successful_requests": 2, + }, + { + **base, + "date": "2024-01-01", + "endpoint": "/v1/embeddings", + "api_key": "key-2", + "group_level": 30, + "distinct_api_keys": 2, + "spend": 3.0, + "prompt_tokens": 30, + "completion_tokens": 0, + "api_requests": 1, + "successful_requests": 1, + }, ] mock_prisma.db.query_raw = AsyncMock(return_value=mock_rows) @@ -313,6 +496,49 @@ async def test_get_api_key_metadata_returns_active_key_metadata(): assert result["active-key-hash-123"]["team_id"] == "team-abc" +@pytest.mark.asyncio +async def test_recovered_key_metadata_preserves_resolved_tags_after_user_details( + monkeypatch: pytest.MonkeyPatch, +) -> None: + resolved: Final = KeyMetadataRow( + api_key="key-hash", + key_alias="key alias", + team_id="team-id", + user_id="user-id", + user_email=None, + key_exists=True, + tags=("production", "internal"), + ) + attach_details: Final = AsyncMock( + return_value={ + "key-hash": { + "key_alias": "key alias", + "team_id": "team-id", + "user_id": "user-id", + "user_email": "user@example.com", + "key_exists": True, + } + } + ) + monkeypatch.setattr(common_daily_activity_module, "attach_user_details", attach_details) + reads: Final = _ProxyDailyActivityReads(MagicMock()) + + result: Final = await reads.recover_key_metadata({"key-hash": resolved}, frozenset(("key-hash",)), None) + + assert result == { + "key-hash": KeyMetadataRow( + api_key="key-hash", + key_alias="key alias", + team_id="team-id", + user_id="user-id", + user_email="user@example.com", + key_exists=True, + tags=("production", "internal"), + ) + } + attach_details.assert_awaited_once() + + @pytest.mark.asyncio async def test_get_api_key_metadata_falls_back_to_deleted_keys(): """Test that get_api_key_metadata should fall back to deleted keys table for missing keys.""" @@ -341,7 +567,6 @@ async def test_get_api_key_metadata_falls_back_to_deleted_keys(): # Verify deleted table was queried with the missing key mock_prisma.db.litellm_deletedverificationtoken.find_many.assert_called_once_with( where={"token": {"in": ["deleted-key-hash-456"]}}, - order={"deleted_at": "desc"}, ) @@ -437,11 +662,13 @@ async def test_get_api_key_metadata_regenerated_key_uses_most_recent_deleted_rec mock_deleted_1.token = "old-key-hash" mock_deleted_1.key_alias = "latest-alias" mock_deleted_1.team_id = "latest-team" + mock_deleted_1.deleted_at = datetime(2024, 1, 2) mock_deleted_2 = MagicMock() mock_deleted_2.token = "old-key-hash" mock_deleted_2.key_alias = "older-alias" mock_deleted_2.team_id = "older-team" + mock_deleted_2.deleted_at = datetime(2024, 1, 1) # Ordered by deleted_at desc, so first record is the most recent mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[mock_deleted_1, mock_deleted_2]) @@ -629,12 +856,15 @@ def test_key_metadata_includes_recovered_user_email(): meta = _key_metadata( { - "dirty-key": { - "key_alias": "batch-worker", - "team_id": "team-1", - "user_id": "alice", - "user_email": "alice@example.com", - } + "dirty-key": KeyMetadataRow( + api_key="dirty-key", + key_alias="batch-worker", + team_id="team-1", + user_id="alice", + user_email="alice@example.com", + key_exists=True, + tags=(), + ) }, "dirty-key", ) @@ -649,11 +879,15 @@ def test_key_metadata_includes_user_id_without_user_email(): meta = _key_metadata( { - "dirty-key": { - "key_alias": "batch-worker", - "team_id": "team-1", - "user_id": "user-123", - } + "dirty-key": KeyMetadataRow( + api_key="dirty-key", + key_alias="batch-worker", + team_id="team-1", + user_id="user-123", + user_email=None, + key_exists=True, + tags=(), + ) }, "dirty-key", ) @@ -694,11 +928,15 @@ def test_update_breakdown_metrics_includes_user_email(): user_id="alice", ) api_key_metadata = { - "dirty-key": { - "key_alias": "batch-worker", - "team_id": "team-1", - "user_email": "alice@example.com", - } + "dirty-key": KeyMetadataRow( + api_key="dirty-key", + key_alias="batch-worker", + team_id="team-1", + user_id=None, + user_email="alice@example.com", + key_exists=True, + tags=(), + ) } update_breakdown_metrics( @@ -870,6 +1108,7 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys(): "endpoint": "/v1/chat/completions", "api_key": None, "group_level": 62, + "distinct_api_keys": None, "spend": 10.0, "prompt_tokens": 100, "completion_tokens": 50, @@ -882,6 +1121,7 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys(): "endpoint": "/v1/chat/completions", "api_key": "deleted-key-hash", "group_level": 30, + "distinct_api_keys": 1, "spend": 10.0, "prompt_tokens": 100, "completion_tokens": 50, @@ -963,10 +1203,21 @@ async def test_aggregated_activity_flags_only_keys_that_key_info_can_still_resol return_value=[{**base, "api_key": key} for key in ("active-key", "deleted-key", "session-key")] ) mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock( - return_value=[SimpleNamespace(token="active-key", key_alias="active", team_id=None, user_id="owner")] + return_value=[ + SimpleNamespace(token="active-key", key_alias="active", team_id=None, user_id="owner", metadata=None) + ] ) mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock( - return_value=[SimpleNamespace(token="deleted-key", key_alias="deleted", team_id=None, user_id="owner")] + return_value=[ + SimpleNamespace( + token="deleted-key", + key_alias="deleted", + team_id=None, + user_id="owner", + metadata=None, + deleted_at=datetime(2024, 1, 2), + ) + ] ) mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[]) @@ -1128,261 +1379,6 @@ async def test_model_groups_breakdown_keys_by_public_name_with_model_fallback(): assert breakdown.models["claude-x"].metrics.spend == 2.0 -class TestAdjustDatesForTimezone: - """ - Regression tests for the timezone double-counting bug. - - Background: the previous implementation expanded the SQL date range by a full - UTC day on whichever side a non-UTC timezone offset pointed. Because spend is - bucketed in whole UTC days in the aggregation table, that expansion caused - single-day queries from non-UTC timezones to include a second full UTC day's - worth of data, producing approximately 2x over-counting. The sum of single-day - spends across a window then exceeded the equivalent multi-day aggregate, which - is mathematically impossible. - - These tests pin the function to a pass-through and assert the additivity - invariant that any future implementation must preserve. - """ - - @pytest.mark.parametrize( - "offset_minutes", - [ - None, - 0, - -330, # IST UTC+5:30 - -540, # JST UTC+9 - -60, # CET UTC+1 - 240, # AST UTC-4 - 300, # EST UTC-5 - 480, # PST UTC-8 - ], - ) - def test_returns_input_dates_unchanged_for_any_offset(self, offset_minutes): - start, end = _adjust_dates_for_timezone("2026-05-29", "2026-05-29", offset_minutes) - assert start == "2026-05-29" - assert end == "2026-05-29" - - def test_single_day_query_does_not_widen_to_two_utc_days(self): - """ - Pins the boundary that caused the original 2x bug: a single IST day must - not be translated into a SQL filter covering two UTC days. - """ - start, end = _adjust_dates_for_timezone("2026-05-29", "2026-05-29", -330) - assert start == end == "2026-05-29", ( - "Single-day IST query expanded to a multi-day UTC range; this is " - "the regression that produced approximately 2x over-counting." - ) - - def test_multi_day_range_endpoints_are_preserved(self): - start, end = _adjust_dates_for_timezone("2026-05-29", "2026-06-02", -330) - assert (start, end) == ("2026-05-29", "2026-06-02") - - @pytest.mark.parametrize("offset_minutes", [-330, 480]) - def test_single_day_sums_match_multi_day_window(self, offset_minutes): - """ - Additivity invariant: querying each day in a window separately and summing - the resulting SQL ranges must cover exactly the same range as querying the - whole window at once. The bug broke this; without it, single-day sums - exceeded the multi-day total by ~50% over a 5-day IST window. - """ - days = ["2026-05-29", "2026-05-30", "2026-05-31", "2026-06-01", "2026-06-02"] - single_day_ranges = [_adjust_dates_for_timezone(d, d, offset_minutes) for d in days] - multi_day_range = _adjust_dates_for_timezone(days[0], days[-1], offset_minutes) - - per_day_starts = [r[0] for r in single_day_ranges] - per_day_ends = [r[1] for r in single_day_ranges] - assert min(per_day_starts) == multi_day_range[0] - assert max(per_day_ends) == multi_day_range[1] - assert per_day_starts == days - assert per_day_ends == days - - -class TestAdjustDatesForTimezoneLiveEnd: - """ - Regression tests for the stale-evening bug: a caller west of UTC whose range - ends on their local "today" was capped at that local date's UTC bucket, so - once UTC rolled past their local midnight (5pm PT), everything sent that - evening sat in the next UTC bucket and the dashboard reported $0 for it - until local midnight. A range that reaches the caller's current day and - opts in via include_current_utc_day must extend to today's UTC bucket; the - only part of that bucket outside the range is the future, which is empty, - so the extension cannot over-count. Callers that do not opt in keep the - pass-through byte for byte. - """ - - PT_EVENING_UTC: Final = datetime(2026, 8, 6, 4, 30, tzinfo=timezone.utc) - - def test_pt_evening_range_ending_today_extends_to_utc_today(self): - start, end = _adjust_dates_for_timezone( - "2026-07-06", "2026-08-05", 420, include_current_utc_day=True, utc_now=self.PT_EVENING_UTC - ) - assert (start, end) == ("2026-07-06", "2026-08-06") - - def test_without_opt_in_live_range_keeps_pass_through(self): - start, end = _adjust_dates_for_timezone("2026-07-06", "2026-08-05", 420, utc_now=self.PT_EVENING_UTC) - assert (start, end) == ("2026-07-06", "2026-08-05") - - def test_pt_historical_range_is_untouched(self): - start, end = _adjust_dates_for_timezone( - "2026-07-01", "2026-08-04", 420, include_current_utc_day=True, utc_now=self.PT_EVENING_UTC - ) - assert (start, end) == ("2026-07-01", "2026-08-04") - - def test_east_of_utc_local_today_already_covers_utc_today(self): - ist_evening_utc: Final = datetime(2026, 8, 5, 17, 0, tzinfo=timezone.utc) - start, end = _adjust_dates_for_timezone( - "2026-07-07", "2026-08-06", -330, include_current_utc_day=True, utc_now=ist_evening_utc - ) - assert (start, end) == ("2026-07-07", "2026-08-06") - - def test_missing_offset_stays_pass_through_even_for_live_range(self): - start, end = _adjust_dates_for_timezone( - "2026-07-06", "2026-08-05", None, include_current_utc_day=True, utc_now=self.PT_EVENING_UTC - ) - assert (start, end) == ("2026-07-06", "2026-08-05") - - def test_utc_caller_range_ending_today_is_unchanged(self): - utc_noon: Final = datetime(2026, 8, 5, 12, 0, tzinfo=timezone.utc) - start, end = _adjust_dates_for_timezone( - "2026-07-06", "2026-08-05", 0, include_current_utc_day=True, utc_now=utc_noon - ) - assert (start, end) == ("2026-07-06", "2026-08-05") - - def test_future_end_date_extends_no_further_than_requested(self): - start, end = _adjust_dates_for_timezone( - "2026-07-06", "2026-08-09", 420, include_current_utc_day=True, utc_now=self.PT_EVENING_UTC - ) - assert (start, end) == ("2026-07-06", "2026-08-09") - - -class TestBuildAggregatedSqlQuery: - """ - Asserts the SQL emitted by the aggregated query path stays anchored to the - user-supplied date range. The original bug shipped a function that returned - expanded dates from _adjust_dates_for_timezone, so the regression surface is - not just the helper but the SQL it feeds into. - """ - - @pytest.mark.parametrize("offset_minutes", [None, 0, -330, 480]) - def test_sql_date_bounds_are_user_supplied_dates(self, offset_minutes): - sql, params = _build_aggregated_sql_query( - table_name="litellm_dailyuserspend", - entity_id_field="user_id", - entity_id="user-1", - start_date="2026-05-29", - end_date="2026-05-29", - model=None, - api_key=None, - timezone_offset_minutes=offset_minutes, - ) - - assert params[0] == "2026-05-29" - assert params[1] == "2026-05-29" - assert "date >= $1" in sql - assert "date <= $2" in sql - - @pytest.mark.parametrize("build", [_build_aggregated_sql_query, _build_entity_rollup_sql_query]) - def test_include_current_utc_day_extends_live_end_bound(self, build): - """ - An offset larger than 24h keeps the caller's local date behind UTC at any - wall-clock hour, so the live-end extension is deterministic: a range ending - on the caller's local today must reach today's UTC bucket (LIT-5818, guards - the #36051 behavior on the aggregated path). - """ - offset_minutes: Final = 1500 - caller_local_today: Final = (datetime.now(timezone.utc) - timedelta(minutes=offset_minutes)).date().isoformat() - utc_today: Final = datetime.now(timezone.utc).date().isoformat() - - _sql, params = build( - table_name="litellm_dailyuserspend", - entity_id_field="user_id", - entity_id="user-1", - start_date="2026-05-01", - end_date=caller_local_today, - model=None, - api_key=None, - timezone_offset_minutes=offset_minutes, - include_current_utc_day=True, - ) - - assert params[0] == "2026-05-01" - assert params[1] == utc_today - - def test_optional_filters_appear_in_params_in_order(self): - sql, params = _build_aggregated_sql_query( - table_name="litellm_dailyuserspend", - entity_id_field="user_id", - entity_id="user-1", - start_date="2026-05-29", - end_date="2026-06-02", - model="bedrock/global.anthropic.claude-opus-4-8", - api_key="sk-test", - timezone_offset_minutes=-330, - ) - - assert params == [ - "2026-05-29", - "2026-06-02", - "user-1", - "bedrock/global.anthropic.claude-opus-4-8", - "sk-test", - ] - assert "model = $4" in sql - assert "api_key = $5" in sql - - -class TestAggregatedEmptyEntityFilter: - _BUILDERS: Final = (_build_aggregated_sql_query, _build_entity_rollup_sql_query) - - @pytest.mark.parametrize("build", _BUILDERS) - def test_empty_entity_list_emits_no_degenerate_in_clause(self, build): - sql, params = build( - table_name="litellm_dailyteamspend", - entity_id_field="team_id", - entity_id=[], - start_date="2026-08-01", - end_date="2026-08-19", - model=None, - api_key=None, - ) - - normalized = " ".join(sql.split()) - assert "IN ()" not in normalized - assert '"team_id" IN' not in normalized - assert params == ["2026-08-01", "2026-08-19"] - - @pytest.mark.parametrize("build", _BUILDERS) - def test_empty_entity_list_matches_nothing_rather_than_everything(self, build): - sql, _ = build( - table_name="litellm_dailyteamspend", - entity_id_field="team_id", - entity_id=[], - start_date="2026-08-01", - end_date="2026-08-19", - model=None, - api_key=None, - ) - - assert "FALSE" in " ".join(sql.split()) - - @pytest.mark.parametrize("build", _BUILDERS) - def test_populated_entity_list_still_filters_on_its_ids(self, build): - sql, params = build( - table_name="litellm_dailyteamspend", - entity_id_field="team_id", - entity_id=["team-alpha", "team-beta"], - start_date="2026-08-01", - end_date="2026-08-19", - model=None, - api_key=None, - ) - - normalized = " ".join(sql.split()) - assert '"team_id" IN ($3, $4)' in normalized - assert "FALSE" not in normalized - assert params == ["2026-08-01", "2026-08-19", "team-alpha", "team-beta"] - - @pytest.mark.asyncio async def test_get_daily_activity_aggregated_empty_result_set(): """Regression test for the empty-range 500. @@ -1405,6 +1401,7 @@ async def test_get_daily_activity_aggregated_empty_result_set(): "mcp_namespaced_tool_name": None, "endpoint": None, "group_level": 127, + "distinct_api_keys": None, "spend": None, "prompt_tokens": None, "completion_tokens": None, @@ -1438,6 +1435,7 @@ async def test_get_daily_activity_aggregated_empty_result_set(): assert result.results == [] assert result.metadata.total_spend == 0.0 + assert result.metadata.entity_total_api_keys is None assert result.metadata.total_prompt_tokens == 0 assert result.metadata.total_completion_tokens == 0 assert result.metadata.total_tokens == 0 @@ -1517,20 +1515,6 @@ class TestEverySavingsDriverSurvivesTheReadPath: assert drivers, "expected the dashboard response to expose at least one savings driver" return drivers - def test_every_driver_is_summed_by_the_rollup_query(self): - sql, _ = _build_aggregated_sql_query( - table_name="litellm_dailyuserspend", - entity_id_field="user_id", - entity_id="user-1", - start_date="2026-07-01", - end_date="2026-07-31", - model=None, - api_key=None, - timezone_offset_minutes=None, - ) - for driver in self._drivers(): - assert f"SUM({driver})" in sql, f"{driver} is never summed, so it reads as zero" - def test_every_driver_is_accumulated_across_rows(self): for driver in self._drivers(): record = _no_spend_record() @@ -1558,20 +1542,6 @@ class TestResponseTimeSurvivesTheReadPath: _FIELDS = ("total_response_time_ms", "timed_requests") - def test_both_halves_are_summed_by_the_rollup_query(self): - sql, _ = _build_aggregated_sql_query( - table_name="litellm_dailyuserspend", - entity_id_field="user_id", - entity_id="user-1", - start_date="2026-09-01", - end_date="2026-09-30", - model=None, - api_key=None, - timezone_offset_minutes=None, - ) - for field in self._FIELDS: - assert f"SUM({field})" in sql, f"{field} is never summed, so the average reads as zero" - def test_accumulating_rows_keeps_sum_and_count_paired(self): first = _no_spend_record() first.total_response_time_ms = 1500 @@ -1608,6 +1578,12 @@ def ptu_cost_attribution_enabled(monkeypatch): def _spend_record(api_key, *, model="gpt-4o-mini-ptu", spend=0.0, ptu_flat_cost=0.0): return SimpleNamespace( api_key=api_key, + user_id=None, + team_id=None, + tag=None, + organization_id=None, + end_user_id=None, + agent_id=None, model=model, model_group=None, mcp_namespaced_tool_name=None, @@ -1630,6 +1606,7 @@ def _spend_record(api_key, *, model="gpt-4o-mini-ptu", spend=0.0, ptu_flat_cost= successful_requests=0, failed_requests=0, ptu_flat_cost=ptu_flat_cost, + request_id=None, ) @@ -1669,9 +1646,7 @@ def _grouping_row( spend=0.0, ptu_flat_cost=0.0, ): - from litellm.proxy.management_endpoints.common_daily_activity import _GroupingSetsRow - - return _GroupingSetsRow( + return GroupingSetsRow( date="2024-01-01", api_key=api_key, model=model, @@ -1680,6 +1655,7 @@ def _grouping_row( mcp_namespaced_tool_name=mcp_namespaced_tool_name, endpoint=endpoint, group_level=group_level, + distinct_api_keys=None, spend=spend, ptu_flat_cost=ptu_flat_cost, prompt_tokens=0, @@ -2186,57 +2162,6 @@ class TestFlagIsNotReadOnTheHotPath: assert reads > 0 -def test_entity_rollup_sql_query_and_api_key_list_filter(): - """The entity rollup companion query keeps its own two grouping sets keyed - by GROUPING(api_key), shares the WHERE builder (list api_key becomes a - parameterized IN, an empty list must match nothing), and the main - aggregated query stays entity-free.""" - from litellm.proxy.management_endpoints.common_daily_activity import ( - _build_entity_rollup_sql_query, - ) - - sql, params = _build_entity_rollup_sql_query( - table_name="litellm_dailyteamspend", - entity_id_field="team_id", - entity_id=None, - start_date="2024-01-01", - end_date="2024-01-31", - model=None, - api_key=["key-1", "key-2"], - ) - assert '"team_id" AS entity_id' in sql - assert "GROUPING(api_key) AS api_key_rolled" in sql - assert '(date, "team_id"),' in sql - assert '(date, "team_id", api_key)' in sql - assert "api_key IN ($3, $4)" in sql - assert "SUM(ptu_flat_cost)::float" in sql - assert params == ["2024-01-01", "2024-01-31", "key-1", "key-2"] - - plain_sql, _ = _build_aggregated_sql_query( - table_name="litellm_dailyteamspend", - entity_id_field="team_id", - entity_id=None, - start_date="2024-01-01", - end_date="2024-01-31", - model=None, - api_key=None, - ) - assert "entity_id" not in plain_sql - assert "GROUPING(date" in plain_sql - - empty_sql, empty_params = _build_aggregated_sql_query( - table_name="litellm_dailyteamspend", - entity_id_field="team_id", - entity_id=None, - start_date="2024-01-01", - end_date="2024-01-31", - model=None, - api_key=[], - ) - assert "FALSE" in empty_sql - assert empty_params == ["2024-01-01", "2024-01-31"] - - @pytest.mark.asyncio async def test_get_daily_activity_aggregated_with_entity_breakdown(): """include_entity_breakdown must run the companion entity rollup query and @@ -2268,19 +2193,44 @@ async def test_get_daily_activity_aggregated_with_entity_breakdown(): "successful_requests": 0, } main_rows = [ - {**base, "date": None, "group_level": 127, "spend": 18.0}, - {**base, "date": "2024-01-01", "group_level": 63, "spend": 18.0}, - {**base, "date": "2024-01-01", "model": "gpt-4o", "group_level": 47, "spend": 18.0}, - {**base, "date": "2024-01-01", "api_key": "key-1", "group_level": 31, "spend": 12.0}, + {**base, "date": None, "group_level": 127, "distinct_api_keys": None, "spend": 18.0}, + {**base, "date": "2024-01-01", "group_level": 63, "distinct_api_keys": None, "spend": 18.0}, + {**base, "date": "2024-01-01", "model": "gpt-4o", "group_level": 47, "distinct_api_keys": None, "spend": 18.0}, + {**base, "date": "2024-01-01", "api_key": "key-1", "group_level": 31, "distinct_api_keys": 1, "spend": 12.0}, ] - entity_base = { - key: value - for key, value in base.items() - if key not in ("model", "model_group", "custom_llm_provider", "mcp_namespaced_tool_name", "endpoint") + entity_base: Final = { + **{ + key: value + for key, value in base.items() + if key not in ("model", "model_group", "custom_llm_provider", "mcp_namespaced_tool_name", "endpoint") + }, + "distinct_api_keys": None, } entity_rows = [ - {**entity_base, "date": "2024-01-01", "entity_id": "team-a", "api_key_rolled": 1, "spend": 12.0}, - {**entity_base, "date": "2024-01-01", "entity_id": "team-b", "api_key_rolled": 1, "spend": 6.0}, + { + **entity_base, + "date": "2024-01-01", + "entity_id": "team-a", + "api_key_rolled": 1, + "distinct_api_keys": 3, + "spend": 12.0, + }, + { + **entity_base, + "date": "2024-01-01", + "entity_id": "team-b", + "api_key_rolled": 1, + "distinct_api_keys": 2, + "spend": 4.0, + }, + { + **entity_base, + "date": "2024-01-01", + "entity_id": None, + "api_key_rolled": 1, + "distinct_api_keys": 1, + "spend": 2.0, + }, { **entity_base, "date": "2024-01-01", @@ -2295,7 +2245,17 @@ async def test_get_daily_activity_aggregated_with_entity_breakdown(): "entity_id": "team-b", "api_key": "key-2", "api_key_rolled": 0, - "spend": 6.0, + "distinct_api_keys": None, + "spend": 4.0, + }, + { + **entity_base, + "date": "2024-01-01", + "entity_id": None, + "api_key": "key-3", + "api_key_rolled": 0, + "distinct_api_keys": None, + "spend": 2.0, }, ] @@ -2320,22 +2280,25 @@ async def test_get_daily_activity_aggregated_with_entity_breakdown(): main_sql = mock_prisma.db.query_raw.call_args_list[0][0][0] entity_sql = mock_prisma.db.query_raw.call_args_list[1][0][0] assert "entity_id" not in main_sql - assert '"team_id" AS entity_id' in entity_sql - assert '(date, "team_id"),' in entity_sql + assert "COALESCE(\"team_id\", '') AS entity_id" in entity_sql + assert "GROUP BY date, COALESCE(\"team_id\", '')" in entity_sql + assert '"team_id" AS entity_id' not in entity_sql assert result.metadata.total_spend == 18.0 + assert result.metadata.entity_total_api_keys == {"team-a": 3, "team-b": 2, "Unassigned": 1} assert len(result.results) == 1 daily = result.results[0] assert daily.metrics.spend == 18.0 entities = daily.breakdown.entities - assert set(entities) == {"team-a", "team-b"} + assert set(entities) == {"team-a", "team-b", "Unassigned"} assert entities["team-a"].metrics.spend == 12.0 assert entities["team-a"].metadata == {"team_alias": "Alpha"} assert entities["team-a"].api_key_breakdown["key-1"].metrics.spend == 12.0 - assert entities["team-b"].metrics.spend == 6.0 + assert entities["Unassigned"].api_key_breakdown["key-3"].metrics.spend == 2.0 + assert entities["team-b"].metrics.spend == 4.0 assert entities["team-b"].metadata == {} - assert entities["team-b"].api_key_breakdown["key-2"].metrics.spend == 6.0 + assert entities["team-b"].api_key_breakdown["key-2"].metrics.spend == 4.0 # Rollups with the entity bit set must still land in their usual buckets assert daily.breakdown.models["gpt-4o"].metrics.spend == 18.0 @@ -2374,17 +2337,17 @@ async def test_get_api_key_metadata_resolves_session_key_via_spend_log_window(): def test_spend_logs_window_pads_min_minus_one_day_and_max_plus_two_days(): - from litellm.proxy.management_endpoints.common_daily_activity import _spend_logs_window + from litellm.proxy.management_endpoints.common_daily_activity import spend_logs_window - window = _spend_logs_window({"2026-09-08", "2026-09-05", "not-a-date"}) + window = spend_logs_window({"2026-09-08", "2026-09-05", "not-a-date"}) assert window == (datetime(2026, 9, 4), datetime(2026, 9, 10)) def test_spend_logs_window_is_none_when_no_date_parses(): - from litellm.proxy.management_endpoints.common_daily_activity import _spend_logs_window + from litellm.proxy.management_endpoints.common_daily_activity import spend_logs_window - assert _spend_logs_window({"garbage", ""}) is None + assert spend_logs_window({"garbage", ""}) is None @pytest.mark.asyncio @@ -2489,3 +2452,63 @@ async def test_get_api_key_metadata_does_not_recover_daily_spend_owner_for_activ assert active_metadata.get("user_email") == "active-owner@example.com" assert active_metadata.get("key_exists") is True recovery_query_raw.assert_not_awaited() + + +def test_raise_public_maps_invalid_date_range_to_400() -> None: + with pytest.raises(HTTPException) as excinfo: + raise_public(InvalidDateRange(reason="Date range must be at most 400 days")) + assert excinfo.value.status_code == 400 + assert excinfo.value.detail == {"error": "Date range must be at most 400 days"} + + +@pytest.mark.parametrize("value", ("2026-9-24", "2026-09-24", "2026-09-4", "2026-02-30", "20260924", "")) +def test_parse_canonical_date_rejects_spellings_that_do_not_round_trip(value: str) -> None: + assert parse_canonical_date(value) is None + + +def test_parse_canonical_date_accepts_the_exact_yyyy_mm_dd_spelling() -> None: + assert parse_canonical_date("2026-09-24") == date(2026, 9, 24) + assert parse_canonical_date("0001-01-01") == date(1, 1, 1) + + +def test_parse_canonical_date_range_reports_missing_then_malformed_dates() -> None: + assert parse_canonical_date_range(None, "2026-09-24") == InvalidDateRange( + reason="Please provide start_date and end_date" + ) + assert parse_canonical_date_range("2026-09-24", "2026-9-26") == InvalidDateRange( + reason="start_date and end_date must be valid YYYY-MM-DD dates" + ) + assert parse_canonical_date_range("2026-09-24", "2026-09-26") == CanonicalDateRange( + start=date(2026, 9, 24), end=date(2026, 9, 26) + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("start_date", ("2026-9-24", "2026-09-24", "2026-09-4")) +async def test_get_daily_activity_rejects_non_canonical_dates_before_querying(start_date: str) -> None: + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_table = MagicMock() + mock_table.count = AsyncMock(return_value=0) + mock_table.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_dailyteamspend = mock_table + + with pytest.raises(HTTPException) as error: + await get_daily_activity( + prisma_client=mock_prisma, + table_name="litellm_dailyteamspend", + entity_id_field="team_id", + entity_id="team-a", + entity_metadata_field=None, + start_date=start_date, + end_date="2026-09-26", + model=None, + api_key=None, + page=1, + page_size=10, + ) + + assert error.value.status_code == 400 + assert error.value.detail == {"error": "start_date and end_date must be valid YYYY-MM-DD dates"} + mock_table.count.assert_not_awaited() + mock_table.find_many.assert_not_awaited() diff --git a/tests/unit/proxy/management_endpoints/test_daily_activity_routes.py b/tests/unit/proxy/management_endpoints/test_daily_activity_routes.py new file mode 100644 index 00000000000..c3f4fdfef54 --- /dev/null +++ b/tests/unit/proxy/management_endpoints/test_daily_activity_routes.py @@ -0,0 +1,1332 @@ +import csv +import io +from collections.abc import AsyncIterator, Iterator, Mapping, Sequence +from dataclasses import dataclass, fields +from itertools import chain +from types import SimpleNamespace +from typing import Final +from unittest.mock import AsyncMock + +import pytest +from fastapi import FastAPI, Request +from fastapi.testclient import TestClient + +from litellm import constants +from litellm.proxy._types import LiteLLM_TeamTable, LiteLLM_UserTable, LitellmUserRoles, Member, UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.management_endpoints.daily_activity_routes import ( + _csv_cell, + get_daily_activity_prisma_client, + get_daily_activity_repository, + router, +) +from litellm.types.proxy.management_endpoints.common_daily_activity import KeySpendMetrics, SpendMetrics +from litellm.types.repositories.daily_activity import ( + AggregatedRows, + DailyActivityScope, + DailyActivityTable, + EntityRollupRow, + ExportRow, + ExportType, + GroupingSetsRow, + KeyMetadataRow, + KeyPage, + KeySpendRow, +) + + +@dataclass(frozen=True, slots=True) +class _Activity: + table: DailyActivityTable + entity_id: str + date: str + api_key: str + model: str + model_group: str + spend: float + flat_cost: float + prompt_tokens: int + completion_tokens: int + cache_read_input_tokens: int + cache_creation_input_tokens: int + compression_saved_tokens: int + compression_savings_spend: float + prompt_caching_savings_spend: float + gateway_injected_caching_savings_spend: float + autorouter_savings_spend: float + api_requests: int + successful_requests: int + failed_requests: int + total_response_time_ms: int + timed_requests: int + + +_ENTITY_CASES: Final[tuple[tuple[str, str, str], ...]] = ( + ("/user", "user_id", "user-a"), + ("/team", "team_ids", "team-a"), + ("/tag", "tags", "blue"), + ("/organization", "organization_ids", "org-a"), + ("/customer", "end_user_ids", "customer-a"), + ("/agent", "agent_ids", "agent-a"), +) +_DATE_PARAMS: Final = {"start_date": "2025-01-01", "end_date": "2025-01-02"} + + +def _activity_for_entity( + table: DailyActivityTable, + entity_id: str, + key_rows: tuple[tuple[str, str, str, float, int], ...], +) -> tuple[_Activity, ...]: + return tuple( + _Activity( + table=table, + entity_id=entity_id, + date=date, + api_key=api_key, + model=model, + model_group="rare-group" if model == "rare-model" else "popular-group", + spend=spend, + flat_cost=0.05, + prompt_tokens=10, + completion_tokens=5, + cache_read_input_tokens=cache_read, + cache_creation_input_tokens=1, + compression_saved_tokens=2, + compression_savings_spend=0.1, + prompt_caching_savings_spend=0.2, + gateway_injected_caching_savings_spend=0.3, + autorouter_savings_spend=0.4, + api_requests=1, + successful_requests=1, + failed_requests=0, + total_response_time_ms=100, + timed_requests=1, + ) + for api_key, date, model, spend, cache_read in key_rows + ) + + +def _seeded_activity() -> tuple[_Activity, ...]: + entity_ids: Final = { + DailyActivityTable.USER: "user-a", + DailyActivityTable.TEAM: "team-a", + DailyActivityTable.TAG: "blue", + DailyActivityTable.ORGANIZATION: "org-a", + DailyActivityTable.CUSTOMER: "customer-a", + DailyActivityTable.AGENT: "agent-a", + } + other_entity_ids: Final = { + DailyActivityTable.USER: "user-b", + DailyActivityTable.TEAM: "team-b", + DailyActivityTable.TAG: "other-blue", + DailyActivityTable.ORGANIZATION: "other-org", + DailyActivityTable.CUSTOMER: "customer-b", + DailyActivityTable.AGENT: "agent-b", + } + key_rows: Final = ( + ("key-alpha", "2025-01-01", "popular", 1.0, 0), + ("key-alpha", "2025-01-02", "popular", 2.0, 0), + ("key-beta", "2025-01-01", "popular", 4.0, 0), + ("key-gamma", "2025-01-01", "popular", 5.0, 0), + ("key-cache", "2025-01-01", "popular", 2.0, 20), + ("key-target", "2025-01-01", "rare-model", 0.5, 0), + ) + return tuple( + chain.from_iterable(_activity_for_entity(table, entity_id, key_rows) for table, entity_id in entity_ids.items()) + ) + tuple( + _Activity( + table=table, + entity_id=other_entity_ids[table], + date="2025-01-01", + api_key=f"key-other-{table.value}", + model="popular", + model_group="popular-group", + spend=100.0, + flat_cost=0.0, + prompt_tokens=10, + completion_tokens=5, + cache_read_input_tokens=5 if table is DailyActivityTable.USER else 0, + cache_creation_input_tokens=0, + compression_saved_tokens=0, + compression_savings_spend=0.0, + prompt_caching_savings_spend=0.0, + gateway_injected_caching_savings_spend=0.0, + autorouter_savings_spend=0.0, + api_requests=1, + successful_requests=1, + failed_requests=0, + total_response_time_ms=100, + timed_requests=1, + ) + for table in entity_ids + ) + + +def _metrics(rows: Sequence[_Activity]) -> Mapping[str, int | float]: + return { + "spend": sum(row.spend for row in rows), + "ptu_flat_cost": sum(row.flat_cost for row in rows), + "prompt_tokens": sum(row.prompt_tokens for row in rows), + "completion_tokens": sum(row.completion_tokens for row in rows), + "cache_read_input_tokens": sum(row.cache_read_input_tokens for row in rows), + "cache_creation_input_tokens": sum(row.cache_creation_input_tokens for row in rows), + "compression_saved_tokens": sum(row.compression_saved_tokens for row in rows), + "compression_savings_spend": sum(row.compression_savings_spend for row in rows), + "prompt_caching_savings_spend": sum(row.prompt_caching_savings_spend for row in rows), + "gateway_injected_caching_savings_spend": sum(row.gateway_injected_caching_savings_spend for row in rows), + "autorouter_savings_spend": sum(row.autorouter_savings_spend for row in rows), + "api_requests": sum(row.api_requests for row in rows), + "successful_requests": sum(row.successful_requests for row in rows), + "failed_requests": sum(row.failed_requests for row in rows), + "total_response_time_ms": sum(row.total_response_time_ms for row in rows), + "timed_requests": sum(row.timed_requests for row in rows), + } + + +def _grouping_row( + rows: Sequence[_Activity], + *, + date: str | None, + api_key: str | None, + group_level: int, + distinct_api_keys: int | None, +) -> GroupingSetsRow: + metric_values: Final = _metrics(rows) + return GroupingSetsRow( + date=date, + api_key=api_key, + **metric_values, + model=None, + model_group=None, + custom_llm_provider=None, + mcp_namespaced_tool_name=None, + endpoint=None, + group_level=group_level, + distinct_api_keys=distinct_api_keys, + ) + + +def _grouping_rows_for_day( + date: str, rows: Sequence[_Activity], top_keys: tuple[str, ...], distinct_api_keys: int +) -> tuple[GroupingSetsRow, ...]: + date_rows: Final = tuple(row for row in rows if row.date == date) + return ( + _grouping_row(date_rows, date=date, api_key=None, group_level=63, distinct_api_keys=distinct_api_keys), + ) + tuple( + _grouping_row( + tuple(row for row in date_rows if row.api_key == api_key), + date=date, + api_key=api_key, + group_level=31, + distinct_api_keys=None, + ) + for api_key in top_keys + if any(row.api_key == api_key for row in date_rows) + ) + + +class _FakeRepository: + def __init__(self, rows: tuple[_Activity, ...]) -> None: + self._rows: Final = rows + self.aggregated = AsyncMock(side_effect=self._aggregated) + self.key_page = AsyncMock(side_effect=self._key_page) + self.key_page_call: tuple[DailyActivityScope, int, int] | None = None + self.search_keys = AsyncMock(side_effect=self._search_keys) + self.model_top_keys = AsyncMock(side_effect=self._model_top_keys) + self.cache_leakage_keys = AsyncMock(side_effect=self._cache_leakage_keys) + self.key_metadata = AsyncMock(side_effect=self._key_metadata) + self.export_rows_error: Exception | None = None + + def _matching_rows(self, scope: DailyActivityScope) -> tuple[_Activity, ...]: + return tuple( + row + for row in self._rows + if row.table == scope.table + and scope.start_date <= row.date <= scope.end_date + and (scope.entity_ids is None or row.entity_id in scope.entity_ids) + and row.entity_id not in scope.exclude_entity_ids + and (scope.api_keys is None or row.api_key in scope.api_keys) + and (scope.model is None or row.model == scope.model) + ) + + async def _aggregated( + self, + scope: DailyActivityScope, + *, + include_entity_breakdown: bool = False, + api_key_limit: int = constants.USAGE_TOP_API_KEYS_DEFAULT, + ) -> AggregatedRows: + rows: Final = self._matching_rows(scope) + key_spend: Final = tuple( + sorted( + ((key, sum(row.spend for row in rows if row.api_key == key)) for key in {row.api_key for row in rows}), + key=lambda item: (-item[1], item[0]), + ) + ) + distinct_keys: Final = len(key_spend) + top_keys: Final = tuple(key for key, _ in key_spend[:api_key_limit]) + dates: Final = tuple(sorted({row.date for row in rows})) + grouping_rows: Final = ( + _grouping_row(rows, date=None, api_key=None, group_level=127, distinct_api_keys=distinct_keys), + ) + tuple( + grouping_row + for date in dates + for grouping_row in _grouping_rows_for_day(date, rows, top_keys, distinct_keys) + ) + entity_rows: Final = ( + tuple(entity_row for date in dates for entity_row in _entity_rows_for_day(date, rows)) + if include_entity_breakdown + else () + ) + return AggregatedRows( + grouping_rows=grouping_rows, + entity_rows=entity_rows if include_entity_breakdown else None, + distinct_api_keys=distinct_keys, + ) + + async def _search_keys(self, scope: DailyActivityScope, *, search: str, limit: int) -> tuple[str, ...]: + rows: Final = self._matching_rows(scope) + return tuple(key for key in dict.fromkeys(row.api_key for row in rows) if search.casefold() in key.casefold())[ + :limit + ] + + async def _key_page(self, scope: DailyActivityScope, *, offset: int, limit: int) -> KeyPage: + self.key_page_call = (scope, offset, limit) + rows: Final = self._matching_rows(scope) + api_keys: Final = _ranked_keys(rows) + return KeyPage( + rows=tuple(_key_spend_row(api_key, rows) for api_key in api_keys[offset : offset + limit]), + total_api_keys=len(api_keys), + ) + + async def _model_top_keys( + self, scope: DailyActivityScope, *, model_group: str, by_model_group: bool, limit: int + ) -> tuple[KeySpendRow, ...]: + rows: Final = tuple( + row + for row in self._matching_rows(scope) + if (row.model_group if by_model_group else row.model) == model_group + ) + return tuple(_key_spend_row(key, rows) for key in _ranked_keys(rows)[:limit]) + + async def _cache_leakage_keys(self, scope: DailyActivityScope, *, limit: int) -> tuple[KeySpendRow, ...]: + rows: Final = tuple(row for row in self._matching_rows(scope) if row.cache_read_input_tokens > 0) + return tuple(_key_spend_row(key, rows) for key in _ranked_keys(rows)[:limit]) + + async def _key_metadata( + self, api_keys: frozenset[str], window: tuple[object, object] | None + ) -> Mapping[str, KeyMetadataRow]: + return { + key: KeyMetadataRow( + api_key=key, + key_alias=f"alias-{key}", + team_id="team-a", + user_id="user-a", + user_email="user@example.test", + key_exists=True, + tags=(), + ) + for key in api_keys + } + + async def export_rows(self, scope: DailyActivityScope, *, export_type: ExportType) -> AsyncIterator[ExportRow]: + if self.export_rows_error is not None: + raise self.export_rows_error + for row in self._matching_rows(scope): + yield ExportRow( + date=row.date, + entity_id=row.entity_id, + entity_alias="=entity", + api_key=row.api_key, + key_alias="+key", + user_id="user-a", + user_email="user@example.test", + model=row.model, + spend=row.spend, + flat_cost=row.flat_cost, + prompt_tokens=row.prompt_tokens, + completion_tokens=row.completion_tokens, + api_requests=row.api_requests, + successful_requests=row.successful_requests, + failed_requests=row.failed_requests, + cache_read_input_tokens=row.cache_read_input_tokens, + cache_creation_input_tokens=row.cache_creation_input_tokens, + ) + + +def _ranked_keys(rows: Sequence[_Activity]) -> tuple[str, ...]: + return tuple( + key + for key, _ in sorted( + ((key, sum(row.spend for row in rows if row.api_key == key)) for key in {row.api_key for row in rows}), + key=lambda item: (-item[1], item[0]), + ) + ) + + +def _key_spend_row(api_key: str, rows: Sequence[_Activity]) -> KeySpendRow: + matching: Final = tuple(row for row in rows if row.api_key == api_key) + return KeySpendRow( + api_key=api_key, + spend=sum(row.spend for row in matching), + prompt_tokens=sum(row.prompt_tokens for row in matching), + completion_tokens=sum(row.completion_tokens for row in matching), + total_tokens=sum(row.prompt_tokens + row.completion_tokens for row in matching), + api_requests=sum(row.api_requests for row in matching), + successful_requests=sum(row.successful_requests for row in matching), + failed_requests=sum(row.failed_requests for row in matching), + cache_read_input_tokens=sum(row.cache_read_input_tokens for row in matching), + cache_creation_input_tokens=sum(row.cache_creation_input_tokens for row in matching), + ) + + +def _entity_rows_for_day( + date: str, rows: Sequence[_Activity], distinct_api_keys: int | None = None +) -> tuple[EntityRollupRow, ...]: + date_rows: Final = tuple(row for row in rows if row.date == date) + entities: Final = tuple(dict.fromkeys(row.entity_id for row in date_rows)) + return tuple( + entity_row + for entity_id in entities + for entity_row in _entity_rows_for_entity(date, entity_id, date_rows, distinct_api_keys) + ) + + +def _entity_rows_for_entity( + date: str, entity_id: str, rows: Sequence[_Activity], distinct_api_keys: int | None +) -> tuple[EntityRollupRow, ...]: + entity_rows: Final = tuple(row for row in rows if row.entity_id == entity_id) + return ( + _entity_rollup( + entity_rows, + date=date, + entity_id=entity_id, + api_key=None, + api_key_rolled=1, + distinct_api_keys=distinct_api_keys, + ), + ) + tuple( + _entity_rollup( + tuple(row for row in entity_rows if row.api_key == api_key), + date=date, + entity_id=entity_id, + api_key=api_key, + api_key_rolled=0, + distinct_api_keys=None, + ) + for api_key in dict.fromkeys(row.api_key for row in entity_rows) + ) + + +def _entity_rollup( + rows: Sequence[_Activity], + *, + date: str, + entity_id: str, + api_key: str | None, + api_key_rolled: int, + distinct_api_keys: int | None, +) -> EntityRollupRow: + return EntityRollupRow( + date=date, + api_key=api_key, + **_metrics(rows), + entity_id=entity_id, + api_key_rolled=api_key_rolled, + distinct_api_keys=distinct_api_keys, + ) + + +class _PrismaTable: + def __init__(self, rows: tuple[object, ...] = ()) -> None: + self._rows: Final = rows + + async def find_many(self, *, where: Mapping[str, object] | None = None, **kwargs: object) -> tuple[object, ...]: + if where is None: + return self._rows + return tuple(row for row in self._rows if _matches(row, where)) + + async def find_unique(self, *, where: Mapping[str, object], **kwargs: object) -> object | None: + return next((row for row in self._rows if _matches(row, where)), None) + + +def _matches(row: object, where: Mapping[str, object]) -> bool: + return all( + getattr(row, field_name, None) in value["in"] + if isinstance(value, Mapping) and "in" in value + else getattr(row, field_name, None) == value + for field_name, value in where.items() + ) + + +def _prisma_client() -> object: + user: Final = LiteLLM_UserTable( + user_id="user-a", + user_email="user@example.test", + user_role=LitellmUserRoles.INTERNAL_USER.value, + teams=["team-a"], + ) + team: Final = LiteLLM_TeamTable( + team_id="team-a", + team_alias="Team A", + members_with_roles=[Member(user_id="user-a", role="user")], + ) + other_user: Final = LiteLLM_UserTable( + user_id="user-b", + user_email="other-user@example.test", + user_role=LitellmUserRoles.INTERNAL_USER.value, + teams=["team-b"], + ) + other_team: Final = LiteLLM_TeamTable( + team_id="team-b", + team_alias="Team B", + members_with_roles=[Member(user_id="user-b", role="user")], + ) + db: Final = SimpleNamespace( + litellm_usertable=_PrismaTable((user, other_user)), + litellm_teamtable=_PrismaTable((team, other_team)), + litellm_verificationtoken=_PrismaTable((SimpleNamespace(token="key-alpha", user_id="user-a"),)), + litellm_organizationmembership=_PrismaTable( + ( + SimpleNamespace(user_id="user-a", organization_id="org-a", user_role="org_admin"), + SimpleNamespace(user_id="user-b", organization_id="other-org", user_role="org_admin"), + ) + ), + litellm_organizationtable=_PrismaTable( + ( + SimpleNamespace(organization_id="org-a", organization_alias="Org A"), + SimpleNamespace(organization_id="other-org", organization_alias="Other Org"), + ) + ), + litellm_endusertable=_PrismaTable( + ( + SimpleNamespace(user_id="customer-a", alias="Customer A"), + SimpleNamespace(user_id="customer-b", alias="Customer B"), + ) + ), + litellm_agentstable=_PrismaTable( + ( + SimpleNamespace(agent_id="agent-a", agent_name="Agent A", created_by="user-a"), + SimpleNamespace(agent_id="agent-b", agent_name="Agent B", created_by="user-b"), + ) + ), + ) + return SimpleNamespace(db=db, writer_db=db) + + +@pytest.fixture +def daily_activity_client() -> Iterator[tuple[TestClient, _FakeRepository]]: + repository: Final = _FakeRepository(_seeded_activity()) + prisma_client: Final = _prisma_client() + app: Final = FastAPI() + app.include_router(router) + + def resolve_auth(request: Request) -> UserAPIKeyAuth: + role: Final = LitellmUserRoles(request.headers.get("x-user-role", LitellmUserRoles.PROXY_ADMIN.value)) + user_id: Final[str | None] = request.headers.get("x-user-id") or ( + "admin" if role != LitellmUserRoles.INTERNAL_USER else None + ) + return UserAPIKeyAuth( + user_id=user_id, + user_role=role, + api_key=request.headers.get("x-api-key"), + ) + + app.dependency_overrides[get_daily_activity_repository] = lambda: repository + app.dependency_overrides[get_daily_activity_prisma_client] = lambda: prisma_client + app.dependency_overrides[user_api_key_auth] = resolve_auth + with TestClient(app) as client: + yield client, repository + + +def _entity_params(query_name: str, entity_id: str) -> dict[str, str]: + return {**_DATE_PARAMS, query_name: entity_id} + + +@pytest.mark.parametrize(("prefix", "query_name", "entity_id"), _ENTITY_CASES) +def test_aggregated_routes_return_scoped_results( + daily_activity_client: tuple[TestClient, _FakeRepository], + prefix: str, + query_name: str, + entity_id: str, +) -> None: + client, _ = daily_activity_client + response: Final = client.get( + f"{prefix}/daily/activity/aggregated", + params=_entity_params(query_name, entity_id), + ) + assert response.status_code == 200, response.text + body: Final = response.json() + assert body["metadata"]["total_spend"] == pytest.approx(14.5), response.text + assert body["metadata"]["total_api_keys"] == 5, response.text + assert len(body["results"]) == 2, response.text + + +@pytest.mark.parametrize(("prefix", "query_name", "entity_id"), _ENTITY_CASES) +def test_admin_aggregates_all_entities_when_filter_is_omitted( + daily_activity_client: tuple[TestClient, _FakeRepository], prefix: str, query_name: str, entity_id: str +) -> None: + client, _ = daily_activity_client + response: Final = client.get( + f"{prefix}/daily/activity/aggregated", + params=_DATE_PARAMS, + ) + assert response.status_code == 200, response.text + assert response.json()["metadata"]["total_spend"] == pytest.approx(114.5), response.text + assert response.json()["metadata"]["total_api_keys"] == 6, response.text + + +@pytest.mark.parametrize(("prefix", "query_name", "entity_id"), _ENTITY_CASES) +def test_search_folds_each_entity_key_across_days( + daily_activity_client: tuple[TestClient, _FakeRepository], prefix: str, query_name: str, entity_id: str +) -> None: + client, _ = daily_activity_client + response: Final = client.get( + f"{prefix}/daily/activity/aggregated/search", + params={**_entity_params(query_name, entity_id), "search": "alpha"}, + ) + assert response.status_code == 200, response.text + assert response.json() == { + "api_keys": [ + { + "api_key": "key-alpha", + "metrics": { + "spend": 3.0, + "flat_cost": 0.0, + "prompt_tokens": 20, + "completion_tokens": 10, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 2, + "compression_saved_tokens": 4, + "compression_savings_spend": 0.2, + "prompt_caching_savings_spend": 0.4, + "gateway_injected_caching_savings_spend": 0.6, + "autorouter_savings_spend": 0.8, + "total_tokens": 30, + "successful_requests": 2, + "failed_requests": 0, + "api_requests": 2, + "total_response_time_ms": 200, + "timed_requests": 2, + }, + "metadata": { + "key_alias": "alias-key-alpha", + "team_id": "team-a", + "user_id": "user-a", + "user_email": "user@example.test", + "key_exists": True, + }, + } + ] + } + + +def test_search_finds_keys_outside_the_top_keys_limit_and_skips_empty_aggregate( + daily_activity_client: tuple[TestClient, _FakeRepository], +) -> None: + client, repository = daily_activity_client + aggregate_response: Final = client.get( + "/user/daily/activity/aggregated", + params={**_entity_params("user_id", "user-a"), "api_key_limit": 3}, + ) + assert aggregate_response.status_code == 200, aggregate_response.text + assert repository.aggregated.call_args.kwargs["api_key_limit"] == 3 + top_keys: Final = frozenset( + chain.from_iterable(result["breakdown"]["api_keys"] for result in aggregate_response.json()["results"]) + ) + assert "key-target" not in top_keys + + search_response: Final = client.get( + "/user/daily/activity/aggregated/search", + params={**_entity_params("user_id", "user-a"), "search": "target", "limit": 7}, + ) + assert search_response.status_code == 200, search_response.text + assert repository.search_keys.call_args.kwargs["limit"] == 7 + assert search_response.json()["api_keys"][0]["api_key"] == "key-target" + assert search_response.json()["api_keys"][0]["metrics"]["spend"] == pytest.approx(0.5) + search_metrics: Final = search_response.json()["api_keys"][0]["metrics"] + assert set(search_metrics) == set(SpendMetrics.model_fields) + assert search_metrics["compression_savings_spend"] == pytest.approx(0.1) + assert search_metrics["total_response_time_ms"] == 100 + + repository.aggregated.reset_mock() + empty_response: Final = client.get( + "/user/daily/activity/aggregated/search", + params={**_entity_params("user_id", "user-a"), "search": "absent"}, + ) + assert empty_response.status_code == 200, empty_response.text + assert empty_response.json() == {"api_keys": []} + repository.aggregated.assert_not_awaited() + + +@pytest.mark.parametrize(("prefix", "query_name", "entity_id"), _ENTITY_CASES) +def test_key_page_routes_map_ranked_rows_and_totals( + daily_activity_client: tuple[TestClient, _FakeRepository], + prefix: str, + query_name: str, + entity_id: str, +) -> None: + client, repository = daily_activity_client + response: Final = client.get( + f"{prefix}/daily/activity/aggregated/keys", + params={**_entity_params(query_name, entity_id), "offset": 1, "limit": 2}, + ) + + assert response.status_code == 200, response.text + assert response.json() == { + "api_keys": [ + { + "api_key": "key-beta", + "metrics": { + "spend": 4.0, + "prompt_tokens": 10, + "completion_tokens": 5, + "total_tokens": 15, + "api_requests": 1, + "successful_requests": 1, + "failed_requests": 0, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 1, + }, + "metadata": { + "key_alias": "alias-key-beta", + "team_id": "team-a", + "user_id": "user-a", + "user_email": "user@example.test", + "key_exists": True, + }, + }, + { + "api_key": "key-alpha", + "metrics": { + "spend": 3.0, + "prompt_tokens": 20, + "completion_tokens": 10, + "total_tokens": 30, + "api_requests": 2, + "successful_requests": 2, + "failed_requests": 0, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 2, + }, + "metadata": { + "key_alias": "alias-key-alpha", + "team_id": "team-a", + "user_id": "user-a", + "user_email": "user@example.test", + "key_exists": True, + }, + }, + ], + "total_api_keys": 5, + "offset": 1, + "limit": 2, + } + key_page_call: Final = repository.key_page_call + assert key_page_call is not None + scope: Final = key_page_call[0] + assert scope.entity_ids == (entity_id,) + assert key_page_call[1:] == (1, 2) + + +@pytest.mark.parametrize("params", ({"limit": 101}, {"offset": -1})) +def test_key_page_route_rejects_invalid_bounds( + daily_activity_client: tuple[TestClient, _FakeRepository], + params: Mapping[str, int], +) -> None: + client, repository = daily_activity_client + response: Final = client.get( + "/user/daily/activity/aggregated/keys", + params={**_entity_params("user_id", "user-a"), **params}, + ) + + assert response.status_code == 422, response.text + repository.key_page.assert_not_awaited() + + +@pytest.mark.parametrize( + ("path", "route_params", "limit_name", "invalid_limit"), + ( + ("/user/daily/activity/aggregated", {}, "api_key_limit", 0), + ( + "/user/daily/activity/aggregated", + {}, + "api_key_limit", + constants.USAGE_TOP_API_KEYS_MAX + 1, + ), + ("/user/daily/activity/aggregated/search", {"search": "key"}, "limit", 0), + ( + "/user/daily/activity/aggregated/search", + {"search": "key"}, + "limit", + constants.USAGE_KEY_SEARCH_MAX + 1, + ), + ( + "/user/daily/activity/aggregated/model_top_keys", + {"model_group": "popular-group"}, + "limit", + 0, + ), + ( + "/user/daily/activity/aggregated/model_top_keys", + {"model_group": "popular-group"}, + "limit", + constants.USAGE_MODEL_TOP_KEYS_MAX + 1, + ), + ( + "/user/daily/activity/aggregated/cache_leakage_keys", + {}, + "limit", + 0, + ), + ( + "/user/daily/activity/aggregated/cache_leakage_keys", + {}, + "limit", + constants.USAGE_CACHE_LEAKAGE_KEYS_MAX + 1, + ), + ), +) +def test_usage_limit_routes_reject_values_outside_bounds( + daily_activity_client: tuple[TestClient, _FakeRepository], + path: str, + route_params: Mapping[str, str], + limit_name: str, + invalid_limit: int, +) -> None: + client, repository = daily_activity_client + response: Final = client.get( + path, + params={ + **_entity_params("user_id", "user-a"), + **route_params, + limit_name: invalid_limit, + }, + ) + + assert response.status_code == 422, response.text + repository.aggregated.assert_not_awaited() + repository.search_keys.assert_not_awaited() + repository.model_top_keys.assert_not_awaited() + repository.cache_leakage_keys.assert_not_awaited() + + +@pytest.mark.parametrize(("prefix", "query_name", "entity_id"), _ENTITY_CASES) +def test_model_top_routes_rank_keys_and_include_metadata( + daily_activity_client: tuple[TestClient, _FakeRepository], prefix: str, query_name: str, entity_id: str +) -> None: + client, repository = daily_activity_client + response: Final = client.get( + f"{prefix}/daily/activity/aggregated/model_top_keys", + params={ + **_entity_params(query_name, entity_id), + "model_group": "rare-group", + "limit": 3, + }, + ) + assert response.status_code == 200, response.text + assert repository.model_top_keys.call_args.kwargs["limit"] == 3 + assert response.json()["model"] == "rare-group" + assert response.json()["by_model_group"] is True + assert tuple(row["api_key"] for row in response.json()["api_keys"]) == ("key-target",) + metrics: Final = response.json()["api_keys"][0]["metrics"] + expected_row: Final = KeySpendRow( + api_key="key-target", + spend=0.5, + prompt_tokens=10, + completion_tokens=5, + total_tokens=15, + api_requests=1, + successful_requests=1, + failed_requests=0, + cache_read_input_tokens=0, + cache_creation_input_tokens=1, + ) + assert set(metrics) == set(KeySpendMetrics.model_fields) + assert metrics == {field: getattr(expected_row, field) for field in KeySpendMetrics.model_fields} + + +@pytest.mark.parametrize(("prefix", "query_name", "entity_id"), _ENTITY_CASES) +def test_export_routes_stream_csv_and_preserve_row_counts( + daily_activity_client: tuple[TestClient, _FakeRepository], + prefix: str, + query_name: str, + entity_id: str, +) -> None: + client, _ = daily_activity_client + response: Final = client.get( + f"{prefix}/daily/activity/export", + params={**_entity_params(query_name, entity_id), "export_type": ExportType.DAILY.value}, + ) + assert response.status_code == 200, response.text + assert response.headers["cache-control"] == "no-store" + assert "attachment;" in response.headers["content-disposition"] + records: Final = tuple(csv.reader(io.StringIO(response.text))) + assert tuple(records[0]) == tuple(field.name for field in fields(ExportRow)) + assert len(records) == 7 + assert records[1][2] == "'=entity" + assert records[1][4] == "'+key" + + +@pytest.mark.parametrize(("prefix", "query_name", "entity_id"), _ENTITY_CASES) +def test_export_routes_stream_json_arrays( + daily_activity_client: tuple[TestClient, _FakeRepository], + prefix: str, + query_name: str, + entity_id: str, +) -> None: + client, _ = daily_activity_client + response: Final = client.get( + f"{prefix}/daily/activity/export", + params={**_entity_params(query_name, entity_id), "format": "json"}, + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("application/json") + assert response.headers["cache-control"] == "no-store" + records: Final = response.json() + assert isinstance(records, list) and len(records) == 6, response.text + assert records[0]["entity_alias"] == "=entity" + + +def test_export_first_row_error_returns_json_error_before_streaming( + daily_activity_client: tuple[TestClient, _FakeRepository], +) -> None: + client, repository = daily_activity_client + repository.export_rows_error = RuntimeError("database query failed") + response: Final = client.get( + "/team/daily/activity/export", + params={**_entity_params("team_ids", "team-a"), "format": "csv"}, + ) + assert response.status_code >= 400 + assert response.headers["content-type"].startswith("application/json") + assert response.text != ",".join(field.name for field in fields(ExportRow)) + "\r\n" + assert "database query failed" in response.text + + +def test_csv_export_with_no_rows_contains_only_header( + daily_activity_client: tuple[TestClient, _FakeRepository], +) -> None: + client, _ = daily_activity_client + response: Final = client.get( + "/team/daily/activity/export", + params={**_entity_params("team_ids", "team-a"), "api_key": "missing-key"}, + ) + assert response.status_code == 200, response.text + assert response.text == ",".join(field.name for field in fields(ExportRow)) + "\r\n" + + +def test_json_export_with_no_rows_is_an_empty_array( + daily_activity_client: tuple[TestClient, _FakeRepository], +) -> None: + client, _ = daily_activity_client + response: Final = client.get( + "/team/daily/activity/export", + params={**_entity_params("team_ids", "team-a"), "api_key": "missing-key", "format": "json"}, + ) + assert response.status_code == 200, response.text + assert response.json() == [] + + +def test_user_routes_preserve_scope_denials_and_service_account_guard( + daily_activity_client: tuple[TestClient, _FakeRepository], +) -> None: + client, repository = daily_activity_client + denied: Final = client.get( + "/user/daily/activity/aggregated", + params=_entity_params("user_id", "user-b"), + headers={"x-user-role": LitellmUserRoles.INTERNAL_USER.value, "x-user-id": "user-a"}, + ) + assert denied.status_code == 403, denied.text + repository.aggregated.assert_not_awaited() + + service_account: Final = client.get( + "/user/daily/activity/aggregated", + params=_DATE_PARAMS, + headers={"x-user-role": LitellmUserRoles.INTERNAL_USER.value}, + ) + assert service_account.status_code == 403, service_account.text + repository.aggregated.assert_not_awaited() + + own_scope: Final = client.get( + "/user/daily/activity/aggregated", + params=_DATE_PARAMS, + headers={"x-user-role": LitellmUserRoles.INTERNAL_USER.value, "x-user-id": "user-a"}, + ) + assert own_scope.status_code == 200, own_scope.text + assert own_scope.json()["metadata"]["total_spend"] == pytest.approx(14.5) + assert own_scope.json()["metadata"]["total_api_keys"] == 5 + + +def test_team_scope_applies_membership_and_user_key_filter( + daily_activity_client: tuple[TestClient, _FakeRepository], +) -> None: + client, repository = daily_activity_client + response: Final = client.get( + "/team/daily/activity/aggregated", + params={**_entity_params("team_ids", "team-a"), "timezone": "480"}, + headers={"x-user-role": LitellmUserRoles.INTERNAL_USER.value, "x-user-id": "user-a"}, + ) + assert response.status_code == 200, response.text + assert response.json()["metadata"]["total_spend"] == pytest.approx(3.0) + assert response.json()["metadata"]["total_api_keys"] == 1 + scope: Final = repository.aggregated.call_args.args[0] + assert scope.api_keys == ("key-alpha",) + assert scope.entity_ids == ("team-a",) + assert scope.timezone_offset_minutes == 480 + assert repository.aggregated.call_args.kwargs["include_entity_breakdown"] is True + + +def test_team_scope_does_not_allow_an_unowned_api_key( + daily_activity_client: tuple[TestClient, _FakeRepository], +) -> None: + client, _ = daily_activity_client + response: Final = client.get( + "/team/daily/activity/aggregated", + params={**_entity_params("team_ids", "team-a"), "api_key": "key-beta"}, + headers={"x-user-role": LitellmUserRoles.INTERNAL_USER.value, "x-user-id": "user-a"}, + ) + assert response.status_code == 200, response.text + assert response.json()["metadata"]["total_spend"] == 0 + assert response.json()["metadata"]["total_api_keys"] == 0 + assert "key-beta" not in response.text + + +def _assert_customer_route_denied(client: TestClient, prefix: str, suffix: str, extra_params: dict[str, str]) -> None: + response: Final = client.get( + f"{prefix}/daily/activity/{suffix}", + params={**_entity_params("end_user_ids", "customer-a"), **extra_params}, + headers={"x-user-role": LitellmUserRoles.INTERNAL_USER.value, "x-user-id": "user-a"}, + ) + assert response.status_code == 403, response.text + + +def test_customer_service_routes_deny_non_admins(daily_activity_client: tuple[TestClient, _FakeRepository]) -> None: + client, repository = daily_activity_client + route_params: Final = ( + ("aggregated", {}), + ("aggregated/search", {"search": "alpha"}), + ("aggregated/model_top_keys", {"model_group": "rare-group"}), + ("export", {"export_type": ExportType.DAILY.value}), + ) + for prefix in ("/customer", "/end_user"): + for suffix, extra_params in route_params: + _assert_customer_route_denied(client, prefix, suffix, extra_params) + repository.aggregated.assert_not_awaited() + + +def test_customer_end_user_aliases_are_hidden_from_openapi( + daily_activity_client: tuple[TestClient, _FakeRepository], +) -> None: + client, _ = daily_activity_client + paths: Final = client.get("/openapi.json").json()["paths"] + assert "/customer/daily/activity/aggregated" in paths + assert "/end_user/daily/activity/aggregated" not in paths + response: Final = client.get( + "/end_user/daily/activity/aggregated", + params=_entity_params("end_user_ids", "customer-a"), + ) + assert response.status_code == 200, response.text + + +@pytest.mark.parametrize( + ("prefix", "query_name", "entity_id"), + _ENTITY_CASES[:4] + (_ENTITY_CASES[-1],), +) +@pytest.mark.parametrize( + ("family", "extra_params"), + ( + ("aggregated", {}), + ("aggregated/search", {"search": "key"}), + ("aggregated/model_top_keys", {"model_group": "popular-group"}), + ("export", {"export_type": ExportType.DAILY.value}), + ), +) +def test_non_admin_routes_return_only_permitted_entities_and_keys( + daily_activity_client: tuple[TestClient, _FakeRepository], + prefix: str, + query_name: str, + entity_id: str, + family: str, + extra_params: Mapping[str, str], +) -> None: + client, _ = daily_activity_client + headers: Final = {"x-user-role": LitellmUserRoles.INTERNAL_USER.value, "x-user-id": "user-a"} + response: Final = client.get( + f"{prefix}/daily/activity/{family}", + params={**_entity_params(query_name, entity_id), **extra_params}, + headers=headers, + ) + assert response.status_code == 200, response.text + assert "key-other-" not in response.text, response.text + if prefix in ("/team", "/tag"): + assert "key-beta" not in response.text, response.text + assert "key-alpha" in response.text, response.text + + +def test_empty_scope_filters_fail_closed(daily_activity_client: tuple[TestClient, _FakeRepository]) -> None: + client, _ = daily_activity_client + response: Final = client.get( + "/tag/daily/activity/aggregated", + params=_entity_params("tags", "blue"), + headers={"x-user-role": LitellmUserRoles.INTERNAL_USER.value, "x-user-id": "user-empty"}, + ) + assert response.status_code == 200, response.text + assert response.json()["metadata"]["total_spend"] == 0 + assert response.json()["results"] == [] + + +@pytest.mark.parametrize( + ("prefix", "query_name", "entity_id", "other_entity_id", "exclude_query_name"), + ( + ("/team", "team_ids", "team-a", "team-b", "exclude_team_ids"), + ("/organization", "organization_ids", "org-a", "other-org", "exclude_organization_ids"), + ("/customer", "end_user_ids", "customer-a", "customer-b", "exclude_end_user_ids"), + ("/agent", "agent_ids", "agent-a", "agent-b", "exclude_agent_ids"), + ), +) +def test_exclusion_filters_apply_after_entity_scope( + daily_activity_client: tuple[TestClient, _FakeRepository], + prefix: str, + query_name: str, + entity_id: str, + other_entity_id: str, + exclude_query_name: str, +) -> None: + client, _ = daily_activity_client + response: Final = client.get( + f"{prefix}/daily/activity/aggregated", + params={ + **_DATE_PARAMS, + query_name: f"{entity_id},{other_entity_id}", + exclude_query_name: entity_id, + }, + ) + assert response.status_code == 200, response.text + assert response.json()["metadata"]["total_spend"] == pytest.approx(100) + assert response.json()["metadata"]["total_api_keys"] == 1 + + +def test_csv_formula_escaping_covers_all_supported_leading_characters() -> None: + dangerous_values: Final = ("=sum", "+sum", "-sum", "@sum", "\tsum", "\rsum") + assert tuple(_csv_cell(value) for value in dangerous_values) == tuple(f"'{value}" for value in dangerous_values) + + +def test_user_cache_leakage_route_returns_cache_keys_and_metadata( + daily_activity_client: tuple[TestClient, _FakeRepository], +) -> None: + client, repository = daily_activity_client + response: Final = client.get( + "/user/daily/activity/aggregated/cache_leakage_keys", + params={**_entity_params("user_id", "user-a"), "limit": 4}, + ) + assert response.status_code == 200, response.text + assert repository.cache_leakage_keys.call_args.kwargs["limit"] == 4 + assert tuple(row["api_key"] for row in response.json()["api_keys"]) == ("key-cache",) + metrics: Final = response.json()["api_keys"][0]["metrics"] + expected_row: Final = KeySpendRow( + api_key="key-cache", + spend=2.0, + prompt_tokens=10, + completion_tokens=5, + total_tokens=15, + api_requests=1, + successful_requests=1, + failed_requests=0, + cache_read_input_tokens=20, + cache_creation_input_tokens=1, + ) + assert set(metrics) == set(KeySpendMetrics.model_fields) + assert metrics == {field: getattr(expected_row, field) for field in KeySpendMetrics.model_fields} + assert response.json()["api_keys"][0]["metadata"]["key_alias"] == "alias-key-cache" + repository.cache_leakage_keys.assert_awaited_once() + repository.key_metadata.assert_awaited_once() + assert repository.key_metadata.call_args.args[0] == frozenset(("key-cache",)) + assert repository.key_metadata.call_args.args[1] is not None + + +def test_user_cache_leakage_route_respects_requested_user_scope( + daily_activity_client: tuple[TestClient, _FakeRepository], +) -> None: + client, repository = daily_activity_client + response: Final = client.get( + "/user/daily/activity/aggregated/cache_leakage_keys", + params=_entity_params("user_id", "user-missing"), + ) + assert response.status_code == 200, response.text + assert response.json() == {"api_keys": []} + scope: Final = repository.cache_leakage_keys.call_args.args[0] + assert scope.entity_ids == ("user-missing",) + + +def test_export_json_stream_has_all_seeded_rows(daily_activity_client: tuple[TestClient, _FakeRepository]) -> None: + client, _ = daily_activity_client + response: Final = client.get( + "/user/daily/activity/export", + params={ + **_entity_params("user_id", "user-a"), + "export_type": ExportType.DAILY.value, + "format": "json", + }, + ) + assert response.status_code == 200, response.text + assert response.headers["content-type"].startswith("application/json") + assert len(response.json()) == 6 + + +def test_user_internal_role_is_scoped_to_api_key(daily_activity_client: tuple[TestClient, _FakeRepository]) -> None: + client, _ = daily_activity_client + response: Final = client.get( + "/tag/daily/activity/aggregated", + params=_entity_params("tags", "blue"), + headers={ + "x-user-role": LitellmUserRoles.INTERNAL_USER.value, + "x-user-id": "user-a", + }, + ) + assert response.status_code == 200, response.text + assert response.json()["metadata"]["total_spend"] == pytest.approx(3.0) + + +def test_user_aggregate_keeps_current_day_query_semantics( + daily_activity_client: tuple[TestClient, _FakeRepository], +) -> None: + client, repository = daily_activity_client + response: Final = client.get( + "/user/daily/activity/aggregated", + params={**_DATE_PARAMS, "user_id": "user-a", "timezone": 480, "include_current_utc_day": "true"}, + ) + assert response.status_code == 200, response.text + assert response.json()["metadata"]["total_spend"] == pytest.approx(14.5) + scope: Final = repository.aggregated.call_args.args[0] + assert scope.entity_ids == ("user-a",) + assert scope.timezone_offset_minutes == 480 + assert scope.include_current_utc_day is True + + +@pytest.mark.parametrize( + ("start_date", "end_date", "message"), + ( + ("2020-01-01", "2026-12-31", "at most 400 days"), + ("0000-01-01", "9999-12-31", "valid YYYY-MM-DD"), + ("2024-06-01", "2024-01-01", "on or after"), + ("not-a-date", "2024-01-31", "valid YYYY-MM-DD"), + ("2026-9-24", "2026-09-26", "valid YYYY-MM-DD"), + ("2026-09-24", "2026-09-26", "valid YYYY-MM-DD"), + ("2026-09-01", "2026-09-4", "valid YYYY-MM-DD"), + (None, "2024-01-31", "start_date and end_date"), + ), +) +def test_team_aggregated_route_rejects_bad_date_ranges( + daily_activity_client: tuple[TestClient, _FakeRepository], + start_date: str | None, + end_date: str | None, + message: str, +) -> None: + client, repository = daily_activity_client + response: Final = client.get( + "/team/daily/activity/aggregated", + params={"start_date": start_date, "end_date": end_date, "team_ids": "team-a"}, + ) + assert response.status_code == 400, response.text + assert message in str(response.json()["detail"]), response.text + repository.aggregated.assert_not_awaited() + + +def test_user_aggregate_rejects_missing_dates(daily_activity_client: tuple[TestClient, _FakeRepository]) -> None: + client, repository = daily_activity_client + missing_dates: Final = client.get("/user/daily/activity/aggregated", params={"user_id": "user-a"}) + assert missing_dates.status_code == 400, missing_dates.text + assert missing_dates.json()["detail"] == {"error": "Please provide start_date and end_date"} + repository.aggregated.assert_not_awaited() + + +@pytest.mark.parametrize( + ("start_date", "end_date", "message"), + ( + ("2020-01-01", "2026-12-31", "at most 400 days"), + ("not-a-date", "2024-01-31", "valid YYYY-MM-DD"), + ("2024-06-01", "2024-01-01", "on or after"), + ), +) +def test_user_key_page_rejects_bad_date_ranges( + daily_activity_client: tuple[TestClient, _FakeRepository], + start_date: str, + end_date: str, + message: str, +) -> None: + client, repository = daily_activity_client + response: Final = client.get( + "/user/daily/activity/aggregated/keys", + params={"start_date": start_date, "end_date": end_date, "user_id": "user-a"}, + ) + assert response.status_code == 400, response.text + assert message in str(response.json()["detail"]), response.text + repository.key_page.assert_not_awaited() + + +_NON_CANONICAL_DATE_RANGES: Final[tuple[tuple[str, str], ...]] = ( + ("2026-9-24", "2026-09-26"), + ("2026-09-24", "2026-09-26"), + ("2026-09-01", "2026-09-4"), +) + + +@pytest.mark.parametrize(("start_date", "end_date"), _NON_CANONICAL_DATE_RANGES) +def test_user_aggregate_rejects_non_canonical_dates( + daily_activity_client: tuple[TestClient, _FakeRepository], start_date: str, end_date: str +) -> None: + client, repository = daily_activity_client + response: Final = client.get( + "/user/daily/activity/aggregated", + params={"start_date": start_date, "end_date": end_date, "user_id": "user-a"}, + ) + assert response.status_code == 400, response.text + assert response.json()["detail"] == {"error": "start_date and end_date must be valid YYYY-MM-DD dates"} + repository.aggregated.assert_not_awaited() + + +def test_user_aggregate_still_accepts_ranges_wider_than_the_team_limit( + daily_activity_client: tuple[TestClient, _FakeRepository], +) -> None: + client, repository = daily_activity_client + response: Final = client.get( + "/user/daily/activity/aggregated", + params={"start_date": "2020-01-01", "end_date": "2026-12-31", "user_id": "user-a"}, + ) + assert response.status_code == 200, response.text + repository.aggregated.assert_awaited_once() + + +@pytest.mark.parametrize(("start_date", "end_date"), _NON_CANONICAL_DATE_RANGES) +@pytest.mark.parametrize(("prefix", "query_name", "entity_id"), _ENTITY_CASES) +def test_export_routes_reject_non_canonical_dates_before_querying( + daily_activity_client: tuple[TestClient, _FakeRepository], + prefix: str, + query_name: str, + entity_id: str, + start_date: str, + end_date: str, +) -> None: + client, repository = daily_activity_client + repository.export_rows_error = AssertionError("export must not query the repository") + response: Final = client.get( + f"{prefix}/daily/activity/export", + params={query_name: entity_id, "start_date": start_date, "end_date": end_date, "export_type": "daily"}, + ) + assert response.status_code == 400, response.text + assert response.json()["detail"] == {"error": "start_date and end_date must be valid YYYY-MM-DD dates"} + assert "content-disposition" not in response.headers + + +def test_export_content_disposition_is_ascii_and_built_from_canonical_dates( + daily_activity_client: tuple[TestClient, _FakeRepository], +) -> None: + client, _ = daily_activity_client + response: Final = client.get( + "/team/daily/activity/export", + params={**_entity_params("team_ids", "team-a"), "export_type": ExportType.DAILY.value}, + ) + assert response.status_code == 200, response.text + disposition: Final = response.headers["content-disposition"] + assert disposition == 'attachment; filename="team-usage-2025-01-01-2025-01-02-daily.csv"' + assert disposition.isascii() diff --git a/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py index d6be455a321..799fe59147c 100644 --- a/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py @@ -15,7 +15,6 @@ from fastapi.testclient import TestClient from pytest_mock import MockerFixture from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler - from litellm.proxy._types import ( LiteLLM_UserTableFiltered, LitellmUserRoles, @@ -26,12 +25,10 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.management_endpoints.internal_user_endpoints import ( - LiteLLM_UserTableWithKeyCount, _authorize_user_list_request, _resolve_org_filter_for_user_search, _resolve_user_email_metadata, _update_internal_user_params, - get_user_key_counts, get_users, new_user, ui_view_users, @@ -2480,184 +2477,6 @@ async def test_get_user_daily_activity_rejects_service_account_caller(monkeypatc mock_get_daily.assert_not_called() -@pytest.mark.asyncio -async def test_get_user_daily_activity_aggregated_rejects_service_account_caller( - monkeypatch, -): - """ - Same security regression as - test_get_user_daily_activity_rejects_service_account_caller, on the - aggregated route. Same shape, raw-SQL builder, same fix. - """ - from unittest.mock import AsyncMock, MagicMock - - from fastapi import HTTPException - - from litellm.proxy.management_endpoints.internal_user_endpoints import ( - get_user_daily_activity_aggregated, - ) - - mock_prisma_client = MagicMock() - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - - mock_get_daily_agg = AsyncMock() - monkeypatch.setattr( - "litellm.proxy.management_endpoints.internal_user_endpoints.get_daily_activity_aggregated", - mock_get_daily_agg, - ) - - service_account_key = UserAPIKeyAuth( - user_id=None, - user_role=LitellmUserRoles.INTERNAL_USER, - ) - - with pytest.raises(HTTPException) as exc_info: - await get_user_daily_activity_aggregated( - start_date="2025-01-01", - end_date="2025-01-31", - model=None, - api_key=None, - user_id=None, - timezone=None, - user_api_key_dict=service_account_key, - ) - - assert exc_info.value.status_code == 403 - assert "Service-account keys" in str(exc_info.value.detail) - mock_get_daily_agg.assert_not_called() - - -@pytest.mark.asyncio -@pytest.mark.parametrize("include_current_utc_day", [False, True]) -async def test_get_user_daily_activity_aggregated_admin_global_view(monkeypatch, include_current_utc_day): - """ - Test that admin users can call the aggregated endpoint without a user_id - to get a global view. Also verifies that the correct arguments are forwarded - to the underlying get_daily_activity_aggregated helper. - """ - from unittest.mock import AsyncMock, MagicMock - - from litellm.proxy.management_endpoints.internal_user_endpoints import ( - get_user_daily_activity_aggregated, - ) - - # Mock the prisma client - mock_prisma_client = MagicMock() - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - - # Mock the downstream helper so we don't need a real DB - mock_response = MagicMock() - mock_get_daily_agg = AsyncMock(return_value=mock_response) - monkeypatch.setattr( - "litellm.proxy.management_endpoints.internal_user_endpoints.get_daily_activity_aggregated", - mock_get_daily_agg, - ) - - # Admin caller - admin_key_dict = UserAPIKeyAuth( - user_id="admin-user-001", - user_role=LitellmUserRoles.PROXY_ADMIN, - ) - - # Admin calls without user_id → global view (entity_id=None) - result = await get_user_daily_activity_aggregated( - start_date="2025-02-01", - end_date="2025-02-28", - model="gpt-4", - api_key=None, - user_id=None, - timezone=480, - include_current_utc_day=include_current_utc_day, - user_api_key_dict=admin_key_dict, - ) - - assert result is mock_response - - # Verify the helper was called with the right parameters - mock_get_daily_agg.assert_called_once_with( - prisma_client=mock_prisma_client, - table_name="litellm_dailyuserspend", - entity_id_field="user_id", - entity_id=None, # global view: no user_id filter - entity_metadata_field=None, - start_date="2025-02-01", - end_date="2025-02-28", - model="gpt-4", - api_key=None, - timezone_offset_minutes=480, - include_current_utc_day=include_current_utc_day, - ) - - -@pytest.mark.asyncio -async def test_get_user_daily_activity_aggregated_non_admin_cannot_view_other_users( - monkeypatch, -): - """ - Same scoping contract as - test_get_user_daily_activity_non_admin_cannot_view_other_users, on the - aggregated route. Non-admins reach this handler now that the route is in - self_managed_routes, so the 403-on-mismatch and default-to-self behaviour - has to hold here too: opening the route must not widen access. - """ - from unittest.mock import AsyncMock, MagicMock, patch - - from fastapi import HTTPException - - from litellm.proxy.management_endpoints.internal_user_endpoints import ( - get_user_daily_activity_aggregated, - ) - - mock_prisma_client = MagicMock() - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) - - non_admin_key_dict = UserAPIKeyAuth( - user_id="regular-user-123", - user_role=LitellmUserRoles.INTERNAL_USER, - ) - - # Case 1: Non-admin targets another user's data — 403, helper never reached - with patch( - "litellm.proxy.management_endpoints.internal_user_endpoints.get_daily_activity_aggregated", - new_callable=AsyncMock, - ) as mock_get_daily_agg: - with pytest.raises(HTTPException) as exc_info: - await get_user_daily_activity_aggregated( - start_date="2025-01-01", - end_date="2025-01-31", - model=None, - api_key=None, - user_id="other-user-456", - timezone=None, - user_api_key_dict=non_admin_key_dict, - ) - - assert exc_info.value.status_code == 403 - assert "Non-admin users can only view their own spend data" in str(exc_info.value.detail) - mock_get_daily_agg.assert_not_called() - - # Case 2: Non-admin omits user_id — scoped to their own user_id, not global - mock_response = MagicMock() - with patch( - "litellm.proxy.management_endpoints.internal_user_endpoints.get_daily_activity_aggregated", - new_callable=AsyncMock, - return_value=mock_response, - ) as mock_get_daily_agg: - result = await get_user_daily_activity_aggregated( - start_date="2025-01-01", - end_date="2025-01-31", - model=None, - api_key=None, - user_id=None, - timezone=None, - user_api_key_dict=non_admin_key_dict, - ) - - assert result is mock_response - mock_get_daily_agg.assert_called_once() - assert mock_get_daily_agg.call_args.kwargs["entity_id"] == "regular-user-123" - - @pytest.mark.asyncio async def test_delete_user_cleans_up_created_by_invitation_links(mocker): """ diff --git a/tests/unit/proxy/management_endpoints/test_organization_endpoints.py b/tests/unit/proxy/management_endpoints/test_organization_endpoints.py index 48586874788..440f93d1387 100644 --- a/tests/unit/proxy/management_endpoints/test_organization_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_organization_endpoints.py @@ -1292,7 +1292,7 @@ async def test_get_organization_daily_activity_non_admin_without_org_admin_role_ ) assert get_daily_activity_mock.call_args.kwargs["entity_id"] == [] - assert org_table_find_many.call_args.kwargs["where"] == {"organization_id": {"in": []}} + org_table_find_many.assert_not_awaited() @pytest.mark.asyncio diff --git a/tests/unit/proxy/management_endpoints/test_team_endpoints.py b/tests/unit/proxy/management_endpoints/test_team_endpoints.py index 32b3c919b07..3e71cc70099 100644 --- a/tests/unit/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_team_endpoints.py @@ -1,10 +1,10 @@ import asyncio import json +from collections.abc import Sequence from contextlib import AbstractContextManager, asynccontextmanager, contextmanager from dataclasses import dataclass from datetime import datetime, timezone from types import SimpleNamespace -from collections.abc import Sequence from typing import Final, Optional, cast from unittest.mock import AsyncMock, MagicMock, PropertyMock, call, patch @@ -53,6 +53,7 @@ from litellm.proxy.management_endpoints.team_endpoints import ( _update_model_table, _validate_and_populate_member_user_info, _validate_team_member_reset_spend_value, + aggregated_date_range_error, delete_team, list_available_teams, reset_team_member_budget_fn, @@ -13596,6 +13597,63 @@ async def test_team_member_add_audits_a_user_created_from_a_list_payload(monkeyp assert mock_audit.call_args.kwargs["team_alias"] == "list-audit" +@pytest.mark.asyncio +async def test_team_member_add_evicts_the_cached_team_roster(monkeypatch): + """Roster checks read the team through get_team_object, so a cached pre-add roster must be dropped.""" + from litellm.proxy._types import TeamMemberAddRequest + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.management_endpoints.team_endpoints import team_member_add + + team_id = "team-roster-evict" + team_row = LiteLLM_TeamTable(team_id=team_id, team_alias="roster-evict", members_with_roles=[]) + cache = UserApiKeyCache() + cache.set_cache(key=f"team_id:{team_id}", value=team_row) + cache.set_cache(key="team_alias:roster-evict", value=team_row) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", AsyncMock()) + monkeypatch.setattr("litellm.proxy.proxy_server.user_api_key_cache", cache) + monkeypatch.setattr("litellm.proxy.proxy_server.proxy_logging_obj", None) + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + monkeypatch.setattr("litellm.proxy.proxy_server.litellm_proxy_admin_name", "default_user_id") + + joined_user = LiteLLM_UserTable(user_id="joiner", max_budget=None, spend=0.0, models=[]) + updated_team = MagicMock() + updated_team.model_dump.return_value = {"team_id": team_id, "members_with_roles": []} + + with ( + patch( + "litellm.proxy.management_endpoints.team_endpoints.get_team_object", + new_callable=AsyncMock, + return_value=team_row, + ), + patch( + "litellm.proxy.management_endpoints.team_endpoints._validate_team_member_add_permissions", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.team_endpoints._validate_and_populate_member_user_info", + new_callable=AsyncMock, + ), + patch( + "litellm.proxy.management_endpoints.team_endpoints._resolve_existing_member_user_ids", + new_callable=AsyncMock, + return_value=frozenset(), + ), + patch( + "litellm.proxy.management_endpoints.team_endpoints._add_team_members_to_team", + new_callable=AsyncMock, + return_value=(updated_team, [joined_user], []), + ), + patch("litellm.proxy.management_endpoints.team_endpoints._schedule_team_member_add_audit_logs"), + ): + await team_member_add( + data=TeamMemberAddRequest(team_id=team_id, member=Member(user_id="joiner", role="user")), + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin-1"), + ) + + assert cache.get_cache(key=f"team_id:{team_id}") is None + assert cache.get_cache(key="team_alias:roster-evict") is None + + class _RecordingAuditLogger(CustomLogger): def __init__(self) -> None: super().__init__() @@ -14545,127 +14603,6 @@ async def test_new_team_batch_enqueued_token_limit_rejected_for_non_admin(): assert "on a team" in str(exc.value.message) -@pytest.mark.asyncio -async def test_get_team_daily_activity_aggregated_scopes_and_flags(mock_db_client): - """The aggregated endpoint must apply the same non-admin key scoping as the - paginated one and request the per-team entity breakdown with the caller's - timezone, so the Team Usage UI gets every day in one response.""" - from litellm.proxy.management_endpoints.team_endpoints import ( - get_team_daily_activity_aggregated, - ) - - user_id = "test_user_123" - team_id = "test_team_456" - user_api_key_dict = UserAPIKeyAuth( - user_id=user_id, user_role=LitellmUserRoles.INTERNAL_USER - ) - - mock_user_info = LiteLLM_UserTable( - user_id=user_id, - teams=[team_id], - max_budget=1000.0, - spend=0.0, - user_email="test@example.com", - user_role="internal_user", - ) - - mock_team_member = Member(user_id=user_id, role="user") - mock_team = MagicMock(spec=LiteLLM_TeamTable) - mock_team.team_id = team_id - mock_team.team_alias = "Test Team" - mock_team.members_with_roles = [mock_team_member] - mock_team.model_dump.return_value = { - "team_id": team_id, - "team_alias": "Test Team", - "members_with_roles": [{"user_id": user_id, "role": "user"}], - } - - user_api_key_1 = MagicMock() - user_api_key_1.token = "user_key_1" - - mock_db_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[mock_team]) - mock_db_client.db.litellm_verificationtoken.find_many = AsyncMock( - return_value=[user_api_key_1] - ) - - with patch( - "litellm.proxy.management_endpoints.team_endpoints.get_user_object", - new_callable=AsyncMock, - ) as mock_get_user_object: - mock_get_user_object.return_value = mock_user_info - - with patch( - "litellm.proxy.management_endpoints.team_endpoints.get_daily_activity_aggregated", - new_callable=AsyncMock, - ) as mock_aggregated: - mock_aggregated.return_value = MagicMock() - - await get_team_daily_activity_aggregated( - team_ids=team_id, - start_date="2024-01-01", - end_date="2024-01-31", - model=None, - api_key=None, - exclude_team_ids=None, - timezone=480, - user_api_key_dict=user_api_key_dict, - ) - - mock_aggregated.assert_called_once() - call_kwargs = mock_aggregated.call_args[1] - assert call_kwargs["api_key"] == ["user_key_1"] - assert call_kwargs["entity_id"] == [team_id] - assert call_kwargs["entity_metadata_field"] == { - team_id: {"team_alias": "Test Team"} - } - assert call_kwargs["include_entity_breakdown"] is True - assert call_kwargs["timezone_offset_minutes"] == 480 - assert call_kwargs["table_name"] == "litellm_dailyteamspend" - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - "start_date,end_date,expected_error", - [ - ("2020-01-01", "2026-12-31", "at most 400 days"), - ("0000-01-01", "9999-12-31", "valid YYYY-MM-DD"), - ("2024-06-01", "2024-01-01", "on or after"), - ("not-a-date", "2024-01-31", "valid YYYY-MM-DD"), - (None, "2024-01-31", "start_date and end_date"), - ], -) -async def test_get_team_daily_activity_aggregated_rejects_bad_ranges( - mock_db_client, start_date, end_date, expected_error -): - """The aggregated endpoint has no pagination bounding its work, so an - unbounded or malformed range must 400 before any query runs.""" - from litellm.proxy.management_endpoints.team_endpoints import ( - get_team_daily_activity_aggregated, - ) - - with patch( - "litellm.proxy.management_endpoints.team_endpoints.get_daily_activity_aggregated", - new_callable=AsyncMock, - ) as mock_aggregated: - with pytest.raises(HTTPException) as exc_info: - await get_team_daily_activity_aggregated( - team_ids=None, - start_date=start_date, - end_date=end_date, - model=None, - api_key=None, - exclude_team_ids=None, - timezone=None, - user_api_key_dict=UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ), - ) - - assert exc_info.value.status_code == 400 - assert expected_error in str(exc_info.value.detail) - mock_aggregated.assert_not_called() - - def _wire_new_team_prisma(mock_db_client): mock_db_client.jsonify_team_object = lambda db_data: db_data mock_db_client.get_data = AsyncMock(return_value=None) @@ -16883,3 +16820,22 @@ def test_list_team_v2_answers_503_no_db_connection_when_the_callers_user_read_hi assert response.status_code == 503, response.text assert response.json() == _DB_OUTAGE_503_BODY + + +@pytest.mark.parametrize( + ("start_date", "end_date"), + ( + ("2026-9-24", "2026-09-26"), + ("2026-09-24", "2026-09-26"), + ("2026-09-01", "2026-09-4"), + ("2026-02-30", "2026-09-26"), + ), +) +def test_aggregated_date_range_error_rejects_non_canonical_dates(start_date: str, end_date: str) -> None: + assert aggregated_date_range_error(start_date, end_date) == "start_date and end_date must be valid YYYY-MM-DD dates" + + +def test_aggregated_date_range_error_accepts_canonical_dates_and_keeps_range_checks() -> None: + assert aggregated_date_range_error("2026-09-24", "2026-09-26") is None + assert aggregated_date_range_error("2026-09-26", "2026-09-24") == "end_date must be on or after start_date" + assert aggregated_date_range_error("2020-01-01", "2026-12-31") == "Date range must be at most 400 days" diff --git a/tests/unit/proxy/management_endpoints/test_ui_sso.py b/tests/unit/proxy/management_endpoints/test_ui_sso.py index 7db37588cad..8ff0b24982f 100644 --- a/tests/unit/proxy/management_endpoints/test_ui_sso.py +++ b/tests/unit/proxy/management_endpoints/test_ui_sso.py @@ -939,6 +939,83 @@ def test_build_sso_user_update_data_normalizes_email(): assert "user_role" not in update_data +def test_build_sso_user_update_data_fills_empty_user_alias_from_display_name(): + """ + An existing SSO user with no alias gets the IdP display name on login. + """ + from litellm.proxy.management_endpoints.types import CustomOpenID + from litellm.proxy.management_endpoints.ui_sso import _build_sso_user_update_data + + sso_result = CustomOpenID( + id="S-1-5-21-adfs-user", + email="jane.doe@example.com", + first_name="Jane", + last_name="Doe", + display_name="Doe, Jane", + provider="generic", + team_ids=[], + ) + + update_data = _build_sso_user_update_data( + result=sso_result, + user_email="jane.doe@example.com", + user_id="S-1-5-21-adfs-user", + existing_user_alias=None, + ) + + assert update_data == {"user_email": "jane.doe@example.com", "user_alias": "Doe, Jane"} + + +def test_build_sso_user_update_data_keeps_existing_user_alias(): + """ + An alias already stored for the user is never overwritten by the IdP display name. + """ + from litellm.proxy.management_endpoints.types import CustomOpenID + from litellm.proxy.management_endpoints.ui_sso import _build_sso_user_update_data + + sso_result = CustomOpenID( + id="S-1-5-21-adfs-user", + email="jane.doe@example.com", + display_name="Doe, Jane", + provider="generic", + team_ids=[], + ) + + update_data = _build_sso_user_update_data( + result=sso_result, + user_email="jane.doe@example.com", + user_id="S-1-5-21-adfs-user", + existing_user_alias="Admin-set alias", + ) + + assert update_data == {"user_email": "jane.doe@example.com"} + + +@pytest.mark.parametrize( + "result, expected_alias", + [ + ( + CustomOpenID(id="user-1", display_name="Doe, Jane", first_name="Jane", last_name="Doe", team_ids=[]), + "Doe, Jane", + ), + (CustomOpenID(id="user-1", first_name="Jane", last_name="Doe", team_ids=[]), "Jane Doe"), + (CustomOpenID(id="user-1", display_name="user-1", first_name="Jane", team_ids=[]), "Jane"), + (CustomOpenID(id="user-1", display_name="user-1", team_ids=[]), None), + (CustomOpenID(id="user-1", display_name=" ", first_name=" Jane ", last_name="Doe", team_ids=[]), "Jane Doe"), + (CustomOpenID(id="user-1", display_name=" ", first_name=" ", team_ids=[]), None), + ({"id": "user-1", "display_name": "Dict User", "first_name": None, "last_name": None}, "Dict User"), + (None, None), + ], +) +def test_get_sso_user_alias(result: CustomOpenID | dict[str, str | None] | None, expected_alias: str | None): + """ + The alias is the IdP display name unless it is just the user id, then the joined first/last name. + """ + from litellm.proxy.management_endpoints.ui_sso import _get_sso_user_alias + + assert _get_sso_user_alias(result) == expected_alias + + def test_generic_response_convertor_normalizes_email(): """ Test that generic_response_convertor normalizes email addresses. @@ -1022,6 +1099,87 @@ async def test_upsert_sso_user_updates_role_for_existing_user(): assert call_args.kwargs["data"]["user_role"] == "proxy_admin" +@pytest.mark.asyncio +async def test_upsert_sso_user_fills_user_alias_for_existing_user(): + """ + An existing user row without an alias is updated with the SSO display name on login. + """ + from litellm.proxy._types import LiteLLM_UserTable + from litellm.proxy.management_endpoints.types import CustomOpenID + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + mock_prisma = MagicMock() + mock_prisma.db.litellm_usertable.update_many = AsyncMock() + + existing_user = LiteLLM_UserTable( + user_id="S-1-5-21-adfs-user", + user_email="jane.doe@example.com", + user_role="internal_user", + user_alias=None, + ) + sso_result = CustomOpenID( + id="S-1-5-21-adfs-user", + email="jane.doe@example.com", + first_name="Jane", + last_name="Doe", + display_name="Doe, Jane", + provider="generic", + team_ids=[], + ) + + await SSOAuthenticationHandler.upsert_sso_user( + result=sso_result, + user_info=existing_user, + user_email="jane.doe@example.com", + user_defined_values=None, + prisma_client=mock_prisma, + ) + + mock_prisma.db.litellm_usertable.update_many.assert_called_once_with( + where={"user_id": "S-1-5-21-adfs-user"}, + data={"user_email": "jane.doe@example.com", "user_alias": "Doe, Jane"}, + ) + + +@pytest.mark.asyncio +async def test_insert_sso_user_sets_user_alias_from_display_name(): + """ + A newly created SSO user is inserted with the IdP display name as user_alias. + """ + from litellm.proxy._types import NewUserResponse, SSOUserDefinedValues + from litellm.proxy.management_endpoints.types import CustomOpenID + from litellm.proxy.management_endpoints.ui_sso import insert_sso_user + + sso_result = CustomOpenID( + id="S-1-5-21-adfs-user", + email="jane.doe@example.com", + first_name="Jane", + last_name="Doe", + display_name="Doe, Jane", + provider="generic", + team_ids=[], + ) + user_defined_values: SSOUserDefinedValues = { + "models": [], + "user_id": "S-1-5-21-adfs-user", + "user_email": "jane.doe@example.com", + "max_budget": None, + "user_role": "internal_user", + "budget_duration": None, + } + + with patch( + "litellm.proxy.management_endpoints.ui_sso.new_user", + return_value=NewUserResponse(user_id="S-1-5-21-adfs-user", key="sk-xxxxx", teams=None), + ) as mock_new_user: + await insert_sso_user(result_openid=sso_result, user_defined_values=user_defined_values) + + new_user_request = mock_new_user.call_args.kwargs["data"] + assert new_user_request.user_id == "S-1-5-21-adfs-user" + assert new_user_request.user_email == "jane.doe@example.com" + assert new_user_request.user_alias == "Doe, Jane" + + @pytest.mark.asyncio async def test_upsert_sso_user_does_not_update_invalid_role(): """ diff --git a/tests/unit/proxy/proxy_server/test_lifecycle.py b/tests/unit/proxy/proxy_server/test_lifecycle.py index 4047473e9d7..ba5501315d9 100644 --- a/tests/unit/proxy/proxy_server/test_lifecycle.py +++ b/tests/unit/proxy/proxy_server/test_lifecycle.py @@ -17,19 +17,16 @@ Pins covered: from __future__ import annotations -import asyncio import inspect import json import logging import os import subprocess from collections.abc import Awaitable, Callable -from typing import List, Optional, Union from unittest.mock import AsyncMock, MagicMock, patch import pytest from apscheduler.schedulers.asyncio import AsyncIOScheduler -from fastapi import FastAPI from pydantic import BaseModel from typing_extensions import TypedDict @@ -682,16 +679,16 @@ class _SampleTD(TypedDict): def test_resolve_typed_dict_type_finds_class_in_optional(): - typ = Optional[_SampleTD] + typ = _SampleTD | None result = _resolve_typed_dict_type(typ) observed = { - "input_repr": "Optional[_SampleTD]", + "input_repr": "_SampleTD | None", "result_is_sample_td": result is _SampleTD, "result_is_class": isinstance(result, type), } assert normalize(observed) == { - "input_repr": "Optional[_SampleTD]", + "input_repr": "_SampleTD | None", "result_is_sample_td": True, "result_is_class": True, } @@ -717,7 +714,7 @@ class _SampleModelB(BaseModel): def test_resolve_pydantic_type_extracts_non_none_args_from_union(): - typ = Union[_SampleModelA, _SampleModelB, None] + typ = _SampleModelA | _SampleModelB | None result = _resolve_pydantic_type(typ) observed = { diff --git a/tests/unit/proxy/test__lazy_features.py b/tests/unit/proxy/test__lazy_features.py new file mode 100644 index 00000000000..d2bb8244c2f --- /dev/null +++ b/tests/unit/proxy/test__lazy_features.py @@ -0,0 +1,211 @@ +import sys +from collections.abc import AsyncGenerator, Mapping +from contextlib import asynccontextmanager +from types import ModuleType +from typing import Final + +import pytest +from fastapi import APIRouter, FastAPI +from fastapi.testclient import TestClient +from pydantic import BaseModel + +from litellm.proxy._lazy_features import ( + LazyFeature, + LazyFeatureMiddleware, + attach_lazy_features, + lazy_tag_to_prefix, + loaded_lazy_modules, +) + +FLAG: Final = "LITELLM_DISABLE_LAZY_ROUTES" +WARMUP_PATH: Final = "/lazy/warm/{name}" + + +class _Operation(BaseModel): + tags: tuple[str, ...] + + +class _WarmupBody(BaseModel): + stub_path: str + paths: Mapping[str, Mapping[str, _Operation]] + + +def _feature_module(monkeypatch: pytest.MonkeyPatch, name: str, path: str) -> LazyFeature: + async def served() -> dict[str, str]: + return {"feature": name} + + router: Final = APIRouter() + router.add_api_route(path, served, methods=["GET"]) + module: Final = ModuleType(f"tests.unit.proxy.lazy_fixture_{name}") + module.router = router # pyright: ignore[reportAttributeAccessIssue] # fixture module built at test time + monkeypatch.setitem(sys.modules, module.__name__, module) + return LazyFeature(name=name, module_path=module.__name__, path_prefixes=(path,)) + + +def _paths(app: FastAPI) -> tuple[str, ...]: + return tuple(str(getattr(route, "path", "")) for route in app.routes) + + +def _has_lazy_middleware(app: FastAPI) -> bool: + return any(middleware.cls is LazyFeatureMiddleware for middleware in app.user_middleware) + + +@pytest.mark.parametrize("value", ("1", "true", "TRUE", "yes", "on")) +def test_flag_registers_every_feature_at_startup(monkeypatch: pytest.MonkeyPatch, value: str) -> None: + monkeypatch.setenv(FLAG, value) + features: Final = ( + _feature_module(monkeypatch, "alpha", "/alpha/list"), + _feature_module(monkeypatch, "beta", "/beta/list"), + ) + app: Final = FastAPI() + + attach_lazy_features(app, features) + + assert WARMUP_PATH not in _paths(app) + assert not _has_lazy_middleware(app) + assert loaded_lazy_modules(app) == set() + with TestClient(app) as client: + at_startup: Final = _paths(app) + assert {"/alpha/list", "/beta/list"} <= set(at_startup) + assert loaded_lazy_modules(app) == {features[0].module_path, features[1].module_path} + assert client.get("/beta/list").json() == {"feature": "beta"} + assert client.post("/lazy/warm/alpha").status_code == 404 + assert _paths(app) == at_startup, "first feature request changed the table" + + +def test_flag_registers_before_the_inner_lifespan_and_after_late_routes(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv(FLAG, "true") + features: Final = (_feature_module(monkeypatch, "epsilon", "/epsilon/{name}"),) + seen_by_inner_lifespan: Final[list[tuple[str, ...]]] = [] # mutable-ok: captured from inside the lifespan + + @asynccontextmanager + async def inner_lifespan(app_: FastAPI) -> AsyncGenerator[None]: + seen_by_inner_lifespan.append(_paths(app_)) + yield + + async def late() -> dict[str, str]: + return {"feature": "late"} + + app: Final = FastAPI(lifespan=inner_lifespan) + attach_lazy_features(app, features) + app.add_api_route("/epsilon/list", late, methods=["GET"]) + + with TestClient(app) as client: + assert client.get("/epsilon/list").json() == {"feature": "late"}, "late eager route must win, as in lazy mode" + assert client.get("/epsilon/x").json() == {"feature": "epsilon"} + assert seen_by_inner_lifespan == [_paths(app)], "startup hooks inside the proxy lifespan must see the full table" + + +def test_flag_lets_a_route_added_during_startup_beat_an_overlapping_feature_route( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setenv(FLAG, "true") + features: Final = (_feature_module(monkeypatch, "zeta", "/zeta/{endpoint:path}"),) + + async def configured() -> dict[str, str]: + return {"feature": "configured"} + + @asynccontextmanager + async def adds_a_pass_through(app_: FastAPI) -> AsyncGenerator[None]: + app_.add_api_route("/zeta/{subpath:path}", configured, methods=["GET"]) + yield + + app: Final = FastAPI(lifespan=adds_a_pass_through) + attach_lazy_features(app, features) + + with TestClient(app) as client: + assert client.get("/zeta/health").json() == {"feature": "configured"}, ( + "lazy mode routes this to startup's route" + ) + + +def test_flag_does_not_bring_back_a_feature_route_removed_during_startup(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv(FLAG, "true") + features: Final = ( + _feature_module(monkeypatch, "eta", "/eta/list"), + _feature_module(monkeypatch, "theta", "/theta/list"), + ) + + @asynccontextmanager + async def drops_eta(app_: FastAPI) -> AsyncGenerator[None]: + app_.router.routes[:] = [route for route in app_.router.routes if getattr(route, "path", "") != "/eta/list"] + yield + + app: Final = FastAPI(lifespan=drops_eta) + attach_lazy_features(app, features) + + with TestClient(app) as client: + assert client.get("/eta/list").status_code == 404 + assert client.get("/theta/list").json() == {"feature": "theta"} + assert "/eta/list" not in _paths(app) + + +@pytest.mark.parametrize("value", (None, "", "0", "false", "off")) +def test_without_the_flag_features_still_mount_on_first_request( + monkeypatch: pytest.MonkeyPatch, value: str | None +) -> None: + if value is None: + monkeypatch.delenv(FLAG, raising=False) + else: + monkeypatch.setenv(FLAG, value) + features: Final = (_feature_module(monkeypatch, "gamma", "/gamma/list"),) + app: Final = FastAPI() + + attach_lazy_features(app, features) + + assert "/gamma/list" not in _paths(app) + assert WARMUP_PATH in _paths(app) + assert _has_lazy_middleware(app) + with TestClient(app) as client: + assert client.get("/gamma/list").json() == {"feature": "gamma"} + assert "/gamma/list" in _paths(app) + + +def test_flag_keeps_registering_after_one_feature_fails_to_import(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv(FLAG, "true") + broken: Final = LazyFeature( + name="broken", module_path="tests.unit.proxy.lazy_fixture_does_not_exist", path_prefixes=("/broken",) + ) + healthy: Final = _feature_module(monkeypatch, "delta", "/delta/list") + app: Final = FastAPI() + + attach_lazy_features(app, (broken, healthy)) + + with TestClient(app) as client: + assert "/delta/list" in _paths(app) + assert loaded_lazy_modules(app) == {broken.module_path, healthy.module_path} + assert client.get("/delta/list").json() == {"feature": "delta"} + assert client.get("/broken").status_code == 404 + + +def test_flag_hides_the_swagger_warmup_plugin(monkeypatch: pytest.MonkeyPatch) -> None: + import litellm.proxy._lazy_openapi_snapshot as snapshot + + monkeypatch.setattr(snapshot, "SNAPSHOT_FILE", snapshot.SNAPSHOT_FILE.with_name("missing-snapshot.json")) + monkeypatch.setenv(FLAG, "false") + assert lazy_tag_to_prefix() != {}, "control: without the flag and without a snapshot the plugin has tags" + monkeypatch.setenv(FLAG, "true") + assert lazy_tag_to_prefix() == {} + + +def test_without_the_flag_the_warmup_route_registers_a_feature_and_returns_its_paths( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.delenv(FLAG, raising=False) + features: Final = ( + _feature_module(monkeypatch, "alpha", "/alpha/list"), + _feature_module(monkeypatch, "beta", "/beta/list"), + ) + app: Final = FastAPI() + attach_lazy_features(app, features) + + with TestClient(app) as client: + assert client.post("/lazy/warm/zeta").status_code == 404 + warmed: Final = client.post("/lazy/warm/alpha") + assert warmed.status_code == 200, warmed.text + body: Final = _WarmupBody.model_validate_json(warmed.text) + assert body.stub_path == "/alpha/list" + assert set(body.paths) == {"/alpha/list"} + assert body.paths["/alpha/list"]["get"].tags == ("alpha",) + assert loaded_lazy_modules(app) == {features[0].module_path} + assert "/alpha/list" in _paths(app) and "/beta/list" not in _paths(app) diff --git a/tests/unit/proxy/test_prisma_migration.py b/tests/unit/proxy/test_prisma_migration.py index 3fc69b34213..b7de849b3b4 100644 --- a/tests/unit/proxy/test_prisma_migration.py +++ b/tests/unit/proxy/test_prisma_migration.py @@ -9,41 +9,23 @@ from litellm.proxy import prisma_migration class TestPrismaMigration: + @pytest.mark.parametrize("env", [{}, {"ENFORCE_PRISMA_MIGRATION_CHECK": "false"}], ids=("unset", "legacy-opt-out")) @patch("litellm.proxy.prisma_migration.subprocess.run") @patch("litellm.proxy.prisma_migration.run_server") - def test_main_enforces_migration_check_by_default( - self, mock_run_server: MagicMock, mock_subprocess_run: MagicMock + def test_main_runs_the_migration_job_with_no_opt_out( + self, mock_run_server: MagicMock, mock_subprocess_run: MagicMock, env: dict[str, str] ) -> None: mock_subprocess_run.return_value = MagicMock(returncode=0, stdout="", stderr="") - with patch.dict(os.environ, {}, clear=True): - assert prisma_migration.main() == 0 - - mock_run_server.assert_called_once_with( - ("--skip_server_startup", "--enforce_prisma_migration_check"), - standalone_mode=False, - ) - - @patch("litellm.proxy.prisma_migration.subprocess.run") - @patch("litellm.proxy.prisma_migration.run_server") - def test_main_disables_migration_check_when_explicitly_false( - self, mock_run_server: MagicMock, mock_subprocess_run: MagicMock - ) -> None: - mock_subprocess_run.return_value = MagicMock(returncode=0, stdout="", stderr="") - - with patch.dict(os.environ, {"ENFORCE_PRISMA_MIGRATION_CHECK": "false"}, clear=True): + with patch.dict(os.environ, env, clear=True): assert prisma_migration.main() == 0 mock_run_server.assert_called_once_with(("--skip_server_startup",), standalone_mode=False) - @pytest.mark.parametrize("env", [{}, {"ENFORCE_PRISMA_MIGRATION_CHECK": "false"}]) @patch("litellm.proxy.prisma_migration.subprocess.run") @patch("litellm.proxy.prisma_migration.run_server") def test_main_exits_zero_when_only_prisma_generate_fails( - self, - mock_run_server: MagicMock, - mock_subprocess_run: MagicMock, - env: dict[str, str], + self, mock_run_server: MagicMock, mock_subprocess_run: MagicMock ) -> None: mock_subprocess_run.return_value = MagicMock( returncode=1, @@ -51,7 +33,7 @@ class TestPrismaMigration: stderr="PermissionError: [Errno 13] Permission denied: '/app/.venv/lib/python3.13/site-packages/prisma/schema.prisma'", ) - with patch.dict(os.environ, env, clear=True): + with patch.dict(os.environ, {}, clear=True): assert prisma_migration.main() == 0 @patch("litellm.proxy.prisma_migration.subprocess.run") @@ -61,7 +43,7 @@ class TestPrismaMigration: ) -> None: mock_run_server.side_effect = SystemExit(1) - with patch.dict(os.environ, {}, clear=True): + with patch.dict(os.environ, {"ENFORCE_PRISMA_MIGRATION_CHECK": "false"}, clear=True): with pytest.raises(SystemExit, match="1"): prisma_migration.main() diff --git a/tests/unit/proxy/test_proxy_cli.py b/tests/unit/proxy/test_proxy_cli.py index 7007189aab9..56980b141b6 100644 --- a/tests/unit/proxy/test_proxy_cli.py +++ b/tests/unit/proxy/test_proxy_cli.py @@ -1,5 +1,6 @@ import inspect import os +from contextlib import nullcontext from pathlib import Path from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch @@ -2128,12 +2129,14 @@ class TestRunServerDbSetup: @patch("subprocess.run") @patch("atexit.register") @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database") + @patch("litellm.proxy.db.prisma_client.PrismaManager.build_request_log_indexes") @patch("litellm.proxy.db.check_migration.check_prisma_schema_diff") @patch("litellm.proxy.db.prisma_client.should_update_prisma_schema") def test_use_prisma_db_push_flag_behavior( self, mock_should_update_schema, mock_check_schema_diff, + mock_build_indexes, mock_setup_database, mock_atexit_register, mock_subprocess_run, @@ -2187,9 +2190,7 @@ class TestRunServerDbSetup: # Test 1: Without --use_prisma_db_push flag (default behavior) # use_prisma_db_push should be False (default), so use_migrate should be True run_server.main(["--local", "--skip_server_startup"], standalone_mode=False) - mock_setup_database.assert_called_with( - use_migrate=True, use_v2_resolver=True - ) + mock_setup_database.assert_called_with(use_migrate=True, use_v2_resolver=True) # Reset mocks mock_setup_database.reset_mock() @@ -2202,18 +2203,18 @@ class TestRunServerDbSetup: ["--local", "--skip_server_startup", "--use_prisma_db_push"], standalone_mode=False, ) - mock_setup_database.assert_called_with( - use_migrate=False, use_v2_resolver=True - ) + mock_setup_database.assert_called_with(use_migrate=False, use_v2_resolver=True) @patch("atexit.register") @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database") # test-quality-ok: run_server always wires the DB; same isolation as the sibling CLI tests above + @patch("litellm.proxy.db.prisma_client.PrismaManager.build_request_log_indexes") # test-quality-ok: run_server always wires the DB; same isolation as the sibling CLI tests above @patch("litellm.proxy.db.check_migration.check_prisma_schema_diff") # test-quality-ok: run_server always wires the DB; same isolation as the sibling CLI tests above @patch("litellm.proxy.db.prisma_client.should_update_prisma_schema") # test-quality-ok: run_server always wires the DB; same isolation as the sibling CLI tests above def test_migrations_run_when_the_prisma_cli_is_not_on_path( self, mock_should_update_schema, mock_check_schema_diff, + mock_build_indexes, mock_setup_database, mock_atexit_register, tmp_path, @@ -2262,24 +2263,24 @@ class TestRunServerDbSetup: run_server.main(["--local", "--skip_server_startup"], standalone_mode=False) assert "prisma CLI is neither on PATH" not in capsys.readouterr().out - mock_setup_database.assert_called_once_with( - use_migrate=True, use_v2_resolver=True - ) + mock_setup_database.assert_called_once_with(use_migrate=True, use_v2_resolver=True) @patch("subprocess.run") @patch("atexit.register") @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database") + @patch("litellm.proxy.db.prisma_client.PrismaManager.build_request_log_indexes") @patch("litellm.proxy.db.check_migration.check_prisma_schema_diff") @patch("litellm.proxy.db.prisma_client.should_update_prisma_schema") def test_startup_fails_when_db_setup_fails( self, mock_should_update_schema, mock_check_schema_diff, + mock_build_indexes, mock_setup_database, mock_atexit_register, mock_subprocess_run, ): - """Test that proxy exits with code 1 when PrismaManager.setup_database returns False and --enforce_prisma_migration_check is set""" + """Test that proxy exits with code 1 when PrismaManager.setup_database returns False, with no opt-in flag""" from litellm.proxy.proxy_cli import run_server mock_subprocess_run.return_value = MagicMock(returncode=0) @@ -2320,28 +2321,21 @@ class TestRunServerDbSetup: } with pytest.raises(SystemExit) as exc_info: - run_server.main( - [ - "--local", - "--skip_server_startup", - "--enforce_prisma_migration_check", - ], - standalone_mode=False, - ) + run_server.main(["--local", "--skip_server_startup"], standalone_mode=False) assert exc_info.value.code == 1 - mock_setup_database.assert_called_once_with( - use_migrate=True, use_v2_resolver=True - ) + mock_setup_database.assert_called_once_with(use_migrate=True, use_v2_resolver=True) @patch("subprocess.run") @patch("atexit.register") @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database") + @patch("litellm.proxy.db.prisma_client.PrismaManager.build_request_log_indexes") @patch("litellm.proxy.db.check_migration.check_prisma_schema_diff") @patch("litellm.proxy.db.prisma_client.should_update_prisma_schema") def test_startup_exits_on_non_postgres_database_url( self, mock_should_update_schema, mock_check_schema_diff, + mock_build_indexes, mock_setup_database, mock_atexit_register, mock_subprocess_run, @@ -2387,12 +2381,14 @@ class TestRunServerDbSetup: @patch("subprocess.run") @patch("atexit.register") @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database") + @patch("litellm.proxy.db.prisma_client.PrismaManager.build_request_log_indexes") @patch("litellm.proxy.db.check_migration.check_prisma_schema_diff") @patch("litellm.proxy.db.prisma_client.should_update_prisma_schema") def test_v2_migration_resolver_opts_in_via_env_var( self, mock_should_update_schema, mock_check_schema_diff, + mock_build_indexes, mock_setup_database, mock_atexit_register, mock_subprocess_run, @@ -2439,11 +2435,101 @@ class TestRunServerDbSetup: ["--local", "--skip_server_startup"], standalone_mode=False ) - mock_setup_database.assert_called_once_with( - use_migrate=True, use_v2_resolver=True - ) + mock_setup_database.assert_called_once_with(use_migrate=True, use_v2_resolver=True) assert "--use_v2_migration_resolver is deprecated" not in capsys.readouterr().out + @pytest.mark.parametrize( + ("arguments", "environment", "warned"), + ( + (("--local", "--skip_server_startup", "--enforce_prisma_migration_check"), {}, True), + (("--local", "--skip_server_startup"), {"ENFORCE_PRISMA_MIGRATION_CHECK": "true"}, False), + (("--local", "--skip_server_startup"), {"ENFORCE_PRISMA_MIGRATION_CHECK": "false"}, False), + ), + ids=("cli-flag", "env-true", "env-false"), + ) + @patch("subprocess.run") + @patch("atexit.register") + @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database", return_value=True) + @patch("litellm.proxy.db.prisma_client.PrismaManager.build_request_log_indexes", return_value=True) + @patch("litellm.proxy.db.check_migration.check_prisma_schema_diff") + @patch("litellm.proxy.db.prisma_client.should_update_prisma_schema", return_value=True) + def test_the_retired_enforce_prisma_migration_check_opt_in_still_parses_and_changes_nothing( + self, + mock_should_update_schema, + mock_check_schema_diff, + mock_build_indexes, + mock_setup_database, + mock_atexit_register, + mock_subprocess_run, + arguments, + environment, + warned, + capsys, + ): + """Deployments still pass the flag or set the env var; the flag is accepted with a + deprecation line and the env var is ignored, and a successful setup boots either way.""" + from litellm.proxy.proxy_cli import run_server + + mock_subprocess_run.return_value = MagicMock(returncode=0) + mock_proxy_module = MagicMock( + app=MagicMock(), + ProxyConfig=MagicMock(), + KeyManagementSettings=MagicMock(), + save_worker_config=MagicMock(), + ) + clean_env = {k: v for k, v in os.environ.items() if k not in ("DATABASE_URL", "DIRECT_URL")} + clean_env["DATABASE_URL"] = "postgresql://test:test@localhost:5432/test" + + with ( + patch.dict(os.environ, {**clean_env, **environment}, clear=True), + patch.dict( + "sys.modules", + {"proxy_server": mock_proxy_module, "litellm.proxy.proxy_server": mock_proxy_module}, + ), + ): + run_server.main(list(arguments), standalone_mode=False) + + mock_setup_database.assert_called_once_with(use_migrate=True, use_v2_resolver=True) + assert ("--enforce_prisma_migration_check is deprecated and has no effect" in capsys.readouterr().out) is warned + + @patch("subprocess.run") + @patch("atexit.register") + @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database", return_value=True) + def test_the_retired_enforce_prisma_migration_check_opt_in_warns_without_a_database( + self, + mock_setup_database, + mock_atexit_register, + mock_subprocess_run, + capsys, + ): + """The deprecation line does not depend on reaching database setup: a deployment that + passes the flag with no DATABASE_URL still learns the flag is dead.""" + from litellm.proxy.proxy_cli import run_server + + mock_subprocess_run.return_value = MagicMock(returncode=0) + mock_proxy_module = MagicMock( + app=MagicMock(), + ProxyConfig=MagicMock(), + KeyManagementSettings=MagicMock(), + save_worker_config=MagicMock(), + ) + clean_env = {k: v for k, v in os.environ.items() if k not in ("DATABASE_URL", "DIRECT_URL")} + + with ( + patch.dict(os.environ, clean_env, clear=True), + patch.dict( + "sys.modules", + {"proxy_server": mock_proxy_module, "litellm.proxy.proxy_server": mock_proxy_module}, + ), + ): + run_server.main( + ["--local", "--skip_server_startup", "--enforce_prisma_migration_check"], + standalone_mode=False, + ) + + mock_setup_database.assert_not_called() + assert "--enforce_prisma_migration_check is deprecated and has no effect" in capsys.readouterr().out + @pytest.mark.parametrize( "use_legacy_flag, env_value, expected", [ @@ -2479,12 +2565,14 @@ class TestRunServerDbSetup: @patch("subprocess.run") @patch("atexit.register") @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database") + @patch("litellm.proxy.db.prisma_client.PrismaManager.build_request_log_indexes") @patch("litellm.proxy.db.check_migration.check_prisma_schema_diff") @patch("litellm.proxy.db.prisma_client.should_update_prisma_schema") def test_legacy_resolver_flag_reaches_database_setup( self, mock_should_update_schema, mock_check_schema_diff, + mock_build_indexes, mock_setup_database, mock_atexit_register, mock_subprocess_run, @@ -2533,9 +2621,76 @@ class TestRunServerDbSetup: standalone_mode=False, ) - mock_setup_database.assert_called_once_with( - use_migrate=True, use_v2_resolver=False + mock_setup_database.assert_called_once_with(use_migrate=True, use_v2_resolver=False) + + @pytest.mark.parametrize( + ("arguments", "migrated", "exits", "waits_for_the_build"), + ( + (("--local", "--skip_server_startup"), True, True, True), + (("--local",), True, False, False), + (("--local",), False, True, False), + ), + ids=("migration-job", "serving-proxy", "serving-proxy-whose-migrations-failed"), + ) + @patch("uvicorn.run") + @patch("subprocess.run") + @patch("atexit.register") + @patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database", return_value=True) + @patch("litellm.proxy.db.prisma_client.PrismaManager.build_request_log_indexes", return_value=False) + @patch("litellm.proxy.db.prisma_client.PrismaManager.start_request_log_index_build") + @patch("litellm.proxy.db.check_migration.check_prisma_schema_diff") + @patch("litellm.proxy.db.prisma_client.should_update_prisma_schema", return_value=True) + def test_the_migration_job_waits_for_the_index_build_and_a_serving_proxy_starts_it_in_the_background( + self, + mock_should_update_schema, + mock_check_schema_diff, + mock_start_build, + mock_build_indexes, + mock_setup_database, + mock_atexit_register, + mock_subprocess_run, + mock_uvicorn_run, + arguments, + migrated, + exits, + waits_for_the_build, + ): + """`--skip_server_startup` is the migration job: it waits for the index build after the + migrations and exits 1 when one could not be built. A serving proxy that ran the + migrations starts the build in the background and serves whatever the build does; one + whose migrations failed exits 1 and starts no build.""" + from litellm.proxy.proxy_cli import run_server + + mock_setup_database.return_value = migrated + mock_subprocess_run.return_value = MagicMock(returncode=0) + mock_proxy_module = MagicMock( + app=MagicMock(), + ProxyConfig=MagicMock(), + KeyManagementSettings=MagicMock(), + save_worker_config=MagicMock(), ) + clean_env = {k: v for k, v in os.environ.items() if k not in ("DATABASE_URL", "DIRECT_URL")} + clean_env["DATABASE_URL"] = "postgresql://test:test@localhost:5432/test" + outcome = pytest.raises(SystemExit) if exits else nullcontext() + + with ( + patch.dict(os.environ, clean_env, clear=True), + patch.dict( + "sys.modules", + {"proxy_server": mock_proxy_module, "litellm.proxy.proxy_server": mock_proxy_module}, + ), + patch( + "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + ) as mock_get_args, + outcome as exc_info, + ): + mock_get_args.return_value = {"app": "litellm.proxy.proxy_server:app", "host": "localhost", "port": 8000} + run_server.main(list(arguments), standalone_mode=False) + + assert (exc_info is not None and exc_info.value.code == 1) is exits + mock_setup_database.assert_called_once_with(use_migrate=True, use_v2_resolver=True) + assert mock_build_indexes.call_count == int(migrated and waits_for_the_build) + assert mock_start_build.call_count == int(migrated and not waits_for_the_build) # --- Module-level helpers for worker startup hook tests --- diff --git a/tests/unit/proxy/test_tracing_endpoints.py b/tests/unit/proxy/test_tracing_endpoints.py index aa1403b8db9..2c587897938 100644 --- a/tests/unit/proxy/test_tracing_endpoints.py +++ b/tests/unit/proxy/test_tracing_endpoints.py @@ -23,6 +23,37 @@ from litellm.tracing.types import TraceScope TEAM_KEY = UserAPIKeyAuth( token="hashed-key", team_id="team-research", org_id="org-1", user_role=LitellmUserRoles.INTERNAL_USER ) +TRACE_RESPONSE: Final = { + "summary": { + "trace_id": "t1", + "name": "trace", + "service": "test", + "input_preview": "", + "start_time": "2026-01-01T00:00:00Z", + "duration_ms": 0, + "status": "ok", + "span_count": 0, + "agent_count": 0, + "agent_invocations": 0, + "llm_calls": 0, + "tool_calls": 0, + "error_count": 0, + "input_tokens": 0, + "output_tokens": 0, + "models": [], + "spend": None, + }, + "agents": [], + "spans": [], +} +SPAN_DETAIL_RESPONSE: Final = { + "span_id": "s1", + "input": "", + "output": "", + "input_ui": {"kind": "text", "text": ""}, + "output_ui": {"kind": "text", "text": ""}, + "attributes": {}, +} @pytest.mark.parametrize( @@ -110,7 +141,9 @@ def test_501_when_tracing_not_enabled( assert response.status_code == 501 assert response.headers["content-type"] == "application/x-protobuf" assert Status.FromString(response.content).message == ( - "Agent tracing is not enabled. Set `tracing:` in general_settings and CLICKHOUSE_URL." if native_available else "" + "Agent tracing is not enabled. Set `tracing:` in general_settings and CLICKHOUSE_URL." + if native_available + else "" ) assert client.get("/v1/traces").status_code == 501 @@ -171,17 +204,16 @@ def test_list_traces_defaults_to_last_24h(client, receiver): def test_get_trace_404_and_200(client, receiver): assert client.get("/v1/traces/missing").status_code == 404 - trace = {"summary": {"trace_id": "t1"}, "agents": [], "spans": []} - receiver.get_trace.return_value = trace + receiver.get_trace.return_value = TRACE_RESPONSE response = client.get("/v1/traces/t1") assert response.status_code == 200 - assert response.json() == trace + assert response.json() == TRACE_RESPONSE receiver.get_trace.assert_awaited_with("t1", {"team_ids": ("team-research",), "api_key_hash": ""}, "") def test_get_span_404_and_200(client, receiver): assert client.get("/v1/traces/t1/spans/s1").status_code == 404 - receiver.get_span.return_value = {"span_id": "s1", "input": "", "output": "", "attributes": {}} + receiver.get_span.return_value = SPAN_DETAIL_RESPONSE response = client.get("/v1/traces/t1/spans/s1") assert response.status_code == 200 assert response.json()["span_id"] == "s1" @@ -207,7 +239,7 @@ def test_get_span_serves_ui_content_from_stored_payloads(client): def test_trace_detail_passes_scoped_reference(client, receiver): - receiver.get_trace.return_value = {"summary": {"trace_id": "t1"}, "agents": [], "spans": []} + receiver.get_trace.return_value = TRACE_RESPONSE assert client.get("/v1/traces/t1?trace_ref=run-one").status_code == 200 receiver.get_trace.assert_awaited_with("t1", {"team_ids": ("team-research",), "api_key_hash": ""}, "run-one") @@ -495,3 +527,96 @@ def test_lens_reads_from_injected_storage_without_receiver() -> None: assert response.status_code == 200, response.text assert response.json()["executions"] == [] storage.lens_sample.assert_awaited_once() + + +@pytest.mark.parametrize( + ("auth", "expected_scope"), + ( + (UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), {"kind": "admin"}), + (UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY), {"kind": "admin"}), + (TEAM_KEY, {"kind": "team", "team_id": "team-research"}), + ( + UserAPIKeyAuth(token="project-key", team_id="team-a", project_id="project-a"), + {"kind": "key", "team_id": "team-a", "api_key_hash": "project-key"}, + ), + (UserAPIKeyAuth(token="solo-key"), {"kind": "key", "team_id": "", "api_key_hash": "solo-key"}), + ), +) +def test_sql_and_help_use_authenticated_scope( + client: TestClient, receiver: MagicMock, auth: UserAPIKeyAuth, expected_scope: dict[str, str] +) -> None: + client.app.dependency_overrides[user_api_key_auth] = lambda: auth + client.app.dependency_overrides[tracing_endpoints.provide_trace_query_secret] = lambda: "test-secret" + receiver.store.storage.query_sql = AsyncMock(return_value='{"data":[{"value":1}]}') + receiver.store.storage.query_help = AsyncMock(return_value='{"guide":"scoped"}') + result: Final = client.post("/v1/traces/query", json={"sql": "SELECT * FROM otel_traces"}) + assert result.status_code == 200, result.text + assert result.json() == {"data": [{"value": 1}]} + receiver.store.storage.query_sql.assert_awaited_once_with( + "SELECT * FROM otel_traces", expected_scope, "test-secret" + ) + help_result: Final = client.get("/v1/traces/query/help") + assert help_result.status_code == 200, help_result.text + assert help_result.json() == {"guide": "scoped"} + receiver.store.storage.query_help.assert_awaited_once_with(expected_scope, "test-secret") + forged: Final = client.post("/v1/traces/query", json={"sql": "SELECT 1", "scope": {"kind": "admin"}}) + assert forged.status_code == 422, forged.text + assert receiver.store.storage.query_sql.await_count == 1 + + +@pytest.mark.parametrize("auth", (UserAPIKeyAuth(), UserAPIKeyAuth(team_id="a", project_id="p"))) +def test_sql_rejects_missing_identity_without_querying( + client: TestClient, receiver: MagicMock, auth: UserAPIKeyAuth +) -> None: + client.app.dependency_overrides[user_api_key_auth] = lambda: auth + client.app.dependency_overrides[tracing_endpoints.provide_trace_query_secret] = lambda: "test-secret" + result: Final = client.post("/v1/traces/query", json={"sql": "SELECT * FROM otel_traces"}) + assert result.status_code == 403, result.text + assert client.get("/v1/traces/query/help").status_code == 403 + receiver.store.storage.query_sql.assert_not_called() + receiver.store.storage.query_help.assert_not_called() + + +@pytest.mark.parametrize( + ("error", "status"), ((ValueError("invalid SQL"), 400), (RuntimeError("reader unavailable"), 503)) +) +def test_sql_reports_rejected_queries_and_unavailable_readers( + client: TestClient, receiver: MagicMock, error: Exception, status: int +) -> None: + client.app.dependency_overrides[tracing_endpoints.provide_trace_query_secret] = lambda: "test-secret" + receiver.store.storage.query_sql = AsyncMock(side_effect=error) + result: Final = client.post("/v1/traces/query", json={"sql": "SELECT 1"}) + assert result.status_code == status, result.text + receiver.store.storage.query_sql.assert_awaited_once_with( + "SELECT 1", {"kind": "team", "team_id": "team-research"}, "test-secret" + ) + + +def test_query_help_does_not_fall_back_when_reader_provisioning_fails(client: TestClient, receiver: MagicMock) -> None: + client.app.dependency_overrides[tracing_endpoints.provide_trace_query_secret] = lambda: "test-secret" + receiver.store.storage.query_help = AsyncMock(side_effect=RuntimeError("reader provisioning failed")) + result: Final = client.get("/v1/traces/query/help") + assert result.status_code == 503, result.text + receiver.store.storage.query_help.assert_awaited_once_with( + {"kind": "team", "team_id": "team-research"}, "test-secret" + ) + + +@pytest.mark.parametrize("secret", (None, "configured-master-key")) +def test_queries_require_a_proxy_secret( + client: TestClient, receiver: MagicMock, monkeypatch: pytest.MonkeyPatch, secret: str | None +) -> None: + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "master_key", secret) + receiver.store.storage.query_sql = AsyncMock(return_value='{"data":[]}') + result: Final = client.post("/v1/traces/query", json={"sql": "SELECT 1"}) + if secret is None: + assert result.status_code == 503, result.text + assert "master key" in result.json()["detail"] + receiver.store.storage.query_sql.assert_not_awaited() + return + assert result.status_code == 200, result.text + receiver.store.storage.query_sql.assert_awaited_once_with( + "SELECT 1", {"kind": "team", "team_id": "team-research"}, secret + ) diff --git a/tests/unit/repositories/test_daily_activity_repository.py b/tests/unit/repositories/test_daily_activity_repository.py new file mode 100644 index 00000000000..4bb833f2bc2 --- /dev/null +++ b/tests/unit/repositories/test_daily_activity_repository.py @@ -0,0 +1,566 @@ +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from datetime import datetime, timezone +from typing import Final + +import pytest +from pydantic import ValidationError + +from litellm import constants +from litellm.repositories.daily_activity_repository import DailyActivityRepository +from litellm.repositories.daily_activity_sql import ( + ExportCursor, + build_cache_leakage_keys_sql, + build_entity_rollup_sql, + build_export_sql, + build_key_page_sql, + build_key_search_sql, + build_model_top_keys_sql, +) +from litellm.types.repositories.daily_activity import ( + DailyActivityScope, + DailyActivityTable, + ExportType, + KeyMetadataRow, + KeyPage, + KeySpendRow, + SpendLogsWindow, +) + + +@dataclass(frozen=True, slots=True) +class _FakeVerificationToken: + token: str + key_alias: str | None + team_id: str | None + user_id: str | None + metadata: object | None + + +@dataclass(frozen=True, slots=True) +class _FakeDeletedVerificationToken(_FakeVerificationToken): + deleted_at: datetime + + +def _scope( + *, + table: DailyActivityTable = DailyActivityTable.USER, + entity_ids: tuple[str, ...] | None = ("user-1",), + api_keys: tuple[str, ...] | None = None, + exclude_entity_ids: tuple[str, ...] = (), + model: str | None = None, +) -> DailyActivityScope: + entity_field: Final = { + DailyActivityTable.USER: "user_id", + DailyActivityTable.TEAM: "team_id", + DailyActivityTable.TAG: "tag", + DailyActivityTable.ORGANIZATION: "organization_id", + DailyActivityTable.CUSTOMER: "end_user_id", + DailyActivityTable.AGENT: "agent_id", + }[table] + return DailyActivityScope( + table=table, + entity_id_field=entity_field, + entity_ids=entity_ids, + exclude_entity_ids=exclude_entity_ids, + api_keys=api_keys, + start_date="2026-01-01", + end_date="2026-01-31", + model=model, + timezone_offset_minutes=None, + ) + + +def _key_spend_row(api_key: str) -> dict[str, object]: + return { + "api_key": api_key, + "spend": 1.0, + "prompt_tokens": 10, + "completion_tokens": 2, + "total_tokens": 12, + "api_requests": 1, + "successful_requests": 1, + "failed_requests": 0, + "cache_read_input_tokens": 3, + "cache_creation_input_tokens": 1, + } + + +def _export_row(api_key: str | None) -> dict[str, object]: + return { + "date": "2026-01-01", + "entity_id": "user-1", + "entity_alias": None, + "api_key": api_key, + "key_alias": None, + "user_id": None, + "user_email": None, + "model": None, + "spend": 1.0, + "flat_cost": 0.0, + "prompt_tokens": 10, + "completion_tokens": 2, + "api_requests": 1, + "successful_requests": 1, + "failed_requests": 0, + "cache_read_input_tokens": 3, + "cache_creation_input_tokens": 1, + } + + +class _FakeTable: + def __init__(self, rows: Sequence[object] = ()) -> None: + self.rows: Final = tuple(rows) + self.find_many_calls: list[Mapping[str, object]] = [] + self.count_calls: list[Mapping[str, object]] = [] + self.pagination_calls: list[tuple[int | None, int | None, tuple[Mapping[str, str], ...] | None]] = [] + + async def find_many( + self, + *, + where: Mapping[str, object], + skip: int | None = None, + take: int | None = None, + order: tuple[Mapping[str, str], ...] | None = None, + ) -> tuple[object, ...]: + self.find_many_calls.append(where) + self.pagination_calls.append((skip, take, order)) + if "token" not in where: + return self.rows + token_filter: Final = where["token"] + if not isinstance(token_filter, Mapping): + return () + token_values: Final = token_filter.get("in") + if not isinstance(token_values, list): + return () + return tuple(row for row in self.rows if isinstance(row, _FakeVerificationToken) and row.token in token_values) + + async def count(self, *, where: Mapping[str, object]) -> int: + self.count_calls.append(where) + return len(self.rows) + + +class _FailingTable(_FakeTable): + def __init__(self, failure: str) -> None: + super().__init__() + self.failure: Final = failure + + async def find_many( + self, + *, + where: Mapping[str, object], + skip: int | None = None, + take: int | None = None, + order: tuple[Mapping[str, str], ...] | None = None, + ) -> tuple[object, ...]: + raise RuntimeError(f"{self.failure}: {where!r} {skip!r} {take!r} {order!r}") + + +class _FakeDatabase: + def __init__(self, responses: Sequence[Sequence[Mapping[str, object]] | None] = ()) -> None: + self.responses = tuple(responses) + self.query_calls: list[tuple[str, tuple[object, ...]]] = [] + self.litellm_verificationtoken = _FakeTable() + self.litellm_deletedverificationtoken = _FakeTable() + self.litellm_dailyuserspend = _FakeTable() + self.litellm_dailyteamspend = _FakeTable() + self.litellm_dailytagspend = _FakeTable() + self.litellm_dailyorganizationspend = _FakeTable() + self.litellm_dailyenduserspend = _FakeTable() + self.litellm_dailyagentspend = _FakeTable() + + async def query_raw(self, query: str, *params: object) -> Sequence[Mapping[str, object]] | None: + self.query_calls.append((query, params)) + response_index: Final = len(self.query_calls) - 1 + if response_index >= len(self.responses): + return () + return self.responses[response_index] + + +class _FakePrismaClient: + def __init__(self, database: _FakeDatabase) -> None: + self.db: Final = database + + +class _ProxyReads: + def __init__(self) -> None: + self.recovery_calls: list[tuple[Mapping[str, KeyMetadataRow], frozenset[str], SpendLogsWindow | None]] = [] + + async def recover_key_metadata( + self, + resolved: Mapping[str, KeyMetadataRow], + api_keys: frozenset[str], + window: SpendLogsWindow | None, + ) -> Mapping[str, KeyMetadataRow]: + self.recovery_calls.append((resolved, api_keys, window)) + return resolved + + +def _repository( + database: _FakeDatabase, proxy_reads: _ProxyReads | None = None +) -> tuple[DailyActivityRepository, _ProxyReads]: + reads: Final = proxy_reads if proxy_reads is not None else _ProxyReads() + return DailyActivityRepository(_FakePrismaClient(database), proxy_reads=reads), reads + + +@pytest.mark.asyncio +async def test_key_methods_send_builder_queries_with_caller_limits() -> None: + database = _FakeDatabase(((_key_spend_row("key-a"),), (_key_spend_row("key-b"),), (_key_spend_row("key-c"),))) + repository, _ = _repository(database) + scope = _scope() + + assert await repository.search_keys(scope, search="key", limit=2) == ("key-a",) + model_keys: Final = await repository.model_top_keys(scope, model_group="model-a", by_model_group=True, limit=2) + leakage_keys: Final = await repository.cache_leakage_keys(scope, limit=2) + + assert tuple(row.api_key for row in model_keys) == ("key-b",) + assert tuple(row.api_key for row in leakage_keys) == ("key-c",) + assert model_keys[0].spend == 1.0 + assert leakage_keys[0].prompt_tokens - leakage_keys[0].cache_read_input_tokens == 7 + assert database.query_calls == [ + ( + build_key_search_sql(scope, search="key", limit=2).sql, + build_key_search_sql(scope, search="key", limit=2).params, + ), + ( + build_model_top_keys_sql(scope, model_group="model-a", by_model_group=True, limit=2).sql, + build_model_top_keys_sql(scope, model_group="model-a", by_model_group=True, limit=2).params, + ), + ( + build_cache_leakage_keys_sql(scope, limit=2).sql, + build_cache_leakage_keys_sql(scope, limit=2).params, + ), + ] + + +@pytest.mark.asyncio +async def test_key_page_maps_rows_and_keeps_total_for_an_empty_page() -> None: + database = _FakeDatabase( + ( + ({"total_api_keys": 2, **_key_spend_row("key-a")},), + ({"total_api_keys": 2, "api_key": None},), + ) + ) + repository, _ = _repository(database) + scope = _scope() + + first_page: Final = await repository.key_page(scope, offset=0, limit=1) + empty_page: Final = await repository.key_page(scope, offset=2, limit=1) + + assert first_page == KeyPage( + rows=( + KeySpendRow( + api_key="key-a", + spend=1.0, + prompt_tokens=10, + completion_tokens=2, + total_tokens=12, + api_requests=1, + successful_requests=1, + failed_requests=0, + cache_read_input_tokens=3, + cache_creation_input_tokens=1, + ), + ), + total_api_keys=2, + ) + assert empty_page == KeyPage(rows=(), total_api_keys=2) + assert database.query_calls == [ + ( + build_key_page_sql(scope, offset=0, limit=1).sql, + build_key_page_sql(scope, offset=0, limit=1).params, + ), + ( + build_key_page_sql(scope, offset=2, limit=1).sql, + build_key_page_sql(scope, offset=2, limit=1).params, + ), + ] + + +@pytest.mark.asyncio +async def test_key_methods_reject_limits_outside_bounds() -> None: + database = _FakeDatabase() + repository, _ = _repository(database) + + with pytest.raises(ValueError, match="limit"): + await repository.search_keys(_scope(), search="key", limit=0) + with pytest.raises(ValueError, match="limit"): + await repository.model_top_keys(_scope(), model_group="model-a", by_model_group=False, limit=0) + with pytest.raises(ValueError, match="limit"): + await repository.cache_leakage_keys(_scope(), limit=0) + with pytest.raises(ValueError, match="limit"): + await repository.search_keys(_scope(), search="key", limit=constants.USAGE_KEY_SEARCH_MAX + 1) + with pytest.raises(ValueError, match="limit"): + await repository.model_top_keys( + _scope(), model_group="model-a", by_model_group=False, limit=constants.USAGE_MODEL_TOP_KEYS_MAX + 1 + ) + with pytest.raises(ValueError, match="limit"): + await repository.cache_leakage_keys(_scope(), limit=constants.USAGE_CACHE_LEAKAGE_KEYS_MAX + 1) + assert database.query_calls == [] + + +@pytest.mark.asyncio +async def test_key_spend_validation_rejects_malformed_rows() -> None: + repository, _ = _repository(_FakeDatabase((({"api_key": "missing-metrics"},),))) + + with pytest.raises(ValidationError): + await repository.search_keys(_scope(), search="key", limit=1) + + +@pytest.mark.asyncio +async def test_key_metadata_prefers_active_rows_and_recovers_all_requested_keys() -> None: + database = _FakeDatabase() + active: Final = _FakeVerificationToken( + token="active", + key_alias="current", + team_id="team-active", + user_id="user-active", + metadata={"tags": ["production", "internal"]}, + ) + deleted_active_duplicate: Final = _FakeDeletedVerificationToken( + token="active", + key_alias="stale", + team_id="team-stale", + user_id="user-stale", + metadata={"tags": []}, + deleted_at=datetime(2026, 1, 3, tzinfo=timezone.utc), + ) + deleted_older: Final = _FakeDeletedVerificationToken( + token="deleted", + key_alias="older", + team_id=None, + user_id=None, + metadata={"tags": "invalid"}, + deleted_at=datetime(2026, 1, 2, tzinfo=timezone.utc), + ) + deleted_newer: Final = _FakeDeletedVerificationToken( + token="deleted", + key_alias="newer", + team_id=None, + user_id=None, + metadata={"tags": ["archived"]}, + deleted_at=datetime(2026, 1, 4, tzinfo=timezone.utc), + ) + malformed_non_list: Final = _FakeVerificationToken( + token="malformed-non-list", + key_alias=None, + team_id=None, + user_id=None, + metadata={"tags": "invalid"}, + ) + malformed_list: Final = _FakeVerificationToken( + token="malformed-list", + key_alias=None, + team_id=None, + user_id=None, + metadata={"tags": [1]}, + ) + database.litellm_verificationtoken = _FakeTable((active, malformed_non_list, malformed_list)) + database.litellm_deletedverificationtoken = _FakeTable((deleted_active_duplicate, deleted_older, deleted_newer)) + proxy_reads: Final = _ProxyReads() + repository, _ = _repository(database, proxy_reads) + window: Final = (datetime(2026, 1, 1), datetime(2026, 2, 1)) + requested: Final = frozenset(("active", "deleted", "malformed-non-list", "malformed-list", "unresolved")) + + result = await repository.key_metadata(requested, window) + + assert result["active"] == KeyMetadataRow( + api_key="active", + key_alias="current", + team_id="team-active", + user_id="user-active", + user_email=None, + key_exists=True, + tags=("production", "internal"), + ) + assert result["deleted"].key_alias == "newer" + assert result["deleted"].key_exists is False + assert result["deleted"].tags == ("archived",) + assert result["malformed-non-list"].tags == () + assert result["malformed-list"].tags == () + assert len(database.litellm_deletedverificationtoken.find_many_calls) == 1 + assert set(database.litellm_deletedverificationtoken.find_many_calls[0]["token"]["in"]) == { + "deleted", + "unresolved", + } + assert proxy_reads.recovery_calls == [ + ( + result, + requested, + window, + ) + ] + + +@pytest.mark.asyncio +async def test_key_metadata_continues_with_active_rows_when_deleted_lookup_fails() -> None: + database = _FakeDatabase() + active: Final = _FakeVerificationToken( + token="active", + key_alias="current", + team_id=None, + user_id=None, + metadata={"tags": []}, + ) + database.litellm_verificationtoken = _FakeTable((active,)) + database.litellm_deletedverificationtoken = _FailingTable("deleted token query failed") + repository, proxy_reads = _repository(database) + + result = await repository.key_metadata(frozenset(("active", "deleted")), None) + + assert result["active"].key_alias == "current" + assert tuple(proxy_reads.recovery_calls[0][0]) == ("active",) + assert proxy_reads.recovery_calls[0][1] == frozenset(("active", "deleted")) + + +@pytest.mark.asyncio +async def test_key_metadata_empty_set_does_not_query_tables() -> None: + database = _FakeDatabase() + repository, proxy_reads = _repository(database) + + assert await repository.key_metadata(frozenset(), None) == {} + assert database.litellm_verificationtoken.find_many_calls == [] + assert proxy_reads.recovery_calls == [] + + +@pytest.mark.asyncio +async def test_key_metadata_propagates_active_token_lookup_failures() -> None: + database = _FakeDatabase() + database.litellm_verificationtoken = _FailingTable("active token query failed") + repository, _ = _repository(database) + + with pytest.raises(RuntimeError, match="active token query failed"): + await repository.key_metadata(frozenset(("active",)), None) + + assert database.litellm_deletedverificationtoken.find_many_calls == [] + + +@pytest.mark.asyncio +async def test_aggregated_normalizes_a_null_raw_query_result() -> None: + database = _FakeDatabase((None,)) + repository, _ = _repository(database) + + result = await repository.aggregated( + _scope(), include_entity_breakdown=False, api_key_limit=constants.USAGE_TOP_API_KEYS_DEFAULT + ) + + assert result.grouping_rows == () + assert result.entity_rows is None + assert result.distinct_api_keys == 0 + assert len(database.query_calls) == 1 + + +@pytest.mark.asyncio +async def test_aggregated_passes_api_key_limit_to_entity_rollup_query() -> None: + database = _FakeDatabase(((), ())) + repository, _ = _repository(database) + scope = _scope(table=DailyActivityTable.TEAM) + + result = await repository.aggregated(scope, include_entity_breakdown=True, api_key_limit=3) + + assert result.entity_rows == () + assert database.query_calls[1] == ( + build_entity_rollup_sql(scope, api_key_limit=3).sql, + build_entity_rollup_sql(scope, api_key_limit=3).params, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("table", "entity_field"), + [ + (DailyActivityTable.USER, "user_id"), + (DailyActivityTable.TEAM, "team_id"), + (DailyActivityTable.TAG, "tag"), + (DailyActivityTable.ORGANIZATION, "organization_id"), + (DailyActivityTable.CUSTOMER, "end_user_id"), + (DailyActivityTable.AGENT, "agent_id"), + ], +) +async def test_daily_rows_selects_the_table_and_applies_filters_and_pagination( + table: DailyActivityTable, entity_field: str +) -> None: + database = _FakeDatabase() + repository, _ = _repository(database) + scope = _scope( + table=table, + entity_ids=("entity-1",), + exclude_entity_ids=("excluded-1",), + api_keys=("key-1",), + model="model-1", + ) + + result = await repository.daily_rows(scope, page=3, page_size=2) + + expected_where: Final = { + "date": {"gte": "2026-01-01", "lte": "2026-01-31"}, + entity_field: {"in": ["entity-1"]}, + "OR": [{entity_field: None}, {entity_field: {"not": {"in": ["excluded-1"]}}}], + "model": "model-1", + "api_key": {"in": ["key-1"]}, + } + tables: Final = { + DailyActivityTable.USER: database.litellm_dailyuserspend, + DailyActivityTable.TEAM: database.litellm_dailyteamspend, + DailyActivityTable.TAG: database.litellm_dailytagspend, + DailyActivityTable.ORGANIZATION: database.litellm_dailyorganizationspend, + DailyActivityTable.CUSTOMER: database.litellm_dailyenduserspend, + DailyActivityTable.AGENT: database.litellm_dailyagentspend, + } + selected_table: Final = tables[table] + + assert result.total_count == 0 + assert result.rows == () + assert selected_table.count_calls == [expected_where] + assert selected_table.find_many_calls == [expected_where] + assert selected_table.pagination_calls == [(4, 2, ({"date": "desc"}, {"id": "asc"}))] + assert sum(len(daily_table.find_many_calls) for daily_table in tables.values()) == 1 + + +@pytest.mark.asyncio +async def test_daily_rows_exclusion_without_entity_filter_keeps_null_entity_rows() -> None: + database = _FakeDatabase() + repository, _ = _repository(database) + scope = _scope(table=DailyActivityTable.TEAM, entity_ids=None, exclude_entity_ids=("litellm-dashboard",)) + + await repository.daily_rows(scope, page=1, page_size=10) + + expected_where: Final = { + "date": {"gte": "2026-01-01", "lte": "2026-01-31"}, + "OR": [{"team_id": None}, {"team_id": {"not": {"in": ["litellm-dashboard"]}}}], + } + assert database.litellm_dailyteamspend.count_calls == [expected_where] + assert database.litellm_dailyteamspend.find_many_calls == [expected_where] + + +@pytest.mark.asyncio +async def test_export_is_lazy_and_uses_the_last_row_as_the_next_cursor(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(constants, "USAGE_EXPORT_BATCH_SIZE", 2) + database = _FakeDatabase( + ( + (_export_row("key-1"), _export_row("key-2")), + (_export_row("key-3"), _export_row("key-4")), + (_export_row("key-5"),), + ) + ) + repository, _ = _repository(database) + rows = repository.export_rows(_scope(), export_type=ExportType.DAILY_WITH_KEYS) + + assert database.query_calls == [] + assert (await rows.__anext__()).api_key == "key-1" + assert len(database.query_calls) == 1 + results = [row async for row in rows] + + assert [row.api_key for row in results] == ["key-2", "key-3", "key-4", "key-5"] + assert len(database.query_calls) == 3 + assert database.query_calls[1][1][-4:] == ("2026-01-01", "user-1", "key-2", 2) + assert database.query_calls[2][1][-4:] == ("2026-01-01", "user-1", "key-4", 2) + assert ( + build_export_sql( + _scope(), + export_type=ExportType.DAILY_WITH_KEYS, + after=ExportCursor("2026-01-01", "user-1", "key-2"), + batch_size=2, + ).params + == database.query_calls[1][1] + ) diff --git a/tests/unit/repositories/test_daily_activity_sql.py b/tests/unit/repositories/test_daily_activity_sql.py new file mode 100644 index 00000000000..c775dfd1f88 --- /dev/null +++ b/tests/unit/repositories/test_daily_activity_sql.py @@ -0,0 +1,420 @@ +from datetime import datetime, timezone +from typing import Final + +import pytest + +from litellm import constants +from litellm.constants import PTU_SENTINEL_API_KEY +from litellm.repositories.daily_activity_sql import ( + ExportCursor, + adjust_dates_for_timezone, + build_aggregated_sql, + build_cache_leakage_keys_sql, + build_entity_rollup_sql, + build_export_sql, + build_key_page_sql, + build_key_search_sql, + build_model_top_keys_sql, + build_where_clause, +) +from litellm.types.proxy.management_endpoints.common_daily_activity import SpendMetrics +from litellm.types.repositories.daily_activity import DailyActivityScope, DailyActivityTable, ExportType + + +def _scope( + *, + table: DailyActivityTable = DailyActivityTable.USER, + entity_ids: tuple[str, ...] | None = ("user-1",), + exclude_entity_ids: tuple[str, ...] = (), + api_keys: tuple[str, ...] | None = None, + model: str | None = None, + timezone_offset_minutes: int | None = None, + include_current_utc_day: bool = False, + start_date: str = "2026-01-01", + end_date: str = "2026-01-31", +) -> DailyActivityScope: + entity_field = { + DailyActivityTable.USER: "user_id", + DailyActivityTable.TEAM: "team_id", + DailyActivityTable.TAG: "tag", + DailyActivityTable.ORGANIZATION: "organization_id", + DailyActivityTable.CUSTOMER: "end_user_id", + DailyActivityTable.AGENT: "agent_id", + }[table] + return DailyActivityScope( + table=table, + entity_id_field=entity_field, + entity_ids=entity_ids, + exclude_entity_ids=exclude_entity_ids, + api_keys=api_keys, + start_date=start_date, + end_date=end_date, + model=model, + timezone_offset_minutes=timezone_offset_minutes, + include_current_utc_day=include_current_utc_day, + ) + + +def test_where_clause_binds_each_filter_as_a_single_array_parameter() -> None: + scope = _scope( + entity_ids=("user-1", "user-2"), + exclude_entity_ids=("user-3",), + api_keys=("key-1", "key-2"), + model="gpt-test", + ) + + sql, params = build_where_clause(scope) + + assert sql == ( + 'date >= $1 AND date <= $2 AND "user_id" = ANY($3::text[]) ' + 'AND ("user_id" IS NULL OR NOT ("user_id" = ANY($4::text[]))) AND model = $5 AND api_key = ANY($6::text[])' + ) + assert params == ( + "2026-01-01", + "2026-01-31", + ["user-1", "user-2"], + ["user-3"], + "gpt-test", + ["key-1", "key-2"], + ) + + +def test_where_clause_exclusion_keeps_null_entity_rows() -> None: + scope = _scope(table=DailyActivityTable.TEAM, entity_ids=None, exclude_entity_ids=("litellm-dashboard",)) + + sql, params = build_where_clause(scope) + + assert sql == 'date >= $1 AND date <= $2 AND ("team_id" IS NULL OR NOT ("team_id" = ANY($3::text[])))' + assert params == ("2026-01-01", "2026-01-31", ["litellm-dashboard"]) + + +@pytest.mark.parametrize( + ("entity_ids", "api_keys", "expected_sql", "expected_params"), + [ + (None, None, "date >= $1 AND date <= $2", ("2026-01-01", "2026-01-31")), + ((), None, "date >= $1 AND date <= $2 AND FALSE", ("2026-01-01", "2026-01-31")), + (None, (), "date >= $1 AND date <= $2 AND FALSE", ("2026-01-01", "2026-01-31")), + ], +) +def test_where_clause_distinguishes_no_filter_from_empty_membership( + entity_ids: tuple[str, ...] | None, + api_keys: tuple[str, ...] | None, + expected_sql: str, + expected_params: tuple[object, ...], +) -> None: + scope = _scope(entity_ids=entity_ids, api_keys=api_keys) + + sql, params = build_where_clause(scope) + + assert sql == expected_sql + assert params == expected_params + + +def test_key_page_sql_orders_exact_spend_and_binds_scope_before_page() -> None: + query = build_key_page_sql(_scope(), offset=7, limit=3) + + assert query.params == ( + "2026-01-01", + "2026-01-31", + ["user-1"], + PTU_SENTINEL_API_KEY, + 3, + 7, + ) + assert "SUM(spend::numeric) AS rank_spend" in query.sql + assert "ORDER BY rank_spend DESC, api_key" in query.sql + assert "(SELECT COUNT(*) FROM ranked)::bigint AS total_api_keys" in query.sql + + +@pytest.mark.parametrize( + ("offset", "limit", "error"), + ( + (0, 0, "limit must be between"), + (0, constants.USAGE_KEY_PAGE_MAX + 1, "limit must be between"), + (-1, 1, "offset must be non-negative"), + ), +) +def test_key_page_sql_rejects_invalid_page_bounds(offset: int, limit: int, error: str) -> None: + with pytest.raises(ValueError, match=error): + build_key_page_sql(_scope(), offset=offset, limit=limit) + + +def test_scope_rejects_an_entity_field_not_allowed_for_its_table() -> None: + with pytest.raises(ValueError, match="Invalid entity_id_field"): + DailyActivityScope( + table=DailyActivityTable.USER, + entity_id_field="team_id", + entity_ids=None, + exclude_entity_ids=(), + api_keys=None, + start_date="2026-01-01", + end_date="2026-01-31", + model=None, + timezone_offset_minutes=None, + ) + + +def test_timezone_adjustment_only_extends_an_opted_in_live_range() -> None: + now = datetime(2026, 8, 6, 4, 30, tzinfo=timezone.utc) + + assert adjust_dates_for_timezone("2026-07-06", "2026-08-05", 420, include_current_utc_day=True, utc_now=now) == ( + "2026-07-06", + "2026-08-06", + ) + assert adjust_dates_for_timezone("2026-07-01", "2026-08-04", 420, include_current_utc_day=True, utc_now=now) == ( + "2026-07-01", + "2026-08-04", + ) + + +@pytest.mark.parametrize("offset_minutes", [None, 0, -330, -540, -60, 240, 300, 480]) +def test_timezone_adjustment_preserves_daily_bucket_dates(offset_minutes: int | None) -> None: + assert adjust_dates_for_timezone("2026-05-29", "2026-05-29", offset_minutes) == ( + "2026-05-29", + "2026-05-29", + ) + + +@pytest.mark.parametrize("offset_minutes", [-330, 480]) +def test_timezone_adjustment_preserves_single_day_additivity(offset_minutes: int) -> None: + days: Final = ("2026-05-29", "2026-05-30", "2026-05-31", "2026-06-01", "2026-06-02") + single_day_ranges: Final = tuple(adjust_dates_for_timezone(day, day, offset_minutes) for day in days) + multi_day_range: Final = adjust_dates_for_timezone(days[0], days[-1], offset_minutes) + + assert tuple(start for start, _ in single_day_ranges) == days + assert tuple(end for _, end in single_day_ranges) == days + assert (min(start for start, _ in single_day_ranges), max(end for _, end in single_day_ranges)) == multi_day_range + + +def test_timezone_adjustment_live_end_handles_offset_and_opt_in_cases() -> None: + pt_evening: Final = datetime(2026, 8, 6, 4, 30, tzinfo=timezone.utc) + ist_evening: Final = datetime(2026, 8, 5, 17, 0, tzinfo=timezone.utc) + utc_noon: Final = datetime(2026, 8, 5, 12, 0, tzinfo=timezone.utc) + + assert adjust_dates_for_timezone( + "2026-07-06", "2026-08-05", 420, include_current_utc_day=True, utc_now=pt_evening + ) == ("2026-07-06", "2026-08-06") + assert adjust_dates_for_timezone("2026-07-06", "2026-08-05", 420, utc_now=pt_evening) == ( + "2026-07-06", + "2026-08-05", + ) + assert adjust_dates_for_timezone( + "2026-07-01", "2026-08-04", 420, include_current_utc_day=True, utc_now=pt_evening + ) == ("2026-07-01", "2026-08-04") + assert adjust_dates_for_timezone( + "2026-07-07", "2026-08-06", -330, include_current_utc_day=True, utc_now=ist_evening + ) == ("2026-07-07", "2026-08-06") + assert adjust_dates_for_timezone( + "2026-07-06", "2026-08-05", None, include_current_utc_day=True, utc_now=pt_evening + ) == ("2026-07-06", "2026-08-05") + assert adjust_dates_for_timezone("2026-07-06", "2026-08-05", 0, include_current_utc_day=True, utc_now=utc_noon) == ( + "2026-07-06", + "2026-08-05", + ) + assert adjust_dates_for_timezone( + "2026-07-06", "2026-08-09", 420, include_current_utc_day=True, utc_now=pt_evening + ) == ("2026-07-06", "2026-08-09") + + +@pytest.mark.parametrize("offset_minutes", [None, 0, -330, 480]) +def test_aggregated_query_uses_the_caller_date_bounds(offset_minutes: int | None) -> None: + query = build_aggregated_sql( + _scope( + timezone_offset_minutes=offset_minutes, + start_date="2026-05-29", + end_date="2026-05-29", + ), + api_key_limit=constants.USAGE_TOP_API_KEYS_DEFAULT, + ) + + assert query.params[:2] == ("2026-05-29", "2026-05-29") + assert "date >= $1" in query.sql + assert "date <= $2" in query.sql + + +def test_aggregate_query_sums_all_savings_drivers_and_response_time() -> None: + query = build_aggregated_sql(_scope(), api_key_limit=constants.USAGE_TOP_API_KEYS_DEFAULT) + fields: Final = tuple(field for field in SpendMetrics.model_fields if field.endswith("_savings_spend")) + ( + "total_response_time_ms", + "timed_requests", + ) + + assert fields + assert all(f"SUM({field})" in query.sql for field in fields) + + +def test_aggregated_query_binds_sentinel_and_api_key_limit_after_scope_values() -> None: + scope = _scope(entity_ids=None, api_keys=("key-1",)) + + query = build_aggregated_sql(scope, api_key_limit=3) + + assert "api_key <> $4" in query.sql + assert "LIMIT $5" in query.sql + assert 'FROM "LiteLLM_DailyUserSpend"' in query.sql + assert query.params == ( + "2026-01-01", + "2026-01-31", + ["key-1"], + PTU_SENTINEL_API_KEY, + 3, + ) + + +@pytest.mark.parametrize("api_key_limit", [0, constants.USAGE_TOP_API_KEYS_MAX + 1]) +def test_aggregated_query_rejects_api_key_limits_outside_bounds(api_key_limit: int) -> None: + with pytest.raises(ValueError, match="api_key_limit"): + build_aggregated_sql(_scope(), api_key_limit=api_key_limit) + + +def test_entity_rollup_bounds_keys_and_reuses_scope_filters() -> None: + query = build_entity_rollup_sql( + _scope(table=DailyActivityTable.TEAM, entity_ids=None, api_keys=("key-1", "key-2")), + api_key_limit=3, + ) + + assert query.sql.count("COALESCE(\"team_id\", '') AS entity_id") == 3 + assert query.sql.count("GROUP BY date, COALESCE(\"team_id\", '')") == 2 + assert '"team_id" AS entity_id' not in query.sql + assert "JOIN top_api_keys USING (api_key)" in query.sql + assert "api_key = ANY($3::text[])" in query.sql + assert query.sql.count("api_key = ANY($3::text[])") == 4 + assert query.sql.count("api_key <> $4") == 2 + assert query.sql.count("ORDER BY SUM(spend::numeric) DESC, api_key") == 1 + assert "k.entity_id = e.entity_id" in query.sql + assert "LIMIT $5" in query.sql + assert query.params == ("2026-01-01", "2026-01-31", ["key-1", "key-2"], PTU_SENTINEL_API_KEY, 3) + + +@pytest.mark.parametrize("api_key_limit", [0, constants.USAGE_TOP_API_KEYS_MAX + 1]) +def test_entity_rollup_rejects_api_key_limits_outside_bounds(api_key_limit: int) -> None: + with pytest.raises(ValueError, match="api_key_limit"): + build_entity_rollup_sql(_scope(), api_key_limit=api_key_limit) + + +def test_search_query_escapes_pattern_metacharacters_and_binds_limit() -> None: + query = build_key_search_sql(_scope(entity_ids=None), search=r"foo%_\bar", limit=4) + + assert "OR api_key IN (" in query.sql + assert 'SELECT v.token FROM "LiteLLM_VerificationToken" v' in query.sql + assert 'LEFT JOIN "LiteLLM_UserTable" u ON u.user_id = v.user_id' in query.sql + assert 'SELECT d.token FROM "LiteLLM_DeletedVerificationToken" d' in query.sql + assert 'LEFT JOIN "LiteLLM_UserTable" u ON u.user_id = d.user_id' in query.sql + assert "d.key_alias ILIKE $3 ESCAPE" in query.sql + assert "d.user_id ILIKE $3 ESCAPE" in query.sql + assert "api_key ILIKE $3 ESCAPE" in query.sql + assert "v.key_alias ILIKE $3 ESCAPE" in query.sql + assert "v.user_id ILIKE $3 ESCAPE" in query.sql + assert "u.user_email ILIKE $3 ESCAPE" in query.sql + assert query.sql.count("ILIKE $3 ESCAPE") == 7 + assert "api_key <> $4" in query.sql + assert "ORDER BY SUM(spend::numeric) DESC, api_key" in query.sql + assert "LIMIT $5" in query.sql + assert query.params == ( + "2026-01-01", + "2026-01-31", + r"%foo\%\_\\bar%", + PTU_SENTINEL_API_KEY, + 4, + ) + + +def test_model_and_cache_key_queries_bind_filters_sentinel_and_limits() -> None: + model_query = build_model_top_keys_sql( + _scope(entity_ids=None), model_group="public-model", by_model_group=True, limit=5 + ) + leakage_query = build_cache_leakage_keys_sql(_scope(entity_ids=None), limit=20) + + assert "COALESCE(NULLIF(model_group, ''), model) = $3" in model_query.sql + assert "api_key <> $4" in model_query.sql + assert "ORDER BY SUM(spend::numeric) DESC, api_key" in model_query.sql + assert model_query.params == ("2026-01-01", "2026-01-31", "public-model", PTU_SENTINEL_API_KEY, 5) + assert "HAVING SUM(prompt_tokens) - SUM(cache_read_input_tokens) > 0" in leakage_query.sql + assert "ORDER BY SUM(prompt_tokens) - SUM(cache_read_input_tokens) DESC, api_key" in leakage_query.sql + assert leakage_query.params == ("2026-01-01", "2026-01-31", PTU_SENTINEL_API_KEY, 20) + + +@pytest.mark.parametrize( + "builder", + [ + lambda: build_key_search_sql(_scope(), search="x", limit=0), + lambda: build_model_top_keys_sql(_scope(), model_group="x", by_model_group=False, limit=0), + lambda: build_cache_leakage_keys_sql(_scope(), limit=0), + lambda: build_export_sql(_scope(), export_type=ExportType.DAILY, after=None, batch_size=0), + ], +) +def test_query_builders_reject_nonpositive_limits(builder) -> None: + with pytest.raises(ValueError, match="limit must be at least 1"): + builder() + + +@pytest.mark.parametrize( + ("export_type", "group_key", "key_filter", "joins"), + [ + (ExportType.DAILY, "''", "", ""), + (ExportType.DAILY_WITH_KEYS, "scoped.api_key", "api_key <> $3", 'LEFT JOIN "LiteLLM_VerificationToken"'), + (ExportType.DAILY_WITH_MODELS, "COALESCE(scoped.model, '')", "api_key <> $3", ""), + ( + ExportType.DAILY_WITH_USERS, + "COALESCE(vt.user_id, dvt.user_id, '')", + "api_key <> $3", + 'LEFT JOIN "LiteLLM_VerificationToken"', + ), + ], +) +def test_export_groups_by_requested_key_and_binds_cursor_after_scope( + export_type: ExportType, group_key: str, key_filter: str, joins: str +) -> None: + query = build_export_sql( + _scope(entity_ids=None), + export_type=export_type, + after=ExportCursor(date="2026-01-12", entity_id="user-2", group_key="group-3"), + batch_size=2, + ) + + assert group_key in query.sql + assert key_filter in query.sql + assert joins in query.sql + assert "(scoped.date, COALESCE(scoped.\"user_id\", '')," in query.sql + order_keys: Final = ( + "scoped.date, COALESCE(scoped.\"user_id\", '')", + *((group_key,) if export_type is not ExportType.DAILY else ()), + ) + assert f"ORDER BY {', '.join(order_keys)}" in query.sql + expected_limit_index: Final = "$6" if export_type is ExportType.DAILY else "$7" + assert f"LIMIT {expected_limit_index}" in query.sql + assert query.params == ( + "2026-01-01", + "2026-01-31", + *((PTU_SENTINEL_API_KEY,) if export_type is not ExportType.DAILY else ()), + "2026-01-12", + "user-2", + "group-3", + 2, + ) + + +@pytest.mark.parametrize("export_type", [ExportType.DAILY_WITH_KEYS, ExportType.DAILY_WITH_USERS]) +def test_export_uses_latest_deleted_key_metadata(export_type: ExportType) -> None: + query = build_export_sql(_scope(entity_ids=None), export_type=export_type, after=None, batch_size=2) + + assert 'FROM "LiteLLM_DeletedVerificationToken"' in query.sql + assert "ORDER BY deleted_at DESC" in query.sql + assert "COALESCE(vt.user_id, dvt.user_id)" in query.sql + + +@pytest.mark.parametrize("export_type", tuple(ExportType)) +def test_export_without_cursor_omits_cursor_predicate_and_parameters(export_type: ExportType) -> None: + query = build_export_sql( + _scope(entity_ids=None), + export_type=export_type, + after=None, + batch_size=2, + ) + + assert "WHERE TRUE AND (scoped.date" not in query.sql + assert query.params == ( + "2026-01-01", + "2026-01-31", + *((PTU_SENTINEL_API_KEY,) if export_type is not ExportType.DAILY else ()), + 2, + ) diff --git a/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py index 643944673a2..2699f9445c9 100644 --- a/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/unit/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -8,7 +8,8 @@ from unittest.mock import AsyncMock, MagicMock import pytest from fastapi import HTTPException -from mcp.types import CallToolResult, TextContent, Tool as MCPTool +from mcp.types import CallToolResult, TextContent +from mcp.types import Tool as MCPTool from openai.types.responses.tool_param import Mcp from litellm.proxy._experimental.mcp_server.faults.list_outcomes import AggregateToolListing @@ -110,9 +111,7 @@ def test_extract_tool_calls_from_chat_response_handles_tool_calls(): object="chat.completion", ) - tool_calls = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_chat_response( - response - ) + tool_calls = LiteLLM_Proxy_MCP_Handler._extract_tool_calls_from_chat_response(response) assert len(tool_calls) == 1 assert tool_calls[0]["function"]["name"] == "foo" @@ -182,9 +181,7 @@ def test_transform_mcp_tools_to_openai_uses_chat_format(monkeypatch): fake_transform_responses, ) - chat_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai( - ["tool"], target_format="chat" - ) + chat_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(["tool"], target_format="chat") resp_tools = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(["tool"]) assert chat_tools == [{"chat": True}] @@ -304,9 +301,7 @@ async def test_execute_tool_calls_strips_prefix_when_alias_differs_from_server_n ) from litellm.proxy._experimental.mcp_server import mcp_server_manager as _msm - _msm.global_mcp_server_manager._get_mcp_server_from_tool_name = MagicMock( - return_value=fake_server - ) + _msm.global_mcp_server_manager._get_mcp_server_from_tool_name = MagicMock(return_value=fake_server) tool_name = "my_deepwiki-read_wiki_structure" tool_calls = [ @@ -380,7 +375,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey fake_manager = types.SimpleNamespace( get_registry=MagicMock(return_value={}), - call_tool=AsyncMock(side_effect=HTTPException(status_code=500, detail="boom")) + call_tool=AsyncMock(side_effect=HTTPException(status_code=500, detail="boom")), ) monkeypatch.setattr( "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", @@ -388,9 +383,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey ) tool_name = "deepwiki-read_wiki_structure" - tool_calls = [ - {"id": "call-err", "function": {"name": tool_name, "arguments": "{}"}} - ] + tool_calls = [{"id": "call-err", "function": {"name": tool_name, "arguments": "{}"}}] user_auth = types.SimpleNamespace(api_key="test_key", user_id="test_user") @@ -408,10 +401,7 @@ async def test_execute_tool_calls_logs_failure_via_post_call_failure_hook(monkey post_call_failure_hook.assert_awaited_once() assert post_call_failure_hook.await_args is not None - assert ( - post_call_failure_hook.await_args.kwargs.get("route") - == "/responses/mcp/call_tool" - ) + assert post_call_failure_hook.await_args.kwargs.get("route") == "/responses/mcp/call_tool" @pytest.mark.asyncio @@ -434,9 +424,7 @@ async def test_execute_tool_calls_passes_litellm_call_id_and_trace_id_to_functio # NOTE: Don't patch via dotted string path here because `litellm.responses` # is a function attribute on the `litellm` package (shadowing the submodule), # which breaks monkeypatch's importpath resolution. - handler_module = importlib.import_module( - "litellm.responses.mcp.litellm_proxy_mcp_handler" - ) + handler_module = importlib.import_module("litellm.responses.mcp.litellm_proxy_mcp_handler") monkeypatch.setattr(handler_module, "function_setup", fake_function_setup) tool_name = "deepwiki-read_wiki_structure" @@ -516,7 +504,9 @@ async def test_execute_tool_calls_applies_post_call_hook_content(monkeypatch): logging_obj = MagicMock() logging_obj.model_call_details = {} - logging_obj.async_post_mcp_tool_call_hook = AsyncMock(return_value=CallToolResult(content=[TextContent(type="text", text="[REDACTED]")], is_error=True)) + logging_obj.async_post_mcp_tool_call_hook = AsyncMock( + return_value=CallToolResult(content=[TextContent(type="text", text="[REDACTED]")], is_error=True) + ) logging_obj.async_success_handler = AsyncMock() handler_module = importlib.import_module("litellm.responses.mcp.litellm_proxy_mcp_handler") monkeypatch.setattr(handler_module, "function_setup", lambda *_args, **_kwargs: (logging_obj, None)) @@ -679,9 +669,7 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch user_auth = types.SimpleNamespace(api_key="test_key", user_id="test_user") tools, _server_names = await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager( user_api_key_auth=user_auth, - mcp_tools_with_litellm_proxy=[ - {"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"} - ], + mcp_tools_with_litellm_proxy=[{"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"}], ) forwarded: Final = LiteLLM_Proxy_MCP_Handler._transform_mcp_tools_to_openai(tools) @@ -700,9 +688,7 @@ async def test_get_mcp_tools_from_manager_enables_list_tools_logging(monkeypatch def test_get_parent_request_tags_from_metadata(): - tags = LiteLLM_Proxy_MCP_Handler._get_parent_request_tags( - {"metadata": {"tags": ["team-a", "prod"]}} - ) + tags = LiteLLM_Proxy_MCP_Handler._get_parent_request_tags({"metadata": {"tags": ["team-a", "prod"]}}) assert tags == ["team-a", "prod"] @@ -739,9 +725,7 @@ async def test_get_mcp_tools_from_manager_forwards_request_tags(monkeypatch): await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager( user_api_key_auth=types.SimpleNamespace(api_key="k", user_id="u"), - mcp_tools_with_litellm_proxy=[ - {"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"} - ], + mcp_tools_with_litellm_proxy=[{"type": "mcp", "server_url": "litellm_proxy/mcp/deepwiki"}], request_tags=["team-a"], ) @@ -761,9 +745,7 @@ async def test_execute_tool_calls_exposes_sanitized_client_headers_to_logging(mo captured.update(kwargs) return None, None - handler_module = importlib.import_module( - "litellm.responses.mcp.litellm_proxy_mcp_handler" - ) + handler_module = importlib.import_module("litellm.responses.mcp.litellm_proxy_mcp_handler") monkeypatch.setattr(handler_module, "function_setup", fake_function_setup) tool_name = "deepwiki-read_wiki_structure" @@ -789,9 +771,7 @@ async def test_execute_tool_calls_propagates_request_tags_to_function_setup(monk captured.update(kwargs) return None, None - handler_module = importlib.import_module( - "litellm.responses.mcp.litellm_proxy_mcp_handler" - ) + handler_module = importlib.import_module("litellm.responses.mcp.litellm_proxy_mcp_handler") monkeypatch.setattr(handler_module, "function_setup", fake_function_setup) tool_name = "deepwiki-read_wiki_structure" @@ -1171,7 +1151,9 @@ def test_create_follow_up_input_keeps_each_reasoning_item_before_its_function_ca "function_call_output", "function_call_output", ] - assert [cast(dict[str, Any], item).get("id") or cast(dict[str, Any], item).get("call_id") for item in follow_up[1:5]] == [ + assert [ + cast(dict[str, Any], item).get("id") or cast(dict[str, Any], item).get("call_id") for item in follow_up[1:5] + ] == [ "rs_1", "call-1", "rs_2", @@ -1229,16 +1211,20 @@ async def test_mcp_follow_up_call_is_stateless_when_store_is_false( async def fake_aresponses(**kwargs: Any) -> ResponsesAPIResponse: captured_calls.append(kwargs) - return first_response if len(captured_calls) == 1 else ResponsesAPIResponse( - id="resp_follow_up", - created_at=1234567891, - model="gpt-5", - object="response", - status="completed", - output=[], - parallel_tool_calls=False, - tool_choice="auto", - tools=[], + return ( + first_response + if len(captured_calls) == 1 + else ResponsesAPIResponse( + id="resp_follow_up", + created_at=1234567891, + model="gpt-5", + object="response", + status="completed", + output=[], + parallel_tool_calls=False, + tool_choice="auto", + tools=[], + ) ) async def fake_process(**kwargs: Any) -> tuple[list[Any], dict[str, str]]: @@ -1279,12 +1265,14 @@ async def test_mcp_follow_up_call_is_stateless_when_store_is_false( @pytest.mark.asyncio async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: pytest.MonkeyPatch): - from litellm.proxy._experimental.mcp_server import operations - from litellm.proxy._experimental.mcp_server import mcp_server_manager + from litellm.proxy._experimental.mcp_server import mcp_server_manager, operations headers: Final = { - "x-app-id": "app-a", "x-nuid": "user-a", "x-user-id": "identity-a", - "x-mcp-deepwiki-authorization": "upstream-sentinel", "authorization": "proxy-sentinel", + "x-app-id": "app-a", + "x-nuid": "user-a", + "x-user-id": "identity-a", + "x-mcp-deepwiki-authorization": "upstream-sentinel", + "authorization": "proxy-sentinel", } manager: Final = types.SimpleNamespace( get_registry=MagicMock(return_value={}), @@ -1298,12 +1286,21 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py monkeypatch.setattr(operations, "_get_allowed_mcp_servers", AsyncMock(return_value=[])) monkeypatch.setattr(operations, "function_setup", setup) response: Final = ResponsesAPIResponse( - id="resp_test", created_at=1234567891, model="test-model", object="response", - status="completed", output=[], parallel_tool_calls=False, tool_choice="auto", tools=[], + id="resp_test", + created_at=1234567891, + model="test-model", + object="response", + status="completed", + output=[], + parallel_tool_calls=False, + tool_choice="auto", + tools=[], ) monkeypatch.setattr(responses_main, "aresponses", AsyncMock(return_value=response)) result: Final = await responses_main.aresponses_api_with_mcp( - input="hi", model="test-model", tools=[{"type": "mcp", "server_url": "litellm_proxy"}], + input="hi", + model="test-model", + tools=[{"type": "mcp", "server_url": "litellm_proxy"}], secret_fields={"raw_headers": headers}, ) assert result is response @@ -1311,3 +1308,84 @@ async def test_responses_discovery_logs_sanitized_caller_headers(monkeypatch: py logged: Final = setup.call_args.kwargs["metadata"]["headers"] assert logged == {"x-app-id": "app-a", "x-nuid": "user-a", "x-user-id": "identity-a"} assert headers["x-mcp-deepwiki-authorization"] == "upstream-sentinel" + + +def _toolset_gateway_manager(toolset_id: str, server_id: str) -> types.SimpleNamespace: + return types.SimpleNamespace( + get_registry=MagicMock(return_value={}), + get_allowed_mcp_servers=AsyncMock(return_value=[]), + get_mcp_servers_from_ids=MagicMock(return_value=[]), + get_mcp_server_by_name=MagicMock(return_value=None), + get_toolset_by_name_cached=AsyncMock(return_value=types.SimpleNamespace(toolset_id=toolset_id)), + resolve_toolset_tool_permissions=AsyncMock(return_value={server_id: ["add"]}), + ) + + +async def _tools_listing_kwargs_for_toolset_url(monkeypatch, team_toolset_id: str) -> dict[str, object]: + from litellm.proxy._experimental.mcp_server.ui_session_utils import granted_toolset_ids + from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LitellmUserRoles, UserAPIKeyAuth + + mock_get_tools = AsyncMock(return_value=AggregateToolListing(tools=[], outcomes={})) + monkeypatch.setattr("litellm.proxy._experimental.mcp_server.server._get_tools_from_mcp_servers", mock_get_tools) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + _toolset_gateway_manager("ts-granted", "srv-1"), + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MagicMock()) + + async def team_permission(context: UserAPIKeyAuth) -> LiteLLM_ObjectPermissionTable: + return LiteLLM_ObjectPermissionTable(object_permission_id="op-team", mcp_toolsets=[team_toolset_id]) + + async def granted_through_team(context: UserAPIKeyAuth) -> frozenset[str]: + return await granted_toolset_ids(context, team_object_permission=team_permission, require_key_access=False) + + team_key: Final = UserAPIKeyAuth(api_key="sk-team", team_id="team-1", user_role=LitellmUserRoles.INTERNAL_USER) + await LiteLLM_Proxy_MCP_Handler._get_mcp_tools_from_manager( + user_api_key_auth=team_key, + mcp_tools_with_litellm_proxy=[{"type": "mcp", "server_url": "litellm_proxy/mcp/team-toolset"}], + granted_toolsets=granted_through_team, + ) + assert mock_get_tools.await_args is not None + return mock_get_tools.await_args.kwargs + + +@pytest.mark.asyncio +async def test_toolset_gateway_url_scopes_a_team_granted_toolset_for_a_key_without_its_own_grant(monkeypatch): + kwargs: Final = await _tools_listing_kwargs_for_toolset_url(monkeypatch, team_toolset_id="ts-granted") + scoped = kwargs["user_api_key_auth"].object_permission + assert scoped is not None + assert scoped.mcp_servers == ["srv-1"] + assert scoped.mcp_tool_permissions == {"srv-1": ["add"]} + assert kwargs["mcp_servers"] is None + + +@pytest.mark.asyncio +async def test_toolset_gateway_url_skips_a_toolset_the_team_does_not_grant(monkeypatch): + kwargs: Final = await _tools_listing_kwargs_for_toolset_url(monkeypatch, team_toolset_id="ts-other") + assert kwargs["user_api_key_auth"].object_permission is None + assert kwargs["mcp_servers"] is None + + +@pytest.mark.asyncio +async def test_apply_toolset_permissions_pins_the_auth_to_explicit_grants_only(monkeypatch: pytest.MonkeyPatch): + """A toolset gateway URL must not widen to operator-open (allow_all_keys) servers.""" + from litellm.proxy._types import UserAPIKeyAuth + + fake_manager = types.SimpleNamespace( + resolve_toolset_tool_permissions=AsyncMock(return_value={"srv-1": ["add"]}), + ) + monkeypatch.setattr( + "litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager", + fake_manager, + ) + + scoped = await LiteLLM_Proxy_MCP_Handler._apply_toolset_permissions( + resolved_toolset_ids=["ts-1"], + resolved_mcp_servers=[], + user_api_key_auth=UserAPIKeyAuth(api_key="sk-test", user_id="u1"), + ) + + assert scoped.mcp_explicit_grants_only is True + assert scoped.object_permission is not None + assert scoped.object_permission.mcp_servers == ["srv-1"] + assert scoped.object_permission.mcp_tool_permissions == {"srv-1": ["add"]} diff --git a/tests/unit/test_check_migrations_no_data_rewrites.py b/tests/unit/test_check_migrations_no_data_rewrites.py index c5d3cdd9073..fb1053ff887 100644 --- a/tests/unit/test_check_migrations_no_data_rewrites.py +++ b/tests/unit/test_check_migrations_no_data_rewrites.py @@ -10,6 +10,8 @@ import importlib.util import sys from pathlib import Path +import pytest + _CHECKER_PATH = Path(__file__).resolve().parents[1] / "code_coverage_tests" / "check_migrations_no_data_rewrites.py" _SPEC = importlib.util.spec_from_file_location("check_migrations_no_data_rewrites", _CHECKER_PATH) assert _SPEC is not None and _SPEC.loader is not None @@ -205,6 +207,96 @@ class TestDefaultedColumnsOnRequestLogTables: assert 'ADD COLUMN ... DEFAULT on "LiteLLM_SpendLogs" rewrites existing rows at boot' in rendered +class TestIndexesOnLogTables: + """Every CREATE INDEX on a request-log table is rejected: a plain one blocks writes for + the whole build and a concurrent one fails on a partitioned parent, so the migration job + (litellm_proxy_extras/request_log_indexes.py) builds those instead.""" + + def test_the_original_spend_log_index_statement_is_flagged(self, tmp_path): + sql = ( + "-- CreateIndex\n" + 'CREATE INDEX IF NOT EXISTS "LiteLLM_SpendLogs_api_key_startTime_idx" ' + 'ON "LiteLLM_SpendLogs"("api_key", "startTime");' + ) + assert _keywords(tmp_path, sql) == ('CREATE INDEX on "LiteLLM_SpendLogs"',) + + def test_the_original_concurrent_call_id_index_statement_is_flagged(self, tmp_path): + sql = ( + 'CREATE INDEX CONCURRENTLY IF NOT EXISTS "LiteLLM_SpendLogs_litellm_call_id_idx" ' + 'ON "LiteLLM_SpendLogs"("litellm_call_id");' + ) + assert _keywords(tmp_path, sql) == ('CREATE INDEX on "LiteLLM_SpendLogs"',) + + def test_unique_index_with_if_not_exists_on_error_logs_is_flagged(self, tmp_path): + sql = 'CREATE UNIQUE INDEX IF NOT EXISTS "ix" ON "LiteLLM_ErrorLogs" ("request_id");' + assert _keywords(tmp_path, sql) == ('CREATE INDEX on "LiteLLM_ErrorLogs"',) + + def test_a_unique_concurrent_index_on_error_logs_is_flagged(self, tmp_path): + sql = 'CREATE UNIQUE INDEX CONCURRENTLY "ix" ON "LiteLLM_ErrorLogs" ("request_id");' + assert _keywords(tmp_path, sql) == ('CREATE INDEX on "LiteLLM_ErrorLogs"',) + + def test_lowercase_schema_qualified_and_only_forms_are_flagged(self, tmp_path): + sql = ( + 'create index on "public"."LiteLLM_SpendLogs" ("api_key");\n' + 'CREATE INDEX "ix" ON ONLY "LiteLLM_SpendLogs" ("api_key");\n' + 'CREATE INDEX CONCURRENTLY "iy" ON "public"."LiteLLM_SpendLogs" ("api_key");' + ) + assert _keywords(tmp_path, sql) == ('CREATE INDEX on "LiteLLM_SpendLogs"',) * 3 + + def test_a_comment_between_on_and_the_table_is_flagged(self, tmp_path): + sql = 'CREATE INDEX CONCURRENTLY "ix" ON /* table */ "LiteLLM_SpendLogs" ("api_key");' + assert _keywords(tmp_path, sql) == ('CREATE INDEX on "LiteLLM_SpendLogs"',) + + def test_a_concurrent_index_with_comments_and_line_breaks_is_flagged(self, tmp_path): + sql = ( + "-- CreateIndex\n" + 'CREATE INDEX CONCURRENTLY IF NOT EXISTS "ix"\n' + ' ON "LiteLLM_SpendLogs" /* partitioned in some deployments */\n' + ' ("api_key", "startTime");\n' + ) + assert _keywords(tmp_path, sql) == ('CREATE INDEX on "LiteLLM_SpendLogs"',) + + def test_indexes_on_a_non_log_table_pass_concurrent_or_not(self, tmp_path): + sql = ( + 'CREATE INDEX "ix" ON "LiteLLM_VerificationToken" ("token");\n' + 'CREATE INDEX CONCURRENTLY "iy" ON "LiteLLM_VerificationToken" ("token");' + ) + assert _keywords(tmp_path, sql) == () + + def test_an_index_run_by_execute_is_flagged(self, tmp_path): + sql = 'DO $$ BEGIN EXECUTE \'CREATE INDEX "ix" ON "LiteLLM_SpendLogs" ("api_key")\'; END $$;' + assert _keywords(tmp_path, sql) == ('CREATE INDEX on "LiteLLM_SpendLogs"',) + + def test_a_marker_does_not_exempt_the_index(self, tmp_path): + sql = ( + '-- data-migration-ok: table is empty at this point\nCREATE INDEX "ix" ON "LiteLLM_SpendLogs" ("api_key");' + ) + assert _keywords(tmp_path, sql) == ('CREATE INDEX on "LiteLLM_SpendLogs"',) + + def test_a_marker_on_a_rewrite_still_leaves_the_index_below_it_flagged(self, tmp_path): + sql = ( + "-- data-migration-ok: one row\n" + 'UPDATE "LiteLLM_SpendLogs" SET "api_key" = \'k\' WHERE "request_id" = \'r\';\n' + 'CREATE INDEX CONCURRENTLY "ix" ON "LiteLLM_SpendLogs" ("api_key");' + ) + assert _keywords(tmp_path, sql) == ('CREATE INDEX on "LiteLLM_SpendLogs"',) + + def test_render_points_at_the_migration_job_index_list(self, tmp_path): + sql = 'CREATE INDEX CONCURRENTLY "ix" ON "public"."LiteLLM_SpendLogs" ("api_key");' + rendered = _scan(tmp_path, sql)[0].render() + assert "20260101000000_fixture/migration.sql:1" in rendered + assert "blocks writes until the build finishes, or fails on a partitioned table" in rendered + assert "REQUEST_LOG_INDEXES in litellm_proxy_extras/request_log_indexes.py" in rendered + + @pytest.mark.parametrize( + "name", + ("20260823000000_add_spend_logs_api_key_starttime_index", "20260831120001_spend_logs_litellm_call_id_index"), + ) + def test_the_inert_index_migrations_scan_clean_without_a_grandfather(self, name): + assert checker.scan_migration(checker.MIGRATIONS_DIR / name) == () + assert name not in checker.GRANDFATHERED + + class TestInsert: def test_insert_values_is_bounded_and_passes(self, tmp_path): assert _keywords(tmp_path, "INSERT INTO \"Foo\" (\"id\") VALUES ('a'), ('b');") == () diff --git a/ui/litellm-dashboard/eslint-budgets.json b/ui/litellm-dashboard/eslint-budgets.json index 44294b5fa97..83b7190b355 100644 --- a/ui/litellm-dashboard/eslint-budgets.json +++ b/ui/litellm-dashboard/eslint-budgets.json @@ -4,7 +4,7 @@ "complexity": { "max": 140, "target": 80 }, "max-depth": { "max": 70, "target": 30 }, "local/no-large-inline-object-arg": { "max": 551, "target": 300 }, - "local/no-long-condition-chain": { "max": 196, "target": 120 }, + "local/no-long-condition-chain": { "max": 194, "target": 120 }, "testing-library/no-container": { "max": 133, "target": 50 }, "testing-library/no-node-access": { "max": 707, "target": 500 }, "testing-library/prefer-screen-queries": { "max": 18, "target": 18 } diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index daf12d11743..2465d07129c 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -1140,15 +1140,7 @@ "count": 1 }, "react-hooks/set-state-in-effect": { - "count": 3 - } - }, - "src/app/(dashboard)/usage/_components/hooks/usePaginatedDailyActivity.ts": { - "react-hooks/refs": { - "count": 1 - }, - "react-hooks/set-state-in-effect": { - "count": 1 + "count": 2 } }, "src/app/(dashboard)/users/_components/BulkEditUsers.tsx": { @@ -1279,11 +1271,6 @@ "count": 1 } }, - "src/components/EntityUsageExport/utils.ts": { - "max-params": { - "count": 3 - } - }, "src/components/GuardrailSettingsView.tsx": { "no-nested-ternary": { "count": 1 @@ -1786,13 +1773,13 @@ "count": 1 }, "max-params": { - "count": 21 + "count": 15 }, "no-nested-ternary": { "count": 5 }, "no-restricted-syntax": { - "count": 147 + "count": 146 }, "prefer-const": { "count": 31 diff --git a/ui/litellm-dashboard/public/assets/agent-traces-preview.png b/ui/litellm-dashboard/public/assets/agent-traces-preview.png deleted file mode 100644 index 34569e26331..00000000000 Binary files a/ui/litellm-dashboard/public/assets/agent-traces-preview.png and /dev/null differ diff --git a/ui/litellm-dashboard/public/assets/logos/microsoft_365.svg b/ui/litellm-dashboard/public/assets/logos/microsoft_365.svg new file mode 100644 index 00000000000..e053ac831fb --- /dev/null +++ b/ui/litellm-dashboard/public/assets/logos/microsoft_365.svg @@ -0,0 +1 @@ + diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.test.tsx index 8d328c4c329..95abc4e22cb 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.test.tsx @@ -1,9 +1,18 @@ import { fireEvent, render, screen } from "@testing-library/react"; -import { describe, expect, it, vi } from "vitest"; +import { beforeEach, describe, expect, it, vi } from "vitest"; -import type { DailyData, KeyMetricWithMetadata, SpendMetrics } from "@/components/UsagePage/types"; +import type { components } from "@/lib/http/schema"; +import type { KeySpendActivityRow } from "@/components/UsagePage/dailyActivityApi"; +import { EMPTY_DAILY_ACTIVITY_METADATA } from "@/components/UsagePage/dailyActivityApi"; +import type { DailyData, SpendMetrics } from "@/components/UsagePage/types"; import type { DailyActivityRange } from "./useDailyActivityRange"; +const mockCacheLeakageKeysCall = vi.fn(); + +vi.mock("@/components/networking", () => ({ + cacheLeakageKeysCall: (...args: unknown[]) => mockCacheLeakageKeysCall(...args), +})); + vi.mock("@/components/shared/advanced_date_picker", () => ({ __esModule: true, default: () =>
, @@ -11,8 +20,9 @@ vi.mock("@/components/shared/advanced_date_picker", () => ({ import CacheLeakageCard from "./CacheLeakageCard"; -const baseMetrics = (overrides: Partial): SpendMetrics => ({ +const baseMetrics = (overrides: Partial): components["schemas"]["SpendMetrics"] => ({ spend: 0, + flat_cost: 0, prompt_tokens: 0, completion_tokens: 0, total_tokens: 0, @@ -21,27 +31,22 @@ const baseMetrics = (overrides: Partial): SpendMetrics => ({ failed_requests: 0, cache_read_input_tokens: 0, cache_creation_input_tokens: 0, + compression_saved_tokens: 0, + compression_savings_spend: 0, + prompt_caching_savings_spend: 0, + gateway_injected_caching_savings_spend: 0, + autorouter_savings_spend: 0, + total_response_time_ms: 0, + timed_requests: 0, ...overrides, }); -const key = (alias: string, metrics: Partial): KeyMetricWithMetadata => ({ +const keyRow = (hash: string, alias: string, metrics: Partial): KeySpendActivityRow => ({ + api_key: hash, metrics: baseMetrics(metrics), metadata: { key_alias: alias, team_id: null }, }); -const dayWithKeys = (date: string, apiKeys: Record): DailyData => ({ - date, - metrics: baseMetrics({}), - breakdown: { - models: {}, - model_groups: {}, - mcp_servers: {}, - providers: {}, - api_keys: apiKeys, - entities: {}, - }, -}); - const dayWithModels = (date: string, models: Record>): DailyData => ({ date, metrics: baseMetrics({}), @@ -67,27 +72,37 @@ const renderWith = (results: DailyData[], overrides: Partial dateValue: {}, onDateChange: vi.fn(), results, + metadata: EMPTY_DAILY_ACTIVITY_METADATA, loading: false, - isFetchingMore: false, - progress: { currentPage: 1, totalPages: 1 }, - cancelled: false, failed: false, - cancel: vi.fn(), + scope: { + accessToken: "test-token", + startTime: new Date(2025, 0, 1), + endTime: new Date(2025, 0, 31), + userId: null, + apiKey: null, + }, ...overrides, }} />, ); describe("CacheLeakageCard", () => { - it("ranks leaking keys by uncached prompt tokens and shows cache hit ratio", () => { - renderWith([ - dayWithKeys("2026-07-12", { - "hash-caching": key("caching-key", { prompt_tokens: 1000, cache_read_input_tokens: 900 }), - "hash-leaky": key("leaky-key", { prompt_tokens: 10000, cache_read_input_tokens: 0 }), - }), - ]); + beforeEach(() => { + mockCacheLeakageKeysCall.mockReset(); + mockCacheLeakageKeysCall.mockResolvedValue({ api_keys: [] }); + }); - expect(screen.getByText("leaky-key")).toBeInTheDocument(); + it("ranks leaking keys from the server-ranked key list and shows cache hit ratio", async () => { + mockCacheLeakageKeysCall.mockResolvedValue({ + api_keys: [ + keyRow("hash-caching", "caching-key", { prompt_tokens: 1000, cache_read_input_tokens: 900 }), + keyRow("hash-leaky", "leaky-key", { prompt_tokens: 10000, cache_read_input_tokens: 0 }), + ], + }); + renderWith([]); + + expect(await screen.findByText("leaky-key")).toBeInTheDocument(); expect(screen.getByText("0.0%")).toBeInTheDocument(); expect(screen.getByText("90.0%")).toBeInTheDocument(); [ @@ -97,23 +112,42 @@ describe("CacheLeakageCard", () => { ].forEach((info) => expect(screen.getByLabelText(info)).toBeInTheDocument()); }); - it("sorts by the clicked column, worst cache hit rate first", () => { - renderWith([ - dayWithKeys("2026-07-12", { - "hash-a": key("alpha", { + it("asks the server for the key ranking under the activity scope", async () => { + renderWith([], { + scope: { + accessToken: "test-token", + startTime: new Date(2025, 0, 1), + endTime: new Date(2025, 0, 31), + userId: "u1", + apiKey: "hash-1", + }, + }); + + await screen.findByText("No key usage in this range."); + expect(mockCacheLeakageKeysCall).toHaveBeenCalledWith( + expect.objectContaining({ entityIds: ["u1"], apiKey: "hash-1", includeCurrentUtcDay: true }), + ); + }); + + it("sorts by the clicked column, worst cache hit rate first", async () => { + mockCacheLeakageKeysCall.mockResolvedValue({ + api_keys: [ + keyRow("hash-a", "alpha", { prompt_tokens: 10000, cache_read_input_tokens: 9000, prompt_caching_savings_spend: 9.0, }), - "hash-b": key("bravo", { + keyRow("hash-b", "bravo", { prompt_tokens: 500, cache_read_input_tokens: 50, prompt_caching_savings_spend: 0.05, }), - }), - ]); + ], + }); + renderWith([]); const firstDataRow = () => screen.getAllByRole("row")[1]; + expect(await screen.findByText("alpha")).toBeInTheDocument(); expect(firstDataRow()).toHaveTextContent("alpha"); fireEvent.click(screen.getByText("Cache hit rate")); @@ -123,7 +157,7 @@ describe("CacheLeakageCard", () => { expect(firstDataRow()).toHaveTextContent("alpha"); }); - it("switches to the model view and lists models from every provider", () => { + it("switches to the model view and lists models from every provider", async () => { renderWith([ dayWithModels("2026-07-12", { "claude-sonnet-5": { prompt_tokens: 5000, cache_read_input_tokens: 0 }, @@ -131,51 +165,32 @@ describe("CacheLeakageCard", () => { }), ]); - fireEvent.click(screen.getByText("By model")); + fireEvent.click(await screen.findByText("By model")); expect(screen.getByText("Cache leakage by model")).toBeInTheDocument(); expect(screen.getByText("claude-sonnet-5")).toBeInTheDocument(); expect(screen.getByText("vertex_ai/gemini-2.5-pro")).toBeInTheDocument(); }); - it("shows an empty state when no key used tokens in the range", () => { - renderWith([dayWithKeys("2026-07-12", {})]); + it("shows an empty state when no key used tokens in the range", async () => { + renderWith([]); - expect(screen.getByText("No key usage in this range.")).toBeInTheDocument(); + expect(await screen.findByText("No key usage in this range.")).toBeInTheDocument(); expect(screen.queryByRole("table")).not.toBeInTheDocument(); }); - it("tells the user the table is still filling in while fallback pages stream", () => { - const day = dayWithKeys("2026-07-12", { - "hash-leaky": key("leaky-key", { prompt_tokens: 10000, cache_read_input_tokens: 0 }), - }); - renderWith([day], { isFetchingMore: true }); + it("reports a load failure instead of claiming the range is empty", async () => { + mockCacheLeakageKeysCall.mockRejectedValue(new Error("route unavailable")); + renderWith([]); - expect(screen.getByRole("table")).toBeInTheDocument(); - expect( - screen.getByText("Data is still loading; rows and totals will update as the rest of the range arrives."), - ).toBeInTheDocument(); + expect(await screen.findByText("Could not load key usage for this range.")).toBeInTheDocument(); + expect(screen.queryByText("No key usage in this range.")).not.toBeInTheDocument(); }); - it("keeps the streaming note off while a fresh range loads over the previous range's rows", () => { - const day = dayWithKeys("2026-07-12", { - "hash-leaky": key("leaky-key", { prompt_tokens: 10000, cache_read_input_tokens: 0 }), - }); - renderWith([day], { loading: true }); + it("shows a loading state while the key ranking is in flight", () => { + mockCacheLeakageKeysCall.mockReturnValue(new Promise(() => {})); + renderWith([]); - expect( - screen.queryByText("Data is still loading; rows and totals will update as the rest of the range arrives."), - ).not.toBeInTheDocument(); - }); - - it("drops the streaming note once the range has settled", () => { - const day = dayWithKeys("2026-07-12", { - "hash-leaky": key("leaky-key", { prompt_tokens: 10000, cache_read_input_tokens: 0 }), - }); - renderWith([day]); - - expect( - screen.queryByText("Data is still loading; rows and totals will update as the rest of the range arrives."), - ).not.toBeInTheDocument(); + expect(screen.getByText("Loading...")).toBeInTheDocument(); }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.tsx index cfac77788b7..7a63ccd1ee4 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.tsx @@ -8,8 +8,17 @@ import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@ import { Tabs, TabsList, TabsTrigger } from "@/components/ui/tabs"; import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/components/ui/tooltip"; import { formatNumberWithCommas } from "@/utils/dataUtils"; -import { CacheLeakageDimension, CacheLeakageRow, computeCacheLeakage, pct, usd } from "./costOptimizationUtils"; +import { + CacheLeakageDimension, + CacheLeakageRow, + computeCacheLeakage, + leakageRowsFromKeyRows, + netSavingsPerCachedToken, + pct, + usd, +} from "./costOptimizationUtils"; import { DailyActivityRange } from "./useDailyActivityRange"; +import { useCacheLeakageKeys } from "./useCacheLeakageKeys"; interface CacheLeakageCardProps { activity: DailyActivityRange; @@ -80,11 +89,20 @@ const SortableHead = ({ }; const CacheLeakageCard: React.FC = ({ activity }) => { - const { results, loading, isFetchingMore } = activity; + const { results, loading } = activity; const [dimension, setDimension] = useState("key"); const [sort, setSort] = useState({ column: "potentialSavings", dir: "desc" }); - const leakage = useMemo(() => computeCacheLeakage(results, dimension), [results, dimension]); - const rows = useMemo(() => [...leakage.rows].sort((a, b) => compareRows(a, b, sort)), [leakage.rows, sort]); + const leakageRate = useMemo(() => netSavingsPerCachedToken(results), [results]); + const keyLeakage = useCacheLeakageKeys(activity, dimension === "key"); + const unsortedRows = useMemo( + () => + dimension === "key" + ? leakageRowsFromKeyRows(keyLeakage.rows, leakageRate) + : computeCacheLeakage(results, "model").rows, + [dimension, keyLeakage.rows, leakageRate, results], + ); + const rows = useMemo(() => [...unsortedRows].sort((a, b) => compareRows(a, b, sort)), [unsortedRows, sort]); + const rowsLoading = dimension === "key" ? keyLeakage.loading : loading; const onSort = (column: SortColumn) => setSort((prev) => @@ -96,6 +114,10 @@ const CacheLeakageCard: React.FC = ({ activity }) => { const subject = dimension === "model" ? "Models" : "Keys"; const firstColumn = dimension === "model" ? "Model" : "Key"; const emptyNoun = dimension === "model" ? "model" : "key"; + const emptyMessage = + dimension === "key" && keyLeakage.failed + ? "Could not load key usage for this range." + : `No ${emptyNoun} usage in this range.`; return ( @@ -119,14 +141,9 @@ const CacheLeakageCard: React.FC = ({ activity }) => { - {rows.length > 0 && isFetchingMore && ( -

- Data is still loading; rows and totals will update as the rest of the range arrives. -

- )} {rows.length === 0 ? (

- {loading || isFetchingMore ? "Loading..." : `No ${emptyNoun} usage in this range.`} + {rowsLoading ? "Loading..." : emptyMessage}

) : (
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.activity.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.activity.test.tsx index f8336f5ab56..204c4b1a409 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.activity.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.activity.test.tsx @@ -3,8 +3,7 @@ import { fireEvent, render, waitFor, screen } from "@testing-library/react"; import { describe, expect, it, vi } from "vitest"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; -const mockUserDailyActivityCall = vi.fn(); -const mockUserDailyActivityAggregatedCall = vi.fn(); +const mockDailyActivityAggregatedCall = vi.fn(); const { useAuthorizedMock, mockToolSpendResponse } = vi.hoisted(() => ({ useAuthorizedMock: vi.fn(), mockToolSpendResponse: { by_tool: [], daily: [], start_date: null, end_date: null }, @@ -15,8 +14,8 @@ vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ })); vi.mock("@/components/networking", () => ({ - userDailyActivityCall: (...args: unknown[]) => mockUserDailyActivityCall(...args), - userDailyActivityAggregatedCall: (...args: unknown[]) => mockUserDailyActivityAggregatedCall(...args), + dailyActivityAggregatedCall: (...args: unknown[]) => mockDailyActivityAggregatedCall(...args), + cacheLeakageKeysCall: vi.fn().mockResolvedValue({ api_keys: [] }), getToolSpend: vi.fn().mockResolvedValue(mockToolSpendResponse), getGeneralSettingsCall: vi.fn().mockResolvedValue([]), organizationListCall: vi.fn().mockResolvedValue([]), @@ -53,7 +52,7 @@ const singlePage = { describe("CostOptimizationView daily activity", () => { it("fetches daily activity once for the page and shares it with every tab that needs it", async () => { - mockUserDailyActivityAggregatedCall.mockResolvedValue(singlePage); + mockDailyActivityAggregatedCall.mockResolvedValue(singlePage); useAuthorizedMock.mockReturnValue({ accessToken: "test-token", userId: "u1", userRole: "proxy_admin" }); const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); @@ -63,25 +62,17 @@ describe("CostOptimizationView daily activity", () => { , ); - await waitFor(() => expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalledTimes(1)); + await waitFor(() => expect(mockDailyActivityAggregatedCall).toHaveBeenCalledTimes(1)); fireEvent.click(screen.getByRole("tab", { name: "Prompt Caching" })); await screen.findByTestId("caching-settings"); - expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalledTimes(1); - expect(mockUserDailyActivityCall).not.toHaveBeenCalled(); - expect(screen.queryByText(/Currently fetching spend data/)).not.toBeInTheDocument(); + expect(mockDailyActivityAggregatedCall).toHaveBeenCalledTimes(1); }); - it("shows the fetch-progress banner while the paginated fallback streams pages in", async () => { - mockUserDailyActivityAggregatedCall.mockReset(); - mockUserDailyActivityCall.mockReset(); - mockUserDailyActivityAggregatedCall.mockRejectedValue(new Error("aggregated unavailable")); - mockUserDailyActivityCall.mockImplementation((...args: unknown[]) => - args[3] === 1 - ? Promise.resolve({ results: [], metadata: { total_pages: 3, has_more: true, page: 1 } }) - : new Promise(() => {}), - ); + it("surfaces a failure alert when the aggregated fetch fails", async () => { + mockDailyActivityAggregatedCall.mockReset(); + mockDailyActivityAggregatedCall.mockRejectedValue(new Error("aggregated unavailable")); useAuthorizedMock.mockReturnValue({ accessToken: "test-token", userId: "u1", userRole: "proxy_admin" }); const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); @@ -91,7 +82,6 @@ describe("CostOptimizationView daily activity", () => { , ); - expect(await screen.findByText(/Currently fetching spend data: fetched 1 \/ 3 pages/)).toBeInTheDocument(); - expect(screen.getByRole("button", { name: "Stop" })).toBeInTheDocument(); + expect(await screen.findByText(/Fetching spend data failed/)).toBeInTheDocument(); }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.test.tsx index 028367555a1..a2f0ca5edc9 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.test.tsx @@ -11,12 +11,7 @@ vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ vi.mock("@/components/networking", () => ({ organizationListCall: vi.fn().mockResolvedValue([]), - userDailyActivityCall: vi - .fn() - .mockResolvedValue({ results: [], metadata: { total_pages: 1, has_more: false, page: 1 } }), - userDailyActivityAggregatedCall: vi - .fn() - .mockResolvedValue({ results: [], metadata: { total_pages: 1, has_more: false, page: 1 } }), + dailyActivityAggregatedCall: vi.fn().mockResolvedValue({ results: [], metadata: {} }), })); vi.mock("./UsageTab", () => ({ __esModule: true, default: () =>
})); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.tsx index 25bd3de0382..0aa4f88495a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.tsx @@ -4,7 +4,7 @@ import React from "react"; import { Info, PiggyBank } from "lucide-react"; import useCan from "@/app/(dashboard)/hooks/useCan"; -import PaginationStatusAlerts from "@/components/shared/PaginationStatusAlerts"; +import { Alert, AlertDescription } from "@/components/shared/Alert"; import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; import { PageHeader } from "@/components/shared/PageHeader"; import UsageTab from "./UsageTab"; @@ -84,13 +84,14 @@ const CostOptimizationView: React.FC = ({ accessToken

- + {activity.failed && ( + + + Fetching spend data failed, so the savings below may be empty rather than final. Reload the page to try + again. + + + )} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingTab.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingTab.test.tsx index 35464c5852e..a4a06cb6db0 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingTab.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingTab.test.tsx @@ -1,6 +1,8 @@ import { fireEvent, render, waitFor, screen } from "@testing-library/react"; import { describe, expect, it, vi } from "vitest"; +import { EMPTY_DAILY_ACTIVITY_METADATA } from "@/components/UsagePage/dailyActivityApi"; + const mockGetGeneralSettingsCall = vi.fn(); vi.mock("@/components/networking", () => ({ @@ -46,12 +48,16 @@ describe("PromptCachingTab", () => { dateValue: {}, onDateChange: vi.fn(), results: [], + metadata: EMPTY_DAILY_ACTIVITY_METADATA, loading: false, - isFetchingMore: false, - progress: { currentPage: 1, totalPages: 1 }, - cancelled: false, failed: false, - cancel: vi.fn(), + scope: { + accessToken: "test-token", + startTime: null, + endTime: null, + userId: null, + apiKey: null, + }, }; render(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/UsageTab.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/UsageTab.test.tsx index c62208aacc5..1d88e25396b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/UsageTab.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/UsageTab.test.tsx @@ -3,6 +3,7 @@ import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; import type { ToolSpendResponse } from "@/components/networking"; +import { EMPTY_DAILY_ACTIVITY_METADATA } from "@/components/UsagePage/dailyActivityApi"; import type { DailyData, SpendMetrics } from "@/components/UsagePage/types"; const mockGetToolSpend = vi.fn(); @@ -119,12 +120,16 @@ const renderWith = (results: DailyData[], options: RenderOptions = {}) => { dateValue: { from, to }, onDateChange: vi.fn(), results, + metadata: EMPTY_DAILY_ACTIVITY_METADATA, loading: false, - isFetchingMore: false, - progress: { currentPage: 1, totalPages: 1 }, - cancelled: false, failed: false, - cancel: vi.fn(), + scope: { + accessToken: "test-token", + startTime: from, + endTime: to, + userId: null, + apiKey: null, + }, }} />, ); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/UsageTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/UsageTab.tsx index 83b202590cb..2d673d96296 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/UsageTab.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/UsageTab.tsx @@ -44,7 +44,7 @@ const EMPTY_TOOL_SPEND: ToolSpendResponse = { const isoDay = (d: Date): string => d.toISOString().slice(0, 10); const UsageTab: React.FC = ({ accessToken, activity }) => { - const { dateValue, onDateChange, results, loading, isFetchingMore } = activity; + const { dateValue, onDateChange, results, loading } = activity; const startTime = dateValue.from ?? null; const endTime = dateValue.to ?? null; @@ -130,7 +130,7 @@ const UsageTab: React.FC = ({ accessToken, activity }) => { - +
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/costOptimizationUtils.ts b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/costOptimizationUtils.ts index 5f16b1b04fd..b95e7e0d973 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/costOptimizationUtils.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/costOptimizationUtils.ts @@ -1,3 +1,4 @@ +import type { KeySpendActivityRow } from "@/components/UsagePage/dailyActivityApi"; import { DailyData, SpendMetrics } from "@/components/UsagePage/types"; import { ToolSpendDailyEntry, ToolSpendEntry } from "@/components/networking"; import { formatNumberWithCommas } from "@/utils/dataUtils"; @@ -101,6 +102,71 @@ const aggregateByModel = (results: readonly DailyData[]): Map { + const totals = [...aggregateByModel(results).values()].reduce( + (agg, a) => ({ + cachedTokens: agg.cachedTokens + a.cacheReadTokens + a.cacheCreationTokens, + realizedCachingSavings: agg.realizedCachingSavings + a.realizedCachingSavings, + }), + { cachedTokens: 0, realizedCachingSavings: 0 }, + ); + const rate = totals.cachedTokens > 0 ? totals.realizedCachingSavings / totals.cachedTokens : null; + return rate != null && rate > 0 ? rate : null; +}; + +const toLeakageRow = ( + id: string, + a: { + alias: string | null; + teamId: string | null; + promptTokens: number; + cacheReadTokens: number; + cacheCreationTokens: number; + }, + rate: number | null, + dimension: CacheLeakageDimension, +): CacheLeakageRow => { + const uncachedPromptTokens = Math.max(0, a.promptTokens - a.cacheReadTokens - a.cacheCreationTokens); + return { + id, + label: dimension === "model" ? id : a.alias ?? `${id.slice(0, 8)}...`, + sublabel: dimension === "model" ? null : a.teamId, + uncachedPromptTokens, + cacheHitRatio: a.promptTokens > 0 ? a.cacheReadTokens / a.promptTokens : 0, + potentialSavings: rate != null ? uncachedPromptTokens * rate : null, + }; +}; + +const sortAndLimit = (rows: CacheLeakageRow[], rate: number | null, limit: number): CacheLeakageRow[] => + rows + .filter((row) => row.uncachedPromptTokens > 0) + .sort((x, y) => + rate != null + ? (y.potentialSavings ?? 0) - (x.potentialSavings ?? 0) + : y.uncachedPromptTokens - x.uncachedPromptTokens, + ) + .slice(0, limit); + +export const leakageRowsFromKeyRows = ( + rows: readonly KeySpendActivityRow[], + rate: number | null, + limit = 10, +): CacheLeakageRow[] => + sortAndLimit( + rows.map((row) => { + const metrics = { + alias: row.metadata.key_alias ?? null, + teamId: row.metadata.team_id ?? null, + promptTokens: row.metrics.prompt_tokens ?? 0, + cacheReadTokens: row.metrics.cache_read_input_tokens ?? 0, + cacheCreationTokens: row.metrics.cache_creation_input_tokens ?? 0, + }; + return toLeakageRow(row.api_key, metrics, rate, "key"); + }), + rate, + limit, + ); + export const computeCacheLeakage = ( results: readonly DailyData[], dimension: CacheLeakageDimension = "key", @@ -123,27 +189,13 @@ export const computeCacheLeakage = ( // A non-positive rate prices no leakage: there is no saving to extrapolate from const rate = netSavingsPerCachedToken != null && netSavingsPerCachedToken > 0 ? netSavingsPerCachedToken : null; - const rows: CacheLeakageRow[] = [...byEntity.entries()] - .map(([id, a]) => { - const uncachedPromptTokens = Math.max(0, a.promptTokens - a.cacheReadTokens - a.cacheCreationTokens); - return { - id, - label: dimension === "model" ? id : a.alias ?? `${id.slice(0, 8)}...`, - sublabel: dimension === "model" ? null : a.teamId, - uncachedPromptTokens, - cacheHitRatio: a.promptTokens > 0 ? a.cacheReadTokens / a.promptTokens : 0, - potentialSavings: rate != null ? uncachedPromptTokens * rate : null, - }; - }) - .filter((row) => row.uncachedPromptTokens > 0); - - const sorted = rows.sort((x, y) => - rate != null - ? (y.potentialSavings ?? 0) - (x.potentialSavings ?? 0) - : y.uncachedPromptTokens - x.uncachedPromptTokens, + const rows = sortAndLimit( + [...byEntity.entries()].map(([id, a]) => toLeakageRow(id, a, rate, dimension)), + rate, + limit, ); - return { rows: sorted.slice(0, limit), netSavingsPerCachedToken }; + return { rows, netSavingsPerCachedToken }; }; export interface DailyToolSpendPoint { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useCacheLeakageKeys.ts b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useCacheLeakageKeys.ts new file mode 100644 index 00000000000..4febd774eb8 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useCacheLeakageKeys.ts @@ -0,0 +1,66 @@ +import { useEffect, useRef, useState } from "react"; + +import { cacheLeakageKeysCall } from "@/components/networking"; +import type { KeySpendActivityRow } from "@/components/UsagePage/dailyActivityApi"; +import type { DailyActivityRange } from "./useDailyActivityRange"; + +interface CacheLeakageKeysResult { + rows: KeySpendActivityRow[]; + loading: boolean; + failed: boolean; +} + +interface SettledKeys { + key: string; + rows: KeySpendActivityRow[]; + failed: boolean; +} + +export const useCacheLeakageKeys = (range: DailyActivityRange, enabled: boolean): CacheLeakageKeysResult => { + const { accessToken, startTime, endTime, userId, apiKey } = range.scope; + const [settled, setSettled] = useState(null); + const requestIdRef = useRef(0); + + const hasTimeRange = !!startTime && !!endTime; + const scopeReady = enabled && !!accessToken && hasTimeRange; + const scopeKey = scopeReady ? JSON.stringify([accessToken, startTime, endTime, userId, apiKey]) : null; + + useEffect(() => { + if (!scopeKey) return; + if (!accessToken || !startTime || !endTime) return; + + const requestId = ++requestIdRef.current; + const isStale = () => requestIdRef.current !== requestId; + + const request = { + accessToken, + startTime, + endTime, + entityIds: userId ? [userId] : null, + apiKey, + includeCurrentUtcDay: true, + }; + cacheLeakageKeysCall(request) + .then((response) => { + if (isStale()) return; + setSettled({ key: scopeKey, rows: response.api_keys, failed: false }); + }) + .catch((error) => { + if (isStale()) return; + console.error("Failed to fetch cache leakage keys:", error); + setSettled({ key: scopeKey, rows: [], failed: true }); + }); + + return () => { + requestIdRef.current++; + }; + // eslint-disable-next-line react-hooks/exhaustive-deps -- scopeKey serializes the scope + }, [scopeKey]); + + const current = scopeKey !== null && settled?.key === scopeKey ? settled : null; + return { + rows: current?.rows ?? [], + loading: scopeKey !== null && current === null, + failed: current?.failed ?? false, + }; +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.test.tsx index b94fa45ecb7..90d6694ccba 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.test.tsx @@ -1,35 +1,35 @@ import { renderHook } from "@testing-library/react"; import { describe, expect, it, vi } from "vitest"; -const mockUsePaginatedDailyActivity = vi.fn(); +import { EMPTY_DAILY_ACTIVITY_METADATA } from "@/components/UsagePage/dailyActivityApi"; -const mockCancel = vi.fn(); +const mockUseAggregatedDailyActivity = vi.fn(); -vi.mock("@/app/(dashboard)/usage/_components/hooks/usePaginatedDailyActivity", () => ({ - usePaginatedDailyActivity: (args: unknown) => { - mockUsePaginatedDailyActivity(args); +vi.mock("@/app/(dashboard)/usage/_components/hooks/useAggregatedDailyActivity", () => ({ + useAggregatedDailyActivity: (options: unknown) => { + mockUseAggregatedDailyActivity(options); return { - data: { results: [] }, + data: { results: [], metadata: EMPTY_DAILY_ACTIVITY_METADATA }, loading: false, - isFetchingMore: false, - progress: { currentPage: 4, totalPages: 9 }, - cancelled: false, failed: false, - coversRange: true, - cancel: mockCancel, }; }, })); vi.mock("@/components/networking", () => ({ - userDailyActivityCall: vi.fn(), - userDailyActivityAggregatedCall: vi.fn(), + dailyActivityAggregatedCall: vi.fn().mockResolvedValue({ results: [], metadata: {} }), })); -import { userDailyActivityAggregatedCall } from "@/components/networking"; +import { dailyActivityAggregatedCall } from "@/components/networking"; import { useActivityDateRange, useDailyActivityRange } from "./useDailyActivityRange"; -const argsOfLastCall = () => mockUsePaginatedDailyActivity.mock.calls.at(-1)?.[0].args as unknown[]; +interface CapturedOptions { + fetch: () => Promise; + enabled: boolean; + deps: unknown[]; +} + +const lastOptions = () => mockUseAggregatedDailyActivity.mock.calls.at(-1)?.[0] as CapturedOptions; describe("useDailyActivityRange", () => { it("offers date-range state without starting a daily-activity query", () => { @@ -37,49 +37,53 @@ describe("useDailyActivityRange", () => { expect(result.current.dateValue.from).toBeInstanceOf(Date); expect(result.current.dateValue.to).toBeInstanceOf(Date); - expect(mockUsePaginatedDailyActivity).not.toHaveBeenCalled(); + expect(mockUseAggregatedDailyActivity).not.toHaveBeenCalled(); }); - it("queries every user's activity for an admin", () => { + it("fetches every user's activity for an admin through the aggregated endpoint", async () => { renderHook(() => useDailyActivityRange("test-token", "u1", "proxy_admin")); - expect(argsOfLastCall()).toEqual(["test-token", expect.any(Date), expect.any(Date), null, true, null]); + await lastOptions().fetch(); + expect(dailyActivityAggregatedCall).toHaveBeenCalledWith( + "user", + expect.objectContaining({ + accessToken: "test-token", + entityIds: null, + includeCurrentUtcDay: true, + }), + ); }); - it("scopes the query to the caller for a non-admin", () => { + it("scopes the query to the caller for a non-admin", async () => { renderHook(() => useDailyActivityRange("test-token", "u1", "internal_user")); - expect(argsOfLastCall()).toEqual(["test-token", expect.any(Date), expect.any(Date), "u1", true, null]); + await lastOptions().fetch(); + expect(dailyActivityAggregatedCall).toHaveBeenCalledWith("user", expect.objectContaining({ entityIds: ["u1"] })); }); it.each(["org_admin", "Org Admin"])( "scopes the query to the caller for %s, who has no admin view on this endpoint", - (role) => { + async (role) => { renderHook(() => useDailyActivityRange("test-token", "u1", role)); - expect(argsOfLastCall()).toEqual(["test-token", expect.any(Date), expect.any(Date), "u1", true, null]); + await lastOptions().fetch(); + expect(dailyActivityAggregatedCall).toHaveBeenCalledWith("user", expect.objectContaining({ entityIds: ["u1"] })); }, ); - it("fetches through the single-shot aggregated endpoint first so days never fragment across pages", () => { - renderHook(() => useDailyActivityRange("test-token", "u1", "proxy_admin")); - - expect(mockUsePaginatedDailyActivity).toHaveBeenLastCalledWith( - expect.objectContaining({ aggregatedFetchFn: userDailyActivityAggregatedCall }), - ); - }); - - it("forwards the pagination progress and cancel affordances instead of dropping them", () => { - const { result } = renderHook(() => useDailyActivityRange("test-token", "u1", "proxy_admin")); - - expect(result.current.progress).toEqual({ currentPage: 4, totalPages: 9 }); - expect(result.current.cancelled).toBe(false); - expect(result.current.cancel).toBe(mockCancel); - }); - it("stays disabled until an access token is available", () => { renderHook(() => useDailyActivityRange(null, "u1", "proxy_admin")); - expect(mockUsePaginatedDailyActivity).toHaveBeenLastCalledWith(expect.objectContaining({ enabled: false })); + expect(lastOptions().enabled).toBe(false); + }); + + it("exposes the request scope so sibling hooks fetch under the same filters", () => { + const { result } = renderHook(() => useDailyActivityRange("test-token", "u1", "internal_user")); + + expect(result.current.scope).toMatchObject({ + accessToken: "test-token", + userId: "u1", + apiKey: null, + }); }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.ts b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.ts index 605926132e9..1af009397b1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.ts @@ -1,9 +1,15 @@ import { useMemo, useState } from "react"; -import { userDailyActivityAggregatedCall, userDailyActivityCall } from "@/components/networking"; +import { dailyActivityAggregatedCall } from "@/components/networking"; +import { + EMPTY_DAILY_ACTIVITY_METADATA, + toDailyData, + type DailyActivityMetadata, + type DailyActivityRequest, +} from "@/components/UsagePage/dailyActivityApi"; import { DailyData } from "@/components/UsagePage/types"; import { spendScopeUserId } from "@/utils/roles"; -import { usePaginatedDailyActivity } from "@/app/(dashboard)/usage/_components/hooks/usePaginatedDailyActivity"; +import { useAggregatedDailyActivity } from "@/app/(dashboard)/usage/_components/hooks/useAggregatedDailyActivity"; const THIRTY_DAYS_MS = 30 * 24 * 60 * 60 * 1000; @@ -12,30 +18,22 @@ export interface DateRange { to?: Date; } +export interface DailyActivityScope { + accessToken: string | null; + startTime: Date | null; + endTime: Date | null; + userId: string | null; + apiKey: string | null; +} + export interface DailyActivityRange { dateValue: DateRange; onDateChange: (value: DateRange) => void; results: DailyData[]; + metadata: DailyActivityMetadata; loading: boolean; - isFetchingMore: boolean; - progress: { currentPage: number; totalPages: number }; - cancelled: boolean; failed: boolean; - cancel: () => void; -} - -/** - * Which slice of daily activity to read. Both fields are passed straight through to the - * endpoint as filters, so the caller — not this hook — decides what the viewer may see. - * - * `userId: null` asks for the whole proxy, which the backend only honours for admins; - * a non-admin must send its own id or the request is rejected. That role decision lives in - * `useDailyActivityRange` below rather than in here, so a caller scoping to one key is not - * silently re-scoped to a user as well. - */ -export interface DailyActivityScope { - userId: string | null; - apiKey?: string | null; + scope: DailyActivityScope; } export type ActivityDateRange = Pick; @@ -47,39 +45,49 @@ export const useActivityDateRange = (): ActivityDateRange => { return { dateValue, onDateChange: setDateValue }; }; +export interface ScopedActivityInput { + userId: string | null; + apiKey?: string | null; +} + export const useScopedDailyActivityRange = ( accessToken: string | null, - scope: DailyActivityScope, + scope: ScopedActivityInput, { dateValue, onDateChange }: ActivityDateRange, ): DailyActivityRange => { const startTime = dateValue.from ?? null; const endTime = dateValue.to ?? null; const { userId, apiKey = null } = scope; - const activityQueryOptions = { - fetchFn: userDailyActivityCall, - aggregatedFetchFn: userDailyActivityAggregatedCall, - // Positional, and read by two functions whose signatures diverge at index 3: the paginated - // call takes `page` there (injected by the hook) and the aggregated one does not. Anything - // appended here must therefore be appended to BOTH networking signatures, in this order. - args: [accessToken, startTime, endTime, userId, true, apiKey], - enabled: !!accessToken && !!startTime && !!endTime, - }; - const { data, loading, isFetchingMore, progress, cancelled, failed, coversRange, cancel } = - usePaginatedDailyActivity(activityQueryOptions); - const readUnavailable = failed || cancelled; - const waitingForRange = activityQueryOptions.enabled && !coversRange && !readUnavailable; + const request = useMemo( + () => + accessToken && startTime && endTime + ? { + accessToken, + startTime, + endTime, + entityIds: userId ? [userId] : null, + apiKey, + includeCurrentUtcDay: true, + } + : null, + [accessToken, startTime, endTime, userId, apiKey], + ); + + const { data, loading, failed } = useAggregatedDailyActivity({ + fetch: () => dailyActivityAggregatedCall("user", request as DailyActivityRequest), + enabled: request !== null, + deps: [accessToken, startTime, endTime, userId, apiKey], + }); return { dateValue, onDateChange, - results: data.results as DailyData[], - loading: loading || waitingForRange, - isFetchingMore, - progress, - cancelled, + results: toDailyData(data), + metadata: data.metadata ?? EMPTY_DAILY_ACTIVITY_METADATA, + loading, failed, - cancel, + scope: { accessToken, startTime, endTime, userId, apiKey }, }; }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/ActivityScope.tsx b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/ActivityScope.tsx index 188c1e6db92..91a74abe8e1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/ActivityScope.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/ActivityScope.tsx @@ -1,10 +1,18 @@ "use client"; -import { useEffect, useId, useState } from "react"; +import { useEffect, useId, useState, type ReactNode } from "react"; import { useQuery } from "@tanstack/react-query"; -import { Plus, X, ArrowUpRight } from "lucide-react"; +import { Plus, X, ChevronRight, RotateCw } from "lucide-react"; import { apiClient } from "@/components/networking"; import { Button } from "@/components/ui/button"; +import { + Combobox, + ComboboxInput, + ComboboxContent, + ComboboxList, + ComboboxItem, + ComboboxEmpty, +} from "@/components/ui/combobox"; import { Input } from "@/components/ui/input"; import { TracePanel } from "./TracePanel"; import { type Sample, type Settings, runTime, durationLabel } from "./lensData"; @@ -15,7 +23,14 @@ export type ActivitySelection = Pick & Partial< Pick< Settings, - "service" | "filters" | "lookback_hours" | "sample_percent" | "sample_size" | "team_id" | "execution_ids" + | "service" + | "agent_name" + | "filters" + | "lookback_hours" + | "sample_percent" + | "sample_size" + | "team_id" + | "execution_ids" > >; @@ -28,10 +43,8 @@ export function RunList({ executions }: { executions: Sample["executions"] }) {

{run.name}

- {runTime(run.start_time)} · {run.source === "traces" ? `${run.span_count} steps` : "LLM request"} -

-

- {run.trace_id} + {runTime(run.start_time)} ·{" "} + {run.source === "traces" ? `${run.span_count} ${run.span_count === 1 ? "step" : "steps"}` : "LLM request"}

))} @@ -43,12 +56,22 @@ export function ActivityScope({ value, onChange, accessToken, + mode = "scope", + onPreviewReady, + manualSelection = false, + nameField, }: { value: ActivitySelection; onChange: (selection: ActivitySelection) => void; accessToken: string; + mode?: "scope" | "activity"; + onPreviewReady?: (ready: boolean) => void; + manualSelection?: boolean; + nameField?: ReactNode; }) { const id = useId(); + const hasFilters = !!value.filters?.length || !!value.team_id; + const [advanced, setAdvanced] = useState(hasFilters || !!value.service || value.source !== "traces"); const [offset, setOffset] = useState(0); const [scope, setScope] = useState(value); const [trace, setTrace] = useState<{ id: string; ref?: string } | null>(null); @@ -63,7 +86,7 @@ export function ActivityScope({ return () => clearTimeout(timer); }, [serialized]); const historyHours = value.lookback_hours ?? 24; - const validWindow = Number.isInteger(historyHours) && historyHours >= 1 && historyHours <= 720; + const validWindow = Number.isInteger(historyHours) && historyHours >= 1 && historyHours <= 8760; const percent = scope.sample_percent ?? 100; const cap = scope.sample_size; const validCap = cap == null || (Number.isInteger(cap) && cap > 0); @@ -96,12 +119,19 @@ export function ActivityScope({ lookback_hours: value.lookback_hours, }; const discoveryOptions = { - queryKey: ["lens-activity-options", value.source, value.lookback_hours, accessToken], + queryKey: ["lens-activity-options", value.source, value.lookback_hours, asOf, accessToken], queryFn: () => load(discoveryScope), staleTime: 60000, enabled: validWindow, }; const discovery = useQuery(discoveryOptions); + const agentOptions = { + queryKey: ["lens-agents", accessToken, asOf], + queryFn: () => apiClient.get("/lens/agents", { accessToken }), + enabled: value.source !== "requests", + staleTime: 60000, + }; + const agents = useQuery(agentOptions); const previewOptions = { queryKey: ["lens-activity-preview", scope, offset, asOf, accessToken], queryFn: () => load(scope, offset), @@ -109,18 +139,38 @@ export function ActivityScope({ staleTime: 30000, }; const preview = useQuery(previewOptions); + const empty = preview.data?.eligible === 0; + useEffect(() => { + if (!empty || !valid) return; + const timer = window.setTimeout(() => setAsOf(new Date().toISOString()), 15000); + return () => window.clearTimeout(timer); + }, [empty, valid, asOf]); + const refreshPreview = () => { + setOffset(0); + setAsOf(new Date().toISOString()); + }; const runs = discovery.data?.executions ?? []; const services = [...new Set(runs.map((r) => r.service).filter(Boolean))].sort(); + const selectedName = value.source === "requests" ? value.service : value.agent_name; + const names = value.source === "requests" ? services : agents.data ?? []; + const selectName = (name: string) => + onChange({ ...value, [value.source === "requests" ? "service" : "agent_name"]: name, execution_ids: [] }); const attributes = runs.flatMap((r) => r.metadata ?? []); const keys = [...new Set(attributes.map((a) => a.key).filter((key) => !key.startsWith("litellm.")))].sort(); const pending = serialized !== JSON.stringify(scope) || preview.isFetching; const ready = !pending && valid; + const hasSelection = !manualSelection || !!value.execution_ids?.length; + const hasMatches = !preview.error && (preview.data?.selected ?? 0) > 0; + const canReview = ready && hasMatches && hasSelection; + useEffect(() => { + onPreviewReady?.(canReview); + }, [canReview, onPreviewReady]); const filters = value.filters ?? []; const edit = (index: number, field: "key" | "value", text: string) => onChange({ ...value, filters: filters.map((f, i) => (i === index ? { ...f, [field]: text } : f)) }); const changeSource = (source: Settings["source"]) => { - const selection = { ...value, source, service: "", filters: [], execution_ids: [] }; + const selection = { ...value, source, service: "", agent_name: "", filters: [], execution_ids: [] }; onChange(selection); }; const windowLabel = validWindow @@ -128,193 +178,201 @@ export function ActivityScope({ : "Choose a valid history window"; const previewTitle = () => { if (pending) return "Finding matching activity…"; - if (!validWindow) return "Choose a history window between 1 and 720 hours"; + if (!validWindow) return "Choose a history window between 1 hour and 365 days"; if (!valid) return "Complete your condition to preview matches"; if (!preview.data) return "Preview unavailable"; - return `${preview.data.eligible} matching ${value.source === "requests" ? "requests" : "runs"}`; + const noun = value.source === "requests" ? "request" : "run"; + return `${preview.data.eligible} matching ${noun}${preview.data.eligible === 1 ? "" : "s"}`; }; return ( -
-
- -

- {value.source === "requests" - ? "Each request is one model call, not an entire agent run." - : "An agent run contains the steps recorded under one trace ID. Separate sessions are not joined automatically."} -

- -

- { - { - requests: "The model alias configured on your LiteLLM gateway. Leave blank for all models.", - both: "Matches the application name on agent runs or the model group on requests. Leave blank to include both without a name filter.", - traces: - "The service.name recorded by your agent’s OpenTelemetry instrumentation. Leave blank for all applications.", - }[value.source ?? "traces"] - } -

-
-

- Narrow by metadata (optional) -

-

- Match a recorded tag, swarm, or environment. Every condition must match exactly. -

- {filters.map((f, index) => ( -
- edit(index, "key", e.target.value)} - /> - is - edit(index, "value", e.target.value)} - /> - - {[...new Set(attributes.filter((a) => a.key === f.key).map((a) => a.value))].sort().map((v) => ( - - -
- ))} - - {keys.map((key) => ( - - -

- Suggestions come from up to 100 recent runs. You can also type a recorded key or value. -

-
- - onChange({ ...value, lookback_hours })} - /> -

- Time window used by each scan. Activity becomes eligible two minutes after it finishes. -

-
- + {value.source !== "requests" && agents.isError && ( +

+ Could not load agents.{" "} + +

+ )} +
setAdvanced(event.currentTarget.open)} className="group"> + + Advanced filters{filters.length ? ` (${filters.length})` : ""} + +
+ {value.source !== "requests" && ( + + )} + +

+ Match any recorded metadata, such as a user ID, environment, or tag. All conditions must match. +

+ {filters.map((f, index) => ( +
+
+ edit(index, "key", e.target.value)} + /> + +
+ edit(index, "value", e.target.value)} + /> + + {[...new Set(attributes.filter((a) => a.key === f.key).map((a) => a.value))].sort().map((v) => ( + +
+ ))} + + {keys.map((key) => ( + + + +
+
+ + ) : ( + <> + onChange({ ...value, lookback_hours })} /> - - -
-

100% with no limit selects all matching activity.

- {!!value.execution_ids?.length && ( - )}
- - onChange({ - ...value, - execution_ids: checked - ? [...(value.execution_ids ?? []), runId] - : (value.execution_ids ?? []).filter((id) => id !== runId), - }) - } - selectedIds={value.execution_ids ?? []} - selectedCount={ - value.execution_ids?.length - ? Math.min( - Math.ceil((value.execution_ids.length * (value.sample_percent ?? 100)) / 100), - value.sample_size ?? Infinity, - ) - : preview.data?.selected ?? 0 - } - title={previewTitle()} - windowLabel={windowLabel} - ready={ready} - error={preview.error} - data={preview.data} - onOpen={(run) => setTrace({ id: run.trace_id, ref: run.trace_ref })} - /> + {mode === "activity" && ( + + onChange({ + ...value, + execution_ids: checked + ? [...(value.execution_ids ?? []), runId] + : (value.execution_ids ?? []).filter((id) => id !== runId), + }) + } + manualSelection={manualSelection} + selectedIds={value.execution_ids ?? []} + selectedCount={ + manualSelection + ? Math.min( + Math.ceil(((value.execution_ids?.length ?? 0) * (value.sample_percent ?? 100)) / 100), + value.sample_size ?? Infinity, + ) + : preview.data?.selected ?? 0 + } + title={previewTitle()} + windowLabel={windowLabel} + ready={ready} + error={preview.error} + data={preview.data} + onRetry={refreshPreview} + onOpen={(run) => setTrace({ id: run.trace_id, ref: run.trace_ref })} + /> + )} {trace && ( void; onSelect: (id: string, checked: boolean) => void; selectedIds: string[]; + manualSelection: boolean; selectedCount: number; title: string; windowLabel: string; ready: boolean; error: Error | null; data: Sample | undefined; + onRetry: () => void; onOpen: (run: Sample["executions"][number]) => void; }) { + const paginated = data?.next_offset != null || offset > 0; + const showSelection = selectedCount !== data?.eligible || paginated; + const selectionData = ready && showSelection ? data : undefined; return (
-

- {title} -

-

{windowLabel} · Preview only, no analysis cost

+
+

+ {title} +

+ +
+

{windowLabel} · No analysis cost

-
+
{ready && error && (

- {error.message} + {error.message}{" "} +

)} {ready && data?.eligible === 0 && (

- No matches. Try removing a condition or check that your agent records this metadata. Very recent runs need - two minutes to settle. + No matches. Try removing a condition or check that your agent records this metadata. Recent trace updates + need two minutes to settle.

)} {ready && data?.executions.map((run) => (
- onSelect(run.id, e.target.checked)} - /> + {manualSelection && ( + onSelect(run.id, e.target.checked)} + /> + )}
{run.source === "traces" && ( - )}
))}
- {ready && data && ( + {selectionData && (

- {selectedCount} selected for analysis · Showing {offset + (data.executions.length ? 1 : 0)}– - {offset + data.executions.length} of {data.eligible} + {selectedCount} selected for analysis + {paginated && ( + <> + {" "} + · Showing {offset + (selectionData.executions.length ? 1 : 0)}– + {offset + selectionData.executions.length} of {selectionData.eligible} + + )}

-
- - -
+ {paginated && ( +
+ + +
+ )}
)}
diff --git a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/AnalysisKey.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/AnalysisKey.integration.test.tsx index 7a0130bd929..f3174b2518d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/AnalysisKey.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/AnalysisKey.integration.test.tsx @@ -12,36 +12,27 @@ describe("Lens billing key", () => { testQueryClient.clear(); vi.clearAllMocks(); }); - it("creates a normal key and only passes its ID to worker settings", async () => { - const user = userEvent.setup(); - const changed = vi.fn(); - vi.mocked(apiClient.get).mockResolvedValue({ keys: [], total_pages: 0 }); - vi.mocked(apiClient.post).mockResolvedValue({ token_id: "b".repeat(64), key: "sk-secret-not-for-settings" }); - renderWithProviders(); - await user.click(screen.getByRole("button", { name: "Create worker key" })); - expect(await screen.findByRole("combobox", { name: "Charge analysis to" })).toHaveValue("Lens: Research"); - expect(apiClient.post).toHaveBeenCalledWith("/key/generate", { - accessToken: "test", - body: { key_alias: "Lens: Research", models: [], metadata: { purpose: "lens" } }, - }); - expect(changed).toHaveBeenCalledExactlyOnceWith("b".repeat(64)); - expect(screen.queryByText("sk-secret-not-for-settings")).not.toBeInTheDocument(); - }); it("pages existing keys without dropping the selected billing key", async () => { const user = userEvent.setup(); const changed = vi.fn(); - vi.mocked(apiClient.get).mockImplementation(async (_path, options) => ({ - keys: - options?.query?.page === "2" - ? [{ token: "c".repeat(64), key_alias: "Second page" }] - : [{ token: "a".repeat(64), key_alias: "First page" }], - total_pages: 2, - })); - renderWithProviders(); + vi.mocked(apiClient.get).mockImplementation(async (path, options) => + path === "/key/info" + ? { info: { models: ["restricted-model"], max_budget: 4, budget_duration: "1d" } } + : { + keys: + options?.query?.page === "2" + ? [{ token: "c".repeat(64), key_alias: "Second page" }] + : [{ token: "a".repeat(64), key_alias: "First page" }], + total_pages: 2, + }, + ); + renderWithProviders(); await user.click(screen.getByRole("combobox", { name: "Charge analysis to" })); await user.click(await screen.findByRole("option", { name: "Load more keys" })); await user.click(await screen.findByRole("option", { name: "Second page" })); expect(changed).toHaveBeenCalledExactlyOnceWith("c".repeat(64)); + expect(await screen.findByText("restricted-model")).toBeInTheDocument(); + expect(screen.getByText("$4.00 / day")).toBeInTheDocument(); }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/AnalysisKey.tsx b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/AnalysisKey.tsx index c26c42f5700..4033b9634b2 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/AnalysisKey.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/AnalysisKey.tsx @@ -1,10 +1,12 @@ "use client"; import { useState } from "react"; -import { useInfiniteQuery } from "@tanstack/react-query"; +import { useInfiniteQuery, useQuery } from "@tanstack/react-query"; import { z } from "zod"; import { apiClient } from "@/components/networking"; -import { Button } from "@/components/ui/button"; +import { SearchSelect } from "@/components/shared/SearchSelect"; +import { Input } from "@/components/ui/input"; +import { AnalysisKeyDetails } from "./AnalysisKeyDetails"; import { Combobox, ComboboxContent, @@ -22,17 +24,14 @@ export function AnalysisKey({ accessToken, value, onChange, - name, }: { accessToken: string; value: string | null; onChange: (key: string | null) => void; - name: string; }) { const [query, setQuery] = useState(""); const [selected, setSelected] = useState(value ? { token: value } : null); - const [creating, setCreating] = useState(false); - const [error, setError] = useState(""); + const queryOptions = { queryKey: ["lens-analysis-keys", accessToken, query], initialPageParam: 1, @@ -61,28 +60,6 @@ export function AnalysisKey({ const choice = keys.find((key) => key.token === value) ?? selected; const loading = keyPages.isFetching; - const create = async () => { - setCreating(true); - setError(""); - try { - const result = await apiClient.post("/key/generate", { - accessToken, - body: { - key_alias: `Lens: ${name}`, - models: [], - metadata: { purpose: "lens" }, - }, - }); - if (!result.token_id) throw new Error("The proxy did not return the new key's ID"); - const key = { token: result.token_id, key_alias: `Lens: ${name}` }; - setSelected(key); - onChange(key.token); - } catch (cause) { - setError(cause instanceof Error ? cause.message : "Could not create a key"); - } finally { - setCreating(false); - } - }; const changeKey = (key: Key | null, details: { cancel: () => void }) => { if (key?.token === "load-more") { details.cancel(); @@ -99,7 +76,7 @@ export function AnalysisKey({ return (

Charge analysis to

-
+
-
-

- Spend appears under this key in API Keys. Its permissions and limits apply. -

- {(error || keyPages.error) && ( + {choice && } + {keyPages.error && (

- {error || keyPages.error?.message} + {keyPages.error.message}

)}
); } + +export type AnalysisAccess = { model: string | null; budget: string }; + +export function AnalysisAccessFields({ + accessToken, + value, + onChange, +}: { + accessToken: string; + value: AnalysisAccess; + onChange: (value: AnalysisAccess) => void; +}) { + const models = useQuery({ + queryKey: ["lens-models", accessToken], + queryFn: () => apiClient.get<{ data: { id: string }[] }>("/models", { accessToken }), + }); + return ( +
+
+ + ({ label: id, value: id }))} + value={value.model} + onValueChange={(model) => onChange({ ...value, model })} + placeholder={models.isLoading ? "Loading models…" : "Select a model"} + /> +
+
+ + onChange({ ...value, budget: e.target.value })} + /> +

Shared across all investigations.

+
+ {models.error && ( +

+ {models.error.message} +

+ )} +
+ ); +} + +export async function createAnalysisKey(accessToken: string, access: AnalysisAccess): Promise { + if (!access.model || !Number.isFinite(Number(access.budget)) || Number(access.budget) <= 0) + throw new Error("Choose a model and a monthly limit greater than zero"); + const result = await apiClient.post("/key/generate", { + accessToken, + body: { + key_alias: "Lens analysis", + models: [access.model], + max_budget: Number(access.budget), + budget_duration: "1mo", + metadata: { purpose: "lens" }, + }, + }); + if (!result.token_id) throw new Error("The proxy did not return the new key's ID"); + return result.token_id; +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/AnalysisKeyDetails.tsx b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/AnalysisKeyDetails.tsx new file mode 100644 index 00000000000..e94099ac6e6 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/AnalysisKeyDetails.tsx @@ -0,0 +1,99 @@ +"use client"; + +import { useQuery } from "@tanstack/react-query"; +import { z } from "zod"; +import { apiClient } from "@/components/networking"; +import { Button } from "@/components/ui/button"; +import { runTime } from "./lensData"; + +const keyInfoFields = { + key_alias: z.string().nullable().optional(), + models: z.array(z.string()), + max_budget: z.number().nullable(), + budget_duration: z.string().nullable().optional(), + rpm_limit: z.number().nullable().optional(), + tpm_limit: z.number().nullable().optional(), + expires: z.string().nullable().optional(), + status: z.string().optional(), +}; +const keyInfoSchema = z.object({ info: z.object(keyInfoFields) }); + +function budgetLabel(amount: number | null, duration?: string | null): string { + if (amount === null) return "No key budget"; + const periods: Record = { + "1mo": "month", + "30d": "month", + "1d": "day", + "24h": "day", + "7d": "week", + "1h": "hour", + }; + const dollars = new Intl.NumberFormat("en-US", { + style: "currency", + currency: "USD", + maximumFractionDigits: 2, + }).format(amount); + return duration ? `${dollars} / ${periods[duration] ?? duration}` : `${dollars} total`; +} + +export function useAnalysisKeyInfo(accessToken: string, keyId?: string) { + return useQuery({ + queryKey: ["lens-key-info", accessToken, keyId], + enabled: !!keyId, + queryFn: async () => + keyInfoSchema.parse(await apiClient.get("/key/info", { accessToken, query: { key: keyId } })).info, + }); +} + +export function AnalysisKeyDetails({ + accessToken, + keyId, + showName = false, +}: { + accessToken: string; + keyId: string; + showName?: boolean; +}) { + const key = useAnalysisKeyInfo(accessToken, keyId); + if (key.isLoading) return

Loading key permissions…

; + if (key.error || !key.data) + return ( +
+ Could not load key permissions + +
+ ); + const info = key.data; + return ( +
+
+ {showName && ( + <> +
Billing key
+
{info.key_alias || "Assigned virtual key"}
+ + )} +
Models
+
{info.models.length ? info.models.join(", ") : "All models"}
+
Key limit
+
{budgetLabel(info.max_budget, info.budget_duration)}
+
+ {info.status && info.status !== "active" && ( +

+ This key is {info.status}. Choose an active key. +

+ )} +
+ Other limits +
+

Requests per minute: {info.rpm_limit ?? "No key limit"}

+

Tokens per minute: {info.tpm_limit ?? "No key limit"}

+

Expires: {info.expires ? runTime(info.expires) : "No expiry"}

+

Team, organization, and model limits still apply.

+
+
+
+ ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/DurationInput.tsx b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/DurationInput.tsx index 7e1227eac8b..ce37ee952cb 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/DurationInput.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/DurationInput.tsx @@ -2,6 +2,7 @@ import { useId, useState } from "react"; import { Input } from "@/components/ui/input"; +import { ChevronDown } from "lucide-react"; export function DurationInput({ label, @@ -47,18 +48,24 @@ export function DurationInput({ value={Number.isFinite(value) ? value / scale : ""} onChange={(event) => onChange(event.target.value === "" ? NaN : Number(event.target.value) * scale)} /> - +
+ +
); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/LensFinding.tsx b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/LensFinding.tsx new file mode 100644 index 00000000000..e1f0a47330b --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/lens/_components/LensFinding.tsx @@ -0,0 +1,146 @@ +import { ArrowUpRight } from "lucide-react"; +import { Button } from "@/components/ui/button"; +import { Textarea } from "@/components/ui/textarea"; +import { Sheet, SheetContent, SheetHeader, SheetTitle, SheetDescription } from "@/components/ui/sheet"; +import { evidenceTarget, runTime, type Finding, type Sample } from "./lensData"; + +export function LensFinding({ + finding, + sampledRuns, + readOnly, + reason, + busy, + onClose, + onReason, + onEvidence, + changeFinding, +}: { + finding?: Finding; + sampledRuns: Sample["executions"]; + readOnly: boolean; + reason: string; + busy: boolean; + onClose: () => void; + onReason: (reason: string) => void; + onEvidence: (evidence: { id: string; span: string }) => void; + changeFinding: (status: Finding["status"]) => Promise; +}) { + const evidenceGroups = finding + ? [...new Set(finding.evidence.map((e) => e.execution_id))].map((id) => ({ + id, + run: sampledRuns.find((r) => r.id === id), + quotes: finding.evidence.filter((e) => e.execution_id === id), + })) + : []; + return ( + { + if (!open) onClose(); + }} + > + + {finding && ( + <> + + {finding.title} + + {finding.kind === "issue" ? `${finding.priority} priority` : "Pattern"} ·{" "} + {finding.occurrences?.length ?? 0} linked {finding.occurrences?.length === 1 ? "run" : "runs"} + + +
+
+

What happened

+

{finding.description}

+
+ {finding.suggestion && ( +
+

What to do next

+

{finding.suggestion}

+
+ )} + {finding.limitation && ( +
+ Evidence limits +

{finding.limitation}

+
+ )} +
+

Evidence by run

+

+ Exact quotes from the recorded activity. Counterexamples are labeled separately from supporting + evidence. +

+
+ {evidenceGroups.map((group) => ( +
+ + {group.run?.name ?? evidenceTarget(group.id)?.id.slice(0, 12) ?? "Recorded run"} + + {group.quotes.length} {group.quotes.length === 1 ? "quote" : "quotes"} + {group.run ? ` · ${runTime(group.run.start_time)}` : ""} + + +
+ {group.quotes.map((e, i) => ( +
+ {e.role === "counterexample" && ( +

Counterexample

+ )} +
+ {e.quote} +
+ +
+ ))} +
+
+ ))} +
+
+ {!readOnly && ( +
+