mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
Merge origin/main into litellm_otel_v2_tenant_internal_spans
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
commit
729744a9ed
653 changed files with 58152 additions and 10077 deletions
|
|
@ -3222,7 +3222,7 @@ workflows:
|
|||
cron: "17 0,6,12,18 * * *"
|
||||
filters:
|
||||
branches:
|
||||
only: litellm_internal_staging
|
||||
only: main
|
||||
jobs: *migration_jobs
|
||||
integration:
|
||||
unless: << pipeline.parameters.run_migration_tests >>
|
||||
|
|
@ -3231,7 +3231,7 @@ workflows:
|
|||
name: integration-<< matrix.suite >>
|
||||
matrix:
|
||||
parameters:
|
||||
suite: [management, accounting, database, providers, extensions, sdk, cost, browser]
|
||||
suite: [management, accounting, database, providers, extensions, mcp, sdk, cost, browser]
|
||||
filters:
|
||||
branches:
|
||||
only:
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ run_full() {
|
|||
|
||||
[ -n "${CIRCLE_PULL_REQUEST:-}" ] || run_full "not a pull request"
|
||||
|
||||
candidate_bases="main litellm_internal_staging litellm_oss_staging"
|
||||
candidate_bases="main"
|
||||
merge_base=""
|
||||
for base in $candidate_bases; do
|
||||
git fetch --quiet origin "$base" 2>/dev/null || continue
|
||||
|
|
|
|||
|
|
@ -9,7 +9,6 @@ fi
|
|||
suite="${1:?integration suite required}"
|
||||
results="test-results/integration-${suite}"
|
||||
mkdir -p "$results"
|
||||
shard_timeout=11m
|
||||
integration_identity="$(.venv/bin/python -c 'import uuid; print(uuid.uuid4().hex)')"
|
||||
upstream_pid=""
|
||||
proxy_pid=""
|
||||
|
|
@ -112,6 +111,15 @@ upstream_pid=$!
|
|||
if [ "$suite" = cost ]; then
|
||||
export INTEGRATION_WORKERS=8
|
||||
fi
|
||||
if [ "$suite" = mcp ]; then
|
||||
export INTEGRATION_WORKERS=4 INTEGRATION_COVERAGE=1
|
||||
fi
|
||||
coverage_data="$PWD/$results/coverage/data"
|
||||
proxy_command=(.venv/bin/python -m integration._support.proxy)
|
||||
if [ "${INTEGRATION_COVERAGE:-0}" = 1 ]; then
|
||||
mkdir -p "$(dirname "$coverage_data")"
|
||||
proxy_command=(.venv/bin/python -m coverage run --rcfile=tests/integration/mcp_coverage.toml -m integration._support.proxy)
|
||||
fi
|
||||
start_proxy() {
|
||||
local port="$1"
|
||||
local log_name="$2"
|
||||
|
|
@ -131,10 +139,11 @@ start_proxy() {
|
|||
fi
|
||||
setsid env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" INTEGRATION_RUN_ID="$integration_identity" \
|
||||
DATABASE_URL="$DATABASE_URL" REDIS_HOST="$REDIS_HOST" REDIS_PORT="$REDIS_PORT" \
|
||||
INTEGRATION_UPSTREAM_URL="$INTEGRATION_UPSTREAM_URL" \
|
||||
LITELLM_MASTER_KEY="$LITELLM_MASTER_KEY" LITELLM_SALT_KEY="$LITELLM_SALT_KEY" LITELLM_UI_PATH="$LITELLM_UI_PATH" PROXY_BASE_URL="http://127.0.0.1:$port" \
|
||||
LITELLM_MODE=PRODUCTION STORE_MODEL_IN_DB=True "${cost_map_env[@]}" \
|
||||
AWS_EC2_METADATA_DISABLED=true DO_NOT_TRACK=1 \
|
||||
.venv/bin/python -m integration._support.proxy --config tests/integration/proxy_config.yaml \
|
||||
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 \
|
||||
> "$results/$log_name" 2>&1 &
|
||||
|
|
@ -146,7 +155,7 @@ proxy_pid="$launched_pid"
|
|||
curl --noproxy '*' -sSf -X POST "$INTEGRATION_PROXY_URL/config/update" \
|
||||
-H "Authorization: Bearer $LITELLM_MASTER_KEY" -H 'Content-Type: application/json' \
|
||||
-d '{"router_settings": {"num_retries": 0}}' > "$results/seed-router-settings.json"
|
||||
if [ "$suite" = management ]; then
|
||||
if [ "$suite" = management ] || [ "$suite" = mcp ]; then
|
||||
export INTEGRATION_PEER_URL=http://127.0.0.1:4001
|
||||
start_proxy 4001 peer.log
|
||||
peer_pid="$launched_pid"
|
||||
|
|
@ -176,7 +185,7 @@ if [ "$suite" = browser ]; then
|
|||
exit 0
|
||||
fi
|
||||
|
||||
timeout --signal=TERM --kill-after=20s "$shard_timeout" env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" \
|
||||
env -i PATH="$PATH" HOME="$HOME" PYTHONPATH="$PYTHONPATH" \
|
||||
INTEGRATION_RUN_ID="$integration_identity" \
|
||||
DATABASE_URL="$DATABASE_URL" REDIS_HOST="$REDIS_HOST" REDIS_PORT="$REDIS_PORT" \
|
||||
INTEGRATION_PROXY_URL="$INTEGRATION_PROXY_URL" INTEGRATION_PEER_URL="$INTEGRATION_PEER_URL" \
|
||||
|
|
@ -187,3 +196,23 @@ timeout --signal=TERM --kill-after=20s "$shard_timeout" env -i PATH="$PATH" HOME
|
|||
INTEGRATION_ORDER_SEED="$INTEGRATION_ORDER_SEED" \
|
||||
LITELLM_LOCAL_MODEL_COST_MAP=True AWS_EC2_METADATA_DISABLED=true DO_NOT_TRACK=1 \
|
||||
.venv/bin/python tests/integration/run.py "$suite" --results "$results"
|
||||
|
||||
if [ "${INTEGRATION_COVERAGE:-0}" = 1 ]; then
|
||||
for covered_pid in "$proxy_pid" "$peer_pid"; do
|
||||
[ -n "$covered_pid" ] || continue
|
||||
kill -TERM -- "-$covered_pid"
|
||||
for _ in {1..300}; do
|
||||
kill -0 "$covered_pid" 2>/dev/null || break
|
||||
sleep 0.1
|
||||
done
|
||||
wait "$covered_pid" 2>/dev/null || true
|
||||
done
|
||||
proxy_pid=""
|
||||
peer_pid=""
|
||||
COVERAGE_FILE="$coverage_data" .venv/bin/python -m coverage combine --rcfile=tests/integration/mcp_coverage.toml
|
||||
COVERAGE_FILE="$coverage_data" .venv/bin/python -m coverage report --rcfile=tests/integration/mcp_coverage.toml \
|
||||
> "$results/coverage/coverage.txt"
|
||||
COVERAGE_FILE="$coverage_data" .venv/bin/python -m coverage html --rcfile=tests/integration/mcp_coverage.toml \
|
||||
-d "$results/coverage/html"
|
||||
tail -n 1 "$results/coverage/coverage.txt"
|
||||
fi
|
||||
|
|
|
|||
|
|
@ -31,8 +31,8 @@ def main() -> None:
|
|||
result: Final = json.loads(Path(sys.argv[1]).read_text())
|
||||
assert not result.get("errors"), result.get("errors")
|
||||
expected: Final = json.loads(
|
||||
(Path(__file__).resolve().parents[2] / "tests/integration/contracts.json").read_text()
|
||||
)["browser"]
|
||||
(Path(__file__).resolve().parents[2] / "tests/e2e/ui/tests/integrationCritical/expected.json").read_text()
|
||||
)
|
||||
assert expected and result["stats"]["expected"] == len(expected)
|
||||
assert all(result["stats"][name] == 0 for name in ("unexpected", "flaky", "skipped"))
|
||||
|
||||
|
|
|
|||
|
|
@ -8,7 +8,6 @@
|
|||
#
|
||||
# Protected branches (always allowed):
|
||||
# - main
|
||||
# - litellm_internal_staging
|
||||
# - dependabot/*
|
||||
# - gh-readonly-queue/*
|
||||
#
|
||||
|
|
@ -22,7 +21,7 @@ ZERO_OID_SHA256="000000000000000000000000000000000000000000000000000000000000000
|
|||
ALLOWED_TYPES="feature|bugfix|hotfix|release|chore"
|
||||
BRANCH_PATTERN="^(${ALLOWED_TYPES})/.+"
|
||||
|
||||
PROTECTED_NAMES="main litellm_internal_staging"
|
||||
PROTECTED_NAMES="main"
|
||||
PROTECTED_PREFIXES="dependabot/ gh-readonly-queue/"
|
||||
|
||||
is_protected() {
|
||||
|
|
@ -78,8 +77,7 @@ if [ -n "$invalid" ]; then
|
|||
chore/bump-deps
|
||||
hotfix/auth-bypass
|
||||
|
||||
Protected (always allowed): main, litellm_internal_staging,
|
||||
dependabot/*, gh-readonly-queue/*.
|
||||
Protected (always allowed): main, dependabot/*, gh-readonly-queue/*.
|
||||
|
||||
See https://conventional-branch.github.io/
|
||||
|
||||
|
|
|
|||
1
.github/e2e-stack/select_tests.py
vendored
1
.github/e2e-stack/select_tests.py
vendored
|
|
@ -11,6 +11,7 @@ UNSUPPORTED: Final = re.compile(
|
|||
r"|^tests/e2e/guardrails/test_presidio_masking_e2e\.py$"
|
||||
r"|^tests/e2e/logging/test_otel_v2_langfuse_generation_output_e2e\.py$"
|
||||
r"|^tests/e2e/logging/test_langsmith_batch_serialization_e2e\.py$"
|
||||
r"|^tests/e2e/secret_manager/"
|
||||
)
|
||||
HARNESS: Final = re.compile(
|
||||
r"^tests/e2e/[A-Za-z0-9_.-]+\.(py|ini)$"
|
||||
|
|
|
|||
15
.github/merge-smoke-tests.json
vendored
Normal file
15
.github/merge-smoke-tests.json
vendored
Normal file
|
|
@ -0,0 +1,15 @@
|
|||
{
|
||||
"cases": {
|
||||
"CHAT-JSON": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_returns_json_reply_over_injected_transport",
|
||||
"CHAT-TEXT-STREAM": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_streams_text_deltas_over_injected_transport",
|
||||
"CHAT-TOOL-STREAM": "tests/test_litellm/llms/openai/test_openai.py::test_acompletion_streams_tool_call_arguments_over_injected_transport",
|
||||
"MODEL-ALLOW": "tests/test_litellm/proxy/auth/test_auth_checks.py::test_can_object_call_model_allows_listed_model_for_key",
|
||||
"MODEL-DENY": "tests/test_litellm/proxy/auth/test_auth_checks.py::test_can_object_call_model_denials_return_forbidden[key-key_model_access_denied]",
|
||||
"COST-EXPLICIT": "tests/test_litellm/test_cost_calculator.py::test_completion_cost_charges_explicit_per_token_rates_over_registered_ones",
|
||||
"COST-ZERO": "tests/test_litellm/test_cost_calculator.py::test_completion_cost_is_zero_when_explicit_rates_are_zero",
|
||||
"LOG-CONTENT-ON": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_keeps_message_content_when_message_logging_is_on",
|
||||
"LOG-CONTENT-OFF": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_standard_logging_payload_redacts_message_content_when_message_logging_is_off",
|
||||
"CALLBACK-SUCCESS": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_async_success_handler_delivers_standard_logging_payload_to_custom_logger",
|
||||
"CALLBACK-FAILURE": "tests/test_litellm/litellm_core_utils/test_litellm_logging.py::test_async_failure_handler_delivers_failure_payload_to_custom_logger"
|
||||
}
|
||||
}
|
||||
58
.github/scripts/assert_ci_coverage.py
vendored
58
.github/scripts/assert_ci_coverage.py
vendored
|
|
@ -235,9 +235,7 @@ class Slice:
|
|||
return True # a `-k` this parser cannot model is assumed to claim everything
|
||||
if any(term.lower() in relative_path.lower() for term in self.excluded):
|
||||
return False
|
||||
return not self.required or any(
|
||||
term.lower() in name.lower() for term in self.required for name in inner_names
|
||||
)
|
||||
return not self.required or any(term.lower() in name.lower() for term in self.required for name in inner_names)
|
||||
|
||||
|
||||
def _strings(node: object) -> Iterable[str]:
|
||||
|
|
@ -307,9 +305,7 @@ def _matchable_names(relative_path: str) -> frozenset[str]:
|
|||
except (OSError, SyntaxError):
|
||||
return frozenset({relative_path})
|
||||
return frozenset({relative_path}) | frozenset(
|
||||
node.name
|
||||
for node in ast.walk(tree)
|
||||
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef))
|
||||
node.name for node in ast.walk(tree) if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef))
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -331,9 +327,7 @@ def _deselected_everywhere(allowlist: Allowlist) -> tuple[Finding, ...]:
|
|||
slices: Final = _slices()
|
||||
named_by_workflow: Final = _workflow_named_tokens()
|
||||
globbed: Final = tuple(
|
||||
path
|
||||
for path in _test_files()
|
||||
if any(_token_covers(glob, path) for slice_ in slices for glob in slice_.globs)
|
||||
path for path in _test_files() if any(_token_covers(glob, path) for slice_ in slices for glob in slice_.globs)
|
||||
)
|
||||
return tuple(
|
||||
Finding(
|
||||
|
|
@ -363,11 +357,7 @@ def _shard_children(root: str, repo_root: pathlib.Path = REPO_ROOT) -> tuple[str
|
|||
child.relative_to(repo_root).as_posix()
|
||||
for child in (repo_root / root).iterdir()
|
||||
if not child.name.startswith(".")
|
||||
and (
|
||||
_holds_tests(child)
|
||||
if child.is_dir()
|
||||
else child.name.startswith("test_") and child.suffix == ".py"
|
||||
)
|
||||
and (_holds_tests(child) if child.is_dir() else child.name.startswith("test_") and child.suffix == ".py")
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -499,13 +489,32 @@ def _check_shards() -> int:
|
|||
return 0
|
||||
|
||||
|
||||
def _integration_groups(runner: pathlib.Path) -> dict[str, tuple[str, ...]]:
|
||||
module: Final = ast.parse(runner.read_text())
|
||||
literal: Final = next(
|
||||
node.value
|
||||
for node in module.body
|
||||
if isinstance(node, ast.AnnAssign) and isinstance(node.target, ast.Name) and node.target.id == "GROUPS"
|
||||
)
|
||||
mapping: Final = literal.args[0] if isinstance(literal, ast.Call) else literal
|
||||
return {group: tuple(folders) for group, folders in ast.literal_eval(mapping).items()}
|
||||
|
||||
|
||||
def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozenset[str], tuple[Finding, ...]]:
|
||||
manifest: Final = repo_root / "tests/integration/contracts.json"
|
||||
if not manifest.exists():
|
||||
runner: Final = repo_root / "tests/integration/run.py"
|
||||
if not runner.exists():
|
||||
return frozenset(), ()
|
||||
entries: Final = json.loads(manifest.read_text())
|
||||
paths: Final = frozenset(node.split("::", 1)[0] for node in entries["tests"])
|
||||
browser_paths: Final = frozenset(node.split("::", 1)[0] for node in entries.get("browser", {}))
|
||||
groups: Final = _integration_groups(runner)
|
||||
integration_root: Final = repo_root / "tests/integration"
|
||||
paths: Final = frozenset(
|
||||
str(path.relative_to(repo_root))
|
||||
for folders in groups.values()
|
||||
for folder in folders
|
||||
for path in (integration_root / folder).glob("test_*.py")
|
||||
)
|
||||
browser_manifest: Final = repo_root / "tests/e2e/ui/tests/integrationCritical/expected.json"
|
||||
browser_nodes: Final = json.loads(browser_manifest.read_text()) if browser_manifest.exists() else ()
|
||||
browser_paths: Final = frozenset(node.split("::", 1)[0] for node in browser_nodes)
|
||||
circle_path: Final = repo_root / ".circleci/config.yml"
|
||||
circle: Final = yaml.safe_load(circle_path.read_text()) if circle_path.exists() else {}
|
||||
steps: Final = circle.get("jobs", {}).get("integration_contracts", {}).get("steps", ())
|
||||
|
|
@ -526,15 +535,14 @@ def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozens
|
|||
)
|
||||
required: Final = (frozenset({"browser"}) if browser_paths else frozenset()) | frozenset(
|
||||
group
|
||||
for group, folders in entries["groups"].items()
|
||||
for group, folders in groups.items()
|
||||
if any(any(path.startswith(f"tests/integration/{folder}/") for folder in folders) for path in paths)
|
||||
)
|
||||
ungrouped: Final = frozenset(
|
||||
path
|
||||
for path in paths
|
||||
if sum(
|
||||
any(path.startswith(f"tests/integration/{folder}/") for folder in folders)
|
||||
for folders in entries["groups"].values()
|
||||
any(path.startswith(f"tests/integration/{folder}/") for folder in folders) for folders in groups.values()
|
||||
)
|
||||
!= 1
|
||||
)
|
||||
|
|
@ -547,10 +555,6 @@ def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozens
|
|||
Finding(path, "integration contract is also selected by GitHub Actions")
|
||||
for path in paths
|
||||
if any(_token_covers(token, path) for token in gha_tokens)
|
||||
) + tuple(
|
||||
Finding(path, "canonical integration test file is missing")
|
||||
for path in paths
|
||||
if not (repo_root / path).is_file()
|
||||
)
|
||||
browser_commands: Final = tuple(
|
||||
scalar.value
|
||||
|
|
@ -592,7 +596,7 @@ def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozens
|
|||
) + tuple(Finding(path, "canonical node must have exactly one integration group") for path in sorted(ungrouped))
|
||||
if not paths or not invoked or not scheduled:
|
||||
return frozenset(), findings + (
|
||||
Finding(str(manifest.relative_to(repo_root)), "dedicated CircleCI runner is missing"),
|
||||
Finding(str(runner.relative_to(repo_root)), "dedicated CircleCI runner is missing"),
|
||||
)
|
||||
return paths | browser_paths, findings + group_findings + browser_findings + exclusion_findings
|
||||
|
||||
|
|
|
|||
493
.github/scripts/run_merge_smoke.py
vendored
Normal file
493
.github/scripts/run_merge_smoke.py
vendored
Normal file
|
|
@ -0,0 +1,493 @@
|
|||
#!/usr/bin/env python3
|
||||
"""Merge smoke harness: bounded checks run inside a loopback-only Linux network namespace."""
|
||||
|
||||
# ruff: noqa: T201 # CLI harness: stdout/stderr lines are the reported result
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import contextlib
|
||||
import http.client
|
||||
import json
|
||||
import os
|
||||
import secrets
|
||||
import signal
|
||||
import socket
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
from collections import Counter
|
||||
from collections.abc import Sequence
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from types import MappingProxyType
|
||||
from typing import Final, NoReturn, TextIO, cast
|
||||
|
||||
import pytest
|
||||
|
||||
EXPECTED_CASES: Final = (
|
||||
"CHAT-JSON",
|
||||
"CHAT-TEXT-STREAM",
|
||||
"CHAT-TOOL-STREAM",
|
||||
"MODEL-ALLOW",
|
||||
"MODEL-DENY",
|
||||
"COST-EXPLICIT",
|
||||
"COST-ZERO",
|
||||
"LOG-CONTENT-ON",
|
||||
"LOG-CONTENT-OFF",
|
||||
"CALLBACK-SUCCESS",
|
||||
"CALLBACK-FAILURE",
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CheckResult:
|
||||
ok: bool
|
||||
detail: str = ""
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class _Args:
|
||||
command: str = ""
|
||||
no_child: bool = False
|
||||
expect: str = ""
|
||||
litellm_bin: str | None = None
|
||||
lite_bin: str | None = None
|
||||
diagnostics_dir: str = ""
|
||||
ready_deadline: float = 120.0
|
||||
shutdown_deadline: float = 20.0
|
||||
poll_interval: float = 0.5
|
||||
manifest: str = ""
|
||||
rootdir: str | None = None
|
||||
|
||||
|
||||
def fail(reason: str) -> NoReturn:
|
||||
print(f"merge-smoke: FAIL {reason}", file=sys.stderr)
|
||||
sys.exit(1)
|
||||
|
||||
|
||||
def ok(step: str) -> None:
|
||||
print(f"merge-smoke: OK {step}")
|
||||
|
||||
|
||||
def tail(path: Path, lines: int = 20) -> str:
|
||||
try:
|
||||
return "\n".join(path.read_text(errors="replace").splitlines()[-lines:])
|
||||
except OSError as exc:
|
||||
return f"<cannot read {path}: {exc}>"
|
||||
|
||||
|
||||
def cmd_verify_isolation(args: _Args) -> int:
|
||||
if os.geteuid() == 0:
|
||||
fail("verify-isolation must run unprivileged (geteuid()==0)")
|
||||
try:
|
||||
socket.create_connection(("192.0.2.1", 9), timeout=3)
|
||||
except OSError as exc:
|
||||
print(f"external connect blocked as expected: errno={exc.errno} {exc}")
|
||||
else:
|
||||
fail("external TCP connect to 192.0.2.1:9 succeeded; namespace is not isolated")
|
||||
listener: Final = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
||||
listener.bind(("127.0.0.1", 0))
|
||||
listener.listen(1)
|
||||
port: Final = cast(int, listener.getsockname()[1])
|
||||
client: Final = socket.create_connection(("127.0.0.1", port), timeout=5)
|
||||
accepted: Final = listener.accept()
|
||||
accepted[0].close()
|
||||
client.close()
|
||||
listener.close()
|
||||
print(f"loopback connect ok on 127.0.0.1:{port}")
|
||||
if not args.no_child:
|
||||
proc: Final = subprocess.run(
|
||||
[sys.executable, str(Path(__file__).resolve()), "verify-isolation", "--no-child"],
|
||||
timeout=30,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
)
|
||||
if proc.returncode != 0:
|
||||
fail(f"child process did not inherit isolation: {proc.stderr.strip()}")
|
||||
print("child process inherits isolation")
|
||||
ok("verify-isolation")
|
||||
return 0
|
||||
|
||||
|
||||
def cmd_interpreter(args: _Args) -> int:
|
||||
print(sys.version)
|
||||
print(sys.executable)
|
||||
actual: Final = f"{sys.version_info.major}.{sys.version_info.minor}"
|
||||
if actual != args.expect:
|
||||
fail(f"interpreter is {actual}, expected {args.expect}")
|
||||
ok(f"interpreter {actual}")
|
||||
return 0
|
||||
|
||||
|
||||
def _run_cli(argv: Sequence[str], label: str) -> CheckResult:
|
||||
try:
|
||||
proc: Final = subprocess.run(list(argv), timeout=120, capture_output=True, text=True)
|
||||
except subprocess.TimeoutExpired:
|
||||
return CheckResult(ok=False, detail=f"{label} timed out after 120s")
|
||||
sys.stdout.write(proc.stdout)
|
||||
sys.stderr.write(proc.stderr)
|
||||
if proc.returncode != 0:
|
||||
return CheckResult(ok=False, detail=f"{label} exited {proc.returncode}")
|
||||
return CheckResult(ok=True)
|
||||
|
||||
|
||||
def cmd_cli(args: _Args) -> int:
|
||||
venv_bin: Final = Path(sys.executable).parent
|
||||
litellm_bin: Final = Path(args.litellm_bin) if args.litellm_bin else venv_bin / "litellm"
|
||||
lite_bin: Final = Path(args.lite_bin) if args.lite_bin else venv_bin / "lite"
|
||||
commands: Final = (
|
||||
("import litellm", [sys.executable, "-c", "import litellm"]),
|
||||
("litellm --version", [str(litellm_bin), "--version"]),
|
||||
("lite version", [str(lite_bin), "version"]),
|
||||
)
|
||||
for label, argv in commands:
|
||||
result = _run_cli(argv, label)
|
||||
if not result.ok:
|
||||
fail(result.detail)
|
||||
ok(label)
|
||||
return 0
|
||||
|
||||
|
||||
def _free_port() -> int:
|
||||
sock: Final = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
|
||||
sock.bind(("127.0.0.1", 0))
|
||||
port: Final = cast(int, sock.getsockname()[1])
|
||||
sock.close()
|
||||
return port
|
||||
|
||||
|
||||
_CONFIG_TEMPLATE: Final = """model_list:
|
||||
- model_name: smoke-model
|
||||
litellm_params:
|
||||
model: openai/smoke-model
|
||||
api_base: http://127.0.0.1:9/v1
|
||||
api_key: synthetic-key
|
||||
general_settings:
|
||||
master_key: os.environ/LITELLM_MASTER_KEY
|
||||
"""
|
||||
|
||||
|
||||
def _listen_inode(port: int) -> str | None:
|
||||
target: Final = f"{port:04X}"
|
||||
for table in ("/proc/net/tcp", "/proc/net/tcp6"):
|
||||
try:
|
||||
rows = Path(table).read_text().splitlines()[1:]
|
||||
except OSError:
|
||||
continue
|
||||
for row in rows:
|
||||
cols = row.split()
|
||||
if len(cols) > 9 and cols[3] == "0A" and cols[1].rsplit(":", 1)[-1] == target:
|
||||
return cols[9]
|
||||
return None
|
||||
|
||||
|
||||
def _ancestors(pid: int) -> frozenset[int]:
|
||||
chain: Final[set[int]] = set()
|
||||
pending: Final[list[int]] = [pid]
|
||||
while pending:
|
||||
current = pending.pop()
|
||||
if current <= 0 or current in chain:
|
||||
continue
|
||||
chain.add(current)
|
||||
try:
|
||||
stat = Path(f"/proc/{current}/stat").read_text()
|
||||
except OSError:
|
||||
continue
|
||||
pending.append(int(stat.rpartition(")")[2].split()[1]))
|
||||
return frozenset(chain)
|
||||
|
||||
|
||||
def _socket_owner_pid(inode: str) -> int | None:
|
||||
for proc_dir in Path("/proc").iterdir():
|
||||
if not proc_dir.name.isdigit():
|
||||
continue
|
||||
fd_dir = proc_dir / "fd"
|
||||
try:
|
||||
for fd in fd_dir.iterdir():
|
||||
try:
|
||||
if os.readlink(fd) == f"socket:[{inode}]":
|
||||
return int(proc_dir.name)
|
||||
except OSError:
|
||||
continue
|
||||
except OSError:
|
||||
continue
|
||||
return None
|
||||
|
||||
|
||||
def _verify_port_owner(port: int, proc: subprocess.Popen[bytes]) -> CheckResult:
|
||||
inode: Final = _listen_inode(port)
|
||||
if inode is None:
|
||||
return CheckResult(ok=False, detail=f"no LISTEN socket found for port {port} in /proc/net/tcp")
|
||||
owner: Final = _socket_owner_pid(inode)
|
||||
if owner is None:
|
||||
return CheckResult(ok=False, detail=f"no process owns the listen socket inode {inode} for port {port}")
|
||||
if owner != proc.pid and proc.pid not in _ancestors(owner):
|
||||
return CheckResult(
|
||||
ok=False, detail=f"port {port} owned by pid {owner} outside the launched process group {proc.pid}"
|
||||
)
|
||||
if proc.poll() is not None:
|
||||
return CheckResult(ok=False, detail=f"proxy exited with code {proc.returncode} after readiness")
|
||||
return CheckResult(ok=True)
|
||||
|
||||
|
||||
def cmd_proxy_startup(args: _Args) -> int:
|
||||
diagnostics: Final = Path(args.diagnostics_dir)
|
||||
diagnostics.mkdir(parents=True, exist_ok=True)
|
||||
venv_bin: Final = Path(sys.executable).parent
|
||||
litellm_bin: Final = Path(args.litellm_bin) if args.litellm_bin else venv_bin / "litellm"
|
||||
port: Final = _free_port()
|
||||
master_key: Final = "sk-smoke-" + secrets.token_hex(16)
|
||||
config_path: Final = diagnostics / "config.yaml"
|
||||
config_path.write_text(_CONFIG_TEMPLATE)
|
||||
log_path: Final = diagnostics / "proxy.log"
|
||||
result_path: Final = diagnostics / "result.json"
|
||||
outcome: Final[dict[str, object]] = {
|
||||
"port": port,
|
||||
"time_to_ready_s": None,
|
||||
"shutdown_s": None,
|
||||
"readiness": None,
|
||||
"outcome": "failed",
|
||||
}
|
||||
log_file: Final = log_path.open("w")
|
||||
env: Final = {
|
||||
**os.environ,
|
||||
"LITELLM_MASTER_KEY": master_key,
|
||||
"LITELLM_LOCAL_MODEL_COST_MAP": "True",
|
||||
}
|
||||
started: Final = time.monotonic()
|
||||
proc: Final = subprocess.Popen(
|
||||
[str(litellm_bin), "--config", str(config_path), "--host", "127.0.0.1", "--port", str(port)],
|
||||
stdout=log_file,
|
||||
stderr=subprocess.STDOUT,
|
||||
start_new_session=True,
|
||||
env=env,
|
||||
)
|
||||
body: str | None = None
|
||||
last_status: int | None = None
|
||||
while time.monotonic() - started < args.ready_deadline:
|
||||
if proc.poll() is not None:
|
||||
log_file.close()
|
||||
result_path.write_text(json.dumps(outcome))
|
||||
fail(f"proxy exited early with code {proc.returncode}\n{tail(log_path)}")
|
||||
try:
|
||||
conn = http.client.HTTPConnection("127.0.0.1", port, timeout=5)
|
||||
conn.request("GET", "/health/readiness")
|
||||
resp = conn.getresponse()
|
||||
last_status = resp.status
|
||||
candidate = resp.read().decode()
|
||||
conn.close()
|
||||
except (http.client.HTTPException, ConnectionError, OSError):
|
||||
time.sleep(args.poll_interval)
|
||||
continue
|
||||
if last_status == 200:
|
||||
body = candidate
|
||||
break
|
||||
time.sleep(args.poll_interval)
|
||||
outcome["time_to_ready_s"] = round(time.monotonic() - started, 3)
|
||||
if body is None:
|
||||
_terminate(proc, log_file)
|
||||
result_path.write_text(json.dumps(outcome))
|
||||
detail = f"last status {last_status}" if last_status is not None else "no response"
|
||||
fail(f"readiness not reached within {args.ready_deadline}s ({detail})\n{tail(log_path)}")
|
||||
outcome["readiness"] = body
|
||||
try:
|
||||
readiness = cast(object, json.loads(body))
|
||||
except json.JSONDecodeError:
|
||||
readiness = None
|
||||
if readiness != {"status": "healthy", "db": "Not connected"}:
|
||||
_terminate(proc, log_file)
|
||||
result_path.write_text(json.dumps(outcome))
|
||||
fail(f"unexpected readiness body: {body}")
|
||||
owner_check: Final = _verify_port_owner(port, proc)
|
||||
if not owner_check.ok:
|
||||
_terminate(proc, log_file)
|
||||
result_path.write_text(json.dumps(outcome))
|
||||
fail(owner_check.detail)
|
||||
shutdown_started: Final = time.monotonic()
|
||||
os.killpg(proc.pid, signal.SIGTERM)
|
||||
try:
|
||||
proc.wait(timeout=args.shutdown_deadline)
|
||||
except subprocess.TimeoutExpired:
|
||||
os.killpg(proc.pid, signal.SIGKILL)
|
||||
proc.wait(timeout=10)
|
||||
outcome["shutdown_s"] = round(time.monotonic() - shutdown_started, 3)
|
||||
log_file.close()
|
||||
result_path.write_text(json.dumps(outcome))
|
||||
fail(f"forced kill after {args.shutdown_deadline}s\n{tail(log_path)}")
|
||||
outcome["shutdown_s"] = round(time.monotonic() - shutdown_started, 3)
|
||||
try:
|
||||
os.killpg(proc.pid, 0)
|
||||
except ProcessLookupError:
|
||||
pass
|
||||
else:
|
||||
os.killpg(proc.pid, signal.SIGKILL)
|
||||
log_file.close()
|
||||
result_path.write_text(json.dumps(outcome))
|
||||
fail("process group survived SIGTERM")
|
||||
log_file.close()
|
||||
outcome["outcome"] = "ok"
|
||||
result_path.write_text(json.dumps(outcome))
|
||||
ok(f"proxy-startup ready={outcome['time_to_ready_s']}s shutdown={outcome['shutdown_s']}s")
|
||||
return 0
|
||||
|
||||
|
||||
def _terminate(proc: subprocess.Popen[bytes], log_file: TextIO) -> None:
|
||||
with contextlib.suppress(ProcessLookupError):
|
||||
os.killpg(proc.pid, signal.SIGTERM)
|
||||
try:
|
||||
proc.wait(timeout=10)
|
||||
except subprocess.TimeoutExpired:
|
||||
with contextlib.suppress(ProcessLookupError):
|
||||
os.killpg(proc.pid, signal.SIGKILL)
|
||||
with contextlib.suppress(subprocess.TimeoutExpired):
|
||||
proc.wait(timeout=10)
|
||||
log_file.close()
|
||||
|
||||
|
||||
def _load_manifest(path: Path) -> MappingProxyType[str, str]:
|
||||
def no_duplicates(pairs: list[tuple[object, object]]) -> dict[object, object]:
|
||||
seen: dict[object, object] = {}
|
||||
for key, value in pairs:
|
||||
if key in seen:
|
||||
raise ValueError(f"duplicate key in manifest: {key}")
|
||||
seen[key] = value
|
||||
return seen
|
||||
|
||||
raw_value: object = cast(object, json.loads(path.read_text(), object_pairs_hook=no_duplicates))
|
||||
if not isinstance(raw_value, dict):
|
||||
raise ValueError("manifest must be an object")
|
||||
loaded: Final = cast(dict[object, object], raw_value)
|
||||
cases_value: object = loaded.get("cases")
|
||||
if not isinstance(cases_value, dict):
|
||||
raise ValueError("manifest must be an object with a 'cases' object")
|
||||
cases_any: Final = cast(dict[object, object], cases_value)
|
||||
cases: Final = {k: v for k, v in cases_any.items() if isinstance(k, str) and isinstance(v, str)}
|
||||
if len(cases) != len(cases_any):
|
||||
raise ValueError("manifest 'cases' must map string ids to string node ids")
|
||||
return MappingProxyType(cases)
|
||||
|
||||
|
||||
@dataclass(slots=True, eq=False)
|
||||
class _Recorder:
|
||||
collect_failed: list[str] = field(default_factory=list)
|
||||
collected: tuple[str, ...] = ()
|
||||
reports: dict[str, list[tuple[str, str, bool]]] = field(default_factory=dict)
|
||||
|
||||
def pytest_collectreport(self, report: pytest.CollectReport) -> None:
|
||||
if report.failed:
|
||||
self.collect_failed.append(report.nodeid)
|
||||
|
||||
def pytest_collection_finish(self, session: pytest.Session) -> None:
|
||||
self.collected = tuple(item.nodeid for item in session.items)
|
||||
|
||||
def pytest_runtest_logreport(self, report: pytest.TestReport) -> None:
|
||||
self.reports.setdefault(report.nodeid, []).append((report.when, report.outcome, hasattr(report, "wasxfail")))
|
||||
|
||||
|
||||
def cmd_pytest(args: _Args) -> int:
|
||||
try:
|
||||
cases: Final = _load_manifest(Path(args.manifest))
|
||||
except (OSError, ValueError, json.JSONDecodeError) as exc:
|
||||
fail(f"manifest invalid: {exc}")
|
||||
if tuple(cases) != EXPECTED_CASES:
|
||||
fail(f"manifest case ids must be exactly {list(EXPECTED_CASES)} in order, got {list(cases)}")
|
||||
node_ids: Final = tuple(cases.values())
|
||||
if len(set(node_ids)) != len(node_ids):
|
||||
fail("manifest node ids are not unique")
|
||||
argv: Final = [
|
||||
*node_ids,
|
||||
"-p",
|
||||
"no:cacheprovider",
|
||||
"-p",
|
||||
"no:xdist",
|
||||
"-p",
|
||||
"no:rerunfailures",
|
||||
"-p",
|
||||
"no:randomly",
|
||||
"-rA",
|
||||
"-q",
|
||||
*(["--rootdir", args.rootdir] if args.rootdir else []),
|
||||
]
|
||||
|
||||
recorder: Final = _Recorder()
|
||||
code: Final = pytest.main(argv, plugins=[recorder])
|
||||
name_of: Final = MappingProxyType({node_id: case_id for case_id, node_id in cases.items()})
|
||||
problems: Final[list[str]] = []
|
||||
if code != 0:
|
||||
problems.append(f"pytest exit code {code}")
|
||||
for failed_id in recorder.collect_failed:
|
||||
problems.append(f"collection failed: {name_of.get(failed_id, failed_id)}")
|
||||
expected: Final = Counter(node_ids)
|
||||
collected: Final = Counter(recorder.collected)
|
||||
for node_id in expected - collected:
|
||||
problems.append(f"missing case {name_of[node_id]} ({node_id})")
|
||||
for node_id in collected - expected:
|
||||
problems.append(f"unexpected test collected: {node_id}")
|
||||
for node_id, count in collected.items():
|
||||
if count > 1:
|
||||
problems.append(f"duplicated test id: {node_id}")
|
||||
if len(recorder.collected) != len(EXPECTED_CASES):
|
||||
problems.append(f"collected {len(recorder.collected)} tests, expected {len(EXPECTED_CASES)}")
|
||||
rows: Final[list[tuple[str, bool]]] = []
|
||||
for case_id, node_id in cases.items():
|
||||
reports = recorder.reports.get(node_id, [])
|
||||
case_ok = (
|
||||
bool(reports)
|
||||
and all(outcome == "passed" and not wasxfail for _, outcome, wasxfail in reports)
|
||||
and {when for when, _, _ in reports} >= {"setup", "call", "teardown"}
|
||||
)
|
||||
rows.append((case_id, case_ok))
|
||||
if not reports:
|
||||
problems.append(f"{case_id} ({node_id}) produced no runtest reports")
|
||||
continue
|
||||
for when, outcome, wasxfail in reports:
|
||||
if outcome != "passed":
|
||||
problems.append(f"{case_id} ({node_id}) {when} outcome={outcome}")
|
||||
if wasxfail:
|
||||
problems.append(f"{case_id} ({node_id}) {when} was xfail/xpass")
|
||||
missing_phases = {"setup", "call", "teardown"} - {when for when, _, _ in reports}
|
||||
for phase in sorted(missing_phases):
|
||||
problems.append(f"{case_id} ({node_id}) missing {phase} report")
|
||||
for case_id, passed in rows:
|
||||
print(f"{case_id} {'PASS' if passed else 'FAIL'} {cases[case_id]}")
|
||||
if problems:
|
||||
for problem in problems:
|
||||
print(f"merge-smoke: {problem}", file=sys.stderr)
|
||||
fail("pytest verdict failed")
|
||||
ok("pytest 11 cases")
|
||||
return 0
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser: Final = argparse.ArgumentParser(description=__doc__)
|
||||
subs: Final = parser.add_subparsers(dest="command", required=True)
|
||||
p_iso: Final = subs.add_parser("verify-isolation")
|
||||
p_iso.add_argument("--no-child", action="store_true")
|
||||
p_interp: Final = subs.add_parser("interpreter")
|
||||
p_interp.add_argument("--expect", required=True)
|
||||
p_cli: Final = subs.add_parser("cli")
|
||||
p_cli.add_argument("--litellm-bin", default=None)
|
||||
p_cli.add_argument("--lite-bin", default=None)
|
||||
p_proxy: Final = subs.add_parser("proxy-startup")
|
||||
p_proxy.add_argument("--diagnostics-dir", required=True)
|
||||
p_proxy.add_argument("--litellm-bin", default=None)
|
||||
p_proxy.add_argument("--ready-deadline", type=float, default=120)
|
||||
p_proxy.add_argument("--shutdown-deadline", type=float, default=20)
|
||||
p_proxy.add_argument("--poll-interval", type=float, default=0.5)
|
||||
p_test: Final = subs.add_parser("pytest")
|
||||
p_test.add_argument("--manifest", required=True)
|
||||
p_test.add_argument("--rootdir", default=None)
|
||||
args: Final = parser.parse_args(namespace=_Args())
|
||||
handlers: Final = {
|
||||
"verify-isolation": cmd_verify_isolation,
|
||||
"interpreter": cmd_interpreter,
|
||||
"cli": cmd_cli,
|
||||
"proxy-startup": cmd_proxy_startup,
|
||||
"pytest": cmd_pytest,
|
||||
}
|
||||
return handlers[args.command](args)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
2
.github/scripts/verify_linux_native_wheel.py
vendored
2
.github/scripts/verify_linux_native_wheel.py
vendored
|
|
@ -214,7 +214,7 @@ def main(
|
|||
native_module: Final = load_native_module(native_path)
|
||||
native_module_loads: Final = native_module is not None
|
||||
panic_test_hook_absent: Final = native_module is not None and not hasattr(native_module, "_panic_for_test")
|
||||
native_size_limit: Final = 35_000_000
|
||||
native_size_limit: Final = 40_000_000
|
||||
native_size_within_limit: Final = native_member.file_size <= native_size_limit
|
||||
validations: Final = (
|
||||
(f"Python tag is {EXPECTED_PYTHON_TAG}", python_tag == EXPECTED_PYTHON_TAG),
|
||||
|
|
|
|||
32
.github/workflows/_test-unit-base.yml
vendored
32
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -4,7 +4,13 @@ on:
|
|||
workflow_call:
|
||||
inputs:
|
||||
test-path:
|
||||
description: "Pytest path(s) to run"
|
||||
description: >-
|
||||
Space-separated pytest paths to run. A path that no longer exists is
|
||||
dropped with a warning instead of being passed to pytest, because one
|
||||
missing path makes pytest-xdist collect nothing and report exit 5, which
|
||||
the step treats as a drained shard. Options are passed through as
|
||||
written, so use the `--flag=value` form: a bare `--ignore path` would
|
||||
have its path existence-checked like any other token.
|
||||
required: true
|
||||
type: string
|
||||
workers:
|
||||
|
|
@ -165,14 +171,22 @@ jobs:
|
|||
DIST: ${{ inputs.dist }}
|
||||
COVERAGE_CORE: sysmon
|
||||
run: |
|
||||
found_path=false
|
||||
for path in ${TEST_PATH}; do
|
||||
if [ -e "${path%%::*}" ]; then
|
||||
found_path=true
|
||||
break
|
||||
fi
|
||||
pytest_args=()
|
||||
existing_paths=0
|
||||
for token in ${TEST_PATH:?}; do
|
||||
case "${token}" in
|
||||
-*) pytest_args+=("${token}") ;;
|
||||
*)
|
||||
if [ -e "${token%%::*}" ]; then
|
||||
pytest_args+=("${token}")
|
||||
existing_paths=$((existing_paths + 1))
|
||||
else
|
||||
echo "::warning::${token} does not exist; drop it from this shard's test-path"
|
||||
fi
|
||||
;;
|
||||
esac
|
||||
done
|
||||
if [ "$found_path" = false ]; then
|
||||
if [ "${existing_paths}" -eq 0 ]; then
|
||||
echo "No path in TEST_PATH exists (${TEST_PATH}); nothing to run"
|
||||
exit 0
|
||||
fi
|
||||
|
|
@ -181,7 +195,7 @@ jobs:
|
|||
xdist_args=(-n "${WORKERS}" --dist="${DIST}")
|
||||
fi
|
||||
set +e
|
||||
uv run --no-sync pytest ${TEST_PATH:?} \
|
||||
uv run --no-sync pytest "${pytest_args[@]}" \
|
||||
--tb=short -vv \
|
||||
--maxfail="${MAX_FAILURES}" \
|
||||
"${xdist_args[@]}" \
|
||||
|
|
|
|||
2
.github/workflows/check-ui-api-types.yml
vendored
2
.github/workflows/check-ui-api-types.yml
vendored
|
|
@ -4,8 +4,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
|
||||
permissions:
|
||||
|
|
|
|||
3
.github/workflows/ci-coverage.yml
vendored
3
.github/workflows/ci-coverage.yml
vendored
|
|
@ -4,13 +4,10 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
|
|||
8
.github/workflows/codeql.yml
vendored
8
.github/workflows/codeql.yml
vendored
|
|
@ -67,12 +67,18 @@ jobs:
|
|||
# further up the stack are modified. The suppression is scoped to this one
|
||||
# file/rule pair via SARIF post-filtering so every other callsite of
|
||||
# py/weak-sensitive-data-hashing in the repository continues to be analyzed.
|
||||
- name: Filter SARIF (OCI sha256)
|
||||
# The same query fires on the HIBP k-anonymity lookup in
|
||||
# litellm/proxy/auth/password_policy.py, where the password's SHA-1 is only
|
||||
# a lookup key into the haveibeenpwned range API (the protocol mandates
|
||||
# SHA-1) and the digest itself never leaves the proxy beyond its first 5
|
||||
# characters.
|
||||
- name: Filter SARIF (OCI sha256, HIBP sha1)
|
||||
if: matrix.language == 'python'
|
||||
uses: advanced-security/filter-sarif@2da736ff05ef065cb2894ac6892e47b5eac2c3c0 # v1.1
|
||||
with:
|
||||
patterns: |
|
||||
-litellm/llms/oci/common_utils.py:py/weak-sensitive-data-hashing
|
||||
-litellm/proxy/auth/password_policy.py:py/weak-sensitive-data-hashing
|
||||
input: sarif-results/python.sarif
|
||||
output: sarif-results/python.sarif
|
||||
|
||||
|
|
|
|||
2
.github/workflows/codspeed.yml
vendored
2
.github/workflows/codspeed.yml
vendored
|
|
@ -4,7 +4,6 @@ on:
|
|||
push:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
paths:
|
||||
- "litellm/**"
|
||||
- "tests/benchmarks/**"
|
||||
|
|
@ -17,7 +16,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
paths:
|
||||
- "litellm/**"
|
||||
- "tests/benchmarks/**"
|
||||
|
|
|
|||
33
.github/workflows/compat-matrix-image.yml
vendored
Normal file
33
.github/workflows/compat-matrix-image.yml
vendored
Normal file
|
|
@ -0,0 +1,33 @@
|
|||
name: Compat Matrix Image
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
paths:
|
||||
- tests/e2e/claude_code/cron_vm/**
|
||||
- .github/workflows/compat-matrix-image.yml
|
||||
workflow_dispatch:
|
||||
|
||||
permissions: {}
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
compat-matrix-image:
|
||||
name: compat-matrix-image
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 15
|
||||
permissions:
|
||||
contents: read
|
||||
steps:
|
||||
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Build the Render cron image
|
||||
run: docker build -f tests/e2e/claude_code/cron_vm/Dockerfile -t compat-matrix:${{ github.sha }} tests/e2e
|
||||
|
||||
- name: Run the pinned binaries as the cron user
|
||||
run: |
|
||||
docker run --rm compat-matrix:${{ github.sha }} bash -c 'set -e; whoami; claude --version; gh --version; uv --version'
|
||||
2
.github/workflows/cost-map-guard.yml
vendored
2
.github/workflows/cost-map-guard.yml
vendored
|
|
@ -4,8 +4,6 @@ on: # zizmor: ignore[dangerous-triggers] runs the base branch's code only; the P
|
|||
pull_request_target:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
|
||||
permissions:
|
||||
|
|
|
|||
186
.github/workflows/create-release.yml
vendored
186
.github/workflows/create-release.yml
vendored
|
|
@ -1,186 +0,0 @@
|
|||
name: Create Release
|
||||
|
||||
on:
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
tag:
|
||||
description: "Release tag (e.g. 1.84.0, 1.84.0rc1, 1.84.0.dev42, 1.84.0-dev.2, 1.84.0.post1; legacy v1.83.10-stable still accepted)"
|
||||
required: true
|
||||
type: string
|
||||
commit_hash:
|
||||
description: "Full 40-char commit SHA to target"
|
||||
required: true
|
||||
type: string
|
||||
|
||||
permissions: {}
|
||||
|
||||
jobs:
|
||||
release:
|
||||
name: Create Release
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: write
|
||||
steps:
|
||||
- name: Validate inputs
|
||||
env:
|
||||
TAG: ${{ inputs.tag }}
|
||||
COMMIT_HASH: ${{ inputs.commit_hash }}
|
||||
run: |
|
||||
if ! echo "${COMMIT_HASH}" | grep -qE '^[0-9a-f]{40}$'; then
|
||||
echo "::error::commit_hash must be a full 40-character commit SHA"
|
||||
exit 1
|
||||
fi
|
||||
if ! echo "${TAG}" | grep -qE '^v?[0-9]+\.[0-9]+\.[0-9]+'; then
|
||||
echo "::error::tag must start with X.Y.Z (optional leading v), e.g. 1.84.0, 1.84.0rc1, 1.84.0.dev42, or v1.83.10-stable"
|
||||
exit 1
|
||||
fi
|
||||
|
||||
- name: Create release
|
||||
env:
|
||||
TAG: ${{ inputs.tag }}
|
||||
COMMIT_HASH: ${{ inputs.commit_hash }}
|
||||
uses: actions/github-script@60a0d83039c74a4aee543508d2ffcb1c3799cdea # v7.0.1
|
||||
with:
|
||||
script: |
|
||||
const tag = process.env.TAG;
|
||||
const commitHash = process.env.COMMIT_HASH;
|
||||
|
||||
// Mark RC / dev / nightly / alpha / beta tags as GitHub pre-releases.
|
||||
// Accept both PEP 440 (`.dev`) and SemVer (`-dev`) separators so tags
|
||||
// like `1.84.0.dev2` and `1.84.0-dev.2` are both detected.
|
||||
// PEP 440 post-releases (e.g. `1.84.0.post1`) and legacy `-stable[.patch.N]`
|
||||
// are stable maintenance releases, not pre-releases.
|
||||
const isPrerelease = /(?:rc|nightly|alpha|beta|[-.]dev)/i.test(tag);
|
||||
|
||||
// A stable release should only claim the repo "latest" badge when its
|
||||
// version is >= the current latest. Otherwise a backport (e.g. 1.84.6)
|
||||
// would steal "latest" from a newer line (e.g. 1.88.1).
|
||||
const versionKey = (rawTag) => {
|
||||
const m = String(rawTag).match(/^v?(\d+)\.(\d+)\.(\d+)/);
|
||||
if (!m) return null;
|
||||
const maintenance = String(rawTag).match(/(?:\.post|\.patch\.)(\d+)/i);
|
||||
return [Number(m[1]), Number(m[2]), Number(m[3]), maintenance ? Number(maintenance[1]) : 0];
|
||||
};
|
||||
const isAtLeast = (a, b) => {
|
||||
for (let i = 0; i < a.length; i++) {
|
||||
if (a[i] !== b[i]) return a[i] > b[i];
|
||||
}
|
||||
return true;
|
||||
};
|
||||
|
||||
const cosignSection = [
|
||||
`## Verify Docker Image Signature`,
|
||||
``,
|
||||
`All LiteLLM Docker images are signed with [cosign](https://docs.sigstore.dev/cosign/overview/). Every release is signed with the same key introduced in [commit \`0112e53\`](https://github.com/BerriAI/litellm/commit/0112e53046018d726492c814b3644b7d376029d0).`,
|
||||
``,
|
||||
`**Verify using the pinned commit hash (recommended):**`,
|
||||
``,
|
||||
`A commit hash is cryptographically immutable, so this is the strongest way to ensure you are using the original signing key:`,
|
||||
``,
|
||||
'```bash',
|
||||
`cosign verify \\`,
|
||||
` --key https://raw.githubusercontent.com/BerriAI/litellm/0112e53046018d726492c814b3644b7d376029d0/cosign.pub \\`,
|
||||
` ghcr.io/berriai/litellm:${tag}`,
|
||||
'```',
|
||||
``,
|
||||
`**Verify using the release tag (convenience):**`,
|
||||
``,
|
||||
`Tags are protected in this repository and resolve to the same key. This option is easier to read but relies on tag protection rules:`,
|
||||
``,
|
||||
'```bash',
|
||||
`cosign verify \\`,
|
||||
` --key https://raw.githubusercontent.com/BerriAI/litellm/${tag}/cosign.pub \\`,
|
||||
` ghcr.io/berriai/litellm:${tag}`,
|
||||
'```',
|
||||
``,
|
||||
`Expected output:`,
|
||||
``,
|
||||
'```',
|
||||
`The following checks were performed on each of these signatures:`,
|
||||
` - The cosign claims were validated`,
|
||||
` - The signatures were verified against the specified public key`,
|
||||
'```',
|
||||
``,
|
||||
`---`,
|
||||
``,
|
||||
].join('\n');
|
||||
|
||||
try {
|
||||
let makeLatest = "false";
|
||||
const newVersion = versionKey(tag);
|
||||
if (!isPrerelease && newVersion) {
|
||||
let latestVersion = null;
|
||||
try {
|
||||
const latest = await github.rest.repos.getLatestRelease({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
});
|
||||
latestVersion = versionKey(latest.data.tag_name);
|
||||
} catch (error) {
|
||||
if (error.status !== 404) throw error;
|
||||
}
|
||||
makeLatest = (!latestVersion || isAtLeast(newVersion, latestVersion)) ? "true" : "false";
|
||||
}
|
||||
|
||||
try {
|
||||
await github.rest.git.createRef({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
ref: `refs/tags/${tag}`,
|
||||
sha: commitHash,
|
||||
});
|
||||
} catch (error) {
|
||||
if (error.status !== 422) throw error;
|
||||
const existing = await github.rest.git.getRef({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
ref: `tags/${tag}`,
|
||||
});
|
||||
if (existing.data.object.sha !== commitHash) {
|
||||
throw new Error(`Tag ${tag} already exists at ${existing.data.object.sha}, expected ${commitHash}`);
|
||||
}
|
||||
}
|
||||
|
||||
const response = await github.rest.repos.createRelease({
|
||||
draft: true,
|
||||
generate_release_notes: true,
|
||||
name: tag,
|
||||
owner: context.repo.owner,
|
||||
prerelease: isPrerelease,
|
||||
repo: context.repo.repo,
|
||||
tag_name: tag,
|
||||
});
|
||||
|
||||
const updatedBody = cosignSection + (response.data.body ?? '');
|
||||
await github.rest.repos.updateRelease({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
release_id: response.data.id,
|
||||
tag_name: tag,
|
||||
body: updatedBody,
|
||||
draft: false,
|
||||
});
|
||||
|
||||
if (!isPrerelease) {
|
||||
await github.rest.repos.updateRelease({
|
||||
owner: context.repo.owner,
|
||||
repo: context.repo.repo,
|
||||
release_id: response.data.id,
|
||||
tag_name: tag,
|
||||
make_latest: makeLatest,
|
||||
});
|
||||
}
|
||||
|
||||
} catch (error) {
|
||||
core.setFailed(error.message);
|
||||
}
|
||||
|
||||
create-branch:
|
||||
name: Create Release Branch
|
||||
needs: release
|
||||
permissions:
|
||||
contents: write
|
||||
uses: ./.github/workflows/create-release-branch.yml
|
||||
with:
|
||||
tag: ${{ inputs.tag }}
|
||||
commit_hash: ${{ inputs.commit_hash }}
|
||||
|
|
@ -1,49 +0,0 @@
|
|||
name: Create Daily Staging Branch
|
||||
|
||||
on:
|
||||
schedule:
|
||||
- cron: "0 0,12 * * *" # Runs every 12 hours at midnight and noon UTC
|
||||
workflow_dispatch: # Allow manual trigger
|
||||
|
||||
jobs:
|
||||
create-staging-branch:
|
||||
if: github.repository == 'BerriAI/litellm'
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: write
|
||||
|
||||
steps:
|
||||
- name: Create daily staging branch
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
run: |
|
||||
BRANCH_NAME="litellm_oss_staging_$(date +'%m_%d_%Y')"
|
||||
echo "Creating branch: $BRANCH_NAME"
|
||||
if gh api "repos/${{ github.repository }}/git/ref/heads/$BRANCH_NAME" --silent 2>/dev/null; then
|
||||
echo "Branch $BRANCH_NAME already exists. Skipping creation."
|
||||
exit 0
|
||||
fi
|
||||
MAIN_SHA=$(gh api "repos/${{ github.repository }}/git/ref/heads/main" --jq '.object.sha')
|
||||
gh api "repos/${{ github.repository }}/git/refs" -f ref="refs/heads/$BRANCH_NAME" -f sha="$MAIN_SHA" --silent
|
||||
echo "Successfully created branch: $BRANCH_NAME at $MAIN_SHA"
|
||||
|
||||
create-internal-dev-branch:
|
||||
if: github.repository == 'BerriAI/litellm'
|
||||
runs-on: ubuntu-latest
|
||||
permissions:
|
||||
contents: write
|
||||
|
||||
steps:
|
||||
- name: Create internal dev branch
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
run: |
|
||||
BRANCH_NAME="litellm_internal_dev_$(date +'%m_%d_%Y')"
|
||||
echo "Creating branch: $BRANCH_NAME"
|
||||
if gh api "repos/${{ github.repository }}/git/ref/heads/$BRANCH_NAME" --silent 2>/dev/null; then
|
||||
echo "Branch $BRANCH_NAME already exists. Skipping creation."
|
||||
exit 0
|
||||
fi
|
||||
MAIN_SHA=$(gh api "repos/${{ github.repository }}/git/ref/heads/main" --jq '.object.sha')
|
||||
gh api "repos/${{ github.repository }}/git/refs" -f ref="refs/heads/$BRANCH_NAME" -f sha="$MAIN_SHA" --silent
|
||||
echo "Successfully created branch: $BRANCH_NAME at $MAIN_SHA"
|
||||
|
|
@ -4,8 +4,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths:
|
||||
- "uv.lock"
|
||||
|
|
|
|||
1
.github/workflows/image-scan.yml
vendored
1
.github/workflows/image-scan.yml
vendored
|
|
@ -4,7 +4,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_branch
|
||||
- "litellm_**"
|
||||
paths:
|
||||
|
|
|
|||
20
.github/workflows/issue_fixed_comment.yml
vendored
20
.github/workflows/issue_fixed_comment.yml
vendored
|
|
@ -6,8 +6,12 @@ on:
|
|||
workflow_dispatch:
|
||||
inputs:
|
||||
issue_number:
|
||||
description: "Closed issue number to comment on manually."
|
||||
required: true
|
||||
description: "Closed issue number to comment on and close the superseded pull requests of. Ignored by a sweep."
|
||||
required: false
|
||||
sweep:
|
||||
description: "Close every open pull request whose linked issues were all fixed on the default branch. Reads every open pull request, so run it at most once an hour."
|
||||
type: boolean
|
||||
default: false
|
||||
pull_request:
|
||||
paths:
|
||||
- .github/workflows/issue_fixed_comment.yml
|
||||
|
|
@ -39,16 +43,17 @@ jobs:
|
|||
with:
|
||||
bun-version: "1.4.0"
|
||||
|
||||
- name: Test the closer lookup, the release placement and the comment
|
||||
- name: Test the closer lookup, the release placement, the comment and the superseded pull request close
|
||||
run: bun test scripts/comment-fixed-issue.test.ts
|
||||
|
||||
comment-fixed-issue:
|
||||
if: github.event_name != 'pull_request' && github.repository == 'BerriAI/litellm'
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 5
|
||||
timeout-minutes: 15
|
||||
permissions:
|
||||
contents: read
|
||||
issues: write
|
||||
pull-requests: write
|
||||
steps:
|
||||
- name: Checkout scripts
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
|
|
@ -59,13 +64,16 @@ jobs:
|
|||
- name: Setup Bun
|
||||
uses: oven-sh/setup-bun@0c5077e51419868618aeaa5fe8019c62421857d6 # v2.2.0
|
||||
with:
|
||||
# Exact version, never latest: the next step holds an issues: write token
|
||||
# Exact version, never latest: the next step holds issues: write and pull-requests: write tokens
|
||||
bun-version: "1.4.0"
|
||||
|
||||
- name: Name the release that carries the fix
|
||||
- name: Name the release that carries the fix and close the pull requests it supersedes
|
||||
shell: bash
|
||||
run: bun run scripts/comment-fixed-issue.ts | tee -a "${GITHUB_STEP_SUMMARY}"
|
||||
env:
|
||||
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
ISSUE_NUMBER: ${{ github.event.issue.number || github.event.inputs.issue_number }}
|
||||
SWEEP: ${{ github.event.inputs.sweep }}
|
||||
DEFAULT_BRANCH: ${{ github.event.repository.default_branch }}
|
||||
DRY_RUN: ${{ vars.ISSUE_FIXED_COMMENT_ENABLED != 'true' }}
|
||||
CLOSE_PRS_DRY_RUN: ${{ vars.ISSUE_FIXED_CLOSE_PRS_ENABLED != 'true' }}
|
||||
|
|
|
|||
2
.github/workflows/osv-scan.yml
vendored
2
.github/workflows/osv-scan.yml
vendored
|
|
@ -4,8 +4,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
schedule:
|
||||
- cron: "23 6 * * *"
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
name: Publish basedpyright base counts
|
||||
|
||||
# Every commit on main or litellm_internal_staging can become a future merge-base.
|
||||
# Every commit on main can become a future merge-base.
|
||||
# Publishing its per-rule basedpyright counts as an artifact lets
|
||||
# scripts/type_check_gate.py download them in seconds instead of paying a
|
||||
# 60-110s second basedpyright pass on every fresh worktree or moved merge-base.
|
||||
|
|
@ -11,7 +11,6 @@ on:
|
|||
push:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
ref:
|
||||
|
|
|
|||
6
.github/workflows/test-code-quality.yml
vendored
6
.github/workflows/test-code-quality.yml
vendored
|
|
@ -4,13 +4,10 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
@ -83,6 +80,9 @@ jobs:
|
|||
- name: test_e2e_changed_gate
|
||||
run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_e2e_changed_gate.py tests/code_coverage_tests/test_e2e_idp_stack.py
|
||||
|
||||
- name: Check merge smoke harness
|
||||
run: uv run --no-sync pytest -q --noconftest -p no:cacheprovider -c /dev/null tests/code_coverage_tests/test_merge_smoke.py
|
||||
|
||||
- name: router_code_coverage
|
||||
run: uv run --no-sync python ./tests/code_coverage_tests/router_code_coverage.py
|
||||
|
||||
|
|
|
|||
2
.github/workflows/test-linting.yml
vendored
2
.github/workflows/test-linting.yml
vendored
|
|
@ -4,8 +4,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
|
||||
permissions:
|
||||
|
|
|
|||
2
.github/workflows/test-litellm-ui-build.yml
vendored
2
.github/workflows/test-litellm-ui-build.yml
vendored
|
|
@ -7,8 +7,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
|
||||
concurrency:
|
||||
|
|
|
|||
2
.github/workflows/test-litellm-ui-lint.yml
vendored
2
.github/workflows/test-litellm-ui-lint.yml
vendored
|
|
@ -6,8 +6,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
|
||||
concurrency:
|
||||
|
|
|
|||
3
.github/workflows/test-litellm-ui-unit.yml
vendored
3
.github/workflows/test-litellm-ui-unit.yml
vendored
|
|
@ -7,13 +7,10 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
|
|
|
|||
|
|
@ -4,8 +4,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
|
||||
permissions:
|
||||
|
|
|
|||
95
.github/workflows/test-merge-smoke.yml
vendored
Normal file
95
.github/workflows/test-merge-smoke.yml
vendored
Normal file
|
|
@ -0,0 +1,95 @@
|
|||
name: Merge smoke checks
|
||||
|
||||
on:
|
||||
pull_request:
|
||||
branches: [main, litellm_internal_staging, litellm_oss_staging, "litellm_**"]
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
||||
concurrency:
|
||||
group: merge-smoke-${{ github.event.pull_request.number || github.run_id }}
|
||||
cancel-in-progress: true
|
||||
|
||||
jobs:
|
||||
dashboard-build:
|
||||
name: Dashboard build
|
||||
runs-on: ubuntu-24.04
|
||||
timeout-minutes: 30
|
||||
steps:
|
||||
- name: Checkout
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Build the dashboard stage
|
||||
run: docker build --target ui-builder -f Dockerfile .
|
||||
|
||||
core-checks:
|
||||
name: Core checks (Python ${{ matrix.python-version }})
|
||||
runs-on: ubuntu-24.04
|
||||
timeout-minutes: 30
|
||||
strategy:
|
||||
fail-fast: false
|
||||
matrix:
|
||||
python-version: ["3.10", "3.11", "3.12", "3.13", "3.14"]
|
||||
env:
|
||||
LITELLM_LOCAL_MODEL_COST_MAP: "True"
|
||||
steps:
|
||||
- name: Checkout
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
persist-credentials: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: ${{ matrix.python-version }}
|
||||
|
||||
- name: Set up uv
|
||||
uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
- name: Install dependencies
|
||||
run: .github/scripts/uv_sync_with_retries.sh --frozen --extra proxy --extra cli --group dev --group proxy-dev --python ${{ matrix.python-version }}
|
||||
|
||||
- name: Create the loopback-only network namespace
|
||||
run: |
|
||||
sudo ip netns add smoke
|
||||
sudo ip netns exec smoke ip link set lo up
|
||||
cat > "${RUNNER_TEMP}/in-netns" <<'WRAP'
|
||||
#!/usr/bin/env bash
|
||||
set -euo pipefail
|
||||
exec sudo --preserve-env=LITELLM_LOCAL_MODEL_COST_MAP ip netns exec smoke setpriv --reuid "$(id -u)" --regid "$(id -g)" --init-groups -- env HOME="${HOME}" PATH="${PATH}" "$@"
|
||||
WRAP
|
||||
chmod +x "${RUNNER_TEMP}/in-netns"
|
||||
echo "IN_NETNS=${RUNNER_TEMP}/in-netns" >> "${GITHUB_ENV}"
|
||||
|
||||
- name: Verify namespace isolation
|
||||
run: $IN_NETNS .venv/bin/python .github/scripts/run_merge_smoke.py verify-isolation
|
||||
|
||||
- name: Verify interpreter version
|
||||
run: $IN_NETNS .venv/bin/python .github/scripts/run_merge_smoke.py interpreter --expect ${{ matrix.python-version }}
|
||||
|
||||
- name: Import and CLI checks
|
||||
run: $IN_NETNS .venv/bin/python .github/scripts/run_merge_smoke.py cli
|
||||
|
||||
- name: Proxy startup check
|
||||
run: $IN_NETNS .venv/bin/python .github/scripts/run_merge_smoke.py proxy-startup --diagnostics-dir "${RUNNER_TEMP}/smoke-diagnostics"
|
||||
|
||||
- name: Run curated smoke cases
|
||||
run: $IN_NETNS .venv/bin/python .github/scripts/run_merge_smoke.py pytest --manifest .github/merge-smoke-tests.json
|
||||
|
||||
- name: Upload smoke diagnostics
|
||||
if: always()
|
||||
uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1
|
||||
with:
|
||||
name: merge-smoke-diagnostics-py${{ matrix.python-version }}
|
||||
path: ${{ runner.temp }}/smoke-diagnostics
|
||||
if-no-files-found: ignore
|
||||
|
||||
- name: Remove the network namespace
|
||||
if: always()
|
||||
run: sudo ip netns delete smoke
|
||||
3
.github/workflows/test-postgres.yml
vendored
3
.github/workflows/test-postgres.yml
vendored
|
|
@ -4,13 +4,10 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
|
|
|
|||
2
.github/workflows/test-redis-compat.yml
vendored
2
.github/workflows/test-redis-compat.yml
vendored
|
|
@ -4,8 +4,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths:
|
||||
- "litellm/_redis.py"
|
||||
|
|
|
|||
12
.github/workflows/test-rust.yml
vendored
12
.github/workflows/test-rust.yml
vendored
|
|
@ -29,8 +29,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths:
|
||||
- "litellm-rust/**"
|
||||
|
|
@ -105,6 +103,16 @@ jobs:
|
|||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- uses: ./.github/actions/setup-uv-with-retries
|
||||
with:
|
||||
version: "0.10.9"
|
||||
|
||||
- name: Install Python dependencies for the bridge tests
|
||||
working-directory: .
|
||||
run: |
|
||||
uv sync --frozen --no-install-project
|
||||
echo "PYTHONPATH=$PWD/.venv/lib/$(ls .venv/lib)/site-packages" >> "$GITHUB_ENV"
|
||||
|
||||
- run: rustup toolchain install --no-self-update
|
||||
|
||||
- uses: taiki-e/install-action@d438492cf8a250514fa2d34b30bc3c0dc37c65ff # v2.87.8
|
||||
|
|
|
|||
2
.github/workflows/test-semgrep.yml
vendored
2
.github/workflows/test-semgrep.yml
vendored
|
|
@ -4,8 +4,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
|
||||
permissions:
|
||||
|
|
|
|||
2
.github/workflows/test-terraform-modules.yml
vendored
2
.github/workflows/test-terraform-modules.yml
vendored
|
|
@ -9,8 +9,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths:
|
||||
- "terraform/litellm/aws/**"
|
||||
|
|
|
|||
|
|
@ -8,8 +8,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths:
|
||||
- "terraform/provider/**"
|
||||
|
|
|
|||
|
|
@ -4,13 +4,10 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
|
|||
3
.github/workflows/test-unit-proxy-db.yml
vendored
3
.github/workflows/test-unit-proxy-db.yml
vendored
|
|
@ -4,13 +4,10 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
|
|
|
|||
11
.github/workflows/test-unit.yml
vendored
11
.github/workflows/test-unit.yml
vendored
|
|
@ -4,13 +4,10 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
push:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
workflow_dispatch:
|
||||
|
||||
permissions:
|
||||
|
|
@ -107,26 +104,18 @@ jobs:
|
|||
tests/test_litellm/batches
|
||||
tests/test_litellm/secret_managers
|
||||
tests/test_litellm/a2a_protocol
|
||||
tests/test_litellm/anthropic_interface
|
||||
tests/test_litellm/chat_completions
|
||||
tests/test_litellm/completion_extras
|
||||
tests/test_litellm/compression
|
||||
tests/test_litellm/containers
|
||||
tests/test_litellm/endpoints
|
||||
tests/test_litellm/models
|
||||
tests/test_litellm/repositories
|
||||
tests/test_litellm/images
|
||||
tests/test_litellm/interactions
|
||||
tests/test_litellm/messages
|
||||
tests/test_litellm/ocr
|
||||
tests/test_litellm/passthrough
|
||||
tests/test_litellm/rag
|
||||
tests/test_litellm/realtime_api
|
||||
tests/test_litellm/rerank_api
|
||||
tests/test_litellm/rust_bridge
|
||||
tests/test_litellm/sandbox
|
||||
tests/test_litellm/skills
|
||||
tests/test_litellm/test_router
|
||||
tests/test_litellm/vector_stores
|
||||
tests/test_litellm/videos
|
||||
tests/test_litellm/test_*.py
|
||||
|
|
|
|||
2
.github/workflows/test-vscode-extension.yml
vendored
2
.github/workflows/test-vscode-extension.yml
vendored
|
|
@ -6,8 +6,6 @@ on:
|
|||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
paths:
|
||||
- "vscode-extension/**"
|
||||
|
|
|
|||
4
.github/workflows/zizmor.yml
vendored
4
.github/workflows/zizmor.yml
vendored
|
|
@ -2,12 +2,10 @@ name: GitHub Actions Security Analysis
|
|||
|
||||
on:
|
||||
push:
|
||||
branches: [main, litellm_internal_staging]
|
||||
branches: [main]
|
||||
pull_request:
|
||||
branches:
|
||||
- main
|
||||
- litellm_internal_staging
|
||||
- litellm_oss_staging
|
||||
- "litellm_**"
|
||||
|
||||
concurrency:
|
||||
|
|
|
|||
|
|
@ -1,10 +1,10 @@
|
|||
# syntax=docker/dockerfile:1.7
|
||||
|
||||
# Base image for building
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
|
||||
|
||||
# Runtime image
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
|
||||
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
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
|
||||
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
|
||||
|
||||
FROM $UV_IMAGE AS uvbin
|
||||
|
|
|
|||
|
|
@ -26,6 +26,7 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
|
|||
"/v2/login",
|
||||
"/v3/login",
|
||||
"/logout",
|
||||
"/session/logout",
|
||||
"/token",
|
||||
"/onboarding/",
|
||||
"/audit",
|
||||
|
|
|
|||
|
|
@ -1,10 +1,10 @@
|
|||
# syntax=docker/dockerfile:1.7
|
||||
|
||||
# Base image for building
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
|
||||
|
||||
# Runtime image
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
|
||||
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
|
||||
|
|
|
|||
|
|
@ -1,8 +1,8 @@
|
|||
# syntax=docker/dockerfile:1.7
|
||||
|
||||
# Base images
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
|
||||
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.
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-enterprise"
|
||||
version = "0.1.69"
|
||||
version = "0.1.70"
|
||||
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.69"
|
||||
version = "0.1.70"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-enterprise==",
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:e624c5d5e42382ce7165ddafcbbf8e6769a24cbd02ea6114b880b05ae5ba2a8d
|
||||
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
|
||||
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
|
||||
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
|
||||
# Checksum from https://www.pgbouncer.org/downloads/ (the Wolfi repo only carries 1.24.x)
|
||||
ARG PGBOUNCER_VERSION=1.25.2
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[project]
|
||||
name = "litellm-proxy-extras"
|
||||
version = "0.4.100"
|
||||
version = "0.4.101"
|
||||
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.9"
|
||||
|
|
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
|
|||
module-root = ""
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.4.100"
|
||||
version = "0.4.101"
|
||||
version_files = [
|
||||
"pyproject.toml:^version",
|
||||
"../pyproject.toml:litellm-proxy-extras==",
|
||||
|
|
|
|||
260
litellm-rust/Cargo.lock
generated
260
litellm-rust/Cargo.lock
generated
|
|
@ -88,6 +88,45 @@ version = "1.2.0"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "03918c3dbd7701a85c6b9887732e2921175f26c350b4563841d0958c21d57e6d"
|
||||
|
||||
[[package]]
|
||||
name = "asn1-rs"
|
||||
version = "0.7.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b7f43a50ac4fdca5df8e885c21b835997f0a1cdee65494a6847694a98652d9d8"
|
||||
dependencies = [
|
||||
"asn1-rs-derive",
|
||||
"asn1-rs-impl",
|
||||
"displaydoc",
|
||||
"nom",
|
||||
"num-traits",
|
||||
"rusticata-macros",
|
||||
"thiserror 2.0.19",
|
||||
"time",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "asn1-rs-derive"
|
||||
version = "0.6.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "3109e49b1e4909e9db6515a30c633684d68cdeaa252f215214cb4fa1a5bfee2c"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.119",
|
||||
"synstructure 0.13.2",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "asn1-rs-impl"
|
||||
version = "0.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "7b18050c2cd6fe86c3a76584ef5e0baf286d038cda203eb6223df2cc413565f7"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.119",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "assert-json-diff"
|
||||
version = "2.0.2"
|
||||
|
|
@ -810,7 +849,7 @@ version = "0.8.0"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "08807e080ed7f9d5433fa9b275196cfc35414f66a0c79d864dc51a0d825231a3"
|
||||
dependencies = [
|
||||
"bit-vec",
|
||||
"bit-vec 0.8.0",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -819,6 +858,15 @@ version = "0.8.0"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "5e764a1d40d510daf35e07be9eb06e75770908c27d411ee6c92109c9840eaaf7"
|
||||
|
||||
[[package]]
|
||||
name = "bit-vec"
|
||||
version = "0.9.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b71798fca2c1fe1086445a7258a4bc81e6e49dcd24c8d0dd9a1e57395b603f51"
|
||||
dependencies = [
|
||||
"serde",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "bitflags"
|
||||
version = "1.3.2"
|
||||
|
|
@ -1120,6 +1168,15 @@ version = "0.4.0"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "338089f42c427b86394a5ee60ff321da23a5c89c9d89514c829687b26359fcff"
|
||||
|
||||
[[package]]
|
||||
name = "crc32c"
|
||||
version = "0.6.8"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "3a47af21622d091a8f0fb295b88bc886ac74efcc613efc19f5d0b21de5c89e47"
|
||||
dependencies = [
|
||||
"rustc_version",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "crc32fast"
|
||||
version = "1.5.1"
|
||||
|
|
@ -1378,6 +1435,20 @@ dependencies = [
|
|||
"thiserror 2.0.19",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "der-parser"
|
||||
version = "10.0.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "07da5016415d5a3c4dd39b11ed26f915f52fc4e0dc197d87908bc916e51bc1a6"
|
||||
dependencies = [
|
||||
"asn1-rs",
|
||||
"displaydoc",
|
||||
"nom",
|
||||
"num-bigint 0.4.8",
|
||||
"num-traits",
|
||||
"rusticata-macros",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "deranged"
|
||||
version = "0.5.8"
|
||||
|
|
@ -2589,21 +2660,6 @@ dependencies = [
|
|||
"wasm-bindgen",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "jsonwebtoken"
|
||||
version = "11.1.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e75fe14a82d81e5f5af639997db37d8b96045938a7ac6ab18cdbe1c7467e05e1"
|
||||
dependencies = [
|
||||
"base64 0.22.1",
|
||||
"getrandom 0.2.17",
|
||||
"js-sys",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"signature",
|
||||
"zeroize",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "lazy_static"
|
||||
version = "1.5.0"
|
||||
|
|
@ -2946,6 +3002,7 @@ name = "litellm-core-utils"
|
|||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"fancy-regex 0.19.2",
|
||||
"litellm-tracing",
|
||||
"litellm-types",
|
||||
"rstest",
|
||||
"serde",
|
||||
|
|
@ -2956,6 +3013,14 @@ dependencies = [
|
|||
"url",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-cost"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"criterion",
|
||||
"proptest",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-framing"
|
||||
version = "0.1.0"
|
||||
|
|
@ -3046,6 +3111,20 @@ dependencies = [
|
|||
"url",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-model-catalog"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"criterion",
|
||||
"indexmap 2.14.0",
|
||||
"litellm-model-catalog",
|
||||
"rstest",
|
||||
"schemars 1.2.2",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"thiserror 2.0.19",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-python-bridge"
|
||||
version = "0.1.0"
|
||||
|
|
@ -3053,6 +3132,7 @@ dependencies = [
|
|||
"aws-sdk-secretsmanager",
|
||||
"bytes",
|
||||
"criterion",
|
||||
"fancy-regex 0.19.2",
|
||||
"futures-util",
|
||||
"litellm-auth",
|
||||
"litellm-auth-aws",
|
||||
|
|
@ -3071,6 +3151,7 @@ dependencies = [
|
|||
"litellm-callbacks-legacy-python",
|
||||
"litellm-core",
|
||||
"litellm-core-utils",
|
||||
"litellm-host",
|
||||
"litellm-host-python",
|
||||
"litellm-http",
|
||||
"litellm-llms",
|
||||
|
|
@ -3078,6 +3159,7 @@ dependencies = [
|
|||
"litellm-secrets-aws",
|
||||
"litellm-secrets-types",
|
||||
"litellm-token-counter",
|
||||
"litellm-tracing",
|
||||
"litellm-types",
|
||||
"pyo3",
|
||||
"pyo3-async-runtimes",
|
||||
|
|
@ -3093,6 +3175,7 @@ dependencies = [
|
|||
"tokio",
|
||||
"tokio-tungstenite",
|
||||
"url",
|
||||
"veil",
|
||||
"wiremock",
|
||||
]
|
||||
|
||||
|
|
@ -3117,10 +3200,11 @@ version = "0.1.0"
|
|||
dependencies = [
|
||||
"aws-sdk-kms",
|
||||
"base64 0.22.1",
|
||||
"futures-util",
|
||||
"google-cloud-auth",
|
||||
"google-cloud-kms-v1",
|
||||
"jsonwebtoken",
|
||||
"litellm-core-utils",
|
||||
"litellm-python-compat",
|
||||
"litellm-secrets-aws",
|
||||
"litellm-secrets-azure",
|
||||
"litellm-secrets-cyberark",
|
||||
|
|
@ -3150,11 +3234,12 @@ dependencies = [
|
|||
"litellm-auth-aws",
|
||||
"litellm-core-utils",
|
||||
"litellm-secrets-types",
|
||||
"litellm-tracing",
|
||||
"rstest",
|
||||
"serde_json",
|
||||
"tempfile",
|
||||
"thiserror 2.0.19",
|
||||
"tokio",
|
||||
"tracing",
|
||||
"veil",
|
||||
"wiremock",
|
||||
]
|
||||
|
|
@ -3186,15 +3271,17 @@ dependencies = [
|
|||
"base64 0.22.1",
|
||||
"litellm-core-utils",
|
||||
"litellm-secrets-types",
|
||||
"litellm-tracing",
|
||||
"moka",
|
||||
"percent-encoding",
|
||||
"rcgen",
|
||||
"reqwest 0.12.28",
|
||||
"rstest",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"tempfile",
|
||||
"thiserror 2.0.19",
|
||||
"tokio",
|
||||
"tracing",
|
||||
"veil",
|
||||
"wiremock",
|
||||
]
|
||||
|
|
@ -3204,6 +3291,7 @@ name = "litellm-secrets-google"
|
|||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"base64 0.22.1",
|
||||
"crc32c",
|
||||
"google-cloud-auth",
|
||||
"google-cloud-gax",
|
||||
"google-cloud-kms-v1",
|
||||
|
|
@ -3229,7 +3317,6 @@ version = "0.1.0"
|
|||
dependencies = [
|
||||
"litellm-core-utils",
|
||||
"litellm-secrets-types",
|
||||
"moka",
|
||||
"rstest",
|
||||
"rustify",
|
||||
"rustify_derive",
|
||||
|
|
@ -3248,6 +3335,7 @@ name = "litellm-secrets-types"
|
|||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"litellm-auth-types",
|
||||
"moka",
|
||||
"rstest",
|
||||
"serde",
|
||||
"serde_json",
|
||||
|
|
@ -3309,6 +3397,19 @@ dependencies = [
|
|||
"tiktoken-rs",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-tracing"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"fancy-regex 0.19.2",
|
||||
"percent-encoding",
|
||||
"rstest",
|
||||
"serde_json",
|
||||
"tokio",
|
||||
"tracing",
|
||||
"tracing-subscriber",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-types"
|
||||
version = "0.1.0"
|
||||
|
|
@ -3549,6 +3650,15 @@ dependencies = [
|
|||
"libc",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "oid-registry"
|
||||
version = "0.8.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "12f40cff3dde1b6087cc5d5f5d4d65712f34016a03ed60e9c08dcc392736b5b7"
|
||||
dependencies = [
|
||||
"asn1-rs",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "once_cell"
|
||||
version = "1.21.4"
|
||||
|
|
@ -3682,6 +3792,16 @@ version = "0.2.3"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "2ee67f1008b1ba2321834326597b8e186293b049a023cdef258527550b9935b4"
|
||||
|
||||
[[package]]
|
||||
name = "pem"
|
||||
version = "4.0.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d354a98a3d1251555de99e8fdd8afda05573c31b82f59063a7b0a29b5527f120"
|
||||
dependencies = [
|
||||
"base64 0.23.1",
|
||||
"serde_core",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "percent-encoding"
|
||||
version = "2.3.2"
|
||||
|
|
@ -3860,7 +3980,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||
checksum = "4b45fcc2344c680f5025fe57779faef368840d0bd1f42f216291f0dc4ace4744"
|
||||
dependencies = [
|
||||
"bit-set",
|
||||
"bit-vec",
|
||||
"bit-vec 0.8.0",
|
||||
"bitflags 2.13.1",
|
||||
"num-traits",
|
||||
"rand 0.9.5",
|
||||
|
|
@ -4249,6 +4369,20 @@ dependencies = [
|
|||
"crossbeam-utils",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rcgen"
|
||||
version = "0.14.10"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "8774e05a7d0de114588e6a28fe7e71694b82614ed569d86d8b389dfbc98b8ad8"
|
||||
dependencies = [
|
||||
"pem",
|
||||
"ring",
|
||||
"rustls-pki-types",
|
||||
"time",
|
||||
"x509-parser",
|
||||
"yasna",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "redis"
|
||||
version = "1.7.0"
|
||||
|
|
@ -4531,6 +4665,15 @@ dependencies = [
|
|||
"semver",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rusticata-macros"
|
||||
version = "4.1.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "faf0c4a6ece9950b9abdb62b1cfcf2a68b3b67a10ba445b3bb85be2a293d0632"
|
||||
dependencies = [
|
||||
"nom",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "rustify"
|
||||
version = "0.7.0"
|
||||
|
|
@ -4748,10 +4891,23 @@ checksum = "687274d293b6cdc6e73e0fee520bf2049650090d7164f87672d212a3c530cf4a"
|
|||
dependencies = [
|
||||
"dyn-clone",
|
||||
"ref-cast",
|
||||
"schemars_derive",
|
||||
"serde",
|
||||
"serde_json",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "schemars_derive"
|
||||
version = "1.2.2"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d98c67716b46af2f0b8cf752abc930f6f9aecfbf671ecfb531db8a31dbe4e2ba"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"serde_derive_internals",
|
||||
"syn 3.0.0",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "scopeguard"
|
||||
version = "1.2.0"
|
||||
|
|
@ -4840,6 +4996,17 @@ dependencies = [
|
|||
"syn 3.0.0",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "serde_derive_internals"
|
||||
version = "0.30.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "f852137cce035d6a4df67ccce505ff6b3e9fd3a10e3e52b24dc71e650bb1a9bd"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 3.0.0",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "serde_json"
|
||||
version = "1.0.150"
|
||||
|
|
@ -4983,15 +5150,6 @@ dependencies = [
|
|||
"libc",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "signature"
|
||||
version = "2.2.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "77549399552de45a898a580c1b41d445bf730df867cc44e6c0233bbc4b8329de"
|
||||
dependencies = [
|
||||
"rand_core 0.6.4",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "simd-adler32"
|
||||
version = "0.3.10"
|
||||
|
|
@ -6293,6 +6451,24 @@ version = "0.6.3"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "1ffae5123b2d3fc086436f8834ae3ab053a283cfac8fe0a0b8eaae044768a4c4"
|
||||
|
||||
[[package]]
|
||||
name = "x509-parser"
|
||||
version = "0.18.1"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "d43b0f71ce057da06bc0851b23ee24f3f86190b07203dd8f567d0b706a185202"
|
||||
dependencies = [
|
||||
"asn1-rs",
|
||||
"data-encoding",
|
||||
"der-parser",
|
||||
"lazy_static",
|
||||
"nom",
|
||||
"oid-registry",
|
||||
"ring",
|
||||
"rusticata-macros",
|
||||
"thiserror 2.0.19",
|
||||
"time",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "xmlparser"
|
||||
version = "0.13.6"
|
||||
|
|
@ -6305,6 +6481,16 @@ version = "0.8.18"
|
|||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "aee1b19627c7c60102ab80d3a9cbe18de90bfe03bfa6c3715447681f0e8c8af6"
|
||||
|
||||
[[package]]
|
||||
name = "yasna"
|
||||
version = "0.6.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "b5f6765e852b9b4dc8e2a76843e4d64d1cea8e79bcde0b6901aea8e7c7f08282"
|
||||
dependencies = [
|
||||
"bit-vec 0.9.1",
|
||||
"time",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "yoke"
|
||||
version = "0.8.3"
|
||||
|
|
@ -6374,20 +6560,6 @@ name = "zeroize"
|
|||
version = "1.9.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "e13c156562582aa81c60cb29407084cdb54c4164760106ab78e6c5b0858cf64e"
|
||||
dependencies = [
|
||||
"zeroize_derive",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "zeroize_derive"
|
||||
version = "1.5.0"
|
||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||
checksum = "3c50655cbb0fe3fc43170059e702f1ce5e19b84cec58dc87b037a09935c2f328"
|
||||
dependencies = [
|
||||
"proc-macro2",
|
||||
"quote",
|
||||
"syn 2.0.119",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "zerotrie"
|
||||
|
|
|
|||
|
|
@ -9,6 +9,8 @@ license = "MIT"
|
|||
repository = "https://github.com/BerriAI/litellm"
|
||||
|
||||
[workspace.dependencies]
|
||||
litellm-tracing = { path = "crates/tracing" }
|
||||
tracing = "0.1"
|
||||
litellm-core = { path = "crates/core" }
|
||||
litellm-host = { path = "crates/host" }
|
||||
litellm-callbacks-legacy-python = { path = "crates/callbacks-legacy-python" }
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ repository.workspace = true
|
|||
|
||||
[dependencies]
|
||||
fancy-regex.workspace = true
|
||||
litellm-tracing.workspace = true
|
||||
litellm-types.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
|
|
|
|||
|
|
@ -1,109 +1 @@
|
|||
use fancy_regex::Regex;
|
||||
|
||||
pub const REDACTED: &str = "REDACTED";
|
||||
|
||||
const DEFAULT_MINIMUM_CUSTOM_KEY_LENGTH: usize = 16;
|
||||
|
||||
fn minimum_custom_key_length() -> usize {
|
||||
std::env::var("MINIMUM_CUSTOM_KEY_LENGTH")
|
||||
.ok()
|
||||
.and_then(|value| value.trim().parse().ok())
|
||||
.unwrap_or(DEFAULT_MINIMUM_CUSTOM_KEY_LENGTH)
|
||||
}
|
||||
|
||||
fn secret_patterns(minimum_custom_key_length: usize) -> String {
|
||||
let sk_suffix_length = minimum_custom_key_length.saturating_sub("sk-".len());
|
||||
[
|
||||
r"-----BEGIN[A-Z \-]*PRIVATE KEY-----[\s\S]*?-----END[A-Z \-]*PRIVATE KEY-----",
|
||||
r"\bya29\.[A-Za-z0-9_.~+/-]+",
|
||||
r#"(?:client_secret|azure_password|azure_username)\s+[^\s,'"})\]{}>]+"#,
|
||||
r"(?:AKIA|ASIA)[0-9A-Z]{16}",
|
||||
r"Bearer\s+[A-Za-z0-9\-._~+/]{10,}=*",
|
||||
r"Basic\s+[A-Za-z0-9+/]{10,}={0,2}",
|
||||
&format!(r"sk-[A-Za-z0-9\-_]{{{sk_suffix_length},}}"),
|
||||
r#"(?<=[?&])(?:api[_-]?key|\w*(?:token|password|passwd|client_secret|secret_key|_secret))=[^\s&'"]+"#,
|
||||
r#"(?:api[_-]?key)['"]?\s*[:=]\s*['"]?[^\s,'"})\]{}>]{8,}"#,
|
||||
r#"(?:x-api-key|api-key)['"]?\s*[:=]\s*['"]?[^\s,'"})\]{}>]+"#,
|
||||
r"x-ak-[A-Za-z0-9\-_]{20,}",
|
||||
r"AIza[0-9A-Za-z\-_]{35}",
|
||||
r#"(?<=[?&])key=[^\s&'"]{8,}"#,
|
||||
r#"(?:^|(?<=\W))\w*(?:password|passwd|client_secret|secret_key|_secret)['"]?\s*[:=]\s*['"]?[^\s,'"})\]{}>]+"#,
|
||||
r#"(?<=://)[^\s'":]{0,4096}:[^\s'"]{1,4096}(?=@)"#,
|
||||
r"dapi[0-9a-f]{32}",
|
||||
r#"litellm\.[A-Za-z0-9_]*_key['"]?\s*[:=]\s*['"]?[^\s,'"})\]{}>]+"#,
|
||||
r#"private_key['"]?\s*[:=]\s*['"]?(?:-----BEGIN[A-Z \-]*PRIVATE KEY-----[\s\S]*?-----END[A-Z \-]*PRIVATE KEY-----|[^\s,'"})\]{}>]+)"#,
|
||||
concat!(
|
||||
r"(?:master_key|xai_key|database_url|db_url|connection_string|",
|
||||
r"aws_secret_access_key|aws_session_token|aws_access_key_id|",
|
||||
r"signing_key|encryption_key|",
|
||||
r"auth_token|access_token|refresh_token|",
|
||||
r"slack_webhook_url|webhook_url|",
|
||||
r"database_connection_string|",
|
||||
r"huggingface_token|jwt_secret)",
|
||||
r#"['"]?\s*[:=]\s*['"]?[^\s,'"})\]{}>]+"#,
|
||||
),
|
||||
r"\beyJ[A-Za-z0-9_-]{10,}\.[A-Za-z0-9_-]+\.[A-Za-z0-9_-]*",
|
||||
r"(?<=[?&])sig=[A-Za-z0-9%+/=]+",
|
||||
r#"\{[^{}]*"type"\s*:\s*"service_account"[^{}]*(?:\{[^{}]*\}[^{}]*)*\}"#,
|
||||
]
|
||||
.join("|")
|
||||
}
|
||||
|
||||
/// Python's `_ENABLE_SECRET_REDACTION` pattern set, compiled once per configuration.
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct SecretRedactor {
|
||||
pattern: Regex,
|
||||
}
|
||||
|
||||
impl SecretRedactor {
|
||||
pub fn new(minimum_custom_key_length: usize) -> Self {
|
||||
let pattern = Regex::new(&format!(
|
||||
"(?i){}",
|
||||
secret_patterns(minimum_custom_key_length)
|
||||
))
|
||||
.expect("secret redaction patterns compile");
|
||||
Self { pattern }
|
||||
}
|
||||
|
||||
/// `None` when `LITELLM_DISABLE_REDACT_SECRETS` turns redaction off.
|
||||
pub fn from_env() -> Option<Self> {
|
||||
let disabled = std::env::var("LITELLM_DISABLE_REDACT_SECRETS")
|
||||
.is_ok_and(|value| value.eq_ignore_ascii_case("true"));
|
||||
(!disabled).then(|| Self::new(minimum_custom_key_length()))
|
||||
}
|
||||
|
||||
pub fn redact(&self, value: &str) -> String {
|
||||
self.pattern.replace_all(value, REDACTED).into_owned()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::bearer("auth failed: Bearer abcdefghijklmnop", "auth failed: REDACTED")]
|
||||
#[case::sk_key("key sk-abcdefghijklmnopqrstuvwxyz rejected", "key REDACTED rejected")]
|
||||
#[case::short_sk_key_is_kept("sk-abc", "sk-abc")]
|
||||
#[case::query_param("GET /v1?api_key=secret123&x=1", "GET /v1?REDACTED&x=1")]
|
||||
#[case::dict_repr("{'api_key': 'abcdefghij'}", "{'REDACTED'}")]
|
||||
#[case::url_credentials("postgres://user:pass@host/db", "postgres://REDACTED@host/db")]
|
||||
#[case::case_insensitive("BEARER ABCDEFGHIJKLMNOP", "REDACTED")]
|
||||
#[case::aws_key("AKIAABCDEFGHIJKLMNOP", "REDACTED")]
|
||||
#[case::sas_signature("https://x.blob/a?sv=1&sig=abc%2B=", "https://x.blob/a?sv=1&REDACTED")]
|
||||
#[case::password_needs_word_boundary("db_password=hunter2", "REDACTED")]
|
||||
#[case::plain_text_is_kept(r#"{"message": "rejected"}"#, r#"{"message": "rejected"}"#)]
|
||||
fn redacts_the_same_spans_as_the_python_patterns(#[case] input: &str, #[case] expected: &str) {
|
||||
assert_eq!(
|
||||
SecretRedactor::new(DEFAULT_MINIMUM_CUSTOM_KEY_LENGTH).redact(input),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sk_threshold_follows_the_minimum_custom_key_length() {
|
||||
let redactor = SecretRedactor::new(8);
|
||||
assert_eq!(redactor.redact("sk-abcde"), REDACTED);
|
||||
assert_eq!(redactor.redact("sk-abcd"), "sk-abcd");
|
||||
}
|
||||
}
|
||||
pub use litellm_tracing::{REDACTED, SecretRedactor};
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ repository.workspace = true
|
|||
autotests = false
|
||||
|
||||
[dependencies]
|
||||
litellm-secrets.workspace = true
|
||||
litellm-types.workspace = true
|
||||
litellm-core-utils.workspace = true
|
||||
litellm-host.workspace = true
|
||||
|
|
@ -36,7 +37,6 @@ url.workspace = true
|
|||
veil.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-secrets.workspace = true
|
||||
litellm-auth-gcp.workspace = true
|
||||
litellm-llms = { workspace = true, features = ["test-support"] }
|
||||
rstest.workspace = true
|
||||
|
|
|
|||
|
|
@ -1,11 +1,9 @@
|
|||
use litellm_auth::{InputSource, SecretValue, Sourced};
|
||||
use litellm_llms::base_llm::{
|
||||
inference::secrets::Secrets,
|
||||
ocr::{
|
||||
handler::OcrClient,
|
||||
transformation::{OcrConnection, OcrCredentialInputs, PreparedOcrRequest},
|
||||
},
|
||||
use litellm_llms::base_llm::ocr::{
|
||||
handler::OcrClient,
|
||||
transformation::{OcrConnection, OcrCredentialInputs, PreparedOcrRequest},
|
||||
};
|
||||
use litellm_secrets::source::Secrets;
|
||||
|
||||
use super::provider_config::OcrProvider;
|
||||
use crate::ocr::types::{LiteLLMOcrRequest, ResolvedOcrRequest};
|
||||
|
|
|
|||
|
|
@ -11,7 +11,6 @@ use litellm_http::{
|
|||
HttpClientPool, HttpSettings, Resolution,
|
||||
media::{PublicDnsResolver, UrlPolicy},
|
||||
};
|
||||
use litellm_llms::base_llm::inference::secrets::{SecretSource, Secrets};
|
||||
use litellm_llms::base_llm::ocr::{
|
||||
error::Error as OcrError,
|
||||
handler::OcrClient,
|
||||
|
|
@ -20,6 +19,7 @@ use litellm_llms::base_llm::ocr::{
|
|||
BaseOcrConfig, LiteLLMOcrResponse, OCR_RESPONSE_MAX_BYTES, OcrTransportConfig,
|
||||
},
|
||||
};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use rstest::rstest;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
|
|
@ -32,27 +32,27 @@ use super::{
|
|||
use crate::ocr::route::{LocalOcrHost, OcrOp, OcrOpResult, ocr_machine};
|
||||
|
||||
struct RecordingSecretSource {
|
||||
names: Arc<Mutex<Vec<&'static str>>>,
|
||||
names: Arc<Mutex<Vec<String>>>,
|
||||
values: &'static [(&'static str, &'static str)],
|
||||
api_base: String,
|
||||
}
|
||||
|
||||
impl SecretSource for RecordingSecretSource {
|
||||
fn resolve<'a>(
|
||||
fn get_secret_str<'a>(
|
||||
&'a self,
|
||||
names: &'a [&'static str],
|
||||
) -> BoxFuture<'a, Result<Secrets, litellm_secrets::Error>> {
|
||||
*self.names.lock().unwrap() = names.to_vec();
|
||||
let values = self.values;
|
||||
let api_base = self.api_base.clone();
|
||||
name: &'a str,
|
||||
) -> BoxFuture<'a, Result<Option<litellm_secrets::SecretValue>, litellm_secrets::Error>> {
|
||||
self.names.lock().unwrap().push(name.to_owned());
|
||||
Box::pin(async move {
|
||||
Ok(Arc::new(move |name: &str| match name {
|
||||
"MISTRAL_AZURE_API_BASE" => Some(api_base.clone()),
|
||||
_ => values
|
||||
Ok(match name {
|
||||
"MISTRAL_AZURE_API_BASE" => Some(self.api_base.clone()),
|
||||
_ => self
|
||||
.values
|
||||
.iter()
|
||||
.find(|(key, _)| *key == name)
|
||||
.map(|(_, value)| value.to_string()),
|
||||
}) as Secrets)
|
||||
}
|
||||
.map(litellm_secrets::SecretValue::new))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
@ -289,7 +289,7 @@ async fn ocr_client_uses_the_injected_http_pool_configuration() {
|
|||
UrlPolicy::default(),
|
||||
VertexAuth::default(),
|
||||
OcrSettings::default(),
|
||||
Arc::new(litellm_llms::base_llm::inference::secrets::EnvironmentSecrets),
|
||||
Arc::new(litellm_secrets::source::EnvironmentSecrets::default()),
|
||||
)
|
||||
.unwrap();
|
||||
crate::ocr::client::perform(&client, wire_request("mistral/model", &base, json!({})))
|
||||
|
|
|
|||
14
litellm-rust/crates/cost/Cargo.toml
Normal file
14
litellm-rust/crates/cost/Cargo.toml
Normal file
|
|
@ -0,0 +1,14 @@
|
|||
[package]
|
||||
name = "litellm-cost"
|
||||
version = "0.1.0"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
criterion.workspace = true
|
||||
proptest.workspace = true
|
||||
|
||||
[[bench]]
|
||||
name = "calculate"
|
||||
harness = false
|
||||
13
litellm-rust/crates/cost/README.md
Normal file
13
litellm-rust/crates/cost/README.md
Normal file
|
|
@ -0,0 +1,13 @@
|
|||
# litellm-cost
|
||||
|
||||
This crate calculates text token charges from rates and usage supplied by its caller. It is standalone and has no Python bridge or proxy integration
|
||||
|
||||
Call `compile(&pricing)` once for an immutable plan, then `plan.calculate(&request)` for each supported request. `calculate(&pricing, &request)` compiles on each call. A successful result exposes pre-multiplier component costs, selected rates, the multiplier, and derived `input()`, `output()`, and `total()` values
|
||||
|
||||
The caller states whether `prompt_tokens` includes cache tokens. Threshold selection uses total input tokens for either convention and selects one rate for the whole request. Thresholds are sorted when compiled, and duplicate thresholds or tier overrides fail deterministically. `Fast` selects priority rates; unknown tiers use standard rates
|
||||
|
||||
`Rate::Missing`, `Rate::Null`, and `Rate::Value(0.0)` remain distinct. Missing cache rates fall back to the selected input rate, and an absent one-hour write rate falls back to the selected write rate. Missing input or output rates return typed errors, including for zero usage. Python's sparse-entry behavior remains outside this native contract
|
||||
|
||||
The supported off-peak shape is one non-wrapping UTC daily window. The caller supplies the applicable regional multiplier after provider-specific selection. Negative or non-finite rates, ambiguous rules, inconsistent cache counts, incomplete write splits, invalid windows and overflow return errors. Callers must decline unsupported inputs before native execution if their public contract accepts those shapes
|
||||
|
||||
This crate does not select models, read catalogs, fetch provider prices, normalize multimodal usage, process provider-reported costs, or calculate non-token charges. It does not change proxy behavior. The reference fixture was generated by `tests/generate_python_reference.py` against the Python implementation at the commit recorded in `tests/python_reference.tsv`, using synthetic rates and fixed usage
|
||||
66
litellm-rust/crates/cost/benches/calculate.rs
Normal file
66
litellm-rust/crates/cost/benches/calculate.rs
Normal file
|
|
@ -0,0 +1,66 @@
|
|||
use criterion::{Criterion, criterion_group, criterion_main};
|
||||
use litellm_cost::{
|
||||
Pricing, PromptConvention, Rate, Rates, Request, ServiceTier, ThresholdPolicy, ThresholdRates,
|
||||
Usage, calculate, compile,
|
||||
};
|
||||
use std::hint::black_box;
|
||||
|
||||
fn bench(c: &mut Criterion) {
|
||||
let pricing = Pricing {
|
||||
standard: Rates {
|
||||
input: Rate::Value(0.000002),
|
||||
output: Rate::Value(0.000008),
|
||||
cache_read: Rate::Value(0.0000005),
|
||||
cache_write: Rate::Missing,
|
||||
cache_write_1h: Rate::Missing,
|
||||
},
|
||||
tiers: &[],
|
||||
thresholds: &[],
|
||||
off_peak: None,
|
||||
};
|
||||
let request = Request {
|
||||
usage: Usage {
|
||||
prompt_tokens: 1000,
|
||||
completion_tokens: 200,
|
||||
cache_read_tokens: 250,
|
||||
cache_write_tokens: 0,
|
||||
cache_write_5m_tokens: None,
|
||||
cache_write_1h_tokens: None,
|
||||
prompt_convention: PromptConvention::IncludesCache,
|
||||
},
|
||||
service_tier: ServiceTier::Standard,
|
||||
threshold_policy: ThresholdPolicy::Exclusive,
|
||||
region_multiplier: None,
|
||||
billed_at_utc_minute: None,
|
||||
};
|
||||
let plan = compile(&pricing).unwrap();
|
||||
c.bench_function("native_compiled_calculation", |b| {
|
||||
b.iter(|| black_box(plan.calculate(black_box(&request)).unwrap()))
|
||||
});
|
||||
c.bench_function("native_full_wrapper", |b| {
|
||||
b.iter(|| black_box(calculate(black_box(&pricing), black_box(&request)).unwrap()))
|
||||
});
|
||||
c.bench_function("native_rate_compilation", |b| {
|
||||
b.iter(|| black_box(compile(black_box(&pricing)).unwrap()))
|
||||
});
|
||||
let threshold = ThresholdRates {
|
||||
above_prompt_tokens: 1000,
|
||||
standard: Rates {
|
||||
input: Rate::Value(0.000004),
|
||||
output: Rate::Value(0.000016),
|
||||
..Rates::EMPTY
|
||||
},
|
||||
tiers: &[],
|
||||
};
|
||||
let threshold_pricing = Pricing {
|
||||
thresholds: &[threshold],
|
||||
..pricing
|
||||
};
|
||||
let threshold_plan = compile(&threshold_pricing).unwrap();
|
||||
c.bench_function("native_threshold_boundary", |b| {
|
||||
b.iter(|| black_box(threshold_plan.calculate(black_box(&request)).unwrap()))
|
||||
});
|
||||
}
|
||||
|
||||
criterion_group!(benches, bench);
|
||||
criterion_main!(benches);
|
||||
40
litellm-rust/crates/cost/examples/charge.rs
Normal file
40
litellm-rust/crates/cost/examples/charge.rs
Normal file
|
|
@ -0,0 +1,40 @@
|
|||
use litellm_cost::{
|
||||
Pricing, PromptConvention, Rate, Rates, Request, ServiceTier, ThresholdPolicy, Usage, compile,
|
||||
};
|
||||
|
||||
fn main() {
|
||||
let pricing = Pricing {
|
||||
standard: Rates {
|
||||
input: Rate::Value(2.0),
|
||||
output: Rate::Value(4.0),
|
||||
cache_read: Rate::Value(0.5),
|
||||
cache_write: Rate::Value(3.0),
|
||||
cache_write_1h: Rate::Missing,
|
||||
},
|
||||
tiers: &[],
|
||||
thresholds: &[],
|
||||
off_peak: None,
|
||||
};
|
||||
let request = Request {
|
||||
usage: Usage {
|
||||
prompt_tokens: 100,
|
||||
completion_tokens: 20,
|
||||
cache_read_tokens: 25,
|
||||
cache_write_tokens: 10,
|
||||
cache_write_5m_tokens: None,
|
||||
cache_write_1h_tokens: None,
|
||||
prompt_convention: PromptConvention::IncludesCache,
|
||||
},
|
||||
service_tier: ServiceTier::Standard,
|
||||
threshold_policy: ThresholdPolicy::Exclusive,
|
||||
region_multiplier: None,
|
||||
billed_at_utc_minute: None,
|
||||
};
|
||||
let cost = compile(&pricing).unwrap().calculate(&request).unwrap();
|
||||
println!(
|
||||
"input={} output={} total={}",
|
||||
cost.input(),
|
||||
cost.output(),
|
||||
cost.total()
|
||||
);
|
||||
}
|
||||
405
litellm-rust/crates/cost/src/lib.rs
Normal file
405
litellm-rust/crates/cost/src/lib.rs
Normal file
|
|
@ -0,0 +1,405 @@
|
|||
#[derive(Clone, Copy, Debug, PartialEq)]
|
||||
pub enum Rate {
|
||||
Missing,
|
||||
Null,
|
||||
Value(f64),
|
||||
}
|
||||
|
||||
impl Rate {
|
||||
fn value(self) -> Option<f64> {
|
||||
match self {
|
||||
Self::Value(value) => Some(value),
|
||||
Self::Missing | Self::Null => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn or(self, fallback: Self) -> Self {
|
||||
if self.value().is_some() {
|
||||
self
|
||||
} else {
|
||||
fallback
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq)]
|
||||
pub struct Rates {
|
||||
pub input: Rate,
|
||||
pub output: Rate,
|
||||
pub cache_read: Rate,
|
||||
pub cache_write: Rate,
|
||||
pub cache_write_1h: Rate,
|
||||
}
|
||||
|
||||
impl Rates {
|
||||
pub const EMPTY: Self = Self {
|
||||
input: Rate::Missing,
|
||||
output: Rate::Missing,
|
||||
cache_read: Rate::Missing,
|
||||
cache_write: Rate::Missing,
|
||||
cache_write_1h: Rate::Missing,
|
||||
};
|
||||
|
||||
fn overlay(self, base: Self) -> Self {
|
||||
Self {
|
||||
input: self.input.or(base.input),
|
||||
output: self.output.or(base.output),
|
||||
cache_read: self.cache_read.or(base.cache_read),
|
||||
cache_write: self.cache_write.or(base.cache_write),
|
||||
cache_write_1h: self.cache_write_1h.or(base.cache_write_1h),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub enum ServiceTier {
|
||||
Standard,
|
||||
Flex,
|
||||
Priority,
|
||||
Fast,
|
||||
Ultrafast,
|
||||
Unknown,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub enum ThresholdPolicy {
|
||||
Exclusive,
|
||||
Inclusive,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub enum PromptConvention {
|
||||
IncludesCache,
|
||||
ExcludesCache,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq)]
|
||||
pub struct Usage {
|
||||
pub prompt_tokens: u64,
|
||||
pub completion_tokens: u64,
|
||||
pub cache_read_tokens: u64,
|
||||
pub cache_write_tokens: u64,
|
||||
pub cache_write_5m_tokens: Option<u64>,
|
||||
pub cache_write_1h_tokens: Option<u64>,
|
||||
pub prompt_convention: PromptConvention,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq)]
|
||||
pub struct TierRates {
|
||||
pub tier: ServiceTier,
|
||||
pub rates: Rates,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq)]
|
||||
pub struct ThresholdRates<'a> {
|
||||
pub above_prompt_tokens: u64,
|
||||
pub standard: Rates,
|
||||
pub tiers: &'a [TierRates],
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq)]
|
||||
pub struct OffPeakRates {
|
||||
pub start_utc_minute: u16,
|
||||
pub end_utc_minute: u16,
|
||||
pub rates: Rates,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq)]
|
||||
pub struct Pricing<'a> {
|
||||
pub standard: Rates,
|
||||
pub tiers: &'a [TierRates],
|
||||
pub thresholds: &'a [ThresholdRates<'a>],
|
||||
pub off_peak: Option<OffPeakRates>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq)]
|
||||
pub struct Request {
|
||||
pub usage: Usage,
|
||||
pub service_tier: ServiceTier,
|
||||
pub threshold_policy: ThresholdPolicy,
|
||||
pub region_multiplier: Option<f64>,
|
||||
pub billed_at_utc_minute: Option<u16>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq)]
|
||||
pub struct Cost {
|
||||
pub uncached_input: f64,
|
||||
pub cache_read: f64,
|
||||
pub cache_write_5m: f64,
|
||||
pub cache_write_1h: f64,
|
||||
pub output: f64,
|
||||
pub multiplier: f64,
|
||||
pub rates: EffectiveRates,
|
||||
}
|
||||
|
||||
impl Cost {
|
||||
pub fn input(self) -> f64 {
|
||||
(self.uncached_input + self.cache_read + self.cache_write_5m + self.cache_write_1h)
|
||||
* self.multiplier
|
||||
}
|
||||
|
||||
pub fn output(self) -> f64 {
|
||||
self.output * self.multiplier
|
||||
}
|
||||
|
||||
pub fn total(self) -> f64 {
|
||||
self.input() + self.output()
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq)]
|
||||
pub struct EffectiveRates {
|
||||
pub input: f64,
|
||||
pub output: f64,
|
||||
pub cache_read: f64,
|
||||
pub cache_write_5m: f64,
|
||||
pub cache_write_1h: f64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub enum PricingError {
|
||||
MissingInputRate,
|
||||
MissingOutputRate,
|
||||
InvalidRate,
|
||||
InvalidRegionMultiplier,
|
||||
InvalidBillingTime,
|
||||
InvalidOffPeakWindow,
|
||||
CacheExceedsPrompt,
|
||||
InvalidCacheWriteDetails,
|
||||
TokenCountOverflow,
|
||||
DuplicateTier,
|
||||
DuplicateThreshold,
|
||||
DuplicateThresholdTier,
|
||||
}
|
||||
|
||||
fn selected_tier(tier: ServiceTier) -> ServiceTier {
|
||||
if tier == ServiceTier::Fast {
|
||||
ServiceTier::Priority
|
||||
} else {
|
||||
tier
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
struct CompiledThreshold {
|
||||
above_prompt_tokens: u64,
|
||||
standard: Rates,
|
||||
tiers: Vec<TierRates>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct PricingPlan {
|
||||
standard: Rates,
|
||||
tiers: Vec<TierRates>,
|
||||
thresholds: Vec<CompiledThreshold>,
|
||||
off_peak: Option<OffPeakRates>,
|
||||
}
|
||||
|
||||
fn valid_rates(rates: Rates) -> bool {
|
||||
[
|
||||
rates.input,
|
||||
rates.output,
|
||||
rates.cache_read,
|
||||
rates.cache_write,
|
||||
rates.cache_write_1h,
|
||||
]
|
||||
.into_iter()
|
||||
.all(|rate| {
|
||||
rate.value()
|
||||
.is_none_or(|value| value.is_finite() && value >= 0.0)
|
||||
})
|
||||
}
|
||||
|
||||
fn validate_tiers(tiers: &[TierRates], duplicate: PricingError) -> Result<(), PricingError> {
|
||||
if tiers.iter().any(|entry| !valid_rates(entry.rates)) {
|
||||
return Err(PricingError::InvalidRate);
|
||||
}
|
||||
if tiers.iter().enumerate().any(|(index, entry)| {
|
||||
matches!(
|
||||
entry.tier,
|
||||
ServiceTier::Standard | ServiceTier::Unknown | ServiceTier::Fast
|
||||
) || tiers[..index]
|
||||
.iter()
|
||||
.any(|previous| previous.tier == entry.tier)
|
||||
}) {
|
||||
return Err(duplicate);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn compile(pricing: &Pricing<'_>) -> Result<PricingPlan, PricingError> {
|
||||
if !valid_rates(pricing.standard) {
|
||||
return Err(PricingError::InvalidRate);
|
||||
}
|
||||
validate_tiers(pricing.tiers, PricingError::DuplicateTier)?;
|
||||
if let Some(window) = pricing.off_peak {
|
||||
if window.start_utc_minute >= 1440
|
||||
|| window.end_utc_minute > 1440
|
||||
|| window.start_utc_minute >= window.end_utc_minute
|
||||
{
|
||||
return Err(PricingError::InvalidOffPeakWindow);
|
||||
}
|
||||
if !valid_rates(window.rates) {
|
||||
return Err(PricingError::InvalidRate);
|
||||
}
|
||||
}
|
||||
let mut thresholds: Vec<_> = pricing
|
||||
.thresholds
|
||||
.iter()
|
||||
.map(|entry| {
|
||||
if !valid_rates(entry.standard) {
|
||||
return Err(PricingError::InvalidRate);
|
||||
}
|
||||
validate_tiers(entry.tiers, PricingError::DuplicateThresholdTier)?;
|
||||
Ok(CompiledThreshold {
|
||||
above_prompt_tokens: entry.above_prompt_tokens,
|
||||
standard: entry.standard,
|
||||
tiers: entry.tiers.to_vec(),
|
||||
})
|
||||
})
|
||||
.collect::<Result<_, _>>()?;
|
||||
thresholds.sort_unstable_by_key(|entry| entry.above_prompt_tokens);
|
||||
if thresholds
|
||||
.windows(2)
|
||||
.any(|pair| pair[0].above_prompt_tokens == pair[1].above_prompt_tokens)
|
||||
{
|
||||
return Err(PricingError::DuplicateThreshold);
|
||||
}
|
||||
Ok(PricingPlan {
|
||||
standard: pricing.standard,
|
||||
tiers: pricing.tiers.to_vec(),
|
||||
thresholds,
|
||||
off_peak: pricing.off_peak,
|
||||
})
|
||||
}
|
||||
|
||||
impl PricingPlan {
|
||||
fn resolve_rates(
|
||||
&self,
|
||||
request: &Request,
|
||||
threshold_tokens: u64,
|
||||
) -> Result<Rates, PricingError> {
|
||||
let tier = selected_tier(request.service_tier);
|
||||
let base = self
|
||||
.tiers
|
||||
.iter()
|
||||
.find(|entry| tier != ServiceTier::Standard && entry.tier == tier)
|
||||
.map_or(self.standard, |entry| entry.rates.overlay(self.standard));
|
||||
let threshold = self.thresholds.iter().rev().find(|entry| {
|
||||
threshold_tokens > entry.above_prompt_tokens
|
||||
|| (request.threshold_policy == ThresholdPolicy::Inclusive
|
||||
&& threshold_tokens == entry.above_prompt_tokens)
|
||||
});
|
||||
let selected = threshold.map_or(base, |entry| {
|
||||
let standard = entry.standard.overlay(base);
|
||||
entry
|
||||
.tiers
|
||||
.iter()
|
||||
.find(|specific| tier != ServiceTier::Standard && specific.tier == tier)
|
||||
.map_or(standard, |specific| specific.rates.overlay(standard))
|
||||
});
|
||||
match self.off_peak {
|
||||
None => Ok(selected),
|
||||
Some(window) => {
|
||||
if window.start_utc_minute >= 1440
|
||||
|| window.end_utc_minute > 1440
|
||||
|| window.start_utc_minute >= window.end_utc_minute
|
||||
{
|
||||
return Err(PricingError::InvalidOffPeakWindow);
|
||||
}
|
||||
let minute = request
|
||||
.billed_at_utc_minute
|
||||
.ok_or(PricingError::InvalidBillingTime)?;
|
||||
if minute >= 1440 {
|
||||
return Err(PricingError::InvalidBillingTime);
|
||||
}
|
||||
if (window.start_utc_minute..window.end_utc_minute).contains(&minute) {
|
||||
Ok(window.rates.overlay(selected))
|
||||
} else {
|
||||
Ok(selected)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn checked_rate(rate: Rate, missing: PricingError) -> Result<f64, PricingError> {
|
||||
let value = rate.value().ok_or(missing)?;
|
||||
if !value.is_finite() || value < 0.0 {
|
||||
return Err(PricingError::InvalidRate);
|
||||
}
|
||||
Ok(value)
|
||||
}
|
||||
|
||||
pub fn calculate(&self, request: &Request) -> Result<Cost, PricingError> {
|
||||
let usage = request.usage;
|
||||
let cached = usage
|
||||
.cache_read_tokens
|
||||
.checked_add(usage.cache_write_tokens)
|
||||
.ok_or(PricingError::TokenCountOverflow)?;
|
||||
let (regular, threshold_tokens) = match usage.prompt_convention {
|
||||
PromptConvention::IncludesCache => (
|
||||
usage
|
||||
.prompt_tokens
|
||||
.checked_sub(cached)
|
||||
.ok_or(PricingError::CacheExceedsPrompt)?,
|
||||
usage.prompt_tokens,
|
||||
),
|
||||
PromptConvention::ExcludesCache => (
|
||||
usage.prompt_tokens,
|
||||
usage
|
||||
.prompt_tokens
|
||||
.checked_add(cached)
|
||||
.ok_or(PricingError::TokenCountOverflow)?,
|
||||
),
|
||||
};
|
||||
let writes = match (usage.cache_write_5m_tokens, usage.cache_write_1h_tokens) {
|
||||
(None, None) => (usage.cache_write_tokens, 0),
|
||||
(Some(five), Some(one)) if five.checked_add(one) == Some(usage.cache_write_tokens) => {
|
||||
(five, one)
|
||||
}
|
||||
_ => return Err(PricingError::InvalidCacheWriteDetails),
|
||||
};
|
||||
let rates = self.resolve_rates(request, threshold_tokens)?;
|
||||
let input = Self::checked_rate(rates.input, PricingError::MissingInputRate)?;
|
||||
let output = Self::checked_rate(rates.output, PricingError::MissingOutputRate)?;
|
||||
let read = Self::checked_rate(
|
||||
rates.cache_read.or(rates.input),
|
||||
PricingError::MissingInputRate,
|
||||
)?;
|
||||
let write = Self::checked_rate(
|
||||
rates.cache_write.or(rates.input),
|
||||
PricingError::MissingInputRate,
|
||||
)?;
|
||||
let write_1h = Self::checked_rate(
|
||||
rates.cache_write_1h.or(rates.cache_write).or(rates.input),
|
||||
PricingError::MissingInputRate,
|
||||
)?;
|
||||
let multiplier = request.region_multiplier.unwrap_or(1.0);
|
||||
if !multiplier.is_finite() || multiplier <= 0.0 {
|
||||
return Err(PricingError::InvalidRegionMultiplier);
|
||||
}
|
||||
let cost = Cost {
|
||||
uncached_input: regular as f64 * input,
|
||||
cache_read: usage.cache_read_tokens as f64 * read,
|
||||
cache_write_5m: writes.0 as f64 * write,
|
||||
cache_write_1h: writes.1 as f64 * write_1h,
|
||||
output: usage.completion_tokens as f64 * output,
|
||||
multiplier,
|
||||
rates: EffectiveRates {
|
||||
input,
|
||||
output,
|
||||
cache_read: read,
|
||||
cache_write_5m: write,
|
||||
cache_write_1h: write_1h,
|
||||
},
|
||||
};
|
||||
if !cost.total().is_finite() {
|
||||
return Err(PricingError::TokenCountOverflow);
|
||||
}
|
||||
Ok(cost)
|
||||
}
|
||||
}
|
||||
|
||||
pub fn calculate(pricing: &Pricing<'_>, request: &Request) -> Result<Cost, PricingError> {
|
||||
compile(pricing)?.calculate(request)
|
||||
}
|
||||
458
litellm-rust/crates/cost/tests/calculation.rs
Normal file
458
litellm-rust/crates/cost/tests/calculation.rs
Normal file
|
|
@ -0,0 +1,458 @@
|
|||
use litellm_cost::{
|
||||
OffPeakRates, Pricing, PricingError, PromptConvention, Rate, Rates, Request, ServiceTier,
|
||||
ThresholdPolicy, ThresholdRates, TierRates, Usage, calculate, compile,
|
||||
};
|
||||
|
||||
fn rates(input: Rate, output: Rate) -> Rates {
|
||||
Rates {
|
||||
input,
|
||||
output,
|
||||
..Rates::EMPTY
|
||||
}
|
||||
}
|
||||
|
||||
fn request() -> Request {
|
||||
Request {
|
||||
usage: Usage {
|
||||
prompt_tokens: 100,
|
||||
completion_tokens: 20,
|
||||
cache_read_tokens: 25,
|
||||
cache_write_tokens: 10,
|
||||
cache_write_5m_tokens: None,
|
||||
cache_write_1h_tokens: None,
|
||||
prompt_convention: PromptConvention::IncludesCache,
|
||||
},
|
||||
service_tier: ServiceTier::Standard,
|
||||
threshold_policy: ThresholdPolicy::Exclusive,
|
||||
region_multiplier: None,
|
||||
billed_at_utc_minute: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn pricing(standard: Rates) -> Pricing<'static> {
|
||||
Pricing {
|
||||
standard,
|
||||
tiers: &[],
|
||||
thresholds: &[],
|
||||
off_peak: None,
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn breakdown_and_total_agree() {
|
||||
let standard = Rates {
|
||||
cache_read: Rate::Value(0.5),
|
||||
cache_write: Rate::Value(3.0),
|
||||
..rates(Rate::Value(2.0), Rate::Value(4.0))
|
||||
};
|
||||
let result = calculate(&pricing(standard), &request()).unwrap();
|
||||
assert_eq!(result.uncached_input, 65.0 * 2.0);
|
||||
assert_eq!(result.cache_read, 25.0 * 0.5);
|
||||
assert_eq!(result.cache_write_5m, 10.0 * 3.0);
|
||||
assert_eq!(result.output(), 20.0 * 4.0);
|
||||
assert_eq!(result.total(), result.input() + result.output());
|
||||
assert_eq!(result.rates.cache_read, 0.5);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn absent_null_and_zero_cache_rates_are_distinct() {
|
||||
let base = rates(Rate::Value(2.0), Rate::Value(4.0));
|
||||
for read in [Rate::Missing, Rate::Null] {
|
||||
let standard = Rates {
|
||||
cache_read: read,
|
||||
..base
|
||||
};
|
||||
assert_eq!(
|
||||
calculate(&pricing(standard), &request()).unwrap().input(),
|
||||
200.0
|
||||
);
|
||||
}
|
||||
let standard = Rates {
|
||||
cache_read: Rate::Value(0.0),
|
||||
cache_write: Rate::Value(0.0),
|
||||
..base
|
||||
};
|
||||
assert_eq!(
|
||||
calculate(&pricing(standard), &request()).unwrap().input(),
|
||||
130.0
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn equivalent_prompt_conventions_select_the_same_threshold() {
|
||||
let threshold = ThresholdRates {
|
||||
above_prompt_tokens: 90,
|
||||
standard: rates(Rate::Value(5.0), Rate::Value(8.0)),
|
||||
tiers: &[],
|
||||
};
|
||||
let specification = Pricing {
|
||||
standard: rates(Rate::Value(2.0), Rate::Value(4.0)),
|
||||
tiers: &[],
|
||||
thresholds: &[threshold],
|
||||
off_peak: None,
|
||||
};
|
||||
let included = request();
|
||||
let excluded = Request {
|
||||
usage: Usage {
|
||||
prompt_tokens: 65,
|
||||
prompt_convention: PromptConvention::ExcludesCache,
|
||||
..included.usage
|
||||
},
|
||||
..included
|
||||
};
|
||||
let plan = compile(&specification).unwrap();
|
||||
assert_eq!(plan.calculate(&included), plan.calculate(&excluded));
|
||||
assert_eq!(plan.calculate(&included).unwrap().rates.input, 5.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn split_writes_and_invalid_accounting() {
|
||||
let standard = Rates {
|
||||
cache_read: Rate::Value(0.5),
|
||||
cache_write: Rate::Value(3.0),
|
||||
cache_write_1h: Rate::Value(5.0),
|
||||
..rates(Rate::Value(2.0), Rate::Value(4.0))
|
||||
};
|
||||
let base = request();
|
||||
let split = Request {
|
||||
usage: Usage {
|
||||
cache_write_5m_tokens: Some(4),
|
||||
cache_write_1h_tokens: Some(6),
|
||||
..base.usage
|
||||
},
|
||||
..base
|
||||
};
|
||||
let result = calculate(&pricing(standard), &split).unwrap();
|
||||
assert_eq!(result.cache_write_5m, 12.0);
|
||||
assert_eq!(result.cache_write_1h, 30.0);
|
||||
let overlapping = Request {
|
||||
usage: Usage {
|
||||
prompt_tokens: 30,
|
||||
..split.usage
|
||||
},
|
||||
..split
|
||||
};
|
||||
assert_eq!(
|
||||
calculate(&pricing(standard), &overlapping),
|
||||
Err(PricingError::CacheExceedsPrompt)
|
||||
);
|
||||
let incomplete = Request {
|
||||
usage: Usage {
|
||||
cache_write_1h_tokens: None,
|
||||
..split.usage
|
||||
},
|
||||
..split
|
||||
};
|
||||
assert_eq!(
|
||||
calculate(&pricing(standard), &incomplete),
|
||||
Err(PricingError::InvalidCacheWriteDetails)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn threshold_tiers_and_boundaries() {
|
||||
let priority = TierRates {
|
||||
tier: ServiceTier::Priority,
|
||||
rates: rates(Rate::Value(3.0), Rate::Missing),
|
||||
};
|
||||
let threshold = ThresholdRates {
|
||||
above_prompt_tokens: 100,
|
||||
standard: rates(Rate::Value(5.0), Rate::Value(8.0)),
|
||||
tiers: &[
|
||||
TierRates {
|
||||
tier: ServiceTier::Priority,
|
||||
rates: rates(Rate::Value(7.0), Rate::Missing),
|
||||
},
|
||||
TierRates {
|
||||
tier: ServiceTier::Flex,
|
||||
rates: rates(Rate::Value(6.0), Rate::Missing),
|
||||
},
|
||||
],
|
||||
};
|
||||
let specification = Pricing {
|
||||
standard: rates(Rate::Value(2.0), Rate::Value(4.0)),
|
||||
tiers: &[priority],
|
||||
thresholds: &[threshold],
|
||||
off_peak: None,
|
||||
};
|
||||
let base = request();
|
||||
let no_cache = Request {
|
||||
usage: Usage {
|
||||
cache_read_tokens: 0,
|
||||
cache_write_tokens: 0,
|
||||
..base.usage
|
||||
},
|
||||
..base
|
||||
};
|
||||
let fast = Request {
|
||||
service_tier: ServiceTier::Fast,
|
||||
..no_cache
|
||||
};
|
||||
let inclusive = Request {
|
||||
threshold_policy: ThresholdPolicy::Inclusive,
|
||||
..fast
|
||||
};
|
||||
let flex = Request {
|
||||
service_tier: ServiceTier::Flex,
|
||||
..inclusive
|
||||
};
|
||||
assert_eq!(calculate(&specification, &no_cache).unwrap().input(), 200.0);
|
||||
assert_eq!(calculate(&specification, &fast).unwrap().input(), 300.0);
|
||||
assert_eq!(
|
||||
calculate(&specification, &inclusive).unwrap().input(),
|
||||
700.0
|
||||
);
|
||||
assert_eq!(calculate(&specification, &flex).unwrap().input(), 600.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn compile_rejects_ambiguous_rates() {
|
||||
let duplicate = ThresholdRates {
|
||||
above_prompt_tokens: 100,
|
||||
standard: Rates::EMPTY,
|
||||
tiers: &[],
|
||||
};
|
||||
let specification = Pricing {
|
||||
standard: rates(Rate::Value(1.0), Rate::Value(1.0)),
|
||||
tiers: &[],
|
||||
thresholds: &[duplicate, duplicate],
|
||||
off_peak: None,
|
||||
};
|
||||
assert_eq!(
|
||||
compile(&specification).err(),
|
||||
Some(PricingError::DuplicateThreshold)
|
||||
);
|
||||
let invalid = pricing(rates(Rate::Value(f64::NAN), Rate::Value(1.0)));
|
||||
assert_eq!(compile(&invalid).err(), Some(PricingError::InvalidRate));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn off_peak_is_one_non_wrapping_utc_window() {
|
||||
let specification = Pricing {
|
||||
standard: rates(Rate::Value(2.0), Rate::Value(4.0)),
|
||||
tiers: &[],
|
||||
thresholds: &[],
|
||||
off_peak: Some(OffPeakRates {
|
||||
start_utc_minute: 60,
|
||||
end_utc_minute: 120,
|
||||
rates: rates(Rate::Value(1.0), Rate::Value(2.0)),
|
||||
}),
|
||||
};
|
||||
let base = request();
|
||||
let start = Request {
|
||||
billed_at_utc_minute: Some(60),
|
||||
..base
|
||||
};
|
||||
let end = Request {
|
||||
billed_at_utc_minute: Some(120),
|
||||
..base
|
||||
};
|
||||
assert_eq!(
|
||||
calculate(&specification, &base),
|
||||
Err(PricingError::InvalidBillingTime)
|
||||
);
|
||||
assert_eq!(calculate(&specification, &start).unwrap().input(), 100.0);
|
||||
assert_eq!(calculate(&specification, &end).unwrap().input(), 200.0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn missing_rates_and_free_rates_remain_distinct() {
|
||||
let base = request();
|
||||
let empty = Request {
|
||||
usage: Usage {
|
||||
prompt_tokens: 0,
|
||||
completion_tokens: 0,
|
||||
cache_read_tokens: 0,
|
||||
cache_write_tokens: 0,
|
||||
..base.usage
|
||||
},
|
||||
..base
|
||||
};
|
||||
assert_eq!(
|
||||
calculate(&pricing(Rates::EMPTY), &empty),
|
||||
Err(PricingError::MissingInputRate)
|
||||
);
|
||||
assert_eq!(
|
||||
calculate(&pricing(rates(Rate::Value(0.0), Rate::Missing)), &empty),
|
||||
Err(PricingError::MissingOutputRate)
|
||||
);
|
||||
assert_eq!(
|
||||
calculate(&pricing(rates(Rate::Value(0.0), Rate::Value(0.0))), &empty)
|
||||
.unwrap()
|
||||
.total(),
|
||||
0.0
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn matches_executed_python_reference_cases() {
|
||||
for row in include_str!("python_reference.tsv")
|
||||
.lines()
|
||||
.filter(|line| !line.starts_with('#'))
|
||||
{
|
||||
let fields: Vec<_> = row.split('\t').collect();
|
||||
let count = |index: usize| fields[index].parse::<u64>().unwrap();
|
||||
let number = |index: usize| fields[index].parse::<f64>().unwrap();
|
||||
let optional_rate = |index: usize| {
|
||||
if fields[index].is_empty() {
|
||||
Rate::Missing
|
||||
} else {
|
||||
Rate::Value(number(index))
|
||||
}
|
||||
};
|
||||
let threshold = ThresholdRates {
|
||||
above_prompt_tokens: if fields[9].is_empty() { 0 } else { count(9) },
|
||||
standard: rates(optional_rate(10), optional_rate(11)),
|
||||
tiers: &[],
|
||||
};
|
||||
let thresholds = if fields[9].is_empty() {
|
||||
&[][..]
|
||||
} else {
|
||||
std::slice::from_ref(&threshold)
|
||||
};
|
||||
let specification = Pricing {
|
||||
standard: Rates {
|
||||
cache_read: optional_rate(7),
|
||||
cache_write: optional_rate(8),
|
||||
..rates(Rate::Value(number(5)), Rate::Value(number(6)))
|
||||
},
|
||||
tiers: &[],
|
||||
thresholds,
|
||||
off_peak: None,
|
||||
};
|
||||
let base = request();
|
||||
let input = Request {
|
||||
usage: Usage {
|
||||
prompt_tokens: count(1),
|
||||
completion_tokens: count(2),
|
||||
cache_read_tokens: count(3),
|
||||
cache_write_tokens: count(4),
|
||||
..base.usage
|
||||
},
|
||||
..base
|
||||
};
|
||||
let actual = calculate(&specification, &input).unwrap();
|
||||
assert_eq!(actual.input(), number(12), "{}", fields[0]);
|
||||
assert_eq!(actual.output(), number(13), "{}", fields[0]);
|
||||
}
|
||||
}
|
||||
|
||||
proptest::proptest! {
|
||||
#[test]
|
||||
fn equivalent_usage_conventions_and_breakdown_agree(
|
||||
regular in 0_u64..1000,
|
||||
read in 0_u64..1000,
|
||||
write in 0_u64..1000,
|
||||
output in 0_u64..1000,
|
||||
) {
|
||||
let threshold = ThresholdRates {
|
||||
above_prompt_tokens: 1000,
|
||||
standard: rates(Rate::Value(5.0), Rate::Value(8.0)),
|
||||
tiers: &[],
|
||||
};
|
||||
let specification = Pricing {
|
||||
standard: rates(Rate::Value(2.0), Rate::Value(4.0)),
|
||||
tiers: &[],
|
||||
thresholds: &[threshold],
|
||||
off_peak: None,
|
||||
};
|
||||
let base = request();
|
||||
let included = Request {
|
||||
usage: Usage {
|
||||
prompt_tokens: regular + read + write,
|
||||
completion_tokens: output,
|
||||
cache_read_tokens: read,
|
||||
cache_write_tokens: write,
|
||||
..base.usage
|
||||
},
|
||||
..base
|
||||
};
|
||||
let excluded = Request {
|
||||
usage: Usage {
|
||||
prompt_tokens: regular,
|
||||
prompt_convention: PromptConvention::ExcludesCache,
|
||||
..included.usage
|
||||
},
|
||||
..included
|
||||
};
|
||||
let plan = compile(&specification).unwrap();
|
||||
let left = plan.calculate(&included).unwrap();
|
||||
let right = plan.calculate(&excluded).unwrap();
|
||||
proptest::prop_assert_eq!(left, right);
|
||||
proptest::prop_assert_eq!(left.total(), left.input() + left.output());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn regional_multiplier_applies_after_input_components_are_summed() {
|
||||
let standard = Rates {
|
||||
cache_read: Rate::Value(0.5),
|
||||
cache_write: Rate::Value(3.0),
|
||||
..rates(Rate::Value(2.0), Rate::Value(4.0))
|
||||
};
|
||||
let base = request();
|
||||
let regional = Request {
|
||||
region_multiplier: Some(1.1),
|
||||
..base
|
||||
};
|
||||
let result = calculate(&pricing(standard), ®ional).unwrap();
|
||||
assert_eq!(result.input(), (65.0 * 2.0 + 25.0 * 0.5 + 10.0 * 3.0) * 1.1);
|
||||
assert_eq!(result.output(), 20.0 * 4.0 * 1.1);
|
||||
let invalid = Request {
|
||||
region_multiplier: Some(f64::NAN),
|
||||
..base
|
||||
};
|
||||
assert_eq!(
|
||||
calculate(&pricing(standard), &invalid),
|
||||
Err(PricingError::InvalidRegionMultiplier)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn compilation_sorts_thresholds_and_rejects_duplicate_tiers() {
|
||||
let high = ThresholdRates {
|
||||
above_prompt_tokens: 200,
|
||||
standard: rates(Rate::Value(7.0), Rate::Missing),
|
||||
tiers: &[],
|
||||
};
|
||||
let low = ThresholdRates {
|
||||
above_prompt_tokens: 100,
|
||||
standard: rates(Rate::Value(5.0), Rate::Missing),
|
||||
tiers: &[],
|
||||
};
|
||||
let specification = Pricing {
|
||||
standard: rates(Rate::Value(2.0), Rate::Value(4.0)),
|
||||
tiers: &[],
|
||||
thresholds: &[high, low],
|
||||
off_peak: None,
|
||||
};
|
||||
let base = request();
|
||||
let above_both = Request {
|
||||
usage: Usage {
|
||||
prompt_tokens: 201,
|
||||
cache_read_tokens: 0,
|
||||
cache_write_tokens: 0,
|
||||
..base.usage
|
||||
},
|
||||
..base
|
||||
};
|
||||
assert_eq!(
|
||||
compile(&specification)
|
||||
.unwrap()
|
||||
.calculate(&above_both)
|
||||
.unwrap()
|
||||
.rates
|
||||
.input,
|
||||
7.0
|
||||
);
|
||||
let duplicate = TierRates {
|
||||
tier: ServiceTier::Flex,
|
||||
rates: Rates::EMPTY,
|
||||
};
|
||||
let invalid = Pricing {
|
||||
tiers: &[duplicate, duplicate],
|
||||
thresholds: &[],
|
||||
..specification
|
||||
};
|
||||
assert_eq!(compile(&invalid).err(), Some(PricingError::DuplicateTier));
|
||||
}
|
||||
96
litellm-rust/crates/cost/tests/generate_python_reference.py
Normal file
96
litellm-rust/crates/cost/tests/generate_python_reference.py
Normal file
|
|
@ -0,0 +1,96 @@
|
|||
import subprocess
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
|
||||
from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token
|
||||
from litellm.types.utils import Usage
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Case:
|
||||
name: str
|
||||
prompt: int
|
||||
completion: int
|
||||
cache_read: int
|
||||
cache_write: int
|
||||
input_rate: float
|
||||
output_rate: float
|
||||
cache_read_rate: float | None = None
|
||||
cache_write_rate: float | None = None
|
||||
threshold: int | None = None
|
||||
threshold_input_rate: float | None = None
|
||||
threshold_output_rate: float | None = None
|
||||
|
||||
|
||||
CASES = (
|
||||
Case("ordinary", 100, 20, 0, 0, 2.0, 4.0),
|
||||
Case("cache_fallback", 100, 20, 25, 10, 2.0, 4.0),
|
||||
Case("cache_specific", 100, 20, 25, 10, 2.0, 4.0, 0.5, 3.0),
|
||||
Case("free_cache", 100, 20, 25, 10, 2.0, 4.0, 0.0, 0.0),
|
||||
Case("threshold_below", 99, 20, 0, 0, 2.0, 4.0, threshold=100, threshold_input_rate=5.0, threshold_output_rate=8.0),
|
||||
Case("threshold_at", 100, 20, 0, 0, 2.0, 4.0, threshold=100, threshold_input_rate=5.0, threshold_output_rate=8.0),
|
||||
Case(
|
||||
"threshold_above", 101, 20, 0, 0, 2.0, 4.0, threshold=100, threshold_input_rate=5.0, threshold_output_rate=8.0
|
||||
),
|
||||
Case(
|
||||
"cache_threshold_above",
|
||||
101,
|
||||
20,
|
||||
25,
|
||||
10,
|
||||
2.0,
|
||||
4.0,
|
||||
threshold=100,
|
||||
threshold_input_rate=5.0,
|
||||
threshold_output_rate=8.0,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def reference(case: Case) -> tuple[float, float]:
|
||||
info = {"input_cost_per_token": case.input_rate, "output_cost_per_token": case.output_rate}
|
||||
if case.cache_read_rate is not None:
|
||||
info["cache_read_input_token_cost"] = case.cache_read_rate
|
||||
if case.cache_write_rate is not None:
|
||||
info["cache_creation_input_token_cost"] = case.cache_write_rate
|
||||
if case.threshold is not None:
|
||||
info[f"input_cost_per_token_above_{case.threshold}_tokens"] = case.threshold_input_rate
|
||||
info[f"output_cost_per_token_above_{case.threshold}_tokens"] = case.threshold_output_rate
|
||||
details = {"cached_tokens": case.cache_read, "cache_write_tokens": case.cache_write}
|
||||
usage = Usage(prompt_tokens=case.prompt, completion_tokens=case.completion, prompt_tokens_details=details)
|
||||
return generic_cost_per_token(
|
||||
model="synthetic",
|
||||
usage=usage,
|
||||
custom_llm_provider="openai",
|
||||
model_info=info,
|
||||
)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
revision = subprocess.check_output(("git", "rev-parse", "HEAD"), text=True).strip()
|
||||
rows = ("# Python reference commit: " + revision,) + tuple(
|
||||
"\t".join(
|
||||
str(value) if value is not None else ""
|
||||
for value in (
|
||||
case.name,
|
||||
case.prompt,
|
||||
case.completion,
|
||||
case.cache_read,
|
||||
case.cache_write,
|
||||
case.input_rate,
|
||||
case.output_rate,
|
||||
case.cache_read_rate,
|
||||
case.cache_write_rate,
|
||||
case.threshold,
|
||||
case.threshold_input_rate,
|
||||
case.threshold_output_rate,
|
||||
*reference(case),
|
||||
)
|
||||
)
|
||||
for case in CASES
|
||||
)
|
||||
Path(__file__).with_name("python_reference.tsv").write_text("\n".join(rows) + "\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
9
litellm-rust/crates/cost/tests/python_reference.tsv
Normal file
9
litellm-rust/crates/cost/tests/python_reference.tsv
Normal file
|
|
@ -0,0 +1,9 @@
|
|||
# Python reference commit: dc4be2fd987c993aefcf16444e34f960c12c8627
|
||||
ordinary 100 20 0 0 2.0 4.0 200.0 80.0
|
||||
cache_fallback 100 20 25 10 2.0 4.0 200.0 80.0
|
||||
cache_specific 100 20 25 10 2.0 4.0 0.5 3.0 172.5 80.0
|
||||
free_cache 100 20 25 10 2.0 4.0 0.0 0.0 130.0 80.0
|
||||
threshold_below 99 20 0 0 2.0 4.0 100 5.0 8.0 198.0 80.0
|
||||
threshold_at 100 20 0 0 2.0 4.0 100 5.0 8.0 200.0 80.0
|
||||
threshold_above 101 20 0 0 2.0 4.0 100 5.0 8.0 505.0 160.0
|
||||
cache_threshold_above 101 20 25 10 2.0 4.0 100 5.0 8.0 505.0 160.0
|
||||
|
Can't render this file because it has a wrong number of fields in line 2.
|
|
|
@ -27,7 +27,9 @@ pub use execution::{
|
|||
pub use fork_gate::RuntimeAlreadyStarted;
|
||||
pub use gil::{release_count, release_gil};
|
||||
pub use handle::{Execution, ExecutionBody, ExecutionStep};
|
||||
pub use marshal::{Pythonized, from_py, from_py_argument, panic_to_pyerr, to_py};
|
||||
pub use marshal::{
|
||||
Pythonized, from_py, from_py_argument, json_loads, json_object_field, panic_to_pyerr, to_py,
|
||||
};
|
||||
|
||||
/// Starts the interpreter and imports the standard modules the tests share, once, so
|
||||
/// parallel test threads never race a first import of `asyncio`.
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ use std::panic::{AssertUnwindSafe, catch_unwind};
|
|||
use pyo3::exceptions::PyValueError;
|
||||
use pyo3::panic::PanicException;
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::PyBytes;
|
||||
use serde::Serialize;
|
||||
use serde::de::DeserializeOwned;
|
||||
|
||||
|
|
@ -32,6 +33,19 @@ where
|
|||
.map_err(PyErr::from)
|
||||
}
|
||||
|
||||
pub fn json_object_field(py: Python<'_>, document: &str, name: &str) -> PyResult<Py<PyAny>> {
|
||||
py.import("json")?
|
||||
.call_method1("loads", (document,))?
|
||||
.call_method1("get", (name,))
|
||||
.map(Bound::unbind)
|
||||
}
|
||||
|
||||
pub fn json_loads(py: Python<'_>, document: &[u8]) -> PyResult<Py<PyAny>> {
|
||||
py.import("json")?
|
||||
.call_method1("loads", (PyBytes::new(py, document),))
|
||||
.map(Bound::unbind)
|
||||
}
|
||||
|
||||
pub struct Pythonized<T>(pub T);
|
||||
|
||||
impl<'py, T> IntoPyObject<'py> for Pythonized<T>
|
||||
|
|
|
|||
|
|
@ -1 +0,0 @@
|
|||
pub mod secrets;
|
||||
|
|
@ -1,19 +0,0 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use futures_util::future::BoxFuture;
|
||||
use litellm_core_utils::settings::{Lookup, ProcessEnvironment};
|
||||
use litellm_secrets::Error;
|
||||
|
||||
pub type Secrets = Arc<dyn Lookup + Send + Sync>;
|
||||
|
||||
pub trait SecretSource: Send + Sync {
|
||||
fn resolve<'a>(&'a self, names: &'a [&'static str]) -> BoxFuture<'a, Result<Secrets, Error>>;
|
||||
}
|
||||
|
||||
pub struct EnvironmentSecrets;
|
||||
|
||||
impl SecretSource for EnvironmentSecrets {
|
||||
fn resolve<'a>(&'a self, _names: &'a [&'static str]) -> BoxFuture<'a, Result<Secrets, Error>> {
|
||||
Box::pin(async { Ok(Arc::new(ProcessEnvironment) as Secrets) })
|
||||
}
|
||||
}
|
||||
|
|
@ -2,6 +2,5 @@ pub mod anthropic_messages;
|
|||
pub mod audio_transcription;
|
||||
pub mod base_model_iterator;
|
||||
pub mod chat;
|
||||
pub mod inference;
|
||||
pub mod ocr;
|
||||
pub mod responses;
|
||||
|
|
|
|||
|
|
@ -13,7 +13,6 @@ use litellm_http::{
|
|||
use serde::{Serialize, de::DeserializeOwned};
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::base_llm::inference::secrets::SecretSource;
|
||||
use crate::base_llm::ocr::{
|
||||
error::Error,
|
||||
settings::OcrSettings,
|
||||
|
|
@ -22,6 +21,7 @@ use crate::base_llm::ocr::{
|
|||
PreparedOcrRequest, decode_request_value, decode_response,
|
||||
},
|
||||
};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
|
||||
/// The route's view of one call, handed to provider code that has to reach the
|
||||
/// caller's hooks mid-flight (guardrails on the outgoing body, raw response events).
|
||||
|
|
@ -95,7 +95,7 @@ impl OcrClient {
|
|||
document_fetcher: MediaFetcher::for_test(document_http),
|
||||
vertex_auth: VertexAuth::default(),
|
||||
settings: OcrSettings::default(),
|
||||
secrets: Arc::new(crate::base_llm::inference::secrets::EnvironmentSecrets),
|
||||
secrets: Arc::new(litellm_secrets::source::EnvironmentSecrets::default()),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ use litellm_core_utils::{
|
|||
settings::ProcessEnvironment,
|
||||
};
|
||||
use litellm_http::outbound::{OutboundRequest, RequestSigner};
|
||||
use litellm_secrets::source::Secrets;
|
||||
use serde::{
|
||||
Deserialize, Serialize,
|
||||
de::{DeserializeOwned, IntoDeserializer},
|
||||
|
|
@ -14,13 +15,10 @@ use serde::{
|
|||
use serde_json::{Map, Value};
|
||||
use serde_with::serde_as;
|
||||
|
||||
use crate::base_llm::{
|
||||
inference::secrets::Secrets,
|
||||
ocr::{
|
||||
error::Error,
|
||||
handler::{CallHooks, OcrClient, read_response_bytes, transform_request_body},
|
||||
settings::OcrSettings,
|
||||
},
|
||||
use crate::base_llm::ocr::{
|
||||
error::Error,
|
||||
handler::{CallHooks, OcrClient, read_response_bytes, transform_request_body},
|
||||
settings::OcrSettings,
|
||||
};
|
||||
|
||||
pub const OCR_RESPONSE_MAX_BYTES: usize = 64 * 1024 * 1024;
|
||||
|
|
|
|||
25
litellm-rust/crates/model-catalog/Cargo.toml
Normal file
25
litellm-rust/crates/model-catalog/Cargo.toml
Normal file
|
|
@ -0,0 +1,25 @@
|
|||
[package]
|
||||
name = "litellm-model-catalog"
|
||||
version = "0.1.0"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
|
||||
[features]
|
||||
schema = ["dep:schemars"]
|
||||
|
||||
[dependencies]
|
||||
indexmap = { version = "2.14.0", features = ["serde"] }
|
||||
schemars = { version = "1.0", optional = true }
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
thiserror.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
criterion.workspace = true
|
||||
rstest.workspace = true
|
||||
litellm-model-catalog = { path = ".", features = ["schema"] }
|
||||
|
||||
[[bench]]
|
||||
name = "catalog"
|
||||
harness = false
|
||||
25
litellm-rust/crates/model-catalog/README.md
Normal file
25
litellm-rust/crates/model-catalog/README.md
Normal file
|
|
@ -0,0 +1,25 @@
|
|||
# Model catalog
|
||||
|
||||
`litellm-model-catalog` builds an immutable snapshot from caller supplied JSON bytes. It has no network, Python, registration, or refresh behavior. The caller supplies optional source, revision, and ETag provenance. Parse and validation are separate so small synthetic catalogs can use explicit integrity limits
|
||||
|
||||
The parser treats `sample_spec` and `fallback_generalizations` as reserved top level metadata. `fallback_rules()` exposes the typed rule array when present; this crate does not execute regex generalizations. Model entries retain all JSON fields except `aliases`, including unknown fields. `field()` returns `None` for an absent key and a JSON null, false, or zero value for a present key. The returned values are borrowed, so callers cannot mutate the snapshot
|
||||
|
||||
Each entry also deserializes into `ModelInfo`, a typed mirror of `model_prices_and_context_window.schema.json`'s `modelEntry` definition, reachable via `ModelEntry::info()`. All schema fields are optional on `ModelInfo`, including `litellm_provider` which the schema marks required, so small synthetic catalogs still parse. Unknown fields are not part of `ModelInfo`; they remain on `fields()`. Building with the `schema` feature adds `schemars` derives and exposes `model_entry_json_schema()` for emitting the entry's JSON Schema. Parse and validation failures are reported by the `Error` enum in `error.rs`, while catalog logic lives in `catalog.rs`
|
||||
|
||||
The integration tests read the repository's catalog and schema files at test time, assert every entry round-trips through `ModelInfo`, and verify that the generated schema's properties match the repository schema
|
||||
|
||||
Aliases point to their canonical entries. An alias that exactly matches any canonical key is skipped; the first canonical entry claiming an alias wins. Invalid alias lists and nonstring names are skipped and reported by `alias_issues()`. Exact lookup wins. For a case insensitive miss, the last key with the same lowercase spelling wins, following Python's lowercase map built after aliases are appended. This uses Rust Unicode lowercasing, which can differ from Python for unusual Unicode model IDs
|
||||
|
||||
`validate()` counts canonical entries before alias expansion and excludes both reserved keys. It enforces an explicit minimum and backup shrink ratio, with Python defaults of 50 models and 0.5. Parsing rejects nonobject model entries and known fields with the wrong JSON type, but ignores unknown fields. It does not enforce every constraint in the JSON schema, calculate prices, resolve providers, or check provenance authenticity. The caller decides how to handle validation failures
|
||||
|
||||
This snapshot does not represent Python's live mutable `litellm.model_cost`, nested dict and list mutation, or mutation of dicts previously returned by Python APIs. It has no bridge or runtime integration
|
||||
|
||||
## Benchmarks
|
||||
|
||||
`cargo bench -p litellm-model-catalog --bench catalog` measures parsing plus alias indexing and exact lookup. For a local Python baseline on the same fixture, use:
|
||||
|
||||
```sh
|
||||
python3 -m timeit -s 'import json, pathlib; body = pathlib.Path("../model_prices_and_context_window.json").read_bytes()' 'json.loads(body)'
|
||||
```
|
||||
|
||||
Run these commands from `litellm-rust`. Python's command measures JSON loading only, without alias expansion or snapshot construction. The Rust benchmark does not include future Python object materialization, so these numbers are not an end to end runtime comparison
|
||||
21
litellm-rust/crates/model-catalog/benches/catalog.rs
Normal file
21
litellm-rust/crates/model-catalog/benches/catalog.rs
Normal file
|
|
@ -0,0 +1,21 @@
|
|||
use criterion::{Criterion, criterion_group, criterion_main};
|
||||
use litellm_model_catalog::{Catalog, Provenance};
|
||||
use std::hint::black_box;
|
||||
|
||||
fn benchmarks(c: &mut Criterion) {
|
||||
let body = include_bytes!("../../../../model_prices_and_context_window.json");
|
||||
c.bench_function("parse_current_catalog", |b| {
|
||||
b.iter(|| Catalog::parse(black_box(body), Provenance::default()).unwrap())
|
||||
});
|
||||
let catalog = Catalog::parse(body, Provenance::default()).unwrap();
|
||||
let key = catalog
|
||||
.model_names()
|
||||
.next()
|
||||
.expect("catalog must have a benchmark key");
|
||||
c.bench_function("lookup_catalog_key", |b| {
|
||||
b.iter(|| black_box(&catalog).lookup(black_box(key)))
|
||||
});
|
||||
}
|
||||
|
||||
criterion_group!(benches, benchmarks);
|
||||
criterion_main!(benches);
|
||||
241
litellm-rust/crates/model-catalog/src/catalog.rs
Normal file
241
litellm-rust/crates/model-catalog/src/catalog.rs
Normal file
|
|
@ -0,0 +1,241 @@
|
|||
use crate::error::Error;
|
||||
use crate::model_info::{FallbackGeneralizations, FallbackRule, ModelInfo};
|
||||
use indexmap::IndexMap;
|
||||
use serde::Deserialize;
|
||||
use serde_json::{Map, Value};
|
||||
use std::collections::HashMap;
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Eq)]
|
||||
pub struct Provenance {
|
||||
pub source: Option<String>,
|
||||
pub revision: Option<String>,
|
||||
pub etag: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq)]
|
||||
pub struct IntegrityLimits {
|
||||
pub backup_model_count: usize,
|
||||
pub min_model_count: usize,
|
||||
pub min_backup_ratio: f64,
|
||||
}
|
||||
|
||||
impl IntegrityLimits {
|
||||
pub fn python_defaults(backup_model_count: usize) -> Self {
|
||||
Self {
|
||||
backup_model_count,
|
||||
min_model_count: 50,
|
||||
min_backup_ratio: 0.5,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum AliasIssue {
|
||||
InvalidList { model: String },
|
||||
InvalidName { model: String },
|
||||
CanonicalCollision { model: String, alias: String },
|
||||
AliasCollision { model: String, alias: String },
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct ModelEntry {
|
||||
fields: Map<String, Value>,
|
||||
info: ModelInfo,
|
||||
}
|
||||
|
||||
impl ModelEntry {
|
||||
pub fn field(&self, name: &str) -> Option<&Value> {
|
||||
self.fields.get(name)
|
||||
}
|
||||
pub fn fields(&self) -> &Map<String, Value> {
|
||||
&self.fields
|
||||
}
|
||||
/// The entry deserialized into the typed mirror of the catalog schema.
|
||||
pub fn info(&self) -> &ModelInfo {
|
||||
&self.info
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug)]
|
||||
pub struct ModelMatch<'a> {
|
||||
pub matched_key: &'a str,
|
||||
pub canonical_key: &'a str,
|
||||
pub entry: &'a ModelEntry,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct Catalog {
|
||||
entries: IndexMap<String, ModelEntry>,
|
||||
aliases: IndexMap<String, String>,
|
||||
lowercase_keys: HashMap<String, String>,
|
||||
sample_spec: Option<Value>,
|
||||
fallback_generalizations: Option<FallbackGeneralizations>,
|
||||
provenance: Provenance,
|
||||
alias_issues: Vec<AliasIssue>,
|
||||
}
|
||||
|
||||
impl Catalog {
|
||||
pub fn parse(body: &[u8], provenance: Provenance) -> Result<Self, Error> {
|
||||
let root: IndexMap<String, Value> = serde_json::from_slice(body)?;
|
||||
if root.is_empty() {
|
||||
return Err(Error::Empty);
|
||||
}
|
||||
|
||||
let mut entries = IndexMap::with_capacity(root.len());
|
||||
let mut alias_lists = Vec::new();
|
||||
let mut alias_issues = Vec::new();
|
||||
let mut sample_spec = None;
|
||||
let mut fallback_generalizations = None;
|
||||
for (name, value) in root {
|
||||
match name.as_str() {
|
||||
"sample_spec" => {
|
||||
sample_spec = Some(value);
|
||||
continue;
|
||||
}
|
||||
"fallback_generalizations" => {
|
||||
fallback_generalizations =
|
||||
Some(serde_json::from_value::<FallbackGeneralizations>(value)?);
|
||||
continue;
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
let Value::Object(ref object) = value else {
|
||||
return Err(Error::EntryNotObject { model: name });
|
||||
};
|
||||
let info = ModelInfo::deserialize(object)?;
|
||||
let Value::Object(mut fields) = value else {
|
||||
unreachable!("value checked is_object above")
|
||||
};
|
||||
if let Some(aliases) = fields.remove("aliases")
|
||||
&& !aliases.is_null()
|
||||
{
|
||||
match aliases {
|
||||
Value::Array(names) => alias_lists.push((name.clone(), names)),
|
||||
_ => alias_issues.push(AliasIssue::InvalidList {
|
||||
model: name.clone(),
|
||||
}),
|
||||
}
|
||||
}
|
||||
entries.insert(name, ModelEntry { fields, info });
|
||||
}
|
||||
|
||||
let mut aliases = IndexMap::new();
|
||||
for (model, names) in alias_lists {
|
||||
for name in names {
|
||||
let Value::String(alias) = name else {
|
||||
alias_issues.push(AliasIssue::InvalidName {
|
||||
model: model.clone(),
|
||||
});
|
||||
continue;
|
||||
};
|
||||
if entries.contains_key(&alias) {
|
||||
alias_issues.push(AliasIssue::CanonicalCollision {
|
||||
model: model.clone(),
|
||||
alias,
|
||||
});
|
||||
} else if aliases.contains_key(&alias) {
|
||||
alias_issues.push(AliasIssue::AliasCollision {
|
||||
model: model.clone(),
|
||||
alias,
|
||||
});
|
||||
} else {
|
||||
aliases.insert(alias, model.clone());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let lowercase_keys = entries
|
||||
.keys()
|
||||
.chain(aliases.keys())
|
||||
.map(|key| (key.to_lowercase(), key.clone()))
|
||||
.collect();
|
||||
Ok(Self {
|
||||
entries,
|
||||
aliases,
|
||||
lowercase_keys,
|
||||
sample_spec,
|
||||
fallback_generalizations,
|
||||
provenance,
|
||||
alias_issues,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn validate(&self, limits: IntegrityLimits) -> Result<(), Error> {
|
||||
if !limits.min_backup_ratio.is_finite() || !(0.0..=1.0).contains(&limits.min_backup_ratio) {
|
||||
return Err(Error::InvalidRatio);
|
||||
}
|
||||
let actual = self.entries.len();
|
||||
if actual < limits.min_model_count {
|
||||
return Err(Error::BelowMinimum {
|
||||
actual,
|
||||
minimum: limits.min_model_count,
|
||||
});
|
||||
}
|
||||
if limits.backup_model_count > 0
|
||||
&& (actual as f64) < (limits.backup_model_count as f64) * limits.min_backup_ratio
|
||||
{
|
||||
return Err(Error::Shrunk {
|
||||
actual,
|
||||
backup: limits.backup_model_count,
|
||||
ratio: limits.min_backup_ratio,
|
||||
});
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn lookup(&self, key: &str) -> Option<ModelMatch<'_>> {
|
||||
let matched_key = if self.entries.contains_key(key) || self.aliases.contains_key(key) {
|
||||
key
|
||||
} else {
|
||||
self.lowercase_keys.get(&key.to_lowercase())?.as_str()
|
||||
};
|
||||
let canonical_key = self
|
||||
.aliases
|
||||
.get(matched_key)
|
||||
.map(String::as_str)
|
||||
.unwrap_or(matched_key);
|
||||
let (canonical_key, entry) = self.entries.get_key_value(canonical_key)?;
|
||||
let matched_key = self
|
||||
.entries
|
||||
.get_key_value(matched_key)
|
||||
.map(|(key, _)| key.as_str())
|
||||
.or_else(|| {
|
||||
self.aliases
|
||||
.get_key_value(matched_key)
|
||||
.map(|(key, _)| key.as_str())
|
||||
})?;
|
||||
Some(ModelMatch {
|
||||
matched_key,
|
||||
canonical_key,
|
||||
entry,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn model_count(&self) -> usize {
|
||||
self.entries.len()
|
||||
}
|
||||
pub fn model_names(&self) -> impl Iterator<Item = &str> {
|
||||
self.entries.keys().map(String::as_str)
|
||||
}
|
||||
pub fn alias_count(&self) -> usize {
|
||||
self.aliases.len()
|
||||
}
|
||||
pub fn aliases(&self) -> &IndexMap<String, String> {
|
||||
&self.aliases
|
||||
}
|
||||
pub fn alias_issues(&self) -> &[AliasIssue] {
|
||||
&self.alias_issues
|
||||
}
|
||||
pub fn sample_spec(&self) -> Option<&Value> {
|
||||
self.sample_spec.as_ref()
|
||||
}
|
||||
pub fn fallback_generalizations(&self) -> Option<&FallbackGeneralizations> {
|
||||
self.fallback_generalizations.as_ref()
|
||||
}
|
||||
pub fn fallback_rules(&self) -> Option<&[FallbackRule]> {
|
||||
Some(self.fallback_generalizations.as_ref()?.rules.as_slice())
|
||||
}
|
||||
pub fn provenance(&self) -> &Provenance {
|
||||
&self.provenance
|
||||
}
|
||||
}
|
||||
28
litellm-rust/crates/model-catalog/src/error.rs
Normal file
28
litellm-rust/crates/model-catalog/src/error.rs
Normal file
|
|
@ -0,0 +1,28 @@
|
|||
use thiserror::Error;
|
||||
|
||||
/// Failures from parsing or validating a catalog snapshot.
|
||||
#[derive(Debug, Error)]
|
||||
pub enum Error {
|
||||
/// The body is not valid JSON, or a model entry fails typed deserialization.
|
||||
#[error("invalid JSON: {0}")]
|
||||
Json(#[from] serde_json::Error),
|
||||
/// The catalog has no entries at all.
|
||||
#[error("catalog is empty")]
|
||||
Empty,
|
||||
/// A non-reserved top level value is not a JSON object.
|
||||
#[error("model {model:?} must be an object")]
|
||||
EntryNotObject { model: String },
|
||||
/// Canonical entry count is under the configured minimum.
|
||||
#[error("catalog has {actual} models, below minimum {minimum}")]
|
||||
BelowMinimum { actual: usize, minimum: usize },
|
||||
/// Canonical entry count is under the configured backup shrink ratio.
|
||||
#[error("catalog has {actual} models, below {ratio} of backup count {backup}")]
|
||||
Shrunk {
|
||||
actual: usize,
|
||||
backup: usize,
|
||||
ratio: f64,
|
||||
},
|
||||
/// The configured minimum backup ratio is not finite or outside `[0, 1]`.
|
||||
#[error("minimum backup ratio must be finite and between zero and one")]
|
||||
InvalidRatio,
|
||||
}
|
||||
16
litellm-rust/crates/model-catalog/src/lib.rs
Normal file
16
litellm-rust/crates/model-catalog/src/lib.rs
Normal file
|
|
@ -0,0 +1,16 @@
|
|||
mod catalog;
|
||||
mod error;
|
||||
mod model_info;
|
||||
#[cfg(feature = "schema")]
|
||||
mod schema;
|
||||
|
||||
pub use catalog::{AliasIssue, Catalog, IntegrityLimits, ModelEntry, ModelMatch, Provenance};
|
||||
pub use error::Error;
|
||||
pub use model_info::{
|
||||
AudioFormat, FallbackGeneralizations, FallbackRule, InputModality, Mode, ModelInfo,
|
||||
OffPeakPricing, OffPeakWindow, OutputModality, ReasoningEffort, SearchContextCostPerQuery,
|
||||
TieredRate, UtcHours, VertexAiAudioApi, WebSearchBillingUnit, Weekday,
|
||||
};
|
||||
|
||||
#[cfg(feature = "schema")]
|
||||
pub use schema::model_entry_json_schema;
|
||||
665
litellm-rust/crates/model-catalog/src/model_info.rs
Normal file
665
litellm-rust/crates/model-catalog/src/model_info.rs
Normal file
|
|
@ -0,0 +1,665 @@
|
|||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
use std::collections::BTreeMap;
|
||||
|
||||
/// Primary API surface / task type of the model.
|
||||
#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
||||
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum Mode {
|
||||
AudioSpeech,
|
||||
AudioTranscription,
|
||||
Chat,
|
||||
Completion,
|
||||
Embedding,
|
||||
Evaluation,
|
||||
Guardrail,
|
||||
ImageEdit,
|
||||
ImageGeneration,
|
||||
Moderation,
|
||||
Ocr,
|
||||
Realtime,
|
||||
Rerank,
|
||||
Responses,
|
||||
Search,
|
||||
VectorStore,
|
||||
VideoGeneration,
|
||||
}
|
||||
|
||||
/// Reasoning effort level accepted or applied by the model.
|
||||
#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
||||
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum ReasoningEffort {
|
||||
None,
|
||||
Minimal,
|
||||
Low,
|
||||
Medium,
|
||||
High,
|
||||
Xhigh,
|
||||
Max,
|
||||
}
|
||||
|
||||
/// Gemini audio generation API the model is served through.
|
||||
#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
||||
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum VertexAiAudioApi {
|
||||
LyriaPredict,
|
||||
LyriaInteractions,
|
||||
}
|
||||
|
||||
/// Whether web search is billed per query or per prompt.
|
||||
#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
||||
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum WebSearchBillingUnit {
|
||||
PerQuery,
|
||||
PerPrompt,
|
||||
}
|
||||
|
||||
/// Audio container format the model can return.
|
||||
#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
||||
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum AudioFormat {
|
||||
Mp3,
|
||||
Wav,
|
||||
}
|
||||
|
||||
/// Input modality the model accepts.
|
||||
#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
||||
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum InputModality {
|
||||
Text,
|
||||
Image,
|
||||
Audio,
|
||||
Video,
|
||||
}
|
||||
|
||||
/// Output modality the model can produce.
|
||||
#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
||||
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum OutputModality {
|
||||
Text,
|
||||
Image,
|
||||
Audio,
|
||||
Video,
|
||||
Code,
|
||||
}
|
||||
|
||||
/// UTC "HH:MM-HH:MM" window, or a list of them; a window may wrap past midnight.
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
||||
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
|
||||
#[serde(untagged)]
|
||||
pub enum UtcHours {
|
||||
Single(String),
|
||||
Multiple(Vec<String>),
|
||||
}
|
||||
|
||||
/// ISO-8601 weekday number (1 = Monday .. 7 = Sunday) or English day name.
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
||||
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
|
||||
#[serde(untagged)]
|
||||
pub enum Weekday {
|
||||
Number(u8),
|
||||
Name(String),
|
||||
}
|
||||
|
||||
/// One off-peak window entry inside `windows`.
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
||||
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct OffPeakWindow {
|
||||
pub hours_utc: UtcHours,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub weekdays: Option<Vec<Weekday>>,
|
||||
}
|
||||
|
||||
/// Rates that replace the same-named base fields inside the stated UTC windows.
|
||||
#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
|
||||
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct OffPeakPricing {
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub hours_utc: Option<UtcHours>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub windows: Option<Vec<OffPeakWindow>>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub weekday_timezone: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_reasoning_token: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost: Option<f64>,
|
||||
}
|
||||
|
||||
/// USD cost per web search query, keyed by search context size.
|
||||
#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
|
||||
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct SearchContextCostPerQuery {
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub search_context_size_low: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub search_context_size_medium: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub search_context_size_high: Option<f64>,
|
||||
}
|
||||
|
||||
/// One tier of a context-length or result-count tiered rate.
|
||||
#[derive(Clone, Debug, Default, Deserialize, PartialEq, Serialize)]
|
||||
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct TieredRate {
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub range: Option<[f64; 2]>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub max_results_range: Option<[f64; 2]>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_reasoning_token: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_query: Option<f64>,
|
||||
}
|
||||
|
||||
/// One regex rule generalizing unknown model ids to known families.
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
||||
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
|
||||
pub struct FallbackRule {
|
||||
pub name: String,
|
||||
pub pattern: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub description: Option<String>,
|
||||
#[serde(flatten)]
|
||||
pub extra: BTreeMap<String, Value>,
|
||||
}
|
||||
|
||||
/// Regex rules that generalize unknown model ids to known families; not a model entry.
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
||||
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct FallbackGeneralizations {
|
||||
pub rules: Vec<FallbackRule>,
|
||||
}
|
||||
|
||||
/// Typed mirror of one catalog model entry.
|
||||
#[derive(Clone, Debug, Deserialize, PartialEq, Serialize)]
|
||||
#[cfg_attr(feature = "schema", derive(schemars::JsonSchema))]
|
||||
pub struct ModelInfo {
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub annotation_cost_per_page: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub annotation_cost_per_page_batches: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub audio_transcription_config: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub bedrock_converse_supports_strict_tools: Option<bool>,
|
||||
/// Highest reasoning effort the Bedrock output_config accepts for this model.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub bedrock_output_config_effort_ceiling: Option<ReasoningEffort>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_audio_token_cost: Option<f64>,
|
||||
/// USD per token written to the provider's prompt cache.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost_above_128k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost_above_1hr: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost_above_1hr_above_200k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost_above_200k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost_above_256k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost_above_272k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost_above_272k_tokens_batches: Option<f64>,
|
||||
/// Flex service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost_above_272k_tokens_flex: Option<f64>,
|
||||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost_above_272k_tokens_priority: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost_batches: Option<f64>,
|
||||
/// Flex service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost_flex: Option<f64>,
|
||||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_creation_input_token_cost_priority: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_audio_token_cost: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_image_token_cost: Option<f64>,
|
||||
/// USD per prompt token served from the provider's prompt cache.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_above_128k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_above_200k_tokens: Option<f64>,
|
||||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_above_200k_tokens_priority: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_above_256k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_above_272k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_above_272k_tokens_batches: Option<f64>,
|
||||
/// Flex service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_above_272k_tokens_flex: Option<f64>,
|
||||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_above_272k_tokens_priority: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_above_512k_tokens: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_batches: Option<f64>,
|
||||
/// Flex service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_flex: Option<f64>,
|
||||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub cache_read_input_token_cost_priority: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub citation_cost_per_token: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub code_interpreter_cost_per_session: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub comment: Option<String>,
|
||||
/// Reasoning effort the provider applies when the request omits reasoning_effort. Gates whether a non-default temperature or the top_p/logprobs sampling params are accepted, which hold only when the effort resolves to 'none'.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub default_reasoning_effort: Option<ReasoningEffort>,
|
||||
/// Date the provider deprecates the model, YYYY-MM-DD.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub deprecation_date: Option<String>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub gemini_audio_only_live: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub gemini_native_audio: Option<bool>,
|
||||
/// USD per Grounding with Google Maps request; billed per query or per prompt per web_search_billing_unit.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub google_maps_grounding_cost_per_query: Option<f64>,
|
||||
/// USD cost per billable guardrail unit, keyed by the provider's usage counter name (e.g. Bedrock's contentPolicyUnits).
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub guardrail_cost_per_unit: Option<BTreeMap<String, f64>>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_audio_per_second: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_audio_per_second_above_128k_tokens: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_audio_token: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_audio_token_batches: Option<f64>,
|
||||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_audio_token_priority: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_character: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_character_above_128k_tokens: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_image: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_image_above_128k_tokens: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_image_token: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_image_token_batches: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_pixel: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_query: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_request: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_second: Option<f64>,
|
||||
/// USD per prompt token.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_above_128k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_above_200k_tokens: Option<f64>,
|
||||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_above_200k_tokens_priority: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_above_256k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_above_272k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_above_272k_tokens_batches: Option<f64>,
|
||||
/// Flex service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_above_272k_tokens_flex: Option<f64>,
|
||||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_above_272k_tokens_priority: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_above_512k_tokens: Option<f64>,
|
||||
/// USD per prompt token via the provider's batch API.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_batches: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_cache_hit: Option<f64>,
|
||||
/// Flex service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_flex: Option<f64>,
|
||||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_token_priority: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_video_per_second: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_video_per_second_above_128k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_video_per_second_above_15s_interval: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_video_per_second_above_8s_interval: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_video_token: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_cost_per_video_token_batches: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub input_dbu_cost_per_token: Option<f64>,
|
||||
/// LiteLLM provider slug; one of https://docs.litellm.ai/docs/providers.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub litellm_provider: Option<String>,
|
||||
/// Maximum prompt/context tokens the model accepts.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub max_input_tokens: Option<u64>,
|
||||
/// Maximum tokens the model can generate in one response.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub max_output_tokens: Option<u64>,
|
||||
/// Legacy field: max output tokens if the provider specifies it, else max input tokens.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub max_tokens: Option<u64>,
|
||||
/// Free-form notes about the entry (e.g. pricing derivation).
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub metadata: Option<BTreeMap<String, Value>>,
|
||||
/// Primary API surface / task type of the model.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub mode: Option<Mode>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub ocr_cost_per_credit: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub ocr_cost_per_page: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub ocr_cost_per_page_batches: Option<f64>,
|
||||
/// Rates that replace the same-named base fields while the request falls inside the stated UTC windows.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub off_peak_pricing: Option<OffPeakPricing>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_audio_token: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_character: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_character_above_128k_tokens: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_image: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_image_1024: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_image_1536: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_image_512: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_image_token: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_pixel: Option<f64>,
|
||||
/// USD per reasoning/thinking token, when billed separately.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_reasoning_token: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_second: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_second_1080p: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_second_2k: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_second_480p: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_second_4k: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_second_720p: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_second_768p: Option<f64>,
|
||||
/// USD per generated token.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_above_128k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_above_200k_tokens: Option<f64>,
|
||||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_above_200k_tokens_priority: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_above_256k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_above_272k_tokens: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_above_272k_tokens_batches: Option<f64>,
|
||||
/// Flex service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_above_272k_tokens_flex: Option<f64>,
|
||||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_above_272k_tokens_priority: Option<f64>,
|
||||
/// Rate applied once the prompt exceeds the token threshold in the field name.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_above_512k_tokens: Option<f64>,
|
||||
/// USD per generated token via the provider's batch API.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_batches: Option<f64>,
|
||||
/// Flex service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_flex: Option<f64>,
|
||||
/// Priority service-tier rate for the same-named base field.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_token_priority: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_video_per_second: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_cost_per_video_token: Option<f64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_dbu_cost_per_token: Option<f64>,
|
||||
/// Embedding dimension for embedding models.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub output_vector_size: Option<u64>,
|
||||
/// Smallest prefix the provider will actually cache; absent means the provider default applies.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub prompt_cache_min_tokens: Option<u64>,
|
||||
/// Provider-internal routing hints (e.g. bedrock_invocation_schema).
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub provider_specific_entry: Option<BTreeMap<String, Value>>,
|
||||
/// Exact reasoning_effort levels this deployment accepts; wins over supports_* flags.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub reasoning_effort_levels: Option<Vec<ReasoningEffort>>,
|
||||
/// Multiplier applied to all token costs when served from a non-global Vertex AI endpoint (e.g. 1.10 = +10%).
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub regional_endpoint_uplift_multiplier: Option<f64>,
|
||||
/// Multiplier applied to all token costs for EU data residency (e.g. 1.10 = +10%).
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub regional_processing_uplift_multiplier_eu: Option<f64>,
|
||||
/// Multiplier applied to all token costs for US data residency (e.g. 1.10 = +10%).
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub regional_processing_uplift_multiplier_us: Option<f64>,
|
||||
/// Provider default requests-per-minute limit.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub rpm: Option<u64>,
|
||||
/// USD cost per web search query, keyed by search context size.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub search_context_cost_per_query: Option<SearchContextCostPerQuery>,
|
||||
/// URL of the provider pricing/model page this entry was taken from.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub source: Option<String>,
|
||||
/// Audio container formats the model can return.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supported_audio_formats: Option<Vec<AudioFormat>>,
|
||||
/// OpenAI-style API routes this model can be called through, e.g. /v1/chat/completions.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supported_endpoints: Option<Vec<String>>,
|
||||
/// Input modalities the model accepts.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supported_modalities: Option<Vec<InputModality>>,
|
||||
/// Output modalities the model can produce.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supported_output_modalities: Option<Vec<OutputModality>>,
|
||||
/// Cloud regions the model is available in ('global' or region ids).
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supported_regions: Option<Vec<String>>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_adaptive_thinking: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_anthropic_compaction: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_anthropic_thinking_payload: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_assistant_prefill: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_audio_input: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_audio_output: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_computer_use: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_embedding_image_input: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_fast_mode: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_forced_tool_use: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_function_calling: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_image_input: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_image_size: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_legacy_thinking: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_low_reasoning_effort: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_max_reasoning_effort: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_mid_conversation_system: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_minimal_reasoning_effort: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_multimodal: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_native_streaming: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_native_structured_output: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_none_reasoning_effort: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_nova_canvas_image_edit: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_output_config: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_parallel_function_calling: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_parallel_tool_use_config: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_pdf_input: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_prompt_cache_breakpoint: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_prompt_caching: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_reasoning: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_response_schema: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_sampling_params: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_speed: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_system_messages: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_thinking_cache_preservation: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_tool_choice: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_tool_search: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_url_context: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_video_input: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_vision: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_web_search: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub supports_xhigh_reasoning_effort: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub thinking_always_on: Option<bool>,
|
||||
/// Context-length or result-count tiered rates; each tier's costs apply within its range.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub tiered_pricing: Option<Vec<TieredRate>>,
|
||||
/// Provider default tokens-per-minute limit.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub tpm: Option<u64>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub use_openai_responses_path: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub uses_embed_content: Option<bool>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub vertex_ai_audio_api: Option<VertexAiAudioApi>,
|
||||
/// Whether web search is billed per query or per prompt.
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub web_search_billing_unit: Option<WebSearchBillingUnit>,
|
||||
}
|
||||
7
litellm-rust/crates/model-catalog/src/schema.rs
Normal file
7
litellm-rust/crates/model-catalog/src/schema.rs
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
use crate::model_info::ModelInfo;
|
||||
|
||||
/// JSON Schema for one catalog model entry, mirroring
|
||||
/// `model_prices_and_context_window.schema.json`'s `modelEntry` definition.
|
||||
pub fn model_entry_json_schema() -> schemars::Schema {
|
||||
schemars::schema_for!(ModelInfo)
|
||||
}
|
||||
253
litellm-rust/crates/model-catalog/tests/catalog.rs
Normal file
253
litellm-rust/crates/model-catalog/tests/catalog.rs
Normal file
|
|
@ -0,0 +1,253 @@
|
|||
use std::path::{Path, PathBuf};
|
||||
|
||||
use litellm_model_catalog::{AliasIssue, Catalog, Error, IntegrityLimits, Provenance};
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::json;
|
||||
|
||||
const ALPHA_FIXTURE: &[u8] = br#"{
|
||||
"sample_spec":{"explanation":"example"},
|
||||
"fallback_generalizations":{"rules":[{"name":"family","pattern":"^new-","model_info":{"mode":"chat"}}]},
|
||||
"Alpha":{"litellm_provider":"test","aliases":["short"],"price":0,"enabled":false,
|
||||
"optional":null,"unknown":{"nested":[1,{"x":true}]}}
|
||||
}"#;
|
||||
|
||||
#[fixture]
|
||||
fn repo_root() -> PathBuf {
|
||||
Path::new(env!("CARGO_MANIFEST_DIR")).join("../../..")
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn current_catalog(repo_root: PathBuf) -> Catalog {
|
||||
let body = std::fs::read(repo_root.join("model_prices_and_context_window.json")).unwrap();
|
||||
Catalog::parse(&body, Provenance::default()).unwrap()
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn backup_catalog(repo_root: PathBuf) -> Catalog {
|
||||
let body = std::fs::read(repo_root.join("litellm/model_prices_and_context_window_backup.json"))
|
||||
.unwrap();
|
||||
Catalog::parse(&body, Provenance::default()).unwrap()
|
||||
}
|
||||
|
||||
#[fixture]
|
||||
fn fixture_catalog() -> Catalog {
|
||||
Catalog::parse(
|
||||
ALPHA_FIXTURE,
|
||||
Provenance {
|
||||
source: Some("fixture".into()),
|
||||
revision: Some("rev".into()),
|
||||
etag: None,
|
||||
},
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn preserves_fields_and_metadata(fixture_catalog: Catalog) {
|
||||
let catalog = fixture_catalog;
|
||||
let entry = catalog.lookup("SHORT").unwrap();
|
||||
assert_eq!(entry.canonical_key, "Alpha");
|
||||
assert_eq!(entry.matched_key, "short");
|
||||
assert_eq!(entry.entry.field("price"), Some(&json!(0)));
|
||||
assert_eq!(entry.entry.field("enabled"), Some(&json!(false)));
|
||||
assert_eq!(entry.entry.field("optional"), Some(&json!(null)));
|
||||
assert_eq!(entry.entry.field("missing"), None);
|
||||
assert_eq!(
|
||||
entry.entry.field("unknown"),
|
||||
Some(&json!({"nested":[1,{"x":true}]}))
|
||||
);
|
||||
assert_eq!(entry.entry.field("aliases"), None);
|
||||
assert_eq!(entry.entry.info().litellm_provider.as_deref(), Some("test"));
|
||||
assert_eq!(
|
||||
catalog.sample_spec(),
|
||||
Some(&json!({"explanation":"example"}))
|
||||
);
|
||||
assert_eq!(catalog.fallback_rules().unwrap().len(), 1);
|
||||
assert_eq!(catalog.provenance().revision.as_deref(), Some("rev"));
|
||||
assert_eq!(catalog.model_count(), 1);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn snapshot_does_not_borrow_source() {
|
||||
let mut source = ALPHA_FIXTURE.to_vec();
|
||||
let catalog = Catalog::parse(&source, Provenance::default()).unwrap();
|
||||
source.fill(b' ');
|
||||
|
||||
let entry = catalog.lookup("short").unwrap();
|
||||
assert_eq!(entry.canonical_key, "Alpha");
|
||||
assert_eq!(entry.entry.field("price"), Some(&json!(0)));
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case("Shared", "First")]
|
||||
#[case("Second", "Second")]
|
||||
#[case("shared", "Second")]
|
||||
#[case("FIRST", "First")]
|
||||
#[case("sHaReD", "Second")]
|
||||
fn alias_collisions_and_case_fallback_follow_python_order(
|
||||
#[case] lookup: &str,
|
||||
#[case] expected: &str,
|
||||
) {
|
||||
let catalog = Catalog::parse(
|
||||
br#"{
|
||||
"First":{"aliases":["Shared","Second","first"],"value":1},
|
||||
"Second":{"aliases":["Shared","sHaReD"],"value":2},
|
||||
"SHARED":{"value":3}
|
||||
}"#,
|
||||
Provenance::default(),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(catalog.lookup(lookup).unwrap().canonical_key, expected);
|
||||
assert_eq!(catalog.alias_count(), 3);
|
||||
assert!(
|
||||
catalog
|
||||
.alias_issues()
|
||||
.contains(&AliasIssue::CanonicalCollision {
|
||||
model: "First".into(),
|
||||
alias: "Second".into(),
|
||||
})
|
||||
);
|
||||
assert!(
|
||||
catalog
|
||||
.alias_issues()
|
||||
.contains(&AliasIssue::AliasCollision {
|
||||
model: "Second".into(),
|
||||
alias: "Shared".into(),
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
enum ValidationOutcome {
|
||||
Ok,
|
||||
Shrunk,
|
||||
BelowMinimum,
|
||||
InvalidRatio,
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case(
|
||||
IntegrityLimits {
|
||||
backup_model_count: 2,
|
||||
min_model_count: 1,
|
||||
min_backup_ratio: 0.5,
|
||||
},
|
||||
ValidationOutcome::Ok
|
||||
)]
|
||||
#[case(
|
||||
IntegrityLimits {
|
||||
backup_model_count: 3,
|
||||
min_model_count: 1,
|
||||
min_backup_ratio: 0.5,
|
||||
},
|
||||
ValidationOutcome::Shrunk
|
||||
)]
|
||||
#[case(
|
||||
IntegrityLimits {
|
||||
backup_model_count: 0,
|
||||
min_model_count: 2,
|
||||
min_backup_ratio: 0.5,
|
||||
},
|
||||
ValidationOutcome::BelowMinimum
|
||||
)]
|
||||
#[case(
|
||||
IntegrityLimits {
|
||||
backup_model_count: 0,
|
||||
min_model_count: 0,
|
||||
min_backup_ratio: f64::NAN,
|
||||
},
|
||||
ValidationOutcome::InvalidRatio
|
||||
)]
|
||||
fn integrity_uses_canonical_count_and_strict_shrink_boundary(
|
||||
#[case] limits: IntegrityLimits,
|
||||
#[case] expected: ValidationOutcome,
|
||||
) {
|
||||
let catalog = Catalog::parse(
|
||||
br#"{"sample_spec":{},"fallback_generalizations":{"rules":[]},"a":{"aliases":["b","c"]}}"#,
|
||||
Provenance::default(),
|
||||
)
|
||||
.unwrap();
|
||||
let actual = catalog.validate(limits);
|
||||
match expected {
|
||||
ValidationOutcome::Ok => assert!(actual.is_ok()),
|
||||
ValidationOutcome::Shrunk => {
|
||||
assert!(matches!(actual, Err(Error::Shrunk { actual: 1, .. })))
|
||||
}
|
||||
ValidationOutcome::BelowMinimum => {
|
||||
assert!(matches!(actual, Err(Error::BelowMinimum { actual: 1, .. })))
|
||||
}
|
||||
ValidationOutcome::InvalidRatio => assert!(matches!(actual, Err(Error::InvalidRatio))),
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
enum MalformedOutcome {
|
||||
Empty,
|
||||
Json,
|
||||
EntryNotObject,
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::empty(b"{}", MalformedOutcome::Empty)]
|
||||
#[case::invalid_json(b"{", MalformedOutcome::Json)]
|
||||
#[case::entry_not_object(br#"{"a":1}"#, MalformedOutcome::EntryNotObject)]
|
||||
#[case::fallback_rules_missing(
|
||||
br#"{"fallback_generalizations":{},"a":{}}"#,
|
||||
MalformedOutcome::Json
|
||||
)]
|
||||
fn malformed_input_and_aliases_have_typed_outcomes(
|
||||
#[case] body: &[u8],
|
||||
#[case] expected: MalformedOutcome,
|
||||
) {
|
||||
let actual = Catalog::parse(body, Provenance::default());
|
||||
match expected {
|
||||
MalformedOutcome::Empty => assert!(matches!(actual, Err(Error::Empty))),
|
||||
MalformedOutcome::Json => assert!(matches!(actual, Err(Error::Json(_)))),
|
||||
MalformedOutcome::EntryNotObject => {
|
||||
assert!(matches!(actual, Err(Error::EntryNotObject { .. })))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn invalid_aliases_are_reported_not_fatal() {
|
||||
let catalog = Catalog::parse(
|
||||
br#"{"a":{"aliases":"bad"},"b":{"aliases":[9,"ok"]}}"#,
|
||||
Provenance::default(),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
catalog.alias_issues(),
|
||||
&[
|
||||
AliasIssue::InvalidList { model: "a".into() },
|
||||
AliasIssue::InvalidName { model: "b".into() },
|
||||
]
|
||||
);
|
||||
assert_eq!(catalog.lookup("ok").unwrap().canonical_key, "b");
|
||||
assert!(catalog.lookup("missing").is_none());
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn parses_current_and_packaged_catalogs_without_pinning_counts(
|
||||
current_catalog: Catalog,
|
||||
backup_catalog: Catalog,
|
||||
) {
|
||||
assert!(current_catalog.model_count() > 0);
|
||||
assert!(backup_catalog.model_count() > 0);
|
||||
assert!(current_catalog.sample_spec().is_some());
|
||||
assert!(backup_catalog.sample_spec().is_some());
|
||||
assert!(
|
||||
current_catalog
|
||||
.validate(IntegrityLimits::python_defaults(
|
||||
backup_catalog.model_count()
|
||||
))
|
||||
.is_ok()
|
||||
);
|
||||
for name in current_catalog.model_names() {
|
||||
let entry = current_catalog.lookup(name).unwrap().entry;
|
||||
assert_eq!(
|
||||
entry.info().litellm_provider.is_some(),
|
||||
entry.field("litellm_provider").is_some()
|
||||
);
|
||||
}
|
||||
}
|
||||
121
litellm-rust/crates/model-catalog/tests/spec_parity.rs
Normal file
121
litellm-rust/crates/model-catalog/tests/spec_parity.rs
Normal file
|
|
@ -0,0 +1,121 @@
|
|||
use std::collections::{BTreeSet, HashSet};
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use indexmap::IndexMap;
|
||||
use litellm_model_catalog::{
|
||||
Catalog, FallbackGeneralizations, ModelInfo, Provenance, model_entry_json_schema,
|
||||
};
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
#[fixture]
|
||||
fn repo_root() -> PathBuf {
|
||||
Path::new(env!("CARGO_MANIFEST_DIR")).join("../../..")
|
||||
}
|
||||
|
||||
fn json_eq(left: &Value, right: &Value) -> bool {
|
||||
match (left, right) {
|
||||
(Value::Number(left), Value::Number(right)) => left.as_f64() == right.as_f64(),
|
||||
(Value::Array(left), Value::Array(right)) => {
|
||||
left.len() == right.len() && left.iter().zip(right).all(|(a, b)| json_eq(a, b))
|
||||
}
|
||||
(Value::Object(left), Value::Object(right)) => {
|
||||
left.len() == right.len()
|
||||
&& left
|
||||
.iter()
|
||||
.all(|(key, value)| right.get(key).is_some_and(|other| json_eq(value, other)))
|
||||
}
|
||||
_ => left == right,
|
||||
}
|
||||
}
|
||||
|
||||
fn keys(value: &Map<String, Value>) -> BTreeSet<String> {
|
||||
value.keys().cloned().collect()
|
||||
}
|
||||
|
||||
fn symmetric_difference(left: &BTreeSet<String>, right: &BTreeSet<String>) -> BTreeSet<String> {
|
||||
left.symmetric_difference(right).cloned().collect()
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case("model_prices_and_context_window.json")]
|
||||
#[case("litellm/model_prices_and_context_window_backup.json")]
|
||||
fn every_entry_round_trips_through_model_info(repo_root: PathBuf, #[case] filename: &str) {
|
||||
let body = std::fs::read(repo_root.join(filename)).unwrap();
|
||||
let document: IndexMap<String, Value> = serde_json::from_slice(&body).unwrap();
|
||||
for (model_name, value) in document {
|
||||
if matches!(
|
||||
model_name.as_str(),
|
||||
"sample_spec" | "fallback_generalizations"
|
||||
) {
|
||||
continue;
|
||||
}
|
||||
let object = value
|
||||
.as_object()
|
||||
.unwrap_or_else(|| panic!("{model_name} is not an object"));
|
||||
let info: ModelInfo = serde_json::from_value(value.clone())
|
||||
.unwrap_or_else(|error| panic!("{model_name} does not deserialize: {error}"));
|
||||
let serialized = serde_json::to_value(info).unwrap();
|
||||
let serialized_object = serialized
|
||||
.as_object()
|
||||
.unwrap_or_else(|| panic!("{model_name} did not serialize as an object"));
|
||||
let mut expected = object.clone();
|
||||
expected.remove("aliases");
|
||||
let expected_keys = keys(&expected);
|
||||
let serialized_keys = keys(serialized_object);
|
||||
assert_eq!(
|
||||
expected_keys,
|
||||
serialized_keys,
|
||||
"{model_name} key difference: {:?}",
|
||||
symmetric_difference(&expected_keys, &serialized_keys)
|
||||
);
|
||||
assert!(
|
||||
json_eq(&Value::Object(expected), &serialized),
|
||||
"{model_name} changed during ModelInfo round-trip"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn fallback_generalizations_are_typed(repo_root: PathBuf) {
|
||||
let body = std::fs::read(repo_root.join("model_prices_and_context_window.json")).unwrap();
|
||||
let document: Map<String, Value> = serde_json::from_slice(&body).unwrap();
|
||||
let Some(raw_rules) = document.get("fallback_generalizations") else {
|
||||
return;
|
||||
};
|
||||
let _: FallbackGeneralizations = serde_json::from_value(raw_rules.clone()).unwrap();
|
||||
let catalog = Catalog::parse(&body, Provenance::default()).unwrap();
|
||||
assert!(
|
||||
catalog
|
||||
.fallback_rules()
|
||||
.is_some_and(|rules| !rules.is_empty())
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn generated_schema_properties_match_repo_schema(repo_root: PathBuf) {
|
||||
let body =
|
||||
std::fs::read(repo_root.join("model_prices_and_context_window.schema.json")).unwrap();
|
||||
let document: Value = serde_json::from_slice(&body).unwrap();
|
||||
let repo_entry_properties = document["$defs"]["modelEntry"]["properties"]
|
||||
.as_object()
|
||||
.unwrap();
|
||||
let generated = serde_json::to_value(model_entry_json_schema()).unwrap();
|
||||
let generated_properties = generated["properties"].as_object().unwrap();
|
||||
let expected = keys(repo_entry_properties);
|
||||
let actual = keys(generated_properties);
|
||||
assert_eq!(
|
||||
expected,
|
||||
actual,
|
||||
"modelEntry property difference: {:?}",
|
||||
symmetric_difference(&expected, &actual)
|
||||
);
|
||||
|
||||
let repo_root_properties = document["properties"].as_object().unwrap();
|
||||
let actual_root: HashSet<String> = repo_root_properties.keys().cloned().collect();
|
||||
let expected_root: HashSet<String> = ["sample_spec", "fallback_generalizations"]
|
||||
.into_iter()
|
||||
.map(str::to_owned)
|
||||
.collect();
|
||||
assert_eq!(actual_root, expected_root);
|
||||
}
|
||||
|
|
@ -19,6 +19,9 @@ huggingface = ["litellm-token-counter/huggingface"]
|
|||
tiktoken = ["litellm-token-counter/tiktoken"]
|
||||
|
||||
[dependencies]
|
||||
fancy-regex.workspace = true
|
||||
litellm-tracing.workspace = true
|
||||
litellm-host.workspace = true
|
||||
bytes.workspace = true
|
||||
futures-util.workspace = true
|
||||
litellm-cache.workspace = true
|
||||
|
|
@ -42,7 +45,7 @@ litellm-core-utils.workspace = true
|
|||
litellm-auth-gcp.workspace = true
|
||||
litellm-http.workspace = true
|
||||
litellm-llms.workspace = true
|
||||
litellm-secrets = { workspace = true, features = ["aws"] }
|
||||
litellm-secrets = { workspace = true, features = ["aws", "azure", "google", "hashicorp", "cyberark"] }
|
||||
litellm-secrets-types.workspace = true
|
||||
litellm-types.workspace = true
|
||||
litellm-host-python.workspace = true
|
||||
|
|
@ -52,6 +55,7 @@ pyo3-async-runtimes.workspace = true
|
|||
reqwest.workspace = true
|
||||
redis = { version = "1.7.0", features = ["tls-rustls"] }
|
||||
serde_json.workspace = true
|
||||
veil.workspace = true
|
||||
thiserror.workspace = true
|
||||
tokio = { workspace = true, features = ["rt", "sync"] }
|
||||
url.workspace = true
|
||||
|
|
|
|||
|
|
@ -1,5 +1,10 @@
|
|||
Native OCR uses `SecretSource` with `EnvironmentSecrets`, preserving process-environment reads. Readable Python secret managers still make OCR decline to the existing Python implementation. `ResolvedSecrets` and the separate `secret_manager_binding()` snapshot are inactive foundations for a later rollout
|
||||
Native OCR uses `litellm_secrets::source::SecretSource`. Built-in secret managers resolve to retained Rust backends. Custom Python managers and overrides keep the callback path. Readable managers still require the Rust secret-manager binding to be enabled
|
||||
|
||||
Cache and secret-manager catalog entries remain Python-only, including when `LITELLM_RUST=1`. The new cache runtime is not connected to SDK or gateway caching
|
||||
The shared proxy initializer captures native configuration without loading the extension or doing native I/O. `_SecretManagerRuntime.from_client` constructs a backend on first use and keeps its handle on the Python client. The secret-manager dispatcher selects Python or Rust through `catalog.py`. Native reads call that handle; Rust routes extract the backend directly. Configuration changes replace the handle, while calls already bound to the previous backend keep using it. Handles cannot be reused after fork. Directly constructed LiteLLM managers are adapted on first native use. Manually supplied SDK clients keep their Python behavior because their credentials cannot be inferred safely. Provider implementations contain no bridge registration
|
||||
|
||||
Retention describes ownership and lifetime. `callbacks-legacy-python::PublicCall` owns Python references for one call to preserve identity. A native cache or secret-manager handle owns shared Rust state across calls to preserve connection pools and caches. Both use existing `Py<T>` and shared Rust ownership, with execution and GIL transitions handled by `litellm-host-python`
|
||||
|
||||
|
||||
Cache and secret-manager catalog entries remain Python-only, including when `LITELLM_RUST=1`. This wiring does not change rollout policy
|
||||
|
||||
OCR provider requests use the shared `litellm-http` pool. AWS and Google secret-manager SDK clients keep their SDK transports, which do not yet inherit the pool's proxy, TLS, certificate, timeout, or observability configuration. Preserve those SDK transports and configure them equivalently instead of forcing them through reqwest
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
use crate::logger::run_sync_value;
|
||||
use litellm_cache_gcs::{DEFAULT_ENDPOINT, GcsConfig};
|
||||
use litellm_cache_redis_semantic::RedisSemanticConfig;
|
||||
use litellm_host_python::{release_gil, run_sync_value};
|
||||
use litellm_host_python::release_gil;
|
||||
use litellm_http::ClientVariant;
|
||||
use pyo3::prelude::*;
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
use crate::logger::run_async;
|
||||
use litellm_cache_response::PartialHits;
|
||||
use litellm_host_python::{ExecutionStep, from_py, release_gil, run_async, to_py};
|
||||
use litellm_host_python::{ExecutionStep, from_py, release_gil, to_py};
|
||||
use pyo3::{
|
||||
PyTraverseError, PyVisit,
|
||||
exceptions::{PyRuntimeError, PyValueError},
|
||||
|
|
|
|||
|
|
@ -1,10 +1,11 @@
|
|||
use crate::logger::run_sync_value;
|
||||
use litellm_auth_aws::AwsAuthConfig;
|
||||
use litellm_cache_gcs::{DEFAULT_ENDPOINT, GcsConfig};
|
||||
use litellm_cache_qdrant_semantic::{OpenAiEmbedderConfig, Quantization};
|
||||
use litellm_cache_redis::{RedisNode, RedisTopology};
|
||||
use litellm_cache_redis_semantic::RedisSemanticConfig;
|
||||
use litellm_cache_s3::{S3CacheConfig, S3Endpoint};
|
||||
use litellm_host_python::{release_gil, run_sync_value};
|
||||
use litellm_host_python::release_gil;
|
||||
use litellm_http::ClientVariant;
|
||||
use pyo3::{
|
||||
PyTraverseError, PyVisit,
|
||||
|
|
|
|||
|
|
@ -470,7 +470,7 @@ impl NativeResponseCache {
|
|||
match self {
|
||||
Self::Exact(_) | Self::QdrantSemantic(_) => {
|
||||
let service = self.clone();
|
||||
litellm_host_python::run_async(
|
||||
crate::logger::run_async(
|
||||
py,
|
||||
async move {
|
||||
service
|
||||
|
|
@ -495,7 +495,7 @@ impl NativeResponseCache {
|
|||
match self {
|
||||
Self::Exact(_) | Self::QdrantSemantic(_) => {
|
||||
let service = self.clone();
|
||||
litellm_host_python::run_async(
|
||||
crate::logger::run_async(
|
||||
py,
|
||||
async move { service.async_lookup(&request, now()).await },
|
||||
super::cache_error,
|
||||
|
|
@ -550,7 +550,7 @@ impl NativeResponseCache {
|
|||
match self {
|
||||
Self::Exact(_) | Self::QdrantSemantic(_) => {
|
||||
let service = self.clone();
|
||||
litellm_host_python::run_async(
|
||||
crate::logger::run_async(
|
||||
py,
|
||||
async move { service.async_store(&request, response, now()).await },
|
||||
super::cache_error,
|
||||
|
|
@ -619,7 +619,7 @@ impl NativeResponseCache {
|
|||
match self {
|
||||
Self::Exact(_) | Self::QdrantSemantic(_) => {
|
||||
let service = self.clone();
|
||||
litellm_host_python::run_async(
|
||||
crate::logger::run_async(
|
||||
py,
|
||||
async move { service.async_store_batch(entries, now()).await },
|
||||
super::cache_error,
|
||||
|
|
|
|||
|
|
@ -1,7 +1,8 @@
|
|||
use crate::logger::run_async;
|
||||
use std::{collections::VecDeque, time::Duration};
|
||||
|
||||
use litellm_cache::Error;
|
||||
use litellm_host_python::{Execution, ExecutionBody, ExecutionStep, run_async};
|
||||
use litellm_host_python::{Execution, ExecutionBody, ExecutionStep};
|
||||
use pyo3::{
|
||||
PyTraverseError, PyVisit,
|
||||
exceptions::{PyException, PyRuntimeError},
|
||||
|
|
|
|||
|
|
@ -102,7 +102,7 @@ pub(crate) fn call_config(
|
|||
.without_missing_files(&|path: &Path| path.exists());
|
||||
let resolution = Resolution::from(&settings);
|
||||
for unsupported in unreported(&REPORTED_UNSUPPORTED, resolution.unsupported) {
|
||||
PythonSettings::warn(py, &unsupported.to_string())?;
|
||||
crate::logger::capture(py).scope(|| litellm_tracing::warn!("{unsupported}"));
|
||||
}
|
||||
Ok(resolution.config)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -4,13 +4,10 @@ mod credentials;
|
|||
mod diagnostics;
|
||||
mod errors;
|
||||
mod http;
|
||||
mod logger;
|
||||
mod marshal;
|
||||
mod python_settings;
|
||||
mod routes;
|
||||
#[allow(
|
||||
dead_code,
|
||||
reason = "secret-manager foundations await rollout activation"
|
||||
)]
|
||||
mod secrets;
|
||||
mod tokenizer;
|
||||
|
||||
|
|
@ -25,6 +22,8 @@ mod _native {
|
|||
#[pymodule_export]
|
||||
use crate::errors::{RustBridgeDeclined, RustUpstreamError};
|
||||
#[pymodule_export]
|
||||
use crate::logger::NativeDiagnosticProcessor;
|
||||
#[pymodule_export]
|
||||
use crate::routes::audio_transcription::{atranscription, transcription};
|
||||
#[pymodule_export]
|
||||
use crate::routes::chat_completions::{
|
||||
|
|
@ -53,7 +52,11 @@ mod _native {
|
|||
let dict = module.dict();
|
||||
dict.set_item("_CacheTestHandle", py.get_type::<CacheTestHandle>())?;
|
||||
dict.set_item("_CacheTestResolver", py.get_type::<CacheTestResolver>())?;
|
||||
dict.set_item("_ResponseCacheRuntime", py.get_type::<ResolvedCache>())
|
||||
dict.set_item("_ResponseCacheRuntime", py.get_type::<ResolvedCache>())?;
|
||||
dict.set_item(
|
||||
"_SecretManagerRuntime",
|
||||
py.get_type::<crate::secrets::runtime::NativeSecretManager>(),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -87,6 +90,7 @@ mod tests {
|
|||
"chat_completions",
|
||||
"achat_completions",
|
||||
"ResponsesWebSocketConnection",
|
||||
"NativeDiagnosticProcessor",
|
||||
"TokenCounter",
|
||||
"Tokenizer",
|
||||
"gil_stats",
|
||||
|
|
|
|||
46
litellm-rust/crates/python-bridge/src/logger/execution.rs
Normal file
46
litellm-rust/crates/python-bridge/src/logger/execution.rs
Normal file
|
|
@ -0,0 +1,46 @@
|
|||
use std::future::Future;
|
||||
|
||||
use pyo3::prelude::*;
|
||||
use serde::Serialize;
|
||||
|
||||
pub(crate) fn run_sync<T, E, F>(
|
||||
py: Python<'_>,
|
||||
future: F,
|
||||
map_error: fn(E) -> PyErr,
|
||||
) -> PyResult<Py<PyAny>>
|
||||
where
|
||||
T: Serialize + Send + 'static,
|
||||
E: Send + 'static,
|
||||
F: Future<Output = Result<T, E>> + Send + 'static,
|
||||
{
|
||||
litellm_host_python::run_sync(py, super::capture(py).instrument(future), map_error)
|
||||
}
|
||||
|
||||
pub(crate) fn run_async<T, E, F>(
|
||||
py: Python<'_>,
|
||||
future: F,
|
||||
map_error: fn(E) -> PyErr,
|
||||
) -> PyResult<Bound<'_, PyAny>>
|
||||
where
|
||||
T: Serialize + Send + 'static,
|
||||
E: Send + 'static,
|
||||
F: Future<Output = Result<T, E>> + Send + 'static,
|
||||
{
|
||||
litellm_host_python::run_async(py, super::capture(py).instrument(future), map_error)
|
||||
}
|
||||
|
||||
pub(crate) fn run_sync_value<T, F>(py: Python<'_>, future: F) -> PyResult<T>
|
||||
where
|
||||
T: Send + 'static,
|
||||
F: Future<Output = PyResult<T>> + Send + 'static,
|
||||
{
|
||||
litellm_host_python::run_sync_value(py, super::capture(py).instrument(future))
|
||||
}
|
||||
|
||||
pub(crate) fn run_async_value<T, F>(py: Python<'_>, future: F) -> PyResult<Bound<'_, PyAny>>
|
||||
where
|
||||
T: for<'py> IntoPyObject<'py> + Send + 'static,
|
||||
F: Future<Output = PyResult<T>> + Send + 'static,
|
||||
{
|
||||
litellm_host_python::run_async_value(py, super::capture(py).instrument(future))
|
||||
}
|
||||
41
litellm-rust/crates/python-bridge/src/logger/machine.rs
Normal file
41
litellm-rust/crates/python-bridge/src/logger/machine.rs
Normal file
|
|
@ -0,0 +1,41 @@
|
|||
use std::sync::OnceLock;
|
||||
|
||||
use litellm_host::{
|
||||
host::HostResult,
|
||||
machine::{HostFailure, Interrupted, Machine, Step},
|
||||
route::Route,
|
||||
};
|
||||
use litellm_tracing::Logger;
|
||||
use pyo3::Python;
|
||||
|
||||
pub(crate) struct LoggedMachine<M> {
|
||||
machine: M,
|
||||
logger: OnceLock<Logger>,
|
||||
}
|
||||
|
||||
impl<M> LoggedMachine<M> {
|
||||
pub(crate) fn new(machine: M) -> Self {
|
||||
Self {
|
||||
machine,
|
||||
logger: OnceLock::new(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<M: Machine> Machine for LoggedMachine<M> {
|
||||
type Route = M::Route;
|
||||
type Complete = M::Complete;
|
||||
|
||||
fn resume(&mut self, result: Option<HostResult<Self::Route>>) -> Step<'_, Self> {
|
||||
let logger = self.logger.get_or_init(|| Python::attach(super::capture));
|
||||
Box::pin(logger.instrument(logger.scope(|| self.machine.resume(result))))
|
||||
}
|
||||
|
||||
fn interrupt(
|
||||
&mut self,
|
||||
failure: HostFailure<<Self::Route as Route>::Error>,
|
||||
) -> Interrupted<'_, Self> {
|
||||
let logger = self.logger.get_or_init(|| Python::attach(super::capture));
|
||||
Box::pin(logger.instrument(logger.scope(|| self.machine.interrupt(failure))))
|
||||
}
|
||||
}
|
||||
166
litellm-rust/crates/python-bridge/src/logger/mod.rs
Normal file
166
litellm-rust/crates/python-bridge/src/logger/mod.rs
Normal file
|
|
@ -0,0 +1,166 @@
|
|||
mod execution;
|
||||
mod machine;
|
||||
|
||||
pub(crate) use execution::{run_async, run_async_value, run_sync, run_sync_value};
|
||||
pub(crate) use machine::LoggedMachine;
|
||||
|
||||
use litellm_host_python::Pythonized;
|
||||
use litellm_tracing::{DiagnosticInput, Level, Logger, Metadata, Policy, Processor, Record, Sink};
|
||||
use pyo3::exceptions::PyRuntimeError;
|
||||
use pyo3::prelude::*;
|
||||
|
||||
const MODULE: &str = "litellm.rust_bridge.logger";
|
||||
type NativeDiagnosticOutput = (String, Option<String>, Option<String>, Vec<String>, bool);
|
||||
|
||||
#[pyclass]
|
||||
pub(crate) struct NativeDiagnosticProcessor {
|
||||
inner: Processor,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl NativeDiagnosticProcessor {
|
||||
#[new]
|
||||
fn new(minimum_custom_key_length: usize) -> Self {
|
||||
Self {
|
||||
inner: Processor::new(minimum_custom_key_length),
|
||||
}
|
||||
}
|
||||
|
||||
fn redact_text(&self, text: &str) -> PyResult<String> {
|
||||
self.inner.redact_text(text).map_err(processing_error)
|
||||
}
|
||||
|
||||
fn redact_structured_text(&self, key: Option<&str>, text: &str) -> PyResult<String> {
|
||||
self.inner
|
||||
.redact_structured_text(key, text)
|
||||
.map_err(processing_error)
|
||||
}
|
||||
|
||||
fn redact_client_message(&self, text: &str) -> PyResult<String> {
|
||||
self.inner
|
||||
.redact_client_message(text)
|
||||
.map_err(processing_error)
|
||||
}
|
||||
|
||||
#[pyo3(signature = (message, exception, stack, leaves, policy))]
|
||||
fn process_diagnostic(
|
||||
&self,
|
||||
message: String,
|
||||
exception: Option<String>,
|
||||
stack: Option<String>,
|
||||
leaves: Vec<(Option<String>, String)>,
|
||||
policy: (bool, i64, i64),
|
||||
) -> PyResult<NativeDiagnosticOutput> {
|
||||
let input = DiagnosticInput {
|
||||
message,
|
||||
exception,
|
||||
stack,
|
||||
leaves,
|
||||
};
|
||||
let policy = Policy {
|
||||
redact: policy.0,
|
||||
base64_limit: policy.1,
|
||||
text_limit: policy.2,
|
||||
};
|
||||
self.inner
|
||||
.process_diagnostic(&input, policy)
|
||||
.map(|output| {
|
||||
(
|
||||
output.message,
|
||||
output.exception,
|
||||
output.stack,
|
||||
output.leaves,
|
||||
output.changed,
|
||||
)
|
||||
})
|
||||
.map_err(processing_error)
|
||||
}
|
||||
|
||||
fn scrub_access_arguments(&self, arguments: Vec<String>) -> PyResult<Vec<String>> {
|
||||
self.inner
|
||||
.scrub_access_arguments(&arguments)
|
||||
.map_err(processing_error)
|
||||
}
|
||||
}
|
||||
|
||||
fn processing_error(_: fancy_regex::Error) -> PyErr {
|
||||
PyRuntimeError::new_err("diagnostic processing failed")
|
||||
}
|
||||
|
||||
struct PythonSink {
|
||||
correlation: (String, String),
|
||||
}
|
||||
|
||||
fn level(level: &Level) -> u8 {
|
||||
match *level {
|
||||
Level::ERROR => 40,
|
||||
Level::WARN => 30,
|
||||
Level::INFO => 20,
|
||||
Level::DEBUG | Level::TRACE => 10,
|
||||
}
|
||||
}
|
||||
|
||||
fn report<T: Default>(py: Python<'_>, result: PyResult<T>) -> T {
|
||||
match result {
|
||||
Ok(value) => value,
|
||||
Err(error) => {
|
||||
error.write_unraisable(py, None);
|
||||
T::default()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Sink for PythonSink {
|
||||
fn enabled(&self, metadata: &Metadata<'_>) -> bool {
|
||||
if !metadata.target().starts_with("litellm_") && !metadata.target().starts_with("_native::")
|
||||
{
|
||||
return false;
|
||||
}
|
||||
Python::try_attach(|py| {
|
||||
report(
|
||||
py,
|
||||
py.import(MODULE)
|
||||
.and_then(|module| module.call_method1("enabled", (level(metadata.level()),)))
|
||||
.and_then(|enabled| enabled.extract()),
|
||||
)
|
||||
})
|
||||
.unwrap_or(false)
|
||||
}
|
||||
|
||||
fn emit(&self, record: &Record) {
|
||||
Python::try_attach(|py| {
|
||||
report(
|
||||
py,
|
||||
py.import(MODULE).and_then(|module| {
|
||||
module
|
||||
.call_method1(
|
||||
"emit",
|
||||
(
|
||||
level(record.metadata.level()),
|
||||
&record.message,
|
||||
record.metadata.file().unwrap_or_default(),
|
||||
record.metadata.line().unwrap_or_default(),
|
||||
record.metadata.target(),
|
||||
Pythonized(&record.fields),
|
||||
(&self.correlation.0, &self.correlation.1),
|
||||
),
|
||||
)
|
||||
.map(|_| ())
|
||||
}),
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn capture(py: Python<'_>) -> Logger {
|
||||
report(
|
||||
py,
|
||||
py.import(MODULE)
|
||||
.and_then(|module| module.call_method0("context"))
|
||||
.and_then(|value| value.extract())
|
||||
.map(|correlation| Logger::new(PythonSink { correlation })),
|
||||
)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
295
litellm-rust/crates/python-bridge/src/logger/tests.rs
Normal file
295
litellm-rust/crates/python-bridge/src/logger/tests.rs
Normal file
|
|
@ -0,0 +1,295 @@
|
|||
use std::{process::Command, task::Poll};
|
||||
|
||||
use litellm_host::{
|
||||
host::HostResult,
|
||||
machine::{HostFailure, Interrupted, Machine, MachineStep, Step},
|
||||
route::Route,
|
||||
};
|
||||
|
||||
use pyo3::{prelude::*, types::PyDict};
|
||||
|
||||
struct DiagnosticMachine;
|
||||
|
||||
impl Route for DiagnosticMachine {
|
||||
type Response = ();
|
||||
type Error = String;
|
||||
type Op = ();
|
||||
type OpResult = ();
|
||||
type Chunk = ();
|
||||
type StreamHead = ();
|
||||
}
|
||||
|
||||
impl Machine for DiagnosticMachine {
|
||||
type Route = Self;
|
||||
type Complete = ();
|
||||
|
||||
fn resume(&mut self, _: Option<HostResult<Self>>) -> Step<'_, Self> {
|
||||
litellm_tracing::warn!("machine started");
|
||||
Box::pin(async {
|
||||
tokio::task::yield_now().await;
|
||||
litellm_tracing::warn!("machine warning");
|
||||
Ok(MachineStep::Complete(()))
|
||||
})
|
||||
}
|
||||
|
||||
fn interrupt(&mut self, _: HostFailure<String>) -> Interrupted<'_, Self> {
|
||||
Box::pin(async {
|
||||
litellm_tracing::warn!("machine interrupted");
|
||||
Ok(())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
fn machine_warning(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
|
||||
let mut machine = super::LoggedMachine::new(DiagnosticMachine);
|
||||
let mut future = Box::pin(async move {
|
||||
machine
|
||||
.resume(None)
|
||||
.await
|
||||
.map_err(pyo3::exceptions::PyValueError::new_err)?;
|
||||
machine
|
||||
.interrupt(HostFailure::Error("stop".into()))
|
||||
.await
|
||||
.map_err(pyo3::exceptions::PyValueError::new_err)
|
||||
});
|
||||
assert!(matches!(
|
||||
litellm_host_python::poll_async_value(py, future.as_mut())?,
|
||||
Poll::Pending
|
||||
));
|
||||
litellm_host_python::run_async_value(py, future)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
fn warning(py: Python<'_>) {
|
||||
super::capture(py).scope(|| {
|
||||
litellm_tracing::warn!(attempt = 3, retry = true, "native warning");
|
||||
});
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
fn levels(py: Python<'_>) {
|
||||
super::capture(py).scope(|| {
|
||||
litellm_tracing::trace!("trace");
|
||||
litellm_tracing::debug!("debug");
|
||||
litellm_tracing::info!("info");
|
||||
litellm_tracing::warn!("warn");
|
||||
litellm_tracing::error!("error");
|
||||
litellm_tracing::warn!(target: "unrelated_transport", "private wire data");
|
||||
});
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
fn asynchronous_warning(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
|
||||
super::run_async_value(py, async {
|
||||
tokio::task::yield_now().await;
|
||||
litellm_tracing::warn!("async warning");
|
||||
Ok(())
|
||||
})
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
fn synchronous_warning(py: Python<'_>) -> PyResult<()> {
|
||||
super::run_sync_value(py, async {
|
||||
tokio::task::yield_now().await;
|
||||
litellm_tracing::warn!("sync warning");
|
||||
Ok(())
|
||||
})
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
fn synchronous_failure(py: Python<'_>) -> PyResult<()> {
|
||||
super::run_sync_value(py, async {
|
||||
litellm_tracing::warn!("failure diagnostic");
|
||||
Err(pyo3::exceptions::PyValueError::new_err("request failed"))
|
||||
})
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
fn http_warning(py: Python<'_>) -> PyResult<()> {
|
||||
crate::http::call_config(py, &PyDict::new(py), false).map(|_| ())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn native_events_reach_python_with_levels_context_reentry_and_http_deduplication() {
|
||||
if std::env::var_os("LITELLM_LOGGER_TEST_PROCESS").is_none() {
|
||||
let output = Command::new(std::env::current_exe().unwrap())
|
||||
.args([
|
||||
"--exact",
|
||||
std::thread::current().name().unwrap(),
|
||||
"--nocapture",
|
||||
])
|
||||
.env("LITELLM_LOGGER_TEST_PROCESS", "1")
|
||||
.output()
|
||||
.unwrap();
|
||||
assert!(
|
||||
output.status.success(),
|
||||
"{}\n{}",
|
||||
String::from_utf8_lossy(&output.stdout),
|
||||
String::from_utf8_lossy(&output.stderr)
|
||||
);
|
||||
return;
|
||||
}
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = PyDict::new(py);
|
||||
locals
|
||||
.set_item(
|
||||
"repo_root",
|
||||
concat!(env!("CARGO_MANIFEST_DIR"), "/../../.."),
|
||||
)
|
||||
.unwrap();
|
||||
locals
|
||||
.set_item(
|
||||
"machine_warning",
|
||||
wrap_pyfunction!(machine_warning, py).unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
locals
|
||||
.set_item(
|
||||
"synchronous_failure",
|
||||
wrap_pyfunction!(synchronous_failure, py).unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
locals
|
||||
.set_item("levels", wrap_pyfunction!(levels, py).unwrap())
|
||||
.unwrap();
|
||||
locals
|
||||
.set_item("warning", wrap_pyfunction!(warning, py).unwrap())
|
||||
.unwrap();
|
||||
locals
|
||||
.set_item(
|
||||
"asynchronous_warning",
|
||||
wrap_pyfunction!(asynchronous_warning, py).unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
locals
|
||||
.set_item(
|
||||
"synchronous_warning",
|
||||
wrap_pyfunction!(synchronous_warning, py).unwrap(),
|
||||
)
|
||||
.unwrap();
|
||||
locals
|
||||
.set_item("http_warning", wrap_pyfunction!(http_warning, py).unwrap())
|
||||
.unwrap();
|
||||
let importable = py
|
||||
.eval(
|
||||
c"__import__('importlib.util', fromlist=['util']).find_spec('dotenv') is not None",
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap()
|
||||
.is_truthy()
|
||||
.unwrap();
|
||||
if !importable {
|
||||
eprintln!("SKIP: litellm package dependencies are not importable in this interpreter");
|
||||
return;
|
||||
}
|
||||
py.run(c"
|
||||
import asyncio
|
||||
import logging
|
||||
import sys
|
||||
sys.path.insert(0, repo_root)
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger, session_id_var, trace_id_var
|
||||
|
||||
class Capture(logging.Handler):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.records = []
|
||||
def emit(self, record):
|
||||
self.records.append(record)
|
||||
warning()
|
||||
|
||||
class Broken(logging.Handler):
|
||||
def emit(self, record):
|
||||
raise ValueError('handler failed')
|
||||
|
||||
capture = Capture()
|
||||
old_handlers = verbose_logger.handlers
|
||||
old_level = verbose_logger.level
|
||||
old_correlation = litellm.request_correlation_in_logs
|
||||
old_curve = litellm.ssl_ecdh_curve
|
||||
old_unraisable = sys.unraisablehook
|
||||
failures = []
|
||||
try:
|
||||
verbose_logger.handlers = [capture]
|
||||
litellm.request_correlation_in_logs = True
|
||||
verbose_logger.setLevel(logging.ERROR)
|
||||
warning()
|
||||
assert capture.records == []
|
||||
verbose_logger.setLevel(logging.WARNING)
|
||||
warning()
|
||||
assert len(capture.records) == 1
|
||||
record = capture.records[0]
|
||||
assert record.getMessage() == 'native warning'
|
||||
assert record.levelno == logging.WARNING
|
||||
assert record.rust_fields == {'attempt': 3, 'retry': True}
|
||||
assert record.pathname.endswith('logger/tests.rs')
|
||||
assert record.lineno > 0
|
||||
assert record.rust_target.endswith('logger::tests')
|
||||
verbose_logger.setLevel(logging.ERROR)
|
||||
warning()
|
||||
assert len(capture.records) == 1
|
||||
verbose_logger.setLevel(logging.WARNING)
|
||||
|
||||
async def request(name):
|
||||
session = session_id_var.set(name)
|
||||
trace = trace_id_var.set('trace-' + name)
|
||||
try:
|
||||
await asynchronous_warning()
|
||||
await machine_warning()
|
||||
synchronous_warning()
|
||||
assert session_id_var.get() == name
|
||||
assert trace_id_var.get() == 'trace-' + name
|
||||
finally:
|
||||
trace_id_var.reset(trace)
|
||||
session_id_var.reset(session)
|
||||
|
||||
async def concurrent():
|
||||
await asyncio.gather(request('first'), request('second'))
|
||||
|
||||
asyncio.run(concurrent())
|
||||
assert sorted((r.getMessage(), r.session_id, r.trace_id) for r in capture.records[1:]) == sorted(
|
||||
(message, name, 'trace-' + name)
|
||||
for name in ('first', 'second')
|
||||
for message in ('async warning', 'sync warning', 'machine started', 'machine warning', 'machine interrupted')
|
||||
)
|
||||
|
||||
verbose_logger.setLevel(logging.DEBUG)
|
||||
before_levels = len(capture.records)
|
||||
levels()
|
||||
assert [(r.getMessage(), r.levelno) for r in capture.records[before_levels:]] == [
|
||||
('trace', logging.DEBUG), ('debug', logging.DEBUG), ('info', logging.INFO),
|
||||
('warn', logging.WARNING), ('error', logging.ERROR),
|
||||
]
|
||||
|
||||
before = len(capture.records)
|
||||
litellm.ssl_ecdh_curve = 'logger-test-unsupported-curve'
|
||||
http_warning()
|
||||
http_warning()
|
||||
assert len(capture.records) == before + 1
|
||||
assert 'logger-test-unsupported-curve' in capture.records[-1].getMessage()
|
||||
assert capture.records[-1].pathname.endswith('http.rs')
|
||||
|
||||
verbose_logger.handlers = [Broken()]
|
||||
sys.unraisablehook = failures.append
|
||||
warning()
|
||||
assert len(failures) == 1
|
||||
assert str(failures[0].exc_value) == 'handler failed'
|
||||
try:
|
||||
synchronous_failure()
|
||||
except ValueError as error:
|
||||
assert str(error) == 'request failed'
|
||||
else:
|
||||
raise AssertionError('request failure was lost')
|
||||
assert len(failures) == 2
|
||||
finally:
|
||||
sys.unraisablehook = old_unraisable
|
||||
verbose_logger.handlers = old_handlers
|
||||
verbose_logger.setLevel(old_level)
|
||||
litellm.request_correlation_in_logs = old_correlation
|
||||
litellm.ssl_ecdh_curve = old_curve
|
||||
", Some(&locals), Some(&locals)).unwrap();
|
||||
});
|
||||
}
|
||||
|
|
@ -44,11 +44,6 @@ impl PythonSettings {
|
|||
pub(crate) fn snapshot(self, value: Bound<'_, PyAny>) -> Snapshot<'_> {
|
||||
Snapshot { group: self, value }
|
||||
}
|
||||
|
||||
pub(crate) fn warn(py: Python<'_>, message: &str) -> PyResult<()> {
|
||||
py.import(MODULE)?.getattr("warn")?.call1((message,))?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
|
|
|||
|
|
@ -1,7 +1,8 @@
|
|||
use crate::logger::{run_async, run_sync};
|
||||
use litellm_core::audio_transcription::{
|
||||
Error, audio_transcription as run_audio_transcription, types::AudioTranscriptionRequest,
|
||||
};
|
||||
use litellm_host_python::{from_py_argument, run_async, run_sync};
|
||||
use litellm_host_python::from_py_argument;
|
||||
use pyo3::prelude::*;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
|
|
|
|||
|
|
@ -1,8 +1,9 @@
|
|||
use crate::logger::{run_async, run_sync};
|
||||
use litellm_core::chat_completions::{
|
||||
Error, chat_completions as run_chat_completions, chat_completions_decline_reason,
|
||||
types::ChatCompletionsRequest,
|
||||
};
|
||||
use litellm_host_python::{from_py_argument, run_async, run_sync};
|
||||
use litellm_host_python::from_py_argument;
|
||||
use litellm_types::utils::ChatCompletionsResponse;
|
||||
use pyo3::prelude::*;
|
||||
use serde_json::{Map, Value};
|
||||
|
|
|
|||
|
|
@ -43,7 +43,7 @@ fn run_messages(
|
|||
py,
|
||||
SURFACE,
|
||||
PublicCall::capture(&request, &args, &kwargs)?,
|
||||
messages_machine(),
|
||||
crate::logger::LoggedMachine::new(messages_machine()),
|
||||
MessagesRouteHost::new(request.unbind()),
|
||||
asynchronous,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -3,17 +3,14 @@ mod errors;
|
|||
mod host;
|
||||
mod project;
|
||||
|
||||
use std::sync::{Arc, LazyLock};
|
||||
use std::sync::LazyLock;
|
||||
|
||||
use host::OcrRouteHost;
|
||||
use litellm_auth_gcp::VertexAuth;
|
||||
use litellm_callbacks_legacy_python::{LegacySurface, PublicCall, run_legacy_call};
|
||||
use litellm_core::ocr::route::ocr_machine;
|
||||
use litellm_core_utils::settings::ProcessEnvironment;
|
||||
use litellm_llms::base_llm::{
|
||||
inference::secrets::{EnvironmentSecrets, SecretSource},
|
||||
ocr::{handler::OcrClient, settings::OcrSettings},
|
||||
};
|
||||
use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings};
|
||||
use pyo3::{
|
||||
prelude::*,
|
||||
types::{PyDict, PyTuple},
|
||||
|
|
@ -21,14 +18,11 @@ use pyo3::{
|
|||
|
||||
use crate::{
|
||||
coercion::FieldSpec,
|
||||
errors::RustBridgeDeclined,
|
||||
http,
|
||||
python_settings::{PythonSettings, Snapshot},
|
||||
secrets,
|
||||
};
|
||||
|
||||
const SECRET_MANAGER_READABLE: FieldSpec<bool> =
|
||||
FieldSpec::new("readable", |field| field.schema_bool());
|
||||
|
||||
const VERTEX_PROJECT: FieldSpec<Option<String>> =
|
||||
FieldSpec::new("vertex_project", |field| field.falsy_optional_string());
|
||||
const VERTEX_LOCATION: FieldSpec<Option<String>> =
|
||||
|
|
@ -58,7 +52,7 @@ fn run_ocr(
|
|||
kwargs: Bound<'_, PyDict>,
|
||||
asynchronous: bool,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
let secrets = process_environment_secrets(&PythonSettings::SecretManager.read(py)?)?;
|
||||
let secrets = secrets::source(py)?;
|
||||
let config = http::call_config(py, &kwargs, asynchronous)?;
|
||||
let client = OcrClient::new(
|
||||
http::pool(),
|
||||
|
|
@ -73,21 +67,12 @@ fn run_ocr(
|
|||
py,
|
||||
if asynchronous { ASYNC_SURFACE } else { SURFACE },
|
||||
PublicCall::capture(&request, &args, &kwargs)?,
|
||||
ocr_machine(client),
|
||||
crate::logger::LoggedMachine::new(ocr_machine(client)),
|
||||
OcrRouteHost::new(request.unbind()),
|
||||
asynchronous,
|
||||
)
|
||||
}
|
||||
|
||||
fn process_environment_secrets(snapshot: &Snapshot<'_>) -> PyResult<Arc<dyn SecretSource>> {
|
||||
if snapshot.read(&SECRET_MANAGER_READABLE)? {
|
||||
return Err(RustBridgeDeclined::new_err(
|
||||
"a readable secret manager is configured and the Rust route only reads the process environment",
|
||||
));
|
||||
}
|
||||
Ok(Arc::new(EnvironmentSecrets))
|
||||
}
|
||||
|
||||
fn ocr_settings(py: Python<'_>) -> PyResult<OcrSettings> {
|
||||
project_provider_defaults(&PythonSettings::ProviderDefaults.read(py)?)
|
||||
}
|
||||
|
|
@ -123,38 +108,10 @@ pub(crate) fn aocr(
|
|||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use pyo3::{prelude::*, types::PyDict};
|
||||
|
||||
use super::process_environment_secrets;
|
||||
use crate::errors::RustBridgeDeclined;
|
||||
use pyo3::prelude::*;
|
||||
|
||||
use crate::python_settings::PythonSettings;
|
||||
|
||||
fn secret_manager<'py>(py: Python<'py>, readable: bool) -> Bound<'py, PyAny> {
|
||||
let locals = PyDict::new(py);
|
||||
locals.set_item("readable", readable).unwrap();
|
||||
py.run(
|
||||
c"import types\nmanager = types.SimpleNamespace(readable=readable)",
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
locals.get_item("manager").unwrap().unwrap()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_readable_secret_manager_sends_the_call_back_to_python() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let declined = process_environment_secrets(
|
||||
&PythonSettings::SecretManager.snapshot(secret_manager(py, true)),
|
||||
)
|
||||
.err()
|
||||
.expect("the Rust route declines");
|
||||
assert!(declined.is_instance_of::<RustBridgeDeclined>(py));
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_defaults_distinguish_falsey_values_and_exact_true() {
|
||||
Python::initialize();
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue