Merge branch 'BerriAI:main' into patch-1

This commit is contained in:
Deepanshu Pal 2026-10-02 10:43:18 +05:30 • committed by GitHub
commit bd18174611
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
370 changed files with 44889 additions and 10754 deletions

View file

@ -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=$!
}

View file

@ -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 ;;

View file

@ -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

View file

@ -268,6 +268,31 @@ For MCP OAuth, an upstream may advertise dynamic client registration but refuse
</details>
<details>
<summary><b>Agents</b> - Run Claude Code, Codex, OpenCode or Deep Agents on any model (Python SDK)</summary>
### 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)
</details>
### 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` |

View file

@ -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

View file

@ -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}

View file

@ -6,8 +6,6 @@ services:
context: .
dockerfile: docker/Dockerfile.non_root
target: runtime
args:
PROXY_EXTRAS_SOURCE: "local"
depends_on:
- squid
user: "101:101"

View file

@ -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

View file

@ -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==",

View file

@ -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,))

View file

@ -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

View file

@ -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;

View file

@ -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;

View file

@ -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>`."""
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<index>[^"]+)"\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}")

View file

@ -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:

View file

@ -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==",

View file

@ -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",
]

View file

@ -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"] }

View file

@ -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

View file

@ -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 `<digits>_<description>.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 },
}

View file

@ -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<Vec<Entry>, 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::<u64>().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<Vec<Entry>, 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<u64> = 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 { .. })));
}
}

View file

@ -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

View file

@ -0,0 +1,5 @@
# Migrations
`litellm-migrate` exports the `Migration` struct and the `migrate!` macro that embeds a directory of `<digits>_<description>.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

View file

@ -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,
}

View file

@ -0,0 +1 @@
SELECT 10;

View file

@ -0,0 +1 @@
SELECT 1;

View file

@ -0,0 +1 @@
SELECT 2;

View file

@ -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);
}

View file

@ -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",

View file

@ -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<Self> {
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<Bound<'py, PyAny>> {
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<Bound<'py, PyAny>> {
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<Bound<'py, PyAny>> {
litellm_host_python::Pythonized(litellm_traces::NORMALIZED_FIELD_DEFINITIONS).into_pyobject(py)
}

View file

@ -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

View file

@ -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

View file

@ -0,0 +1,3 @@
fn main() {
println!("cargo:rerun-if-changed=migrations");
}

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -1 +0,0 @@
ALTER TABLE {database}.otel_traces ADD COLUMN IF NOT EXISTS EngineReceivedMs UInt64 DEFAULT 0

View file

@ -1 +0,0 @@
ALTER TABLE {database}.spend_logs ADD COLUMN IF NOT EXISTS EngineReceivedMs UInt64 DEFAULT 0

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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<QueryAccessError>),
}

View file

@ -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};

View file

@ -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<String, String>) -> bool {
true
}
fn consumed_attributes(&self, attributes: &BTreeMap<String, String>) -> [&'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<String, String>,
) -> Result<NormalizedSpan, DecodeError> {
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(),
})
}
}

View file

@ -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<ContentBlock>),
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::<Vec<_>>()
.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<String, Value>);
#[derive(Deserialize)]
struct ResponseMetadata {
id: Option<String>,
}
#[derive(Deserialize)]
struct RawMessage {
kwargs: Option<Box<RawMessage>>,
#[serde(rename = "type")]
kind: Option<String>,
role: Option<String>,
content: Option<MessageContent>,
tool_calls: Option<Vec<RawToolCall>>,
name: Option<Value>,
response_metadata: Option<ResponseMetadata>,
}
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<RawMessage>),
Nested(Vec<Vec<RawMessage>>),
}
impl<'de> Deserialize<'de> for MessageBatch {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
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<Value>| {
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<Option<T>, 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<RawMessage>,
}
#[derive(Deserialize)]
struct Generation {
message: Option<GenerationMessage>,
}
#[derive(Default, Deserialize)]
struct Payload {
#[serde(default, deserialize_with = "lenient")]
messages: Option<MessageBatch>,
#[serde(default, deserialize_with = "lenient")]
generations: Option<Vec<Vec<Generation>>>,
}
#[derive(Deserialize)]
struct Command {
update: CommandUpdate,
}
#[derive(Deserialize)]
struct CommandUpdate {
messages: Vec<Value>,
}
#[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<W: ?Sized + io::Write>(
&mut self,
writer: &mut W,
first: bool,
) -> io::Result<()> {
if first {
Ok(())
} else {
writer.write_all(b", ")
}
}
fn begin_object_key<W: ?Sized + io::Write>(
&mut self,
writer: &mut W,
first: bool,
) -> io::Result<()> {
if first {
Ok(())
} else {
writer.write_all(b", ")
}
}
fn begin_object_value<W: ?Sized + io::Write>(&mut self, writer: &mut W) -> io::Result<()> {
writer.write_all(b": ")
}
}
fn encode<T: Serialize>(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::<Vec<_>>(),
)
}
fn span_type(
name: &str,
parent_span_id: &str,
attributes: &BTreeMap<String, String>,
) -> 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::<Value>(raw_completion).unwrap_or(Value::Null);
let raw = completion.get("output").cloned().unwrap_or(completion);
let selected = serde_json::from_value::<Command>(raw.clone())
.ok()
.and_then(|command| command.update.messages.into_iter().last())
.unwrap_or(raw);
let output = serde_json::from_value::<ContentValue>(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<String, String>) -> SpanIo {
let raw_prompt = attr(attributes, "gen_ai.prompt");
let raw_completion = attr(attributes, "gen_ai.completion");
let prompt = serde_json::from_str::<Payload>(raw_prompt).unwrap_or_default();
let completion = serde_json::from_str::<Payload>(raw_completion).unwrap_or_default();
if kind == ObservationType::Llm
&& serde_json::from_str::<Value>(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<String, String>) -> bool {
scope_name == "langsmith" || attributes.contains_key("langsmith.span.kind")
}
fn consumed_attributes(&self, _attributes: &BTreeMap<String, String>) -> [&'static str; 2] {
["gen_ai.prompt", "gen_ai.completion"]
}
fn normalize(
&self,
name: &str,
parent_span_id: &str,
attributes: &BTreeMap<String, String>,
) -> Result<NormalizedSpan, DecodeError> {
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, "[]");
}
}

View file

@ -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<String, String>) -> bool;
fn consumed_attributes(&self, attributes: &BTreeMap<String, String>) -> [&'static str; 2];
fn normalize(
&self,
name: &str,
parent_span_id: &str,
attributes: &BTreeMap<String, String>,
) -> Result<NormalizedSpan, DecodeError>;
}
mod genai;
mod langsmith;
mod openinference;
use genai::GenAiNormalizer;
use langsmith::LangSmithNormalizer;
use openinference::OpenInferenceNormalizer;
fn attr<'a>(attributes: &'a BTreeMap<String, String>, key: &str) -> &'a str {
attributes.get(key).map(String::as_str).unwrap_or_default()
}
fn first<'a>(attributes: &'a BTreeMap<String, String>, left: &str, right: &str) -> &'a str {
let value = attr(attributes, left);
if value.is_empty() {
attr(attributes, right)
} else {
value
}
}
fn tokens(attributes: &BTreeMap<String, String>, key: &str) -> Result<u32, DecodeError> {
let value = attr(attributes, key).trim();
if value.is_empty() {
return Ok(0);
}
match value.parse::<i128>() {
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<String, String>) -> 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<String, String>,
) -> Result<Normalization, DecodeError> {
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());
}
}

View file

@ -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<String, String>) -> bool {
attributes.contains_key("openinference.span.kind")
}
fn consumed_attributes(&self, _attributes: &BTreeMap<String, String>) -> [&'static str; 2] {
["input.value", "output.value"]
}
fn normalize(
&self,
_name: &str,
parent_span_id: &str,
attributes: &BTreeMap<String, String>,
) -> Result<NormalizedSpan, DecodeError> {
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(),
})
}
}

View file

@ -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<DecodedEvent>,
pub normalized: NormalizedSpan,
pub consumed_attributes: [&'static str; 2],
}
pub fn decode_otlp(

View file

@ -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<Vec<DecodedSpan>, DecodeError> {
let mut budget = Budget::new(MAX_DECODED_SPAN_BYTES);
@ -125,10 +125,26 @@ fn decoded_span(
budget: &mut Budget,
) -> Result<DecodedSpan, DecodeError> {
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::<Result<Vec<_>, DecodeError>>()?,
normalized,
consumed_attributes: normalization.consumed_attributes,
})
}

View file

@ -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<T> {
data: Vec<T>,
}
#[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<PathPart>,
types: BTreeSet<&'static str>,
expression: String,
}
#[derive(Deserialize, Serialize)]
struct ColumnSchema {
name: String,
#[serde(rename = "type")]
kind: String,
#[serde(flatten)]
details: BTreeMap<String, Value>,
}
#[derive(Serialize)]
struct TableSchema {
name: &'static str,
columns: Vec<ColumnSchema>,
}
#[derive(Serialize)]
struct MetadataCatalog {
table: &'static str,
column: &'static str,
fields: Vec<MetadataField>,
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<String>,
}
#[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<AttributeField>,
truncated: bool,
discovery_sql: String,
scope: &'static str,
#[serde(skip_serializing_if = "Option::is_none")]
error: Option<String>,
}
pub async fn query_sql(
client: &Client,
connection: &Connection,
sql: &str,
) -> Result<String, Error> {
execute_read(client, connection, sql, &BTreeMap::new()).await
}
async fn rows<T: serde::de::DeserializeOwned>(
client: &Client,
connection: &Connection,
sql: &str,
) -> Result<Vec<T>, Error> {
let body = query_sql(client, connection, sql).await?;
serde_json::from_str::<Rows<T>>(&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::<Vec<_>>()
.join(", ");
format!("JSONExtractRaw(metadata, {arguments})")
}
fn discover(
value: &Value,
path: Vec<PathPart>,
fields: &mut BTreeMap<Vec<PathPart>, 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::<Value>(&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<String, Error> {
let tables = stream::iter(["otel_traces", "agent_traces_by_key", "spend_logs"])
.then(|table| async move {
Ok::<_, Error>(TableSchema {
name: table,
columns: rows::<ColumnSchema>(
client,
connection,
&format!("DESCRIBE TABLE {table}"),
)
.await?,
})
})
.try_collect::<Vec<_>>()
.await?;
let metadata = match rows::<MetadataRow>(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::<AttributeRow>(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::<Vec<_>>()
.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::<Vec<_>>(),
"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)
);
}
}

View file

@ -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<String, Error> {
template.render().map_err(|_| Error::InvalidResponse)
}

View file

@ -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<String, Connection>,
slots: Arc<Semaphore>,
}
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<OwnedSemaphorePermit, QueryAccessError> {
self.slots
.clone()
.try_acquire_owned()
.map_err(|_| QueryAccessError::Busy)
}
pub async fn connection(
&self,
client: &Client,
scope: &QueryScope,
secret: &str,
) -> Result<Connection, QueryAccessError> {
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<Connection, QueryAccessError> {
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<String, QueryAccessError> {
let mut mac = Hmac::<Sha256>::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('\'', "\\'"))
}

View file

@ -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}",

View file

@ -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<Self, Error> {
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"),

View file

@ -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 %}

View file

@ -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<ClickHouseDatabase>,
) -> 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<ClickHouseDatabase>,
) -> 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::<BTreeMap<_, _>>();
let sample: serde_json::Value = serde_json::from_str(
&execute_read(
&database.client,
&connection,
LensQuery::Sample.sql(),
&parameters,
)
.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(),
&parameters,
)
.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<ClickHouseDatabase>,
#[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"}],
"<custom>&{{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, '<custom>&{{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<_>>(),
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<ClickHouseDatabase>,
#[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::<Result<Vec<_>, _>>()?;
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::<Result<Vec<_>, _>>()?;
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(())
}

View file

@ -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"));
}

View file

@ -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<ClickHouse>,
client: Client,
writer: Connection,
readers: QueryReaders,
}
#[fixture]
async fn database() -> Result<Database, Box<dyn std::error::Error>> {
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<Database, Box<dyn std::error::Error>>,
#[case] scope: QueryScope,
#[case] expected: Vec<&str>,
) -> Result<(), Box<dyn std::error::Error>> {
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::<Vec<_>>()
),
"{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<Database, Box<dyn std::error::Error>>,
) -> Result<(), Box<dyn std::error::Error>> {
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<Database, Box<dyn std::error::Error>>,
) -> Result<(), Box<dyn std::error::Error>> {
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<Database, Box<dyn std::error::Error>>,
) -> Result<(), Box<dyn std::error::Error>> {
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::<Result<Vec<_>, _>>()?;
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(())
}

View file

@ -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

View file

@ -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,

View file

@ -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",
}
)

View file

@ -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",
)

View file

@ -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

689
litellm/harness/endpoint.py Normal file
View file

@ -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)

45
litellm/harness/errors.py Normal file
View file

@ -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

View file

@ -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")

View file

@ -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")

View file

@ -0,0 +1,161 @@
"""
Generic handler for CLI harnesses (Claude Code, Codex, OpenCode).
The config (`litellm/llms/<harness>/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 <private_dir>/<dir> 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

View file

@ -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

View file

@ -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

1070
litellm/harness/runtime.py Normal file

File diff suppressed because it is too large Load diff

View file

@ -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",
)

View file

@ -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: ...

View file

@ -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 ("<hex> ./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 <args>` 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 <args>` 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)

View file

@ -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)

View file

@ -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)
]

443
litellm/harness/sync.py Normal file
View file

@ -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)

214
litellm/harness/types.py Normal file
View file

@ -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,
)

View file

@ -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
)

View file

@ -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,

View file

@ -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

View file

@ -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>/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."""

View file

@ -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))

View file

@ -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)

View file

@ -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)

View file

View file

@ -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 = "<synthetic>"
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)

View file

View file

View file

@ -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)

View file

View file

@ -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
`<workdir>/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/<name>" or "f/<name>" 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)

View file

@ -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/<group>` 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=")

View file

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