mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_e2e_testing
This commit is contained in:
commit
7761ee1a23
242 changed files with 22408 additions and 2560 deletions
|
|
@ -133,6 +133,26 @@ commands:
|
|||
done
|
||||
echo "record/replay proxy did not become ready" >&2
|
||||
exit 1
|
||||
start_fake_openai_endpoint:
|
||||
description: "Start the canned OpenAI mock (tests/_fake_openai_endpoint_server.py) on host port 8190 and wait until healthy. Models whose api_base points here (via FAKE_OPENAI_API_BASE) get well-formed chat/text/embedding responses with realistic usage, so the E2E run neither pays for nor depends on the live provider. A request whose model is '429' returns HTTP 429 for rate-limit/cooldown tests. Run after uv deps are synced."
|
||||
steps:
|
||||
- run:
|
||||
name: Start fake OpenAI endpoint
|
||||
background: true
|
||||
command: |
|
||||
uv run --no-sync python tests/_fake_openai_endpoint_server.py --host 0.0.0.0 --port 8190
|
||||
- run:
|
||||
name: Wait for fake OpenAI endpoint
|
||||
command: |
|
||||
for i in $(seq 1 30); do
|
||||
if curl -sf http://localhost:8190/health >/dev/null 2>&1; then
|
||||
echo "fake OpenAI endpoint is up"
|
||||
exit 0
|
||||
fi
|
||||
sleep 1
|
||||
done
|
||||
echo "fake OpenAI endpoint did not become ready" >&2
|
||||
exit 1
|
||||
setup_litellm_enterprise_pip:
|
||||
steps:
|
||||
- run:
|
||||
|
|
@ -168,6 +188,8 @@ jobs:
|
|||
name: win/default
|
||||
shell: powershell.exe
|
||||
working_directory: ~/project
|
||||
environment:
|
||||
UV_PYTHON: "3.11"
|
||||
steps:
|
||||
- checkout
|
||||
- run:
|
||||
|
|
@ -200,7 +222,7 @@ jobs:
|
|||
if (-not (Select-String -Path $PROFILE -SimpleMatch $uvBin -Quiet)) {
|
||||
Add-Content -Path $PROFILE -Value "`$env:Path = `"$uvBin;`$env:Path`""
|
||||
}
|
||||
uv sync --frozen --group dev --python (Get-Command python).Source
|
||||
uv sync --frozen --group dev --python 3.11
|
||||
- run:
|
||||
name: Run Windows-specific test
|
||||
command: |
|
||||
|
|
@ -594,6 +616,8 @@ jobs:
|
|||
working_directory: ~/project
|
||||
resource_class: large
|
||||
parallelism: 4
|
||||
environment:
|
||||
FAKE_OPENAI_API_BASE: http://127.0.0.1:8190
|
||||
steps:
|
||||
- checkout
|
||||
- setup_google_dns
|
||||
|
|
@ -609,6 +633,7 @@ jobs:
|
|||
paths:
|
||||
- ~/.cache/uv
|
||||
key: v1-uv-cache-{{ checksum "uv.lock" }}
|
||||
- start_fake_openai_endpoint
|
||||
# Run pytest and generate JUnit XML report
|
||||
- setup_litellm_enterprise_pip
|
||||
- run:
|
||||
|
|
@ -1549,6 +1574,7 @@ jobs:
|
|||
name: Install Dependencies
|
||||
command: |
|
||||
uv sync --frozen --all-groups --all-extras --python 3.12
|
||||
- start_fake_openai_endpoint
|
||||
- start_postgres:
|
||||
db_name: litellm_test
|
||||
- attach_workspace:
|
||||
|
|
@ -1586,6 +1612,7 @@ jobs:
|
|||
-e DATABASE_URL="postgresql://postgres:postgres@host.docker.internal:5432/litellm_test" \
|
||||
-e DEFAULT_NUM_WORKERS_LITELLM_PROXY=1 \
|
||||
-e DISABLE_SCHEMA_UPDATE="True" \
|
||||
-e FAKE_OPENAI_API_BASE=http://host.docker.internal:8190 \
|
||||
--name my-app \
|
||||
--add-host=host.docker.internal:host-gateway \
|
||||
-v $(pwd)/litellm/proxy/example_config_yaml/bad_schema.prisma:/app/schema.prisma \
|
||||
|
|
@ -1648,6 +1675,7 @@ jobs:
|
|||
zstd -d litellm-docker-database.tar.zst --stdout | docker load
|
||||
docker tag litellm-docker-database:ci my-app:latest
|
||||
- start_openai_record_replay_proxy
|
||||
- start_fake_openai_endpoint
|
||||
- run:
|
||||
name: Run Docker container
|
||||
command: |
|
||||
|
|
@ -1655,6 +1683,7 @@ jobs:
|
|||
-p 4000:4000 \
|
||||
-e DATABASE_URL=postgresql://postgres:postgres@host.docker.internal:5432/circle_test \
|
||||
-e USE_PRISMA_MIGRATE=True \
|
||||
-e FAKE_OPENAI_API_BASE=http://host.docker.internal:8190 \
|
||||
-e AZURE_API_KEY=$AZURE_API_KEY \
|
||||
-e REDIS_HOST=$REDIS_HOST \
|
||||
-e REDIS_PASSWORD=$REDIS_PASSWORD \
|
||||
|
|
@ -1817,6 +1846,7 @@ jobs:
|
|||
zstd -d litellm-docker-database.tar.zst --stdout | docker load
|
||||
docker images | grep litellm-docker-database
|
||||
- start_openai_record_replay_proxy
|
||||
- start_fake_openai_endpoint
|
||||
- run:
|
||||
name: Run Docker container
|
||||
# intentionally give bad redis credentials here
|
||||
|
|
@ -1830,6 +1860,7 @@ jobs:
|
|||
-e REDIS_PORT=$REDIS_PORT \
|
||||
-e LITELLM_MASTER_KEY="sk-1234" \
|
||||
-e OPENAI_API_KEY=$OPENAI_API_KEY \
|
||||
-e FAKE_OPENAI_API_BASE=http://host.docker.internal:8190 \
|
||||
-e LITELLM_LICENSE=$LITELLM_LICENSE \
|
||||
-e OTEL_EXPORTER="in_memory" \
|
||||
-e AWS_ACCESS_KEY_ID=$AWS_ACCESS_KEY_ID \
|
||||
|
|
@ -1889,6 +1920,7 @@ jobs:
|
|||
-e REDIS_PORT=$REDIS_PORT \
|
||||
-e LITELLM_MASTER_KEY="sk-1234" \
|
||||
-e OPENAI_API_KEY=$OPENAI_API_KEY \
|
||||
-e FAKE_OPENAI_API_BASE=http://host.docker.internal:8190 \
|
||||
-e LITELLM_LICENSE="bad-license" \
|
||||
--add-host host.docker.internal:host-gateway \
|
||||
--name my-app-3 \
|
||||
|
|
@ -1938,6 +1970,7 @@ jobs:
|
|||
uv sync --frozen --all-groups --all-extras --python 3.12
|
||||
- start_postgres
|
||||
- start_redis
|
||||
- start_fake_openai_endpoint
|
||||
- attach_workspace:
|
||||
at: ~/project
|
||||
- run:
|
||||
|
|
@ -1961,6 +1994,7 @@ jobs:
|
|||
-e REDIS_PORT=6379 \
|
||||
-e LITELLM_MASTER_KEY="sk-1234" \
|
||||
-e OPENAI_API_KEY=$OPENAI_API_KEY \
|
||||
-e FAKE_OPENAI_API_BASE=http://host.docker.internal:8190 \
|
||||
-e LITELLM_LICENSE=$LITELLM_LICENSE \
|
||||
-e AWS_ACCESS_KEY_ID=$AWS_ACCESS_KEY_ID \
|
||||
-e AWS_SECRET_ACCESS_KEY=$AWS_SECRET_ACCESS_KEY \
|
||||
|
|
@ -2020,6 +2054,7 @@ jobs:
|
|||
command: |
|
||||
uv sync --frozen --all-groups --all-extras --python 3.12
|
||||
- start_postgres
|
||||
- start_fake_openai_endpoint
|
||||
- attach_workspace:
|
||||
at: ~/project
|
||||
- run:
|
||||
|
|
@ -2039,6 +2074,7 @@ jobs:
|
|||
-e REDIS_PASSWORD=$REDIS_PASSWORD \
|
||||
-e REDIS_PORT=$REDIS_PORT \
|
||||
-e LITELLM_MASTER_KEY="sk-1234" \
|
||||
-e FAKE_OPENAI_API_BASE=http://host.docker.internal:8190 \
|
||||
-e LITELLM_LICENSE=$LITELLM_LICENSE \
|
||||
-e USE_DDTRACE=True \
|
||||
-e DD_API_KEY=$DD_API_KEY \
|
||||
|
|
@ -2060,6 +2096,7 @@ jobs:
|
|||
-e REDIS_PASSWORD=$REDIS_PASSWORD \
|
||||
-e REDIS_PORT=$REDIS_PORT \
|
||||
-e LITELLM_MASTER_KEY="sk-1234" \
|
||||
-e FAKE_OPENAI_API_BASE=http://host.docker.internal:8190 \
|
||||
-e LITELLM_LICENSE=$LITELLM_LICENSE \
|
||||
-e USE_DDTRACE=True \
|
||||
-e DD_API_KEY=$DD_API_KEY \
|
||||
|
|
@ -2112,6 +2149,7 @@ jobs:
|
|||
command: |
|
||||
uv sync --frozen --all-groups --all-extras --python 3.12
|
||||
- start_postgres
|
||||
- start_fake_openai_endpoint
|
||||
- attach_workspace:
|
||||
at: ~/project
|
||||
- run:
|
||||
|
|
@ -2129,6 +2167,7 @@ jobs:
|
|||
-e DATABASE_URL=postgresql://postgres:postgres@host.docker.internal:5432/circle_test \
|
||||
-e STORE_MODEL_IN_DB="True" \
|
||||
-e LITELLM_MASTER_KEY="sk-1234" \
|
||||
-e FAKE_OPENAI_API_BASE=http://host.docker.internal:8190 \
|
||||
-e LITELLM_LICENSE=$LITELLM_LICENSE \
|
||||
--add-host host.docker.internal:host-gateway \
|
||||
--name my-app \
|
||||
|
|
@ -2187,6 +2226,7 @@ jobs:
|
|||
command: |
|
||||
docker build -t my-app:latest -f docker/build_from_pip/Dockerfile.build_from_pip .
|
||||
- start_postgres
|
||||
- start_fake_openai_endpoint
|
||||
- run:
|
||||
name: Run Docker container
|
||||
# intentionally give bad redis credentials here
|
||||
|
|
@ -2200,6 +2240,7 @@ jobs:
|
|||
-e REDIS_PORT=$REDIS_PORT \
|
||||
-e LITELLM_MASTER_KEY="sk-1234" \
|
||||
-e OPENAI_API_KEY=$OPENAI_API_KEY \
|
||||
-e FAKE_OPENAI_API_BASE=http://host.docker.internal:8190 \
|
||||
-e LITELLM_LICENSE=$LITELLM_LICENSE \
|
||||
-e OTEL_EXPORTER="in_memory" \
|
||||
-e AWS_ACCESS_KEY_ID=$AWS_ACCESS_KEY_ID \
|
||||
|
|
|
|||
2
.github/pull_request_template.md
vendored
2
.github/pull_request_template.md
vendored
|
|
@ -4,7 +4,7 @@
|
|||
|
||||
## Linear ticket
|
||||
|
||||
<!-- if you are an internal contributor, add the Linear ticket e.g. "Resolves LIT-1234" to magically link the Linear ticket to the GitHub PR -->
|
||||
<!-- if you are an internal contributor (e.g., your username is postfixed with -berri or -berriai), add "Resolves " followed by the Linear ticket e.g. "Resolves LIT-1234" to magically link the Linear ticket to the GitHub PR -->
|
||||
|
||||
## Pre-Submission checklist
|
||||
|
||||
|
|
|
|||
50
.github/scripts/_agent_shin_actions.py
vendored
Normal file
50
.github/scripts/_agent_shin_actions.py
vendored
Normal file
|
|
@ -0,0 +1,50 @@
|
|||
"""Dry-run wrapper(s) around Agent Shin GitHub mutations.
|
||||
|
||||
The rollout scripts currently need only one mutation wrapped, so this module
|
||||
exposes a single ``maybe_post_comment`` helper. It takes a ``dry_run: bool``
|
||||
keyword argument and the body is intentionally trivial:
|
||||
|
||||
if dry_run:
|
||||
print(...) # log what we would do, return
|
||||
return
|
||||
real_mutation(...) # otherwise, actually do it
|
||||
|
||||
That shape means a dry-run preview differs from the real run in exactly one
|
||||
line per side effect: the call site. So when you `python3 script.py` locally
|
||||
without ``--close``, you can be confident the actions printed are the ones the
|
||||
GitHub Action would have performed (modulo ordering on retry/error paths,
|
||||
which are deliberately simple). Any further mutation a rollout script needs
|
||||
should get the same ``maybe_*`` treatment instead of calling the raw
|
||||
``triage_with_llm`` mutation directly.
|
||||
|
||||
Importing from this module pulls in the real mutation from ``triage_with_llm``
|
||||
— call sites in the rollout scripts should NEVER import ``post_comment``
|
||||
directly; that would skip the dry-run gate and is the bug class this module
|
||||
exists to prevent.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
import textwrap
|
||||
|
||||
# Import the module itself rather than the bare names so monkeypatching
|
||||
# `triage_with_llm.post_comment` (or any of the other mutations) in tests is
|
||||
# reflected here — `from triage_with_llm import post_comment` would bind the
|
||||
# original function to a local name and bypass the patch, defeating the whole
|
||||
# point of these wrappers.
|
||||
import triage_with_llm
|
||||
|
||||
|
||||
def _log(line: str) -> None:
|
||||
"""Print a single dry-run line to stdout (one log statement per side effect)."""
|
||||
print(line, file=sys.stdout, flush=True)
|
||||
|
||||
|
||||
def maybe_post_comment(repo: str, number: int, body: str, *, dry_run: bool) -> None:
|
||||
"""Post a comment on ``repo#number`` — or, in dry-run, log what we would post."""
|
||||
if dry_run:
|
||||
_log(f"[DRY RUN] comment {repo}#{number}:")
|
||||
_log(textwrap.indent(body, " "))
|
||||
return
|
||||
triage_with_llm.post_comment(repo, number, body)
|
||||
211
.github/scripts/agent_shin_shared.py
vendored
Normal file
211
.github/scripts/agent_shin_shared.py
vendored
Normal file
|
|
@ -0,0 +1,211 @@
|
|||
"""Constants and helpers shared by Agent Shin's triage scripts.
|
||||
|
||||
Both `triage_with_llm.py` (the LLM-judge entrypoint) and
|
||||
`close_low_quality_prs.py` (the daily Greptile-score sweep) need to
|
||||
agree on the same notions of:
|
||||
|
||||
* What counts as a Greptile-authored review comment
|
||||
(``GREPTILE_BOT_LOGINS``) and how to extract a confidence score from
|
||||
its body (``SCORE_PATTERN`` / :func:`extract_greptile_score`).
|
||||
* How long the 2-hour grace window is (``GRACE_PERIOD_SECONDS``) and
|
||||
the HTML marker stamped into a grace-warning comment so the *other*
|
||||
script can see "Agent Shin already warned" and behave accordingly
|
||||
(``GRACE_COMMENT_MARKER``).
|
||||
* Who Agent Shin is on GitHub (``AGENT_SHIN_DEFAULT_BOT_LOGIN``).
|
||||
* How GitHub-style ISO-8601 timestamps round-trip into timezone-aware
|
||||
:class:`datetime.datetime` (:func:`parse_iso8601`).
|
||||
|
||||
Keeping these in one module means a future change (new Greptile output
|
||||
format, a longer grace window, a new allowlisted account) is a single edit
|
||||
instead of two — the original split version had to call out in comments
|
||||
that the two copies "must stay in sync" precisely because nothing
|
||||
enforced it.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime as dt
|
||||
import json
|
||||
import os
|
||||
import re
|
||||
import subprocess
|
||||
from typing import Iterable
|
||||
|
||||
GREPTILE_BOT_LOGINS = frozenset({"greptile-apps", "greptile-apps[bot]"})
|
||||
|
||||
SCORE_PATTERN = re.compile(
|
||||
r"confidence\s*score\s*[:\-]?\s*(\d+)\s*/\s*5",
|
||||
re.IGNORECASE,
|
||||
)
|
||||
|
||||
GRACE_COMMENT_MARKER = "<!-- agent-shin:grace-warning -->"
|
||||
|
||||
# Hidden HTML marker stamped on every Agent Shin auto-close comment (the LLM
|
||||
# judge's grace/review-gate close and the daily Greptile sweep's close).
|
||||
# `was_closed_by_agent_shin` requires this marker — not just the closing actor —
|
||||
# before `@agent-shin reconsider` may reopen, because the `github-actions[bot]`
|
||||
# identity is shared with every other workflow in the repo and is not unique to
|
||||
# Agent Shin. Both close paths must stamp it or the reconsider path silently
|
||||
# rejects the contributor.
|
||||
AGENT_SHIN_CLOSE_MARKER = "<!-- agent-shin:closed -->"
|
||||
|
||||
# 2 hours between the grace warning and the auto-close. Short enough to
|
||||
# dogfood the "fix it before it closes" loop in one sitting; bump back up
|
||||
# (e.g. 86400 for a day) for the public rollout.
|
||||
GRACE_PERIOD_SECONDS = 7200
|
||||
|
||||
AGENT_SHIN_DEFAULT_BOT_LOGIN = "github-actions[bot]"
|
||||
|
||||
|
||||
def _logins(*names: str) -> frozenset[str]:
|
||||
"""Build a login set normalized for case-insensitive membership checks.
|
||||
|
||||
Callers compare via ``login.lower() in <set>``, so the stored values
|
||||
must be lowercase. Normalizing here lets the literals keep each
|
||||
account's canonical GitHub casing (e.g. ``SwiftWinds``) for
|
||||
readability without breaking the lookup.
|
||||
"""
|
||||
return frozenset(name.lower() for name in names)
|
||||
|
||||
|
||||
# Dogfood rollout gate. While this set is non-empty, Agent Shin acts ONLY on
|
||||
# PRs/issues authored by these logins and skips everyone else. For an
|
||||
# allowlisted author the usual internal/external classification is bypassed, so
|
||||
# an internal account (e.g. a maintainer's own work login) still gets triaged
|
||||
# while the bot is being tested on a small set of accounts. Empty the set to
|
||||
# lift the restriction and restore full triage for the public rollout. Logins
|
||||
# are compared case-insensitively.
|
||||
ALLOWLIST_LOGINS = _logins("mateo-berri", "SwiftWinds")
|
||||
|
||||
# `gh {pr,issue} list` has no "fetch everything" flag — `--limit` is the only
|
||||
# control and it defaults to 30. Pass a ceiling far above any realistic open
|
||||
# backlog (low thousands today) so gh paginates the API until the queue is
|
||||
# exhausted rather than silently truncating. The bulk sweeps MUST see the whole
|
||||
# backlog: gh lists newest-first, so a low cap drops the *oldest* PRs/issues —
|
||||
# exactly the stale ones a low-quality sweep is meant to catch.
|
||||
GH_LIST_ALL_LIMIT = 100_000
|
||||
|
||||
|
||||
def extract_greptile_score(comments: Iterable[dict]) -> tuple[int, dict] | None:
|
||||
"""Return (score, comment) for the most recent Greptile-authored comment
|
||||
that contains a "Confidence Score: X/5". Returns None if no such comment.
|
||||
|
||||
"Most recent" is determined by the comment's `updated_at` (falling back to
|
||||
`created_at`), so re-reviews override earlier passes.
|
||||
"""
|
||||
candidates: list[tuple[str, int, dict]] = []
|
||||
for comment in comments:
|
||||
user = (comment.get("user") or {}).get("login", "")
|
||||
if user not in GREPTILE_BOT_LOGINS:
|
||||
continue
|
||||
body = comment.get("body") or ""
|
||||
match = SCORE_PATTERN.search(body)
|
||||
if not match:
|
||||
continue
|
||||
score = int(match.group(1))
|
||||
timestamp = comment.get("updated_at") or comment.get("created_at") or ""
|
||||
candidates.append((timestamp, score, comment))
|
||||
|
||||
if not candidates:
|
||||
return None
|
||||
|
||||
candidates.sort(key=lambda triple: triple[0])
|
||||
_, score, comment = candidates[-1]
|
||||
return score, comment
|
||||
|
||||
|
||||
def parse_iso8601(value: str) -> dt.datetime:
|
||||
"""Parse a GitHub ISO-8601 timestamp into a timezone-aware datetime."""
|
||||
return dt.datetime.fromisoformat(value.replace("Z", "+00:00"))
|
||||
|
||||
|
||||
def gh(*args: str) -> str:
|
||||
"""Run a `gh` CLI command and return stdout. Raises on non-zero exit.
|
||||
|
||||
Shared by both Agent Shin entrypoints so a future change here
|
||||
(timeout handling, logging, retry on transient failures) only needs
|
||||
to be made once.
|
||||
"""
|
||||
result = subprocess.run(
|
||||
["gh", *args],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
check=True,
|
||||
)
|
||||
return result.stdout
|
||||
|
||||
|
||||
def list_open_items(kind: str, *, repo: str | None, fields: str) -> list[dict]:
|
||||
"""Return EVERY open PR (``kind="pr"``) or issue (``kind="issue"``) in ``repo``.
|
||||
|
||||
Wraps ``gh {pr,issue} list`` with ``--limit GH_LIST_ALL_LIMIT`` so the full
|
||||
backlog is fetched instead of the default 30 (or any other arbitrary cap).
|
||||
Both bulk sweeps — the daily Greptile closer and the one-shot rollout
|
||||
heads-up — rely on this seeing the whole queue, including the oldest items.
|
||||
|
||||
``fields`` is the comma-separated ``--json`` field list the caller needs
|
||||
(e.g. ``"number"`` for the rollout, the full set for the closer).
|
||||
"""
|
||||
if kind not in ("pr", "issue"):
|
||||
raise ValueError(f"kind must be 'pr' or 'issue', got {kind!r}")
|
||||
repo_args = ["--repo", repo] if repo else []
|
||||
raw = gh(
|
||||
kind,
|
||||
"list",
|
||||
"--state",
|
||||
"open",
|
||||
"--limit",
|
||||
str(GH_LIST_ALL_LIMIT),
|
||||
"--json",
|
||||
fields,
|
||||
*repo_args,
|
||||
)
|
||||
return json.loads(raw)
|
||||
|
||||
|
||||
def seconds_since_latest_marker_comment(
|
||||
comments: Iterable[dict],
|
||||
*,
|
||||
marker: str,
|
||||
bot_login: str | None = None,
|
||||
now: dt.datetime | None = None,
|
||||
) -> float | None:
|
||||
"""Return seconds since the bot's most recent comment containing ``marker``.
|
||||
|
||||
Filters comments by author so a contributor who quotes the HTML
|
||||
marker (e.g. via GitHub's "Quote reply" feature, which preserves
|
||||
HTML comments in the raw markdown of the quoted text) is not
|
||||
mistaken for a bot warning — that would silently reset cooldown
|
||||
timers and suppress legitimate notifications.
|
||||
|
||||
``bot_login`` defaults to the `AGENT_SHIN_BOT_LOGIN` env override or
|
||||
``AGENT_SHIN_DEFAULT_BOT_LOGIN`` so callers normally don't need to
|
||||
pass it. ``now`` is injectable for tests / callers (like the daily
|
||||
sweep) that want every age calculation pinned to one snapshot.
|
||||
"""
|
||||
expected_login = (
|
||||
bot_login
|
||||
or os.environ.get("AGENT_SHIN_BOT_LOGIN")
|
||||
or AGENT_SHIN_DEFAULT_BOT_LOGIN
|
||||
).lower()
|
||||
latest: dt.datetime | None = None
|
||||
for comment in comments:
|
||||
author = ((comment.get("user") or {}).get("login") or "").lower()
|
||||
if author != expected_login:
|
||||
continue
|
||||
body = comment.get("body") or ""
|
||||
if marker not in body:
|
||||
continue
|
||||
created = comment.get("created_at")
|
||||
if not created:
|
||||
continue
|
||||
try:
|
||||
ts = parse_iso8601(created)
|
||||
except ValueError:
|
||||
continue
|
||||
if latest is None or ts > latest:
|
||||
latest = ts
|
||||
if latest is None:
|
||||
return None
|
||||
reference = now if now is not None else dt.datetime.now(dt.timezone.utc)
|
||||
return (reference - latest).total_seconds()
|
||||
573
.github/scripts/close_low_quality_prs.py
vendored
Normal file
573
.github/scripts/close_low_quality_prs.py
vendored
Normal file
|
|
@ -0,0 +1,573 @@
|
|||
#!/usr/bin/env python3
|
||||
"""
|
||||
Auto-close low-quality pull requests.
|
||||
|
||||
Closes open PRs (including drafts, regardless of age) that satisfy ALL of:
|
||||
1. Have a Greptile (`greptile-apps`) review comment whose latest
|
||||
"Confidence Score: X/5" is below the configured threshold (default: 4).
|
||||
2. Are authored by an external OSS contributor (internal BerriAI
|
||||
contributors are exempt).
|
||||
3. Do not carry an opt-out label (default: "do not close").
|
||||
|
||||
`--min-age-days` is retained as an opt-in safety net for one-off backfill
|
||||
runs (default: 0). The team's intent is that the count of open PRs equals
|
||||
the count of PRs internal collaborators need to action on, so neither age
|
||||
nor draft status acts as a free pass.
|
||||
|
||||
For each match, the script posts an explanatory comment and closes the PR.
|
||||
Because OSS contributors *cannot* reopen a PR closed by the bot/maintainer
|
||||
(GitHub limitation), the close-comment instructs them to push their fixes
|
||||
and **open a fresh PR**, or to comment `@agent-shin reconsider` on the
|
||||
closed PR to have the LLM judge re-evaluate (and reopen on pass).
|
||||
|
||||
Requires the `gh` CLI to be authenticated.
|
||||
|
||||
Usage examples:
|
||||
# Dry run (default) - prints what would be closed
|
||||
python3 close_low_quality_prs.py
|
||||
|
||||
# Actually close matching PRs
|
||||
python3 close_low_quality_prs.py --close
|
||||
|
||||
# Restrict to PRs at least N days old (one-off backfill safety net)
|
||||
python3 close_low_quality_prs.py --min-age-days 7 --min-score 4 --close
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import datetime as dt
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
from typing import Iterable
|
||||
|
||||
# Add this script's directory to `sys.path` so the sibling
|
||||
# `agent_shin_shared` module is importable when the script is invoked
|
||||
# directly (e.g. `python3 .github/scripts/close_low_quality_prs.py ...`).
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
from agent_shin_shared import ( # noqa: E402 -- sys.path adjusted above
|
||||
AGENT_SHIN_CLOSE_MARKER,
|
||||
ALLOWLIST_LOGINS,
|
||||
GRACE_COMMENT_MARKER,
|
||||
GRACE_PERIOD_SECONDS,
|
||||
GREPTILE_BOT_LOGINS,
|
||||
SCORE_PATTERN,
|
||||
extract_greptile_score,
|
||||
gh,
|
||||
list_open_items,
|
||||
parse_iso8601,
|
||||
seconds_since_latest_marker_comment,
|
||||
)
|
||||
|
||||
# `GREPTILE_BOT_LOGINS` and `SCORE_PATTERN` (Greptile's GitHub App login
|
||||
# variants and the "Confidence Score: X/5" regex) are imported from
|
||||
# `agent_shin_shared` so the LLM judge in `triage_with_llm.py` and this
|
||||
# daily Greptile sweep read the score through the same set of logins
|
||||
# and the same regex.
|
||||
|
||||
# `author_association` values for internal BerriAI contributors who should be
|
||||
# exempt from auto-triage.
|
||||
INTERNAL_AUTHOR_ASSOCIATIONS = frozenset({"OWNER", "MEMBER", "COLLABORATOR"})
|
||||
|
||||
# Default labels that exempt a PR from auto-close. Defined at module scope (not
|
||||
# as a mutable argparse default) so that `--optout-label foo` REPLACES the
|
||||
# defaults instead of appending to them — the argparse `action="append"` +
|
||||
# `default=[...]` combination silently mutates the shared default list.
|
||||
DEFAULT_OPTOUT_LABELS = ("do not close", "keep open", "wip")
|
||||
|
||||
# `GRACE_COMMENT_MARKER` (HTML marker appended to grace-period warning
|
||||
# comments — used by either script to recognize that a warning was
|
||||
# already posted) and `GRACE_PERIOD_SECONDS` (length of the grace
|
||||
# period between the warning and the actual auto-close, 2 hours) are
|
||||
# imported from `agent_shin_shared` so the Agent Shin LLM judge and
|
||||
# this daily Greptile sweep agree on the same marker and duration.
|
||||
|
||||
|
||||
def fetch_open_prs(repo: str | None) -> list[dict]:
|
||||
"""Fetch all open PRs (number, createdAt, isDraft, labels, author).
|
||||
|
||||
Includes drafts: `gh pr list --state open` returns both ready-for-review
|
||||
and draft PRs by default. This is the desired behavior — drafts are not
|
||||
a free pass; the internal-collaborator open-PR queue should reflect every
|
||||
PR that needs human attention regardless of draft status.
|
||||
"""
|
||||
fields = "number,title,createdAt,isDraft,labels,author,url"
|
||||
return list_open_items("pr", repo=repo, fields=fields)
|
||||
|
||||
|
||||
def fetch_pr_author_association(pr_number: int, repo: str | None) -> str:
|
||||
"""Return the GitHub `author_association` for a PR, uppercase.
|
||||
|
||||
Values: OWNER, MEMBER, COLLABORATOR, CONTRIBUTOR, FIRST_TIME_CONTRIBUTOR,
|
||||
FIRST_TIMER, MANNEQUIN, NONE. Returns "" on lookup failure.
|
||||
"""
|
||||
endpoint = (
|
||||
f"repos/{repo}/pulls/{pr_number}"
|
||||
if repo
|
||||
else f"repos/{{owner}}/{{repo}}/pulls/{pr_number}"
|
||||
)
|
||||
try:
|
||||
data = json.loads(gh("api", endpoint))
|
||||
except subprocess.CalledProcessError:
|
||||
return ""
|
||||
return (data.get("author_association") or "").upper()
|
||||
|
||||
|
||||
def is_external_pr_author(pr: dict, repo: str | None) -> bool:
|
||||
"""Return True if the PR author is an external OSS contributor.
|
||||
|
||||
Internal = `OWNER` / `MEMBER` / `COLLABORATOR` association, or a bot login.
|
||||
"""
|
||||
login = ((pr.get("author") or {}).get("login") or "").lower()
|
||||
if login.endswith("[bot]") or login in {"dependabot", "github-actions"}:
|
||||
return False
|
||||
association = fetch_pr_author_association(pr["number"], repo)
|
||||
# Fail-safe: if the API lookup failed (empty string), treat the author as
|
||||
# internal so we don't auto-close their PR. Auto-close is destructive, so
|
||||
# an unknown association should never make a PR eligible for closing.
|
||||
if not association or association in INTERNAL_AUTHOR_ASSOCIATIONS:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def fetch_pr_comments(pr_number: int, repo: str | None) -> list[dict]:
|
||||
"""Fetch issue-level comments on a PR (where Greptile posts its summary)."""
|
||||
endpoint = (
|
||||
f"repos/{repo}/issues/{pr_number}/comments?per_page=100"
|
||||
if repo
|
||||
else f"repos/{{owner}}/{{repo}}/issues/{pr_number}/comments?per_page=100"
|
||||
)
|
||||
raw = gh("api", "--paginate", endpoint)
|
||||
comments: list[dict] = []
|
||||
for line in raw.strip().splitlines():
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
try:
|
||||
parsed = json.loads(line)
|
||||
except json.JSONDecodeError:
|
||||
# A malformed line should not blow up the whole sweep. Skip and
|
||||
# carry on so the remaining PRs in this run still get evaluated.
|
||||
continue
|
||||
if isinstance(parsed, list):
|
||||
comments.extend(parsed)
|
||||
else:
|
||||
comments.append(parsed)
|
||||
return comments
|
||||
|
||||
|
||||
def has_optout_label(pr: dict, optout_labels: set[str]) -> bool:
|
||||
labels = {label.get("name", "").lower() for label in pr.get("labels", [])}
|
||||
return bool(labels & {lbl.lower() for lbl in optout_labels})
|
||||
|
||||
|
||||
def seconds_since_last_grace_warning(
|
||||
comments: Iterable[dict],
|
||||
*,
|
||||
bot_login: str | None = None,
|
||||
now: dt.datetime | None = None,
|
||||
) -> float | None:
|
||||
"""Return seconds since the bot's most recent grace-period warning, or
|
||||
None if no such warning has ever been posted on this PR.
|
||||
|
||||
Thin wrapper over
|
||||
`agent_shin_shared.seconds_since_latest_marker_comment` — the
|
||||
centralized helper handles the bot-author filter, marker match,
|
||||
timestamp parsing, and `now` injection. Keeping this wrapper
|
||||
preserves the closer's "already-fetched comments + injectable now"
|
||||
interface so callers (and tests) don't need to change.
|
||||
"""
|
||||
return seconds_since_latest_marker_comment(
|
||||
comments,
|
||||
marker=GRACE_COMMENT_MARKER,
|
||||
bot_login=bot_login,
|
||||
now=now,
|
||||
)
|
||||
|
||||
|
||||
def format_grace_warning_comment(score: int, threshold: int) -> str:
|
||||
"""Comment posted on the FIRST low-Greptile-score detection — gives
|
||||
the contributor a 2-hour grace window before the auto-close fires on
|
||||
the next daily cron run.
|
||||
|
||||
Mirrors `format_grace_warning_pr_comment` in
|
||||
`triage_with_llm.py` in spirit (2-hour grace + escape hatches), but
|
||||
framed around Greptile's confidence score instead of the LLM judge's
|
||||
rubric since the close trigger here is the Greptile signal.
|
||||
"""
|
||||
return (
|
||||
"🚅 Hi, thanks for the PR! I'm **Agent Shin**, the automated triage bot for this "
|
||||
"repository.\n"
|
||||
"\n"
|
||||
"Heads up: Greptile's most recent review scored this PR "
|
||||
f"**{score}/5**, below our merge bar of **{threshold}/5**.\n"
|
||||
"\n"
|
||||
"If the score isn't lifted in the next **2 hours**, I'll auto-close this PR. That's "
|
||||
"**not** us saying the change isn't worthwhile. We want the open-PR list to mirror "
|
||||
"what a maintainer can act on *right now*, so contributors like you don't get lost in "
|
||||
"a backlog. Take your time; everything below still works after the close.\n"
|
||||
"\n"
|
||||
"**During the grace period:** push fixes that address Greptile's feedback, then comment "
|
||||
"`@greptileai` to request a fresh review. If "
|
||||
f"the new score is **{threshold}/5 or higher**, the PR stays open and no further "
|
||||
"action is needed on your side.\n"
|
||||
"\n"
|
||||
"**If the PR does get auto-closed in 2 hours, you still have an easy recovery path:**\n"
|
||||
"\n"
|
||||
"- Comment `@greptileai` to request a fresh review. **This still works even after "
|
||||
f"the PR is closed**, and a score of {threshold}/5 or higher is one of the signals "
|
||||
"that lifts the PR back into the review queue. A low Greptile score isn't a blocker.\n"
|
||||
"- Comment `@agent-shin reconsider` after pushing fixes; I'll re-run the rubric and "
|
||||
"reopen the PR if both gates (description rubric + Greptile score) now pass.\n"
|
||||
"\n"
|
||||
f"{GRACE_COMMENT_MARKER}"
|
||||
)
|
||||
|
||||
|
||||
def post_grace_warning(
|
||||
pr: dict,
|
||||
score: int,
|
||||
threshold: int,
|
||||
repo: str | None,
|
||||
dry_run: bool,
|
||||
) -> None:
|
||||
"""Post the 2-hour grace-period warning comment on `pr`.
|
||||
|
||||
The warning carries `GRACE_COMMENT_MARKER` so subsequent runs can
|
||||
detect that the contributor has already been told about the
|
||||
pending close. Does NOT close the PR — the close happens on the
|
||||
next eligible run after `GRACE_PERIOD_SECONDS` elapses (handled
|
||||
by `close_pr`).
|
||||
"""
|
||||
pr_number = pr["number"]
|
||||
repo_args = ["--repo", repo] if repo else []
|
||||
|
||||
if dry_run:
|
||||
print(
|
||||
f" [DRY RUN] Would post grace warning to PR #{pr_number} "
|
||||
f"(greptile={score}/5): {pr['title']}"
|
||||
)
|
||||
return
|
||||
|
||||
comment_body = format_grace_warning_comment(score, threshold)
|
||||
gh("pr", "comment", str(pr_number), "--body", comment_body, *repo_args)
|
||||
print(f" Posted grace warning on PR #{pr_number} (greptile={score}/5)")
|
||||
|
||||
|
||||
def format_close_comment(score: int, threshold: int) -> str:
|
||||
"""Comment posted when a low-Greptile-score PR is auto-closed.
|
||||
|
||||
Carries `AGENT_SHIN_CLOSE_MARKER` so the `@agent-shin reconsider` path
|
||||
(guarded by `was_closed_by_agent_shin`) recognizes this as an Agent Shin
|
||||
close and is allowed to reopen the PR once it passes again; without the
|
||||
marker that recovery path the comment advertises silently rejects the
|
||||
contributor.
|
||||
"""
|
||||
score_sentence = (
|
||||
f"Greptile's most recent review scored this PR **{score}/5**, below "
|
||||
f"our merge bar of **{threshold}/5**, and the 2-hour grace period since "
|
||||
"the warning has elapsed.\n\n"
|
||||
)
|
||||
return (
|
||||
f"Closing as part of automated PR triage.\n\n"
|
||||
f"{score_sentence}"
|
||||
"We close low-confidence PRs aggressively to keep the review queue "
|
||||
"manageable for maintainers and contributors alike. **This is not a "
|
||||
"rejection of the idea.** To bring this back:\n\n"
|
||||
"1. Push the fixes that address Greptile's feedback (continue using "
|
||||
"your existing branch is fine).\n"
|
||||
"2. **Open a new PR** with the updated branch. Greptile will review "
|
||||
"it again, and if it scores "
|
||||
f"**{threshold}/5 or higher** a maintainer will take another look.\n\n"
|
||||
"_Why open a new PR instead of reopening this one?_ GitHub does not "
|
||||
"let external contributors reopen a PR that was closed by a bot or "
|
||||
"maintainer, so a fresh PR is the most reliable path forward. If you "
|
||||
"would prefer this exact PR re-evaluated, comment "
|
||||
"`@agent-shin reconsider` once you've pushed the fixes; Agent Shin "
|
||||
"will re-run triage and reopen this PR if it now meets the bar. "
|
||||
"You can also comment `@greptileai` to request a fresh Greptile "
|
||||
"review; that works **even after the PR is closed**.\n\n"
|
||||
"Thanks for contributing to LiteLLM. We know auto-closures can sting; "
|
||||
"the goal is to keep the project healthy, not to dismiss your work."
|
||||
f"\n\n{AGENT_SHIN_CLOSE_MARKER}"
|
||||
)
|
||||
|
||||
|
||||
def close_pr(
|
||||
pr: dict,
|
||||
score: int,
|
||||
threshold: int,
|
||||
age_days: int,
|
||||
repo: str | None,
|
||||
dry_run: bool,
|
||||
label: str | None,
|
||||
) -> None:
|
||||
"""Post the explanatory comment and close the PR."""
|
||||
pr_number = pr["number"]
|
||||
repo_args = ["--repo", repo] if repo else []
|
||||
|
||||
if dry_run:
|
||||
print(
|
||||
f" [DRY RUN] Would close PR #{pr_number} "
|
||||
f"(age={age_days}d, greptile={score}/5): {pr['title']}"
|
||||
)
|
||||
return
|
||||
|
||||
comment_body = format_close_comment(score, threshold)
|
||||
gh("pr", "comment", str(pr_number), "--body", comment_body, *repo_args)
|
||||
|
||||
if label:
|
||||
try:
|
||||
gh("pr", "edit", str(pr_number), "--add-label", label, *repo_args)
|
||||
except subprocess.CalledProcessError as exc:
|
||||
stderr = (exc.stderr or "").strip()
|
||||
print(f" warn: failed to add label '{label}' to #{pr_number}: {stderr}")
|
||||
|
||||
gh("pr", "close", str(pr_number), *repo_args)
|
||||
print(f" Closed PR #{pr_number} (greptile={score}/5, age={age_days}d)")
|
||||
|
||||
|
||||
def evaluate_pr(
|
||||
pr: dict,
|
||||
now: dt.datetime,
|
||||
min_age_days: int,
|
||||
min_score: int,
|
||||
repo: str | None,
|
||||
optout_labels: set[str],
|
||||
allowlist: frozenset[str] = ALLOWLIST_LOGINS,
|
||||
) -> tuple[str, int | None, int | None]:
|
||||
"""Decide what to do with `pr` on this triage run.
|
||||
|
||||
Returns (action, score_or_none, age_days_or_none) where action is one of:
|
||||
"skip-too-young", "skip-optout-label", "skip-not-allowlisted",
|
||||
"skip-internal", "skip-no-greptile-score", "skip-score-ok",
|
||||
"warn-grace", "skip-in-grace-period", or "close".
|
||||
|
||||
Drafts are NOT skipped — the goal is "open PR count == PRs internal
|
||||
collaborators need to action on", and a draft that Greptile scored <4/5
|
||||
is still in that queue. Authors can opt out via the `wip` label (see
|
||||
`DEFAULT_OPTOUT_LABELS`) if they need to keep a long-lived draft open.
|
||||
|
||||
Grace-period semantics: the first time a PR fails the rubric, the
|
||||
action is `warn-grace` — the caller should post a warning comment but
|
||||
NOT close the PR. On a subsequent run, if the warning is still less
|
||||
than `GRACE_PERIOD_SECONDS` old AND the PR still fails, the action is
|
||||
`skip-in-grace-period`. Once the warning ages out and the rubric is
|
||||
still failing, the action is `close`.
|
||||
"""
|
||||
if has_optout_label(pr, optout_labels):
|
||||
return ("skip-optout-label", None, None)
|
||||
|
||||
created = parse_iso8601(pr["createdAt"])
|
||||
age_days = (now - created).days
|
||||
# `min_age_days` defaults to 0 (close as soon as Greptile scores low).
|
||||
# Set a positive value via --min-age-days for one-off backfill runs that
|
||||
# want to skip very-young PRs.
|
||||
if min_age_days > 0 and age_days < min_age_days:
|
||||
return ("skip-too-young", None, age_days)
|
||||
|
||||
# While the allowlist is active it is the sole author gate: only those
|
||||
# logins are acted on and the external-only restriction is bypassed for
|
||||
# them. Otherwise auto-close only external OSS contributors — internal
|
||||
# contributors (BerriAI org members) handle their own backlog.
|
||||
login = ((pr.get("author") or {}).get("login") or "").lower()
|
||||
if allowlist:
|
||||
if login not in allowlist:
|
||||
return ("skip-not-allowlisted", None, age_days)
|
||||
elif not is_external_pr_author(pr, repo):
|
||||
return ("skip-internal", None, age_days)
|
||||
|
||||
comments = fetch_pr_comments(pr["number"], repo)
|
||||
extraction = extract_greptile_score(comments)
|
||||
if extraction is None:
|
||||
return ("skip-no-greptile-score", None, age_days)
|
||||
|
||||
score, _ = extraction
|
||||
if score >= min_score:
|
||||
return ("skip-score-ok", score, age_days)
|
||||
|
||||
grace_age = seconds_since_last_grace_warning(comments, now=now)
|
||||
if grace_age is None:
|
||||
return ("warn-grace", score, age_days)
|
||||
if grace_age < GRACE_PERIOD_SECONDS:
|
||||
return ("skip-in-grace-period", score, age_days)
|
||||
|
||||
return ("close", score, age_days)
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument(
|
||||
"--repo",
|
||||
type=str,
|
||||
default=None,
|
||||
help="Repository (owner/repo). Auto-detected if omitted.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--min-age-days",
|
||||
type=int,
|
||||
default=0,
|
||||
help=(
|
||||
"Minimum age (in days) before a PR is eligible. Default 0 = "
|
||||
"close as soon as Greptile flags it. Set a positive value for "
|
||||
"one-off backfill runs that want to spare very-young PRs."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--min-score",
|
||||
type=int,
|
||||
default=4,
|
||||
choices=range(1, 6),
|
||||
help="Greptile score below which a PR is closed (default: 4 -> closes <4/5).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--optout-label",
|
||||
action="append",
|
||||
default=None,
|
||||
help=(
|
||||
"Label(s) that exempt a PR from auto-close. Repeat to add more. "
|
||||
"Case-insensitive. When omitted, defaults to "
|
||||
f"{list(DEFAULT_OPTOUT_LABELS)!r}; passing this flag REPLACES the "
|
||||
"defaults (argparse `append` with a mutable default would append "
|
||||
"instead, which we explicitly avoid)."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--close-label",
|
||||
type=str,
|
||||
default=None,
|
||||
help=(
|
||||
"Optional label to add to PRs that get auto-closed "
|
||||
"(e.g. 'auto-closed-low-quality'). Must already exist on the repo."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--close",
|
||||
action="store_true",
|
||||
help="Actually close matching PRs (default is dry-run).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--limit",
|
||||
type=int,
|
||||
default=None,
|
||||
help="Maximum number of PRs to close in one run (safety net).",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
dry_run = not args.close
|
||||
if dry_run:
|
||||
print("=== DRY RUN MODE (pass --close to actually close PRs) ===\n")
|
||||
|
||||
print("Fetching open PRs...")
|
||||
prs = fetch_open_prs(args.repo)
|
||||
print(f"Found {len(prs)} open PRs.\n")
|
||||
|
||||
now = dt.datetime.now(dt.timezone.utc)
|
||||
optout_labels = set(args.optout_label or DEFAULT_OPTOUT_LABELS)
|
||||
|
||||
closed = 0
|
||||
summary = {
|
||||
"close": 0,
|
||||
"warn-grace": 0,
|
||||
"skip-in-grace-period": 0,
|
||||
"skip-too-young": 0,
|
||||
"skip-optout-label": 0,
|
||||
"skip-not-allowlisted": 0,
|
||||
"skip-internal": 0,
|
||||
"skip-no-greptile-score": 0,
|
||||
"skip-score-ok": 0,
|
||||
}
|
||||
|
||||
# `warned` tracks grace-warning comments posted in this run so the
|
||||
# `--limit` safety net bounds *all* destructive write actions, not
|
||||
# just closures. Without this cap, a backlog of PRs failing the
|
||||
# threshold simultaneously could flood contributors with comments.
|
||||
warned = 0
|
||||
for pr in sorted(prs, key=lambda p: p["createdAt"]):
|
||||
try:
|
||||
action, score, age_days = evaluate_pr(
|
||||
pr,
|
||||
now,
|
||||
args.min_age_days,
|
||||
args.min_score,
|
||||
args.repo,
|
||||
optout_labels,
|
||||
)
|
||||
summary[action] = summary.get(action, 0) + 1
|
||||
|
||||
if action == "warn-grace":
|
||||
assert score is not None
|
||||
print(
|
||||
f"#{pr['number']}: \"{pr['title']}\" "
|
||||
f"(age={age_days}d, greptile={score}/5) -> warn-grace"
|
||||
)
|
||||
post_grace_warning(
|
||||
pr,
|
||||
score=score,
|
||||
threshold=args.min_score,
|
||||
repo=args.repo,
|
||||
dry_run=dry_run,
|
||||
)
|
||||
if not dry_run:
|
||||
warned += 1
|
||||
if args.limit is not None and (warned + closed) >= args.limit:
|
||||
print(
|
||||
f"\nReached --limit={args.limit} "
|
||||
f"(closed={closed}, warned={warned}); stopping."
|
||||
)
|
||||
break
|
||||
continue
|
||||
|
||||
if action != "close":
|
||||
continue
|
||||
|
||||
assert score is not None and age_days is not None
|
||||
print(
|
||||
f"#{pr['number']}: \"{pr['title']}\" "
|
||||
f"(age={age_days}d, greptile={score}/5) -> close"
|
||||
)
|
||||
close_pr(
|
||||
pr,
|
||||
score=score,
|
||||
threshold=args.min_score,
|
||||
age_days=age_days,
|
||||
repo=args.repo,
|
||||
dry_run=dry_run,
|
||||
label=args.close_label,
|
||||
)
|
||||
|
||||
if not dry_run:
|
||||
closed += 1
|
||||
if args.limit is not None and (warned + closed) >= args.limit:
|
||||
print(
|
||||
f"\nReached --limit={args.limit} "
|
||||
f"(closed={closed}, warned={warned}); stopping."
|
||||
)
|
||||
break
|
||||
except Exception as exc: # noqa: BLE001 - per-PR errors don't abort the sweep
|
||||
summary["error"] = summary.get("error", 0) + 1
|
||||
print(
|
||||
f"!! PR #{pr.get('number')}: {exc}",
|
||||
file=sys.stderr,
|
||||
)
|
||||
continue
|
||||
|
||||
print("\n=== Summary ===")
|
||||
for key, value in summary.items():
|
||||
print(f" {key:28s} {value}")
|
||||
if dry_run:
|
||||
print(f"\nTotal would close: {summary['close']}")
|
||||
else:
|
||||
print(f"\nTotal closed: {closed}")
|
||||
print(
|
||||
f"Total {'would warn (grace)' if dry_run else 'warned (grace)'}: "
|
||||
f"{summary['warn-grace']}"
|
||||
)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
282
.github/scripts/triage-requirements.txt
vendored
Normal file
282
.github/scripts/triage-requirements.txt
vendored
Normal file
|
|
@ -0,0 +1,282 @@
|
|||
# Hash-pinned dependency set for the Agent Shin triage scripts.
|
||||
# Installed in privileged triage workflows, so every package is pinned to an
|
||||
# exact version with SHA-256 hashes and installed with pip --require-hashes.
|
||||
#
|
||||
# Regenerate after bumping openai:
|
||||
# echo 'openai==<version>' \
|
||||
# | uv pip compile - --generate-hashes --python-version 3.12 \
|
||||
# --no-annotate --no-header -o .github/scripts/triage-requirements.txt
|
||||
|
||||
annotated-types==0.7.0 \
|
||||
--hash=sha256:1f02e8b43a8fbbc3f3e0d4f0f4bfc8131bcb4eebe8849b8e5c773f3a1c582a53 \
|
||||
--hash=sha256:aff07c09a53a08bc8cfccb9c85b05f1aa9a2a6f23728d790723543408344ce89
|
||||
anyio==4.14.0 \
|
||||
--hash=sha256:b47c1f9ccf73e67021df785332508f99379c68fa7d0684e8e3492cb1d4b23f89 \
|
||||
--hash=sha256:dd9b7a2a9799ed6552fde617b2c5df02b7fdd7d88392fc48101e51bae46164d9
|
||||
certifi==2026.6.17 \
|
||||
--hash=sha256:024c88eeec92ca068db80f02b8b07c9cef7b9fe261d1d535abfd5abd6f6af432 \
|
||||
--hash=sha256:2227dcbaafe0d2f59279d1762ddddc37783ed4354594f194ffc31d20f41fc3db
|
||||
distro==1.9.0 \
|
||||
--hash=sha256:2fa77c6fd8940f116ee1d6b94a2f90b13b5ea8d019b98bc8bafdcabcdd9bdbed \
|
||||
--hash=sha256:7bffd925d65168f85027d8da9af6bddab658135b840670a223589bc0c8ef02b2
|
||||
h11==0.16.0 \
|
||||
--hash=sha256:4e35b956cf45792e4caa5885e69fba00bdbc6ffafbfa020300e549b208ee5ff1 \
|
||||
--hash=sha256:63cf8bbe7522de3bf65932fda1d9c2772064ffb3dae62d55932da54b31cb6c86
|
||||
httpcore==1.0.9 \
|
||||
--hash=sha256:2d400746a40668fc9dec9810239072b40b4484b640a8c38fd654a024c7a1bf55 \
|
||||
--hash=sha256:6e34463af53fd2ab5d807f399a9b45ea31c3dfa2276f15a2c3f00afff6e176e8
|
||||
httpx==0.28.1 \
|
||||
--hash=sha256:75e98c5f16b0f35b567856f597f06ff2270a374470a5c2392242528e3e3e42fc \
|
||||
--hash=sha256:d909fcccc110f8c7faf814ca82a9a4d816bc5a6dbfea25d6591d6985b8ba59ad
|
||||
idna==3.18 \
|
||||
--hash=sha256:7f952cbe720b688055e3f87de14f5c3e5fdaa8bc3928985c4077ca689de849a2 \
|
||||
--hash=sha256:ffb385a7e039654cef1ab9ef32c6fafe283c0c0467bba1d9029738ce4a14a848
|
||||
jiter==0.15.0 \
|
||||
--hash=sha256:01a8222cf05ab1128e239421156c207949808acaaea2bdfd33130ae666786e86 \
|
||||
--hash=sha256:032396229564bca02440396bd327710719f724f5e7b7e9f7a8eb3faa4a2c2281 \
|
||||
--hash=sha256:04b400bbf8c9efb03d9bdd976475c919c1d85593b04b9fff7ae234065daf87ae \
|
||||
--hash=sha256:05906b93d72f03339e6bb7cf8dc10ebda64a0266126eed6beba79e20abcf5fd4 \
|
||||
--hash=sha256:066f8f33f18b2419cd8213b2436fa7fbc9c499f315971cfa3ce1f9820c001b1b \
|
||||
--hash=sha256:0ab068bce62a45aa3e7367eceaffb5dde60b7eb853be8dece45132e3d0ff4879 \
|
||||
--hash=sha256:0be6f5ad41a809f303f416d17cec92a7a725902fb9b4f3de3d19362ac0ef8554 \
|
||||
--hash=sha256:0e90a1c315a0226ec822d973817967f9223b7701546c8c2a7913e7ab0926294d \
|
||||
--hash=sha256:0f862193b8696249d22ec433e85fd2ab0ad9596bc3e45e6c0bc55e8aeba97be2 \
|
||||
--hash=sha256:1303d4d68a9b051ea90502402063ecf3807da00ad2affa19ca1ae3b90b3c5f67 \
|
||||
--hash=sha256:144f8e72cb53dab146347b91cceac01f5481237f2b93b4a339a1ee8f8878b67c \
|
||||
--hash=sha256:182226cbc930c9fab81bc2e41a4da672f89539906dadb05e75670ac07b94f71f \
|
||||
--hash=sha256:1c11465f97e2abf45a014b83b730222f8f1c5335e802c7055a67d50de6f1f4e3 \
|
||||
--hash=sha256:1c15024a3d892223b18f597c86d59387249dc396590844ce6b9f6131d1093bae \
|
||||
--hash=sha256:1d54fb5b31dea401a41af3f8a7d2512e9b6a6a005491e6166c7e4ffab9639a9c \
|
||||
--hash=sha256:25ffbe229aa8cd98c28879d8aa1a6e34ae77992ab984a65fba800859dab16269 \
|
||||
--hash=sha256:2a77aadd57cac1682e4401a72724d2796d89a4ba129b1a5812aa94ee480826eb \
|
||||
--hash=sha256:2ae901f3a55bfafdde31d289590fa25e3245735a2b1e8c7cc15871710a002871 \
|
||||
--hash=sha256:2b0074e2f56eb2dacca1689760fd2852a068f85a0547a157b82cb4cafeb6768b \
|
||||
--hash=sha256:2c8aea7781d2a372227871de4e1a1332aa96f5a89fd76c5e835dafdbad102887 \
|
||||
--hash=sha256:2c9cb907439d20bd0c7d7565ca01ee52234203208433749bae5b516907526928 \
|
||||
--hash=sha256:2fb6a5d26af81fc0f00f9360a891e05cf755e149bba391c4d563adc54812973d \
|
||||
--hash=sha256:2fd73e3da91a0a722d67165e849ce2cdc10de0e0d48738c142be8c6c5f310f4c \
|
||||
--hash=sha256:30ce1a5d16b5641dc935d50ef775af6a0871e3d14ab05d6fc54dff371b78e558 \
|
||||
--hash=sha256:30ce785d2adb8e32c3f7741442370a74834ec4c01f3c48f0750227a0b4ef27d6 \
|
||||
--hash=sha256:30f2218e6a9e5c18bc10fe6d41ac189c442c88eacf11bad9f28ef95a9bef00e6 \
|
||||
--hash=sha256:351a341c2105aa430b7047e30f1bf7975f6313b00165d3fc07be2edaf741f279 \
|
||||
--hash=sha256:37a10c377ce3a4a85f4a67f28b7afe093154cde77eaf248a72e856aa08b4d865 \
|
||||
--hash=sha256:392b8ab019e5502d08aff85c6272209c24bc2cbe706ea82a56368f524236614a \
|
||||
--hash=sha256:3e4540b8e74e4268811ac05db226a6a128ff572e7e0ce3f1163b693cadb184cd \
|
||||
--hash=sha256:40b2c7e92c44a84d748d21706c68dc6ff8161d80b59c99d774721a0d2317d7c7 \
|
||||
--hash=sha256:411fa4dfa5a7ae3d11491027ffb9beadec3996010a986862db70d91abba1c750 \
|
||||
--hash=sha256:4251acc80e2b7c9b7b8823456ea0fceeb0734dac2df7636d3c711b38476b5a76 \
|
||||
--hash=sha256:42bfb257930800cf43e7c62c832402c704ab60797c992faf88d20e903eac8f32 \
|
||||
--hash=sha256:4363818355dbc70ae1a8e9eaba9de350d93ede4ff6992b8f8eb8cbb6e5122d42 \
|
||||
--hash=sha256:4ab395feec8d249ec4044e228e98a7033f043426a265df439dc3698823f0a4e4 \
|
||||
--hash=sha256:50164d7610c00e7cd913a873fce30b6beeebf4b37e53983e33f22de4c900f6b8 \
|
||||
--hash=sha256:50e51156192722a9c58db112837d3f8ef96fb3c5ecc14e95f409134b08b158ec \
|
||||
--hash=sha256:510c8b3c17a0ed9ac69850c0438dada3c9b82d9c4d589fcb62002a5a9cf3a866 \
|
||||
--hash=sha256:5157de9f76eb4bc5ea74a1219366a25f945ad305641d74e04f59c54087091aa9 \
|
||||
--hash=sha256:54d5d6090cdc1b7c9e780dfb04949a990adb1e301a2fc0bbcee7de4638d33f9a \
|
||||
--hash=sha256:553fcac2ef2cb990877f9fc0833b8b629a3e6a5670b6b5fd58219b41a653ddc4 \
|
||||
--hash=sha256:5607e6013ed7e6b0ec9661e467b7ffde0aa7ab36833a04850f26fcf88ed4845b \
|
||||
--hash=sha256:5d6a60072b44c3c2b797a7ddcbcbbf2b34ea3cfd4721580fbfd2a09d9d9b84ba \
|
||||
--hash=sha256:5f30bae8bc1c2d613e28e5af3e8cceb09b742f1c8a8a5f839fb67afaffc03b61 \
|
||||
--hash=sha256:62ebd14e47e9aed9df4472afcb2663668ce4d74891cd54f86bf6e44029d6dc89 \
|
||||
--hash=sha256:631f13a3d04e97d4e083993b10f4b99530e3a10d953e2eb5e196b7dc7f812ce0 \
|
||||
--hash=sha256:6550fa135c7deb8ead6af49ed7ff648532ea8334a1447fe34a36315ef79c5c29 \
|
||||
--hash=sha256:66b1880df2d01e206e8339769d1c7c1753bcb653efd6289e203f6f24ebada0c0 \
|
||||
--hash=sha256:6eac374c5c975709b69c10f09afd199df74150172156ad10c8d4fd785b7da995 \
|
||||
--hash=sha256:71683c38c825452999b5717fcae07ea708e8c93003e808be4319c1b02e3d176e \
|
||||
--hash=sha256:7553333dd0930c104a5a0db8df72bf7219fe663d731383b576bb6ed6351c984d \
|
||||
--hash=sha256:75e8a04e91432dde9f1838373cf93d23726c79d3e908d319acf0e796f85592e7 \
|
||||
--hash=sha256:773b6eb282ce11ee19f05f6b2d4404fa308e5bbd353b0b80a0262caad6db2cd7 \
|
||||
--hash=sha256:774f93f65031856bf14ad9f59bdcab8b8cad501e5ceabd51ba3525f76937a25b \
|
||||
--hash=sha256:7c468136b8bd6bb18c8786e4236a1fa27362f24cb23450ba0cb204ab379b8e6f \
|
||||
--hash=sha256:7ce8902f939970048b233087082e7bb829db29375811c7ad50687b8624c6fd08 \
|
||||
--hash=sha256:7d3d6683288c11cbab50e865f2e2f13950179aa45410e30b2cfbd3fb7b0177bf \
|
||||
--hash=sha256:7f6163c0f10b055245f814dcc59f4818da60dfe72f3e72ab89fc24b6bd5e9c52 \
|
||||
--hash=sha256:8020c99ec13a7db2b6f96cbe82ef4721c88b426a4892f27478044af0284615ef \
|
||||
--hash=sha256:813dfbb17d65328bf86e5f0905dd277ba2265d3ca20556e86c0c7035b7182e5a \
|
||||
--hash=sha256:860a74063284a2ae9bfedd694f299cc2c68e2696c5f3d440cc9d18bb81b9dd04 \
|
||||
--hash=sha256:8c9004af7c8d67cce7f1aae1026fb55607f4aa600710d08ede3a3ce4aeefe7e0 \
|
||||
--hash=sha256:8d2c0c44d569ce0f2850f5c926f8caeb5f245fbc84475aeb36efccc2103e6dbd \
|
||||
--hash=sha256:8f7e9bc0f1135039b22ee6eab588d42df1ce55842b30740a352885eb267bd941 \
|
||||
--hash=sha256:90c5db5527c221249a876160663ab891ace358c17f7b9c93ec1478b7f0550e5c \
|
||||
--hash=sha256:9100ddbec09741cc66feb0fc6773f8bdbd0e3c345689368f260082ff85dcc0cd \
|
||||
--hash=sha256:913d02d29c9606643418d9ccfc3b72492ab25a6bf7889934e09a3490f8d3438b \
|
||||
--hash=sha256:980c256edb05b78a111b99c4de3b1d32e31634b867fd1fc2cf726e7b7bba9854 \
|
||||
--hash=sha256:9f924585cdacf631cd382b657966847bb537bf9ed0a6f9b991da5f05a631480f \
|
||||
--hash=sha256:a254e10b593624d230c365b6d616b22ca0ad65e63a16e6631c2b3466022e6ba8 \
|
||||
--hash=sha256:a2a438005b6f22d0273413484d6094d7c2c5d10ec1b3a3bf128e0d1d3ba53258 \
|
||||
--hash=sha256:a97261f1fccb8e50ecd2890a96e46efdc3f57c80a197324c6777827231eca712 \
|
||||
--hash=sha256:ab596fa3837e91e7e6a31b5f639988bfc6a35d1f915ac3932d946062219d588f \
|
||||
--hash=sha256:abbf258599526ad0326fe51e252e24f2bd6f24f1852681b4b78feda3808f1d18 \
|
||||
--hash=sha256:ac0d9ddea4350974be7a221fc25895f251a8fee748c889bdced2141c0fec1a49 \
|
||||
--hash=sha256:acf4ee4d1fc55917239fe72972fb292dd773055d05eb040d36f4326e02cc2c0e \
|
||||
--hash=sha256:ae1b0d82ac2d987f9ea512b1c9adfcc71a28de3dea3a6039b54d76cffda9901e \
|
||||
--hash=sha256:b15741f501469009ae0ae90b7147958a664a7dede40aa7ff174a8a4645f546d0 \
|
||||
--hash=sha256:b15d3ec9b0449c40e85319bdb4caa8b77ab526e74f5532ed94bec15e2f66822c \
|
||||
--hash=sha256:b3b3b775e33d3bfaec9899edc526ae97b0da0bf9d071a46124ba419149a414f8 \
|
||||
--hash=sha256:b6c0ffae686c39bf3737be60793783267628783ea42545632c10b291105aee45 \
|
||||
--hash=sha256:c210f8b35dc6f30aafd4b4365ca89b9d1189f21ab49b8e68fa6322a847aef138 \
|
||||
--hash=sha256:c2f6bb8b5216ab9e7873bc08b5d7bef2b8abbb578a3069bf1cd14a45d71d771d \
|
||||
--hash=sha256:c60e71b6d10cfc284c9bf36bd885e8d44c46f688ce50aa91b5edd90181dea687 \
|
||||
--hash=sha256:c6694a173ecabc12eb60efbc0b474464ead1951ff65cd8b1e72100715c64512b \
|
||||
--hash=sha256:c77496cb10bd7549690fbbab3e5ec05857b83e49276f4a9423a766ddd2afcd4c \
|
||||
--hash=sha256:c84c1b7be454b0c16f8499b4ebfbfd82ea5cca6527cceefcbbc06a7557b5ed2e \
|
||||
--hash=sha256:cc0bc345cf2df9d1c00ac443f50d543c1ccfa8b0422cb85b1ab70d681c0b255b \
|
||||
--hash=sha256:ceb8fc27d38793f9c97149be8302720c5b22e5c195a37bf2c45dc36c4600a512 \
|
||||
--hash=sha256:cf4bd113a69c0a740e27cb962ce10630c36d2b8f59d759a651b955ee9d18a823 \
|
||||
--hash=sha256:d1aa62e277fc1cbd80e6deacae6f4d983b41b3d7728e0645c5d741a6149bba45 \
|
||||
--hash=sha256:d1e7b1776f0797956c509e123d0952d10d293a9492dea9f288ab9570ec01d1a5 \
|
||||
--hash=sha256:d636d5095155afd364247f65070fab7beda13498d7ff4de331046e704ab9657f \
|
||||
--hash=sha256:d726e3ceeb337191324b49de298142f27c3ad10886341555d1d5315b5f252c6a \
|
||||
--hash=sha256:d72d8af5c1013656a8870c866660627d1a75bc185814ee022c8533caa1de88ae \
|
||||
--hash=sha256:d8d2955167274e15d79a7a020afdd9b39c990eb80b2d89fca695d92dcfdd38ec \
|
||||
--hash=sha256:d92a5cd21fdb083931d546c207aa29633787c5dc5b02daab2d32b843f88a2c53 \
|
||||
--hash=sha256:e58585a58209d72691ce2d62a9147445f5a87beb0bde97fde284c96ae392a3d1 \
|
||||
--hash=sha256:e7196e56f1cd69af1dbb07dff02dcfb260a50b45a82d409d92a06fedb32473b5 \
|
||||
--hash=sha256:eda3071db3346334beae1360b46da4606da57bf3528c167b3c38533afaf9f2c5 \
|
||||
--hash=sha256:edebcf7d1f601199084bb6e844d7dc67e03e04f6ac786b0332d616635c4ff7a4 \
|
||||
--hash=sha256:ef1fd24d9413f6209e00d3d5a453e67acfe004a25cc6c8e8484faed4311ab9e8 \
|
||||
--hash=sha256:f0b271b462769543716f92d3a4f90527df6ef5ed05ee95ec4137f513e21e1b77 \
|
||||
--hash=sha256:f18f85e4218d1b40f000f42a92239a7a61a902cd42c65e6c360dbd17dcb20894 \
|
||||
--hash=sha256:f1e1754960f38ec40613a07e5e372df67acb3b890fb383b6fb3de3e49ddbf3c7 \
|
||||
--hash=sha256:f2143ab06181d2b029eedcb6af3cebe95f11bbac62441781860f98ee9330a6a6 \
|
||||
--hash=sha256:f3d37768fce7f88dd2a8c6091f2325dea27d30d30d5c6e7a1c0f0af77723b708 \
|
||||
--hash=sha256:fa248c9eb220197d363f688818dac2fd4b2f0cd7d843ca7105d652034823427d
|
||||
openai==2.33.0 \
|
||||
--hash=sha256:03ac37d70e8c9e3a8124214e3afa785e2cbc12e627fbd98177a086ef2fd87ad5 \
|
||||
--hash=sha256:f850c435e2a4685bba3295bd54912dd26315d9c1b7733068186134d6e0599f9a
|
||||
pydantic==2.13.4 \
|
||||
--hash=sha256:45a282cde31d808236fd7ea9d919b128653c8b38b393d1c4ab335c62924d9aba \
|
||||
--hash=sha256:c40756b57adaa8b1efeeced5c196f3f3b7c435f90e84ea7f443901bec8099ef6
|
||||
pydantic-core==2.46.4 \
|
||||
--hash=sha256:00c603d540afdd6b80eb39f078f33ebd46211f02f33e34a32d9f053bba711de0 \
|
||||
--hash=sha256:0186750b482eefa11d7f435892b09c5c606193ef3375bcf94aa00ae6bfb66262 \
|
||||
--hash=sha256:041bde0a48fd37cf71cab1c9d56d3e8625a3793fef1f7dd232b3ff37e978ecda \
|
||||
--hash=sha256:0c563b08bca408dc7f65f700633d8442fffb2421fc47b8101377e9fd65051ff0 \
|
||||
--hash=sha256:0cbe8b01f948de4286c74cdd6c667aceb38f5c1e26f0693b3983d9d74887c65e \
|
||||
--hash=sha256:0ce40cd7b21210e99342afafbd4d0f76d784eb5b1d60f3bdc566be4983c6c73b \
|
||||
--hash=sha256:0e96592440881c74a213e5ad528e2b24d3d4f940de2766bed9010ab1d9e51594 \
|
||||
--hash=sha256:10e17cbb10a330363733efc4d7c4d0dd827ac0909b8f6a6542298fed1ea62f29 \
|
||||
--hash=sha256:133878133d271ade3d41d1bfb2a45ec38dbdbda40bc065921c6b04e4630127e2 \
|
||||
--hash=sha256:14d4edf427bdcf950a8a02d7cb44a08614388dd6e1bdcbf4f67504fa7887da9c \
|
||||
--hash=sha256:14f4c5d6db102bd796a627bbb3a17b4cf4574b9ae861d8b7c9a9661c6dd3362d \
|
||||
--hash=sha256:17299feefe090f2caa5b8e37222bb5f663e4935a8bfa6931d4102e5df1a9f398 \
|
||||
--hash=sha256:184c081504d17f1c1066e430e117142b2c77d9448a97f7b65c6ac9fd9aee238d \
|
||||
--hash=sha256:18e5ceec2ab67e6d5f1a9085e5a24c9c4e2ac4545730bfe668680bca05e555f3 \
|
||||
--hash=sha256:19e51f073cd3df251856a8a4189fbdf1de4012c3ebacfb1884f94f1eb406079f \
|
||||
--hash=sha256:1a7dd0b3ee80d90150e3495a3a13ac34dbcbfd4f012996a6a1d8900e91b5c0fb \
|
||||
--hash=sha256:1d8ba486450b14f3b1d63bc521d410ec7565e52f887b9fb671791886436a42f7 \
|
||||
--hash=sha256:2108ba5c1c1eca18030634489dc544844144ee36357f2f9f780b93e7ddbb44b5 \
|
||||
--hash=sha256:228ee9bae8bef5b1e97ec58302f80357c37199e0d0a99174e138d28e6957b9d9 \
|
||||
--hash=sha256:23ace664830ee0bfe014a0c7bc248b1f7f25ed7ad103852c317624a1083af462 \
|
||||
--hash=sha256:2412e734dcb48da14d4e4006b82b46b74f2518b8a26ee7e58c6844a6cd6d03c4 \
|
||||
--hash=sha256:29c61fc04a3d840155ff08e475a04809278972fe6aef51e2720554e96367e34b \
|
||||
--hash=sha256:2f84c03c8607173d16b5a854ec68a2f9079ae03237a54fb506d13af47e1d018d \
|
||||
--hash=sha256:3009f12e4e90b7f88b4f9adb1b0c4a3d58fe7820f3238c190047209d148026df \
|
||||
--hash=sha256:3245406455a5d98187ec35530fd772b1d799b26667980872c8d4614991e2c4a2 \
|
||||
--hash=sha256:3447661d99f75a3683a4cf5c87da72f2161964611864dbbeac7fbb118bb4bfc0 \
|
||||
--hash=sha256:372429a130e469c9cd698925ce5fc50940b7a1336b0d82038e63d5bbc4edc519 \
|
||||
--hash=sha256:395aebd9183f9d112f569aeb5b2214d1a10a33bec8456447f7fbdfa51d38d4cd \
|
||||
--hash=sha256:3a233125ac121aa3ffba9a2b59edfc4a985a76092dc8279586ab4b71390875e7 \
|
||||
--hash=sha256:3be77f45df024d789a672ae34f8b06fb346c4f9f46ea714956660ea4862e89ac \
|
||||
--hash=sha256:3bf92c5d0e00fefaab325a4d27828fe6b6e2a21848686b5b60d2d9eeb09d76c6 \
|
||||
--hash=sha256:3ecbc122d18468d06ca279dc26a8c2e2d5acb10943bb35e36ae92096dc3b5565 \
|
||||
--hash=sha256:3fb702cd90b0446a3a1c5e470bfa0dd23c0233b676a9099ddcc964fa6ca13898 \
|
||||
--hash=sha256:428e04521a40150c85216fc8b85e8d39fece235a9cf5e383761238c7fa9b96fb \
|
||||
--hash=sha256:432c179df7874eeb73307aad2df0755e1ae0efa61ff0ea89b93e194411ae3928 \
|
||||
--hash=sha256:4a05d69cba51d852c5c3e92758653245a50c0b646ced0cf05bd793ed592839d6 \
|
||||
--hash=sha256:4c63ebc82684aa89d9a3bcbd13d515b3be44250dc68dd3bd81526c1cb31286c3 \
|
||||
--hash=sha256:4fc73cb559bdb54b1134a706a2802a4cddd27a0633f5abb7e53056268751ac6a \
|
||||
--hash=sha256:4fcbe087dbc2068af7eda3aa87634eba216dbda64d1ae73c8684b621d33f6596 \
|
||||
--hash=sha256:56cb4851bcaf3d117eddcef4fe66afd750a50274b0da8e22be256d10e5611987 \
|
||||
--hash=sha256:5855698a4856556d86e8e6cd8434bc3ac0314ee8e12089ae0e143f64c6256e4e \
|
||||
--hash=sha256:5a4330cdbc57162e4b3aa303f588ba752257694c9c9be3e7ebb11b4aca659b5d \
|
||||
--hash=sha256:5b712b53160b79a5850310b912a5ef8e57e56947c8ad690c227f5c9d7e561712 \
|
||||
--hash=sha256:5d5902252db0d3cedf8d4a1bc68f70eeb430f7e4c7104c8c476753519b423008 \
|
||||
--hash=sha256:617d7e2ca7dcb8c5cf6bcb8c59b8832c94b36196bbf1cbd1bfb56ed341905edd \
|
||||
--hash=sha256:62f875393d7f270851f20523dd2e29f082bcc82292d66db2b64ea71f64b6e1c1 \
|
||||
--hash=sha256:633147d34cf4550417f12e2b1a0383973bdf5cdfde212cb09e9a581cf10820be \
|
||||
--hash=sha256:66ce7632c22d837c95301830e111ad0128a32b8207533b60896a96c4915192ea \
|
||||
--hash=sha256:6b3ace8194b0e5204818c92802dcdca7fc6d88aabbb799d7c795540d9cd6d292 \
|
||||
--hash=sha256:6f2eeda33a839975441c86a4119e1383c50b47faf0cbb5176985565c6bb02c33 \
|
||||
--hash=sha256:7027560ee92211647d0d34e3f7cd6f50da56399d26a9c8ad0da286d3869a53f3 \
|
||||
--hash=sha256:7283d57845ecf5a163403eb0702dfc220cc4fbdd18919cb5ccea4f95ee1cdab4 \
|
||||
--hash=sha256:7a5f930472650a82629163023e630d160863fce524c616f4e5186e5de9d9a49b \
|
||||
--hash=sha256:7bfb192b3f4b9e8a89b6277b6ce787564f62cfd272055f6e685726b111dc7826 \
|
||||
--hash=sha256:811ff8e9c313ab425368bcbb36e5c4ebd7108c2bbf4e4089cfbb0b01eff63fac \
|
||||
--hash=sha256:8233f2947cf85404441fd7e0085f53b10c93e0ee78611099b5c7237e36aacbf7 \
|
||||
--hash=sha256:82cf5301172168103724d49a1444d3378cb20cdee30b116a1bd6031236298a5d \
|
||||
--hash=sha256:8358a950c8909158e3df31538a7e4edc2d7265a7c54b47f0864d9e5bae9dcebf \
|
||||
--hash=sha256:85bb3611ff1802f3ee7fdd7dbff26b56f343fb432d57a4728fdd49b6ef35e2f4 \
|
||||
--hash=sha256:86e1a4418c6cd97d60c95c71164158eaf7324fae7b0923264016baa993eba6fc \
|
||||
--hash=sha256:8b9bab013d1c7a79d3501ff86d0bc9c31bf587db4551677b96bec07df78c6b15 \
|
||||
--hash=sha256:8c5dac79fa1614d1e06ca695109c6105923bd9c7d1d6c918d4e637b7e6b32fd3 \
|
||||
--hash=sha256:8d0820e8192167f80d88d64038e609c31452eeca865b4e1d9950a27a4609b00b \
|
||||
--hash=sha256:8daafc69c93ee8a0204506a3b6b30f586ef54028f52aeeeb5c4cfc5184fd5914 \
|
||||
--hash=sha256:9037063db01f09b09e237c282b6792bd4da634b5402c4e7f0c61effed7701a04 \
|
||||
--hash=sha256:905a0ed8ea6f2d61c1738835f99b699348d7857379083e5fc497fa0c967a407c \
|
||||
--hash=sha256:90884113d8b48f760e9587002789ddd741e76ab9f89518cd1e43b1f1a52ec44b \
|
||||
--hash=sha256:91a06d2e259ecfbd8c901d70c3c507900458498142b3026a296b7de4d1322cc9 \
|
||||
--hash=sha256:926c9541b14b12b1681dca8a0b75feb510b06c6341b70a8e500c2fdcff837cce \
|
||||
--hash=sha256:9401557acd873c3a7f3eb9383edef8ac4968f9510e340f4808d427e75667e7b4 \
|
||||
--hash=sha256:9551187363ffc0de2a00b2e47c25aeaeb1020b69b668762966df15fc5659dd5a \
|
||||
--hash=sha256:962ccbab7b642487b1d8b7df90ef677e03134cf1fd8880bf698649b22a69371f \
|
||||
--hash=sha256:97e7cf2be5c77b7d1a9713a05605d49460d02c6078d38d8bef3cbe323c548424 \
|
||||
--hash=sha256:9aa768456404a8bf48a4406685ac2bec8e72b62c69313734fa3b73cf33b3a894 \
|
||||
--hash=sha256:9bc519fbf2b7578398853d815009ae5e4d4603d12f4e3f91da8c06852d3da3e9 \
|
||||
--hash=sha256:9d56801be94b86a9da183e5f3766e6310752b99ff647e38b09a9500d88e46e76 \
|
||||
--hash=sha256:9f444c499b3eefd3a92e348059471ea0c3a6e303d9c1cec09fa748fd9f895201 \
|
||||
--hash=sha256:9fa8ae11da9e2b3126c6426f147e0fba88d96d65921799bb30c6abd1cb2c97fb \
|
||||
--hash=sha256:a0f62d0a58f4e7da165457e995725421e0064f2255d8eccebc49f41bbc23b109 \
|
||||
--hash=sha256:a396dcc17e5a0b164dbe026896245a4fa9ff402edca1dff0be3d53a517f74de4 \
|
||||
--hash=sha256:aaa2a54443eff1950ba5ddc6b6ccda0d9c84a364276a62f969bdf2a390650848 \
|
||||
--hash=sha256:ad785e92e6dc634c21555edc8bd6b64957ab844541bcb96a1366c202951ae526 \
|
||||
--hash=sha256:af8244b2bef6aaad6d92cda81372de7f8c8d36c9f0c3ea36e827c60e7d9467a0 \
|
||||
--hash=sha256:b078afbc25f3a1436c7a1d2cd3e322497ee99615ba97c563566fdf46aff1ee01 \
|
||||
--hash=sha256:b2f69dec1725e79a012d920df1707de5caf7ed5e08f3be4435e25803efc47458 \
|
||||
--hash=sha256:b8458003118a712e66286df6a707db01c52c0f52f7db8e4a38f0da1d3b94fc4e \
|
||||
--hash=sha256:bb63e0198ca18aad131c089b9204c23079c3afa95487e561f4c522d519e55aba \
|
||||
--hash=sha256:bfec22eab3c8cc2ceec0248aec886624116dc079afa027ecc8ad4a7e62010f8a \
|
||||
--hash=sha256:c1747f85cee84c26985853c6f3d9bd3e75da5212912443fa111c113b9c246f39 \
|
||||
--hash=sha256:c1b3f518abeca3aa13c712fd202306e145abf59a18b094a6bafb2d2bbf59192c \
|
||||
--hash=sha256:c50f2528cf200c5eed56faf3f4e22fcd5f38c157a8b78576e6ba3168ec35f000 \
|
||||
--hash=sha256:c68fcd102d71ea85c5b2dfac3f4f8476eff42a9e078fd5faefff6d145063536b \
|
||||
--hash=sha256:c7a7bd4e39e8e4c12c39cd480356842b6a8a06e41b23a55a5e3e191718838ddf \
|
||||
--hash=sha256:c94f0688e7b8d0a67abf40e57a7eaaecd17cc9586706a31b76c031f63df052b4 \
|
||||
--hash=sha256:cbaf13819775b7f769bf4a1f066cb6df7a28d4480081a589828ef190226881cd \
|
||||
--hash=sha256:cd2213145bcc2ba85884d0ac63d222fece9209678f77b9b4d76f054c561adb28 \
|
||||
--hash=sha256:ce5c1d2a8b27468f433ca974829c44060b8097eedc39933e3c206a90ee49c4a9 \
|
||||
--hash=sha256:d396ec2b979760aaf3218e76c24e65bd0aca24983298653b3a9d7a45f9e47b30 \
|
||||
--hash=sha256:d51026d73fcfd93610abc7b27789c26b313920fcfb20e27462d74a7f8b06e983 \
|
||||
--hash=sha256:d80ee3d731373b24cebbc10d689ca4ee1875caf0d5703a245db18efd4dd37fc1 \
|
||||
--hash=sha256:d995260fdf4e1db774581b4900e0f832abe3c7c84996726bbc161b19c8f29e76 \
|
||||
--hash=sha256:da4b951fe36dc7c3a1ccb4e3cd1747c3542b8c9ceede8fc86cae054e764485f5 \
|
||||
--hash=sha256:daa27d92c36f24388fe3ad306b174781c747627f134452e4f128ea00ce1fe8c4 \
|
||||
--hash=sha256:db06ffe51636ffe9ca531fe9023dd64bdd794be8754cb5df57c5498ae5b518a7 \
|
||||
--hash=sha256:e0d65b8c354be7fb5f720c3caa8bc940bc2d20ce749c8e06135f07f8ed95dd7c \
|
||||
--hash=sha256:e68b7a074f65a2fd746c52a7ce6142ab7006074ac269ace0c25cd8ba171f8066 \
|
||||
--hash=sha256:e739fee756ba1010f8bcccb534252e85a35fe45ae92c295a06059ce58b74ccd3 \
|
||||
--hash=sha256:e846ae7835bf0703ae43f534ab79a867146dadd59dc9ca5c8b53d5c8f7c9ef02 \
|
||||
--hash=sha256:e9c26f834c65f5752f3f06cb08cb86a913ceb7274d0db6e267808a708b46bc89 \
|
||||
--hash=sha256:ea793e075b70290d89d8142074262885d3f7da19634845135751bd6344f73b50 \
|
||||
--hash=sha256:f027324c56cd5406ca49c124b0db10e56c69064fec039acc571c29020cc87c76 \
|
||||
--hash=sha256:f13a646d65d09fbf1bc6b3a9635d30095c8e7e5cc419ff35ecc563c5fd04cd49 \
|
||||
--hash=sha256:f47286a97f0bc9b8859519809077b91b2cefe4ae47fcbf5e466a009c1c5d742b \
|
||||
--hash=sha256:f747929cf940cddb5b3668a390056ddd5ba2e5010615ea2dcf4f9c4f3ab8791d \
|
||||
--hash=sha256:f99626688942fb746e545232e7726926f3be91b5975f8b55327665fafda991c7 \
|
||||
--hash=sha256:f9fa868638bf362d3d138ea55829cefb3d5f4b0d7f142234382a15e2485dbec4 \
|
||||
--hash=sha256:fbdb89b3e1c94a30cc5edfce477c6e6a5dc4d8f84665b455c27582f211a1c72c \
|
||||
--hash=sha256:fc010ab034c8c7452522748bf937df58020d256ccae0874463d1f4d01758af8e \
|
||||
--hash=sha256:fc3e9034a63de20e15e8ade85358bc6efc614008cab72898b4b4952bea0509ff \
|
||||
--hash=sha256:fd8b3d9fd264be37976686c7f65cd52a83f5e84f4bfd2adf9c1d469676bbb6ae
|
||||
sniffio==1.3.1 \
|
||||
--hash=sha256:2f6da418d1f1e0fddd844478f41680e794e6051915791a034ff65e5f100525a2 \
|
||||
--hash=sha256:f4324edc670a0f49750a81b895f35c3adb843cca46f0530f79fc1babb23789dc
|
||||
tqdm==4.68.3 \
|
||||
--hash=sha256:00dfa48452b6b6cfae3dd9885636c23d3422d1ec97c66d96818cbd5e0821d482 \
|
||||
--hash=sha256:39832cc2def2789a6f29df83f172db7416cea70052c0907a57801c5f2fdccb03
|
||||
typing-extensions==4.15.0 \
|
||||
--hash=sha256:0cea48d173cc12fa28ecabc3b837ea3cf6f38c6d1136f85cbaaf598984861466 \
|
||||
--hash=sha256:f0fa19c6845758ab08074a0cfa8b7aecb71c999ca73d62883bc25cc018c4e548
|
||||
typing-inspection==0.4.2 \
|
||||
--hash=sha256:4ed1cacbdc298c220f1bd249ed5287caa16f34d44ef4e9c3d0cbad5b521545e7 \
|
||||
--hash=sha256:ba561c48a67c5958007083d386c3295464928b01faa735ab8547c5692e87f464
|
||||
557
.github/scripts/triage_rollout_heads_up.py
vendored
Normal file
557
.github/scripts/triage_rollout_heads_up.py
vendored
Normal file
|
|
@ -0,0 +1,557 @@
|
|||
#!/usr/bin/env python3
|
||||
"""One-shot 7-day heads-up sweep for the Agent Shin rollout.
|
||||
|
||||
Posts a friendly "the OSS triage bot kicks in next Monday" comment on every
|
||||
open external PR/issue that currently *would* fail the new rubric — i.e.,
|
||||
every PR/issue Agent Shin would close once the rollout completes. The point
|
||||
is to give contributors a full week to fix their description before the bot
|
||||
ever takes a destructive action, so nobody is surprised by an auto-close.
|
||||
|
||||
The script is designed to run **exactly once** at rollout, fired by a manual
|
||||
``workflow_dispatch`` (``dry_run=false``) on the heads-up workflow. Re-runs
|
||||
are safe: every comment is stamped with the hidden ``HEADS_UP_MARKER`` and
|
||||
PRs/issues that already carry the marker are skipped.
|
||||
|
||||
Dry-run vs. real run
|
||||
--------------------
|
||||
Defaults to dry-run. Passing ``--close`` flips into real mode. Every GitHub
|
||||
mutation goes through ``_agent_shin_actions``, which has a one-line
|
||||
``if dry_run: log else: do_it`` per call, so the only difference between a
|
||||
dry-run preview and the real run is the call site that actually hits the
|
||||
GitHub API.
|
||||
|
||||
Local preview::
|
||||
|
||||
python3 .github/scripts/triage_rollout_heads_up.py --repo BerriAI/litellm
|
||||
|
||||
Real run (the manual rollout dispatch uses this)::
|
||||
|
||||
python3 .github/scripts/triage_rollout_heads_up.py --repo BerriAI/litellm --close
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import datetime as dt
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
# Make the sibling triage_with_llm + _agent_shin_actions importable when this
|
||||
# script is invoked directly (the GitHub workflow does `python3 .github/scripts/...`).
|
||||
_SCRIPTS_DIR = Path(__file__).resolve().parent
|
||||
if str(_SCRIPTS_DIR) not in sys.path:
|
||||
sys.path.insert(0, str(_SCRIPTS_DIR))
|
||||
|
||||
from _agent_shin_actions import maybe_post_comment # noqa: E402
|
||||
from agent_shin_shared import ( # noqa: E402
|
||||
AGENT_SHIN_DEFAULT_BOT_LOGIN,
|
||||
ALLOWLIST_LOGINS,
|
||||
list_open_items,
|
||||
)
|
||||
from triage_with_llm import ( # noqa: E402
|
||||
DEFAULT_MODEL,
|
||||
call_llm_judge,
|
||||
fetch_issue,
|
||||
fetch_pr,
|
||||
gh,
|
||||
is_internal_contributor,
|
||||
review_gate,
|
||||
triage,
|
||||
)
|
||||
|
||||
# Hidden marker so re-runs skip PRs/issues we've already notified. Distinct from
|
||||
# the within-grace / ready / regressed markers so it can't be confused with the
|
||||
# steady-state lifecycle comments.
|
||||
HEADS_UP_MARKER = "<!-- agent-shin:rollout-heads-up -->"
|
||||
|
||||
# Placeholder until the litellm-docs PR ships. The rollout blog post explains
|
||||
# the new rubric, the 7-day grace, and how to recover after an auto-close.
|
||||
# TODO(docs): replace with the canonical URL once the litellm-docs PR merges.
|
||||
ROLLOUT_BLOG_URL = "https://docs.litellm.ai/docs/agent_shin_triage_rollout"
|
||||
|
||||
# Default cutoff is one week from "now". Computed at runtime so the wording
|
||||
# stays correct even if the rollout is merged later than planned. The user can
|
||||
# override with --close-on YYYY-MM-DD when running the script manually.
|
||||
DEFAULT_GRACE_DAYS = 7
|
||||
|
||||
# The daily auto-close sweeps (close_low_quality_prs.yml at 09:00 UTC and
|
||||
# review_gate.yml at 09:30 UTC) are what actually close a still-failing item,
|
||||
# so the deadline we promise contributors has to name that wall-clock moment.
|
||||
ACTIVATION_TIME_UTC = "09:00 UTC"
|
||||
|
||||
|
||||
def _format_cutoff(cutoff: dt.date) -> str:
|
||||
"""Human-readable, timezone-explicit cutoff, e.g. ``Monday, June 1, 2026
|
||||
(09:00 UTC)`` — the moment a still-failing PR/issue gets closed."""
|
||||
return (
|
||||
f"{cutoff.strftime('%A, %B')} {cutoff.day}, {cutoff.year} "
|
||||
f"({ACTIVATION_TIME_UTC})"
|
||||
)
|
||||
|
||||
|
||||
def _rubric_section_pr() -> str:
|
||||
return (
|
||||
"**Going forward, every external PR needs ONE of:**\n"
|
||||
"\n"
|
||||
"- A linked GitHub issue using a closing keyword: "
|
||||
"`Fixes #1234`, `Closes #1234`, or `Resolves #1234`, OR\n"
|
||||
"- All three of: a clear **problem description**, **expected vs. "
|
||||
"actual behavior**, and **end-to-end QA proof** (at least one of a "
|
||||
"short screen recording / video, before/after screenshots, or the "
|
||||
"exact commands you ran with their real output; mocked or stubbed "
|
||||
"runs don't count).\n"
|
||||
"\n"
|
||||
"PRs also need a **Greptile confidence score of 4/5 or higher** before "
|
||||
"the bot will tag them `ready for review`. You can `@greptileai` to "
|
||||
"request a fresh review at any time, including after the PR is closed."
|
||||
)
|
||||
|
||||
|
||||
def _rubric_section_issue() -> str:
|
||||
return (
|
||||
"**Going forward, every external issue needs:**\n"
|
||||
"\n"
|
||||
"- For **bug reports**: end-to-end evidence of the bug (at least one "
|
||||
"of a screen recording / video, a screenshot, or the exact commands "
|
||||
"you ran with their real output / traceback) plus expected vs. actual "
|
||||
"behavior. Written steps with no run output don't count, and mocked "
|
||||
"or stubbed runs don't count.\n"
|
||||
"- For **feature requests**: a clear description of the proposed "
|
||||
"feature plus a use case + concrete example (config, API call, UI "
|
||||
"flow, or scenario showing what's blocked today)."
|
||||
)
|
||||
|
||||
|
||||
def _description_only_note(kind: str) -> str:
|
||||
noun = "PR" if kind == "pr" else "issue"
|
||||
return (
|
||||
f"⚠️ **The requirements must live in the {noun} *description*, not in "
|
||||
"comments.** Some PRs/issues collect 100+ comments from humans and "
|
||||
"bots; reading the entire thread on every triage run would balloon "
|
||||
"GitHub API usage (we'd start getting 429'd) and blow out the LLM "
|
||||
"judge's context. The bot only reads the description, so anything "
|
||||
"you add as a comment will be invisible to it."
|
||||
)
|
||||
|
||||
|
||||
def _missing_section(verdict: dict, greptile_score: int | None) -> str:
|
||||
"""Bullet list of what's currently missing on this PR/issue.
|
||||
|
||||
Combines the LLM judge's `missing` list (rubric items) with a Greptile
|
||||
shortfall (for PRs) so the contributor sees one list of things to fix.
|
||||
"""
|
||||
missing = list(verdict.get("missing") or [])
|
||||
if greptile_score is not None and greptile_score < 4:
|
||||
missing.insert(
|
||||
0,
|
||||
f"Greptile's most recent review scored this PR {greptile_score}/5 "
|
||||
"(below the 4/5 bar Agent Shin will require).",
|
||||
)
|
||||
if not missing:
|
||||
return (
|
||||
"_The bot couldn't articulate a specific missing piece; see the "
|
||||
"rubric link above and double-check the description includes all "
|
||||
"of it before the rollout._"
|
||||
)
|
||||
bullets = "\n".join(f"- {m}" for m in missing)
|
||||
return f"**What this one is currently missing:**\n\n{bullets}"
|
||||
|
||||
|
||||
def _recovery_section(kind: str) -> str:
|
||||
if kind == "pr":
|
||||
return (
|
||||
"**If the bot closes this PR after the rollout:** update the "
|
||||
"description with the missing pieces, then either open a fresh "
|
||||
"PR or comment `@agent-shin reconsider` on the closed PR. If "
|
||||
"Greptile re-scores you at 4/5 or higher I'll reopen and tag "
|
||||
"the PR `ready for review`. (`@greptileai` works on closed PRs "
|
||||
"too; a fresh review is one of the signals that lifts you back "
|
||||
"into the queue.) This is **not** us losing interest in your "
|
||||
"change; far from it. We just need open PRs to be a list of "
|
||||
"things a maintainer can act on, so we can get to yours faster."
|
||||
)
|
||||
return (
|
||||
"**If the bot closes this issue after the rollout:** edit the issue "
|
||||
"description to add the missing pieces, then comment `@agent-shin "
|
||||
"reconsider` on the closed issue. I'll re-evaluate and, if the rubric "
|
||||
"is met, reopen it. (GitHub doesn't let external authors reopen an "
|
||||
"issue a maintainer or bot closed, so the comment is the reliable "
|
||||
"path.) This is **not** us saying the bug isn't real or the request "
|
||||
"isn't useful; it's so the remaining open issues are a list of things "
|
||||
"a maintainer can act on."
|
||||
)
|
||||
|
||||
|
||||
def format_heads_up_comment(
|
||||
*, kind: str, verdict: dict, greptile_score: int | None, cutoff: dt.date
|
||||
) -> str:
|
||||
"""Compose the friendly 7-day heads-up comment posted on a failing PR/issue."""
|
||||
noun = "PR" if kind == "pr" else "issue"
|
||||
rubric = _rubric_section_pr() if kind == "pr" else _rubric_section_issue()
|
||||
cutoff_str = _format_cutoff(cutoff)
|
||||
explanation = (verdict.get("explanation") or "").strip()
|
||||
explanation_block = (
|
||||
f"> _(The judge's note for this one: {explanation})_\n\n" if explanation else ""
|
||||
)
|
||||
|
||||
return (
|
||||
"🚅 **Heads-up: we're turning on the OSS triage bot in "
|
||||
f"{DEFAULT_GRACE_DAYS} days, on {cutoff_str}.**\n"
|
||||
"\n"
|
||||
"We're rolling out **Agent Shin**, an LLM-as-judge triage bot for "
|
||||
f"external {noun}s. Once it's live, the bot reads each open "
|
||||
f"{noun}'s description, scores it against a small rubric, and "
|
||||
f"auto-closes any {noun} that's missing the basics, with a single "
|
||||
f"comment explaining what's missing and how to recover. Full "
|
||||
f"context: [Agent Shin rollout blog post]({ROLLOUT_BLOG_URL}).\n"
|
||||
"\n"
|
||||
f"{rubric}\n"
|
||||
"\n"
|
||||
f"{_description_only_note(kind)}\n"
|
||||
"\n"
|
||||
f"{_missing_section(verdict, greptile_score)}\n"
|
||||
"\n"
|
||||
f"{explanation_block}"
|
||||
"**Timeline (you have a week):**\n"
|
||||
"\n"
|
||||
f"- We turn the bot on in {DEFAULT_GRACE_DAYS} days, on "
|
||||
f"**{cutoff_str}**. You have until then to update this {noun}'s "
|
||||
"description with the missing pieces above.\n"
|
||||
f"- If this {noun} still fails the rubric at **{cutoff_str}**, "
|
||||
"we'll close it.\n"
|
||||
f"- From then on the bot runs daily, and every {noun} that fails "
|
||||
"the rubric gets a **2-hour lifetime**: one warning comment, then "
|
||||
"auto-close 2 hours later.\n"
|
||||
"\n"
|
||||
f"{_recovery_section(kind)}\n"
|
||||
"\n"
|
||||
f"{HEADS_UP_MARKER}"
|
||||
)
|
||||
|
||||
|
||||
def _list_open_numbers(repo: str, kind: str) -> list[int]:
|
||||
"""Return every open PR or issue number in ``repo``.
|
||||
|
||||
Delegates to ``list_open_items`` so the full backlog is fetched (no cap)
|
||||
and the `gh {pr,issue} list` invocation stays in one shared place. ``gh
|
||||
issue list`` would include PRs, but ``list_open_items`` uses the dedicated
|
||||
command per kind, so the two never mix.
|
||||
"""
|
||||
return [
|
||||
item["number"] for item in list_open_items(kind, repo=repo, fields="number")
|
||||
]
|
||||
|
||||
|
||||
def _has_heads_up_marker(item: dict) -> bool:
|
||||
"""Cheap fast-path: check the PR/issue body itself for the marker.
|
||||
|
||||
The marker is appended to the *comment* we post, not the body, so this
|
||||
will only fire if the body literally contains the marker text. We still
|
||||
do the comment-marker check separately below; this body check just lets
|
||||
us short-circuit for PRs/issues that quote the marker for any reason.
|
||||
"""
|
||||
body = item.get("body") or ""
|
||||
return HEADS_UP_MARKER in body
|
||||
|
||||
|
||||
def _comments_have_marker(repo: str, number: int) -> bool:
|
||||
"""True if the bot already posted a comment carrying the marker.
|
||||
|
||||
Used for idempotency: a re-run skips items the previous run notified.
|
||||
Filters by author (matching the sibling marker-checks in
|
||||
``triage_with_llm._has_marker`` and
|
||||
``agent_shin_shared.seconds_since_latest_marker_comment``) so a
|
||||
contributor who quotes the heads-up via GitHub's "Quote reply" — which
|
||||
preserves HTML comments in the raw markdown — can't trick the
|
||||
idempotency check into silently skipping a real heads-up.
|
||||
|
||||
Comments live on the unified issues endpoint regardless of whether the
|
||||
item is a PR or an issue, so no ``kind`` argument is required here.
|
||||
"""
|
||||
expected_login = (
|
||||
os.environ.get("AGENT_SHIN_BOT_LOGIN") or AGENT_SHIN_DEFAULT_BOT_LOGIN
|
||||
).lower()
|
||||
raw = gh(
|
||||
"api",
|
||||
"--paginate",
|
||||
f"repos/{repo}/issues/{number}/comments?per_page=100",
|
||||
)
|
||||
for line in raw.splitlines():
|
||||
line = line.strip()
|
||||
if not line:
|
||||
continue
|
||||
try:
|
||||
payload = json.loads(line)
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
comments = payload if isinstance(payload, list) else [payload]
|
||||
for comment in comments:
|
||||
author = ((comment.get("user") or {}).get("login") or "").lower()
|
||||
if author != expected_login:
|
||||
continue
|
||||
if HEADS_UP_MARKER in (comment.get("body") or ""):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _evaluate_pr(*, repo: str, number: int, model: str, judge: Any = None) -> dict:
|
||||
"""Run the future PR rubric (review_gate) in dry-run and return the result."""
|
||||
return review_gate(
|
||||
repo=repo,
|
||||
number=number,
|
||||
close=False, # we only want the verdict, never act here
|
||||
model=model,
|
||||
judge=judge,
|
||||
)
|
||||
|
||||
|
||||
def _evaluate_issue(*, repo: str, number: int, model: str, judge: Any = None) -> dict:
|
||||
"""Run the future issue rubric (triage kind='issue') in dry-run."""
|
||||
return triage(
|
||||
repo=repo,
|
||||
kind="issue",
|
||||
number=number,
|
||||
close=False,
|
||||
model=model,
|
||||
judge=judge,
|
||||
)
|
||||
|
||||
|
||||
def _would_be_closed(kind: str, result: dict) -> bool:
|
||||
"""True if the future triage would auto-close this PR/issue based on the
|
||||
rubric (regardless of grace-period gating).
|
||||
|
||||
For PRs we trust ``review_gate``'s ``passing`` field — it combines the LLM
|
||||
verdict and the Greptile score. For issues we read the LLM verdict
|
||||
directly. Both fields are ``None``/missing on skip paths
|
||||
(skip-internal-author, skip-llm-error, etc.) where the future bot would
|
||||
NOT close the item — those return False.
|
||||
"""
|
||||
if kind == "pr":
|
||||
passing = result.get("passing")
|
||||
if passing is None:
|
||||
return False # skipped — nothing for the heads-up to warn about
|
||||
return passing is False
|
||||
verdict = result.get("verdict") or {}
|
||||
return (verdict.get("verdict") or "").lower() == "fail"
|
||||
|
||||
|
||||
def _process_one(
|
||||
*,
|
||||
repo: str,
|
||||
kind: str,
|
||||
number: int,
|
||||
model: str,
|
||||
cutoff: dt.date,
|
||||
dry_run: bool,
|
||||
judge: Any = None,
|
||||
skip_marker_check: bool = False,
|
||||
allowlist: frozenset[str] = ALLOWLIST_LOGINS,
|
||||
) -> dict:
|
||||
"""Evaluate one PR/issue and post a heads-up if it would be auto-closed.
|
||||
|
||||
Returns a per-item dict for the summary table.
|
||||
"""
|
||||
base = {"kind": kind, "number": number}
|
||||
fetcher = fetch_pr if kind == "pr" else fetch_issue
|
||||
item = fetcher(repo, number)
|
||||
|
||||
if (item.get("state") or "") != "open":
|
||||
return {**base, "action": "skip-not-open"}
|
||||
if allowlist:
|
||||
login = (item.get("user") or {}).get("login") or ""
|
||||
if login.lower() not in allowlist:
|
||||
return {**base, "action": "skip-not-allowlisted"}
|
||||
elif is_internal_contributor(item):
|
||||
return {**base, "action": "skip-internal-author"}
|
||||
if not skip_marker_check and _has_heads_up_marker(item):
|
||||
return {**base, "action": "skip-already-marked-in-body"}
|
||||
if not skip_marker_check and _comments_have_marker(repo, number):
|
||||
return {**base, "action": "skip-already-notified"}
|
||||
|
||||
if kind == "pr":
|
||||
result = _evaluate_pr(repo=repo, number=number, model=model, judge=judge)
|
||||
else:
|
||||
result = _evaluate_issue(repo=repo, number=number, model=model, judge=judge)
|
||||
|
||||
if not _would_be_closed(kind, result):
|
||||
return {**base, "action": "skip-passing", "evaluator": result.get("action")}
|
||||
|
||||
verdict = result.get("verdict") or {}
|
||||
greptile_score = result.get("greptile_score") if kind == "pr" else None
|
||||
comment = format_heads_up_comment(
|
||||
kind=kind, verdict=verdict, greptile_score=greptile_score, cutoff=cutoff
|
||||
)
|
||||
maybe_post_comment(repo, number, comment, dry_run=dry_run)
|
||||
return {
|
||||
**base,
|
||||
"action": "heads-up-posted" if not dry_run else "would-post-heads-up",
|
||||
"verdict": (verdict.get("verdict") or "").lower(),
|
||||
"greptile_score": greptile_score,
|
||||
}
|
||||
|
||||
|
||||
def _print_summary(results: list[dict]) -> None:
|
||||
"""Tally per-action counts so a dry-run preview tells you at a glance how
|
||||
many comments the real run would post."""
|
||||
counts: dict[str, int] = {}
|
||||
for r in results:
|
||||
counts[r["action"]] = counts.get(r["action"], 0) + 1
|
||||
print("\n=== rollout heads-up summary ===")
|
||||
for action in sorted(counts):
|
||||
print(f" {action:35s} {counts[action]}")
|
||||
print(f" total {len(results)}")
|
||||
|
||||
|
||||
def run(
|
||||
*,
|
||||
repo: str,
|
||||
close: bool,
|
||||
cutoff: dt.date,
|
||||
model: str,
|
||||
kinds: tuple[str, ...] = ("pr", "issue"),
|
||||
judge: Any = None,
|
||||
only_numbers: dict[str, list[int]] | None = None,
|
||||
skip_marker_check: bool = False,
|
||||
) -> list[dict]:
|
||||
"""Sweep ``repo`` and post heads-up comments. Returns the per-item results."""
|
||||
dry_run = not close
|
||||
if dry_run:
|
||||
print(
|
||||
f"[DRY RUN] sweeping {repo}; --close not passed, no comments will be posted."
|
||||
)
|
||||
else:
|
||||
print(f"[REAL RUN] sweeping {repo}; comments WILL be posted.")
|
||||
print(f"Cutoff date in comment body: {cutoff.isoformat()}")
|
||||
|
||||
results: list[dict] = []
|
||||
for kind in kinds:
|
||||
if only_numbers and kind in only_numbers:
|
||||
numbers = list(only_numbers[kind])
|
||||
else:
|
||||
numbers = _list_open_numbers(repo, kind)
|
||||
print(f"\n--- {kind}s: {len(numbers)} open ---")
|
||||
for n in numbers:
|
||||
try:
|
||||
result = _process_one(
|
||||
repo=repo,
|
||||
kind=kind,
|
||||
number=n,
|
||||
model=model,
|
||||
cutoff=cutoff,
|
||||
dry_run=dry_run,
|
||||
judge=judge,
|
||||
skip_marker_check=skip_marker_check,
|
||||
)
|
||||
except (
|
||||
Exception
|
||||
) as exc: # noqa: BLE001 - per-item errors don't abort the sweep
|
||||
result = {
|
||||
"kind": kind,
|
||||
"number": n,
|
||||
"action": "error",
|
||||
"error": str(exc),
|
||||
}
|
||||
print(f"!! {kind}#{n}: {exc}", file=sys.stderr)
|
||||
print(f" {kind}#{n}: {result['action']}")
|
||||
results.append(result)
|
||||
_print_summary(results)
|
||||
return results
|
||||
|
||||
|
||||
def main() -> int:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--repo", required=True, help="owner/repo")
|
||||
parser.add_argument(
|
||||
"--close",
|
||||
action="store_true",
|
||||
help=(
|
||||
"Actually post comments. Without this flag the script is in "
|
||||
"dry-run mode and only logs what it would do."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--close-on",
|
||||
type=dt.date.fromisoformat,
|
||||
default=None,
|
||||
help=(
|
||||
"Cutoff date shown in the heads-up comment as the rollout date "
|
||||
f"(default: today + {DEFAULT_GRACE_DAYS} days)."
|
||||
),
|
||||
)
|
||||
parser.add_argument(
|
||||
"--model",
|
||||
default=os.environ.get("TRIAGE_MODEL") or DEFAULT_MODEL,
|
||||
help=f"Model for the rubric LLM judge (default: {DEFAULT_MODEL}).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--kind",
|
||||
choices=("pr", "issue", "both"),
|
||||
default="both",
|
||||
help="Restrict the sweep to PRs or issues only (default: both).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--only-pr",
|
||||
type=int,
|
||||
action="append",
|
||||
default=[],
|
||||
help="Limit the PR sweep to these PR numbers (repeat for several).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--only-issue",
|
||||
type=int,
|
||||
action="append",
|
||||
default=[],
|
||||
help="Limit the issue sweep to these issue numbers (repeat for several).",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ignore-existing-marker",
|
||||
action="store_true",
|
||||
help=(
|
||||
"Re-post on PRs/issues that already carry the heads-up marker. "
|
||||
"Useful for testing the comment wording on a known PR."
|
||||
),
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
cutoff = args.close_on or (
|
||||
dt.datetime.now(dt.timezone.utc).date() + dt.timedelta(days=DEFAULT_GRACE_DAYS)
|
||||
)
|
||||
|
||||
kinds: tuple[str, ...]
|
||||
if args.kind == "pr":
|
||||
kinds = ("pr",)
|
||||
elif args.kind == "issue":
|
||||
kinds = ("issue",)
|
||||
else:
|
||||
kinds = ("pr", "issue")
|
||||
|
||||
only: dict[str, list[int]] = {}
|
||||
if args.only_pr:
|
||||
only["pr"] = args.only_pr
|
||||
if args.only_issue:
|
||||
only["issue"] = args.only_issue
|
||||
|
||||
# The script must NOT hit the LLM in dry-run if no key is set — we still
|
||||
# want a useful preview that says "skip-no-llm-key" for items that would
|
||||
# have been judged. Production runs require OPENAI_API_KEY.
|
||||
if args.close and not os.environ.get("OPENAI_API_KEY"):
|
||||
parser.error("OPENAI_API_KEY must be set for --close (real-run) mode.")
|
||||
|
||||
run(
|
||||
repo=args.repo,
|
||||
close=args.close,
|
||||
cutoff=cutoff,
|
||||
model=args.model,
|
||||
kinds=kinds,
|
||||
only_numbers=only or None,
|
||||
skip_marker_check=args.ignore_existing_marker,
|
||||
)
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
1778
.github/scripts/triage_with_llm.py
vendored
Normal file
1778
.github/scripts/triage_with_llm.py
vendored
Normal file
File diff suppressed because it is too large
Load diff
2
.github/workflows/check-ui-api-types.yml
vendored
2
.github/workflows/check-ui-api-types.yml
vendored
|
|
@ -54,7 +54,7 @@ jobs:
|
|||
run: uv run --no-sync prisma generate --schema litellm/proxy/schema.prisma
|
||||
|
||||
- name: Set up Node.js
|
||||
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0
|
||||
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
|
||||
with:
|
||||
node-version: "20"
|
||||
cache: "npm"
|
||||
|
|
|
|||
92
.github/workflows/close_low_quality_prs.yml
vendored
Normal file
92
.github/workflows/close_low_quality_prs.yml
vendored
Normal file
|
|
@ -0,0 +1,92 @@
|
|||
name: Close Low-Quality PRs
|
||||
|
||||
# Auto-close any open PR (including drafts, regardless of age) authored by an
|
||||
# external OSS contributor that Greptile reviewed with a confidence score
|
||||
# below 4/5. Closures are explained in a comment that tells the contributor
|
||||
# to push fixes and open a fresh PR (since OSS authors cannot reopen a PR
|
||||
# closed by a bot/maintainer) or comment `@agent-shin reconsider` to have
|
||||
# Agent Shin re-evaluate.
|
||||
#
|
||||
# Manual one-off run:
|
||||
# gh workflow run "Close Low-Quality PRs" -f close=true
|
||||
#
|
||||
# Dry-run preview (no PRs are touched):
|
||||
# gh workflow run "Close Low-Quality PRs" -f close=false
|
||||
|
||||
on:
|
||||
schedule:
|
||||
# Daily at 09:00 UTC. Pairs well with the stale-issue workflow at midnight.
|
||||
- cron: "0 9 * * *"
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
close:
|
||||
description: "Actually close matching PRs (false = dry run)."
|
||||
required: false
|
||||
default: "false"
|
||||
type: choice
|
||||
options:
|
||||
- "true"
|
||||
- "false"
|
||||
min_age_days:
|
||||
description: "Minimum PR age in days (default 0 = no age filter)."
|
||||
required: false
|
||||
default: "0"
|
||||
min_score:
|
||||
description: "Greptile score below which a PR is closed (1-5)."
|
||||
required: false
|
||||
default: "4"
|
||||
limit:
|
||||
description: "Maximum number of PRs to close in a single run."
|
||||
required: false
|
||||
default: "25"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
pull-requests: write
|
||||
issues: write
|
||||
|
||||
jobs:
|
||||
close-low-quality-prs:
|
||||
if: github.repository == 'BerriAI/litellm'
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout triage script
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
sparse-checkout: .github/scripts
|
||||
persist-credentials: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Run low-quality PR closer
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
# Scheduled runs are ALWAYS dry-run, even when AGENT_SHIN_ENABLED is
|
||||
# "true", so the team can QA the closer's verdicts in step summaries
|
||||
# before any contributor sees a PR closed. Real closures only happen
|
||||
# on manual workflow_dispatch with close=true (and the variable set).
|
||||
CLOSE_FLAG: ${{ github.event.inputs.close || 'false' }}
|
||||
AGENT_SHIN_ENABLED: ${{ vars.AGENT_SHIN_ENABLED }}
|
||||
MIN_AGE_DAYS: ${{ github.event.inputs.min_age_days || '0' }}
|
||||
MIN_SCORE: ${{ github.event.inputs.min_score || '4' }}
|
||||
LIMIT: ${{ github.event.inputs.limit || '25' }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
ARGS=(
|
||||
--repo "${{ github.repository }}"
|
||||
--min-age-days "${MIN_AGE_DAYS}"
|
||||
--min-score "${MIN_SCORE}"
|
||||
--limit "${LIMIT}"
|
||||
)
|
||||
if [ "${AGENT_SHIN_ENABLED:-false}" != "true" ]; then
|
||||
echo "::notice::AGENT_SHIN_ENABLED is not 'true' -> forcing dry-run regardless of close input."
|
||||
elif [ "${GITHUB_EVENT_NAME:-}" = "workflow_dispatch" ] && [ "${CLOSE_FLAG}" = "true" ]; then
|
||||
ARGS+=(--close)
|
||||
echo "::notice::Running in close-on-fail mode."
|
||||
else
|
||||
echo "::notice::AGENT_SHIN_ENABLED is true but this trigger is dry-run (scheduled event or close=false)."
|
||||
fi
|
||||
python3 .github/scripts/close_low_quality_prs.py "${ARGS[@]}"
|
||||
6
.github/workflows/codeql.yml
vendored
6
.github/workflows/codeql.yml
vendored
|
|
@ -43,14 +43,14 @@ jobs:
|
|||
persist-credentials: false
|
||||
|
||||
- name: Initialize CodeQL
|
||||
uses: github/codeql-action/init@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3
|
||||
uses: github/codeql-action/init@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3.34.1
|
||||
with:
|
||||
languages: ${{ matrix.language }}
|
||||
build-mode: ${{ matrix.build-mode }}
|
||||
config-file: ./.github/codeql/codeql-config.yml
|
||||
|
||||
- name: Perform CodeQL Analysis
|
||||
uses: github/codeql-action/analyze@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3
|
||||
uses: github/codeql-action/analyze@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3.34.1
|
||||
with:
|
||||
category: "/language:${{ matrix.language }}"
|
||||
output: sarif-results
|
||||
|
|
@ -77,7 +77,7 @@ jobs:
|
|||
output: sarif-results/python.sarif
|
||||
|
||||
- name: Upload SARIF
|
||||
uses: github/codeql-action/upload-sarif@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3
|
||||
uses: github/codeql-action/upload-sarif@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3.34.1
|
||||
with:
|
||||
sarif_file: sarif-results
|
||||
category: "/language:${{ matrix.language }}"
|
||||
|
|
|
|||
4
.github/workflows/test-litellm-ui-build.yml
vendored
4
.github/workflows/test-litellm-ui-build.yml
vendored
|
|
@ -25,7 +25,7 @@ jobs:
|
|||
persist-credentials: false
|
||||
|
||||
- name: Setup Node.js
|
||||
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0
|
||||
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
|
||||
with:
|
||||
node-version: "20"
|
||||
cache: "npm"
|
||||
|
|
@ -77,7 +77,7 @@ jobs:
|
|||
|
||||
- name: Setup Node.js
|
||||
if: steps.changed.outputs.has_files == 'true'
|
||||
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0
|
||||
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
|
||||
with:
|
||||
node-version: "20"
|
||||
cache: "npm"
|
||||
|
|
|
|||
10
.github/workflows/test-unit-proxy-endpoints.yml
vendored
10
.github/workflows/test-unit-proxy-endpoints.yml
vendored
|
|
@ -11,8 +11,6 @@ on:
|
|||
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
|
|
@ -20,6 +18,10 @@ concurrency:
|
|||
|
||||
jobs:
|
||||
proxy-endpoints:
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
uses: ./.github/workflows/_test-unit-base.yml
|
||||
with:
|
||||
test-path: >-
|
||||
|
|
@ -52,6 +54,10 @@ jobs:
|
|||
# is independent and its coverage artifact is uploaded separately.
|
||||
# See: https://www.notion.so/36c43b8acdab81ee845fd5365128a2fc
|
||||
proxy-server:
|
||||
permissions:
|
||||
contents: read
|
||||
id-token: write
|
||||
pull-requests: write
|
||||
uses: ./.github/workflows/_test-unit-base.yml
|
||||
with:
|
||||
test-path: tests/test_litellm/proxy/proxy_server
|
||||
|
|
|
|||
7
.github/workflows/test_server_root_path.yml
vendored
7
.github/workflows/test_server_root_path.yml
vendored
|
|
@ -32,17 +32,16 @@ jobs:
|
|||
df -h /
|
||||
|
||||
- name: Set up Docker Buildx
|
||||
uses: docker/setup-buildx-action@8d2750c68a42422c14e847fe6c8ac0403b4cbd6f # v3.12
|
||||
uses: docker/setup-buildx-action@8d2750c68a42422c14e847fe6c8ac0403b4cbd6f # v3.12.0
|
||||
|
||||
- name: Build Docker image
|
||||
uses: docker/build-push-action@0adf9959216b96bec444f325f1e493d4aa344497 #v6.14
|
||||
uses: docker/build-push-action@0adf9959216b96bec444f325f1e493d4aa344497 # v6.14.0
|
||||
with:
|
||||
context: .
|
||||
file: ./docker/Dockerfile.non_root
|
||||
tags: litellm-test:${{ github.sha }}
|
||||
load: true
|
||||
cache-from: type=gha
|
||||
cache-to: type=gha,mode=max
|
||||
push: false
|
||||
|
||||
- name: Start LiteLLM container with SERVER_ROOT_PATH
|
||||
run: |
|
||||
|
|
|
|||
96
.github/workflows/triage_issue_with_llm.yml
vendored
Normal file
96
.github/workflows/triage_issue_with_llm.yml
vendored
Normal file
|
|
@ -0,0 +1,96 @@
|
|||
name: Agent Shin — Issue triage
|
||||
|
||||
# LLM-as-judge triage for external GitHub issues.
|
||||
#
|
||||
# DRY-RUN BY DEFAULT. See .github/workflows/triage_pr_with_llm.yml for the
|
||||
# enablement procedure — same repo variable (`AGENT_SHIN_ENABLED=true`)
|
||||
# unlocks the PR and issue triage flows together.
|
||||
|
||||
on:
|
||||
issues:
|
||||
types: [opened, reopened]
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
issue_number:
|
||||
description: "Issue number to triage manually."
|
||||
required: true
|
||||
close:
|
||||
description: "If true and AGENT_SHIN_ENABLED=true, actually close on fail."
|
||||
required: false
|
||||
default: "false"
|
||||
type: choice
|
||||
options:
|
||||
- "true"
|
||||
- "false"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
issues: write
|
||||
|
||||
jobs:
|
||||
triage:
|
||||
if: github.repository == 'BerriAI/litellm'
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout triage script
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
sparse-checkout: .github/scripts
|
||||
persist-credentials: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Install LLM client
|
||||
run: pip install --no-cache-dir --require-hashes -r .github/scripts/triage-requirements.txt
|
||||
|
||||
- name: Run Agent Shin
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
# Only expose the LLM key when the bot is enabled or a collaborator
|
||||
# triggers it manually, so an external user can't force paid LLM
|
||||
# calls by churning issues while the bot is still in dry-run.
|
||||
# The Python script calls the LLM whenever this var is set
|
||||
# (regardless of `--close`); stripping `--close` doesn't suppress
|
||||
# the API call, only the destructive side effects.
|
||||
OPENAI_API_KEY: ${{ (vars.AGENT_SHIN_ENABLED == 'true' || github.event_name == 'workflow_dispatch') && secrets.OPENAI_API_KEY || '' }}
|
||||
OPENAI_BASE_URL: ${{ vars.OPENAI_BASE_URL }}
|
||||
TRIAGE_MODEL: ${{ vars.TRIAGE_MODEL }}
|
||||
AGENT_SHIN_ENABLED: ${{ vars.AGENT_SHIN_ENABLED }}
|
||||
DISPATCH_CLOSE: ${{ github.event.inputs.close }}
|
||||
ISSUE_NUMBER: ${{ github.event.issue.number || github.event.inputs.issue_number }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
ARGS=(--repo "${{ github.repository }}" --issue "${ISSUE_NUMBER}")
|
||||
# Fail-safe gating: only the EXACT string "true" enables the
|
||||
# destructive --close path. The workflow_dispatch input is a
|
||||
# `choice` dropdown of "true"/"false" so the UI is constrained,
|
||||
# but the API (`gh workflow run -f close=...`) accepts any
|
||||
# string, and a `!= "false"` check would treat "True", "yes",
|
||||
# "1", "TRUE", typos, and accidental whitespace as enabling
|
||||
# closure. Mirror the Greptile closer's `= "true"` pattern.
|
||||
if [ "${AGENT_SHIN_ENABLED:-false}" = "true" ] && [ "${DISPATCH_CLOSE:-false}" = "true" ]; then
|
||||
ARGS+=(--close)
|
||||
echo "::notice::Agent Shin is ENABLED and running in close-on-fail mode."
|
||||
elif [ "${AGENT_SHIN_ENABLED:-false}" = "true" ]; then
|
||||
echo "::notice::Agent Shin is ENABLED but this trigger is dry-run (workflow_dispatch close != 'true')."
|
||||
else
|
||||
echo "::notice::Agent Shin is in DRY-RUN mode (AGENT_SHIN_ENABLED is not 'true'). No comments will be posted; no issues will be closed."
|
||||
fi
|
||||
# Automatic `issues` events stay dry-run regardless until the team
|
||||
# explicitly invokes workflow_dispatch with close=true.
|
||||
if [ "${GITHUB_EVENT_NAME:-}" = "issues" ]; then
|
||||
# filter out --close rather than substituting to "" (which would
|
||||
# leave an empty positional arg that argparse rejects)
|
||||
FILTERED=()
|
||||
for arg in "${ARGS[@]}"; do
|
||||
if [ "${arg}" != "--close" ]; then
|
||||
FILTERED+=("${arg}")
|
||||
fi
|
||||
done
|
||||
ARGS=("${FILTERED[@]}")
|
||||
echo "::notice::issues trigger -> forcing dry-run."
|
||||
fi
|
||||
python3 .github/scripts/triage_with_llm.py "${ARGS[@]}"
|
||||
172
.github/workflows/triage_reconsider.yml
vendored
Normal file
172
.github/workflows/triage_reconsider.yml
vendored
Normal file
|
|
@ -0,0 +1,172 @@
|
|||
name: Agent Shin — reconsider
|
||||
|
||||
# Comment-trigger workflow: when the PR/issue author (or an internal
|
||||
# collaborator) comments `@agent-shin reconsider` on a CLOSED PR/issue,
|
||||
# Agent Shin re-runs LLM-judge triage on the current title+body and:
|
||||
#
|
||||
# - on PASS: posts a "re-evaluated and reopened" comment + reopens.
|
||||
# - on FAIL: posts a "still missing X" comment and leaves it closed,
|
||||
# so the contributor can iterate again.
|
||||
#
|
||||
# This exists because GitHub does NOT let an external (non-write-access)
|
||||
# OSS contributor reopen a PR/issue closed by a bot or maintainer. Without
|
||||
# this comment trigger, a contributor whose PR Agent Shin auto-closed
|
||||
# would have no path back into the review queue except opening a fresh PR
|
||||
# (which loses the original PR's history). The bot, on the other hand,
|
||||
# has write access via GH_TOKEN and can reopen on their behalf.
|
||||
#
|
||||
# DRY-RUN BY DEFAULT — gated on `vars.AGENT_SHIN_ENABLED == 'true'` just
|
||||
# like the other Agent Shin workflows. The workflow also gates on the
|
||||
# commenter being either the PR/issue author or an internal collaborator
|
||||
# (OWNER/MEMBER/COLLABORATOR) so random commenters cannot DOS the LLM
|
||||
# judge or force a reopen.
|
||||
|
||||
on:
|
||||
issue_comment:
|
||||
types: [created]
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
issues: write
|
||||
pull-requests: write
|
||||
|
||||
jobs:
|
||||
reconsider:
|
||||
if: |
|
||||
github.repository == 'BerriAI/litellm'
|
||||
&& contains(github.event.comment.body, '@agent-shin reconsider')
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Authorize commenter
|
||||
# Only the PR/issue author OR an internal collaborator may trigger
|
||||
# a reconsider. Outside random commenters could otherwise spam the
|
||||
# phrase to burn LLM budget or, if a fail-open bug were ever
|
||||
# introduced, force a reopen on someone else's behalf.
|
||||
#
|
||||
# We expose the authorization decision as a step output and gate
|
||||
# every subsequent (potentially destructive) step on it. A `run:`
|
||||
# step with `exit 0` would NOT stop the job — only `if:` gating
|
||||
# on a known-true output is safe here.
|
||||
id: auth
|
||||
env:
|
||||
COMMENTER: ${{ github.event.comment.user.login }}
|
||||
AUTHOR: ${{ github.event.issue.user.login }}
|
||||
ASSOCIATION: ${{ github.event.comment.author_association }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
if [ "${COMMENTER}" = "${AUTHOR}" ]; then
|
||||
echo "::notice::Authorized: commenter is the PR/issue author."
|
||||
echo "authorized=true" >> "$GITHUB_OUTPUT"
|
||||
exit 0
|
||||
fi
|
||||
case "${ASSOCIATION}" in
|
||||
OWNER|MEMBER|COLLABORATOR)
|
||||
echo "::notice::Authorized: commenter is an internal collaborator (${ASSOCIATION})."
|
||||
echo "authorized=true" >> "$GITHUB_OUTPUT"
|
||||
;;
|
||||
*)
|
||||
echo "::notice::Commenter '${COMMENTER}' (${ASSOCIATION}) is not authorized to trigger reconsider; skipping subsequent steps."
|
||||
echo "authorized=false" >> "$GITHUB_OUTPUT"
|
||||
;;
|
||||
esac
|
||||
|
||||
- name: React 👀 to acknowledge the reconsider
|
||||
# Add an eyes reaction to the triggering comment the moment we accept
|
||||
# it, so the contributor gets instant feedback that the bot saw their
|
||||
# `@agent-shin reconsider` before the slower triage steps run. Gated on
|
||||
# AGENT_SHIN_ENABLED so dry-run leaves no visible trace. Best-effort:
|
||||
# a reactions API hiccup must never fail the actual reconsider.
|
||||
if: steps.auth.outputs.authorized == 'true' && vars.AGENT_SHIN_ENABLED == 'true'
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
COMMENT_ID: ${{ github.event.comment.id }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
gh api --method POST \
|
||||
-H "Accept: application/vnd.github+json" \
|
||||
"repos/${{ github.repository }}/issues/comments/${COMMENT_ID}/reactions" \
|
||||
-f content=eyes \
|
||||
|| echo "::warning::failed to add 👀 reaction (non-fatal)"
|
||||
|
||||
- name: Checkout triage script
|
||||
if: steps.auth.outputs.authorized == 'true'
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
sparse-checkout: .github/scripts
|
||||
persist-credentials: false
|
||||
|
||||
- name: Set up Python
|
||||
if: steps.auth.outputs.authorized == 'true'
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Install LLM client
|
||||
if: steps.auth.outputs.authorized == 'true'
|
||||
run: pip install --no-cache-dir --require-hashes -r .github/scripts/triage-requirements.txt
|
||||
|
||||
- name: Run Agent Shin reconsider
|
||||
if: steps.auth.outputs.authorized == 'true'
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
# Only expose the LLM key when the bot is enabled, so a PR/issue
|
||||
# author can't force paid LLM calls by spamming `@agent-shin
|
||||
# reconsider` while the bot is still in dry-run. The Python script
|
||||
# calls the LLM whenever this var is set (regardless of `--close`);
|
||||
# stripping `--close` doesn't suppress the API call, only the
|
||||
# destructive side effects. Mirror the gating used by every other
|
||||
# Agent Shin workflow (triage_pr_with_llm.yml, review_gate.yml, ...).
|
||||
OPENAI_API_KEY: ${{ vars.AGENT_SHIN_ENABLED == 'true' && secrets.OPENAI_API_KEY || '' }}
|
||||
OPENAI_BASE_URL: ${{ vars.OPENAI_BASE_URL }}
|
||||
TRIAGE_MODEL: ${{ vars.TRIAGE_MODEL }}
|
||||
AGENT_SHIN_ENABLED: ${{ vars.AGENT_SHIN_ENABLED }}
|
||||
# `issue_comment` events fire for both issues and PR comments.
|
||||
# `issue.pull_request` is set iff this is a PR comment, so we use
|
||||
# its presence to decide whether to invoke `--pr N` or `--issue N`.
|
||||
IS_PR: ${{ github.event.issue.pull_request != null }}
|
||||
NUMBER: ${{ github.event.issue.number }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
if [ "${IS_PR}" = "true" ]; then
|
||||
ARGS=(--repo "${{ github.repository }}" --pr "${NUMBER}" --reconsider)
|
||||
else
|
||||
ARGS=(--repo "${{ github.repository }}" --issue "${NUMBER}" --reconsider)
|
||||
fi
|
||||
# Reconsider's destructive actions (post comment + reopen) are
|
||||
# gated on `--close`, mirroring the regular triage workflows.
|
||||
# When AGENT_SHIN_ENABLED is not the EXACT string "true", we
|
||||
# still run the script so its verdict + would-X action lands in
|
||||
# the step summary for QA — but without `--close`, the script
|
||||
# returns `would-reopen` / `would-reconsider-still-failing`
|
||||
# instead of touching GitHub state.
|
||||
#
|
||||
# Use the positive `= "true"` gate (not `!= "true" -> exit`) so
|
||||
# the workflow guardrails in
|
||||
# tests/test_litellm/test_github_triage_workflows.py see the
|
||||
# canonical fail-safe enable pattern. Unknown values like
|
||||
# "True", "yes", "1", or typos fall through to the dry-run
|
||||
# branch, which is the safe default.
|
||||
if [ "${AGENT_SHIN_ENABLED:-false}" = "true" ]; then
|
||||
ARGS+=(--close)
|
||||
echo "::notice::Agent Shin reconsider ENABLED — running real triage (close=true)."
|
||||
else
|
||||
echo "::notice::AGENT_SHIN_ENABLED is not 'true' -> reconsider stays in dry-run (no comment, no reopen)."
|
||||
fi
|
||||
python3 .github/scripts/triage_with_llm.py "${ARGS[@]}"
|
||||
|
||||
- name: React 👍 when the reconsider finishes
|
||||
# Once the reconsider run has completed successfully, add a thumbs-up so
|
||||
# the contributor sees the bot is done (the 👀 stays, signalling
|
||||
# seen -> handled). `success()` keeps this from firing if the run
|
||||
# errored, and the AGENT_SHIN_ENABLED gate keeps dry-run inert.
|
||||
if: success() && steps.auth.outputs.authorized == 'true' && vars.AGENT_SHIN_ENABLED == 'true'
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
COMMENT_ID: ${{ github.event.comment.id }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
gh api --method POST \
|
||||
-H "Accept: application/vnd.github+json" \
|
||||
"repos/${{ github.repository }}/issues/comments/${COMMENT_ID}/reactions" \
|
||||
-f content=+1 \
|
||||
|| echo "::warning::failed to add 👍 reaction (non-fatal)"
|
||||
92
.github/workflows/triage_rollout_heads_up.yml
vendored
Normal file
92
.github/workflows/triage_rollout_heads_up.yml
vendored
Normal file
|
|
@ -0,0 +1,92 @@
|
|||
name: Agent Shin — rollout heads-up (one-shot)
|
||||
|
||||
# Fires the 7-day heads-up comment on every open external PR/issue that the
|
||||
# new triage bot would auto-close. The real sweep is a deliberate one-shot:
|
||||
# trigger it at rollout via a manual `workflow_dispatch` with `dry_run=false`.
|
||||
# The script is idempotent (skips items that already carry the
|
||||
# `<!-- agent-shin:rollout-heads-up -->` marker), so a re-run is harmless.
|
||||
#
|
||||
# The automatic push trigger runs DRY-RUN only, so merging the script to
|
||||
# `litellm_internal_staging` never posts a comment; it just confirms the
|
||||
# workflow is wired up. Posting real comments requires the manual dispatch,
|
||||
# which is also the only trigger that exposes `OPENAI_API_KEY`. The heads-up
|
||||
# is intentionally NOT gated on `AGENT_SHIN_ENABLED`: it has to warn
|
||||
# contributors while that flag is still off, ahead of the flip that turns on
|
||||
# auto-closing.
|
||||
#
|
||||
# The workflow is a thin shell over `.github/scripts/triage_rollout_heads_up.py`.
|
||||
# Dry-run vs. real run differ in EXACTLY one CLI flag (`--close`), added only
|
||||
# on a manual dispatch with `dry_run=false`.
|
||||
|
||||
on:
|
||||
push:
|
||||
branches:
|
||||
- litellm_internal_staging
|
||||
paths:
|
||||
# The presence of this script on staging IS the rollout merge marker.
|
||||
# Editing the file later would re-fire the workflow; that's safe because
|
||||
# the script skips PRs/issues that already have the heads-up marker.
|
||||
- ".github/scripts/triage_rollout_heads_up.py"
|
||||
workflow_dispatch:
|
||||
inputs:
|
||||
dry_run:
|
||||
description: "Dry run (true = preview only, false = actually post comments)."
|
||||
required: false
|
||||
default: "true"
|
||||
type: choice
|
||||
options:
|
||||
- "true"
|
||||
- "false"
|
||||
|
||||
permissions:
|
||||
contents: read
|
||||
issues: write
|
||||
pull-requests: write
|
||||
|
||||
jobs:
|
||||
heads-up:
|
||||
if: github.repository == 'BerriAI/litellm'
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Checkout triage scripts
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
with:
|
||||
sparse-checkout: .github/scripts
|
||||
persist-credentials: false
|
||||
|
||||
- name: Set up Python
|
||||
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
|
||||
with:
|
||||
python-version: "3.12"
|
||||
|
||||
- name: Install LLM client
|
||||
run: pip install --no-cache-dir --require-hashes -r .github/scripts/triage-requirements.txt
|
||||
|
||||
- name: Run heads-up sweep
|
||||
env:
|
||||
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
|
||||
# Only the manual dispatch (the real-run trigger) needs the LLM key.
|
||||
# The automatic push trigger runs dry-run and never posts, so it gets
|
||||
# no key. Mirrors the sibling triage workflows, which expose the key
|
||||
# only on an enabled/dispatched run rather than unconditionally.
|
||||
OPENAI_API_KEY: ${{ github.event_name == 'workflow_dispatch' && secrets.OPENAI_API_KEY || '' }}
|
||||
OPENAI_BASE_URL: ${{ vars.OPENAI_BASE_URL }}
|
||||
TRIAGE_MODEL: ${{ vars.TRIAGE_MODEL }}
|
||||
# The real run is a deliberate manual dispatch with dry_run=false.
|
||||
# Use the EXACT "false" comparison so any unexpected input value
|
||||
# fail-closes to dry-run (mirrors the AGENT_SHIN_ENABLED pattern in
|
||||
# the sibling workflows). The automatic push trigger always stays
|
||||
# dry-run, so merging the script never posts.
|
||||
DRY_RUN_INPUT: ${{ github.event.inputs.dry_run }}
|
||||
run: |
|
||||
set -euo pipefail
|
||||
ARGS=(--repo "${{ github.repository }}")
|
||||
if [ "${GITHUB_EVENT_NAME:-}" = "workflow_dispatch" ] && [ "${DRY_RUN_INPUT:-true}" = "false" ]; then
|
||||
ARGS+=(--close)
|
||||
echo "::notice::Manual rollout dispatch with dry_run=false -> heads-up comments WILL be posted."
|
||||
elif [ "${GITHUB_EVENT_NAME:-}" = "workflow_dispatch" ]; then
|
||||
echo "::notice::Manual dispatch in dry-run mode -> previewing only, no comments will be posted."
|
||||
else
|
||||
echo "::notice::Automatic push trigger -> dry-run preview only. Fire the real rollout sweep with a manual workflow_dispatch (dry_run=false)."
|
||||
fi
|
||||
python3 .github/scripts/triage_rollout_heads_up.py "${ARGS[@]}"
|
||||
13
.github/workflows/zizmor.yml
vendored
13
.github/workflows/zizmor.yml
vendored
|
|
@ -2,9 +2,9 @@ name: GitHub Actions Security Analysis
|
|||
|
||||
on:
|
||||
push:
|
||||
branches: [main]
|
||||
branches: [main, litellm_internal_staging]
|
||||
pull_request:
|
||||
branches: [main]
|
||||
branches: [main, litellm_internal_staging]
|
||||
|
||||
concurrency:
|
||||
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
|
||||
|
|
@ -18,9 +18,7 @@ jobs:
|
|||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 5
|
||||
permissions:
|
||||
security-events: write
|
||||
contents: read
|
||||
actions: read
|
||||
steps:
|
||||
- name: Checkout repository
|
||||
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
|
||||
|
|
@ -28,4 +26,9 @@ jobs:
|
|||
persist-credentials: false
|
||||
|
||||
- name: Run zizmor
|
||||
uses: zizmorcore/zizmor-action@71321a20a9ded102f6e9ce5718a2fcec2c4f70d8 # v0.5.2
|
||||
uses: zizmorcore/zizmor-action@5f14fd08f7cf1cb1609c1e344975f152c7ee938d # v0.5.6
|
||||
with:
|
||||
version: "1.24.1"
|
||||
min-severity: medium
|
||||
advanced-security: false
|
||||
annotations: true
|
||||
|
|
|
|||
|
|
@ -345,6 +345,7 @@ curl -X POST 'http://0.0.0.0:4000/v1/chat/completions' \
|
|||
| [OVHCloud AI Endpoints (`ovhcloud`)](https://docs.litellm.ai/docs/providers/ovhcloud) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Perplexity AI (`perplexity`)](https://docs.litellm.ai/docs/providers/perplexity) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Petals (`petals`)](https://docs.litellm.ai/docs/providers/petals) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Pinstripes (`pinstripes`)](https://docs.litellm.ai/docs/providers/pinstripes) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Predibase (`predibase`)](https://docs.litellm.ai/docs/providers/predibase) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Recraft (`recraft`)](https://docs.litellm.ai/docs/providers/recraft) | | | | | ✅ | | | | | |
|
||||
| [Replicate (`replicate`)](https://docs.litellm.ai/docs/providers/replicate) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
|
|
|
|||
|
|
@ -1,31 +1,31 @@
|
|||
{
|
||||
"reportAny": {
|
||||
"baseline": 24954,
|
||||
"baseline": 24989,
|
||||
"slack": 2500
|
||||
},
|
||||
"reportArgumentType": {
|
||||
"baseline": 1863,
|
||||
"baseline": 1934,
|
||||
"slack": 180
|
||||
},
|
||||
"reportAssignmentType": {
|
||||
"baseline": 220,
|
||||
"slack": 3
|
||||
"slack": 22
|
||||
},
|
||||
"reportAttributeAccessIssue": {
|
||||
"baseline": 335,
|
||||
"slack": 3
|
||||
"baseline": 346,
|
||||
"slack": 35
|
||||
},
|
||||
"reportCallIssue": {
|
||||
"baseline": 77,
|
||||
"baseline": 87,
|
||||
"slack": 10
|
||||
},
|
||||
"reportConstantRedefinition": {
|
||||
"baseline": 39,
|
||||
"slack": 3
|
||||
"slack": 4
|
||||
},
|
||||
"reportDeprecated": {
|
||||
"baseline": 217,
|
||||
"slack": 10
|
||||
"slack": 22
|
||||
},
|
||||
"reportDuplicateImport": {
|
||||
"baseline": 28,
|
||||
|
|
@ -41,11 +41,11 @@
|
|||
},
|
||||
"reportGeneralTypeIssues": {
|
||||
"baseline": 151,
|
||||
"slack": 3
|
||||
"slack": 15
|
||||
},
|
||||
"reportIncompatibleMethodOverride": {
|
||||
"baseline": 52,
|
||||
"slack": 10
|
||||
"slack": 5
|
||||
},
|
||||
"reportIncompatibleVariableOverride": {
|
||||
"baseline": 8,
|
||||
|
|
@ -73,7 +73,7 @@
|
|||
},
|
||||
"reportMissingParameterType": {
|
||||
"baseline": 3933,
|
||||
"slack": 10
|
||||
"slack": 390
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"baseline": 10612,
|
||||
|
|
@ -97,7 +97,7 @@
|
|||
},
|
||||
"reportOptionalMemberAccess": {
|
||||
"baseline": 724,
|
||||
"slack": 10
|
||||
"slack": 72
|
||||
},
|
||||
"reportOptionalOperand": {
|
||||
"baseline": 3,
|
||||
|
|
@ -120,8 +120,8 @@
|
|||
"slack": 3
|
||||
},
|
||||
"reportReturnType": {
|
||||
"baseline": 118,
|
||||
"slack": 10
|
||||
"baseline": 126,
|
||||
"slack": 13
|
||||
},
|
||||
"reportTypedDictNotRequiredAccess": {
|
||||
"baseline": 20,
|
||||
|
|
@ -136,19 +136,19 @@
|
|||
"slack": 3000
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"baseline": 76,
|
||||
"baseline": 75,
|
||||
"slack": 10
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"baseline": 27322,
|
||||
"baseline": 27037,
|
||||
"slack": 2500
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"baseline": 13636,
|
||||
"baseline": 13612,
|
||||
"slack": 1000
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"baseline": 21776,
|
||||
"baseline": 21445,
|
||||
"slack": 2000
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
|
|
@ -156,7 +156,7 @@
|
|||
"slack": 10
|
||||
},
|
||||
"reportUnnecessaryComparison": {
|
||||
"baseline": 680,
|
||||
"baseline": 683,
|
||||
"slack": 10
|
||||
},
|
||||
"reportUnnecessaryContains": {
|
||||
|
|
@ -164,12 +164,12 @@
|
|||
"slack": 3
|
||||
},
|
||||
"reportUnnecessaryIsInstance": {
|
||||
"baseline": 807,
|
||||
"slack": 10
|
||||
"baseline": 808,
|
||||
"slack": 80
|
||||
},
|
||||
"reportUntypedBaseClass": {
|
||||
"baseline": 110,
|
||||
"slack": 3
|
||||
"slack": 11
|
||||
},
|
||||
"reportUntypedFunctionDecorator": {
|
||||
"baseline": 22,
|
||||
|
|
@ -185,10 +185,10 @@
|
|||
},
|
||||
"reportUnusedImport": {
|
||||
"baseline": 670,
|
||||
"slack": 10
|
||||
"slack": 50
|
||||
},
|
||||
"reportUnusedVariable": {
|
||||
"baseline": 865,
|
||||
"slack": 10
|
||||
"slack": 50
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ model_list:
|
|||
litellm_params:
|
||||
model: openai/fake
|
||||
api_key: fake-key
|
||||
api_base: https://exampleopenaiendpoint-production.up.railway.app/
|
||||
api_base: os.environ/FAKE_OPENAI_API_BASE
|
||||
|
||||
general_settings:
|
||||
alerting: ["slack"]
|
||||
|
|
@ -1,9 +1,9 @@
|
|||
# What is this?
|
||||
## This hook is used to check for LiteLLM managed files in the request body, and replace them with model-specific file id
|
||||
|
||||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Literal, Optional, Union, cast
|
||||
|
||||
from fastapi import HTTPException
|
||||
|
|
@ -1472,8 +1472,8 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
error_message += f" (showing {MAX_BATCHES_IN_ERROR} most recent): {', '.join(batch_statuses)}. "
|
||||
|
||||
error_message += (
|
||||
f"To delete this file before complete cost tracking, please delete or cancel the referencing batch(es) first. "
|
||||
f"Alternatively, wait for all batches to complete and for cost to be computed (batch_processed=true)."
|
||||
"To delete this file before complete cost tracking, please delete or cancel the referencing batch(es) first. "
|
||||
"Alternatively, wait for all batches to complete and for cost to be computed (batch_processed=true)."
|
||||
)
|
||||
|
||||
# Record blocked deletion metric
|
||||
|
|
@ -1550,9 +1550,22 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
|
||||
if specific_model_file_id_mapping:
|
||||
exception_dict = {}
|
||||
for model_id, file_id in specific_model_file_id_mapping.items():
|
||||
for model_id, provider_file_id in specific_model_file_id_mapping.items():
|
||||
try:
|
||||
return await llm_router.afile_content(model=model_id, file_id=file_id, **data) # type: ignore
|
||||
# Cloud-storage providers (e.g. Bedrock S3) validate file ids
|
||||
# against the deployment's configured bucket, which they only
|
||||
# trust from this immutable server-side snapshot, never from
|
||||
# request params.
|
||||
credentials = llm_router.get_deployment_credentials_with_provider(
|
||||
model_id=model_id
|
||||
)
|
||||
if credentials is not None:
|
||||
data["_litellm_internal_model_credentials"] = cast(
|
||||
Dict, MappingProxyType(dict(credentials))
|
||||
)
|
||||
else:
|
||||
data.pop("_litellm_internal_model_credentials", None)
|
||||
return await llm_router.afile_content(model=model_id, file_id=provider_file_id, **data) # type: ignore
|
||||
except Exception as e:
|
||||
exception_dict[model_id] = str(e)
|
||||
raise Exception(
|
||||
|
|
|
|||
|
|
@ -213,6 +213,15 @@ standard_logging_payload_excluded_fields: Optional[List[str]] = (
|
|||
log_raw_request_response: bool = False
|
||||
redact_messages_in_exceptions: Optional[bool] = False
|
||||
redact_user_api_key_info: Optional[bool] = False
|
||||
# When True (default — preserves historical behavior), the Router appends
|
||||
# internal config names (model_group, fallback model groups, deployment
|
||||
# timeouts, fallback failure details) onto exception messages and surfaces
|
||||
# them to clients via ProxyException.message. Set to False if you do NOT
|
||||
# want the proxy's internal model_name / fallback wiring visible to clients.
|
||||
# Deprecation: planned to flip to False (redact by default) in a future
|
||||
# major release; opt in early with `litellm.expose_router_debug_in_errors
|
||||
# = False`.
|
||||
expose_router_debug_in_errors: bool = True
|
||||
filter_invalid_headers: Optional[bool] = False
|
||||
add_user_information_to_llm_headers: Optional[bool] = (
|
||||
None # adds user_id, team_id, token hash (params from StandardLoggingMetadata) to request headers
|
||||
|
|
@ -235,6 +244,17 @@ modify_params = bool(os.getenv("LITELLM_MODIFY_PARAMS", False))
|
|||
use_chat_completions_url_for_anthropic_messages: bool = bool(
|
||||
os.getenv("LITELLM_USE_CHAT_COMPLETIONS_URL_FOR_ANTHROPIC_MESSAGES", False)
|
||||
) # When True, routes OpenAI /v1/messages requests to chat/completions instead of the Responses API
|
||||
# When True, strip the OpenAI-flavored `usage.total_tokens` field that
|
||||
# LiteLLM injects into non-streaming /v1/messages responses, bringing the
|
||||
# wire response into line with the Anthropic spec (matches the streaming
|
||||
# SSE path, which already omits total_tokens). Default False to preserve
|
||||
# backward compatibility for clients that read the LiteLLM-shaped
|
||||
# `usage.total_tokens` today. Planned to flip to True in a future major
|
||||
# release; opt in early via Python:
|
||||
# `litellm.strip_anthropic_total_tokens = True`
|
||||
# Or via `litellm_settings.strip_anthropic_total_tokens: true` in
|
||||
# config.yaml.
|
||||
strip_anthropic_total_tokens: bool = False
|
||||
route_all_chat_openai_to_responses: bool = (
|
||||
os.getenv("LITELLM_ROUTE_ALL_CHAT_OPENAI_TO_RESPONSES", "false").lower() == "true"
|
||||
) # When True, routes all OpenAI /chat/completions requests through the Responses API bridge
|
||||
|
|
@ -413,7 +433,7 @@ anthropic_beta_headers_url: str = os.getenv(
|
|||
"LITELLM_ANTHROPIC_BETA_HEADERS_URL",
|
||||
"https://raw.githubusercontent.com/BerriAI/litellm/main/litellm/anthropic_beta_headers_config.json",
|
||||
)
|
||||
suppress_debug_info = False
|
||||
suppress_debug_info: bool = False
|
||||
dynamodb_table_name: Optional[str] = None
|
||||
s3_callback_params: Optional[Dict] = None
|
||||
s3_audit_callback_params: Optional[Dict] = None
|
||||
|
|
|
|||
|
|
@ -100,6 +100,8 @@ class Cache:
|
|||
gcs_path: Optional[str] = None,
|
||||
redis_semantic_cache_embedding_model: str = "text-embedding-ada-002",
|
||||
redis_semantic_cache_index_name: Optional[str] = None,
|
||||
valkey_semantic_cache_embedding_model: str = "text-embedding-ada-002",
|
||||
valkey_semantic_cache_index_name: str | None = None,
|
||||
redis_flush_size: Optional[int] = None,
|
||||
redis_startup_nodes: Optional[List] = None,
|
||||
disk_cache_dir: Optional[str] = None,
|
||||
|
|
@ -208,6 +210,21 @@ class Cache:
|
|||
index_name=redis_semantic_cache_index_name,
|
||||
**kwargs,
|
||||
)
|
||||
elif type == LiteLLMCacheType.VALKEY_SEMANTIC:
|
||||
# Imported here, not at module top, so the optional redis dependency
|
||||
# is only required when this backend is actually selected.
|
||||
from .valkey_semantic_cache import ValkeySemanticCache
|
||||
|
||||
self.cache = ValkeySemanticCache(
|
||||
host=host,
|
||||
port=port,
|
||||
password=password,
|
||||
similarity_threshold=similarity_threshold,
|
||||
embedding_model=valkey_semantic_cache_embedding_model,
|
||||
index_name=valkey_semantic_cache_index_name,
|
||||
startup_nodes=redis_startup_nodes,
|
||||
**kwargs,
|
||||
)
|
||||
elif type == LiteLLMCacheType.QDRANT_SEMANTIC:
|
||||
self.cache = QdrantSemanticCache(
|
||||
qdrant_api_base=qdrant_api_base,
|
||||
|
|
@ -267,12 +284,50 @@ class Cache:
|
|||
if (
|
||||
self.type == LiteLLMCacheType.REDIS
|
||||
or self.type == LiteLLMCacheType.REDIS_SEMANTIC
|
||||
or self.type == LiteLLMCacheType.VALKEY_SEMANTIC
|
||||
) and default_in_redis_ttl is not None:
|
||||
self.ttl = default_in_redis_ttl
|
||||
|
||||
if self.namespace is not None and isinstance(self.cache, RedisCache):
|
||||
self.cache.namespace = self.namespace
|
||||
|
||||
# Params whose values carry prompt content. Excluded from semantic-cache
|
||||
# scope keys so differently worded prompts share a bucket and match via
|
||||
# vector similarity rather than being split into per-wording buckets.
|
||||
_SEMANTIC_CACHE_SCOPE_EXCLUDED_PARAMS: frozenset = frozenset(
|
||||
{"messages", "prompt", "input"}
|
||||
)
|
||||
|
||||
# Server-set identity (from proxy auth) used to isolate semantic-cache
|
||||
# buckets per tenant. Required once the prompt is out of the scope key, so a
|
||||
# similar prompt from another key/team/org stays in a separate bucket.
|
||||
_SEMANTIC_CACHE_TENANT_SCOPE_FIELDS: tuple[str, ...] = (
|
||||
"user_api_key",
|
||||
"user_api_key_team_id",
|
||||
"user_api_key_org_id",
|
||||
)
|
||||
|
||||
def _is_semantic_cache(self) -> bool:
|
||||
return self.type in (
|
||||
LiteLLMCacheType.REDIS_SEMANTIC,
|
||||
LiteLLMCacheType.QDRANT_SEMANTIC,
|
||||
LiteLLMCacheType.VALKEY_SEMANTIC,
|
||||
)
|
||||
|
||||
def _get_semantic_cache_tenant_scope(self, kwargs: dict) -> str:
|
||||
metadata: dict = kwargs.get("metadata") or {}
|
||||
litellm_params: dict = kwargs.get("litellm_params") or {}
|
||||
metadata_in_litellm_params: dict = litellm_params.get("metadata") or {}
|
||||
|
||||
scope = ""
|
||||
for field in self._SEMANTIC_CACHE_TENANT_SCOPE_FIELDS:
|
||||
value = metadata.get(field)
|
||||
if value is None:
|
||||
value = metadata_in_litellm_params.get(field)
|
||||
if value is not None:
|
||||
scope += f"{field}: {value}"
|
||||
return scope
|
||||
|
||||
def get_cache_key(self, **kwargs) -> str:
|
||||
"""
|
||||
Get the cache key for the given arguments.
|
||||
|
|
@ -293,7 +348,15 @@ class Cache:
|
|||
|
||||
combined_kwargs = ModelParamHelper._get_all_llm_api_params()
|
||||
litellm_param_kwargs = all_litellm_params
|
||||
is_semantic_cache = self._is_semantic_cache()
|
||||
scope_excluded_params = (
|
||||
self._SEMANTIC_CACHE_SCOPE_EXCLUDED_PARAMS
|
||||
if is_semantic_cache
|
||||
else frozenset()
|
||||
)
|
||||
for param in kwargs:
|
||||
if param in scope_excluded_params:
|
||||
continue
|
||||
if param in combined_kwargs:
|
||||
param_value: Optional[str] = self._get_param_value(param, kwargs)
|
||||
if param_value is not None:
|
||||
|
|
@ -309,6 +372,9 @@ class Cache:
|
|||
param_value = kwargs[param]
|
||||
cache_key += f"{str(param)}: {str(param_value)}"
|
||||
|
||||
if is_semantic_cache:
|
||||
cache_key += self._get_semantic_cache_tenant_scope(kwargs)
|
||||
|
||||
hashed_cache_key = Cache._get_hashed_cache_key(cache_key)
|
||||
hashed_cache_key = self._add_namespace_to_cache_key(hashed_cache_key, **kwargs)
|
||||
verbose_logger.debug(
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ Supports syncing responses to Google Cloud Storage Buckets using HTTP requests.
|
|||
import json
|
||||
import asyncio
|
||||
from typing import Optional
|
||||
from urllib.parse import quote
|
||||
|
||||
from litellm._logging import print_verbose, verbose_logger
|
||||
from litellm.integrations.gcs_bucket.gcs_bucket_base import GCSBucketBase
|
||||
|
|
@ -48,7 +49,7 @@ class GCSCache(BaseCache):
|
|||
headers = self._construct_headers()
|
||||
object_name = self.key_prefix + key
|
||||
bucket_name = self.bucket_name
|
||||
url = f"https://storage.googleapis.com/upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={object_name}"
|
||||
url = f"https://storage.googleapis.com/upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={quote(object_name, safe='')}"
|
||||
data = json.dumps(value)
|
||||
self.sync_client.post(url=url, data=data, headers=headers)
|
||||
except Exception as e:
|
||||
|
|
@ -59,7 +60,7 @@ class GCSCache(BaseCache):
|
|||
headers = self._construct_headers()
|
||||
object_name = self.key_prefix + key
|
||||
bucket_name = self.bucket_name
|
||||
url = f"https://storage.googleapis.com/upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={object_name}"
|
||||
url = f"https://storage.googleapis.com/upload/storage/v1/b/{bucket_name}/o?uploadType=media&name={quote(object_name, safe='')}"
|
||||
data = json.dumps(value)
|
||||
await self.async_client.post(url=url, data=data, headers=headers)
|
||||
except Exception as e:
|
||||
|
|
@ -72,7 +73,7 @@ class GCSCache(BaseCache):
|
|||
headers = self._construct_headers()
|
||||
object_name = self.key_prefix + key
|
||||
bucket_name = self.bucket_name
|
||||
url = f"https://storage.googleapis.com/storage/v1/b/{bucket_name}/o/{object_name}?alt=media"
|
||||
url = f"https://storage.googleapis.com/storage/v1/b/{bucket_name}/o/{quote(object_name, safe='')}?alt=media"
|
||||
response = self.sync_client.get(url=url, headers=headers)
|
||||
if response.status_code == 200:
|
||||
cached_response = json.loads(response.text)
|
||||
|
|
@ -91,7 +92,7 @@ class GCSCache(BaseCache):
|
|||
headers = self._construct_headers()
|
||||
object_name = self.key_prefix + key
|
||||
bucket_name = self.bucket_name
|
||||
url = f"https://storage.googleapis.com/storage/v1/b/{bucket_name}/o/{object_name}?alt=media"
|
||||
url = f"https://storage.googleapis.com/storage/v1/b/{bucket_name}/o/{quote(object_name, safe='')}?alt=media"
|
||||
response = await self.async_client.get(url=url, headers=headers)
|
||||
if response.status_code == 200:
|
||||
return json.loads(response.text)
|
||||
|
|
|
|||
|
|
@ -903,6 +903,43 @@ class RedisCache(BaseCache):
|
|||
)
|
||||
raise e
|
||||
|
||||
@_redis_circuit_breaker_guard
|
||||
async def async_set_max(
|
||||
self,
|
||||
key: str,
|
||||
value: float,
|
||||
ttl: int | None = None,
|
||||
) -> float | None:
|
||||
"""Atomically set ``key`` to ``value`` only when ``value`` is greater
|
||||
than the stored value (or the key is unset), refreshing the TTL.
|
||||
|
||||
Monotonic by construction: it never lowers the stored value, so a repair
|
||||
that writes an authoritative-but-slightly-stale total cannot clobber a
|
||||
concurrent increment that has already pushed the counter higher. The
|
||||
GET/compare/SET runs in a single Lua call, so it is also atomic across
|
||||
racing callers and pods. Returns the resulting value.
|
||||
"""
|
||||
_redis_client = self.init_async_client()
|
||||
_used_ttl = self.get_ttl(ttl=ttl)
|
||||
key = self.check_and_fix_namespace(key=key)
|
||||
lua = (
|
||||
"local cur = redis.call('GET', KEYS[1]) "
|
||||
"if cur == false or tonumber(cur) < tonumber(ARGV[1]) then "
|
||||
"redis.call('SET', KEYS[1], ARGV[1]) "
|
||||
"if tonumber(ARGV[2]) > 0 then redis.call('EXPIRE', KEYS[1], ARGV[2]) end "
|
||||
"return ARGV[1] end "
|
||||
"return cur"
|
||||
)
|
||||
result = cast(
|
||||
"str | bytes | int | float | None",
|
||||
await _redis_client.eval(lua, 1, key, str(value), str(int(_used_ttl or 0))),
|
||||
)
|
||||
if result is None:
|
||||
return None
|
||||
if isinstance(result, bytes):
|
||||
result = result.decode()
|
||||
return float(result)
|
||||
|
||||
async def flush_cache_buffer(self):
|
||||
print_verbose(
|
||||
f"flushing to redis....reached size of buffer {len(self.redis_batch_writing_buffer)}"
|
||||
|
|
|
|||
353
litellm/caching/valkey_semantic_cache.py
Normal file
353
litellm/caching/valkey_semantic_cache.py
Normal file
|
|
@ -0,0 +1,353 @@
|
|||
"""
|
||||
Valkey Semantic Cache implementation for LiteLLM
|
||||
|
||||
Backs semantic caching with Valkey (for example AWS ElastiCache for Valkey)
|
||||
running the valkey-search module.
|
||||
|
||||
RedisVL cannot drive valkey-search: it gates on a RediSearch module version
|
||||
that valkey-search does not report, and its SemanticCache index uses a TEXT
|
||||
field that valkey-search does not implement. This backend therefore talks to
|
||||
valkey-search directly over redis-py, building a vector index from the field
|
||||
types valkey-search does support (TAG for cache-key isolation and VECTOR for
|
||||
the prompt embedding) and running KNN queries for retrieval. Prompt extraction,
|
||||
embedding generation, and cached-response parsing are reused from
|
||||
RedisSemanticCache since those are backend agnostic.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import hashlib
|
||||
import os
|
||||
import struct
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
from redis import Redis
|
||||
from redis.asyncio import Redis as AsyncRedis
|
||||
from redis.commands.search.field import TagField, VectorField
|
||||
from redis.commands.search.indexDefinition import IndexDefinition, IndexType
|
||||
from redis.commands.search.query import Query
|
||||
|
||||
from litellm._logging import print_verbose
|
||||
from litellm._uuid import uuid
|
||||
|
||||
from .redis_semantic_cache import RedisSemanticCache
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _ValkeyCacheHit:
|
||||
response: str
|
||||
distance: float
|
||||
|
||||
|
||||
class ValkeySemanticCache(RedisSemanticCache):
|
||||
"""Valkey-backed semantic cache for LLM responses."""
|
||||
|
||||
DEFAULT_VALKEY_INDEX_NAME: str = "litellm_semantic_cache_index"
|
||||
EMBEDDING_FIELD_NAME: str = "embedding"
|
||||
PROMPT_FIELD_NAME: str = "prompt"
|
||||
RESPONSE_FIELD_NAME: str = "response"
|
||||
DISTANCE_FIELD_NAME: str = "vector_distance"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
host: str | None = None,
|
||||
port: str | None = None,
|
||||
password: str | None = None,
|
||||
redis_url: str | None = None,
|
||||
similarity_threshold: float | None = None,
|
||||
embedding_model: str = "text-embedding-ada-002",
|
||||
index_name: str | None = None,
|
||||
ssl: bool = False,
|
||||
startup_nodes: list | None = None,
|
||||
sync_client: Redis | None = None,
|
||||
async_client: AsyncRedis | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
if similarity_threshold is None:
|
||||
raise ValueError("similarity_threshold must be provided, passed None")
|
||||
|
||||
if startup_nodes:
|
||||
raise ValueError(
|
||||
"valkey-semantic does not support cluster-mode-enabled (multi-shard) "
|
||||
"endpoints. The async cluster client cannot route the FT.* search "
|
||||
"commands reliably. Point it at a cluster-mode-disabled endpoint "
|
||||
"instead (a primary with replicas is fine; only horizontal sharding "
|
||||
"is unsupported), or pass a single redis_url. On AWS, vector search "
|
||||
"needs ElastiCache for Valkey 8.2+ on a node-based cluster."
|
||||
)
|
||||
|
||||
self.similarity_threshold = similarity_threshold
|
||||
self.embedding_model = embedding_model
|
||||
self.index_name = index_name or self.DEFAULT_VALKEY_INDEX_NAME
|
||||
self.key_prefix = f"{self.index_name}:"
|
||||
self._index_dim: int | None = None
|
||||
|
||||
resolved_url = None
|
||||
if sync_client is None or async_client is None:
|
||||
resolved_url = redis_url or self._build_valkey_url(
|
||||
host, port, password, ssl
|
||||
)
|
||||
self.sync_client = (
|
||||
sync_client if sync_client is not None else Redis.from_url(resolved_url) # type: ignore[arg-type]
|
||||
)
|
||||
self.async_client = (
|
||||
async_client
|
||||
if async_client is not None
|
||||
else AsyncRedis.from_url(resolved_url) # type: ignore[arg-type]
|
||||
)
|
||||
|
||||
print_verbose(f"Valkey semantic-cache initializing index - {self.index_name}")
|
||||
|
||||
@staticmethod
|
||||
def _build_valkey_url(
|
||||
host: str | None, port: str | None, password: str | None, ssl: bool = False
|
||||
) -> str:
|
||||
host = host or os.environ.get("VALKEY_HOST") or os.environ.get("REDIS_HOST")
|
||||
port = port or os.environ.get("VALKEY_PORT") or os.environ.get("REDIS_PORT")
|
||||
password = (
|
||||
password
|
||||
or os.environ.get("VALKEY_PASSWORD")
|
||||
or os.environ.get("REDIS_PASSWORD")
|
||||
)
|
||||
|
||||
if not host or not port:
|
||||
raise ValueError(
|
||||
"Missing required Valkey configuration. Provide host and port "
|
||||
"(or VALKEY_HOST/VALKEY_PORT), or pass redis_url."
|
||||
)
|
||||
|
||||
credentials = f":{password}@" if password else ""
|
||||
scheme = "rediss" if ssl else "redis"
|
||||
return f"{scheme}://{credentials}{host}:{port}"
|
||||
|
||||
@classmethod
|
||||
def _scope_tag(cls, key: str) -> str:
|
||||
# valkey-search TAG fields tokenize on punctuation and do not honour
|
||||
# backslash escaping, so an arbitrary cache key cannot be matched
|
||||
# verbatim. Hashing to hex yields a token that is always exact-match
|
||||
# safe and still uniquely isolates a caller's scope.
|
||||
return hashlib.sha256(str(key).encode("utf-8")).hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def _embedding_to_bytes(embedding: list[float]) -> bytes:
|
||||
return struct.pack(f"<{len(embedding)}f", *embedding)
|
||||
|
||||
def _index_schema(self, dim: int) -> tuple[TagField, VectorField]:
|
||||
return (
|
||||
TagField(self.CACHE_KEY_FIELD_NAME),
|
||||
VectorField(
|
||||
self.EMBEDDING_FIELD_NAME,
|
||||
"HNSW",
|
||||
{"TYPE": "FLOAT32", "DIM": dim, "DISTANCE_METRIC": "COSINE"},
|
||||
),
|
||||
)
|
||||
|
||||
def _index_definition(self) -> IndexDefinition:
|
||||
return IndexDefinition(prefix=[self.key_prefix], index_type=IndexType.HASH)
|
||||
|
||||
@staticmethod
|
||||
def _is_index_exists_error(exc: Exception) -> bool:
|
||||
return "already exists" in str(exc).lower()
|
||||
|
||||
@staticmethod
|
||||
def _extract_index_dim(info: dict) -> int | None:
|
||||
# FT.INFO nests the vector field's "dimensions" one level inside its
|
||||
# "index" block, so flatten each field descriptor a single level and
|
||||
# scan for the dimensions marker.
|
||||
for field in info.get("attributes") or []:
|
||||
if not isinstance(field, (list, tuple)):
|
||||
continue
|
||||
flat = [
|
||||
sub
|
||||
for item in field
|
||||
for sub in (item if isinstance(item, (list, tuple)) else [item])
|
||||
]
|
||||
for i, marker in enumerate(flat):
|
||||
if marker in (b"dimensions", "dimensions") and i + 1 < len(flat):
|
||||
return int(flat[i + 1])
|
||||
return None
|
||||
|
||||
def _assert_dim_matches(self, info: dict, dim: int) -> None:
|
||||
existing_dim = self._extract_index_dim(info)
|
||||
if existing_dim is not None and existing_dim != dim:
|
||||
raise ValueError(
|
||||
f"Valkey semantic-cache index '{self.index_name}' already exists with "
|
||||
f"embedding dimension {existing_dim}, but the configured embedding "
|
||||
f"model produced dimension {dim}. Use a different "
|
||||
f"valkey_semantic_cache_index_name or drop the existing index."
|
||||
)
|
||||
|
||||
def _ensure_index_sync(self, dim: int) -> None:
|
||||
if self._index_dim == dim:
|
||||
return
|
||||
try:
|
||||
self.sync_client.ft(self.index_name).create_index(
|
||||
self._index_schema(dim), definition=self._index_definition()
|
||||
)
|
||||
except Exception as exc:
|
||||
if not self._is_index_exists_error(exc):
|
||||
raise
|
||||
self._assert_dim_matches(self.sync_client.ft(self.index_name).info(), dim)
|
||||
self._index_dim = dim
|
||||
|
||||
async def _ensure_index_async(self, dim: int) -> None:
|
||||
if self._index_dim == dim:
|
||||
return
|
||||
try:
|
||||
await self.async_client.ft(self.index_name).create_index(
|
||||
self._index_schema(dim), definition=self._index_definition()
|
||||
)
|
||||
except Exception as exc:
|
||||
if not self._is_index_exists_error(exc):
|
||||
raise
|
||||
info = await self.async_client.ft(self.index_name).info()
|
||||
self._assert_dim_matches(info, dim)
|
||||
self._index_dim = dim
|
||||
|
||||
def _doc_key(self, key: str) -> str:
|
||||
return f"{self.key_prefix}{self._scope_tag(key)}:{uuid.uuid4()}"
|
||||
|
||||
def _doc_mapping(
|
||||
self, key: str, prompt: str, value_str: str, embedding: list[float]
|
||||
) -> dict:
|
||||
return {
|
||||
self.CACHE_KEY_FIELD_NAME: self._scope_tag(key),
|
||||
self.PROMPT_FIELD_NAME: prompt,
|
||||
self.RESPONSE_FIELD_NAME: value_str,
|
||||
self.EMBEDDING_FIELD_NAME: self._embedding_to_bytes(embedding),
|
||||
}
|
||||
|
||||
def _knn_query(self, key: str) -> Query:
|
||||
scope = self._scope_tag(key)
|
||||
query_string = (
|
||||
f"(@{self.CACHE_KEY_FIELD_NAME}:{{{scope}}})"
|
||||
f"=>[KNN 1 @{self.EMBEDDING_FIELD_NAME} $vec AS {self.DISTANCE_FIELD_NAME}]"
|
||||
)
|
||||
return (
|
||||
Query(query_string)
|
||||
.return_fields(self.RESPONSE_FIELD_NAME, self.DISTANCE_FIELD_NAME)
|
||||
.dialect(2)
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def _first_hit(cls, search_result: Any) -> _ValkeyCacheHit | None:
|
||||
docs = getattr(search_result, "docs", [])
|
||||
if not docs:
|
||||
return None
|
||||
doc = docs[0]
|
||||
return _ValkeyCacheHit(
|
||||
response=str(getattr(doc, cls.RESPONSE_FIELD_NAME)),
|
||||
distance=float(getattr(doc, cls.DISTANCE_FIELD_NAME)),
|
||||
)
|
||||
|
||||
def _resolve_hit(self, hit: _ValkeyCacheHit | None, key: str, **kwargs: Any) -> Any:
|
||||
if hit is None:
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
return None
|
||||
|
||||
similarity = 1 - hit.distance
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = similarity
|
||||
|
||||
if similarity < self.similarity_threshold:
|
||||
return None
|
||||
return self._get_cache_logic(cached_response=hit.response)
|
||||
|
||||
def set_cache(self, key: str, value: Any, **kwargs: Any) -> None:
|
||||
print_verbose(f"Valkey semantic-cache set_cache, kwargs: {kwargs}")
|
||||
try:
|
||||
prompt = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
print_verbose("No prompt provided for semantic caching")
|
||||
return
|
||||
|
||||
embedding = self._get_embedding(prompt)
|
||||
self._ensure_index_sync(len(embedding))
|
||||
|
||||
doc_key = self._doc_key(key)
|
||||
self.sync_client.hset(
|
||||
doc_key, mapping=self._doc_mapping(key, prompt, str(value), embedding)
|
||||
)
|
||||
ttl = self._get_ttl(**kwargs)
|
||||
if ttl is not None:
|
||||
self.sync_client.expire(doc_key, ttl)
|
||||
except Exception as e:
|
||||
print_verbose(f"Error in Valkey semantic-cache set_cache: {str(e)}")
|
||||
|
||||
def get_cache(self, key: str, **kwargs: Any) -> Any:
|
||||
print_verbose(f"Valkey semantic-cache get_cache, kwargs: {kwargs}")
|
||||
try:
|
||||
prompt = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
return None
|
||||
|
||||
embedding = self._get_embedding(prompt)
|
||||
self._ensure_index_sync(len(embedding))
|
||||
|
||||
search_result = self.sync_client.ft(self.index_name).search(
|
||||
self._knn_query(key),
|
||||
query_params={"vec": self._embedding_to_bytes(embedding)},
|
||||
)
|
||||
return self._resolve_hit(self._first_hit(search_result), key, **kwargs)
|
||||
except Exception as e:
|
||||
print_verbose(f"Error in Valkey semantic-cache get_cache: {str(e)}")
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
|
||||
async def async_set_cache(self, key: str, value: Any, **kwargs: Any) -> None:
|
||||
print_verbose(f"Async Valkey semantic-cache set_cache, kwargs: {kwargs}")
|
||||
try:
|
||||
prompt = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
print_verbose("No prompt provided for semantic caching")
|
||||
return
|
||||
|
||||
embedding = await self._get_async_embedding(prompt, **kwargs)
|
||||
await self._ensure_index_async(len(embedding))
|
||||
|
||||
doc_key = self._doc_key(key)
|
||||
await self.async_client.hset(
|
||||
doc_key, mapping=self._doc_mapping(key, prompt, str(value), embedding)
|
||||
)
|
||||
ttl = self._get_ttl(**kwargs)
|
||||
if ttl is not None:
|
||||
await self.async_client.expire(doc_key, ttl)
|
||||
except Exception as e:
|
||||
print_verbose(f"Error in async Valkey semantic-cache set_cache: {str(e)}")
|
||||
|
||||
async def async_get_cache(self, key: str, **kwargs: Any) -> Any:
|
||||
print_verbose(f"Async Valkey semantic-cache get_cache, kwargs: {kwargs}")
|
||||
try:
|
||||
prompt = self._get_prompt_from_kwargs(**kwargs)
|
||||
if prompt is None:
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
return None
|
||||
|
||||
embedding = await self._get_async_embedding(prompt, **kwargs)
|
||||
await self._ensure_index_async(len(embedding))
|
||||
|
||||
search_result = await self.async_client.ft(self.index_name).search(
|
||||
self._knn_query(key),
|
||||
query_params={"vec": self._embedding_to_bytes(embedding)},
|
||||
)
|
||||
return self._resolve_hit(self._first_hit(search_result), key, **kwargs)
|
||||
except Exception as e:
|
||||
print_verbose(f"Error in async Valkey semantic-cache get_cache: {str(e)}")
|
||||
kwargs.setdefault("metadata", {})["semantic-similarity"] = 0.0
|
||||
|
||||
async def async_set_cache_pipeline(
|
||||
self, cache_list: list[tuple[str, Any]], **kwargs: Any
|
||||
) -> None:
|
||||
try:
|
||||
await asyncio.gather(
|
||||
*[
|
||||
self.async_set_cache(key, value, **kwargs)
|
||||
for key, value in cache_list
|
||||
]
|
||||
)
|
||||
except Exception as e:
|
||||
print_verbose(
|
||||
f"Error in Valkey semantic-cache async_set_cache_pipeline: {str(e)}"
|
||||
)
|
||||
|
||||
async def _index_info(self) -> dict:
|
||||
return await self.async_client.ft(self.index_name).info()
|
||||
|
|
@ -802,6 +802,7 @@ openai_compatible_endpoints: List = [
|
|||
"https://api.inference.wandb.ai/v1",
|
||||
"https://api.clarifai.com/v2/ext/openai/v1",
|
||||
"https://api.libertai.io/v1",
|
||||
"https://pinstripes.io/v1",
|
||||
]
|
||||
|
||||
|
||||
|
|
@ -865,6 +866,7 @@ openai_compatible_providers: List = [
|
|||
"clarifai",
|
||||
"docker_model_runner",
|
||||
"ragflow",
|
||||
"pinstripes", # Pinstripes - JSON-configured provider
|
||||
]
|
||||
openai_text_completion_compatible_providers: List = (
|
||||
[ # providers that support `/v1/completions`
|
||||
|
|
|
|||
|
|
@ -888,6 +888,23 @@ def _map_traffic_type_to_service_tier(traffic_type: Optional[str]) -> Optional[s
|
|||
return service_tier
|
||||
|
||||
|
||||
def _normalize_service_tier(service_tier: object) -> str | None:
|
||||
"""
|
||||
Reduce a service_tier value to a concrete billable tier string or None.
|
||||
|
||||
"auto" is a routing preference and any non-string value is not a billable
|
||||
tier, so both defer to standard pricing (or to the tier the provider reports
|
||||
on the response usage) instead of crashing the downstream cost-key lookup,
|
||||
which calls service_tier.lower()
|
||||
"""
|
||||
if (
|
||||
not isinstance(service_tier, str)
|
||||
or service_tier.lower() == ServiceTier.AUTO.value
|
||||
):
|
||||
return None
|
||||
return service_tier
|
||||
|
||||
|
||||
def _get_usage_object(
|
||||
completion_response: Any,
|
||||
) -> Optional[Usage]:
|
||||
|
|
@ -1227,15 +1244,7 @@ def completion_cost(
|
|||
if service_tier is None and optional_params is not None:
|
||||
service_tier = optional_params.get("service_tier")
|
||||
|
||||
# A request-level service_tier only prices the request when it is a
|
||||
# concrete billable tier string. "auto" is a routing preference and any
|
||||
# non-string value is not a billable tier, so defer to the tier the
|
||||
# provider reports on the response/usage instead of crashing or mispricing
|
||||
if (
|
||||
not isinstance(service_tier, str)
|
||||
or service_tier.lower() == ServiceTier.AUTO.value
|
||||
):
|
||||
service_tier = None
|
||||
service_tier = _normalize_service_tier(service_tier)
|
||||
|
||||
# Extract service_tier from completion_response if not provided
|
||||
if service_tier is None and completion_response is not None:
|
||||
|
|
@ -1244,6 +1253,8 @@ def completion_cost(
|
|||
elif isinstance(completion_response, dict):
|
||||
service_tier = completion_response.get("service_tier")
|
||||
|
||||
service_tier = _normalize_service_tier(service_tier)
|
||||
|
||||
# Extract service_tier from usage object if not provided
|
||||
if service_tier is None and cost_per_token_usage_object is not None:
|
||||
if isinstance(cost_per_token_usage_object, BaseModel):
|
||||
|
|
@ -1253,6 +1264,8 @@ def completion_cost(
|
|||
elif isinstance(cost_per_token_usage_object, dict):
|
||||
service_tier = cost_per_token_usage_object.get("service_tier")
|
||||
|
||||
service_tier = _normalize_service_tier(service_tier)
|
||||
|
||||
selected_model = _select_model_name_for_cost_calc(
|
||||
model=model,
|
||||
completion_response=completion_response,
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@
|
|||
from collections import OrderedDict
|
||||
from contextlib import contextmanager
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, Iterator, Mapping, cast
|
||||
from typing import TYPE_CHECKING, Any, Callable, Iterator, Mapping, Sequence, cast
|
||||
|
||||
from opentelemetry.context import attach, get_current
|
||||
from opentelemetry.sdk.trace import TracerProvider
|
||||
|
|
@ -546,6 +546,58 @@ class OpenTelemetryV2(CustomLogger):
|
|||
return span
|
||||
|
||||
|
||||
def select_global_otel_v2_logger(
|
||||
in_memory_loggers: Sequence[object],
|
||||
registered: "OpenTelemetryV2 | None" = None,
|
||||
) -> "OpenTelemetryV2":
|
||||
"""The single ``OpenTelemetryV2`` whose provider should become the OTel global.
|
||||
|
||||
The callback factory designates one logger as canonical the moment it builds
|
||||
the first one (``_init_otel_logger_on_litellm_proxy`` sets
|
||||
``proxy_server.open_telemetry_logger``), and every other v2 entry point —
|
||||
guardrail, identity seeding, phase spans — already routes through that same
|
||||
``registered`` owner. Reuse it here too so the global provider has one source
|
||||
of truth instead of a second, independently-derived guess; this is the logger
|
||||
a preset (arize, langfuse, …) folds the ``OTEL_*`` base exporter and its own
|
||||
exporter into, so the FastAPI server span and the gen-ai spans share one
|
||||
provider and one trace.
|
||||
|
||||
Fall back to ``in_memory_loggers`` for the SDK path, where no proxy global is
|
||||
set (selecting from there, not ``service_callback``, which a preset logger does
|
||||
not always reach), and build a generic logger from ``OTEL_*`` only when none was
|
||||
configured at all. Each fallback still avoids the second generic logger that
|
||||
orphaned the gen-ai spans onto a different backend than the server span.
|
||||
"""
|
||||
if registered is not None:
|
||||
return registered
|
||||
existing = next(
|
||||
(cb for cb in in_memory_loggers if isinstance(cb, OpenTelemetryV2)), None
|
||||
)
|
||||
return existing if existing is not None else OpenTelemetryV2()
|
||||
|
||||
|
||||
def publish_global_otel_v2_provider(
|
||||
in_memory_loggers: Sequence[object],
|
||||
set_global_provider: Callable[[TracerProvider], None],
|
||||
registered: "OpenTelemetryV2 | None" = None,
|
||||
) -> "OpenTelemetryV2":
|
||||
"""Select the single v2 logger and publish its provider as the OTel global.
|
||||
|
||||
The proxy calls this once at startup, after callbacks are initialized, so the
|
||||
preset logger already exists; it passes ``registered`` (the canonical owner the
|
||||
factory designated as ``proxy_server.open_telemetry_logger``) so the global
|
||||
provider reuses the same logger the rest of the v2 code emits through (see
|
||||
:func:`select_global_otel_v2_logger`). Both ``registered`` and
|
||||
``set_global_provider`` (the proxy passes
|
||||
``opentelemetry.trace.set_tracer_provider``) are injected so the publish step is
|
||||
unit-testable without reading or mutating real global OTel state. Returns the
|
||||
logger whose provider was published.
|
||||
"""
|
||||
logger = select_global_otel_v2_logger(in_memory_loggers, registered=registered)
|
||||
set_global_provider(logger._tracer_provider)
|
||||
return logger
|
||||
|
||||
|
||||
def _registered_v2_logger() -> "OpenTelemetryV2 | None":
|
||||
try:
|
||||
from litellm.proxy import proxy_server
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""Typed configuration for the OpenTelemetry instrumentation."""
|
||||
|
||||
from enum import Enum
|
||||
from typing import Any, List
|
||||
|
||||
from pydantic import AliasChoices, BaseModel, Field, field_validator, model_validator
|
||||
|
|
@ -23,6 +24,23 @@ class CaptureMessageContent(str):
|
|||
SPAN_AND_EVENT = "span_and_event"
|
||||
|
||||
|
||||
class ExporterOwner(str, Enum):
|
||||
"""The preset that contributed an exporter. Values match the callback names
|
||||
in ``presets.PRESET_BY_CALLBACK`` so per-request dynamic-credential routing
|
||||
can match an exporter's owner against the credential source's callback name.
|
||||
A ``str`` enum so the value compares equal to the bare callback-name string."""
|
||||
|
||||
# Arize AX (the hosted platform) and Arize Phoenix (the open-source / Phoenix
|
||||
# Cloud tracer) are distinct backends with separate config and auth, so they
|
||||
# are separate owners. The member value stays the public callback name.
|
||||
ARIZE_AX = "arize"
|
||||
ARIZE_PHOENIX = "arize_phoenix"
|
||||
LANGFUSE_OTEL = "langfuse_otel"
|
||||
WEAVE_OTEL = "weave_otel"
|
||||
LEVO = "levo"
|
||||
AGENTOPS = "agentops"
|
||||
|
||||
|
||||
class _OTelV2Flag(BaseSettings):
|
||||
model_config = SettingsConfigDict(extra="ignore")
|
||||
|
||||
|
|
@ -49,6 +67,15 @@ class ExporterSpec(BaseModel):
|
|||
)
|
||||
endpoint: str | None = None
|
||||
headers: str | None = None
|
||||
owner: ExporterOwner | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"The preset that contributed this exporter. Per-request dynamic OTLP "
|
||||
"credentials are applied only to the exporter whose owner matches the "
|
||||
"credential source, so one tenant's vendor key never lands on a "
|
||||
"different backend's exporter."
|
||||
),
|
||||
)
|
||||
options: dict[str, str] | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
|
|
|
|||
|
|
@ -88,13 +88,23 @@ class TenantTracerCache:
|
|||
return get_tracer(provider, self._tracer_name)
|
||||
|
||||
def _config_with_headers(self, headers: Mapping[str, str]) -> OpenTelemetryV2Config:
|
||||
"""Clone the config, replacing OTLP exporter headers with ``headers``."""
|
||||
"""Clone the config, stamping ``headers`` onto the credential's own exporter.
|
||||
|
||||
``headers`` are the per-request credentials of ``self._callback_name`` (the
|
||||
integration that built this cache), so they apply only to the exporter that
|
||||
integration contributed (``spec.owner``). A request that carries one
|
||||
tenant's Arize key must never rewrite the headers of a co-configured
|
||||
Langfuse or self-hosted collector exporter, which would leak that key to a
|
||||
different backend.
|
||||
"""
|
||||
header_str = ",".join(f"{key}={value}" for key, value in headers.items())
|
||||
header_update: dict[str, str] = {"headers": header_str}
|
||||
exporters = [
|
||||
(
|
||||
spec
|
||||
if spec.kind.lower() in _NON_OTLP_KINDS
|
||||
else spec.model_copy(update={"headers": header_str})
|
||||
spec.model_copy(update=header_update)
|
||||
if spec.owner == self._callback_name
|
||||
and spec.kind.lower() not in _NON_OTLP_KINDS
|
||||
else spec
|
||||
)
|
||||
for spec in self._config.exporters
|
||||
]
|
||||
|
|
|
|||
|
|
@ -16,7 +16,11 @@ from pydantic import Field
|
|||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
|
||||
from litellm.integrations.otel.model.config import (
|
||||
ExporterOwner,
|
||||
ExporterSpec,
|
||||
OpenTelemetryV2Config,
|
||||
)
|
||||
from litellm.integrations.otel.plumbing.providers import register_exporter_factory
|
||||
|
||||
_AGENTOPS_ENDPOINT = "https://otlp.agentops.cloud/v1/traces"
|
||||
|
|
@ -59,6 +63,7 @@ def agentops_preset(
|
|||
options=(
|
||||
{"api_key": settings.api_key} if settings.api_key else None
|
||||
),
|
||||
owner=ExporterOwner.AGENTOPS,
|
||||
),
|
||||
],
|
||||
"resource_attributes": {
|
||||
|
|
|
|||
|
|
@ -4,7 +4,11 @@ from pydantic import Field
|
|||
from pydantic_settings import BaseSettings, SettingsConfigDict
|
||||
|
||||
from litellm.integrations.arize.arize import ArizeLogger as _V1ArizeLogger
|
||||
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
|
||||
from litellm.integrations.otel.model.config import (
|
||||
ExporterOwner,
|
||||
ExporterSpec,
|
||||
OpenTelemetryV2Config,
|
||||
)
|
||||
from litellm.integrations.otel.presets.utils import ensure_mappers
|
||||
from litellm.types.utils import StandardCallbackDynamicParams
|
||||
|
||||
|
|
@ -34,6 +38,7 @@ def arize_preset(
|
|||
kind=arize_cfg.protocol or "otlp_grpc",
|
||||
endpoint=arize_cfg.endpoint or "https://otlp.arize.com/v1",
|
||||
headers=headers,
|
||||
owner=ExporterOwner.ARIZE_AX,
|
||||
),
|
||||
],
|
||||
"mapper_names": ensure_mappers(base.mapper_names, "openinference"),
|
||||
|
|
|
|||
|
|
@ -3,7 +3,11 @@
|
|||
from litellm.integrations.langfuse.langfuse_otel import (
|
||||
LangfuseOtelLogger as _V1Langfuse,
|
||||
)
|
||||
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
|
||||
from litellm.integrations.otel.model.config import (
|
||||
ExporterOwner,
|
||||
ExporterSpec,
|
||||
OpenTelemetryV2Config,
|
||||
)
|
||||
from litellm.integrations.otel.presets.utils import ensure_mappers
|
||||
from litellm.types.utils import StandardCallbackDynamicParams
|
||||
|
||||
|
|
@ -23,6 +27,7 @@ def langfuse_preset(
|
|||
kind=kind,
|
||||
endpoint=cfg.endpoint,
|
||||
headers=cfg.headers,
|
||||
owner=ExporterOwner.LANGFUSE_OTEL,
|
||||
),
|
||||
],
|
||||
"mapper_names": ensure_mappers(base.mapper_names, "langfuse"),
|
||||
|
|
|
|||
|
|
@ -1,7 +1,11 @@
|
|||
"""Levo preset — OTLP/HTTP to a Levo collector with org+workspace headers."""
|
||||
|
||||
from litellm.integrations.levo.levo import LevoLogger as _V1Levo
|
||||
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
|
||||
from litellm.integrations.otel.model.config import (
|
||||
ExporterOwner,
|
||||
ExporterSpec,
|
||||
OpenTelemetryV2Config,
|
||||
)
|
||||
|
||||
|
||||
def levo_preset(
|
||||
|
|
@ -18,6 +22,7 @@ def levo_preset(
|
|||
kind="otlp_http",
|
||||
endpoint=cfg.endpoint,
|
||||
headers=cfg.otlp_auth_headers,
|
||||
owner=ExporterOwner.LEVO,
|
||||
),
|
||||
],
|
||||
}
|
||||
|
|
|
|||
|
|
@ -6,7 +6,11 @@ from pydantic_settings import BaseSettings, SettingsConfigDict
|
|||
from litellm.integrations.arize.arize_phoenix import (
|
||||
ArizePhoenixLogger as _V1Phoenix,
|
||||
)
|
||||
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
|
||||
from litellm.integrations.otel.model.config import (
|
||||
ExporterOwner,
|
||||
ExporterSpec,
|
||||
OpenTelemetryV2Config,
|
||||
)
|
||||
from litellm.integrations.otel.presets.utils import ensure_mappers
|
||||
|
||||
|
||||
|
|
@ -37,6 +41,7 @@ def phoenix_preset(
|
|||
kind=cfg.protocol if hasattr(cfg, "protocol") else "otlp_http",
|
||||
endpoint=cfg.endpoint,
|
||||
headers=headers,
|
||||
owner=ExporterOwner.ARIZE_PHOENIX,
|
||||
),
|
||||
],
|
||||
"mapper_names": ensure_mappers(base.mapper_names, "openinference"),
|
||||
|
|
|
|||
|
|
@ -1,6 +1,10 @@
|
|||
"""Weave (W&B) preset."""
|
||||
|
||||
from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config
|
||||
from litellm.integrations.otel.model.config import (
|
||||
ExporterOwner,
|
||||
ExporterSpec,
|
||||
OpenTelemetryV2Config,
|
||||
)
|
||||
from litellm.integrations.otel.presets.utils import ensure_mappers
|
||||
from litellm.integrations.weave.weave_otel import (
|
||||
_get_weave_authorization_header,
|
||||
|
|
@ -23,6 +27,7 @@ def weave_preset(
|
|||
kind=weave_cfg.protocol or "otlp_http",
|
||||
endpoint=weave_cfg.endpoint,
|
||||
headers=weave_cfg.otlp_auth_headers,
|
||||
owner=ExporterOwner.WEAVE_OTEL,
|
||||
),
|
||||
],
|
||||
# Weave consumes OpenInference + a small Weave-specific overlay.
|
||||
|
|
|
|||
|
|
@ -15,8 +15,23 @@ BEDROCK_MANAGED_S3_PREFIXES = (
|
|||
BEDROCK_MANAGED_S3_UPLOAD_PREFIX,
|
||||
BEDROCK_MANAGED_S3_OUTPUT_PREFIX,
|
||||
)
|
||||
MANAGED_CLOUD_STORAGE_SCHEMES = ("s3://", "gs://")
|
||||
_MAPPING_PROXY_TYPE: type = type(MappingProxyType({}))
|
||||
|
||||
|
||||
def is_managed_cloud_storage_uri(file_id: str) -> bool:
|
||||
"""
|
||||
True if file_id is a raw cloud-storage object URI (e.g. ``s3://bucket/key``).
|
||||
|
||||
These are internal provider artifacts. On the multi-tenant proxy they must be
|
||||
retrieved through their managed unified file id so owner/team access is enforced;
|
||||
a raw URI supplied by a caller bypasses that check.
|
||||
"""
|
||||
return isinstance(file_id, str) and file_id.startswith(
|
||||
MANAGED_CLOUD_STORAGE_SCHEMES
|
||||
)
|
||||
|
||||
|
||||
_SAFE_OBJECT_COMPONENT_PATTERN = re.compile(r"[^A-Za-z0-9._-]+")
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -170,6 +170,16 @@ def get_error_message(error_obj) -> Optional[str]:
|
|||
|
||||
|
||||
####### EXCEPTION MAPPING ################
|
||||
def _get_body_error_code(error_str: str) -> int | None:
|
||||
"""Return error.code from a JSON error body, or None if not parseable."""
|
||||
try:
|
||||
body = json.loads(error_str)
|
||||
code = body.get("error", {}).get("code")
|
||||
return int(code) if code is not None else None
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _get_response_headers(original_exception: Exception) -> Optional[httpx.Headers]:
|
||||
"""
|
||||
Extract and return the response headers from an exception, if present.
|
||||
|
|
@ -1415,6 +1425,29 @@ def exception_type( # type: ignore
|
|||
),
|
||||
),
|
||||
)
|
||||
elif (
|
||||
isinstance(getattr(original_exception, "status_code", None), int)
|
||||
and 500 <= original_exception.status_code < 600
|
||||
and _get_body_error_code(error_str) == 429
|
||||
):
|
||||
# upstream gateway wraps a 429 inside a 5xx envelope
|
||||
# e.g. HTTP 500/503 with {"error":{"code":429,...}}.
|
||||
# Scoped to 5xx so HTTP 400/401 with body code:429
|
||||
# still maps to BadRequestError / AuthenticationError.
|
||||
exception_mapping_worked = True
|
||||
raise RateLimitError(
|
||||
message=f"litellm.RateLimitError: {custom_llm_provider}Exception - {error_str}",
|
||||
model=model,
|
||||
llm_provider=custom_llm_provider,
|
||||
litellm_debug_info=extra_information,
|
||||
response=httpx.Response(
|
||||
status_code=429,
|
||||
request=httpx.Request(
|
||||
method="POST",
|
||||
url=" https://cloud.google.com/vertex-ai/",
|
||||
),
|
||||
),
|
||||
)
|
||||
elif (
|
||||
"500 Internal Server Error" in error_str
|
||||
or "The model is overloaded." in error_str
|
||||
|
|
|
|||
|
|
@ -388,6 +388,9 @@ def get_llm_provider(
|
|||
elif endpoint == "https://api.inference.wandb.ai/v1":
|
||||
custom_llm_provider = "wandb"
|
||||
dynamic_api_key = get_secret_str("WANDB_API_KEY")
|
||||
elif endpoint == "https://pinstripes.io/v1":
|
||||
custom_llm_provider = "pinstripes"
|
||||
dynamic_api_key = get_secret_str("PINSTRIPES_API_KEY")
|
||||
|
||||
if api_base is not None and not isinstance(api_base, str):
|
||||
raise Exception(
|
||||
|
|
@ -641,7 +644,7 @@ def _get_openai_compatible_provider_info(
|
|||
api_base,
|
||||
dynamic_api_key,
|
||||
) = litellm.BedrockMantleChatConfig()._get_openai_compatible_provider_info(
|
||||
api_base, api_key, litellm_params=litellm_params
|
||||
api_base, api_key, litellm_params=litellm_params, model=model
|
||||
)
|
||||
elif custom_llm_provider == "nvidia_nim":
|
||||
# nvidia_nim is openai compatible, we just need to set this to custom_openai and have the api_base be https://api.endpoints.anyscale.com/v1
|
||||
|
|
|
|||
|
|
@ -44,8 +44,8 @@ class HealthCheckHelpers:
|
|||
model_params["litellm_logging_obj"] = litellm_logging_obj
|
||||
model_params["fallbacks"] = fallback_models
|
||||
model_params["max_tokens"] = model_params.get(
|
||||
"max_tokens", 10
|
||||
) # gpt-5-nano throws errors for max_tokens=1
|
||||
"max_tokens", 16
|
||||
) # GPT-5 models require max_output_tokens >= 16
|
||||
await acompletion(**model_params)
|
||||
return {}
|
||||
|
||||
|
|
|
|||
|
|
@ -2975,7 +2975,12 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
)
|
||||
self.model_call_details["end_time"] = end_time
|
||||
self.model_call_details.setdefault("original_response", None)
|
||||
self.model_call_details["response_cost"] = 0
|
||||
# A stream interrupted mid-flight still billed the provider for the
|
||||
# chunks already delivered; the router stashes that recovered usage as
|
||||
# ``combined_usage_object`` and pre-computes its cost, so preserve it
|
||||
# here instead of zeroing the spend on an otherwise-failed request.
|
||||
if self.model_call_details.get("combined_usage_object") is None:
|
||||
self.model_call_details["response_cost"] = 0
|
||||
|
||||
if hasattr(exception, "headers") and isinstance(exception.headers, dict):
|
||||
self.model_call_details.setdefault("litellm_params", {})
|
||||
|
|
@ -5416,19 +5421,20 @@ class StandardLoggingPayloadSetup:
|
|||
tb_lines[:MAXIMUM_TRACEBACK_LINES_TO_LOG]
|
||||
) # Limit to first 100 lines
|
||||
|
||||
# Prefer the `.message` attribute (set by ProxyException and every
|
||||
# litellm.exceptions.* class) over str(exc); ProxyException does not
|
||||
# call super().__init__() nor define __str__, so str() on it returns
|
||||
# an empty string, which used to silently strip the human-readable
|
||||
# message from spend_logs.metadata.error_information.
|
||||
# Use isinstance, not truthiness: an explicit empty string on
|
||||
# `.message` is a deliberate value and must not be replaced by
|
||||
# `str(exc)`.
|
||||
explicit_message = getattr(original_exception, "message", None)
|
||||
error_message = (
|
||||
explicit_message
|
||||
if isinstance(explicit_message, str) and explicit_message
|
||||
else str(original_exception)
|
||||
)
|
||||
if isinstance(explicit_message, str):
|
||||
error_message = explicit_message
|
||||
else:
|
||||
error_message = str(original_exception) if original_exception else ""
|
||||
|
||||
# Duck-typed read so bare-Exception subclasses like
|
||||
# `litellm.BudgetExceededError` can participate without joining the
|
||||
# RateLimitError hierarchy (which would break `except BudgetExceededError`).
|
||||
# Validated against the enum value sets so a third-party exception that
|
||||
# happens to declare a `.category` or `.rate_limit_type` string attribute
|
||||
# can't leak garbage into the payload or Prometheus label cardinality.
|
||||
rate_limit_category = validate_rate_limit_category(
|
||||
getattr(original_exception, "category", None)
|
||||
)
|
||||
|
|
@ -5441,7 +5447,7 @@ class StandardLoggingPayloadSetup:
|
|||
error_class=error_class,
|
||||
llm_provider=_llm_provider_in_exception,
|
||||
traceback=traceback_info,
|
||||
error_message=error_message if original_exception else "",
|
||||
error_message=error_message,
|
||||
error_rate_limit_category=rate_limit_category,
|
||||
error_rate_limit_type=rate_limit_type,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -191,6 +191,11 @@ def _get_service_tier_cost_key(base_key: str, service_tier: Optional[str]) -> st
|
|||
return base_key
|
||||
|
||||
|
||||
def _parse_above_token_threshold(key: str) -> float:
|
||||
threshold_str = key.split("_above_")[1].split("_tokens")[0]
|
||||
return float(threshold_str.replace("k", "")) * (1000 if "k" in threshold_str else 1)
|
||||
|
||||
|
||||
def _get_token_base_cost(
|
||||
model_info: ModelInfo, usage: Usage, service_tier: Optional[str] = None
|
||||
) -> Tuple[float, float, float, float, float]:
|
||||
|
|
@ -256,15 +261,13 @@ def _get_token_base_cost(
|
|||
|
||||
# Only sort the threshold keys (typically 1-2 keys instead of 66+)
|
||||
threshold: Optional[float] = None
|
||||
for key in sorted(threshold_keys, reverse=True):
|
||||
for key in sorted(threshold_keys, key=_parse_above_token_threshold, reverse=True):
|
||||
value = model_info.get(key)
|
||||
if value is not None:
|
||||
try:
|
||||
# Handle both formats: _above_128k_tokens and _above_128_tokens
|
||||
threshold_str = key.split("_above_")[1].split("_tokens")[0]
|
||||
threshold = float(threshold_str.replace("k", "")) * (
|
||||
1000 if "k" in threshold_str else 1
|
||||
)
|
||||
threshold = _parse_above_token_threshold(key)
|
||||
if usage.prompt_tokens > threshold:
|
||||
# Prefer a service_tier-specific above-threshold key when available,
|
||||
# e.g. input_cost_per_token_priority_above_200k_tokens for Gemini
|
||||
|
|
|
|||
|
|
@ -2290,6 +2290,7 @@ class CustomStreamWrapper:
|
|||
litellm.request_timeout
|
||||
)
|
||||
if self.logging_obj is not None:
|
||||
self._record_partial_usage_for_failure()
|
||||
## LOGGING
|
||||
threading.Thread(
|
||||
target=self.logging_obj.failure_handler,
|
||||
|
|
@ -2303,6 +2304,7 @@ class CustomStreamWrapper:
|
|||
except Exception as e:
|
||||
traceback_exception = traceback.format_exc()
|
||||
if self.logging_obj is not None:
|
||||
self._record_partial_usage_for_failure()
|
||||
## LOGGING
|
||||
threading.Thread(
|
||||
target=self.logging_obj.failure_handler,
|
||||
|
|
@ -2314,6 +2316,33 @@ class CustomStreamWrapper:
|
|||
)
|
||||
self._handle_stream_fallback_error(e)
|
||||
|
||||
def _record_partial_usage_for_failure(self) -> None:
|
||||
"""
|
||||
A stream that breaks mid-flight still billed the provider for the chunks
|
||||
already delivered. Recover that partial usage from the chunks seen so
|
||||
far and stash it, with its cost, on the logging object so the failure
|
||||
handler records the real partial spend instead of zero. A request that
|
||||
later recovers via a router fallback overwrites this with the combined
|
||||
success log on the same request id, so this never double counts.
|
||||
"""
|
||||
if self.logging_obj is None or not self.chunks:
|
||||
return
|
||||
try:
|
||||
partial_response = litellm.stream_chunk_builder(chunks=self.chunks)
|
||||
usage = cast(Optional[Usage], getattr(partial_response, "usage", None))
|
||||
if usage is None:
|
||||
return
|
||||
self.logging_obj.model_call_details["combined_usage_object"] = usage
|
||||
self.logging_obj.model_call_details["response_cost"] = (
|
||||
self.logging_obj._response_cost_calculator(result=partial_response)
|
||||
or 0.0
|
||||
)
|
||||
except Exception as recover_error:
|
||||
verbose_logger.debug(
|
||||
"could not recover partial usage for interrupted stream: %s",
|
||||
recover_error,
|
||||
)
|
||||
|
||||
def _handle_stream_fallback_error(self, e: Exception) -> "NoReturn":
|
||||
"""
|
||||
Common error handling for both __next__ and __anext__.
|
||||
|
|
|
|||
|
|
@ -859,7 +859,17 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
"""
|
||||
new_tools: List[ChatCompletionToolParam] = []
|
||||
tool_name_mapping: Dict[str, str] = {}
|
||||
mapped_tool_params = ["name", "input_schema", "description", "cache_control"]
|
||||
# "type" is the Anthropic tool type (e.g. "custom"); it must not be
|
||||
# merged into the OpenAI function `parameters` schema below, or it
|
||||
# overwrites the real parameters.type ("object") and the provider
|
||||
# rejects the request. See #30557.
|
||||
mapped_tool_params = [
|
||||
"name",
|
||||
"input_schema",
|
||||
"description",
|
||||
"cache_control",
|
||||
"type",
|
||||
]
|
||||
|
||||
for idx, tool in enumerate(tools):
|
||||
# Check if this is an Anthropic-native tool that should be kept as-is
|
||||
|
|
|
|||
|
|
@ -218,8 +218,20 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
|
|||
- Qualifier goes as query parameter
|
||||
- Only the payload goes in the request body
|
||||
|
||||
Payload shape:
|
||||
- ``prompt`` is always present and contains the text-only flatten of the
|
||||
last message's content (existing behavior).
|
||||
- ``content`` is added ONLY when the ``forward_multimodal_content`` litellm
|
||||
param is truthy AND the last message's ``content`` is a list containing a
|
||||
non-text block (e.g. ``image_url``, ``file``, ``input_audio``). The list is
|
||||
forwarded verbatim so the agent's ``@app.entrypoint`` handler can parse the
|
||||
OpenAI-shaped multimodal blocks. This is opt-in because an AgentCore agent
|
||||
must be explicitly written to read ``payload["content"]``; by default the
|
||||
payload stays byte-identical to the legacy ``{"prompt": "..."}`` shape.
|
||||
|
||||
Returns:
|
||||
dict: Payload dict containing the prompt
|
||||
dict: Payload dict containing the prompt and (optionally) the OpenAI
|
||||
content list.
|
||||
"""
|
||||
verbose_logger.debug(
|
||||
f"AgentCore transform_request - optional_params keys: {list(optional_params.keys())}"
|
||||
|
|
@ -231,6 +243,20 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
|
|||
# Create the payload - this is what goes in the body (raw JSON)
|
||||
payload: dict = {"prompt": prompt}
|
||||
|
||||
# Opt-in: when forward_multimodal_content is set, forward the OpenAI content
|
||||
# list verbatim under "content" so an attachment-aware agent can read the raw
|
||||
# blocks (image_url, file, etc.). Default off keeps the payload byte-identical
|
||||
# to the legacy {"prompt": "..."} shape for agents that only read the prompt.
|
||||
if self._should_forward_multimodal_content(optional_params, litellm_params):
|
||||
last_content = messages[-1].get("content")
|
||||
if isinstance(last_content, list) and any(
|
||||
isinstance(block, dict) and block.get("type") not in (None, "text")
|
||||
for block in last_content
|
||||
):
|
||||
# Copy so the payload never aliases messages[-1]["content"]; shallow,
|
||||
# not deep, to avoid cloning large base64 media on the request path.
|
||||
payload["content"] = list(last_content)
|
||||
|
||||
# Get or generate session ID - this goes in the header
|
||||
runtime_session_id = self._get_runtime_session_id(optional_params)
|
||||
headers["X-Amzn-Bedrock-AgentCore-Runtime-Session-Id"] = runtime_session_id
|
||||
|
|
@ -246,6 +272,29 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM):
|
|||
verbose_logger.debug(f"PAYLOAD: {payload}")
|
||||
return payload
|
||||
|
||||
@staticmethod
|
||||
def _should_forward_multimodal_content(
|
||||
optional_params: dict, litellm_params: dict
|
||||
) -> bool:
|
||||
"""Whether to forward raw OpenAI content blocks under ``payload["content"]``.
|
||||
|
||||
Opt-in via the ``forward_multimodal_content`` litellm param (default ``False``)
|
||||
because AgentCore agents must be explicitly written to read the field. The
|
||||
value may arrive as a bool or a config/env string ("true", "1", ...). Checks
|
||||
``optional_params`` first (where other AgentCore params land), then
|
||||
``litellm_params``.
|
||||
"""
|
||||
for source in (optional_params, litellm_params):
|
||||
if not isinstance(source, dict):
|
||||
continue
|
||||
value = source.get("forward_multimodal_content")
|
||||
if value is None:
|
||||
continue
|
||||
if isinstance(value, str):
|
||||
return value.strip().lower() in ("1", "true", "yes", "on")
|
||||
return bool(value)
|
||||
return False
|
||||
|
||||
def _extract_sse_json(self, line: str) -> Optional[Dict]:
|
||||
"""Extract and parse JSON from an SSE data line."""
|
||||
if not line.startswith("data:"):
|
||||
|
|
|
|||
|
|
@ -1,8 +1,6 @@
|
|||
import asyncio
|
||||
import base64
|
||||
import os
|
||||
from types import MappingProxyType
|
||||
from typing import Any, Coroutine, Mapping, Optional, Tuple, Union, cast
|
||||
from collections.abc import Mapping
|
||||
from typing import Any, Coroutine, Optional, Tuple, Union
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -17,7 +15,6 @@ from litellm.types.llms.openai import (
|
|||
FileContentRequest,
|
||||
HttpxBinaryResponseContent,
|
||||
)
|
||||
from litellm.types.utils import SpecialEnums
|
||||
|
||||
from ..base_aws_llm import BaseAWSLLM
|
||||
|
||||
|
|
@ -37,40 +34,9 @@ class BedrockFilesHandler(BaseAWSLLM):
|
|||
)
|
||||
|
||||
def _extract_s3_uri_from_file_id(self, file_id: str) -> str:
|
||||
"""
|
||||
Extract S3 URI from encoded file ID.
|
||||
from .transformation import extract_s3_uri_from_file_id
|
||||
|
||||
The file ID can be in two formats:
|
||||
1. Base64-encoded unified file ID containing: llm_output_file_id,s3://bucket/path
|
||||
2. Direct S3 URI: s3://bucket/litellm-managed-prefix/path
|
||||
|
||||
Args:
|
||||
file_id: Encoded file ID or direct S3 URI
|
||||
|
||||
Returns:
|
||||
S3 URI (e.g., "s3://bucket-name/path/to/file")
|
||||
"""
|
||||
# First, try to decode if it's a base64-encoded unified file ID
|
||||
try:
|
||||
# Add padding if needed
|
||||
padded = file_id + "=" * (-len(file_id) % 4)
|
||||
decoded = base64.urlsafe_b64decode(padded).decode()
|
||||
|
||||
# Check if it's a unified file ID format
|
||||
if decoded.startswith(SpecialEnums.LITELM_MANAGED_FILE_ID_PREFIX.value):
|
||||
# Extract llm_output_file_id from the decoded string
|
||||
if "llm_output_file_id," in decoded:
|
||||
s3_uri = decoded.split("llm_output_file_id,")[1].split(";")[0]
|
||||
return s3_uri
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# If not base64 encoded or doesn't contain llm_output_file_id, accept only
|
||||
# explicit S3 URIs. Bucket and key validation happens before any S3 call.
|
||||
if file_id.startswith("s3://"):
|
||||
return file_id
|
||||
|
||||
raise ValueError("file_id must be a managed LiteLLM S3 file id")
|
||||
return extract_s3_uri_from_file_id(file_id)
|
||||
|
||||
def _parse_s3_uri(
|
||||
self,
|
||||
|
|
@ -95,26 +61,12 @@ class BedrockFilesHandler(BaseAWSLLM):
|
|||
allow_legacy_cloud_file_ids=allow_legacy_cloud_file_ids,
|
||||
)
|
||||
|
||||
def _get_configured_s3_bucket_name(self, litellm_params: dict) -> str:
|
||||
trusted_model_credentials = litellm_params.get(
|
||||
"_litellm_internal_model_credentials"
|
||||
)
|
||||
bucket_name = None
|
||||
if isinstance(trusted_model_credentials, type(MappingProxyType({}))):
|
||||
trusted_model_credentials_mapping = cast(
|
||||
Mapping[str, Any], trusted_model_credentials
|
||||
)
|
||||
candidate_bucket_name = trusted_model_credentials_mapping.get(
|
||||
"s3_bucket_name"
|
||||
)
|
||||
if isinstance(candidate_bucket_name, str):
|
||||
bucket_name = candidate_bucket_name
|
||||
bucket_name = bucket_name or os.getenv("AWS_S3_BUCKET_NAME")
|
||||
if not bucket_name:
|
||||
raise ValueError(
|
||||
"S3 bucket_name is required. Set 's3_bucket_name' in proxy config or AWS_S3_BUCKET_NAME for Bedrock file content retrieval."
|
||||
)
|
||||
return bucket_name
|
||||
def _get_configured_s3_bucket_name(
|
||||
self, litellm_params: Mapping[str, object]
|
||||
) -> str:
|
||||
from .transformation import get_configured_s3_bucket_name
|
||||
|
||||
return get_configured_s3_bucket_name(litellm_params)
|
||||
|
||||
async def afile_content(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -1,23 +1,37 @@
|
|||
import base64
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
from typing import Any, Dict, List, Optional, Tuple, Union
|
||||
from collections.abc import Mapping, MutableMapping
|
||||
from types import MappingProxyType
|
||||
from typing import (
|
||||
Any,
|
||||
Dict,
|
||||
List,
|
||||
Optional,
|
||||
Tuple,
|
||||
Union,
|
||||
)
|
||||
from urllib.parse import unquote
|
||||
|
||||
import httpx
|
||||
from httpx import Headers, Response
|
||||
from openai.types.file_deleted import FileDeleted
|
||||
from pydantic import BaseModel, ConfigDict
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.files.utils import FilesAPIUtils
|
||||
from litellm.litellm_core_utils.cloud_storage_security import (
|
||||
BEDROCK_MANAGED_S3_BATCH_PREFIX,
|
||||
BEDROCK_MANAGED_S3_PREFIXES,
|
||||
BEDROCK_MANAGED_S3_UPLOAD_PREFIX,
|
||||
build_managed_cloud_object_name,
|
||||
encode_s3_object_key_for_url,
|
||||
sanitize_cloud_object_component,
|
||||
should_allow_legacy_cloud_file_ids,
|
||||
split_configured_cloud_bucket_name,
|
||||
validate_managed_cloud_file_id,
|
||||
)
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
|
|
@ -28,18 +42,98 @@ from litellm.llms.base_llm.files.transformation import (
|
|||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
CreateFileRequest,
|
||||
FileContentRequest,
|
||||
FileTypes,
|
||||
HttpxBinaryResponseContent,
|
||||
OpenAICreateFileRequestOptionalParams,
|
||||
OpenAIFileObject,
|
||||
PathLike,
|
||||
)
|
||||
from litellm.types.utils import ExtractedFileData, LlmProviders
|
||||
from litellm.types.utils import ExtractedFileData, LlmProviders, SpecialEnums
|
||||
from litellm.utils import get_llm_provider
|
||||
|
||||
from ..base_aws_llm import BaseAWSLLM
|
||||
from ..common_utils import BedrockError
|
||||
|
||||
# litellm_params key used to hand the SigV4-signed GET headers from
|
||||
# `transform_file_content_request` to `validate_environment` (the only hook
|
||||
# the shared file-content HTTP handler exposes for setting request headers).
|
||||
# Same pattern as the `upload_url` handoff in `transform_create_file_request`.
|
||||
S3_SIGNED_GET_HEADERS_PARAM = "_s3_signed_get_headers"
|
||||
|
||||
|
||||
class _BedrockS3RequestParams(BaseModel):
|
||||
"""Typed view of the credential/region params the S3 GetObject path reads."""
|
||||
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
|
||||
aws_access_key_id: str | None = None
|
||||
aws_secret_access_key: str | None = None
|
||||
aws_session_token: str | None = None
|
||||
aws_region_name: str | None = None
|
||||
aws_session_name: str | None = None
|
||||
aws_profile_name: str | None = None
|
||||
aws_role_name: str | None = None
|
||||
aws_web_identity_token: str | None = None
|
||||
aws_sts_endpoint: str | None = None
|
||||
s3_region_name: str | None = None
|
||||
s3_endpoint_url: str | None = None
|
||||
|
||||
|
||||
class _TrustedS3ModelCredentials(BaseModel):
|
||||
"""The S3 bucket the server trusts file ids against, from the deployment snapshot."""
|
||||
|
||||
model_config = ConfigDict(extra="ignore")
|
||||
|
||||
s3_bucket_name: str | None = None
|
||||
|
||||
|
||||
def extract_s3_uri_from_file_id(file_id: str) -> str:
|
||||
"""
|
||||
Resolve a Bedrock file id to its S3 URI.
|
||||
|
||||
Accepts either a base64-encoded LiteLLM unified file id (whose decoded
|
||||
form carries `llm_output_file_id,s3://...`) or a direct `s3://` URI.
|
||||
"""
|
||||
try:
|
||||
padded = file_id + "=" * (-len(file_id) % 4)
|
||||
decoded = base64.urlsafe_b64decode(padded).decode()
|
||||
|
||||
if decoded.startswith(SpecialEnums.LITELM_MANAGED_FILE_ID_PREFIX.value):
|
||||
if "llm_output_file_id," in decoded:
|
||||
return decoded.split("llm_output_file_id,")[1].split(";")[0]
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
if file_id.startswith("s3://"):
|
||||
return file_id
|
||||
|
||||
raise ValueError("file_id must be a managed LiteLLM S3 file id")
|
||||
|
||||
|
||||
def get_configured_s3_bucket_name(litellm_params: Mapping[str, object]) -> str:
|
||||
"""
|
||||
Resolve the server-configured S3 bucket for Bedrock file operations.
|
||||
|
||||
Only trusts the immutable server-side credential snapshot or the
|
||||
environment; never a request-supplied param, since the bucket is what
|
||||
`validate_managed_cloud_file_id` checks file ids against.
|
||||
"""
|
||||
trusted_model_credentials = litellm_params.get(
|
||||
"_litellm_internal_model_credentials"
|
||||
)
|
||||
bucket_name: str | None = None
|
||||
if isinstance(trusted_model_credentials, MappingProxyType):
|
||||
snapshot: dict[str, object] = {}
|
||||
snapshot.update(trusted_model_credentials) # any-ok: untyped snapshot
|
||||
bucket_name = _TrustedS3ModelCredentials.model_validate(snapshot).s3_bucket_name
|
||||
bucket_name = bucket_name or os.getenv("AWS_S3_BUCKET_NAME")
|
||||
if not bucket_name:
|
||||
raise ValueError(
|
||||
"S3 bucket_name is required. Set 's3_bucket_name' in proxy config or AWS_S3_BUCKET_NAME for Bedrock file content retrieval."
|
||||
)
|
||||
return bucket_name
|
||||
|
||||
|
||||
class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
||||
"""
|
||||
|
|
@ -63,16 +157,21 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
headers: MutableMapping[str, object],
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
litellm_params: MutableMapping[str, object],
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
) -> dict:
|
||||
# No additional headers needed for S3 uploads - AWS credentials handled by BaseAWSLLM
|
||||
return headers
|
||||
result: dict[str, object] = {}
|
||||
result.update(headers)
|
||||
signed_headers = litellm_params.pop(S3_SIGNED_GET_HEADERS_PARAM, None)
|
||||
if isinstance(signed_headers, Mapping):
|
||||
result.update(signed_headers) # any-ok: untyped handoff headers
|
||||
# otherwise no extra headers - AWS credentials are handled by BaseAWSLLM
|
||||
return result
|
||||
|
||||
def _get_content_from_openai_file(self, openai_file_content: FileTypes) -> str:
|
||||
"""
|
||||
|
|
@ -927,23 +1026,114 @@ class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
|
|||
|
||||
def transform_file_content_request(
|
||||
self,
|
||||
file_content_request,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
) -> tuple[str, dict]:
|
||||
raise NotImplementedError(
|
||||
"BedrockFilesConfig does not support file content retrieval"
|
||||
file_content_request: FileContentRequest,
|
||||
optional_params: Mapping[str, object],
|
||||
litellm_params: MutableMapping[str, object],
|
||||
) -> tuple[str, dict[str, str]]:
|
||||
"""
|
||||
Build a SigV4-signed S3 GetObject request for a Bedrock batch file.
|
||||
|
||||
Bedrock batch file ids are `s3://bucket/key` URIs (or unified ids
|
||||
that decode to one); the bucket and key are validated against the
|
||||
server-configured bucket before any request is signed.
|
||||
"""
|
||||
file_id = file_content_request.get("file_id")
|
||||
if not file_id:
|
||||
raise ValueError("file_id is required for Bedrock file content retrieval")
|
||||
|
||||
s3_uri = extract_s3_uri_from_file_id(file_id)
|
||||
bucket_name, object_key = validate_managed_cloud_file_id(
|
||||
file_id=s3_uri,
|
||||
scheme="s3://",
|
||||
configured_bucket_name=get_configured_s3_bucket_name(litellm_params),
|
||||
allowed_object_prefixes=BEDROCK_MANAGED_S3_PREFIXES,
|
||||
allow_legacy_cloud_file_ids=should_allow_legacy_cloud_file_ids(
|
||||
litellm_params
|
||||
),
|
||||
)
|
||||
|
||||
# The shared file-content handler passes optional_params={}, so AWS
|
||||
# credentials/region arrive via litellm_params here (unlike the upload
|
||||
# path). s3_region_name wins over aws_region_name, same priority as
|
||||
# get_complete_file_url above.
|
||||
merged_params: dict[str, object] = {}
|
||||
merged_params.update(litellm_params)
|
||||
merged_params.update(optional_params)
|
||||
request_params = _BedrockS3RequestParams.model_validate(merged_params)
|
||||
|
||||
region_preference = (
|
||||
request_params.s3_region_name or request_params.aws_region_name
|
||||
)
|
||||
region_params: dict[str, str | None] = {"aws_region_name": region_preference}
|
||||
aws_region_name = self._get_aws_region_name(
|
||||
optional_params=region_params, model=""
|
||||
)
|
||||
|
||||
s3_endpoint_url = (
|
||||
request_params.s3_endpoint_url
|
||||
or f"https://s3.{aws_region_name}.amazonaws.com"
|
||||
).rstrip("/")
|
||||
url = f"{s3_endpoint_url}/{bucket_name}/{encode_s3_object_key_for_url(object_key)}"
|
||||
|
||||
litellm_params[S3_SIGNED_GET_HEADERS_PARAM] = self._sign_s3_get_request(
|
||||
api_base=url,
|
||||
aws_region_name=aws_region_name,
|
||||
request_params=request_params,
|
||||
)
|
||||
return url, {}
|
||||
|
||||
def _sign_s3_get_request(
|
||||
self,
|
||||
api_base: str,
|
||||
aws_region_name: str,
|
||||
request_params: _BedrockS3RequestParams,
|
||||
) -> dict[str, str]:
|
||||
"""
|
||||
SigV4-sign an S3 GetObject request, mirroring `_sign_s3_request` (PUT).
|
||||
"""
|
||||
try:
|
||||
import hashlib
|
||||
|
||||
from botocore.auth import SigV4Auth
|
||||
from botocore.awsrequest import AWSRequest
|
||||
except ImportError:
|
||||
raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
|
||||
|
||||
credentials = self.get_credentials( # any-ok: boto3 Credentials is untyped
|
||||
aws_access_key_id=request_params.aws_access_key_id,
|
||||
aws_secret_access_key=request_params.aws_secret_access_key,
|
||||
aws_session_token=request_params.aws_session_token,
|
||||
aws_region_name=aws_region_name,
|
||||
aws_session_name=request_params.aws_session_name,
|
||||
aws_profile_name=request_params.aws_profile_name,
|
||||
aws_role_name=request_params.aws_role_name,
|
||||
aws_web_identity_token=request_params.aws_web_identity_token,
|
||||
aws_sts_endpoint=request_params.aws_sts_endpoint,
|
||||
)
|
||||
|
||||
empty_body_hash = hashlib.sha256(b"").hexdigest()
|
||||
aws_request = AWSRequest( # any-ok: botocore AWSRequest is untyped
|
||||
method="GET",
|
||||
url=api_base,
|
||||
headers={"x-amz-content-sha256": empty_body_hash},
|
||||
)
|
||||
auth = SigV4Auth(credentials, "s3", aws_region_name) # any-ok: botocore untyped
|
||||
auth.add_auth(aws_request) # any-ok: botocore request mutation is untyped
|
||||
return dict(aws_request.headers) # any-ok: botocore headers are untyped
|
||||
|
||||
def transform_file_content_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
litellm_params: dict,
|
||||
) -> HttpxBinaryResponseContent:
|
||||
raise NotImplementedError(
|
||||
"BedrockFilesConfig does not support file content retrieval"
|
||||
)
|
||||
if raw_response.status_code >= 400:
|
||||
raise BedrockError(
|
||||
status_code=raw_response.status_code,
|
||||
message=raw_response.text,
|
||||
headers=raw_response.headers,
|
||||
)
|
||||
return HttpxBinaryResponseContent(response=raw_response)
|
||||
|
||||
|
||||
class BedrockJsonlFilesTransformation:
|
||||
|
|
|
|||
|
|
@ -4,8 +4,10 @@ Amazon Bedrock Mantle - OpenAI-compatible inference engine in Amazon Bedrock.
|
|||
API docs: https://docs.aws.amazon.com/bedrock/latest/userguide/bedrock-mantle.html
|
||||
|
||||
Base URL: https://bedrock-mantle.{region}.api.aws/v1
|
||||
Auth: AWS Bedrock API key as Bearer token (set via BEDROCK_MANTLE_API_KEY env var)
|
||||
or region-aware key via BEDROCK_MANTLE_{REGION}_API_KEY.
|
||||
Auth: Bearer token (litellm_params.api_key, BEDROCK_MANTLE_API_KEY, or the
|
||||
standard AWS_BEARER_TOKEN_BEDROCK) when present; otherwise AWS SigV4
|
||||
(service "bedrock") over the standard credential chain. See
|
||||
BedrockMantleAuthMixin in common_utils.
|
||||
"""
|
||||
|
||||
from typing import Iterator, AsyncIterator, Any, List, Optional, Tuple, Union
|
||||
|
|
@ -13,20 +15,27 @@ from typing import Iterator, AsyncIterator, Any, List, Optional, Tuple, Union
|
|||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
from litellm.llms.bedrock_mantle.common_utils import (
|
||||
BEDROCK_MANTLE_DEFAULT_REGION,
|
||||
BedrockMantleAuthMixin,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
from ..common_utils import mantle_base_segment
|
||||
from ...openai_like.chat.transformation import OpenAILikeChatConfig
|
||||
|
||||
BEDROCK_MANTLE_DEFAULT_REGION = "us-east-1"
|
||||
|
||||
|
||||
class BedrockMantleChatConfig(OpenAILikeChatConfig):
|
||||
class BedrockMantleChatConfig(BedrockMantleAuthMixin, OpenAILikeChatConfig):
|
||||
"""
|
||||
Transformation config for Amazon Bedrock Mantle OpenAI-compatible API.
|
||||
"""
|
||||
|
||||
def __init__(self, aws_signer: BaseAWSLLM | None = None):
|
||||
super().__init__()
|
||||
self._aws_signer = aws_signer or BaseAWSLLM()
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return "bedrock_mantle"
|
||||
|
|
@ -40,6 +49,7 @@ class BedrockMantleChatConfig(OpenAILikeChatConfig):
|
|||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
litellm_params: Optional[GenericLiteLLMParams] = None,
|
||||
model: str | None = None,
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
region = (
|
||||
(litellm_params.aws_region_name if litellm_params else None)
|
||||
|
|
@ -49,12 +59,15 @@ class BedrockMantleChatConfig(OpenAILikeChatConfig):
|
|||
or BEDROCK_MANTLE_DEFAULT_REGION
|
||||
)
|
||||
BaseAWSLLM._validate_aws_region_name(region)
|
||||
# The base path segment is data-driven per model (use_openai_responses_path
|
||||
# flag): gemma-4-* and gpt-5.x are served on /openai/v1, everything else on
|
||||
# /v1. An explicit api_base still wins over the derived default.
|
||||
api_base = (
|
||||
api_base
|
||||
or get_secret_str("BEDROCK_MANTLE_API_BASE")
|
||||
or f"https://bedrock-mantle.{region}.api.aws/v1"
|
||||
or f"https://bedrock-mantle.{region}.api.aws/{mantle_base_segment(model, litellm.model_cost)}"
|
||||
)
|
||||
dynamic_api_key = api_key or get_secret_str("BEDROCK_MANTLE_API_KEY")
|
||||
dynamic_api_key = self._resolve_bearer_token(api_key)
|
||||
return api_base, dynamic_api_key
|
||||
|
||||
def validate_environment(
|
||||
|
|
|
|||
147
litellm/llms/bedrock_mantle/common_utils.py
Normal file
147
litellm/llms/bedrock_mantle/common_utils.py
Normal file
|
|
@ -0,0 +1,147 @@
|
|||
"""Shared auth, region resolution, and routing helpers for the Amazon Bedrock Mantle provider.
|
||||
|
||||
Mantle authenticates with a Bearer token when one is available
|
||||
(litellm_params.api_key, BEDROCK_MANTLE_API_KEY, or the standard
|
||||
AWS_BEARER_TOKEN_BEDROCK); otherwise it falls back to AWS SigV4 (service
|
||||
"bedrock") over the standard credential chain (IAM role / access key / profile /
|
||||
web identity). The Chat Completions and Responses backends share this behaviour
|
||||
through BedrockMantleAuthMixin so the two paths can never drift apart.
|
||||
|
||||
The two routing helpers (mantle_supports_responses, mantle_base_segment) are
|
||||
pure functions of (model, model_cost) so they can be unit-tested without patching
|
||||
global state.
|
||||
"""
|
||||
|
||||
import re
|
||||
from typing import Tuple
|
||||
|
||||
from botocore.exceptions import (
|
||||
CredentialRetrievalError,
|
||||
NoCredentialsError,
|
||||
PartialCredentialsError,
|
||||
ProfileNotFound,
|
||||
)
|
||||
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
BEDROCK_MANTLE_DEFAULT_REGION = "us-east-1"
|
||||
|
||||
# Standard Mantle host: https://bedrock-mantle.<region>.api.aws (group 1 = region).
|
||||
MANTLE_HOST_RE = re.compile(
|
||||
r"^https?://bedrock-mantle\.([^/.]+)\.api\.aws", re.IGNORECASE
|
||||
)
|
||||
|
||||
|
||||
class BedrockMantleAuthMixin:
|
||||
_aws_signer: BaseAWSLLM
|
||||
|
||||
@staticmethod
|
||||
def _resolve_bearer_token(api_key: str | None) -> str | None:
|
||||
return (
|
||||
api_key
|
||||
or get_secret_str("BEDROCK_MANTLE_API_KEY")
|
||||
or get_secret_str("AWS_BEARER_TOKEN_BEDROCK")
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _resolve_region(params: dict) -> str:
|
||||
region = params.get("aws_region_name")
|
||||
if region:
|
||||
BaseAWSLLM._validate_aws_region_name(region)
|
||||
return region
|
||||
base = params.get("api_base") or get_secret_str("BEDROCK_MANTLE_API_BASE")
|
||||
if base:
|
||||
match = MANTLE_HOST_RE.match(base.rstrip("/"))
|
||||
if match:
|
||||
return match.group(1)
|
||||
return (
|
||||
get_secret_str("BEDROCK_MANTLE_REGION")
|
||||
or get_secret_str("AWS_REGION_NAME")
|
||||
or get_secret_str("AWS_REGION")
|
||||
or BEDROCK_MANTLE_DEFAULT_REGION
|
||||
)
|
||||
|
||||
def sign_request(
|
||||
self,
|
||||
headers: dict,
|
||||
optional_params: dict,
|
||||
request_data: dict,
|
||||
api_base: str,
|
||||
api_key: str | None = None,
|
||||
model: str | None = None,
|
||||
stream: bool | None = None,
|
||||
fake_stream: bool | None = None,
|
||||
) -> Tuple[dict, bytes | None]:
|
||||
bearer = self._resolve_bearer_token(api_key)
|
||||
if not bearer:
|
||||
# Pin the credential-scope region to the region of the actual signing URL
|
||||
# so the SigV4 scope and URL host can never disagree, even when a stale
|
||||
# api_base and aws_region_name point at different regions.
|
||||
host_match = MANTLE_HOST_RE.match(api_base.rstrip("/"))
|
||||
optional_params = {
|
||||
**optional_params,
|
||||
"aws_region_name": (
|
||||
host_match.group(1)
|
||||
if host_match
|
||||
else self._resolve_region({**optional_params, "api_base": api_base})
|
||||
),
|
||||
}
|
||||
headers = {k: v for k, v in headers.items() if k.lower() != "authorization"}
|
||||
try:
|
||||
return self._aws_signer._sign_request(
|
||||
service_name="bedrock",
|
||||
headers=headers,
|
||||
optional_params=optional_params,
|
||||
request_data=request_data,
|
||||
api_base=api_base,
|
||||
api_key=bearer,
|
||||
model=model,
|
||||
stream=stream,
|
||||
fake_stream=fake_stream,
|
||||
)
|
||||
except (
|
||||
NoCredentialsError,
|
||||
PartialCredentialsError,
|
||||
ProfileNotFound,
|
||||
CredentialRetrievalError,
|
||||
) as e:
|
||||
raise ValueError(
|
||||
"Bedrock Mantle auth failed: no Bearer token and no usable AWS "
|
||||
"credentials. Set BEDROCK_MANTLE_API_KEY (or AWS_BEARER_TOKEN_BEDROCK) "
|
||||
"or pass api_key for Bearer auth, or provide AWS credentials "
|
||||
"(IAM role / access key / profile / web identity) for SigV4."
|
||||
) from e
|
||||
|
||||
|
||||
def mantle_supports_responses(model: str | None, model_cost: dict) -> bool:
|
||||
"""Whether a Bedrock Mantle model can serve the native Responses API.
|
||||
|
||||
Purely data-driven from the model's price-map capability signal -- either
|
||||
/v1/responses in supported_endpoints, or mode=responses -- both overridable
|
||||
via register_model and proxy model_info, so onboarding a model is a JSON
|
||||
change, never a code change. There is deliberately NO model-name match here:
|
||||
capability is per-model, not per-family (openai.gpt-oss-120b supports
|
||||
Responses while openai.gpt-oss-safeguard-120b does not, despite sharing the
|
||||
gpt-oss substring), so a substring gate would be wrong. A model absent from
|
||||
model_cost simply has no signal and returns False (chat-completions emulation).
|
||||
"""
|
||||
entry = model_cost.get(f"bedrock_mantle/{model}", {})
|
||||
if "/v1/responses" in (entry.get("supported_endpoints") or []):
|
||||
return True
|
||||
return entry.get("mode") == "responses"
|
||||
|
||||
|
||||
def mantle_base_segment(model: str | None, model_cost: dict) -> str:
|
||||
"""Return the base path segment for a Bedrock Mantle model's OpenAI surface.
|
||||
|
||||
Data-driven from the model's price-map use_openai_responses_path flag
|
||||
(overridable via register_model / proxy model_info). Per the AWS model cards,
|
||||
gpt-5.x and the google gemma-4-* family carry that flag and are served on the
|
||||
/openai/v1 base (.../openai/v1/responses and .../openai/v1/chat/completions);
|
||||
every other model including gpt-oss uses the standard /v1 base. The segment is
|
||||
the base for the model's whole OpenAI-compatible surface, so both the chat and
|
||||
responses configs derive from it -- there is no separate model-name rule.
|
||||
"""
|
||||
entry = model_cost.get(f"bedrock_mantle/{model}", {})
|
||||
return "openai/v1" if entry.get("use_openai_responses_path") is True else "v1"
|
||||
|
|
@ -15,26 +15,20 @@ role / access key / profile / web identity), signed via the shared
|
|||
BaseAWSLLM._sign_request after the request body is finalized.
|
||||
"""
|
||||
|
||||
import re
|
||||
from typing import Any, Dict, List, Optional, Tuple
|
||||
|
||||
from botocore.exceptions import (
|
||||
CredentialRetrievalError,
|
||||
NoCredentialsError,
|
||||
PartialCredentialsError,
|
||||
ProfileNotFound,
|
||||
)
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
|
||||
from litellm.llms.bedrock_mantle.common_utils import (
|
||||
MANTLE_HOST_RE,
|
||||
BedrockMantleAuthMixin,
|
||||
)
|
||||
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
BEDROCK_MANTLE_DEFAULT_REGION = "us-east-1"
|
||||
|
||||
# Checked longest/most-specific first so a full endpoint URL collapses to host
|
||||
# in one pass and the appended path never doubles.
|
||||
_BASE_SUFFIXES_TO_STRIP = (
|
||||
|
|
@ -45,18 +39,13 @@ _BASE_SUFFIXES_TO_STRIP = (
|
|||
"/v1",
|
||||
)
|
||||
|
||||
# Standard Mantle host: https://bedrock-mantle.<region>.api.aws (group 1 = region).
|
||||
_MANTLE_HOST_RE = re.compile(
|
||||
r"^https?://bedrock-mantle\.([^/.]+)\.api\.aws", re.IGNORECASE
|
||||
)
|
||||
|
||||
# Per Bedrock Mantle Responses API validation errors.
|
||||
_BEDROCK_MANTLE_SUPPORTED_RESPONSE_TOOL_TYPES = frozenset(
|
||||
{"function", "mcp", "custom", "namespace", "tool_search"}
|
||||
)
|
||||
|
||||
|
||||
class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
||||
class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPIConfig):
|
||||
def __init__(
|
||||
self,
|
||||
aws_signer: Optional[BaseAWSLLM] = None,
|
||||
|
|
@ -70,24 +59,6 @@ class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
def custom_llm_provider(self) -> LlmProviders:
|
||||
return LlmProviders.BEDROCK_MANTLE
|
||||
|
||||
@staticmethod
|
||||
def _resolve_region(params: dict) -> str:
|
||||
region = params.get("aws_region_name")
|
||||
if region:
|
||||
BaseAWSLLM._validate_aws_region_name(region)
|
||||
return region
|
||||
base = params.get("api_base") or get_secret_str("BEDROCK_MANTLE_API_BASE")
|
||||
if base:
|
||||
match = _MANTLE_HOST_RE.match(base.rstrip("/"))
|
||||
if match:
|
||||
return match.group(1)
|
||||
return (
|
||||
get_secret_str("BEDROCK_MANTLE_REGION")
|
||||
or get_secret_str("AWS_REGION_NAME")
|
||||
or get_secret_str("AWS_REGION")
|
||||
or BEDROCK_MANTLE_DEFAULT_REGION
|
||||
)
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
|
|
@ -107,7 +78,7 @@ class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
# For the standard Mantle host (including the default-region base that
|
||||
# responses/main.py auto-injects into litellm_params.api_base), pin to the
|
||||
# single resolved region so aws_region_name wins; preserve custom proxy hosts.
|
||||
if _MANTLE_HOST_RE.match(base):
|
||||
if MANTLE_HOST_RE.match(base):
|
||||
base = f"https://bedrock-mantle.{region}.api.aws"
|
||||
path = "/openai/v1/responses" if self.use_openai_path else "/v1/responses"
|
||||
return f"{base}{path}"
|
||||
|
|
@ -116,13 +87,9 @@ class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams]
|
||||
) -> dict:
|
||||
litellm_params = litellm_params or GenericLiteLLMParams()
|
||||
api_key = (
|
||||
litellm_params.api_key
|
||||
or get_secret_str("BEDROCK_MANTLE_API_KEY")
|
||||
or get_secret_str("AWS_BEARER_TOKEN_BEDROCK")
|
||||
)
|
||||
if api_key:
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
bearer = self._resolve_bearer_token(litellm_params.api_key)
|
||||
if bearer:
|
||||
headers["Authorization"] = f"Bearer {bearer}"
|
||||
if litellm_params.aws_bedrock_project_id:
|
||||
headers["OpenAI-Project"] = litellm_params.aws_bedrock_project_id
|
||||
return headers
|
||||
|
|
@ -182,58 +149,3 @@ class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
params.pop("tools", None)
|
||||
|
||||
return params
|
||||
|
||||
def sign_request(
|
||||
self,
|
||||
headers: dict,
|
||||
optional_params: dict,
|
||||
request_data: dict,
|
||||
api_base: str,
|
||||
api_key: Optional[str] = None,
|
||||
model: Optional[str] = None,
|
||||
stream: Optional[bool] = None,
|
||||
fake_stream: Optional[bool] = None,
|
||||
) -> Tuple[dict, Optional[bytes]]:
|
||||
bearer = (
|
||||
api_key
|
||||
or get_secret_str("BEDROCK_MANTLE_API_KEY")
|
||||
or get_secret_str("AWS_BEARER_TOKEN_BEDROCK")
|
||||
)
|
||||
if not bearer:
|
||||
# SigV4 path. Pin the credential-scope region to the region of the actual
|
||||
# signing URL (api_base, already region-resolved by get_complete_url) so the
|
||||
# SigV4 scope and the URL host can never disagree. Resolve from api_base first,
|
||||
# then fall back to the regular precedence. Also drop any caller Authorization
|
||||
# so _sign_request's restore-original-Authorization step cannot override the
|
||||
# SigV4 header.
|
||||
optional_params = {
|
||||
**optional_params,
|
||||
"aws_region_name": self._resolve_region(
|
||||
{**optional_params, "api_base": api_base}
|
||||
),
|
||||
}
|
||||
headers = {k: v for k, v in headers.items() if k.lower() != "authorization"}
|
||||
try:
|
||||
return self._aws_signer._sign_request(
|
||||
service_name="bedrock",
|
||||
headers=headers,
|
||||
optional_params=optional_params,
|
||||
request_data=request_data,
|
||||
api_base=api_base,
|
||||
api_key=bearer,
|
||||
model=model,
|
||||
stream=stream,
|
||||
fake_stream=fake_stream,
|
||||
)
|
||||
except (
|
||||
NoCredentialsError,
|
||||
PartialCredentialsError,
|
||||
ProfileNotFound,
|
||||
CredentialRetrievalError,
|
||||
) as e:
|
||||
raise ValueError(
|
||||
"Bedrock Mantle auth failed: no Bearer token and no usable AWS "
|
||||
"credentials. Set BEDROCK_MANTLE_API_KEY (or AWS_BEARER_TOKEN_BEDROCK) "
|
||||
"or pass api_key for Bearer auth, or provide AWS credentials "
|
||||
"(IAM role / access key / profile / web identity) for SigV4."
|
||||
) from e
|
||||
|
|
|
|||
|
|
@ -116,6 +116,16 @@ class AiohttpResponseStream(httpx.AsyncByteStream):
|
|||
# For other exceptions, use the normal mapping
|
||||
with map_aiohttp_exceptions():
|
||||
raise
|
||||
finally:
|
||||
# Release the aiohttp connection when iteration ends for any
|
||||
# reason (read timeout, cancellation from a client disconnect,
|
||||
# GeneratorExit). Without this, abnormally terminated streams
|
||||
# permanently hold a slot in the TCPConnector pool; once the
|
||||
# pool is exhausted every request to that host times out (408)
|
||||
# until the proxy is restarted, even after the backend recovers.
|
||||
# On a fully-read response the connection was already released
|
||||
# at EOF and close() is a no-op.
|
||||
self._aiohttp_response.close()
|
||||
|
||||
async def aclose(self) -> None:
|
||||
with map_aiohttp_exceptions():
|
||||
|
|
|
|||
|
|
@ -392,7 +392,10 @@ class FireworksAIConfig(OpenAIGPTConfig):
|
|||
headers: dict,
|
||||
) -> dict:
|
||||
if not model.startswith("accounts/") and "#" not in model:
|
||||
model = f"accounts/fireworks/models/{model}"
|
||||
if model.endswith("-fast"):
|
||||
model = f"accounts/fireworks/routers/{model}"
|
||||
else:
|
||||
model = f"accounts/fireworks/models/{model}"
|
||||
messages = self._transform_messages_helper(
|
||||
messages=messages, model=model, litellm_params=litellm_params
|
||||
)
|
||||
|
|
|
|||
|
|
@ -18,6 +18,9 @@ from typing import (
|
|||
overload,
|
||||
)
|
||||
|
||||
import os
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
|
|
@ -426,6 +429,32 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
|
|||
)
|
||||
return messages, tools
|
||||
|
||||
def _should_preserve_cache_control_for_endpoint(
|
||||
self,
|
||||
custom_llm_provider: str | None,
|
||||
api_base: str | None,
|
||||
) -> bool:
|
||||
"""
|
||||
The generic `openai` provider also reaches OpenAI-compatible endpoints
|
||||
(a LiteLLM proxy, vLLM, an Anthropic-compatible gateway) via a custom
|
||||
api_base. Those can understand cache_control, so it must survive there.
|
||||
Real OpenAI cannot, so it is still stripped for an openai.com host.
|
||||
"""
|
||||
if custom_llm_provider != "openai":
|
||||
return False
|
||||
resolved_api_base = (
|
||||
api_base
|
||||
or litellm.api_base
|
||||
or os.getenv("OPENAI_BASE_URL")
|
||||
or os.getenv("OPENAI_API_BASE")
|
||||
)
|
||||
if not resolved_api_base:
|
||||
return False
|
||||
hostname = urlparse(resolved_api_base).hostname
|
||||
if hostname is None:
|
||||
return False
|
||||
return hostname != "openai.com" and not hostname.endswith(".openai.com")
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -441,11 +470,14 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
|
|||
dict: The transformed request. Sent as the body of the API call.
|
||||
"""
|
||||
messages = self._transform_messages(messages=messages, model=model)
|
||||
messages, tools = self.remove_cache_control_flag_from_messages_and_tools(
|
||||
model=model, messages=messages, tools=optional_params.get("tools", [])
|
||||
)
|
||||
if tools is not None and len(tools) > 0:
|
||||
optional_params["tools"] = tools
|
||||
if not self._should_preserve_cache_control_for_endpoint(
|
||||
litellm_params.get("custom_llm_provider"), litellm_params.get("api_base")
|
||||
):
|
||||
messages, tools = self.remove_cache_control_flag_from_messages_and_tools(
|
||||
model=model, messages=messages, tools=optional_params.get("tools", [])
|
||||
)
|
||||
if tools is not None and len(tools) > 0:
|
||||
optional_params["tools"] = tools
|
||||
|
||||
optional_params.pop("max_retries", None)
|
||||
|
||||
|
|
@ -466,16 +498,19 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
|
|||
transformed_messages = await self._transform_messages(
|
||||
messages=messages, model=model, is_async=True
|
||||
)
|
||||
(
|
||||
transformed_messages,
|
||||
tools,
|
||||
) = self.remove_cache_control_flag_from_messages_and_tools(
|
||||
model=model,
|
||||
messages=transformed_messages,
|
||||
tools=optional_params.get("tools", []),
|
||||
)
|
||||
if tools is not None and len(tools) > 0:
|
||||
optional_params["tools"] = tools
|
||||
if not self._should_preserve_cache_control_for_endpoint(
|
||||
litellm_params.get("custom_llm_provider"), litellm_params.get("api_base")
|
||||
):
|
||||
(
|
||||
transformed_messages,
|
||||
tools,
|
||||
) = self.remove_cache_control_flag_from_messages_and_tools(
|
||||
model=model,
|
||||
messages=transformed_messages,
|
||||
tools=optional_params.get("tools", []),
|
||||
)
|
||||
if tools is not None and len(tools) > 0:
|
||||
optional_params["tools"] = tools
|
||||
if self.__class__._is_base_class:
|
||||
return {
|
||||
"model": model,
|
||||
|
|
|
|||
|
|
@ -159,5 +159,14 @@
|
|||
"max_completion_tokens": "max_tokens"
|
||||
},
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"]
|
||||
},
|
||||
"pinstripes": {
|
||||
"base_url": "https://pinstripes.io/v1",
|
||||
"api_key_env": "PINSTRIPES_API_KEY",
|
||||
"api_base_env": "PINSTRIPES_API_BASE",
|
||||
"param_mappings": {
|
||||
"max_completion_tokens": "max_tokens"
|
||||
},
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses", "/v1/embeddings"]
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -58,11 +58,8 @@ def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]:
|
|||
|
||||
## CALCULATE OUTPUT COST
|
||||
output_cost_per_token = _safe_float_cast(model_info.get("output_cost_per_token"))
|
||||
completion_cost: float = (usage.completion_tokens or 0) * output_cost_per_token
|
||||
|
||||
## ADD REASONING TOKENS COST (if present)
|
||||
reasoning_tokens = getattr(usage, "reasoning_tokens", 0) or 0
|
||||
# Also check completion_tokens_details if reasoning_tokens is not directly available
|
||||
if (
|
||||
reasoning_tokens == 0
|
||||
and hasattr(usage, "completion_tokens_details")
|
||||
|
|
@ -73,9 +70,19 @@ def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]:
|
|||
)
|
||||
|
||||
reasoning_cost_value = model_info.get("output_cost_per_reasoning_token")
|
||||
|
||||
# `completion_tokens` includes `reasoning_tokens` per the OpenAI/Perplexity usage
|
||||
# convention (codified for the central path in PR #18607). When a reasoning rate is
|
||||
# configured we subtract before the output-rate multiplication so the reasoning
|
||||
# tokens are not billed twice.
|
||||
if reasoning_tokens > 0 and reasoning_cost_value is not None:
|
||||
reasoning_cost_per_token = _safe_float_cast(reasoning_cost_value)
|
||||
completion_cost += reasoning_tokens * reasoning_cost_per_token
|
||||
non_reasoning_completion_tokens = max(
|
||||
0, (usage.completion_tokens or 0) - reasoning_tokens
|
||||
)
|
||||
completion_cost: float = non_reasoning_completion_tokens * output_cost_per_token
|
||||
completion_cost += reasoning_tokens * _safe_float_cast(reasoning_cost_value)
|
||||
else:
|
||||
completion_cost = (usage.completion_tokens or 0) * output_cost_per_token
|
||||
|
||||
## ADD SEARCH QUERIES COST (if present)
|
||||
num_search_queries = 0
|
||||
|
|
|
|||
3
litellm/llms/tinyfish/search/__init__.py
Normal file
3
litellm/llms/tinyfish/search/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from litellm.llms.tinyfish.search.transformation import TinyfishSearchConfig
|
||||
|
||||
__all__ = ["TinyfishSearchConfig"]
|
||||
164
litellm/llms/tinyfish/search/transformation.py
Normal file
164
litellm/llms/tinyfish/search/transformation.py
Normal file
|
|
@ -0,0 +1,164 @@
|
|||
"""
|
||||
TinyFish Search API.
|
||||
Endpoint: GET https://api.search.tinyfish.ai
|
||||
Docs: https://docs.tinyfish.ai/search-api
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Literal, TypedDict
|
||||
from urllib.parse import urlencode
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, TypeAdapter, ValidationError
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.search.transformation import (
|
||||
BaseSearchConfig,
|
||||
SearchResponse,
|
||||
SearchResult,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
|
||||
class _TinyfishSearchRequestRequired(TypedDict):
|
||||
query: str
|
||||
|
||||
|
||||
class TinyfishSearchRequest(_TinyfishSearchRequestRequired, total=False):
|
||||
location: str
|
||||
language: str
|
||||
page: int
|
||||
include_thumbnail: bool
|
||||
max_results: int
|
||||
|
||||
|
||||
class _TinyfishResultItem(BaseModel, frozen=True):
|
||||
title: str = ""
|
||||
url: str = ""
|
||||
snippet: str = ""
|
||||
|
||||
|
||||
class _TinyfishApiResponse(BaseModel, frozen=True):
|
||||
results: tuple[_TinyfishResultItem, ...] = ()
|
||||
|
||||
|
||||
_UrlEncodableParams = TypeAdapter(dict[str, str | int | bool])
|
||||
_StrList = TypeAdapter(list[str])
|
||||
_StrFrozenSet = TypeAdapter(frozenset[str])
|
||||
|
||||
_TINYFISH_PARAMS_KEY = "_tinyfish_params"
|
||||
|
||||
|
||||
class TinyfishSearchConfig(BaseSearchConfig):
|
||||
TINYFISH_API_BASE = "https://api.search.tinyfish.ai"
|
||||
|
||||
@staticmethod
|
||||
def ui_friendly_name() -> str:
|
||||
return "TinyFish"
|
||||
|
||||
def get_http_method(self) -> Literal["GET", "POST"]:
|
||||
return "GET"
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict[str, str],
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
**kwargs: object,
|
||||
) -> dict[str, str]:
|
||||
resolved_key = api_key or get_secret_str("TINYFISH_API_KEY")
|
||||
if not resolved_key:
|
||||
raise ValueError(
|
||||
"TINYFISH_API_KEY is not set. Set `TINYFISH_API_KEY` environment variable."
|
||||
)
|
||||
return {**headers, "X-API-Key": resolved_key, "Accept": "application/json"}
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
optional_params: dict[str, object],
|
||||
data: dict[str, object] | list[dict[str, object]] | None = None,
|
||||
**kwargs: object,
|
||||
) -> str:
|
||||
resolved_base = (
|
||||
api_base or get_secret_str("TINYFISH_API_BASE") or self.TINYFISH_API_BASE
|
||||
)
|
||||
if isinstance(data, dict) and _TINYFISH_PARAMS_KEY in data:
|
||||
validated_params = _UrlEncodableParams.validate_python(
|
||||
data[_TINYFISH_PARAMS_KEY]
|
||||
)
|
||||
return f"{resolved_base}?{urlencode(validated_params, doseq=True)}"
|
||||
return resolved_base
|
||||
|
||||
def transform_search_request(
|
||||
self,
|
||||
query: str | list[str],
|
||||
optional_params: dict[str, object],
|
||||
**kwargs: object,
|
||||
) -> dict[str, object]:
|
||||
resolved_query = " ".join(query) if isinstance(query, list) else query
|
||||
|
||||
request_data: TinyfishSearchRequest = {"query": resolved_query}
|
||||
|
||||
country = optional_params.get("country")
|
||||
if isinstance(country, str):
|
||||
request_data["location"] = country
|
||||
|
||||
raw_max = optional_params.get("max_results")
|
||||
if isinstance(raw_max, (int, float, str)):
|
||||
request_data["max_results"] = max(1, min(int(raw_max), 20))
|
||||
|
||||
try:
|
||||
domains = _StrList.validate_python(
|
||||
optional_params.get("search_domain_filter")
|
||||
)
|
||||
except (ValidationError, TypeError):
|
||||
domains = []
|
||||
if domains:
|
||||
request_data["query"] = _append_domain_filters(
|
||||
request_data["query"], domains
|
||||
)
|
||||
|
||||
result_data: dict[str, object] = dict(request_data)
|
||||
|
||||
raw_supported: object = (
|
||||
self.get_supported_perplexity_optional_params() # any-ok: base class returns bare set
|
||||
)
|
||||
supported_perplexity = _StrFrozenSet.validate_python(raw_supported)
|
||||
for param, value in optional_params.items():
|
||||
if param not in supported_perplexity and param not in result_data:
|
||||
result_data[param] = value
|
||||
|
||||
return {_TINYFISH_PARAMS_KEY: result_data}
|
||||
|
||||
def transform_search_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj | None,
|
||||
**kwargs: object,
|
||||
) -> SearchResponse:
|
||||
raw_json: object = raw_response.json() # any-ok: httpx Response.json() -> Any
|
||||
parsed = _TinyfishApiResponse.model_validate(raw_json)
|
||||
|
||||
max_results_str: str = "20"
|
||||
if raw_response.request:
|
||||
raw_param: object = (
|
||||
raw_response.request.url.params.get( # any-ok: httpx QueryParams.get() -> Any
|
||||
"max_results", "20"
|
||||
)
|
||||
)
|
||||
max_results_str = str(raw_param)
|
||||
max_results: int = min(int(max_results_str), 20)
|
||||
|
||||
results = [
|
||||
SearchResult(title=item.title, url=item.url, snippet=item.snippet)
|
||||
for item in parsed.results[:max_results]
|
||||
]
|
||||
|
||||
return SearchResponse(results=results, object="search")
|
||||
|
||||
|
||||
def _append_domain_filters(query: str, domains: list[str]) -> str:
|
||||
domain_clauses = " OR ".join(f"site:{d}" for d in domains)
|
||||
return f"({query}) ({domain_clauses})"
|
||||
|
|
@ -271,7 +271,7 @@ def supports_response_json_schema(model: str) -> bool:
|
|||
|
||||
# Gemini 2.0+ and 2.5+ models support responseJsonSchema
|
||||
# Pattern matches: gemini-2.0-*, gemini-2.5-*, gemini-3-*, etc.
|
||||
gemini_2_plus_pattern = re.compile(r"gemini-([2-9]|[1-9]\d+)\.")
|
||||
gemini_2_plus_pattern = re.compile(r"gemini-(?:[2-9]|[1-9]\d+)(?:\.|\-)")
|
||||
|
||||
return bool(gemini_2_plus_pattern.search(model_lower))
|
||||
|
||||
|
|
|
|||
|
|
@ -10912,13 +10912,13 @@
|
|||
"supports_tool_choice": true
|
||||
},
|
||||
"command-r7b-12-2024": {
|
||||
"input_cost_per_token": 1.5e-07,
|
||||
"input_cost_per_token": 3.75e-08,
|
||||
"litellm_provider": "cohere_chat",
|
||||
"max_input_tokens": 128000,
|
||||
"max_output_tokens": 4096,
|
||||
"max_tokens": 4096,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3.75e-08,
|
||||
"output_cost_per_token": 1.5e-07,
|
||||
"source": "https://docs.cohere.com/v2/docs/command-r7b",
|
||||
"supports_function_calling": true,
|
||||
"supports_tool_choice": true
|
||||
|
|
@ -14612,6 +14612,38 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/deepseek-v4-flash": {
|
||||
"cache_read_input_token_cost": 2.8e-08,
|
||||
"input_cost_per_token": 1.4e-07,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 384000,
|
||||
"max_tokens": 384000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.8e-07,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/deepseek-v4-pro": {
|
||||
"cache_read_input_token_cost": 1.45e-07,
|
||||
"input_cost_per_token": 1.74e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 384000,
|
||||
"max_tokens": 384000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3.48e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/firefunction-v2": {
|
||||
"input_cost_per_token": 9e-07,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
|
|
@ -14687,43 +14719,64 @@
|
|||
"input_cost_per_token": 1.4e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 202800,
|
||||
"max_output_tokens": 202800,
|
||||
"max_tokens": 202800,
|
||||
"max_output_tokens": 131072,
|
||||
"max_tokens": 131072,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 4.4e-06,
|
||||
"source": "https://fireworks.ai/models/fireworks/glm-5p1",
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/glm-5p2": {
|
||||
"cache_read_input_token_cost": 2.6e-07,
|
||||
"input_cost_per_token": 1.4e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 131072,
|
||||
"max_tokens": 131072,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 4.4e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/gpt-oss-120b": {
|
||||
"cache_read_input_token_cost": 1.5e-08,
|
||||
"input_cost_per_token": 1.5e-07,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 131072,
|
||||
"max_output_tokens": 131072,
|
||||
"max_tokens": 131072,
|
||||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 6e-07,
|
||||
"source": "https://fireworks.ai/pricing",
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/gpt-oss-20b": {
|
||||
"input_cost_per_token": 5e-08,
|
||||
"cache_read_input_token_cost": 3.5e-08,
|
||||
"input_cost_per_token": 7e-08,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 131072,
|
||||
"max_output_tokens": 131072,
|
||||
"max_tokens": 131072,
|
||||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2e-07,
|
||||
"source": "https://fireworks.ai/pricing",
|
||||
"output_cost_per_token": 3e-07,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/kimi-k2-instruct": {
|
||||
"input_cost_per_token": 6e-07,
|
||||
|
|
@ -14779,6 +14832,38 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/kimi-k2p6": {
|
||||
"cache_read_input_token_cost": 1.6e-07,
|
||||
"input_cost_per_token": 9.5e-07,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 262144,
|
||||
"max_output_tokens": 262144,
|
||||
"max_tokens": 262144,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 4e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/kimi-k2p7-code": {
|
||||
"cache_read_input_token_cost": 1.9e-07,
|
||||
"input_cost_per_token": 9.5e-07,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 262144,
|
||||
"max_output_tokens": 262144,
|
||||
"max_tokens": 262144,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 4e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/llama-v3p1-405b-instruct": {
|
||||
"input_cost_per_token": 3e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
|
|
@ -14896,6 +14981,38 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/minimax-m2p7": {
|
||||
"cache_read_input_token_cost": 6e-08,
|
||||
"input_cost_per_token": 3e-07,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 196608,
|
||||
"max_output_tokens": 196608,
|
||||
"max_tokens": 196608,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.2e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/minimax-m3": {
|
||||
"cache_read_input_token_cost": 6e-08,
|
||||
"input_cost_per_token": 3e-07,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 512000,
|
||||
"max_output_tokens": 512000,
|
||||
"max_tokens": 512000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.2e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/mixtral-8x22b-instruct-hf": {
|
||||
"input_cost_per_token": 1.2e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
|
|
@ -14948,6 +15065,38 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": false
|
||||
},
|
||||
"fireworks_ai/deepseek-v4-flash": {
|
||||
"cache_read_input_token_cost": 2.8e-08,
|
||||
"input_cost_per_token": 1.4e-07,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 384000,
|
||||
"max_tokens": 384000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.8e-07,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/deepseek-v4-pro": {
|
||||
"cache_read_input_token_cost": 1.45e-07,
|
||||
"input_cost_per_token": 1.74e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 384000,
|
||||
"max_tokens": 384000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3.48e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/glm-4p7": {
|
||||
"cache_read_input_token_cost": 3e-07,
|
||||
"input_cost_per_token": 6e-07,
|
||||
|
|
@ -14968,15 +15117,80 @@
|
|||
"input_cost_per_token": 1.4e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 202800,
|
||||
"max_output_tokens": 202800,
|
||||
"max_tokens": 202800,
|
||||
"max_output_tokens": 131072,
|
||||
"max_tokens": 131072,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 4.4e-06,
|
||||
"source": "https://fireworks.ai/models/fireworks/glm-5p1",
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/glm-5p1-fast": {
|
||||
"cache_read_input_token_cost": 5.2e-07,
|
||||
"input_cost_per_token": 2.8e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 202800,
|
||||
"max_output_tokens": 131072,
|
||||
"max_tokens": 131072,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 8.8e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/glm-5p2": {
|
||||
"cache_read_input_token_cost": 2.6e-07,
|
||||
"input_cost_per_token": 1.4e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 1048576,
|
||||
"max_output_tokens": 131072,
|
||||
"max_tokens": 131072,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 4.4e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/gpt-oss-120b": {
|
||||
"cache_read_input_token_cost": 1.5e-08,
|
||||
"input_cost_per_token": 1.5e-07,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 131072,
|
||||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 6e-07,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/gpt-oss-20b": {
|
||||
"cache_read_input_token_cost": 3.5e-08,
|
||||
"input_cost_per_token": 7e-08,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 131072,
|
||||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3e-07,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/kimi-k2p5": {
|
||||
"cache_read_input_token_cost": 1e-07,
|
||||
|
|
@ -14992,6 +15206,70 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"fireworks_ai/kimi-k2p6": {
|
||||
"cache_read_input_token_cost": 1.6e-07,
|
||||
"input_cost_per_token": 9.5e-07,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 262144,
|
||||
"max_output_tokens": 262144,
|
||||
"max_tokens": 262144,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 4e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"fireworks_ai/kimi-k2p6-fast": {
|
||||
"cache_read_input_token_cost": 3e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 262144,
|
||||
"max_output_tokens": 262144,
|
||||
"max_tokens": 262144,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 8e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"fireworks_ai/kimi-k2p7-code": {
|
||||
"cache_read_input_token_cost": 1.9e-07,
|
||||
"input_cost_per_token": 9.5e-07,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 262144,
|
||||
"max_output_tokens": 262144,
|
||||
"max_tokens": 262144,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 4e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"fireworks_ai/kimi-k2p7-code-fast": {
|
||||
"cache_read_input_token_cost": 3.8e-07,
|
||||
"input_cost_per_token": 1.9e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 262144,
|
||||
"max_output_tokens": 262144,
|
||||
"max_tokens": 262144,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 8e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"fireworks_ai/minimax-m2p1": {
|
||||
"cache_read_input_token_cost": 3e-08,
|
||||
"input_cost_per_token": 3e-07,
|
||||
|
|
@ -15006,6 +15284,54 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"fireworks_ai/minimax-m2p7": {
|
||||
"cache_read_input_token_cost": 6e-08,
|
||||
"input_cost_per_token": 3e-07,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 196608,
|
||||
"max_output_tokens": 196608,
|
||||
"max_tokens": 196608,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.2e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/minimax-m3": {
|
||||
"cache_read_input_token_cost": 6e-08,
|
||||
"input_cost_per_token": 3e-07,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 512000,
|
||||
"max_output_tokens": 512000,
|
||||
"max_tokens": 512000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.2e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/qwen3p7-plus": {
|
||||
"cache_read_input_token_cost": 8e-08,
|
||||
"input_cost_per_token": 4e-07,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 262144,
|
||||
"max_output_tokens": 65536,
|
||||
"max_tokens": 65536,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.6e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"fireworks_ai/nomic-ai/nomic-embed-text-v1": {
|
||||
"input_cost_per_token": 8e-09,
|
||||
"litellm_provider": "fireworks_ai-embedding-models",
|
||||
|
|
@ -39467,6 +39793,22 @@
|
|||
"litellm_provider": "fireworks_ai",
|
||||
"mode": "chat"
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/qwen3p7-plus": {
|
||||
"cache_read_input_token_cost": 8e-08,
|
||||
"input_cost_per_token": 4e-07,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 262144,
|
||||
"max_output_tokens": 65536,
|
||||
"max_tokens": 65536,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 1.6e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/models/qwq-32b": {
|
||||
"max_tokens": 131072,
|
||||
"max_input_tokens": 131072,
|
||||
|
|
@ -39629,6 +39971,54 @@
|
|||
"litellm_provider": "fireworks_ai",
|
||||
"mode": "chat"
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/routers/glm-5p1-fast": {
|
||||
"cache_read_input_token_cost": 5.2e-07,
|
||||
"input_cost_per_token": 2.8e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 202800,
|
||||
"max_output_tokens": 131072,
|
||||
"max_tokens": 131072,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 8.8e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": false
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/routers/kimi-k2p6-fast": {
|
||||
"cache_read_input_token_cost": 3e-07,
|
||||
"input_cost_per_token": 2e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 262144,
|
||||
"max_output_tokens": 262144,
|
||||
"max_tokens": 262144,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 8e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"fireworks_ai/accounts/fireworks/routers/kimi-k2p7-code-fast": {
|
||||
"cache_read_input_token_cost": 3.8e-07,
|
||||
"input_cost_per_token": 1.9e-06,
|
||||
"litellm_provider": "fireworks_ai",
|
||||
"max_input_tokens": 262144,
|
||||
"max_output_tokens": 262144,
|
||||
"max_tokens": 262144,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 8e-06,
|
||||
"source": "https://docs.fireworks.ai/serverless/pricing",
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true
|
||||
},
|
||||
"novita/deepseek/deepseek-v3.2": {
|
||||
"litellm_provider": "novita",
|
||||
"mode": "chat",
|
||||
|
|
@ -41993,6 +42383,7 @@
|
|||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"],
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -42007,6 +42398,7 @@
|
|||
"max_output_tokens": 32768,
|
||||
"max_tokens": 32768,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"],
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -42021,6 +42413,7 @@
|
|||
"max_output_tokens": 65536,
|
||||
"max_tokens": 65536,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": ["/v1/chat/completions"],
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
|
|
@ -42034,6 +42427,7 @@
|
|||
"max_output_tokens": 65536,
|
||||
"max_tokens": 65536,
|
||||
"mode": "chat",
|
||||
"supported_endpoints": ["/v1/chat/completions"],
|
||||
"supports_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_response_schema": true,
|
||||
|
|
@ -42087,6 +42481,8 @@
|
|||
"max_output_tokens": 256000,
|
||||
"max_tokens": 256000,
|
||||
"mode": "chat",
|
||||
"use_openai_responses_path": true,
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"],
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": false,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -42101,6 +42497,8 @@
|
|||
"max_output_tokens": 256000,
|
||||
"max_tokens": 256000,
|
||||
"mode": "chat",
|
||||
"use_openai_responses_path": true,
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"],
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": false,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -42115,6 +42513,8 @@
|
|||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
"use_openai_responses_path": true,
|
||||
"supported_endpoints": ["/v1/chat/completions", "/v1/responses"],
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": false,
|
||||
"supports_reasoning": true,
|
||||
|
|
@ -42413,6 +42813,19 @@
|
|||
],
|
||||
"supports_audio_input": true
|
||||
},
|
||||
"soniox/stt-async-v5": {
|
||||
"litellm_provider": "soniox",
|
||||
"max_output_tokens": 8000,
|
||||
"max_tokens": 8000,
|
||||
"input_cost_per_second": 0.0,
|
||||
"output_cost_per_second": 0.0000277778,
|
||||
"mode": "audio_transcription",
|
||||
"source": "https://soniox.com/pricing",
|
||||
"supported_endpoints": [
|
||||
"/v1/audio/transcriptions"
|
||||
],
|
||||
"supports_audio_input": true
|
||||
},
|
||||
"tensormesh/Qwen/Qwen3.5-397B-A17B-FP8": {
|
||||
"litellm_provider": "tensormesh",
|
||||
"mode": "chat",
|
||||
|
|
|
|||
|
|
@ -23,9 +23,20 @@ import os
|
|||
from urllib.parse import quote
|
||||
|
||||
# Constants
|
||||
LITELLM_MCP_SERVER_NAME = "litellm-mcp-server"
|
||||
#
|
||||
# NOTE: The environment-backed values below are read once, when this module is
|
||||
# first imported, and cached for the lifetime of the process. Changing the
|
||||
# corresponding environment variables after import has no effect unless the
|
||||
# module is reloaded (e.g. ``importlib.reload``). Tests that override these
|
||||
# variables must reload this module — see
|
||||
# ``tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_identity_env.py``.
|
||||
LITELLM_MCP_SERVER_NAME = os.environ.get(
|
||||
"LITELLM_MCP_SERVER_NAME", "litellm-mcp-server"
|
||||
)
|
||||
LITELLM_MCP_SERVER_VERSION = "1.0.0"
|
||||
LITELLM_MCP_SERVER_DESCRIPTION = "MCP Server for LiteLLM"
|
||||
LITELLM_MCP_SERVER_DESCRIPTION = os.environ.get(
|
||||
"LITELLM_MCP_SERVER_DESCRIPTION", "MCP Server for LiteLLM"
|
||||
)
|
||||
MCP_TOOL_PREFIX_SEPARATOR = os.environ.get("MCP_TOOL_PREFIX_SEPARATOR", "-")
|
||||
MCP_TOOL_PREFIX_FORMAT = "{server_name}{separator}{tool_name}"
|
||||
|
||||
|
|
|
|||
|
|
@ -361,6 +361,16 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/realtime?{model}",
|
||||
"/v1/realtime?{model}",
|
||||
"/openai/v1/realtime?{model}",
|
||||
# realtime (GA WebRTC HTTP routes)
|
||||
"/realtime/client_secrets",
|
||||
"/v1/realtime/client_secrets",
|
||||
"/openai/v1/realtime/client_secrets",
|
||||
"/realtime/calls",
|
||||
"/v1/realtime/calls",
|
||||
"/openai/v1/realtime/calls",
|
||||
"/realtime/transcription_sessions",
|
||||
"/v1/realtime/transcription_sessions",
|
||||
"/openai/v1/realtime/transcription_sessions",
|
||||
# responses API
|
||||
"/responses",
|
||||
"/v1/responses",
|
||||
|
|
|
|||
|
|
@ -5,6 +5,7 @@ Unified /v1/messages endpoint - (Anthropic Spec)
|
|||
from fastapi import APIRouter, Depends, HTTPException, Request, Response
|
||||
from fastapi.responses import JSONResponse
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.anthropic_interface.exceptions import AnthropicExceptionMapping
|
||||
from litellm.integrations.custom_guardrail import ModifyResponseException
|
||||
|
|
@ -23,6 +24,40 @@ from litellm.types.utils import TokenCountResponse
|
|||
router = APIRouter()
|
||||
|
||||
|
||||
def _strip_total_tokens_from_anthropic_response(response: Any) -> None:
|
||||
"""Remove the OpenAI-flavored `usage.total_tokens` field that LiteLLM
|
||||
injects into Anthropic /v1/messages responses.
|
||||
|
||||
The Anthropic /v1/messages spec only defines:
|
||||
input_tokens, output_tokens, cache_creation_input_tokens,
|
||||
cache_read_input_tokens, cache_creation.{ephemeral_5m,ephemeral_1h}
|
||||
The streaming SSE path (message_delta.usage) already does not include
|
||||
total_tokens; this brings the non-streaming path into the same shape.
|
||||
|
||||
Handles both shapes returned by `base_process_llm_request`:
|
||||
- plain `dict` (most common — `AnthropicMessagesResponse` is a TypedDict
|
||||
and is `dict` at runtime)
|
||||
- Pydantic model whose `usage` attribute is dict-shaped (e.g. a
|
||||
BaseModel that holds raw Anthropic usage as a `dict[str, int]`)
|
||||
|
||||
Streaming results (StreamingResponse, AsyncIterator, etc.) and Pydantic
|
||||
models with strongly-typed Usage sub-models are left untouched —
|
||||
those paths either have separate serialization handling or impose
|
||||
type constraints the helper does not try to subvert.
|
||||
"""
|
||||
if response is None:
|
||||
return
|
||||
if isinstance(response, dict):
|
||||
usage = response.get("usage")
|
||||
if isinstance(usage, dict) and "total_tokens" in usage:
|
||||
usage.pop("total_tokens", None)
|
||||
return
|
||||
# Pydantic-model fallback: only mutate if `usage` is a dict.
|
||||
usage = getattr(response, "usage", None)
|
||||
if isinstance(usage, dict) and "total_tokens" in usage:
|
||||
usage.pop("total_tokens", None)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/v1/messages",
|
||||
tags=["[beta] Anthropic `/v1/messages`"],
|
||||
|
|
@ -72,6 +107,18 @@ async def anthropic_response(
|
|||
user_api_base=user_api_base,
|
||||
version=version,
|
||||
)
|
||||
# Optionally strip the non-Anthropic `usage.total_tokens` field
|
||||
# LiteLLM adds internally. Anthropic's official /v1/messages spec
|
||||
# only defines input_tokens / output_tokens / cache_*_input_tokens;
|
||||
# total_tokens is an OpenAI convention. Default off
|
||||
# (`litellm.strip_anthropic_total_tokens = False`) to preserve
|
||||
# backward compatibility for clients that currently read it; set
|
||||
# to True to align the wire response with the spec (and with the
|
||||
# streaming SSE path, which already omits total_tokens).
|
||||
# spend_logs / Prometheus still compute total internally — this
|
||||
# only affects the wire response.
|
||||
if litellm.strip_anthropic_total_tokens:
|
||||
_strip_total_tokens_from_anthropic_response(result)
|
||||
return result
|
||||
except ModifyResponseException as e:
|
||||
# Guardrail flagged content in passthrough mode - return 200 with violation message
|
||||
|
|
|
|||
|
|
@ -61,6 +61,7 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start
|
||||
from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec
|
||||
from litellm.proxy.common_utils.http_parsing_utils import (
|
||||
_safe_get_request_headers,
|
||||
|
|
@ -725,6 +726,7 @@ async def common_checks(
|
|||
user_spend = await get_current_spend(
|
||||
counter_key=f"spend:user:{user_object.user_id}",
|
||||
fallback_spend=user_object.spend or 0.0,
|
||||
max_budget=user_budget,
|
||||
)
|
||||
if math.isfinite(user_budget) and user_spend >= user_budget:
|
||||
raise litellm.BudgetExceededError(
|
||||
|
|
@ -1127,6 +1129,8 @@ async def _check_end_user_budget(
|
|||
end_user_spend = await get_current_spend(
|
||||
counter_key=f"spend:end_user:{end_user_obj.user_id}",
|
||||
fallback_spend=end_user_obj.spend or 0.0,
|
||||
max_budget=end_user_budget,
|
||||
fallback_authoritative=True,
|
||||
)
|
||||
if end_user_spend > end_user_budget:
|
||||
raise litellm.BudgetExceededError(
|
||||
|
|
@ -3615,6 +3619,7 @@ async def _virtual_key_max_budget_check(
|
|||
spend = await get_current_spend(
|
||||
counter_key=counter_key,
|
||||
fallback_spend=fallback_spend,
|
||||
max_budget=valid_token.max_budget,
|
||||
)
|
||||
|
||||
####################################
|
||||
|
|
@ -3684,6 +3689,10 @@ async def _virtual_key_multi_budget_check(
|
|||
window_spend = await get_current_spend(
|
||||
counter_key=counter_key,
|
||||
fallback_spend=0.0,
|
||||
max_budget=w["max_budget"],
|
||||
window_entity_type="Key",
|
||||
window_entity_id=valid_token.token,
|
||||
window_start=get_budget_window_start(w),
|
||||
)
|
||||
if math.isfinite(w["max_budget"]) and window_spend >= w["max_budget"]:
|
||||
raise litellm.BudgetExceededError(
|
||||
|
|
@ -3938,6 +3947,7 @@ async def _check_team_member_budget(
|
|||
team_member_spend = await get_current_spend(
|
||||
counter_key=f"spend:team_member:{valid_token.user_id}:{team_object.team_id}",
|
||||
fallback_spend=team_member_spend,
|
||||
max_budget=team_member_budget,
|
||||
)
|
||||
|
||||
if (
|
||||
|
|
@ -4023,6 +4033,7 @@ async def _team_max_budget_check(
|
|||
spend = await get_current_spend(
|
||||
counter_key=f"spend:team:{team_object.team_id}",
|
||||
fallback_spend=team_object.spend or 0.0,
|
||||
max_budget=team_object.max_budget,
|
||||
)
|
||||
|
||||
if math.isfinite(team_object.max_budget) and spend > team_object.max_budget:
|
||||
|
|
@ -4072,6 +4083,10 @@ async def _team_multi_budget_check(
|
|||
window_spend = await get_current_spend(
|
||||
counter_key=counter_key,
|
||||
fallback_spend=0.0,
|
||||
max_budget=w["max_budget"],
|
||||
window_entity_type="Team",
|
||||
window_entity_id=team_object.team_id,
|
||||
window_start=get_budget_window_start(w),
|
||||
)
|
||||
if math.isfinite(w["max_budget"]) and window_spend >= w["max_budget"]:
|
||||
raise litellm.BudgetExceededError(
|
||||
|
|
@ -4377,6 +4392,7 @@ async def _organization_max_budget_check(
|
|||
org_spend = await get_current_spend(
|
||||
counter_key=f"spend:org:{org_id}",
|
||||
fallback_spend=org_table.spend or 0.0,
|
||||
max_budget=org_max_budget,
|
||||
)
|
||||
|
||||
# Check if organization spend exceeds max budget
|
||||
|
|
@ -4454,6 +4470,8 @@ async def _tag_max_budget_check(
|
|||
tag_spend = await get_current_spend(
|
||||
counter_key=f"spend:tag:{tag_name}",
|
||||
fallback_spend=tag_object.spend or 0.0,
|
||||
max_budget=tag_object.litellm_budget_table.max_budget,
|
||||
fallback_authoritative=True,
|
||||
)
|
||||
if tag_spend <= tag_object.litellm_budget_table.max_budget:
|
||||
continue
|
||||
|
|
|
|||
|
|
@ -1267,6 +1267,14 @@ _MODEL_ROUTING_BODY_TARGET_MODEL_ROUTE_MARKERS = (
|
|||
"/vector_stores",
|
||||
)
|
||||
_MODEL_ROUTING_COMPLETION_MODEL_ROUTE_MARKERS = ("/evals",)
|
||||
# Realtime WebRTC routes carry the effective model inside the nested
|
||||
# ``session.model`` field (see realtime_endpoints.endpoints), so the model the
|
||||
# request will actually use is not present at the top level. Extract it here so
|
||||
# can_key_call_model() validates the real target model.
|
||||
_MODEL_ROUTING_SESSION_MODEL_ROUTE_MARKERS = (
|
||||
"/realtime/client_secrets",
|
||||
"/realtime/calls",
|
||||
)
|
||||
_MODEL_ROUTING_ID_FIELDS = (
|
||||
"file_id",
|
||||
"input_file_id",
|
||||
|
|
@ -1449,6 +1457,12 @@ def _extract_model_candidates_from_request(
|
|||
_append_model_candidates(candidates, body_model)
|
||||
if uses_body_target_model_sources or not body_model:
|
||||
_append_model_candidates(candidates, request_data.get("target_model_names"))
|
||||
if _route_matches_any_marker(
|
||||
route=route, markers=_MODEL_ROUTING_SESSION_MODEL_ROUTE_MARKERS
|
||||
):
|
||||
session = request_data.get("session")
|
||||
if isinstance(session, dict):
|
||||
_append_model_candidates(candidates, session.get("model"))
|
||||
if uses_completion_model_sources and isinstance(
|
||||
request_data.get("completion"), dict
|
||||
):
|
||||
|
|
|
|||
|
|
@ -1837,6 +1837,7 @@ async def _user_api_key_auth_builder(
|
|||
team_member_spend = await get_current_spend(
|
||||
counter_key=f"spend:team_member:{valid_token.user_id}:{valid_token.team_id}",
|
||||
fallback_spend=team_member_spend,
|
||||
max_budget=team_member_budget,
|
||||
)
|
||||
if team_member_spend > team_member_budget:
|
||||
raise litellm.BudgetExceededError(
|
||||
|
|
|
|||
|
|
@ -2513,6 +2513,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
debug_enabled = verbose_proxy_logger.isEnabledFor(logging.DEBUG)
|
||||
stream_completed = False
|
||||
client_disconnected = False
|
||||
delivered_chunk = False
|
||||
try:
|
||||
str_so_far = ""
|
||||
async for (
|
||||
|
|
@ -2529,36 +2530,38 @@ class ProxyBaseLLMRequestProcessing:
|
|||
"async_data_generator: received streaming chunk - %s", chunk
|
||||
)
|
||||
|
||||
if fast_path:
|
||||
yield serialize_chunk(chunk)
|
||||
continue
|
||||
if not fast_path:
|
||||
chunk = await proxy_logging_obj.async_post_call_streaming_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=chunk,
|
||||
data=request_data,
|
||||
str_so_far=str_so_far,
|
||||
)
|
||||
|
||||
chunk = await proxy_logging_obj.async_post_call_streaming_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=chunk,
|
||||
data=request_data,
|
||||
str_so_far=str_so_far,
|
||||
)
|
||||
if isinstance(chunk, (ModelResponse, ModelResponseStream)):
|
||||
response_str = litellm.get_response_string(response_obj=chunk)
|
||||
str_so_far += response_str
|
||||
elif hasattr(chunk, "model_dump"):
|
||||
try:
|
||||
d = chunk.model_dump(mode="json", exclude_none=True)
|
||||
if isinstance(d, dict):
|
||||
str_so_far += str(d.get("content", ""))
|
||||
except Exception:
|
||||
pass
|
||||
elif isinstance(chunk, dict):
|
||||
str_so_far += str(chunk.get("content", ""))
|
||||
|
||||
if isinstance(chunk, (ModelResponse, ModelResponseStream)):
|
||||
response_str = litellm.get_response_string(response_obj=chunk)
|
||||
str_so_far += response_str
|
||||
elif hasattr(chunk, "model_dump"):
|
||||
try:
|
||||
d = chunk.model_dump(mode="json", exclude_none=True)
|
||||
if isinstance(d, dict):
|
||||
str_so_far += str(d.get("content", ""))
|
||||
except Exception:
|
||||
pass
|
||||
elif isinstance(chunk, dict):
|
||||
str_so_far += str(chunk.get("content", ""))
|
||||
|
||||
model_name = request_data.get("model", "")
|
||||
chunk = (
|
||||
ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(
|
||||
model_name = request_data.get("model", "")
|
||||
chunk = ProxyBaseLLMRequestProcessing._process_chunk_with_cost_injection(
|
||||
chunk, model_name
|
||||
)
|
||||
)
|
||||
|
||||
# Set before the yield: an async generator suspends at the yield,
|
||||
# so a GeneratorExit on client disconnect is raised there and any
|
||||
# statement after the yield never runs. The slow-path hook is
|
||||
# awaited above, so a cancellation during it still leaves this
|
||||
# False and refunds.
|
||||
delivered_chunk = True
|
||||
yield serialize_chunk(chunk)
|
||||
stream_completed = True
|
||||
except (asyncio.CancelledError, GeneratorExit):
|
||||
|
|
@ -2573,6 +2576,14 @@ class ProxyBaseLLMRequestProcessing:
|
|||
user_api_key_dict
|
||||
)
|
||||
client_disconnected = True
|
||||
if not delivered_chunk:
|
||||
from litellm.proxy.spend_tracking.budget_reservation import (
|
||||
release_budget_reservation_on_cancel,
|
||||
)
|
||||
|
||||
await release_budget_reservation_on_cancel(
|
||||
getattr(user_api_key_dict, "budget_reservation", None)
|
||||
)
|
||||
raise
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(
|
||||
|
|
|
|||
|
|
@ -10,7 +10,10 @@ url_to_redirect_to += "/login"
|
|||
new_ui_login_url = get_custom_url("", "ui/login")
|
||||
|
||||
|
||||
def build_ui_login_form(show_deprecation_banner: bool = False) -> str:
|
||||
def build_ui_login_form(
|
||||
show_deprecation_banner: bool = False,
|
||||
hide_default_credentials_hint: bool = False,
|
||||
) -> str:
|
||||
banner_html = (
|
||||
f"""
|
||||
<div class="deprecation-banner">
|
||||
|
|
@ -23,6 +26,25 @@ def build_ui_login_form(show_deprecation_banner: bool = False) -> str:
|
|||
else ""
|
||||
)
|
||||
|
||||
info_box_html = (
|
||||
""
|
||||
if hide_default_credentials_hint
|
||||
else """
|
||||
<div class="info-box">
|
||||
<div class="info-header">
|
||||
<svg xmlns="http://www.w3.org/2000/svg" width="16" height="16" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round">
|
||||
<circle cx="12" cy="12" r="10"></circle>
|
||||
<line x1="12" y1="16" x2="12" y2="12"></line>
|
||||
<line x1="12" y1="8" x2="12.01" y2="8"></line>
|
||||
</svg>
|
||||
Default Credentials
|
||||
</div>
|
||||
<p>By default, Username is <code>admin</code> and Password is your set LiteLLM Proxy <code>MASTER_KEY</code>.</p>
|
||||
<p>Need to set UI credentials or SSO? <a href="https://docs.litellm.ai/docs/proxy/ui" target="_blank">Check the documentation</a>.</p>
|
||||
</div>
|
||||
"""
|
||||
)
|
||||
|
||||
return f"""
|
||||
<!DOCTYPE html>
|
||||
<html lang="en">
|
||||
|
|
@ -232,18 +254,7 @@ def build_ui_login_form(show_deprecation_banner: bool = False) -> str:
|
|||
</div>
|
||||
<h2>Login</h2>
|
||||
<p class="subtitle">Access your LiteLLM Admin UI.</p>
|
||||
<div class="info-box">
|
||||
<div class="info-header">
|
||||
<svg xmlns="http://www.w3.org/2000/svg" width="16" height="16" viewBox="0 0 24 24" fill="none" stroke="currentColor" stroke-width="2" stroke-linecap="round" stroke-linejoin="round">
|
||||
<circle cx="12" cy="12" r="10"></circle>
|
||||
<line x1="12" y1="16" x2="12" y2="12"></line>
|
||||
<line x1="12" y1="8" x2="12.01" y2="8"></line>
|
||||
</svg>
|
||||
Default Credentials
|
||||
</div>
|
||||
<p>By default, Username is <code>admin</code> and Password is your set LiteLLM Proxy <code>MASTER_KEY</code>.</p>
|
||||
<p>Need to set UI credentials or SSO? <a href="https://docs.litellm.ai/docs/proxy/ui" target="_blank">Check the documentation</a>.</p>
|
||||
</div>
|
||||
{info_box_html}
|
||||
<label for="username">Username<span class="required">*</span></label>
|
||||
<input type="text" id="username" name="username" required placeholder="Enter your username" autocomplete="username">
|
||||
|
||||
|
|
@ -264,6 +275,3 @@ def build_ui_login_form(show_deprecation_banner: bool = False) -> str:
|
|||
</body>
|
||||
</html>
|
||||
"""
|
||||
|
||||
|
||||
html_form = build_ui_login_form(show_deprecation_banner=True)
|
||||
|
|
|
|||
|
|
@ -1613,199 +1613,215 @@ class DBSpendUpdateWriter:
|
|||
start_time = time.time()
|
||||
|
||||
try:
|
||||
for i in range(n_retry_times + 1):
|
||||
try:
|
||||
# Sort the transactions to minimize the probability of deadlocks by reducing the chance of concurrent
|
||||
# trasactions locking the same rows/ranges in different orders.
|
||||
transactions_to_process = dict(
|
||||
sorted(
|
||||
daily_spend_transactions.items(),
|
||||
# Normally to avoid deadlocks we would sort by the index, but since we have sprinkled indexes
|
||||
# on our schema like we're discount Salt Bae, we just sort by all fields that have an index,
|
||||
# in an ad-hoc (but hopefully sensible) order of indexes. The actual ordering matters less than
|
||||
# ensuring that all concurrent transactions sort in the same order.
|
||||
# We could in theory use the dict key, as it contains basically the same fields, but this is more
|
||||
# robust to future changes in the key format.
|
||||
# If _update_daily_spend ever gets the ability to write to multiple tables at once, the sorting
|
||||
# should sort by the table first.
|
||||
key=lambda x: (
|
||||
x[1].get("date") or "",
|
||||
x[1].get(entity_id_field) or "",
|
||||
x[1].get("api_key") or "",
|
||||
x[1].get("model") or "",
|
||||
x[1].get("custom_llm_provider") or "",
|
||||
),
|
||||
)[:BATCH_SIZE]
|
||||
)
|
||||
|
||||
if len(transactions_to_process) == 0:
|
||||
verbose_proxy_logger.debug(
|
||||
f"No new transactions to process for daily {entity_type} spend update"
|
||||
)
|
||||
break
|
||||
|
||||
while daily_spend_transactions:
|
||||
for i in range(n_retry_times + 1):
|
||||
try:
|
||||
async with prisma_client.db.batch_() as batcher:
|
||||
for _, transaction in transactions_to_process.items():
|
||||
entity_id = transaction.get(entity_id_field)
|
||||
# Sort the transactions to minimize the probability of deadlocks by reducing the chance of concurrent
|
||||
# trasactions locking the same rows/ranges in different orders.
|
||||
transactions_to_process = dict(
|
||||
sorted(
|
||||
daily_spend_transactions.items(),
|
||||
# Normally to avoid deadlocks we would sort by the index, but since we have sprinkled indexes
|
||||
# on our schema like we're discount Salt Bae, we just sort by all fields that have an index,
|
||||
# in an ad-hoc (but hopefully sensible) order of indexes. The actual ordering matters less than
|
||||
# ensuring that all concurrent transactions sort in the same order.
|
||||
# We could in theory use the dict key, as it contains basically the same fields, but this is more
|
||||
# robust to future changes in the key format.
|
||||
# If _update_daily_spend ever gets the ability to write to multiple tables at once, the sorting
|
||||
# should sort by the table first.
|
||||
key=lambda x: (
|
||||
x[1].get("date") or "",
|
||||
x[1].get(entity_id_field) or "",
|
||||
x[1].get("api_key") or "",
|
||||
x[1].get("model") or "",
|
||||
x[1].get("custom_llm_provider") or "",
|
||||
),
|
||||
)[:BATCH_SIZE]
|
||||
)
|
||||
|
||||
# Construct the where clause dynamically
|
||||
where_clause = {
|
||||
unique_constraint_name: {
|
||||
if len(transactions_to_process) == 0:
|
||||
verbose_proxy_logger.debug(
|
||||
f"No new transactions to process for daily {entity_type} spend update"
|
||||
)
|
||||
return
|
||||
|
||||
try:
|
||||
async with prisma_client.db.batch_() as batcher:
|
||||
for _, transaction in transactions_to_process.items():
|
||||
entity_id = transaction.get(entity_id_field)
|
||||
|
||||
# Construct the where clause dynamically
|
||||
where_clause = {
|
||||
unique_constraint_name: {
|
||||
entity_id_field: entity_id,
|
||||
"date": transaction["date"],
|
||||
"api_key": transaction["api_key"],
|
||||
"model": transaction["model"],
|
||||
"custom_llm_provider": transaction.get(
|
||||
"custom_llm_provider"
|
||||
)
|
||||
or "",
|
||||
"mcp_namespaced_tool_name": transaction.get(
|
||||
"mcp_namespaced_tool_name"
|
||||
)
|
||||
or "",
|
||||
"endpoint": transaction.get("endpoint")
|
||||
or "",
|
||||
}
|
||||
}
|
||||
|
||||
# Get the table dynamically
|
||||
table = getattr(batcher, table_name)
|
||||
|
||||
# Common data structure for both create and update
|
||||
common_data = {
|
||||
entity_id_field: entity_id,
|
||||
"date": transaction["date"],
|
||||
"api_key": transaction["api_key"],
|
||||
"model": transaction["model"],
|
||||
"custom_llm_provider": transaction.get(
|
||||
"custom_llm_provider"
|
||||
)
|
||||
or "",
|
||||
"model": transaction.get("model"),
|
||||
"model_group": transaction.get("model_group"),
|
||||
"mcp_namespaced_tool_name": transaction.get(
|
||||
"mcp_namespaced_tool_name"
|
||||
)
|
||||
or "",
|
||||
"custom_llm_provider": transaction.get(
|
||||
"custom_llm_provider"
|
||||
),
|
||||
"endpoint": transaction.get("endpoint") or "",
|
||||
}
|
||||
}
|
||||
|
||||
# Get the table dynamically
|
||||
table = getattr(batcher, table_name)
|
||||
|
||||
# Common data structure for both create and update
|
||||
common_data = {
|
||||
entity_id_field: entity_id,
|
||||
"date": transaction["date"],
|
||||
"api_key": transaction["api_key"],
|
||||
"model": transaction.get("model"),
|
||||
"model_group": transaction.get("model_group"),
|
||||
"mcp_namespaced_tool_name": transaction.get(
|
||||
"mcp_namespaced_tool_name"
|
||||
)
|
||||
or "",
|
||||
"custom_llm_provider": transaction.get(
|
||||
"custom_llm_provider"
|
||||
),
|
||||
"endpoint": transaction.get("endpoint") or "",
|
||||
"prompt_tokens": transaction["prompt_tokens"],
|
||||
"completion_tokens": transaction[
|
||||
"completion_tokens"
|
||||
],
|
||||
"spend": transaction["spend"],
|
||||
"api_requests": transaction["api_requests"],
|
||||
"successful_requests": transaction[
|
||||
"successful_requests"
|
||||
],
|
||||
"failed_requests": transaction["failed_requests"],
|
||||
}
|
||||
|
||||
# Add cache-related fields if they exist
|
||||
if "cache_read_input_tokens" in transaction:
|
||||
common_data["cache_read_input_tokens"] = (
|
||||
transaction.get("cache_read_input_tokens", 0)
|
||||
)
|
||||
if "cache_creation_input_tokens" in transaction:
|
||||
common_data["cache_creation_input_tokens"] = (
|
||||
transaction.get(
|
||||
"cache_creation_input_tokens", 0
|
||||
)
|
||||
)
|
||||
|
||||
if entity_type == "tag" and "request_id" in transaction:
|
||||
common_data["request_id"] = transaction.get(
|
||||
"request_id"
|
||||
)
|
||||
|
||||
# Create update data structure
|
||||
update_data = {
|
||||
"prompt_tokens": {
|
||||
"increment": transaction["prompt_tokens"]
|
||||
},
|
||||
"completion_tokens": {
|
||||
"increment": transaction["completion_tokens"]
|
||||
},
|
||||
"spend": {"increment": transaction["spend"]},
|
||||
"api_requests": {
|
||||
"increment": transaction["api_requests"]
|
||||
},
|
||||
"successful_requests": {
|
||||
"increment": transaction["successful_requests"]
|
||||
},
|
||||
"failed_requests": {
|
||||
"increment": transaction["failed_requests"]
|
||||
},
|
||||
}
|
||||
|
||||
# Add cache-related fields to update if they exist
|
||||
if "cache_read_input_tokens" in transaction:
|
||||
update_data["cache_read_input_tokens"] = {
|
||||
"increment": transaction.get(
|
||||
"cache_read_input_tokens", 0
|
||||
)
|
||||
}
|
||||
if "cache_creation_input_tokens" in transaction:
|
||||
update_data["cache_creation_input_tokens"] = {
|
||||
"increment": transaction.get(
|
||||
"cache_creation_input_tokens", 0
|
||||
)
|
||||
"prompt_tokens": transaction["prompt_tokens"],
|
||||
"completion_tokens": transaction[
|
||||
"completion_tokens"
|
||||
],
|
||||
"spend": transaction["spend"],
|
||||
"api_requests": transaction["api_requests"],
|
||||
"successful_requests": transaction[
|
||||
"successful_requests"
|
||||
],
|
||||
"failed_requests": transaction[
|
||||
"failed_requests"
|
||||
],
|
||||
}
|
||||
|
||||
if entity_type == "tag" and "request_id" in transaction:
|
||||
update_data["request_id"] = transaction.get(
|
||||
"request_id"
|
||||
# Add cache-related fields if they exist
|
||||
if "cache_read_input_tokens" in transaction:
|
||||
common_data["cache_read_input_tokens"] = (
|
||||
transaction.get(
|
||||
"cache_read_input_tokens", 0
|
||||
)
|
||||
)
|
||||
if "cache_creation_input_tokens" in transaction:
|
||||
common_data["cache_creation_input_tokens"] = (
|
||||
transaction.get(
|
||||
"cache_creation_input_tokens", 0
|
||||
)
|
||||
)
|
||||
|
||||
if (
|
||||
entity_type == "tag"
|
||||
and "request_id" in transaction
|
||||
):
|
||||
common_data["request_id"] = transaction.get(
|
||||
"request_id"
|
||||
)
|
||||
|
||||
# Create update data structure
|
||||
update_data = {
|
||||
"prompt_tokens": {
|
||||
"increment": transaction["prompt_tokens"]
|
||||
},
|
||||
"completion_tokens": {
|
||||
"increment": transaction[
|
||||
"completion_tokens"
|
||||
]
|
||||
},
|
||||
"spend": {"increment": transaction["spend"]},
|
||||
"api_requests": {
|
||||
"increment": transaction["api_requests"]
|
||||
},
|
||||
"successful_requests": {
|
||||
"increment": transaction[
|
||||
"successful_requests"
|
||||
]
|
||||
},
|
||||
"failed_requests": {
|
||||
"increment": transaction["failed_requests"]
|
||||
},
|
||||
}
|
||||
|
||||
# Add cache-related fields to update if they exist
|
||||
if "cache_read_input_tokens" in transaction:
|
||||
update_data["cache_read_input_tokens"] = {
|
||||
"increment": transaction.get(
|
||||
"cache_read_input_tokens", 0
|
||||
)
|
||||
}
|
||||
if "cache_creation_input_tokens" in transaction:
|
||||
update_data["cache_creation_input_tokens"] = {
|
||||
"increment": transaction.get(
|
||||
"cache_creation_input_tokens", 0
|
||||
)
|
||||
}
|
||||
|
||||
if (
|
||||
entity_type == "tag"
|
||||
and "request_id" in transaction
|
||||
):
|
||||
update_data["request_id"] = transaction.get(
|
||||
"request_id"
|
||||
)
|
||||
|
||||
# Add endpoint to update_data so existing rows get their endpoint field updated
|
||||
update_data["endpoint"] = (
|
||||
transaction.get("endpoint") or ""
|
||||
)
|
||||
|
||||
# Add endpoint to update_data so existing rows get their endpoint field updated
|
||||
update_data["endpoint"] = (
|
||||
transaction.get("endpoint") or ""
|
||||
)
|
||||
table.upsert(
|
||||
where=where_clause,
|
||||
data={
|
||||
"create": common_data,
|
||||
"update": update_data,
|
||||
},
|
||||
)
|
||||
except Exception as batch_error:
|
||||
# Log detailed error information for debugging batch upsert failures
|
||||
# This helps diagnose issues like unique constraint violations
|
||||
spend_log_error(
|
||||
"Daily %s spend batch upsert failed. "
|
||||
"Table: %s, Constraint: %s, Batch size: %d, Error: %s",
|
||||
entity_type,
|
||||
table_name,
|
||||
unique_constraint_name,
|
||||
len(transactions_to_process),
|
||||
str(batch_error),
|
||||
exc=batch_error,
|
||||
)
|
||||
raise
|
||||
|
||||
table.upsert(
|
||||
where=where_clause,
|
||||
data={
|
||||
"create": common_data,
|
||||
"update": update_data,
|
||||
},
|
||||
)
|
||||
except Exception as batch_error:
|
||||
# Log detailed error information for debugging batch upsert failures
|
||||
# This helps diagnose issues like unique constraint violations
|
||||
spend_log_error(
|
||||
"Daily %s spend batch upsert failed. "
|
||||
"Table: %s, Constraint: %s, Batch size: %d, Error: %s",
|
||||
entity_type,
|
||||
table_name,
|
||||
unique_constraint_name,
|
||||
len(transactions_to_process),
|
||||
str(batch_error),
|
||||
exc=batch_error,
|
||||
verbose_proxy_logger.debug(
|
||||
f"Processed {len(transactions_to_process)} daily {entity_type} transactions in {time.time() - start_time:.2f}s"
|
||||
)
|
||||
raise
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f"Processed {len(transactions_to_process)} daily {entity_type} transactions in {time.time() - start_time:.2f}s"
|
||||
)
|
||||
# Remove processed transactions
|
||||
for key in transactions_to_process.keys():
|
||||
daily_spend_transactions.pop(key, None)
|
||||
|
||||
# Remove processed transactions
|
||||
for key in transactions_to_process.keys():
|
||||
daily_spend_transactions.pop(key, None)
|
||||
break
|
||||
|
||||
break
|
||||
|
||||
except DB_CONNECTION_ERROR_TYPES as e:
|
||||
if i >= n_retry_times:
|
||||
_raise_failed_update_spend_exception(
|
||||
e=e,
|
||||
start_time=start_time,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
except DB_CONNECTION_ERROR_TYPES as e:
|
||||
if i >= n_retry_times:
|
||||
_raise_failed_update_spend_exception(
|
||||
e=e,
|
||||
start_time=start_time,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
await asyncio.sleep(
|
||||
# Sleep a random amount to avoid retrying and deadlocking again: when two transactions deadlock they are
|
||||
# cancelled basically at the same time, so if they wait the same time they will also retry at the same time
|
||||
# and thus they are more likely to deadlock again.
|
||||
# Instead, we sleep a random amount so that they retry at slightly different times, lowering the chance of
|
||||
# repeated deadlocks, and therefore of exceeding the retry limit.
|
||||
random.uniform(2**i, 2 ** (i + 1))
|
||||
)
|
||||
await asyncio.sleep(
|
||||
# Sleep a random amount to avoid retrying and deadlocking again: when two transactions deadlock they are
|
||||
# cancelled basically at the same time, so if they wait the same time they will also retry at the same time
|
||||
# and thus they are more likely to deadlock again.
|
||||
# Instead, we sleep a random amount so that they retry at slightly different times, lowering the chance of
|
||||
# repeated deadlocks, and therefore of exceeding the retry limit.
|
||||
random.uniform(2**i, 2 ** (i + 1))
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
if "transactions_to_process" in locals():
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import asyncio
|
||||
import json
|
||||
from litellm._uuid import uuid
|
||||
from typing import TYPE_CHECKING, Any, Optional
|
||||
|
||||
|
|
@ -167,8 +168,11 @@ end
|
|||
self._release_lock_script = script_register(
|
||||
self._COMPARE_AND_DELETE_LOCK_SCRIPT
|
||||
)
|
||||
# acquire_lock stores the pod_id via async_set_cache, which
|
||||
# JSON-encodes the value; compare against the same encoding so
|
||||
# the Lua equality check matches and the lock is released
|
||||
result = await self._release_lock_script(
|
||||
keys=[lock_key], args=[self.pod_id]
|
||||
keys=[lock_key], args=[json.dumps(self.pod_id)]
|
||||
)
|
||||
return int(result or 0)
|
||||
except Exception:
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ import urllib
|
|||
import urllib.parse
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any, Dict, Optional, Union
|
||||
from typing import Any, Callable, Union
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.secret_managers.main import str_to_bool
|
||||
|
|
@ -31,7 +31,7 @@ class IAMEndpoint:
|
|||
port: str
|
||||
user: str
|
||||
name: str
|
||||
schema: Optional[str] = None
|
||||
schema: str | None = None
|
||||
|
||||
def build_url(self, token: str) -> str:
|
||||
url = f"postgresql://{self.user}:{token}@{self.host}:{self.port}/{self.name}"
|
||||
|
|
@ -53,7 +53,7 @@ def parse_iam_endpoint_from_url(url: str) -> IAMEndpoint:
|
|||
if not name:
|
||||
raise ValueError("Cannot parse IAM endpoint from URL: missing database name")
|
||||
port = str(parsed.port) if parsed.port else "5432"
|
||||
schema: Optional[str] = None
|
||||
schema: str | None = None
|
||||
if parsed.query:
|
||||
qs = urllib.parse.parse_qs(parsed.query)
|
||||
schema_vals = qs.get("schema")
|
||||
|
|
@ -94,7 +94,7 @@ class PrismaWrapper:
|
|||
iam_token_db_auth: bool,
|
||||
*,
|
||||
db_url_env_var: str = "DATABASE_URL",
|
||||
iam_endpoint: Optional[IAMEndpoint] = None,
|
||||
iam_endpoint: IAMEndpoint | None = None,
|
||||
recreate_uses_datasource: bool = False,
|
||||
log_prefix: str = "",
|
||||
):
|
||||
|
|
@ -116,9 +116,25 @@ class PrismaWrapper:
|
|||
self._log_prefix = f"{log_prefix} " if log_prefix else ""
|
||||
|
||||
# Background token refresh task management
|
||||
self._token_refresh_task: Optional[asyncio.Task] = None
|
||||
self._token_refresh_task: asyncio.Task | None = None
|
||||
self._reconnection_lock = asyncio.Lock()
|
||||
self._last_refresh_time: Optional[datetime] = None
|
||||
self._last_refresh_time: datetime | None = None
|
||||
|
||||
# Coordination for planned engine restarts (issue #29176). Every
|
||||
# `recreate_prisma_client` SIGTERMs the running query-engine on
|
||||
# purpose. The engine-death watcher (in `PrismaClient`) must be able
|
||||
# to tell that planned kill apart from a real crash, otherwise it
|
||||
# triggers its own reconnect and kills the freshly-spawned engine.
|
||||
# - `_expected_engine_deaths`: PIDs we intentionally killed; the
|
||||
# watcher consumes these instead of reconnecting.
|
||||
# - `_engine_generation`: monotonic counter bumped on every
|
||||
# successful recreate, used by callers as an optimistic-lock token
|
||||
# so racing/cascading recreates collapse into a single restart.
|
||||
# - `on_engine_replaced`: optional callback fired after a recreate so
|
||||
# the owner (PrismaClient) can re-arm its watcher on the new PID.
|
||||
self._expected_engine_deaths: set[int] = set()
|
||||
self._engine_generation: int = 0
|
||||
self.on_engine_replaced: Callable[[], None] | None = None
|
||||
|
||||
def _get_engine_pid(self) -> int:
|
||||
"""Get the PID of the current Prisma engine subprocess, or 0 if unavailable."""
|
||||
|
|
@ -167,7 +183,7 @@ class PrismaWrapper:
|
|||
except (ProcessLookupError, PermissionError, OSError):
|
||||
pass # Exited after SIGTERM — expected
|
||||
|
||||
def _extract_token_from_db_url(self, db_url: Optional[str]) -> Optional[str]:
|
||||
def _extract_token_from_db_url(self, db_url: str | None) -> str | None:
|
||||
"""
|
||||
Extract the token (password) from the DATABASE_URL.
|
||||
|
||||
|
|
@ -188,7 +204,7 @@ class PrismaWrapper:
|
|||
except Exception:
|
||||
return None
|
||||
|
||||
def _parse_token_expiration(self, token: Optional[str]) -> Optional[datetime]:
|
||||
def _parse_token_expiration(self, token: str | None) -> datetime | None:
|
||||
"""
|
||||
Parse the token to extract its expiration time.
|
||||
|
||||
|
|
@ -255,7 +271,7 @@ class PrismaWrapper:
|
|||
# If already past refresh time, return 0 (refresh immediately)
|
||||
return max(0, seconds_until_refresh)
|
||||
|
||||
def is_token_expired(self, token_url: Optional[str]) -> bool:
|
||||
def is_token_expired(self, token_url: str | None) -> bool:
|
||||
"""Check if the token in the given URL is expired."""
|
||||
if token_url is None:
|
||||
return True
|
||||
|
|
@ -272,7 +288,7 @@ class PrismaWrapper:
|
|||
|
||||
return datetime.utcnow() > expiration_time
|
||||
|
||||
def get_rds_iam_token(self) -> Optional[str]:
|
||||
def get_rds_iam_token(self) -> str | None:
|
||||
"""Generate a new RDS IAM token and update the configured DB URL env var.
|
||||
|
||||
When the wrapper was constructed with an explicit `iam_endpoint`
|
||||
|
|
@ -313,8 +329,12 @@ class PrismaWrapper:
|
|||
return _db_url
|
||||
|
||||
async def recreate_prisma_client(
|
||||
self, new_db_url: str, http_client: Optional[Any] = None
|
||||
):
|
||||
self,
|
||||
new_db_url: str,
|
||||
http_client: Any | None = None,
|
||||
*,
|
||||
expected_generation: int | None = None,
|
||||
) -> bool:
|
||||
"""Disconnect and reconnect the Prisma client with a new database URL.
|
||||
|
||||
Kills the old engine subprocess directly (SIGTERM → SIGKILL) rather than
|
||||
|
|
@ -327,14 +347,70 @@ class PrismaWrapper:
|
|||
the reader wrapper opts into `recreate_uses_datasource=True` so the
|
||||
new URL is passed explicitly via `datasource={"url": ...}` (Prisma
|
||||
does not auto-read alternate env vars like DATABASE_URL_READ_REPLICA).
|
||||
|
||||
Serializes all recreations through `self._reconnection_lock` so the
|
||||
IAM-refresh path and the engine-death/transport-error reconnect paths
|
||||
cannot recreate concurrently (issue #29176). `expected_generation`, if
|
||||
given, is an optimistic-lock token: when it no longer matches
|
||||
`self._engine_generation` once the lock is held, another path already
|
||||
replaced the engine, so this call is a no-op and returns ``False``.
|
||||
|
||||
Returns:
|
||||
bool: ``True`` if the client was actually recreated, ``False`` if
|
||||
the recreate was skipped because the engine generation moved on.
|
||||
"""
|
||||
async with self._reconnection_lock:
|
||||
return await self._recreate_prisma_client_locked(
|
||||
new_db_url,
|
||||
http_client=http_client,
|
||||
expected_generation=expected_generation,
|
||||
)
|
||||
|
||||
async def _recreate_prisma_client_locked(
|
||||
self,
|
||||
new_db_url: str,
|
||||
http_client: Any | None = None,
|
||||
*,
|
||||
expected_generation: int | None = None,
|
||||
) -> bool:
|
||||
"""Core recreate logic. Caller MUST hold `self._reconnection_lock`.
|
||||
|
||||
Split out so callers that already hold the lock (e.g.
|
||||
`_safe_refresh_token`, which double-checks token freshness under the
|
||||
lock) don't re-acquire it — `asyncio.Lock` is not reentrant.
|
||||
"""
|
||||
from prisma import Prisma # type: ignore
|
||||
|
||||
if (
|
||||
expected_generation is not None
|
||||
and expected_generation != self._engine_generation
|
||||
):
|
||||
verbose_proxy_logger.info(
|
||||
"%sSkipping Prisma client recreate: engine already replaced "
|
||||
"(generation %s != expected %s).",
|
||||
self._log_prefix,
|
||||
self._engine_generation,
|
||||
expected_generation,
|
||||
)
|
||||
return False
|
||||
|
||||
old_engine_pid = self._get_engine_pid()
|
||||
if old_engine_pid > 0:
|
||||
# Record BEFORE the kill so the engine-death watcher, which may
|
||||
# fire the instant the process dies, recognizes this as a planned
|
||||
# restart and does not launch its own reconnect.
|
||||
#
|
||||
# A stale entry can linger when the watcher re-arms on the new PID
|
||||
# before the old PID's death callback runs (the callback then
|
||||
# early-returns on PID mismatch without consuming it). Such entries
|
||||
# are harmless but would accumulate on a long-running proxy (~one
|
||||
# per IAM refresh), so cap the set — those old PIDs are long dead.
|
||||
if len(self._expected_engine_deaths) >= 64:
|
||||
self._expected_engine_deaths.clear()
|
||||
self._expected_engine_deaths.add(old_engine_pid)
|
||||
await self._kill_engine_process(old_engine_pid)
|
||||
|
||||
kwargs: Dict[str, Any] = {}
|
||||
kwargs: dict[str, Any] = {}
|
||||
if http_client is not None:
|
||||
kwargs["http"] = http_client
|
||||
if self._recreate_uses_datasource:
|
||||
|
|
@ -342,6 +418,15 @@ class PrismaWrapper:
|
|||
self._original_prisma = Prisma(**kwargs)
|
||||
|
||||
await self._original_prisma.connect()
|
||||
self._engine_generation += 1
|
||||
|
||||
# Let the owner (PrismaClient) re-arm its engine-death watcher on the
|
||||
# newly-spawned engine PID. Scheduled, never awaited, so a slow watcher
|
||||
# can't stall the refresh while we hold the reconnection lock.
|
||||
if self.on_engine_replaced is not None:
|
||||
self.on_engine_replaced()
|
||||
|
||||
return True
|
||||
|
||||
async def start_token_refresh_task(self) -> None:
|
||||
"""
|
||||
|
|
@ -441,9 +526,23 @@ class PrismaWrapper:
|
|||
preventing multiple concurrent reconnection attempts.
|
||||
"""
|
||||
async with self._reconnection_lock:
|
||||
# Double-checked under the lock: another trigger (e.g. the
|
||||
# proactive loop racing a __getattr__ fallback) may have already
|
||||
# refreshed while we waited. Recreating again would needlessly kill
|
||||
# the engine that refresh just spawned (issue #29176), so coalesce
|
||||
# by skipping when the current token still has comfortable runway.
|
||||
if self._token_refresh_not_needed(os.getenv(self._db_url_env_var)):
|
||||
verbose_proxy_logger.debug(
|
||||
"%sRDS IAM token still fresh; skipping redundant refresh.",
|
||||
self._log_prefix,
|
||||
)
|
||||
return
|
||||
|
||||
new_db_url = self.get_rds_iam_token()
|
||||
if new_db_url:
|
||||
await self.recreate_prisma_client(new_db_url)
|
||||
# We already hold `_reconnection_lock`; call the locked core
|
||||
# directly (the public method would re-acquire and deadlock).
|
||||
await self._recreate_prisma_client_locked(new_db_url)
|
||||
self._last_refresh_time = datetime.utcnow()
|
||||
verbose_proxy_logger.info(
|
||||
"%sRDS IAM token refreshed successfully. New token valid for ~15 minutes.",
|
||||
|
|
@ -455,6 +554,23 @@ class PrismaWrapper:
|
|||
self._log_prefix,
|
||||
)
|
||||
|
||||
def _token_refresh_not_needed(self, token_url: str | None) -> bool:
|
||||
"""True iff the token in ``token_url`` has more than the refresh buffer
|
||||
of runway left, so a refresh would be redundant.
|
||||
|
||||
Used to coalesce stacked refresh triggers. Deliberately mirrors the
|
||||
proactive loop's schedule (refresh at ``expiration - buffer``): a token
|
||||
with exactly ``buffer`` seconds left is NOT considered fresh, so the
|
||||
legitimate proactive refresh still fires. Unparseable tokens return
|
||||
``False`` (refresh) — skipping them would mean never refreshing.
|
||||
"""
|
||||
token = self._extract_token_from_db_url(token_url)
|
||||
expiration_time = self._parse_token_expiration(token)
|
||||
if expiration_time is None:
|
||||
return False
|
||||
seconds_left = (expiration_time - datetime.utcnow()).total_seconds()
|
||||
return seconds_left > self.TOKEN_REFRESH_BUFFER_SECONDS
|
||||
|
||||
def __getattr__(self, name: str):
|
||||
"""
|
||||
Proxy attribute access to the underlying Prisma client.
|
||||
|
|
@ -598,7 +714,7 @@ class PrismaManager:
|
|||
|
||||
|
||||
def should_update_prisma_schema(
|
||||
disable_updates: Optional[Union[bool, str]] = None,
|
||||
disable_updates: Union[bool, str] | None = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Determines if Prisma Schema updates should be applied during startup.
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ otherwise PrismaClient uses the writer-only PrismaWrapper directly.
|
|||
"""
|
||||
|
||||
import os
|
||||
from typing import Any, Callable, Optional
|
||||
from typing import Any, Callable
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy.db.prisma_client import PrismaWrapper
|
||||
|
|
@ -117,7 +117,7 @@ class RoutingPrismaWrapper:
|
|||
)
|
||||
|
||||
async def disconnect(self, *args: Any, **kwargs: Any) -> None:
|
||||
first_error: Optional[BaseException] = None
|
||||
first_error: BaseException | None = None
|
||||
for client in (self._writer, self._reader):
|
||||
try:
|
||||
await client.disconnect(*args, **kwargs)
|
||||
|
|
@ -144,8 +144,12 @@ class RoutingPrismaWrapper:
|
|||
await self._reader.stop_token_refresh_task()
|
||||
|
||||
async def recreate_prisma_client(
|
||||
self, new_db_url: str, http_client: Optional[Any] = None
|
||||
) -> None:
|
||||
self,
|
||||
new_db_url: str,
|
||||
http_client: Any | None = None,
|
||||
*,
|
||||
expected_generation: int | None = None,
|
||||
) -> bool:
|
||||
"""Recreate both writer and reader Prisma clients.
|
||||
|
||||
The writer reconnect path in PrismaClient calls
|
||||
|
|
@ -155,8 +159,19 @@ class RoutingPrismaWrapper:
|
|||
the writer first (its URL is the one passed in), then best-effort
|
||||
recreate the reader. A reader failure flips `_reader_unavailable=True`
|
||||
so reads transparently fall through to the writer.
|
||||
|
||||
`expected_generation` is forwarded to the writer's optimistic-lock
|
||||
guard. If the writer recreate is skipped (another path already replaced
|
||||
the engine — issue #29176), we skip the reader too rather than churning
|
||||
it needlessly, and return ``False``.
|
||||
"""
|
||||
await self._writer.recreate_prisma_client(new_db_url, http_client=http_client)
|
||||
writer_recreated = await self._writer.recreate_prisma_client(
|
||||
new_db_url,
|
||||
http_client=http_client,
|
||||
expected_generation=expected_generation,
|
||||
)
|
||||
if not writer_recreated:
|
||||
return False
|
||||
try:
|
||||
await self._recreate_reader(http_client=http_client)
|
||||
self._reader_unavailable = False
|
||||
|
|
@ -167,8 +182,9 @@ class RoutingPrismaWrapper:
|
|||
"Reads will fall back to the writer until the reader recovers.",
|
||||
e,
|
||||
)
|
||||
return True
|
||||
|
||||
async def _recreate_reader(self, http_client: Optional[Any] = None) -> None:
|
||||
async def _recreate_reader(self, http_client: Any | None = None) -> None:
|
||||
"""Resolve the reader URL and recreate its Prisma client.
|
||||
|
||||
IAM-enabled readers regenerate their token (host/port/user came from
|
||||
|
|
|
|||
|
|
@ -3,12 +3,12 @@ model_list:
|
|||
litellm_params:
|
||||
model: openai/fake
|
||||
api_key: fake-key
|
||||
api_base: https://exampleopenaiendpoint-production.up.railway.app/
|
||||
api_base: os.environ/FAKE_OPENAI_API_BASE
|
||||
- model_name: gpt-4
|
||||
litellm_params:
|
||||
model: openai/gpt-4
|
||||
api_key: fake-key
|
||||
api_base: https://exampleopenaiendpoint-production.up.railway.app/
|
||||
api_base: os.environ/FAKE_OPENAI_API_BASE
|
||||
|
||||
litellm_settings:
|
||||
callbacks: ["gcs_bucket"]
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ model_list:
|
|||
litellm_params:
|
||||
model: openai/fake
|
||||
api_key: fake-key
|
||||
api_base: https://exampleopenaiendpoint-production.up.railway.app/
|
||||
api_base: os.environ/FAKE_OPENAI_API_BASE
|
||||
tags: ["teamA"]
|
||||
model_info:
|
||||
id: "team-a-model"
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ model_list:
|
|||
litellm_params:
|
||||
model: openai/my-fake-model
|
||||
api_key: my-fake-key
|
||||
api_base: https://exampleopenaiendpoint-production.up.railway.app/
|
||||
api_base: os.environ/FAKE_OPENAI_API_BASE
|
||||
|
||||
litellm_settings:
|
||||
cache: True
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ model_list:
|
|||
litellm_params:
|
||||
model: openai/gpt-5-mini
|
||||
api_key: fake-key
|
||||
api_base: https://exampleopenaiendpoint-production.up.railway.app/
|
||||
api_base: os.environ/FAKE_OPENAI_API_BASE
|
||||
tags: ["teamA"]
|
||||
model_info:
|
||||
id: "team-a-model"
|
||||
|
|
@ -11,7 +11,7 @@ model_list:
|
|||
litellm_params:
|
||||
model: openai/gpt-5-mini
|
||||
api_key: fake-key
|
||||
api_base: https://exampleopenaiendpoint-production.up.railway.app/
|
||||
api_base: os.environ/FAKE_OPENAI_API_BASE
|
||||
tags: ["teamB"]
|
||||
model_info:
|
||||
id: "team-b-model"
|
||||
|
|
@ -24,7 +24,7 @@ model_list:
|
|||
litellm_params:
|
||||
model: openai/429
|
||||
api_key: fake-key
|
||||
api_base: https://exampleopenaiendpoint-production.up.railway.app
|
||||
api_base: os.environ/FAKE_OPENAI_API_BASE
|
||||
- model_name: llava-hf
|
||||
litellm_params:
|
||||
model: openai/llava-hf/llava-v1.6-vicuna-7b-hf
|
||||
|
|
@ -35,12 +35,12 @@ model_list:
|
|||
- model_name: bedrock/*
|
||||
litellm_params:
|
||||
model: bedrock/*
|
||||
api_base: https://exampleopenaiendpoint-production.up.railway.app/
|
||||
api_base: os.environ/FAKE_OPENAI_API_BASE
|
||||
- model_name: openai/*
|
||||
litellm_params:
|
||||
model: openai/*
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
api_base: https://exampleopenaiendpoint-production.up.railway.app/
|
||||
api_base: os.environ/FAKE_OPENAI_API_BASE
|
||||
|
||||
|
||||
litellm_settings:
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ model_list:
|
|||
litellm_params:
|
||||
model: openai/gpt-5-mini
|
||||
api_key: fake-key
|
||||
api_base: https://exampleopenaiendpoint-production.up.railway.app/
|
||||
api_base: os.environ/FAKE_OPENAI_API_BASE
|
||||
|
||||
general_settings:
|
||||
use_redis_transaction_buffer: true
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ model_list:
|
|||
litellm_params:
|
||||
model: openai/my-fake-model
|
||||
api_key: my-fake-key
|
||||
api_base: https://exampleopenaiendpoint-production.up.railway.app/
|
||||
api_base: os.environ/FAKE_OPENAI_API_BASE
|
||||
|
||||
general_settings:
|
||||
store_model_in_db: true
|
||||
|
|
|
|||
|
|
@ -0,0 +1,51 @@
|
|||
from typing import TYPE_CHECKING, Union
|
||||
|
||||
from litellm.types.guardrails import (
|
||||
GuardrailEventHooks,
|
||||
Mode,
|
||||
SupportedGuardrailIntegrations,
|
||||
)
|
||||
|
||||
from .repelloai import RepelloAIGuardrail
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.types.guardrails import Guardrail, LitellmParams
|
||||
|
||||
|
||||
def _event_hook_from_mode(
|
||||
mode: str | list[str] | Mode,
|
||||
) -> Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode]:
|
||||
if isinstance(mode, Mode):
|
||||
return mode
|
||||
if isinstance(mode, list):
|
||||
return [GuardrailEventHooks(item) for item in mode]
|
||||
return GuardrailEventHooks(mode)
|
||||
|
||||
|
||||
def initialize_guardrail(
|
||||
litellm_params: "LitellmParams", guardrail: "Guardrail"
|
||||
) -> RepelloAIGuardrail:
|
||||
import litellm
|
||||
|
||||
_repelloai_callback = RepelloAIGuardrail(
|
||||
guardrail_name=guardrail["guardrail_name"],
|
||||
api_key=litellm_params.api_key,
|
||||
api_base=litellm_params.api_base,
|
||||
asset_id=litellm_params.asset_id,
|
||||
unreachable_fallback=litellm_params.unreachable_fallback,
|
||||
event_hook=_event_hook_from_mode(litellm_params.mode),
|
||||
default_on=litellm_params.default_on or False,
|
||||
)
|
||||
litellm.logging_callback_manager.add_litellm_callback(_repelloai_callback)
|
||||
|
||||
return _repelloai_callback
|
||||
|
||||
|
||||
guardrail_initializer_registry = {
|
||||
SupportedGuardrailIntegrations.REPELLOAI.value: initialize_guardrail,
|
||||
}
|
||||
|
||||
|
||||
guardrail_class_registry = {
|
||||
SupportedGuardrailIntegrations.REPELLOAI.value: RepelloAIGuardrail,
|
||||
}
|
||||
613
litellm/proxy/guardrails/guardrail_hooks/repelloai/repelloai.py
Normal file
613
litellm/proxy/guardrails/guardrail_hooks/repelloai/repelloai.py
Normal file
|
|
@ -0,0 +1,613 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import AsyncGenerator, Literal
|
||||
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
from pydantic import BaseModel
|
||||
from typing_extensions import TypeGuard
|
||||
|
||||
from fastapi import HTTPException
|
||||
from httpx import HTTPError, Response as HttpxResponse
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client, # pyright: ignore[reportUnknownVariableType]
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.common_utils.callback_utils import (
|
||||
add_guardrail_to_applied_guardrails_header, # pyright: ignore[reportUnknownVariableType]
|
||||
)
|
||||
from litellm.proxy.guardrails._content_utils import build_inspection_messages
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.guardrails import GuardrailEventHooks, Mode
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.repelloai import (
|
||||
RepelloAIAnalyzeResponse,
|
||||
)
|
||||
from litellm.types.utils import (
|
||||
CallTypesLiteral,
|
||||
GuardrailStatus,
|
||||
LLMResponseTypes,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
)
|
||||
|
||||
DEFAULT_REPELLOAI_API_BASE = "https://argusapi.repello.ai/sdk/v1"
|
||||
DEFAULT_REPELLOAI_TIMEOUT = 30.0
|
||||
BLOCKED_VERDICT = "blocked"
|
||||
FLAGGED_VERDICT = "flagged"
|
||||
PASSED_VERDICT = "passed"
|
||||
|
||||
# Argus returns these for a permanently broken guardrail (bad key, unknown
|
||||
# asset_id, malformed payload), not a transient outage. They must always
|
||||
# block, never honour fail_open.
|
||||
CONFIG_ERROR_STATUS_CODES = frozenset({400, 401, 403, 404, 422})
|
||||
_SCHEMA_SCALAR_KEYS = frozenset(("name", "description", "title", "const", "default"))
|
||||
_SCHEMA_LIST_KEYS = frozenset(("enum", "examples"))
|
||||
_SCHEMA_EXTRACTED_KEYS = _SCHEMA_SCALAR_KEYS | _SCHEMA_LIST_KEYS
|
||||
|
||||
|
||||
class RepelloAIGuardrailMissingSecrets(Exception):
|
||||
pass
|
||||
|
||||
|
||||
def _is_object_dict(value: object) -> TypeGuard[dict[str, object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip
|
||||
return isinstance(value, dict)
|
||||
|
||||
|
||||
def _is_object_list(value: object) -> TypeGuard[list[object]]: # guard-ok: isinstance narrows correctly; predicate is trivially correct # fmt: skip
|
||||
return isinstance(value, list)
|
||||
|
||||
|
||||
class RepelloAIGuardrail(CustomGuardrail):
|
||||
@staticmethod
|
||||
def _get_field(obj: object, key: str) -> object:
|
||||
if _is_object_dict(obj):
|
||||
return obj.get(key)
|
||||
return getattr(obj, key, None)
|
||||
|
||||
@classmethod
|
||||
def _extract_tool_call_args_from_message(cls, message: object) -> list[str]:
|
||||
args: list[str] = []
|
||||
|
||||
tool_calls = cls._get_field(message, "tool_calls")
|
||||
if _is_object_list(tool_calls):
|
||||
for tool_call in tool_calls:
|
||||
function = cls._get_field(tool_call, "function")
|
||||
arguments = cls._get_field(function, "arguments")
|
||||
if isinstance(arguments, str) and arguments.strip():
|
||||
args.append(arguments)
|
||||
|
||||
function_call = cls._get_field(message, "function_call")
|
||||
arguments = cls._get_field(function_call, "arguments")
|
||||
if isinstance(arguments, str) and arguments.strip():
|
||||
args.append(arguments)
|
||||
|
||||
return args
|
||||
|
||||
@staticmethod
|
||||
def _iter_schema_text(node: object) -> list[str]:
|
||||
texts: list[str] = []
|
||||
stack: list[object] = [node]
|
||||
|
||||
while stack:
|
||||
current = stack.pop()
|
||||
if _is_object_dict(current):
|
||||
for key in _SCHEMA_SCALAR_KEYS:
|
||||
value = current.get(key)
|
||||
if isinstance(value, str) and value:
|
||||
texts.append(value)
|
||||
for key in _SCHEMA_LIST_KEYS:
|
||||
items = current.get(key)
|
||||
if _is_object_list(items):
|
||||
for item in items:
|
||||
if isinstance(item, str) and item:
|
||||
texts.append(item)
|
||||
remaining: list[object] = [
|
||||
v for k, v in current.items() if k not in _SCHEMA_EXTRACTED_KEYS
|
||||
]
|
||||
stack.extend(reversed(remaining))
|
||||
elif _is_object_list(current):
|
||||
stack.extend(reversed(current))
|
||||
|
||||
return texts
|
||||
|
||||
@classmethod
|
||||
def _extract_tool_definition_text(cls, data: dict[str, object]) -> list[str]:
|
||||
texts: list[str] = []
|
||||
|
||||
tools = data.get("tools")
|
||||
for tool in tools if _is_object_list(tools) else []:
|
||||
if not _is_object_dict(tool):
|
||||
continue
|
||||
function = tool.get("function")
|
||||
if _is_object_dict(function):
|
||||
texts.extend(cls._iter_schema_text(function))
|
||||
|
||||
functions = data.get("functions")
|
||||
for function in functions if _is_object_list(functions) else []:
|
||||
if _is_object_dict(function):
|
||||
texts.extend(cls._iter_schema_text(function))
|
||||
|
||||
return texts
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
asset_id: str | None = None,
|
||||
unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed",
|
||||
guardrail_name: str | None = None,
|
||||
event_hook: (
|
||||
GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None
|
||||
) = None,
|
||||
default_on: bool = False,
|
||||
):
|
||||
self.repelloai_api_key = (
|
||||
api_key
|
||||
or get_secret_str("ARGUS_API_KEY")
|
||||
or get_secret_str("REPELLOAI_API_KEY")
|
||||
or ""
|
||||
)
|
||||
if not self.repelloai_api_key:
|
||||
raise RepelloAIGuardrailMissingSecrets(
|
||||
"Couldn't get Repello API key. Set `ARGUS_API_KEY` in the environment "
|
||||
"or pass `api_key` to the guardrail in the config file."
|
||||
)
|
||||
|
||||
self.asset_id = asset_id
|
||||
if not self.asset_id:
|
||||
raise ValueError(
|
||||
"Repello guardrail requires an `asset_id`. Create an asset in the Repello "
|
||||
"dashboard and set `asset_id` on the guardrail in the config file."
|
||||
)
|
||||
|
||||
self.api_base = (
|
||||
api_base
|
||||
or get_secret_str("REPELLOAI_API_BASE")
|
||||
or DEFAULT_REPELLOAI_API_BASE
|
||||
)
|
||||
self.unreachable_fallback: Literal["fail_closed", "fail_open"] = (
|
||||
"fail_open" if unreachable_fallback == "fail_open" else "fail_closed"
|
||||
)
|
||||
self.async_handler = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.GuardrailCallback,
|
||||
params={"timeout": DEFAULT_REPELLOAI_TIMEOUT},
|
||||
)
|
||||
super().__init__( # pyright: ignore[reportUnknownMemberType]
|
||||
guardrail_name=guardrail_name,
|
||||
event_hook=event_hook,
|
||||
default_on=default_on,
|
||||
)
|
||||
|
||||
async def _call_analyze(
|
||||
self,
|
||||
text: str,
|
||||
stage: Literal["prompt", "response"],
|
||||
request_data: dict[str, object],
|
||||
event_type: GuardrailEventHooks,
|
||||
) -> RepelloAIAnalyzeResponse | None:
|
||||
endpoint = f"{self.api_base}/analyze/{stage}"
|
||||
request: dict[str, object] = {
|
||||
"asset_id": self.asset_id or "",
|
||||
"scan_data": {stage: text},
|
||||
}
|
||||
|
||||
status: GuardrailStatus = "success"
|
||||
guardrail_json_response: str | dict[str, object] | list[dict[str, object]] = ""
|
||||
start_time: datetime = datetime.now()
|
||||
repelloai_response: RepelloAIAnalyzeResponse | None = None
|
||||
try:
|
||||
verbose_proxy_logger.debug("RepelloAI Argus request: %s", request)
|
||||
raw_response: HttpxResponse | None = (
|
||||
await self.async_handler.post( # pyright: ignore[reportUnknownMemberType]
|
||||
url=endpoint,
|
||||
headers={"X-API-Key": self.repelloai_api_key},
|
||||
json=request,
|
||||
)
|
||||
)
|
||||
if raw_response is None:
|
||||
raise ValueError("RepelloAI Argus returned no response")
|
||||
response: HttpxResponse = raw_response
|
||||
self._raise_for_config_error(response)
|
||||
response.raise_for_status()
|
||||
try:
|
||||
repelloai_response = TypeAdapter(
|
||||
RepelloAIAnalyzeResponse
|
||||
).validate_json(response.text)
|
||||
except ValidationError as e:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={
|
||||
"error": "RepelloAI Argus guardrail returned invalid JSON",
|
||||
"status_code": response.status_code,
|
||||
},
|
||||
) from e
|
||||
verbose_proxy_logger.debug(
|
||||
"RepelloAI Argus response: %s", repelloai_response
|
||||
)
|
||||
if self._verdict_blocks(repelloai_response):
|
||||
status = "guardrail_intervened"
|
||||
return repelloai_response
|
||||
except HTTPException as e:
|
||||
status = "guardrail_failed_to_respond"
|
||||
guardrail_json_response = str(e.detail) if not isinstance(e.detail, (dict, list)) else e.detail # type: ignore[assignment]
|
||||
raise
|
||||
except HTTPError as e:
|
||||
status = "guardrail_failed_to_respond"
|
||||
guardrail_json_response = str(e)
|
||||
return self._handle_unreachable(e)
|
||||
except Exception as e:
|
||||
status = "guardrail_failed_to_respond"
|
||||
guardrail_json_response = str(e)
|
||||
raise HTTPException(
|
||||
status_code=500, detail={"error": "RepelloAI Argus guardrail failed"}
|
||||
) from e
|
||||
finally:
|
||||
end_time = datetime.now()
|
||||
if repelloai_response is not None:
|
||||
guardrail_json_response = dict(repelloai_response)
|
||||
self.add_standard_logging_guardrail_information_to_request_data( # pyright: ignore[reportUnknownMemberType]
|
||||
guardrail_json_response=guardrail_json_response,
|
||||
guardrail_status=status,
|
||||
request_data=request_data,
|
||||
start_time=start_time.timestamp(),
|
||||
end_time=end_time.timestamp(),
|
||||
duration=(end_time - start_time).total_seconds(),
|
||||
masked_entity_count={},
|
||||
event_type=event_type,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _raise_for_config_error(response: HttpxResponse) -> None:
|
||||
if response.status_code in CONFIG_ERROR_STATUS_CODES:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={
|
||||
"error": "RepelloAI Argus guardrail is misconfigured",
|
||||
"status_code": response.status_code,
|
||||
},
|
||||
)
|
||||
|
||||
def _verdict_blocks(
|
||||
self, repelloai_response: RepelloAIAnalyzeResponse | None
|
||||
) -> bool:
|
||||
if repelloai_response is None:
|
||||
return False
|
||||
verdict = repelloai_response.get("verdict")
|
||||
if verdict == BLOCKED_VERDICT:
|
||||
return True
|
||||
if verdict in (PASSED_VERDICT, FLAGGED_VERDICT):
|
||||
return False
|
||||
verbose_proxy_logger.warning(
|
||||
"RepelloAI Argus returned an unrecognized verdict (%s) - blocking.",
|
||||
verdict,
|
||||
)
|
||||
return True
|
||||
|
||||
def _handle_unreachable(self, error: Exception) -> RepelloAIAnalyzeResponse | None:
|
||||
verbose_proxy_logger.warning("RepelloAI Argus unreachable: %s", str(error))
|
||||
if self.unreachable_fallback == "fail_closed":
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": "RepelloAI Argus guardrail unreachable"},
|
||||
)
|
||||
return None
|
||||
|
||||
def _raise_if_blocked(
|
||||
self, repelloai_response: RepelloAIAnalyzeResponse | None
|
||||
) -> None:
|
||||
if repelloai_response is None:
|
||||
return
|
||||
if self._verdict_blocks(repelloai_response):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=self._format_blocked_detail(repelloai_response),
|
||||
)
|
||||
self._log_flagged_verdict(repelloai_response)
|
||||
|
||||
@classmethod
|
||||
def _format_blocked_detail(
|
||||
cls, repelloai_response: RepelloAIAnalyzeResponse
|
||||
) -> str:
|
||||
policies = repelloai_response.get("policies_violated")
|
||||
if not isinstance(policies, list) or not policies:
|
||||
return "Blocked by RepelloAI Argus guardrail."
|
||||
|
||||
formatted_policies: list[str] = []
|
||||
for policy in policies:
|
||||
policy_name = policy.get("policy_name") or "unknown_policy"
|
||||
details: list[str] = []
|
||||
action_taken = policy.get("action_taken")
|
||||
if action_taken:
|
||||
details.append(f"action: {action_taken}")
|
||||
policy_details = policy.get("details")
|
||||
if isinstance(policy_details, dict):
|
||||
score = policy_details.get("score")
|
||||
if score is not None:
|
||||
details.append(f"score: {score}")
|
||||
suffix = f" ({', '.join(details)})" if details else ""
|
||||
formatted_policies.append(f"{policy_name}{suffix}")
|
||||
|
||||
if not formatted_policies:
|
||||
return "Blocked by RepelloAI Argus guardrail."
|
||||
return f"Blocked by RepelloAI Argus guardrail. Policies violated: {'; '.join(formatted_policies)}."
|
||||
|
||||
@staticmethod
|
||||
def _log_flagged_verdict(repelloai_response: RepelloAIAnalyzeResponse) -> None:
|
||||
if repelloai_response.get("verdict") == FLAGGED_VERDICT:
|
||||
verbose_proxy_logger.warning(
|
||||
"RepelloAI Argus flagged content (allowed): %s",
|
||||
repelloai_response.get("policies_violated"),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _extract_prompt_message_text(data: dict[str, object]) -> list[str]:
|
||||
messages = build_inspection_messages(data)
|
||||
return [
|
||||
content
|
||||
for message in messages
|
||||
if isinstance(content := message.get("content"), str) and content
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _extract_input_text_parts(content: object) -> list[str]:
|
||||
if not _is_object_list(content):
|
||||
return []
|
||||
return [
|
||||
text
|
||||
for part in content
|
||||
if _is_object_dict(part) and part.get("type") == "input_text"
|
||||
if isinstance(text := part.get("text"), str) and text
|
||||
]
|
||||
|
||||
@staticmethod
|
||||
def _extract_prompt_field_text(data: dict[str, object]) -> list[str]:
|
||||
prompt = data.get("prompt")
|
||||
if isinstance(prompt, str) and prompt:
|
||||
return [prompt]
|
||||
if _is_object_list(prompt):
|
||||
return [item for item in prompt if isinstance(item, str) and item]
|
||||
return []
|
||||
|
||||
@classmethod
|
||||
def _extract_prompt_text(cls, data: dict[str, object]) -> str | None:
|
||||
texts = cls._extract_prompt_message_text(data)
|
||||
texts.extend(cls._extract_prompt_field_text(data))
|
||||
|
||||
instructions = data.get("instructions")
|
||||
if isinstance(instructions, str) and instructions:
|
||||
texts.append(instructions)
|
||||
|
||||
raw_messages = data.get("messages")
|
||||
if _is_object_list(raw_messages):
|
||||
for message in raw_messages:
|
||||
texts.extend(cls._extract_tool_call_args_from_message(message))
|
||||
|
||||
raw_input = data.get("input")
|
||||
if _is_object_list(raw_input):
|
||||
for item in raw_input:
|
||||
if _is_object_dict(item):
|
||||
if "role" not in item:
|
||||
continue
|
||||
texts.extend(cls._extract_tool_call_args_from_message(item))
|
||||
texts.extend(cls._extract_input_text_parts(item.get("content")))
|
||||
|
||||
texts.extend(cls._extract_tool_definition_text(data))
|
||||
return "\n".join(text for text in texts if text) if texts else None
|
||||
|
||||
async def async_pre_call_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
cache: litellm.DualCache,
|
||||
data: dict[str, object],
|
||||
call_type: CallTypesLiteral,
|
||||
) -> Exception | str | dict[str, object] | None:
|
||||
verbose_proxy_logger.debug("RepelloAI Argus: pre_call_hook")
|
||||
|
||||
event_type = GuardrailEventHooks.pre_call
|
||||
if (
|
||||
self.should_run_guardrail( # pyright: ignore[reportUnknownMemberType]
|
||||
data=data, event_type=event_type
|
||||
)
|
||||
is not True
|
||||
):
|
||||
return data
|
||||
|
||||
text = self._extract_prompt_text(data)
|
||||
if not text:
|
||||
verbose_proxy_logger.warning(
|
||||
"RepelloAI Argus: no inspectable prompt text in data - skipping."
|
||||
)
|
||||
return data
|
||||
|
||||
repelloai_response = await self._call_analyze(
|
||||
text=text,
|
||||
stage="prompt",
|
||||
request_data=data,
|
||||
event_type=event_type,
|
||||
)
|
||||
self._raise_if_blocked(repelloai_response)
|
||||
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=data, guardrail_name=self.guardrail_name
|
||||
)
|
||||
return data
|
||||
|
||||
async def async_post_call_success_hook(
|
||||
self,
|
||||
data: dict[str, object],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: LLMResponseTypes,
|
||||
):
|
||||
verbose_proxy_logger.debug("RepelloAI Argus: post_call_success_hook")
|
||||
|
||||
event_type = GuardrailEventHooks.post_call
|
||||
if (
|
||||
self.should_run_guardrail( # pyright: ignore[reportUnknownMemberType]
|
||||
data=data, event_type=event_type
|
||||
)
|
||||
is not True
|
||||
):
|
||||
return response
|
||||
|
||||
text = self._extract_response_text(response)
|
||||
if not text:
|
||||
verbose_proxy_logger.warning(
|
||||
"RepelloAI Argus: no inspectable response text - skipping."
|
||||
)
|
||||
return response
|
||||
|
||||
repelloai_response = await self._call_analyze(
|
||||
text=text,
|
||||
stage="response",
|
||||
request_data=data,
|
||||
event_type=event_type,
|
||||
)
|
||||
self._raise_if_blocked(repelloai_response)
|
||||
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=data, guardrail_name=self.guardrail_name
|
||||
)
|
||||
return response
|
||||
|
||||
async def async_post_call_streaming_iterator_hook(
|
||||
self,
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
response: AsyncGenerator[ModelResponseStream, None],
|
||||
request_data: dict[str, object],
|
||||
) -> AsyncGenerator[ModelResponseStream, None]:
|
||||
from litellm import main as litellm_main
|
||||
|
||||
event_type = GuardrailEventHooks.post_call
|
||||
if (
|
||||
self.should_run_guardrail( # pyright: ignore[reportUnknownMemberType]
|
||||
data=request_data, event_type=event_type
|
||||
)
|
||||
is not True
|
||||
):
|
||||
async for chunk in response:
|
||||
yield chunk
|
||||
return
|
||||
|
||||
chunks: list[ModelResponseStream] = []
|
||||
async for chunk in response:
|
||||
chunks.append(chunk)
|
||||
|
||||
assembled = litellm_main.stream_chunk_builder( # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType]
|
||||
chunks=chunks
|
||||
)
|
||||
text = (
|
||||
self._extract_response_text(assembled)
|
||||
if isinstance(assembled, ModelResponse)
|
||||
else None
|
||||
)
|
||||
if text:
|
||||
repelloai_response = await self._call_analyze(
|
||||
text=text,
|
||||
stage="response",
|
||||
request_data=request_data,
|
||||
event_type=event_type,
|
||||
)
|
||||
if repelloai_response is not None:
|
||||
self._log_flagged_verdict(repelloai_response)
|
||||
if self._verdict_blocks(repelloai_response):
|
||||
from litellm.proxy.proxy_server import StreamingCallbackError
|
||||
|
||||
raise StreamingCallbackError("Blocked by RepelloAI Argus guardrail")
|
||||
add_guardrail_to_applied_guardrails_header(
|
||||
request_data=request_data, guardrail_name=self.guardrail_name
|
||||
)
|
||||
else:
|
||||
verbose_proxy_logger.warning(
|
||||
"RepelloAI Argus: no inspectable text in streamed response; skipping scan. "
|
||||
"guardrail=%s assembled_type=%s",
|
||||
self.guardrail_name,
|
||||
type(assembled).__name__,
|
||||
)
|
||||
|
||||
for chunk in chunks:
|
||||
yield chunk
|
||||
|
||||
@staticmethod
|
||||
def _extract_response_text(response: object) -> str | None:
|
||||
if _is_object_dict(response):
|
||||
response_dict = response
|
||||
elif isinstance(response, ModelResponse):
|
||||
response_dict = (
|
||||
response.model_dump() # pyright: ignore[reportUnknownMemberType]
|
||||
)
|
||||
else:
|
||||
output_text = getattr(response, "output_text", None)
|
||||
if isinstance(output_text, str) and output_text:
|
||||
return output_text
|
||||
response_dict = {}
|
||||
|
||||
text = RepelloAIGuardrail._extract_chat_completion_text(response_dict)
|
||||
if text:
|
||||
return text
|
||||
return RepelloAIGuardrail._extract_responses_api_text(response_dict)
|
||||
|
||||
@classmethod
|
||||
def _extract_chat_completion_text(
|
||||
cls, response_dict: dict[str, object]
|
||||
) -> str | None:
|
||||
choices = response_dict.get("choices")
|
||||
if not _is_object_list(choices):
|
||||
return None
|
||||
parts: list[str] = []
|
||||
for choice in choices:
|
||||
if not _is_object_dict(choice):
|
||||
continue
|
||||
message = choice.get("message")
|
||||
if _is_object_dict(message):
|
||||
content = message.get("content")
|
||||
if isinstance(content, str) and content:
|
||||
parts.append(content)
|
||||
parts.extend(cls._extract_tool_call_args_from_message(message))
|
||||
text = choice.get("text")
|
||||
if isinstance(text, str) and text:
|
||||
parts.append(text)
|
||||
return "\n".join(parts) if parts else None
|
||||
|
||||
@staticmethod
|
||||
def _extract_responses_api_text(response_dict: dict[str, object]) -> str | None:
|
||||
output = response_dict.get("output")
|
||||
if not _is_object_list(output):
|
||||
return None
|
||||
texts: list[str] = []
|
||||
for output_item in output:
|
||||
if not _is_object_dict(output_item):
|
||||
continue
|
||||
item_type = output_item.get("type")
|
||||
if item_type == "function_call":
|
||||
arguments = output_item.get("arguments")
|
||||
if isinstance(arguments, str) and arguments:
|
||||
texts.append(arguments)
|
||||
continue
|
||||
if item_type != "message":
|
||||
continue
|
||||
content = output_item.get("content")
|
||||
if not _is_object_list(content):
|
||||
continue
|
||||
for content_item in content:
|
||||
if not _is_object_dict(content_item):
|
||||
continue
|
||||
if content_item.get("type") not in ("output_text", "text"):
|
||||
continue
|
||||
text = content_item.get("text")
|
||||
if isinstance(text, str) and text:
|
||||
texts.append(text)
|
||||
return "".join(texts) if texts else None
|
||||
|
||||
@staticmethod
|
||||
def get_config_model() -> type[GuardrailConfigModel[BaseModel]] | None:
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.repelloai import (
|
||||
RepelloAIGuardrailConfigModel,
|
||||
)
|
||||
|
||||
return RepelloAIGuardrailConfigModel
|
||||
|
|
@ -401,7 +401,7 @@ def _resolve_health_check_max_tokens(
|
|||
3. For non-wildcard reasoning routes: BACKGROUND_HEALTH_CHECK_MAX_TOKENS_REASONING
|
||||
from env (if set)
|
||||
4. BACKGROUND_HEALTH_CHECK_MAX_TOKENS (global, any route including wildcards)
|
||||
5. Non-wildcard default: 5
|
||||
5. Non-wildcard default: 16
|
||||
6. Wildcard and nothing from (1)(4): leave unset (caller omits max_tokens)
|
||||
"""
|
||||
explicit = model_info.get("health_check_max_tokens", None)
|
||||
|
|
@ -432,7 +432,7 @@ def _resolve_health_check_max_tokens(
|
|||
return int(BACKGROUND_HEALTH_CHECK_MAX_TOKENS)
|
||||
|
||||
if not is_wildcard:
|
||||
return 5
|
||||
return 16
|
||||
|
||||
return None
|
||||
|
||||
|
|
|
|||
|
|
@ -162,9 +162,20 @@ class _ProxyDBLogger(CustomLogger):
|
|||
if obj_start is not None:
|
||||
actual_start_time = obj_start
|
||||
|
||||
# A stream that broke mid-flight still billed the provider for the
|
||||
# chunks already delivered. ``post_call_failure_hook`` lifts that
|
||||
# recovered cost onto request_data (the usage rides along in
|
||||
# ``combined_usage_object`` for the token columns), so attribute the
|
||||
# real partial spend to this failure row instead of zero.
|
||||
recovered_response_cost = 0.0
|
||||
if isinstance(request_data.get("combined_usage_object"), litellm.Usage):
|
||||
recovered_response_cost = max(
|
||||
float(request_data.get("response_cost") or 0.0), 0.0
|
||||
)
|
||||
|
||||
await proxy_logging_obj.db_spend_update_writer.update_database(
|
||||
token=user_api_key_dict.api_key,
|
||||
response_cost=0.0,
|
||||
response_cost=recovered_response_cost,
|
||||
user_id=user_api_key_dict.user_id,
|
||||
end_user_id=user_api_key_dict.end_user_id,
|
||||
team_id=user_api_key_dict.team_id,
|
||||
|
|
|
|||
|
|
@ -35,16 +35,25 @@ _PRISMA_TO_PG_TABLE: Dict[str, str] = {
|
|||
|
||||
|
||||
def update_metrics(existing_metrics: SpendMetrics, record: Any) -> SpendMetrics:
|
||||
"""Update metrics with new record data."""
|
||||
existing_metrics.spend += record.spend
|
||||
existing_metrics.prompt_tokens += record.prompt_tokens
|
||||
existing_metrics.completion_tokens += record.completion_tokens
|
||||
existing_metrics.total_tokens += record.prompt_tokens + record.completion_tokens
|
||||
existing_metrics.cache_read_input_tokens += record.cache_read_input_tokens
|
||||
existing_metrics.cache_creation_input_tokens += record.cache_creation_input_tokens
|
||||
existing_metrics.api_requests += record.api_requests
|
||||
existing_metrics.successful_requests += record.successful_requests
|
||||
existing_metrics.failed_requests += record.failed_requests
|
||||
"""Update metrics with new record data.
|
||||
|
||||
Rollup rows can carry None for numeric fields when SUM() spans zero rows
|
||||
(e.g. a key with no spend), so coalesce to 0 before accumulating to avoid
|
||||
a TypeError. Mirrors the handling in ``_record_to_spend_metrics``.
|
||||
"""
|
||||
prompt_tokens = record.prompt_tokens or 0
|
||||
completion_tokens = record.completion_tokens or 0
|
||||
existing_metrics.spend += record.spend or 0.0
|
||||
existing_metrics.prompt_tokens += prompt_tokens
|
||||
existing_metrics.completion_tokens += completion_tokens
|
||||
existing_metrics.total_tokens += prompt_tokens + completion_tokens
|
||||
existing_metrics.cache_read_input_tokens += record.cache_read_input_tokens or 0
|
||||
existing_metrics.cache_creation_input_tokens += (
|
||||
record.cache_creation_input_tokens or 0
|
||||
)
|
||||
existing_metrics.api_requests += record.api_requests or 0
|
||||
existing_metrics.successful_requests += record.successful_requests or 0
|
||||
existing_metrics.failed_requests += record.failed_requests or 0
|
||||
return existing_metrics
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -397,7 +397,7 @@ def _set_object_metadata_field(
|
|||
field_name: Name of the metadata field to set
|
||||
value: Value to set for the field
|
||||
"""
|
||||
if field_name in LiteLLM_ManagementEndpoint_MetadataFields_Premium:
|
||||
if field_name in LiteLLM_ManagementEndpoint_MetadataFields_Premium and value:
|
||||
_premium_user_check(field_name)
|
||||
|
||||
object_data.metadata = object_data.metadata or {}
|
||||
|
|
@ -563,13 +563,11 @@ def _update_metadata_field(updated_kv: dict, field_name: str) -> None:
|
|||
field_name: Name of the metadata field being updated
|
||||
"""
|
||||
if field_name in LiteLLM_ManagementEndpoint_MetadataFields_Premium:
|
||||
value = updated_kv.get(field_name)
|
||||
# Skip the premium check for empty collections ([] or {}).
|
||||
# The UI sends these as defaults even when the user hasn't configured
|
||||
# any enterprise features (see issue #20304). However, we still
|
||||
# proceed with the update so that users can intentionally clear a
|
||||
# previously-set field by sending an empty list/dict.
|
||||
if value is not None and value != [] and value != {}:
|
||||
# The UI sends falsy defaults (False, [], {}) even when the user has not
|
||||
# enabled any enterprise feature (see #20304, #30285); require a license
|
||||
# only for a truthy value. The falsy value is still persisted below so a
|
||||
# previously-set field can be cleared.
|
||||
if updated_kv.get(field_name):
|
||||
_premium_user_check()
|
||||
|
||||
if field_name in updated_kv and updated_kv[field_name] is not None:
|
||||
|
|
|
|||
|
|
@ -1793,7 +1793,8 @@ def prepare_metadata_fields(
|
|||
if k in LiteLLM_ManagementEndpoint_MetadataFields_Premium:
|
||||
from litellm.proxy.utils import _premium_user_check
|
||||
|
||||
_premium_user_check(k)
|
||||
if v:
|
||||
_premium_user_check(k)
|
||||
casted_metadata[k] = v
|
||||
|
||||
except Exception as e:
|
||||
|
|
@ -5121,7 +5122,7 @@ async def list_keys(
|
|||
size: int = Query(10, description="Page size", ge=1, le=100),
|
||||
user_id: Optional[str] = Query(
|
||||
None,
|
||||
description="Filter keys by user ID. Supports partial matching (substring, case-insensitive).",
|
||||
description="Filter keys by user ID. Exact match by default; set substring_matching=true (admin only) for case-insensitive substring matching.",
|
||||
),
|
||||
team_id: Optional[str] = Query(None, description="Filter keys by team ID"),
|
||||
organization_id: Optional[str] = Query(
|
||||
|
|
@ -5130,7 +5131,7 @@ async def list_keys(
|
|||
key_hash: Optional[str] = Query(None, description="Filter keys by key hash"),
|
||||
key_alias: Optional[str] = Query(
|
||||
None,
|
||||
description="Filter keys by key alias. Supports partial matching (substring, case-insensitive).",
|
||||
description="Filter keys by key alias. Exact match by default; set substring_matching=true (admin only) for case-insensitive substring matching.",
|
||||
),
|
||||
return_full_object: bool = Query(False, description="Return full key object"),
|
||||
include_team_keys: bool = Query(
|
||||
|
|
@ -5154,6 +5155,10 @@ async def list_keys(
|
|||
access_group_id: Optional[str] = Query(
|
||||
None, description="Filter keys by access group ID"
|
||||
),
|
||||
substring_matching: bool = Query(
|
||||
False,
|
||||
description="If true (proxy admins only), match user_id/key_alias as case-insensitive substrings instead of exact values. Defaults to false: /key/list matched these exactly before substring search was added, and an exact user_id/key_alias filter must never return another user's keys.",
|
||||
),
|
||||
) -> KeyListResponseObject:
|
||||
"""
|
||||
List all keys for a given user / team / organization.
|
||||
|
|
@ -5235,12 +5240,21 @@ async def list_keys(
|
|||
else:
|
||||
admin_team_ids = None
|
||||
|
||||
use_substring_matching = user_api_key_dict.user_role in [
|
||||
is_proxy_admin = user_api_key_dict.user_role in [
|
||||
LitellmUserRoles.PROXY_ADMIN.value,
|
||||
LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value,
|
||||
]
|
||||
|
||||
if not user_id and not use_substring_matching:
|
||||
# Substring matching is opt-in (admin-only). /key/list matched user_id and
|
||||
# key_alias exactly before substring search was added; auto-applying a
|
||||
# substring match to every admin call broke that contract and let a caller
|
||||
# passing an exact user_id (e.g. an integration scoping to one user with an
|
||||
# admin key) receive other users' keys (user_id="alice" -> "alice2"). Exact
|
||||
# by default restores the prior behavior; the dashboard opts in explicitly.
|
||||
use_substring_matching = substring_matching and is_proxy_admin
|
||||
|
||||
# Admins may omit user_id to list all keys; non-admins are scoped to self.
|
||||
if not user_id and not is_proxy_admin:
|
||||
user_id = user_api_key_dict.user_id
|
||||
|
||||
response = await _list_key_helper(
|
||||
|
|
|
|||
|
|
@ -49,6 +49,8 @@ from litellm.constants import LITELLM_PROXY_ADMIN_NAME
|
|||
from litellm.proxy._experimental.mcp_server.utils import (
|
||||
build_env_var_setup_url,
|
||||
collect_env_var_references,
|
||||
LITELLM_MCP_SERVER_DESCRIPTION,
|
||||
LITELLM_MCP_SERVER_NAME,
|
||||
get_server_prefix,
|
||||
parse_admin_env_vars,
|
||||
)
|
||||
|
|
@ -89,8 +91,6 @@ def does_mcp_server_exist(
|
|||
|
||||
|
||||
DEFAULT_MCP_REGISTRY_VERSION = "1.0.0"
|
||||
LITELLM_MCP_SERVER_NAME = "litellm-mcp-server"
|
||||
LITELLM_MCP_SERVER_DESCRIPTION = "MCP Server for LiteLLM"
|
||||
|
||||
try:
|
||||
importlib.import_module("mcp")
|
||||
|
|
|
|||
|
|
@ -2961,6 +2961,7 @@ async def team_member_update(
|
|||
returned_team_info: TeamInfoResponseObject = await team_info(
|
||||
http_request=http_request,
|
||||
team_id=data.team_id,
|
||||
key_limit=None,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
|
|
@ -3577,6 +3578,9 @@ async def team_info(
|
|||
team_id: str = fastapi.Query(
|
||||
default=None, description="Team ID in the request parameters"
|
||||
),
|
||||
key_limit: int | None = fastapi.Query(
|
||||
default=None, description="Limit the number of keys returned", gt=0
|
||||
),
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
|
|
@ -3632,6 +3636,7 @@ async def team_info(
|
|||
table_name="key",
|
||||
query_type="find_all",
|
||||
expires=datetime.now(),
|
||||
limit=key_limit,
|
||||
)
|
||||
|
||||
if keys is None:
|
||||
|
|
|
|||
|
|
@ -90,7 +90,7 @@ from litellm.proxy.common_utils.admin_ui_utils import (
|
|||
from litellm.proxy.common_utils.html_forms.jwt_display_template import (
|
||||
jwt_display_template,
|
||||
)
|
||||
from litellm.proxy.common_utils.html_forms.ui_login import html_form
|
||||
from litellm.proxy.common_utils.html_forms.ui_login import build_ui_login_form
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.management_endpoints.internal_user_endpoints import new_user
|
||||
from litellm.proxy.management_endpoints.sso import CustomMicrosoftSSO
|
||||
|
|
@ -902,6 +902,7 @@ async def google_login(
|
|||
Example:
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
general_settings,
|
||||
premium_user,
|
||||
prisma_client,
|
||||
user_api_key_cache,
|
||||
|
|
@ -948,7 +949,6 @@ async def google_login(
|
|||
missing_env_vars = show_missing_vars_in_env()
|
||||
if missing_env_vars is not None:
|
||||
return missing_env_vars
|
||||
ui_username = os.getenv("UI_USERNAME")
|
||||
|
||||
# get url from request - always use regular callback, but set state for CLI
|
||||
redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso(
|
||||
|
|
@ -1009,16 +1009,20 @@ async def google_login(
|
|||
samesite="lax",
|
||||
)
|
||||
return sso_redirect
|
||||
elif ui_username is not None:
|
||||
# No Google, Microsoft SSO
|
||||
# Use UI Credentials set in .env
|
||||
from fastapi.responses import HTMLResponse
|
||||
|
||||
return HTMLResponse(content=html_form, status_code=200)
|
||||
else:
|
||||
from fastapi.responses import HTMLResponse
|
||||
from fastapi.responses import HTMLResponse
|
||||
|
||||
return HTMLResponse(content=html_form, status_code=200)
|
||||
hide_default_credentials_hint = (
|
||||
os.getenv("LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT", "false").lower() == "true"
|
||||
or general_settings.get("hide_default_credentials_hint", False) is True
|
||||
)
|
||||
return HTMLResponse(
|
||||
content=build_ui_login_form(
|
||||
show_deprecation_banner=True,
|
||||
hide_default_credentials_hint=hide_default_credentials_hint,
|
||||
),
|
||||
status_code=200,
|
||||
)
|
||||
|
||||
|
||||
def generic_response_convertor(
|
||||
|
|
|
|||
53
litellm/proxy/middleware/security_headers_middleware.py
Normal file
53
litellm/proxy/middleware/security_headers_middleware.py
Normal file
|
|
@ -0,0 +1,53 @@
|
|||
"""
|
||||
Adds anti-framing / content-type security headers to every HTTP response.
|
||||
|
||||
X-Frame-Options and Content-Security-Policy: frame-ancestors 'none' stop the
|
||||
admin UI and login pages from being embedded cross-origin (clickjacking).
|
||||
X-Content-Type-Options: nosniff stops MIME sniffing.
|
||||
|
||||
Strict-Transport-Security is opt-in via LITELLM_ENABLE_HSTS because it only
|
||||
makes sense over HTTPS and would lock browsers out of plain-http deployments.
|
||||
|
||||
Headers are set with setdefault so a route that intentionally sets its own
|
||||
value is never overridden.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
from starlette.datastructures import MutableHeaders
|
||||
from starlette.types import ASGIApp, Message, Receive, Scope, Send
|
||||
|
||||
STATIC_SECURITY_HEADERS = (
|
||||
("X-Frame-Options", "DENY"),
|
||||
("Content-Security-Policy", "frame-ancestors 'none'"),
|
||||
("X-Content-Type-Options", "nosniff"),
|
||||
)
|
||||
HSTS_HEADER = ("Strict-Transport-Security", "max-age=31536000; includeSubDomains")
|
||||
|
||||
|
||||
def _hsts_enabled() -> bool:
|
||||
return os.getenv("LITELLM_ENABLE_HSTS", "false").strip().lower() == "true"
|
||||
|
||||
|
||||
class SecurityHeadersMiddleware:
|
||||
def __init__(self, app: ASGIApp) -> None:
|
||||
self.app = app
|
||||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
if scope["type"] != "http":
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
|
||||
async def send_with_security_headers(message: Message) -> None:
|
||||
if message["type"] == "http.response.start":
|
||||
headers = MutableHeaders(scope=message)
|
||||
applied = (
|
||||
(*STATIC_SECURITY_HEADERS, HSTS_HEADER)
|
||||
if _hsts_enabled()
|
||||
else STATIC_SECURITY_HEADERS
|
||||
)
|
||||
for name, value in applied:
|
||||
headers.setdefault(name, value)
|
||||
await send(message)
|
||||
|
||||
await self.app(scope, receive, send_with_security_headers)
|
||||
|
|
@ -38,6 +38,9 @@ from litellm.proxy.common_utils.openai_endpoint_utils import (
|
|||
get_custom_llm_provider_from_request_headers,
|
||||
get_custom_llm_provider_from_request_query,
|
||||
)
|
||||
from litellm.litellm_core_utils.cloud_storage_security import (
|
||||
is_managed_cloud_storage_uri,
|
||||
)
|
||||
from litellm.proxy.openai_files_endpoints.common_utils import (
|
||||
_is_base64_encoded_unified_file_id,
|
||||
encode_file_id_with_model,
|
||||
|
|
@ -726,6 +729,15 @@ async def get_file_content(
|
|||
}
|
||||
)
|
||||
else:
|
||||
# A raw cloud-storage URI (s3://, gs://) supplied here would skip the
|
||||
# managed-file owner/team check that only runs for unified ids, letting
|
||||
# a caller read another tenant's object by its key. Such objects are only
|
||||
# reachable through their managed unified id.
|
||||
if is_managed_cloud_storage_uri(file_id):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail="Raw cloud storage file ids cannot be retrieved directly. Use the LiteLLM managed file id returned when the file was created.",
|
||||
)
|
||||
# Check for model-based credential routing
|
||||
(
|
||||
should_route,
|
||||
|
|
|
|||
|
|
@ -8,6 +8,9 @@ import litellm
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.litellm_logging import use_custom_pricing_for_model
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
get_content_from_model_response,
|
||||
)
|
||||
from litellm.llms.anthropic import get_anthropic_config
|
||||
from litellm.llms.anthropic.chat.handler import (
|
||||
ModelResponseIterator as AnthropicModelResponseIterator,
|
||||
|
|
@ -136,6 +139,84 @@ class AnthropicPassthroughLoggingHandler:
|
|||
return model
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _stream_was_interrupted(
|
||||
all_chunks: Sequence[Union[str, bytes]],
|
||||
) -> bool:
|
||||
"""
|
||||
Anthropic ends a stream with ``content_block_stop`` -> ``message_delta``
|
||||
-> ``message_stop``; a client disconnect leaves the last event mid
|
||||
``content_block_delta``. Scan from the tail and decide on the first
|
||||
terminal-region event, so the common completed case is O(1) rather than
|
||||
re-deserializing every line of the stream.
|
||||
"""
|
||||
for raw in reversed(all_chunks):
|
||||
text = raw.decode("utf-8") if isinstance(raw, bytes) else raw
|
||||
for line in reversed(text.splitlines()):
|
||||
if not line.startswith("data:"):
|
||||
continue
|
||||
try:
|
||||
data = json.loads(line[len("data:") :].strip())
|
||||
except (json.JSONDecodeError, ValueError):
|
||||
continue
|
||||
if not isinstance(data, dict):
|
||||
continue
|
||||
etype = data.get("type")
|
||||
if etype == "message_delta":
|
||||
return False
|
||||
if etype in (
|
||||
"content_block_delta",
|
||||
"content_block_stop",
|
||||
"message_start",
|
||||
):
|
||||
return True
|
||||
return True
|
||||
|
||||
@staticmethod
|
||||
def _recover_interrupted_stream_output_tokens(
|
||||
response: Union[ModelResponse, TextCompletionResponse],
|
||||
all_chunks: Sequence[Union[str, bytes]],
|
||||
model: str,
|
||||
) -> None:
|
||||
"""
|
||||
An Anthropic stream interrupted before its terminal ``message_delta``
|
||||
(client disconnect) carries only the ``message_start`` ``output_tokens``
|
||||
placeholder (typically 1-3), so completion tokens and spend are
|
||||
undercounted ~20x. Re-tokenize the buffered output text to recover a
|
||||
realistic ``output_tokens`` for usage/cost. Completed streams are
|
||||
untouched because their terminal ``message_delta`` short-circuits here.
|
||||
"""
|
||||
if not isinstance(response, ModelResponse):
|
||||
return
|
||||
if not AnthropicPassthroughLoggingHandler._stream_was_interrupted(all_chunks):
|
||||
return
|
||||
usage = getattr(response, "usage", None)
|
||||
if usage is None:
|
||||
return
|
||||
output_text = get_content_from_model_response(response)
|
||||
if not output_text:
|
||||
return
|
||||
try:
|
||||
recovered_output_tokens = litellm.token_counter(
|
||||
model=model, text=output_text, count_response_tokens=True
|
||||
)
|
||||
except Exception:
|
||||
verbose_proxy_logger.warning(
|
||||
"Could not re-tokenize interrupted stream output; "
|
||||
"keeping placeholder completion token count."
|
||||
)
|
||||
return
|
||||
if recovered_output_tokens <= (usage.completion_tokens or 0):
|
||||
return
|
||||
usage.completion_tokens = recovered_output_tokens
|
||||
usage.total_tokens = (usage.prompt_tokens or 0) + recovered_output_tokens
|
||||
# Anthropic costing reads completion_tokens_details.text_tokens, so the
|
||||
# stale message_start placeholder there must be corrected too or spend
|
||||
# stays undercounted even after completion_tokens is fixed.
|
||||
details = getattr(usage, "completion_tokens_details", None)
|
||||
if details is not None and getattr(details, "text_tokens", None) is not None:
|
||||
details.text_tokens = recovered_output_tokens
|
||||
|
||||
@staticmethod
|
||||
def _create_anthropic_response_logging_payload(
|
||||
litellm_model_response: Union[ModelResponse, TextCompletionResponse],
|
||||
|
|
@ -277,6 +358,11 @@ class AnthropicPassthroughLoggingHandler:
|
|||
"result": None,
|
||||
"kwargs": {},
|
||||
}
|
||||
AnthropicPassthroughLoggingHandler._recover_interrupted_stream_output_tokens(
|
||||
response=complete_streaming_response,
|
||||
all_chunks=all_chunks,
|
||||
model=model,
|
||||
)
|
||||
kwargs = AnthropicPassthroughLoggingHandler._create_anthropic_response_logging_payload(
|
||||
litellm_model_response=complete_streaming_response,
|
||||
model=model,
|
||||
|
|
|
|||
|
|
@ -426,6 +426,9 @@ from litellm.proxy.middleware.prometheus_auth_middleware import PrometheusAuthMi
|
|||
from litellm.proxy.middleware.request_size_limit_middleware import (
|
||||
RequestSizeLimitMiddleware,
|
||||
)
|
||||
from litellm.proxy.middleware.security_headers_middleware import (
|
||||
SecurityHeadersMiddleware,
|
||||
)
|
||||
from litellm.proxy.ocr_endpoints.endpoints import router as ocr_router
|
||||
from litellm.proxy.openai_files_endpoints.files_endpoints import (
|
||||
router as openai_files_router,
|
||||
|
|
@ -840,37 +843,6 @@ async def proxy_startup_event(app: FastAPI):
|
|||
if isinstance(worker_config, dict):
|
||||
await initialize(**worker_config)
|
||||
|
||||
## V2 OTEL: now that config (and therefore the callbacks) is loaded, publish
|
||||
## the chosen V2 logger's TracerProvider as the OTel global. The FastAPI
|
||||
## instrumentation mounted at app-creation binds to the global provider, so
|
||||
## this is what makes server spans and gen-ai spans share one provider and
|
||||
## land in the same trace. Prefer an already-registered preset logger
|
||||
## (arize, langfuse, …) so server spans export to that backend too; otherwise
|
||||
## build a generic one from OTEL_* envs. ``set_tracer_provider`` only takes
|
||||
## effect once, so the first configured logger wins.
|
||||
try:
|
||||
from litellm.integrations.otel.model.config import is_otel_v2_enabled
|
||||
|
||||
if is_otel_v2_enabled():
|
||||
from opentelemetry import trace as _otel_trace
|
||||
|
||||
from litellm.integrations.otel.logger import OpenTelemetryV2
|
||||
|
||||
_otel_v2_logger = (
|
||||
next(
|
||||
(
|
||||
cb
|
||||
for cb in litellm.service_callback
|
||||
if isinstance(cb, OpenTelemetryV2)
|
||||
),
|
||||
None,
|
||||
)
|
||||
or OpenTelemetryV2()
|
||||
)
|
||||
_otel_trace.set_tracer_provider(_otel_v2_logger._tracer_provider)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug("Skipping OTel V2 provider setup: %s", e)
|
||||
|
||||
# check if DATABASE_URL in environment - load from there
|
||||
if prisma_client is None:
|
||||
_db_url: Optional[str] = get_secret("DATABASE_URL", None) # type: ignore
|
||||
|
|
@ -907,6 +879,42 @@ async def proxy_startup_event(app: FastAPI):
|
|||
redis_usage_cache=transaction_buffer_redis_cache,
|
||||
)
|
||||
|
||||
## V2 OTEL: publish the chosen V2 logger's TracerProvider as the OTel global.
|
||||
## This MUST run after callback initialization above: a preset (arize, langfuse,
|
||||
## …) builds its logger there, folding the OTEL_* base exporter and its own
|
||||
## exporter into one logger. The FastAPI instrumentation mounted at app-creation
|
||||
## binds to the global provider, so reusing that one logger is what makes the
|
||||
## server span and the gen-ai spans share one provider and land in the same
|
||||
## trace, exporting to every configured backend. Running before callback init
|
||||
## (when no logger exists yet) would build a second, generic logger whose
|
||||
## provider became the global, orphaning the gen-ai spans onto a different
|
||||
## backend than the server span. A generic logger is built only when none was
|
||||
## configured.
|
||||
try:
|
||||
from litellm.integrations.otel.model.config import is_otel_v2_enabled
|
||||
|
||||
if is_otel_v2_enabled():
|
||||
from opentelemetry import trace as _otel_trace
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import _in_memory_loggers
|
||||
from litellm.integrations.otel.logger import (
|
||||
OpenTelemetryV2,
|
||||
publish_global_otel_v2_provider,
|
||||
)
|
||||
|
||||
registered = (
|
||||
open_telemetry_logger
|
||||
if isinstance(open_telemetry_logger, OpenTelemetryV2)
|
||||
else None
|
||||
)
|
||||
publish_global_otel_v2_provider(
|
||||
_in_memory_loggers, # any-ok: pre-existing untyped List[Any] global
|
||||
_otel_trace.set_tracer_provider,
|
||||
registered=registered,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug("Skipping OTel V2 provider setup: %s", e)
|
||||
|
||||
## Validate use_redis_transaction_buffer requires Redis cache ##
|
||||
ProxyStartupEvent._validate_redis_transaction_buffer_config(
|
||||
general_settings=general_settings,
|
||||
|
|
@ -1757,6 +1765,7 @@ app.add_middleware(
|
|||
|
||||
app.add_middleware(PrometheusAuthMiddleware)
|
||||
app.add_middleware(InFlightRequestsMiddleware)
|
||||
app.add_middleware(SecurityHeadersMiddleware)
|
||||
|
||||
|
||||
def mount_swagger_ui():
|
||||
|
|
@ -2026,7 +2035,43 @@ def cost_tracking():
|
|||
)
|
||||
|
||||
|
||||
async def get_current_spend(counter_key: str, fallback_spend: float) -> float:
|
||||
# Bounds authoritative DB re-reads when enforcing a budget against a
|
||||
# stale-low spend counter: at most one DB read per counter per window.
|
||||
SPEND_DB_FLOOR_CACHE_TTL_SECONDS = 5
|
||||
|
||||
|
||||
def _fail_closed_budget_enforcement() -> bool:
|
||||
return general_settings.get("fail_closed_budget_enforcement") is True
|
||||
|
||||
|
||||
def _raise_budget_unverifiable(counter_key: str) -> None:
|
||||
verbose_proxy_logger.warning(
|
||||
"fail_closed_budget_enforcement: rejecting request — spend for %s could "
|
||||
"not be verified against Redis or the database",
|
||||
counter_key,
|
||||
)
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_503_SERVICE_UNAVAILABLE,
|
||||
detail={
|
||||
"error": (
|
||||
"Budget enforcement unavailable: current spend could not be "
|
||||
"verified against Redis or the database, and "
|
||||
"fail_closed_budget_enforcement is enabled, so the request was "
|
||||
"rejected to avoid exceeding the configured budget. Retry shortly."
|
||||
)
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
async def get_current_spend(
|
||||
counter_key: str,
|
||||
fallback_spend: float,
|
||||
max_budget: float | None = None,
|
||||
window_entity_type: str | None = None,
|
||||
window_entity_id: str | None = None,
|
||||
window_start: datetime | None = None,
|
||||
fallback_authoritative: bool = False,
|
||||
) -> float:
|
||||
"""
|
||||
Read current spend from the cross-pod spend counter.
|
||||
|
||||
|
|
@ -2040,7 +2085,168 @@ async def get_current_spend(counter_key: str, fallback_spend: float) -> float:
|
|||
2. In-memory counter (single-instance or Redis failure)
|
||||
3. Reseed from authoritative DB spend (counter expired, cross-pod stale)
|
||||
4. Caller-supplied fallback (DB unavailable, cold start)
|
||||
|
||||
When ``max_budget`` is supplied, the counter is re-checked against the
|
||||
authoritative recorded spend before a request is admitted. A Redis counter
|
||||
that survived a Redis restart can return a stale-low value loaded from an
|
||||
older RDB snapshot; that read is a hit (not a clean miss), so step 3 never
|
||||
runs and a key can leak spend past ``max_budget`` indefinitely. The
|
||||
authoritative source depends on the counter: primary key/team/user/org
|
||||
counters read the DB row; per-window counters (``window_start`` supplied)
|
||||
aggregate spend logs; end-user/tag counters have no DB row, so the caller's
|
||||
``fallback_spend`` (loaded fresh in auth) is authoritative. The DB read is
|
||||
skipped for healthy primary counters (counter at or above recorded spend)
|
||||
and cached in-process for a few seconds, so a persistently stale counter
|
||||
drives at most one read per counter per window rather than one per request.
|
||||
"""
|
||||
current, verified = await _read_spend_counter_estimate(
|
||||
counter_key=counter_key, fallback_spend=fallback_spend
|
||||
)
|
||||
if fallback_authoritative:
|
||||
verified = True
|
||||
|
||||
if max_budget is None or current >= max_budget:
|
||||
return current
|
||||
|
||||
# Cheap staleness signal for primary counters: the counter reads below the
|
||||
# spend this caller already knows about. Window counters have no such signal
|
||||
# (fallback is 0), so they always re-check, bounded by the cache. Strict mode
|
||||
# (fail_closed_budget_enforcement) always re-checks against the authoritative
|
||||
# source too, so a counter that is stale-low at the same time as the caller's
|
||||
# cached spend cannot slip through; the 5s cache keeps that bounded.
|
||||
is_window = window_start is not None
|
||||
if fallback_spend > current or is_window or _fail_closed_budget_enforcement():
|
||||
authoritative = await _authoritative_floor_spend(
|
||||
counter_key=counter_key,
|
||||
window_entity_type=window_entity_type,
|
||||
window_entity_id=window_entity_id,
|
||||
window_start=window_start,
|
||||
)
|
||||
if authoritative is not None:
|
||||
verified = True
|
||||
if authoritative > current:
|
||||
await _repair_stale_spend_counter(
|
||||
counter_key=counter_key, db_spend=authoritative
|
||||
)
|
||||
return authoritative
|
||||
elif fallback_spend > current:
|
||||
# end-user / tag counters have no DB row; fallback_spend is the
|
||||
# authoritative recorded value loaded in auth.
|
||||
return fallback_spend
|
||||
|
||||
# Opt-in hard guarantee: when the spend backing this admit decision came
|
||||
# only from a per-pod cache (Redis and DB both unreadable), reject rather
|
||||
# than admit on an unverifiable budget. No-op unless the flag is set, so
|
||||
# default behavior is unchanged.
|
||||
if not verified and _fail_closed_budget_enforcement():
|
||||
_raise_budget_unverifiable(counter_key)
|
||||
|
||||
return current
|
||||
|
||||
|
||||
async def _repair_stale_spend_counter(counter_key: str, db_spend: float) -> None:
|
||||
"""Raise a counter that has fallen below the authoritative DB spend (e.g.
|
||||
Redis restarted and reloaded an older snapshot) so every worker reads the
|
||||
corrected value directly instead of re-deriving it per request, and so a
|
||||
worker whose own cached spend is also stale still sees the true total.
|
||||
|
||||
The write is monotonic: it only ever raises the counter, so a repair that
|
||||
carries a slightly-stale DB total cannot clobber a concurrent increment that
|
||||
already pushed the counter higher (which would let racing requests
|
||||
under-count). Redis enforces this atomically via async_set_max; the
|
||||
in-memory copy is guarded by a read-compare-write with no await in between,
|
||||
so it is atomic within the worker.
|
||||
"""
|
||||
cached = spend_counter_cache.in_memory_cache.get_cache(key=counter_key)
|
||||
needs_update = True
|
||||
if cached is not None:
|
||||
try:
|
||||
needs_update = float(cached) < db_spend
|
||||
except (TypeError, ValueError):
|
||||
needs_update = True
|
||||
if needs_update:
|
||||
spend_counter_cache.in_memory_cache.set_cache(key=counter_key, value=db_spend)
|
||||
if spend_counter_cache.redis_cache is not None:
|
||||
try:
|
||||
await spend_counter_cache.redis_cache.async_set_max(
|
||||
key=counter_key, value=db_spend
|
||||
)
|
||||
except Exception:
|
||||
verbose_proxy_logger.debug(
|
||||
"Unable to repair stale spend counter %s in Redis",
|
||||
counter_key,
|
||||
exc_info=True,
|
||||
)
|
||||
|
||||
|
||||
async def reseed_spend_counter_from_db(counter_key: str) -> None:
|
||||
"""Recover a counter that the reservation reconcile found in an inconsistent
|
||||
state (missing, or where applying the reconcile delta would drive it
|
||||
negative) by reseeding it from the DB instead of deleting it.
|
||||
|
||||
The DB row is a LAGGING authoritative floor, not post-request truth: the
|
||||
entity .spend column is flushed in batches (every PROXY_BATCH_WRITE_AT), so
|
||||
it can exclude this request's just-recorded cost and other buffered spend.
|
||||
That is fine here: the monotonic set-max can only RAISE a stale-low counter
|
||||
toward that floor (never lowers it or clobbers a concurrent increment), and
|
||||
the read-time floor (_authoritative_floor_spend) converges to the true total
|
||||
as the buffer flushes. The point is to restore enforcement to a real floor
|
||||
rather than leave the counter deleted and unenforced (the prior fail-open).
|
||||
Counters with no DB row (window/end-user/tag) are left untouched rather than
|
||||
deleted, so enforcement keeps reading whatever value they hold.
|
||||
"""
|
||||
db_spend = await SpendCounterReseed.from_db(
|
||||
prisma_client=prisma_client, counter_key=counter_key
|
||||
)
|
||||
if db_spend is None:
|
||||
return
|
||||
await _repair_stale_spend_counter(counter_key=counter_key, db_spend=db_spend)
|
||||
|
||||
|
||||
async def _authoritative_floor_spend(
|
||||
counter_key: str,
|
||||
window_entity_type: str | None = None,
|
||||
window_entity_id: str | None = None,
|
||||
window_start: datetime | None = None,
|
||||
) -> float | None:
|
||||
marker_key = f"spend_db_floor:{counter_key}"
|
||||
cached = spend_counter_cache.in_memory_cache.get_cache(key=marker_key)
|
||||
if cached is not None:
|
||||
return float(cached)
|
||||
|
||||
db_spend = await SpendCounterReseed.from_db(
|
||||
prisma_client=prisma_client, counter_key=counter_key
|
||||
)
|
||||
if (
|
||||
db_spend is None
|
||||
and window_entity_type is not None
|
||||
and window_entity_id is not None
|
||||
and window_start is not None
|
||||
):
|
||||
db_spend = await SpendCounterReseed.window_from_spend_logs(
|
||||
prisma_client=prisma_client,
|
||||
entity_type=window_entity_type,
|
||||
entity_id=window_entity_id,
|
||||
window_start=window_start,
|
||||
)
|
||||
if db_spend is None:
|
||||
return None
|
||||
|
||||
spend_counter_cache.in_memory_cache.set_cache(
|
||||
key=marker_key,
|
||||
value=db_spend,
|
||||
ttl=SPEND_DB_FLOOR_CACHE_TTL_SECONDS,
|
||||
)
|
||||
return db_spend
|
||||
|
||||
|
||||
async def _read_spend_counter_estimate(
|
||||
counter_key: str, fallback_spend: float
|
||||
) -> tuple[float, bool]:
|
||||
"""Return (spend, authoritative). ``authoritative`` is True when the value
|
||||
came from Redis or a fresh DB read (cross-pod truth), False when it came
|
||||
from the per-pod in-memory copy or the caller's fallback. Only the
|
||||
fail-closed path reads the flag; normal callers ignore it."""
|
||||
# 1. Redis first (cross-pod authoritative). On clean miss, skip
|
||||
# in-memory: per-pod in-memory only has this pod's writes, so it
|
||||
# would mask cross-pod increments.
|
||||
|
|
@ -2049,7 +2255,7 @@ async def get_current_spend(counter_key: str, fallback_spend: float) -> float:
|
|||
try:
|
||||
val = await spend_counter_cache.redis_cache.async_get_cache(key=counter_key)
|
||||
if val is not None:
|
||||
return float(val)
|
||||
return float(val), True
|
||||
redis_clean_miss = True
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
@ -2062,7 +2268,7 @@ async def get_current_spend(counter_key: str, fallback_spend: float) -> float:
|
|||
if not redis_clean_miss:
|
||||
val = spend_counter_cache.in_memory_cache.get_cache(key=counter_key)
|
||||
if val is not None:
|
||||
return float(val)
|
||||
return float(val), False
|
||||
|
||||
# 3. Reseed from DB - fallback_spend lags cross-pod, would allow bypass.
|
||||
db_spend = await SpendCounterReseed.coalesced(
|
||||
|
|
@ -2071,10 +2277,10 @@ async def get_current_spend(counter_key: str, fallback_spend: float) -> float:
|
|||
counter_key=counter_key,
|
||||
)
|
||||
if db_spend is not None:
|
||||
return db_spend
|
||||
return db_spend, True
|
||||
|
||||
# 4. Caller-supplied fallback (DB unavailable).
|
||||
return fallback_spend
|
||||
return fallback_spend, False
|
||||
|
||||
|
||||
async def increment_spend_counters(
|
||||
|
|
@ -8758,7 +8964,7 @@ async def chat_completion(
|
|||
completion_stream=_iterator,
|
||||
model=e.model,
|
||||
custom_llm_provider="cached_response",
|
||||
logging_obj=data.get("litellm_logging_obj", None),
|
||||
logging_obj=_data.get("litellm_logging_obj", None),
|
||||
)
|
||||
selected_data_generator = select_data_generator(
|
||||
response=_streaming_response,
|
||||
|
|
@ -8793,7 +8999,7 @@ async def chat_completion(
|
|||
completion_stream=_iterator,
|
||||
model=data.get("model", ""),
|
||||
custom_llm_provider="cached_response",
|
||||
logging_obj=data.get("litellm_logging_obj", None),
|
||||
logging_obj=_data.get("litellm_logging_obj", None),
|
||||
)
|
||||
selected_data_generator = select_data_generator(
|
||||
response=_streaming_response,
|
||||
|
|
@ -12940,6 +13146,9 @@ async def model_info_v1(
|
|||
# use internal routing keys (model_name_{team_id}_{uuid}) and were omitted
|
||||
# when v1 resolved models only via public model_name strings.
|
||||
all_models: List[dict] = copy.deepcopy(llm_router.model_list)
|
||||
alias_models = copy.deepcopy(llm_router.get_model_list_from_model_alias())
|
||||
all_models.extend(alias_models)
|
||||
|
||||
allowed_model_names = _get_v1_model_info_allowed_model_names(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
llm_router=llm_router,
|
||||
|
|
@ -13507,26 +13716,24 @@ async def fallback_login(request: Request):
|
|||
|
||||
# get url from request
|
||||
redirect_url = get_custom_url(str(request.base_url))
|
||||
ui_username = os.getenv("UI_USERNAME")
|
||||
if redirect_url.endswith("/"):
|
||||
redirect_url += "sso/callback"
|
||||
else:
|
||||
redirect_url += "/sso/callback"
|
||||
|
||||
if ui_username is not None:
|
||||
# No Google, Microsoft SSO
|
||||
# Use UI Credentials set in .env
|
||||
from fastapi.responses import HTMLResponse
|
||||
from fastapi.responses import HTMLResponse
|
||||
|
||||
return HTMLResponse(
|
||||
content=build_ui_login_form(show_deprecation_banner=False), status_code=200
|
||||
)
|
||||
else:
|
||||
from fastapi.responses import HTMLResponse
|
||||
|
||||
return HTMLResponse(
|
||||
content=build_ui_login_form(show_deprecation_banner=False), status_code=200
|
||||
)
|
||||
hide_default_credentials_hint = (
|
||||
os.getenv("LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT", "false").lower() == "true"
|
||||
or general_settings.get("hide_default_credentials_hint", False) is True
|
||||
)
|
||||
return HTMLResponse(
|
||||
content=build_ui_login_form(
|
||||
show_deprecation_banner=False,
|
||||
hide_default_credentials_hint=hide_default_credentials_hint,
|
||||
),
|
||||
status_code=200,
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
|
|
|
|||
|
|
@ -2596,7 +2596,7 @@
|
|||
"default_value": null
|
||||
}
|
||||
],
|
||||
"default_model_placeholder": "soniox/stt-async-v4"
|
||||
"default_model_placeholder": "soniox/stt-async-v5"
|
||||
},
|
||||
{
|
||||
"provider": "TEXT_COMPLETION_CODESTRAL",
|
||||
|
|
|
|||
|
|
@ -130,6 +130,12 @@ async def _prepare_client_secret_session(
|
|||
session_model = req.session.model if req.session else None
|
||||
model: str = session_model or req.model or _DEFAULT_REALTIME_MODEL
|
||||
if session_type != "transcription":
|
||||
await can_key_call_resolved_model(
|
||||
model=model,
|
||||
valid_token=user_api_key_dict,
|
||||
llm_model_list=llm_model_list,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
return model, session_data, session_type
|
||||
|
||||
transcription_model_candidates = _transcription_model_candidates_from_session(
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
|
@ -162,10 +163,14 @@ async def reserve_budget_for_request(
|
|||
if not applied_entries:
|
||||
return None
|
||||
|
||||
input_cost = estimate_request_input_cost(
|
||||
request_body=request_body, route=route, llm_router=llm_router
|
||||
)
|
||||
return {
|
||||
"reserved_cost": reservation_cost,
|
||||
"entries": applied_entries,
|
||||
"finalized": False,
|
||||
"input_cost": min(float(input_cost or 0.0), reservation_cost),
|
||||
}
|
||||
|
||||
|
||||
|
|
@ -195,6 +200,41 @@ async def release_budget_reservation(budget_reservation: Optional[dict]) -> None
|
|||
)
|
||||
|
||||
|
||||
async def release_budget_reservation_on_cancel(
|
||||
budget_reservation: dict | None,
|
||||
) -> None:
|
||||
"""Reconcile a still-open reservation when the request is cancelled mid-flight.
|
||||
|
||||
A client disconnect or timeout cancels the request task, which surfaces as
|
||||
CancelledError / GeneratorExit rather than a normal exception, so neither the
|
||||
success cost callback nor the failure hook runs and the pre-call reservation
|
||||
is never reconciled. Left alone it pins the spend counter above real spend
|
||||
and 429s subsequent requests until the counter's TTL expires.
|
||||
|
||||
Reconcile to the request's input-token cost rather than refunding to zero:
|
||||
by the time a request is cancelled in-flight the provider call was already
|
||||
dispatched, so the input tokens were billed even if no chunk reached the
|
||||
client. Refunding to zero would let a caller abort pre-token to dodge that
|
||||
charge; the worst-case output portion of the reservation is still released.
|
||||
|
||||
asyncio.shield keeps the reconcile running to completion even though the
|
||||
surrounding task is being cancelled. The `finalized` guard makes this a no-op
|
||||
when success/failure handling already reconciled, so calling it on every
|
||||
cancellation path is safe.
|
||||
"""
|
||||
if not budget_reservation or budget_reservation.get("finalized") is True:
|
||||
return
|
||||
incurred_cost = float(budget_reservation.get("input_cost") or 0.0)
|
||||
try:
|
||||
await asyncio.shield(
|
||||
reconcile_budget_reservation(
|
||||
budget_reservation=budget_reservation, actual_cost=incurred_cost
|
||||
)
|
||||
)
|
||||
except (asyncio.CancelledError, Exception):
|
||||
pass
|
||||
|
||||
|
||||
async def invalidate_budget_reservation_counters(
|
||||
budget_reservation: Optional[dict],
|
||||
) -> None:
|
||||
|
|
@ -628,12 +668,14 @@ async def _set_reserved_entries_actual_cost(
|
|||
entries: List[dict],
|
||||
actual_cost: float,
|
||||
default_reserved_cost: float,
|
||||
reseed_on_inconsistent: bool = True,
|
||||
) -> None:
|
||||
for entry in entries:
|
||||
await _set_reserved_entry_actual_cost(
|
||||
entry=entry,
|
||||
actual_cost=actual_cost,
|
||||
default_reserved_cost=default_reserved_cost,
|
||||
reseed_on_inconsistent=reseed_on_inconsistent,
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -641,8 +683,12 @@ async def _set_reserved_entry_actual_cost(
|
|||
entry: dict,
|
||||
actual_cost: float,
|
||||
default_reserved_cost: float,
|
||||
reseed_on_inconsistent: bool = True,
|
||||
) -> None:
|
||||
from litellm.proxy.proxy_server import _increment_spend_counter_cache
|
||||
from litellm.proxy.proxy_server import (
|
||||
_increment_spend_counter_cache,
|
||||
reseed_spend_counter_from_db,
|
||||
)
|
||||
|
||||
counter_key = entry.get("counter_key")
|
||||
if counter_key is None:
|
||||
|
|
@ -656,46 +702,49 @@ async def _set_reserved_entry_actual_cost(
|
|||
adjustment = target_adjustment - applied_adjustment
|
||||
if adjustment == 0:
|
||||
return
|
||||
await _ensure_counter_can_apply_adjustment(
|
||||
if await _counter_can_apply_adjustment(
|
||||
counter_key=counter_key,
|
||||
adjustment=adjustment,
|
||||
)
|
||||
await _increment_spend_counter_cache(
|
||||
counter_key=counter_key,
|
||||
increment=adjustment,
|
||||
)
|
||||
):
|
||||
await _increment_spend_counter_cache(
|
||||
counter_key=counter_key,
|
||||
increment=adjustment,
|
||||
)
|
||||
elif reseed_on_inconsistent:
|
||||
# Post-call reconcile / release: the counter was flushed or reseeded
|
||||
# between reservation and reconcile (Redis restart / cross-pod reset),
|
||||
# so the optimistic delta no longer applies. Recover by reseeding from
|
||||
# the DB's lagging authoritative floor rather than deleting the counter
|
||||
# and failing open — deleting it is what left budgets unenforced after a
|
||||
# Redis reload.
|
||||
await reseed_spend_counter_from_db(counter_key=counter_key)
|
||||
else:
|
||||
# Pre-call admission resize: the in-flight reservation cost is not yet
|
||||
# persisted, so the DB floor would discard it. Keep the original
|
||||
# fail-closed behavior (raise -> reserve_budget_for_request releases and
|
||||
# denies) rather than admitting against an inconsistent counter.
|
||||
raise RuntimeError(
|
||||
f"Cannot resize budget reservation against inconsistent counter {counter_key}"
|
||||
)
|
||||
entry["applied_adjustment"] = target_adjustment
|
||||
|
||||
|
||||
async def _ensure_counter_can_apply_adjustment(
|
||||
async def _counter_can_apply_adjustment(
|
||||
counter_key: str,
|
||||
adjustment: float,
|
||||
) -> None:
|
||||
from litellm.proxy.proxy_server import (
|
||||
_invalidate_spend_counter,
|
||||
spend_counter_cache,
|
||||
)
|
||||
) -> bool:
|
||||
from litellm.proxy.proxy_server import spend_counter_cache
|
||||
|
||||
current_value = await spend_counter_cache.async_get_cache(key=counter_key)
|
||||
if current_value is None:
|
||||
await _invalidate_spend_counter(counter_key=counter_key)
|
||||
raise RuntimeError(
|
||||
f"Cannot apply budget reservation adjustment to missing counter {counter_key}"
|
||||
)
|
||||
return False
|
||||
|
||||
try:
|
||||
current_float = float(current_value)
|
||||
except (TypeError, ValueError):
|
||||
await _invalidate_spend_counter(counter_key=counter_key)
|
||||
raise RuntimeError(
|
||||
f"Cannot apply budget reservation adjustment to non-numeric counter {counter_key}"
|
||||
)
|
||||
return False
|
||||
|
||||
if adjustment < 0 and current_float + adjustment < -1e-12:
|
||||
await _invalidate_spend_counter(counter_key=counter_key)
|
||||
raise RuntimeError(
|
||||
f"Budget reservation adjustment would make counter negative {counter_key}"
|
||||
)
|
||||
return not (adjustment < 0 and current_float + adjustment < -1e-12)
|
||||
|
||||
|
||||
async def _release_applied_entries_best_effort(
|
||||
|
|
@ -735,6 +784,7 @@ async def _resize_applied_reservation(
|
|||
entries=entries,
|
||||
actual_cost=new_reserved_cost,
|
||||
default_reserved_cost=current_reserved_cost,
|
||||
reseed_on_inconsistent=False,
|
||||
)
|
||||
for entry in entries:
|
||||
entry["reserved_cost"] = new_reserved_cost
|
||||
|
|
@ -817,6 +867,61 @@ def estimate_request_max_cost(
|
|||
return max(cast(List[float], estimates))
|
||||
|
||||
|
||||
def estimate_request_input_cost(
|
||||
request_body: dict,
|
||||
route: str,
|
||||
llm_router: Router | None,
|
||||
) -> float | None:
|
||||
"""Cost of the request's input tokens alone.
|
||||
|
||||
Once the provider request is dispatched the input tokens are billed even if
|
||||
the client disconnects before the first chunk, so this is the cost floor a
|
||||
cancelled in-flight request has already incurred. A cancelled reservation is
|
||||
reconciled to this instead of being refunded to zero.
|
||||
"""
|
||||
model = get_model_from_request(request_body, route, llm_router=llm_router)
|
||||
if model is None:
|
||||
return None
|
||||
|
||||
models = [model] if isinstance(model, str) else model
|
||||
estimates = [
|
||||
_estimate_request_input_cost_for_model(
|
||||
request_body=request_body,
|
||||
route=route,
|
||||
model=model_name,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
for model_name in models
|
||||
]
|
||||
estimates = [estimate for estimate in estimates if estimate is not None]
|
||||
if not estimates:
|
||||
return None
|
||||
return max(cast("list[float]", estimates))
|
||||
|
||||
|
||||
def _estimate_request_input_cost_for_model(
|
||||
request_body: dict,
|
||||
route: str,
|
||||
model: str,
|
||||
llm_router: Router | None,
|
||||
) -> float | None:
|
||||
model_info = _get_model_cost_info(model=model, llm_router=llm_router)
|
||||
if model_info is None:
|
||||
return None
|
||||
input_cost_per_token = _to_float(model_info.get("input_cost_per_token"))
|
||||
if input_cost_per_token is None:
|
||||
return None
|
||||
input_tokens = _estimate_input_tokens(
|
||||
request_body=request_body,
|
||||
route=route,
|
||||
model=model,
|
||||
model_info=model_info,
|
||||
)
|
||||
if input_tokens is None:
|
||||
return None
|
||||
return input_tokens * input_cost_per_token
|
||||
|
||||
|
||||
def _estimate_request_max_cost_for_model(
|
||||
request_body: dict,
|
||||
route: str,
|
||||
|
|
|
|||
|
|
@ -263,6 +263,13 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
|
|||
elif isinstance(_usage, dict):
|
||||
usage = _usage
|
||||
|
||||
# A request that failed mid-stream has no usable response_obj usage, but the
|
||||
# streaming handler may have recovered the usage from the chunks already
|
||||
# delivered. Honor that override so the partial usage lands in spend tracking.
|
||||
_combined_usage = kwargs.get("combined_usage_object")
|
||||
if not usage and isinstance(_combined_usage, litellm.Usage):
|
||||
usage = _combined_usage.model_dump()
|
||||
|
||||
id = get_spend_logs_id(call_type or "acompletion", response_obj_dict, kwargs)
|
||||
standard_logging_payload = cast(
|
||||
Optional[StandardLoggingPayload], kwargs.get("standard_logging_object", None)
|
||||
|
|
|
|||
|
|
@ -14,13 +14,11 @@ from litellm.proxy._types import *
|
|||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.repositories.config_repository import ConfigRepository
|
||||
from litellm.repositories.table_repositories import (
|
||||
DailyTagSpendRepository,
|
||||
SSOConfigRepository,
|
||||
UISettingsRepository,
|
||||
)
|
||||
from litellm.types.proxy.management_endpoints.ui_sso import (
|
||||
DefaultTeamSSOParams,
|
||||
InProductNudgeResponse,
|
||||
SSOConfig,
|
||||
)
|
||||
|
||||
|
|
@ -178,11 +176,6 @@ class UISettings(BaseModel):
|
|||
description="If true, org admins cannot generate API keys via /key/generate.",
|
||||
)
|
||||
|
||||
disable_ui_nudges: bool = Field(
|
||||
default=False,
|
||||
description="If true, suppresses in-product UI nudges (survey and Claude Code feedback popups) for all users.",
|
||||
)
|
||||
|
||||
|
||||
class UISettingsResponse(SettingsResponse):
|
||||
"""Response model for UI settings"""
|
||||
|
|
@ -206,7 +199,6 @@ ALLOWED_UI_SETTINGS_FIELDS = {
|
|||
"scope_user_search_to_org",
|
||||
"disable_custom_api_keys",
|
||||
"disable_key_generate_for_org_admin",
|
||||
"disable_ui_nudges",
|
||||
}
|
||||
|
||||
# Flags that must be synced from the persisted UISettings into
|
||||
|
|
@ -1117,34 +1109,6 @@ async def update_mcp_semantic_filter_settings(
|
|||
return result
|
||||
|
||||
|
||||
@router.get(
|
||||
"/in_product_nudges",
|
||||
tags=["UI Settings"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=InProductNudgeResponse,
|
||||
)
|
||||
async def get_in_product_nudges():
|
||||
"""
|
||||
Get in-product nudges configuration.
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(
|
||||
status_code=500,
|
||||
detail={"error": "Database not connected. Please connect a database."},
|
||||
)
|
||||
|
||||
db_record = await DailyTagSpendRepository(prisma_client).table.find_first(
|
||||
where={"tag": "User-Agent: claude-cli"}
|
||||
)
|
||||
|
||||
if db_record:
|
||||
return InProductNudgeResponse(is_claude_code_enabled=True)
|
||||
|
||||
return InProductNudgeResponse(is_claude_code_enabled=False)
|
||||
|
||||
|
||||
UI_SETTINGS_CACHE_KEY = "ui_settings:settings_dict"
|
||||
UI_SETTINGS_CACHE_TTL = 600 # 10 minutes
|
||||
|
||||
|
|
|
|||
|
|
@ -2128,12 +2128,21 @@ class ProxyLogging:
|
|||
# compute preprocessing latency after the logging object is popped.
|
||||
_logging_obj = request_data.get("litellm_logging_obj")
|
||||
if _logging_obj is not None:
|
||||
_first_handoff = getattr(_logging_obj, "model_call_details", {}).get(
|
||||
"first_api_call_start_time"
|
||||
)
|
||||
_model_call_details = getattr(_logging_obj, "model_call_details", {})
|
||||
_first_handoff = _model_call_details.get("first_api_call_start_time")
|
||||
if _first_handoff is not None:
|
||||
request_data["first_api_call_start_time"] = _first_handoff
|
||||
|
||||
# A stream that broke mid-flight still billed the provider for the
|
||||
# chunks already delivered; the streaming handler stashes that
|
||||
# recovered usage and cost here. Lift them onto request_data so the
|
||||
# failure-path spend callbacks (which run after the logging object
|
||||
# is popped) record the real partial spend instead of zero.
|
||||
_recovered_usage = _model_call_details.get("combined_usage_object")
|
||||
if _recovered_usage is not None:
|
||||
request_data["combined_usage_object"] = _recovered_usage
|
||||
request_data["response_cost"] = _model_call_details.get("response_cost")
|
||||
|
||||
# Remove before callbacks iterate — not serialisable
|
||||
request_data.pop("litellm_logging_obj", None)
|
||||
|
||||
|
|
@ -3432,6 +3441,7 @@ class PrismaClient:
|
|||
r.expires = r.expires.isoformat()
|
||||
elif query_type == "find_all" and team_id is not None:
|
||||
response = await VerificationTokenRepository(self).table.find_many(
|
||||
take=limit,
|
||||
where={"team_id": team_id},
|
||||
include={"litellm_budget_table": True},
|
||||
)
|
||||
|
|
@ -4452,6 +4462,14 @@ class PrismaClient:
|
|||
"prisma-query-engine PID %s already dead at watch start.",
|
||||
pid,
|
||||
)
|
||||
if self._consume_expected_death(pid):
|
||||
verbose_proxy_logger.info(
|
||||
"PID %s death was planned (engine already replaced); "
|
||||
"not reconnecting.",
|
||||
pid,
|
||||
)
|
||||
self._cleanup_engine_watcher()
|
||||
return True
|
||||
self._engine_confirmed_dead = True
|
||||
self._reap_all_zombies()
|
||||
self._cleanup_engine_watcher()
|
||||
|
|
@ -4496,12 +4514,39 @@ class PrismaClient:
|
|||
except RuntimeError:
|
||||
pass
|
||||
|
||||
def _consume_expected_death(self, pid: int) -> bool:
|
||||
"""True iff ``pid`` was killed on purpose by a planned recreate.
|
||||
|
||||
`PrismaWrapper.recreate_prisma_client` records the old engine PID in
|
||||
`_expected_engine_deaths` before SIGTERM-ing it (IAM token refresh,
|
||||
guarded reconnect). When the watcher then sees that PID die, this lets
|
||||
it recognize the death as planned and skip its own reconnect, which
|
||||
would otherwise kill the engine the recreate just spawned (#29176).
|
||||
|
||||
Consumes (removes) the PID so a later real crash of a reused PID is
|
||||
still handled. Tolerant of `self.db` stand-ins (tests / older clients)
|
||||
that don't expose a real set.
|
||||
"""
|
||||
expected = getattr(self.db, "_expected_engine_deaths", None)
|
||||
if isinstance(expected, set) and pid in expected:
|
||||
expected.discard(pid)
|
||||
return True
|
||||
return False
|
||||
|
||||
def _on_engine_death_from_thread(self, dead_pid: int) -> None:
|
||||
"""Called on the event loop thread when the waitpid thread detects engine death."""
|
||||
if self._engine_confirmed_dead:
|
||||
return
|
||||
if dead_pid != self._engine_pid:
|
||||
return
|
||||
if self._consume_expected_death(dead_pid):
|
||||
verbose_proxy_logger.info(
|
||||
"prisma-query-engine PID %s exited as part of a planned restart; "
|
||||
"not reconnecting (engine already replaced).",
|
||||
dead_pid,
|
||||
)
|
||||
self._cleanup_engine_watcher()
|
||||
return
|
||||
verbose_proxy_logger.error(
|
||||
"prisma-query-engine PID %s exited (waitpid thread); triggering reconnect.",
|
||||
dead_pid,
|
||||
|
|
@ -4556,6 +4601,14 @@ class PrismaClient:
|
|||
self._engine_pidfd = -1
|
||||
return
|
||||
dead_pid = self._engine_pid
|
||||
if self._consume_expected_death(dead_pid):
|
||||
verbose_proxy_logger.info(
|
||||
"prisma-query-engine PID %s exited (pidfd event) as part of a "
|
||||
"planned restart; not reconnecting (engine already replaced).",
|
||||
dead_pid,
|
||||
)
|
||||
self._cleanup_engine_watcher()
|
||||
return
|
||||
verbose_proxy_logger.error(
|
||||
"prisma-query-engine PID %s exited (pidfd event); triggering reconnect.",
|
||||
dead_pid,
|
||||
|
|
@ -4579,9 +4632,18 @@ class PrismaClient:
|
|||
try:
|
||||
os.kill(self._engine_pid, 0)
|
||||
except ProcessLookupError:
|
||||
dead_pid = self._engine_pid
|
||||
if self._consume_expected_death(dead_pid):
|
||||
verbose_proxy_logger.info(
|
||||
"prisma-query-engine PID %s gone as part of a planned "
|
||||
"restart; not reconnecting (engine already replaced).",
|
||||
dead_pid,
|
||||
)
|
||||
self._cleanup_engine_watcher()
|
||||
return
|
||||
verbose_proxy_logger.error(
|
||||
"prisma-query-engine PID %s gone; triggering reconnect.",
|
||||
self._engine_pid,
|
||||
dead_pid,
|
||||
)
|
||||
self._engine_confirmed_dead = True
|
||||
self._reap_all_zombies()
|
||||
|
|
@ -4668,6 +4730,22 @@ class PrismaClient:
|
|||
self._engine_confirmed_dead = False
|
||||
verbose_proxy_logger.debug("Stopped engine process watcher.")
|
||||
|
||||
def _handle_writer_engine_replaced(self) -> None:
|
||||
"""Re-arm the engine watcher after a planned writer-engine restart.
|
||||
|
||||
Wired as `PrismaWrapper.on_engine_replaced` and invoked from inside
|
||||
`recreate_prisma_client` once the new engine is connected (IAM token
|
||||
refresh, guarded reconnect). The old watcher was tracking the engine
|
||||
we just intentionally killed, so we tear it down and re-arm on the new
|
||||
PID. Scheduling `_start_engine_watcher` as a task (rather than awaiting)
|
||||
keeps us from blocking the recreate while it still holds the wrapper's
|
||||
reconnection lock. Without this re-arm, a planned restart would leave
|
||||
the proxy with no engine-death detection until the next reconnect.
|
||||
"""
|
||||
self._engine_confirmed_dead = False
|
||||
self._cleanup_engine_watcher()
|
||||
asyncio.create_task(self._start_engine_watcher())
|
||||
|
||||
async def _run_reconnect_cycle(
|
||||
self, timeout_seconds: Optional[float] = None
|
||||
) -> None:
|
||||
|
|
@ -4688,6 +4766,17 @@ class PrismaClient:
|
|||
else self._db_watchdog_reconnect_timeout_seconds
|
||||
)
|
||||
|
||||
# Snapshot the writer's engine generation BEFORE any await. Both
|
||||
# reconnect branches forward it to recreate_prisma_client as an
|
||||
# optimistic-lock token: if a concurrent IAM token refresh replaces the
|
||||
# engine after this point, the generation moves and the recreate becomes
|
||||
# a no-op instead of killing the engine the refresh just spawned
|
||||
# (#29176). Captured here — atomically with the dead-engine decision
|
||||
# below — rather than inside the reconnect closures, because those run
|
||||
# after an `asyncio.wait_for(...)` yield during which a refresh could
|
||||
# otherwise slip in and bump the very generation the closure then reads.
|
||||
expected_generation = getattr(self.writer_db, "_engine_generation", None)
|
||||
|
||||
engine_is_dead = self._engine_confirmed_dead or (
|
||||
self._engine_pid > 0 and not self._is_engine_alive()
|
||||
)
|
||||
|
|
@ -4708,7 +4797,16 @@ class PrismaClient:
|
|||
"DATABASE_URL not set; cannot recreate Prisma client."
|
||||
)
|
||||
raise RuntimeError("DATABASE_URL not set")
|
||||
await self.db.recreate_prisma_client(db_url)
|
||||
# Forward the entry-snapshot generation. The engine was
|
||||
# confirmed dead, but a concurrent IAM refresh may have already
|
||||
# respawned it; the guard makes this recreate a no-op in that
|
||||
# case rather than killing the fresh engine (#29176). Unlike the
|
||||
# direct path there is no SELECT 1 probe here, so the generation
|
||||
# guard is the only thing standing between a crash-reconnect and
|
||||
# a refresh that raced it.
|
||||
await self.db.recreate_prisma_client(
|
||||
db_url, expected_generation=expected_generation
|
||||
)
|
||||
await self._start_engine_watcher()
|
||||
|
||||
await asyncio.wait_for(_do_heavy_reconnect(), timeout=effective_timeout)
|
||||
|
|
@ -4730,13 +4828,36 @@ class PrismaClient:
|
|||
"DATABASE_URL not set; cannot reconnect Prisma client."
|
||||
)
|
||||
raise RuntimeError("DATABASE_URL not set")
|
||||
# Probe the writer BEFORE recreating. A concurrent IAM token
|
||||
# refresh may have just replaced the engine (issue #29176); if
|
||||
# the writer answers SELECT 1 the connection is already healthy
|
||||
# and recreating would needlessly kill that fresh engine. If we
|
||||
# do recreate, the entry-snapshot generation lets the wrapper
|
||||
# detect a refresh that landed since cycle entry and skip the
|
||||
# redundant restart.
|
||||
writer = self.writer_db
|
||||
try:
|
||||
await writer.query_raw("SELECT 1")
|
||||
verbose_proxy_logger.info(
|
||||
"Writer healthy on probe; skipping recreate (engine "
|
||||
"likely already replaced by a token refresh)."
|
||||
)
|
||||
await self._start_engine_watcher()
|
||||
return
|
||||
except Exception as probe_err:
|
||||
verbose_proxy_logger.warning(
|
||||
"Writer probe failed (%s); recreating Prisma client.",
|
||||
probe_err,
|
||||
)
|
||||
# Fresh Prisma client + new engine subprocess. The previous
|
||||
# "lightweight" path called `disconnect()` which blocks the
|
||||
# event loop on `subprocess.Popen.wait()`; since that call
|
||||
# ends up killing the engine anyway, we do it non-blockingly
|
||||
# via `_kill_engine_process` inside `recreate_prisma_client`.
|
||||
self._cleanup_engine_watcher()
|
||||
await self.db.recreate_prisma_client(db_url)
|
||||
await self.db.recreate_prisma_client(
|
||||
db_url, expected_generation=expected_generation
|
||||
)
|
||||
await self._start_engine_watcher()
|
||||
# Smoke-test the writer specifically; query_raw on the routing
|
||||
# wrapper sends to the reader, which would not validate the
|
||||
|
|
@ -4897,6 +5018,11 @@ class PrismaClient:
|
|||
return
|
||||
if self._db_health_watchdog_task is not None:
|
||||
return
|
||||
# Let planned writer-engine restarts (IAM token refresh, guarded
|
||||
# reconnect) re-arm the watcher on the new PID instead of being
|
||||
# mistaken for a crash (issue #29176). Set on the writer wrapper since
|
||||
# the watcher tracks the writer engine.
|
||||
self.writer_db.on_engine_replaced = self._handle_writer_engine_replaced
|
||||
self._db_health_watchdog_task = asyncio.create_task(
|
||||
self._db_health_watchdog_loop()
|
||||
)
|
||||
|
|
@ -6328,15 +6454,37 @@ def create_model_info_response(
|
|||
"created": DEFAULT_MODEL_CREATED_AT_TIME,
|
||||
"owned_by": provider,
|
||||
}
|
||||
|
||||
# Surface context-window limits for OpenAI-compatible discovery clients.
|
||||
# Only emitted when known, so wildcard routes and limitless backends stay clean.
|
||||
# Limits are best-effort enrichment, so a single malformed deployment degrades
|
||||
# to the base response rather than 500-ing the whole listing.
|
||||
if llm_router is not None:
|
||||
try:
|
||||
model_group_info = llm_router.get_model_group_info(model_id)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(
|
||||
"create_model_info_response: get_model_group_info failed for %s: %s",
|
||||
model_id,
|
||||
e,
|
||||
)
|
||||
model_group_info = None
|
||||
if model_group_info is not None:
|
||||
if model_group_info.max_input_tokens is not None:
|
||||
base["max_input_tokens"] = int(model_group_info.max_input_tokens)
|
||||
if model_group_info.max_output_tokens is not None:
|
||||
base["max_output_tokens"] = int(model_group_info.max_output_tokens)
|
||||
|
||||
if not include_metadata:
|
||||
return base
|
||||
|
||||
effective_fallback_type = fallback_type if fallback_type is not None else "general"
|
||||
valid_fallback_types = ("general", "context_window", "content_policy")
|
||||
|
||||
valid_fallback_types = ["general", "context_window", "content_policy"]
|
||||
if effective_fallback_type not in valid_fallback_types:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Invalid fallback_type. Must be one of: {list(valid_fallback_types)}",
|
||||
detail=f"Invalid fallback_type. Must be one of: {valid_fallback_types}",
|
||||
)
|
||||
|
||||
fallbacks = get_all_fallbacks(
|
||||
|
|
|
|||
|
|
@ -42,6 +42,8 @@ class BaseRAGIngestion(ABC):
|
|||
vector stores, so it overrides the embedding step to be a no-op.
|
||||
"""
|
||||
|
||||
supports_existing_file_id: bool = False
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
ingest_options: RAGIngestOptions,
|
||||
|
|
@ -280,6 +282,7 @@ class BaseRAGIngestion(ABC):
|
|||
content_type: Optional[str],
|
||||
chunks: List[str],
|
||||
embeddings: Optional[List[List[float]]],
|
||||
existing_file_id: str | None = None,
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
Store content in vector store.
|
||||
|
|
@ -292,6 +295,7 @@ class BaseRAGIngestion(ABC):
|
|||
content_type: MIME type
|
||||
chunks: Text chunks (if chunking was done locally)
|
||||
embeddings: Embeddings (if embedding was done locally)
|
||||
existing_file_id: Provider file ID supplied by the caller, if any
|
||||
|
||||
Returns:
|
||||
Tuple of (vector_store_id, file_id)
|
||||
|
|
@ -326,6 +330,12 @@ class BaseRAGIngestion(ABC):
|
|||
)
|
||||
|
||||
try:
|
||||
if existing_file_id and not self.supports_existing_file_id:
|
||||
raise ValueError(
|
||||
f"{self.__class__.__name__} does not support ingesting an existing file_id. "
|
||||
"Upload file data or provide file_url instead."
|
||||
)
|
||||
|
||||
# Step 2: OCR (optional)
|
||||
extracted_text = await self.ocr(
|
||||
file_content=file_content,
|
||||
|
|
@ -349,6 +359,7 @@ class BaseRAGIngestion(ABC):
|
|||
content_type=content_type,
|
||||
chunks=chunks,
|
||||
embeddings=embeddings,
|
||||
existing_file_id=existing_file_id,
|
||||
)
|
||||
|
||||
return RAGIngestResponse(
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue