mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge branch 'BerriAI:main' into patch-1
This commit is contained in:
commit
bd18174611
370 changed files with 44889 additions and 10754 deletions
|
|
@ -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=$!
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 ;;
|
||||
|
|
|
|||
8
.github/ci-coverage-allowlist.yml
vendored
8
.github/ci-coverage-allowlist.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
25
README.md
25
README.md
|
|
@ -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` |
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -6,8 +6,6 @@ services:
|
|||
context: .
|
||||
dockerfile: docker/Dockerfile.non_root
|
||||
target: runtime
|
||||
args:
|
||||
PROXY_EXTRAS_SOURCE: "local"
|
||||
depends_on:
|
||||
- squid
|
||||
user: "101:101"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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==",
|
||||
|
|
|
|||
|
|
@ -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,))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
463
litellm-proxy-extras/litellm_proxy_extras/request_log_indexes.py
Normal file
463
litellm-proxy-extras/litellm_proxy_extras/request_log_indexes.py
Normal 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}")
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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==",
|
||||
|
|
|
|||
74
litellm-rust/Cargo.lock
generated
74
litellm-rust/Cargo.lock
generated
|
|
@ -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",
|
||||
]
|
||||
|
||||
|
|
|
|||
|
|
@ -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"] }
|
||||
|
|
|
|||
19
litellm-rust/crates/migrate-macros/Cargo.toml
Normal file
19
litellm-rust/crates/migrate-macros/Cargo.toml
Normal 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
|
||||
21
litellm-rust/crates/migrate-macros/src/error.rs
Normal file
21
litellm-rust/crates/migrate-macros/src/error.rs
Normal 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 },
|
||||
}
|
||||
199
litellm-rust/crates/migrate-macros/src/lib.rs
Normal file
199
litellm-rust/crates/migrate-macros/src/lib.rs
Normal 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 { .. })));
|
||||
}
|
||||
}
|
||||
12
litellm-rust/crates/migrate/Cargo.toml
Normal file
12
litellm-rust/crates/migrate/Cargo.toml
Normal 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
|
||||
5
litellm-rust/crates/migrate/README.md
Normal file
5
litellm-rust/crates/migrate/README.md
Normal 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
|
||||
8
litellm-rust/crates/migrate/src/lib.rs
Normal file
8
litellm-rust/crates/migrate/src/lib.rs
Normal 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,
|
||||
}
|
||||
1
litellm-rust/crates/migrate/tests/fixtures/migrations/10_tenth.sql
vendored
Normal file
1
litellm-rust/crates/migrate/tests/fixtures/migrations/10_tenth.sql
vendored
Normal file
|
|
@ -0,0 +1 @@
|
|||
SELECT 10;
|
||||
1
litellm-rust/crates/migrate/tests/fixtures/migrations/1_first.sql
vendored
Normal file
1
litellm-rust/crates/migrate/tests/fixtures/migrations/1_first.sql
vendored
Normal file
|
|
@ -0,0 +1 @@
|
|||
SELECT 1;
|
||||
1
litellm-rust/crates/migrate/tests/fixtures/migrations/2_second.sql
vendored
Normal file
1
litellm-rust/crates/migrate/tests/fixtures/migrations/2_second.sql
vendored
Normal file
|
|
@ -0,0 +1 @@
|
|||
SELECT 2;
|
||||
21
litellm-rust/crates/migrate/tests/migrate.rs
Normal file
21
litellm-rust/crates/migrate/tests/migrate.rs
Normal 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);
|
||||
}
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
3
litellm-rust/crates/traces/build.rs
Normal file
3
litellm-rust/crates/traces/build.rs
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
fn main() {
|
||||
println!("cargo:rerun-if-changed=migrations");
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -1 +0,0 @@
|
|||
ALTER TABLE {database}.otel_traces ADD COLUMN IF NOT EXISTS EngineReceivedMs UInt64 DEFAULT 0
|
||||
|
|
@ -1 +0,0 @@
|
|||
ALTER TABLE {database}.spend_logs ADD COLUMN IF NOT EXISTS EngineReceivedMs UInt64 DEFAULT 0
|
||||
6
litellm-rust/crates/traces/query/lens_agents.sql
Normal file
6
litellm-rust/crates/traces/query/lens_agents.sql
Normal 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
|
||||
8
litellm-rust/crates/traces/query/lens_availability.sql
Normal file
8
litellm-rust/crates/traces/query/lens_availability.sql
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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>),
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
63
litellm-rust/crates/traces/src/normalize/genai.rs
Normal file
63
litellm-rust/crates/traces/src/normalize/genai.rs
Normal 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(),
|
||||
})
|
||||
}
|
||||
}
|
||||
468
litellm-rust/crates/traces/src/normalize/langsmith.rs
Normal file
468
litellm-rust/crates/traces/src/normalize/langsmith.rs
Normal 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, "[]");
|
||||
}
|
||||
}
|
||||
239
litellm-rust/crates/traces/src/normalize/mod.rs
Normal file
239
litellm-rust/crates/traces/src/normalize/mod.rs
Normal 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());
|
||||
}
|
||||
}
|
||||
53
litellm-rust/crates/traces/src/normalize/openinference.rs
Normal file
53
litellm-rust/crates/traces/src/normalize/openinference.rs
Normal 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(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
})
|
||||
}
|
||||
|
|
|
|||
346
litellm-rust/crates/traces/src/query.rs
Normal file
346
litellm-rust/crates/traces/src/query.rs
Normal 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)
|
||||
);
|
||||
}
|
||||
}
|
||||
89
litellm-rust/crates/traces/src/query/guide.rs
Normal file
89
litellm-rust/crates/traces/src/query/guide.rs
Normal 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)
|
||||
}
|
||||
200
litellm-rust/crates/traces/src/query_access.rs
Normal file
200
litellm-rust/crates/traces/src/query_access.rs
Normal 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('\'', "\\'"))
|
||||
}
|
||||
|
|
@ -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}",
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
65
litellm-rust/crates/traces/templates/query_help.jinja
Normal file
65
litellm-rust/crates/traces/templates/query_help.jinja
Normal 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 %}
|
||||
|
|
@ -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(),
|
||||
¶meters,
|
||||
)
|
||||
.await?,
|
||||
)?;
|
||||
assert_eq!(sample["data"].as_array().expect("rows").len(), 1);
|
||||
assert_eq!(sample["data"][0]["trace_id"], "research");
|
||||
assert_eq!(sample["data"][0]["span_count"], 2);
|
||||
let available: serde_json::Value = serde_json::from_str(
|
||||
&execute_read(
|
||||
&database.client,
|
||||
&connection,
|
||||
LensQuery::Availability.sql(),
|
||||
¶meters,
|
||||
)
|
||||
.await?,
|
||||
)?;
|
||||
assert_eq!(available["data"][0]["traces"], 1);
|
||||
assert_eq!(available["data"][0]["requests"], 0);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::empty(false)]
|
||||
#[case::custom_metadata(true)]
|
||||
#[tokio::test]
|
||||
async fn query_help_discovers_live_schema_and_runs_its_examples(
|
||||
#[future(awt)] database: TestResult<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(())
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"));
|
||||
}
|
||||
|
|
|
|||
268
litellm-rust/crates/traces/tests/query_access.rs
Normal file
268
litellm-rust/crates/traces/tests/query_access.rs
Normal 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(())
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
)
|
||||
|
|
|
|||
98
litellm/harness/__init__.py
Normal file
98
litellm/harness/__init__.py
Normal 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",
|
||||
)
|
||||
62
litellm/harness/context.py
Normal file
62
litellm/harness/context.py
Normal 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
689
litellm/harness/endpoint.py
Normal 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
45
litellm/harness/errors.py
Normal 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
|
||||
35
litellm/harness/handlers/__init__.py
Normal file
35
litellm/harness/handlers/__init__.py
Normal 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")
|
||||
42
litellm/harness/handlers/base.py
Normal file
42
litellm/harness/handlers/base.py
Normal 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")
|
||||
161
litellm/harness/handlers/cli_handler.py
Normal file
161
litellm/harness/handlers/cli_handler.py
Normal 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
|
||||
261
litellm/harness/handlers/deepagents_handler.py
Normal file
261
litellm/harness/handlers/deepagents_handler.py
Normal 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
|
||||
37
litellm/harness/options.py
Normal file
37
litellm/harness/options.py
Normal 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
1070
litellm/harness/runtime.py
Normal file
File diff suppressed because it is too large
Load diff
25
litellm/harness/sandbox/__init__.py
Normal file
25
litellm/harness/sandbox/__init__.py
Normal 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",
|
||||
)
|
||||
62
litellm/harness/sandbox/base.py
Normal file
62
litellm/harness/sandbox/base.py
Normal 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: ...
|
||||
298
litellm/harness/sandbox/docker.py
Normal file
298
litellm/harness/sandbox/docker.py
Normal 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)
|
||||
277
litellm/harness/sandbox/local.py
Normal file
277
litellm/harness/sandbox/local.py
Normal 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)
|
||||
183
litellm/harness/sandbox/snapshot.py
Normal file
183
litellm/harness/sandbox/snapshot.py
Normal 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
443
litellm/harness/sync.py
Normal 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
214
litellm/harness/types.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
0
litellm/llms/base_llm/harness/__init__.py
Normal file
0
litellm/llms/base_llm/harness/__init__.py
Normal file
151
litellm/llms/base_llm/harness/transformation.py
Normal file
151
litellm/llms/base_llm/harness/transformation.py
Normal 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."""
|
||||
116
litellm/llms/base_llm/harness/utils.py
Normal file
116
litellm/llms/base_llm/harness/utils.py
Normal 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))
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
0
litellm/llms/claude_code/__init__.py
Normal file
0
litellm/llms/claude_code/__init__.py
Normal file
0
litellm/llms/claude_code/harness/__init__.py
Normal file
0
litellm/llms/claude_code/harness/__init__.py
Normal file
389
litellm/llms/claude_code/harness/transformation.py
Normal file
389
litellm/llms/claude_code/harness/transformation.py
Normal 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)
|
||||
0
litellm/llms/codex/__init__.py
Normal file
0
litellm/llms/codex/__init__.py
Normal file
0
litellm/llms/codex/harness/__init__.py
Normal file
0
litellm/llms/codex/harness/__init__.py
Normal file
349
litellm/llms/codex/harness/transformation.py
Normal file
349
litellm/llms/codex/harness/transformation.py
Normal 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)
|
||||
0
litellm/llms/deepagents/__init__.py
Normal file
0
litellm/llms/deepagents/__init__.py
Normal file
0
litellm/llms/deepagents/harness/__init__.py
Normal file
0
litellm/llms/deepagents/harness/__init__.py
Normal file
559
litellm/llms/deepagents/harness/sandbox_backend.py
Normal file
559
litellm/llms/deepagents/harness/sandbox_backend.py
Normal 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)
|
||||
313
litellm/llms/deepagents/harness/transformation.py
Normal file
313
litellm/llms/deepagents/harness/transformation.py
Normal 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=")
|
||||
0
litellm/llms/opencode/__init__.py
Normal file
0
litellm/llms/opencode/__init__.py
Normal file
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue