Merge remote-tracking branch 'origin/main' into litellm_lit_8128_off_peak_pricing_schema

This commit is contained in:
kerry 2026-09-19 00:27:52 +00:00
commit e91f17ac3a
415 changed files with 30102 additions and 11327 deletions

View file

@ -1785,6 +1785,12 @@ jobs:
- wait_for_service:
url: http://localhost:4000
timeout: "300"
- run:
name: Seed the routing strategy through /config/update
command: |
curl --noproxy '*' -sSf -X POST http://localhost:4000/config/update \
-H 'Authorization: Bearer sk-1234' -H 'Content-Type: application/json' \
-d '{"router_settings": {"routing_strategy": "usage-based-routing-v2"}}'
- run:
name: Run tests
command: |

View file

@ -125,6 +125,9 @@ start_proxy() {
start_proxy 4000 proxy.log
proxy_pid="$launched_pid"
.venv/bin/python .circleci/scripts/wait_integration_services.py
curl --noproxy '*' -sSf -X POST "$INTEGRATION_PROXY_URL/config/update" \
-H "Authorization: Bearer $LITELLM_MASTER_KEY" -H 'Content-Type: application/json' \
-d '{"router_settings": {"num_retries": 0}}' > "$results/seed-router-settings.json"
if [ "$suite" = management ]; then
export INTEGRATION_PEER_URL=http://127.0.0.1:4001
start_proxy 4001 peer.log

View file

@ -1,50 +0,0 @@
"""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)

View file

@ -1,211 +0,0 @@
"""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()

View file

@ -1,573 +0,0 @@
#!/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())

View file

@ -1,282 +0,0 @@
# 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

File diff suppressed because it is too large Load diff

View file

@ -1,92 +0,0 @@
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[@]}"

View file

@ -1,28 +0,0 @@
name: Create Daily oss-agent-shin Branch
on:
schedule:
- cron: "0 0 * * *" # Runs every day at midnight UTC
workflow_dispatch: # Allow manual trigger
jobs:
create-oss-agent-shin-branch:
if: github.repository == 'BerriAI/litellm'
runs-on: ubuntu-latest
permissions:
contents: write
steps:
- name: Create daily oss-agent-shin branch
env:
GH_TOKEN: ${{ secrets.GITHUB_TOKEN }}
run: |
BRANCH_NAME="litellm_oss_agent_shin_$(date +'%m_%d_%Y')"
echo "Creating branch: $BRANCH_NAME"
if gh api "repos/${{ github.repository }}/git/ref/heads/$BRANCH_NAME" --silent 2>/dev/null; then
echo "Branch $BRANCH_NAME already exists. Skipping creation."
exit 0
fi
MAIN_SHA=$(gh api "repos/${{ github.repository }}/git/ref/heads/main" --jq '.object.sha')
gh api "repos/${{ github.repository }}/git/refs" -f ref="refs/heads/$BRANCH_NAME" -f sha="$MAIN_SHA" --silent
echo "Successfully created branch: $BRANCH_NAME at $MAIN_SHA"

View file

@ -93,12 +93,13 @@ jobs:
responses-api-endpoint: ${{ vars.LITELLM_API_BASE }}/v1/responses
prompt-file: .github/prompts/duplicate-issue-check.md
output-schema-file: .github/prompts/duplicate-issue-check.schema.json
sandbox: read-only
# read-only denies network, and the whole method is searching the tracker with gh
codex-args: '["-c", "sandbox_permissions=[\"network-full-access\"]"]'
sandbox: workspace-write
# The whole method is searching the tracker with gh, and network is only switchable in workspace-write
codex-args: '["-c", "sandbox_workspace_write.network_access=true"]'
model: ${{ vars.DUPLICATE_CHECK_MODEL }}
# Issue authors have no write access and the action refuses them by default; the
# prompt is fixed, the sandbox read-only, and the only token is read-only on a public repo
codex-version: "0.154.0"
# Issue authors have no write access and the action refuses them by default; the prompt is
# fixed, writes stay inside the throwaway checkout, and the only token is read-only on a public repo
allow-users: "*"
- name: Summary

View file

@ -41,4 +41,5 @@ jobs:
"$RUNNER_TEMP/osv-scanner" scan source \
--config osv-scanner.toml \
-L uv.lock \
-L ui/litellm-dashboard/package-lock.json
-L ui/litellm-dashboard/package-lock.json \
-L vscode-extension/package-lock.json

View file

@ -0,0 +1,65 @@
name: VS Code Extension
permissions:
contents: read
on:
pull_request:
branches:
- main
- litellm_internal_staging
- litellm_oss_staging
- "litellm_**"
paths:
- "vscode-extension/**"
- ".github/workflows/test-vscode-extension.yml"
push:
branches:
- main
paths:
- "vscode-extension/**"
- ".github/workflows/test-vscode-extension.yml"
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.sha }}
cancel-in-progress: ${{ github.event_name == 'pull_request' }}
jobs:
vscode-extension:
runs-on: ubuntu-latest
timeout-minutes: 10
defaults:
run:
working-directory: vscode-extension
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
fetch-depth: 1
persist-credentials: false
- name: Set up Node.js
uses: actions/setup-node@a0853c24544627f65ddf259abe73b1d18a591444 # v5.0.0
with:
node-version: "24"
cache: npm
cache-dependency-path: vscode-extension/package-lock.json
- name: Install dependencies
run: npm ci
- name: Typecheck
run: npm run typecheck
- name: Unit tests
run: npm test
- name: Package extension
run: npm run package
- name: Upload VSIX
uses: actions/upload-artifact@4cec3d8aa04e39d1a68397de0c4cd6fb9dce8ec1 # v4.6.1
with:
name: litellm-vscode
path: vscode-extension/*.vsix
if-no-files-found: error

View file

@ -1,172 +0,0 @@
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)"

View file

@ -17,6 +17,7 @@ from fastapi import HTTPException
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.constants import DEFAULT_OPENAI_MODERATIONS_MODEL
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails._content_utils import iter_message_text
@ -24,11 +25,9 @@ from litellm.types.utils import CallTypesLiteral
class _ENTERPRISE_OpenAI_Moderation(CustomLogger):
def __init__(self):
self.model_name = (
litellm.openai_moderations_model_name or "text-moderation-latest"
) # pass the model_name you initialized on litellm.Router()
pass
@property
def model_name(self) -> str:
return litellm.openai_moderations_model_name or DEFAULT_OPENAI_MODERATIONS_MODEL
#### CALL HOOKS - proxy only ####

View file

@ -82,6 +82,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
"/anthropic/",
"/azure/",
"/azure_ai/",
"/azure_speech/",
"/aws/",
"/bedrock/",
"/comprehendmedical",
@ -94,6 +95,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
"/vertex-ai/",
"/assemblyai/",
"/eu.assemblyai/",
"/deepgram/",
"/langfuse/",
"/vllm/",
"/mistral/",

View file

@ -66,7 +66,7 @@
"/v1/video" "/v1/videos" "/video" "/videos" "/v1/search" "/search"
"/v1/containers" "/containers" "/v1/evals" "/v1/memory" "/queue/chat"
"/v1beta" "/interactions"
"/anthropic" "/azure" "/azure_ai" "/aws" "/bedrock" "/comprehendmedical" "/transcribe" "/cohere" "/gemini" "/google"
"/anthropic" "/azure" "/azure_ai" "/azure_speech" "/aws" "/bedrock" "/comprehendmedical" "/transcribe" "/cohere" "/gemini" "/google"
"/vertex_ai" "/vertex-ai" "/assemblyai" "/eu.assemblyai" "/langfuse" "/vllm"
"/mistral" "/groq" "/voyage" "/cursor" "/milvus" "/openai_passthrough"
"/toolset"

View file

@ -0,0 +1,35 @@
-- CreateTable
CREATE TABLE IF NOT EXISTS "LiteLLM_DailyGlobalSpend" (
"id" TEXT NOT NULL,
"date" TEXT NOT NULL,
"model" TEXT,
"model_group" TEXT,
"custom_llm_provider" TEXT,
"mcp_namespaced_tool_name" TEXT,
"endpoint" TEXT,
"prompt_tokens" BIGINT NOT NULL DEFAULT 0,
"completion_tokens" BIGINT NOT NULL DEFAULT 0,
"cache_read_input_tokens" BIGINT NOT NULL DEFAULT 0,
"cache_creation_input_tokens" BIGINT NOT NULL DEFAULT 0,
"compression_saved_tokens" BIGINT NOT NULL DEFAULT 0,
"compression_savings_spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0,
"prompt_caching_savings_spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0,
"gateway_injected_caching_savings_spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0,
"autorouter_savings_spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0,
"spend" DOUBLE PRECISION NOT NULL DEFAULT 0.0,
"api_requests" BIGINT NOT NULL DEFAULT 0,
"successful_requests" BIGINT NOT NULL DEFAULT 0,
"failed_requests" BIGINT NOT NULL DEFAULT 0,
"total_response_time_ms" BIGINT NOT NULL DEFAULT 0,
"timed_requests" BIGINT NOT NULL DEFAULT 0,
"created_at" TIMESTAMP(3) NOT NULL DEFAULT CURRENT_TIMESTAMP,
"updated_at" TIMESTAMP(3) NOT NULL,
CONSTRAINT "LiteLLM_DailyGlobalSpend_pkey" PRIMARY KEY ("id")
);
-- CreateIndex
CREATE INDEX IF NOT EXISTS "LiteLLM_DailyGlobalSpend_date_idx" ON "LiteLLM_DailyGlobalSpend"("date");
-- CreateIndex
CREATE UNIQUE INDEX IF NOT EXISTS "LiteLLM_DailyGlobalSpend_date_model_model_group_custom_llm__key" ON "LiteLLM_DailyGlobalSpend"("date", "model", "model_group", "custom_llm_provider", "mcp_namespaced_tool_name", "endpoint");

View file

@ -820,6 +820,37 @@ model LiteLLM_DailyUserSpend {
@@index([endpoint])
}
// Key-free daily rollup of LiteLLM_DailyUserSpend, read by the global usage view
model LiteLLM_DailyGlobalSpend {
id String @id @default(uuid())
date String
model String?
model_group String?
custom_llm_provider String?
mcp_namespaced_tool_name String?
endpoint String?
prompt_tokens BigInt @default(0)
completion_tokens BigInt @default(0)
cache_read_input_tokens BigInt @default(0)
cache_creation_input_tokens BigInt @default(0)
compression_saved_tokens BigInt @default(0)
compression_savings_spend Float @default(0.0)
prompt_caching_savings_spend Float @default(0.0)
gateway_injected_caching_savings_spend Float @default(0.0)
autorouter_savings_spend Float @default(0.0)
spend Float @default(0.0)
api_requests BigInt @default(0)
successful_requests BigInt @default(0)
failed_requests BigInt @default(0)
total_response_time_ms BigInt @default(0)
timed_requests BigInt @default(0)
created_at DateTime @default(now())
updated_at DateTime @updatedAt
@@unique([date, model, model_group, custom_llm_provider, mcp_namespaced_tool_name, endpoint])
@@index([date])
}
// Track daily organization spend metrics per model and key
model LiteLLM_DailyOrganizationSpend {
id String @id @default(uuid())

119
litellm-rust/Cargo.lock generated
View file

@ -559,6 +559,21 @@ dependencies = [
"vsimd",
]
[[package]]
name = "bit-set"
version = "0.8.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "08807e080ed7f9d5433fa9b275196cfc35414f66a0c79d864dc51a0d825231a3"
dependencies = [
"bit-vec",
]
[[package]]
name = "bit-vec"
version = "0.8.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5e764a1d40d510daf35e07be9eb06e75770908c27d411ee6c92109c9840eaaf7"
[[package]]
name = "bitflags"
version = "2.13.1"
@ -1166,6 +1181,17 @@ dependencies = [
"pin-project-lite",
]
[[package]]
name = "fancy-regex"
version = "0.19.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d301f5bf187b3c295fce6468d3875037a0bccc5f6b151c63cac2f85babf21912"
dependencies = [
"bit-set",
"regex-automata",
"regex-syntax",
]
[[package]]
name = "fastrand"
version = "2.5.0"
@ -2001,24 +2027,18 @@ dependencies = [
"tokio",
]
[[package]]
name = "litellm-callbacks"
version = "0.1.0"
dependencies = [
"rstest",
"serde_json",
"tokio",
]
[[package]]
name = "litellm-callbacks-legacy"
version = "0.1.0"
dependencies = [
"litellm-callbacks",
"litellm-auth",
"litellm-host",
"litellm-host-python",
"proptest",
"pyo3",
"rstest",
"serde_json",
"strum",
]
[[package]]
@ -2030,8 +2050,8 @@ dependencies = [
"futures-util",
"litellm-auth",
"litellm-auth-aws",
"litellm-callbacks",
"litellm-core-utils",
"litellm-host",
"litellm-llms",
"litellm-types",
"mime_guess",
@ -2059,7 +2079,9 @@ dependencies = [
name = "litellm-core-utils"
version = "0.1.0"
dependencies = [
"fancy-regex",
"litellm-types",
"rstest",
"serde",
"serde_json",
"serde_path_to_error",
@ -2082,12 +2104,22 @@ dependencies = [
"tokio",
]
[[package]]
name = "litellm-host"
version = "0.1.0"
dependencies = [
"litellm-auth",
"rstest",
"serde_json",
"tokio",
]
[[package]]
name = "litellm-host-python"
version = "0.1.0"
dependencies = [
"futures-util",
"litellm-callbacks",
"litellm-host",
"pyo3",
"pyo3-async-runtimes",
"pythonize",
@ -2111,9 +2143,9 @@ dependencies = [
"litellm-auth-aws",
"litellm-auth-azure",
"litellm-auth-gcp",
"litellm-callbacks",
"litellm-core-utils",
"litellm-framing",
"litellm-host",
"litellm-types",
"reqwest 0.12.28",
"rstest",
@ -2570,6 +2602,25 @@ dependencies = [
"unicode-ident",
]
[[package]]
name = "proptest"
version = "1.11.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "4b45fcc2344c680f5025fe57779faef368840d0bd1f42f216291f0dc4ace4744"
dependencies = [
"bit-set",
"bit-vec",
"bitflags",
"num-traits",
"rand 0.9.5",
"rand_chacha 0.9.0",
"rand_xorshift",
"regex-syntax",
"rusty-fork",
"tempfile",
"unarray",
]
[[package]]
name = "pyo3"
version = "0.29.2"
@ -2651,6 +2702,12 @@ dependencies = [
"serde",
]
[[package]]
name = "quick-error"
version = "1.2.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "a1d01941d82fa2ab50be1e79e6714289dd7cde78eba4c074bc5a4374f650dfe0"
[[package]]
name = "quinn"
version = "0.11.11"
@ -2814,6 +2871,15 @@ dependencies = [
"rand_core 0.10.1",
]
[[package]]
name = "rand_xorshift"
version = "0.4.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "513962919efc330f829edb2535844d1b912b0fbe2ca165d613e4e8788bb05a5a"
dependencies = [
"rand_core 0.9.5",
]
[[package]]
name = "rayon"
version = "1.12.0"
@ -3213,6 +3279,18 @@ version = "1.0.23"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "cf54715a573b99ac80df0bc206da022bcd442c974952c7b9720069370852e21f"
[[package]]
name = "rusty-fork"
version = "0.3.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "cc6bf79ff24e648f6da1f8d1f011e9cac26491b619e6b9280f2b47f1774e6ee2"
dependencies = [
"fnv",
"quick-error",
"tempfile",
"wait-timeout",
]
[[package]]
name = "ryu"
version = "1.0.23"
@ -4068,6 +4146,12 @@ dependencies = [
"syn 2.0.119",
]
[[package]]
name = "unarray"
version = "0.1.4"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "eaea85b334db583fe3274d12b4cd1880032beab409c0d774be044d4480ab9a94"
[[package]]
name = "unicase"
version = "2.9.0"
@ -4181,6 +4265,15 @@ version = "0.8.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "5c3082ca00d5a5ef149bb8b555a72ae84c9c59f7250f013ac822ac2e49b19c64"
[[package]]
name = "wait-timeout"
version = "0.2.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "09ac3b126d3914f9849036f826e054cbabdc8519970b8998ddaf3b5bd3c65f11"
dependencies = [
"libc",
]
[[package]]
name = "walkdir"
version = "2.5.0"

View file

@ -10,7 +10,7 @@ repository = "https://github.com/BerriAI/litellm"
[workspace.dependencies]
litellm-core = { path = "crates/core" }
litellm-callbacks = { path = "crates/callbacks" }
litellm-host = { path = "crates/host" }
litellm-callbacks-legacy = { path = "crates/callbacks-legacy" }
litellm-framing = { path = "crates/framer" }
litellm-auth = { path = "crates/auth" }
@ -26,6 +26,7 @@ litellm-token-counter = { path = "crates/token-counter" }
litellm-host-python = { path = "crates/host-python" }
bytes = "1"
proptest = "1.7.0"
pyo3 = "0.29.2"
pyo3-async-runtimes = { version = "0.29.0", features = ["tokio-runtime"] }
pythonize = "0.29.0"
@ -50,6 +51,7 @@ strum = { version = "0.28.0", features = ["derive"] }
url = "2.5.8"
time = { version = "0.3.53", features = ["parsing"] }
criterion = "0.8.2"
fancy-regex = "0.19.2"
veil = "0.3.0"
[profile.release]

View file

@ -1,6 +1,8 @@
use serde::Deserialize;
use veil::Redact;
#[derive(Redact, Clone)]
#[derive(Redact, Clone, Deserialize)]
#[serde(transparent)]
pub struct SecretValue(#[redact(with = "[REDACTED]")] String);
impl SecretValue {

View file

@ -1,15 +1,17 @@
- Target invariants, not completion claims
- Keep this crate the legacy `@client` wrapper as the native call sees it, and nothing else: the `Logging` contract (`function_setup`, the deployment hooks, `pre_call`/`post_call`, the sync and async success and failure fan-out, the deferred proxy release, the argument sharing those callbacks rely on) plus the kwargs rewrites the wrapper makes on the way in (credential-name inheritance, the budget and retry-count limits)
- The driver in `litellm-host-python`, the routes and core see one `CallbackAdapter`; they never learn which Python objects consume a call
- The driver in `litellm-host-python`, the routes and core see one `PythonLifecycle`; they never learn which Python objects consume a call
- Rust drives the call; every litellm Python internal it still borrows is a variant of `LegacyPython`, grouped by subsystem (`Wrapper`, `Logging`, `DeploymentHooks`)
- The enum only shrinks: when Rust owns a subsystem, delete its group rather than adding a Rust path beside it
- Calling a user's own callback directly is permanent Python surface and gets its own type outside `LegacyPython`
- `PublicCall` is the caller's call as `Logging` sees it: the positional arguments, the keyword view as the legacy path rewrites it (setup, deployment hook, prepare) and the bound request object whose attributes back keywords the caller omitted; routes hand it over through `run_legacy_call` and keep no copy
- `setup` decides once who owns the `Logging` instance and returns it as `CallSetup.bridge_owned`; `PythonLogger` carries it and nothing on the instance records it
- A logger the caller passed as `litellm_logging_obj` is caller-owned and observed in full, because the caller reads it after the call; the proxy is the live case
- A logger `function_setup` built for this call is bridge-owned, so each fan-out phase is skipped when `callbacks_needed` finds no registry, dynamic callback, `logger_fn` or debug switch for it; cost, timing and response metadata still run
- `setup` reuses a `Logging` the caller passed as `litellm_logging_obj` (the proxy and Router are the live cases) and otherwise builds one through `function_setup`, as `@client` does
- Either way every phase calls the same `Logging` method the Python path calls; which callbacks run is `Logging`'s decision, never this crate's
- Callbacks receive the caller's own objects and may mutate them; this crate alone carries that obligation
- Retain complete boundary arguments, opaque unknown values, aliases, omitted/default distinctions and deliberate copies; preserve the established deployment-hook kwargs view
- Re-alias every `passthrough_fields` body key to the caller's object before `pre_call`; a keyword wins over the request attribute even when it is an explicit `None`
- Before `pre_call`, re-alias every body key whose value equals the caller's argument to the caller's own object; this crate compares the two itself, and the argument is resolved by `litellm_host_python::lookup`
- Retain independently captured body/header roots from `pre_call` to `post_call`; in-place mutation reaches the wire, envelope field replacement is visible to later callbacks only
- A later kind of callback host (WASM, in-process Rust) has none of these obligations, so they stay out of `litellm-callbacks`, `litellm-host-python` and the bridge; the only facts that cross from the route are the prepared keyword view and `RequestContext.passthrough_fields`
- A later kind of callback host (WASM, in-process Rust) has none of these obligations, so they stay out of `litellm-host`, `litellm-host-python` and the bridge; the only fact that crosses from the route is the prepared keyword view
- Success and failure handlers receive the exact selected public response or exception; logging projections, redaction and snapshots keep their own copy contracts
- Ordinary failure-handler errors cannot suppress the other eligible family or replace the mapped provider error; a cancellation ends the call with no further dispatch
- Dispatch errors never replay provider work or trigger the opposite outcome; the proxy's acceptance or rejection releases deferred success at most once

View file

@ -7,10 +7,15 @@ repository.workspace = true
autotests = false
[dependencies]
litellm-callbacks.workspace = true
litellm-host.workspace = true
litellm-host-python.workspace = true
pyo3.workspace = true
strum.workspace = true
serde_json.workspace = true
[dev-dependencies]
litellm-auth.workspace = true
proptest.workspace = true
rstest.workspace = true
serde_json.workspace = true

View file

@ -0,0 +1,118 @@
{
"setup": [
"call_type",
"args",
"kwargs",
"start_time",
"asynchronous"
],
"check_limits": [
"kwargs"
],
"finalize": [
"response",
"logger",
"kwargs",
"start_time",
"end_time"
],
"update_logging": [
"logger",
"kwargs",
"model",
"optional_params",
"litellm_params",
"custom_llm_provider"
],
"pre_call": [
"logger",
"input",
"api_key",
"additional_args"
],
"post_call": [
"logger",
"original_response",
"api_key",
"additional_args"
],
"defers_async_logging": [
"logger"
],
"defer_success": [
"logger",
"pending"
],
"sync_success_for_async_call": [
"logger",
"response",
"start",
"end"
],
"failure_handler": [
"logger",
"error",
"start",
"end",
"asynchronous"
],
"submit_success": [
"logger",
"response",
"start",
"end"
],
"async_success_handler": [
"logger",
"response",
"start",
"end"
],
"enqueue_logging": [
"coroutine"
],
"restore_context": [
"logger"
],
"custom_pricing_fields": [],
"is_internal_call": [],
"credential_list": [],
"warn_unknown_credential": [
"name",
"loaded"
],
"before_deployment_call": [
"kwargs",
"call_type"
],
"after_deployment_success": [
"kwargs",
"response",
"call_type"
],
"after_deployment_failure": [
"kwargs",
"error",
"call_type"
],
"stream_opened": [
"logger"
],
"stream_success": [
"logger",
"url_route",
"endpoint_type",
"request_body",
"chunks",
"start",
"end",
"first_chunk"
],
"stream_failure": [
"logger",
"endpoint_type",
"request_body",
"chunks",
"error"
]
}

View file

@ -2,21 +2,26 @@
//! raises is answered with the same `Logging` calls, in the same order, as the Python
//! `@client` path makes them.
use litellm_callbacks::event::{CallEvent, FailureOrigin, RequestContext, Timing, WireRequest};
use litellm_host::event::{
FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest, epoch_seconds,
};
use litellm_host_python::{
AdapterStep, CallbackAdapter, PublicValue, from_py, missing_state, to_py,
LifecycleEvent, LifecycleStep, PythonLifecycle, from_py, missing_state, to_py,
};
use pyo3::{
exceptions::{PyBaseException, PyException},
gc::{PyTraverseError, PyVisit},
prelude::*,
types::PyDict,
types::{PyDict, PyList},
};
use serde_json::Value;
use crate::{
DeploymentHooks, LegacyCallbacks, PublicCall, PythonLogger,
deferred::{PendingLogging, PendingSuccess},
finalize, is_internal_call, prepare, setup,
finalize, is_internal_call,
legacy_python::Streaming,
prepare, setup,
};
/// What the legacy contract needs to know about the route it is logging.
@ -25,6 +30,22 @@ pub struct LegacySurface {
pub call_type: &'static str,
/// What `Logging.pre_call` is told the input was.
pub input_description: &'static str,
/// How a streamed response is billed; `None` for a route that never streams.
pub stream: Option<PassThroughStream>,
}
/// The pass-through billing a streamed response goes through once its chunks are in.
#[derive(Clone, Copy, Debug)]
pub struct PassThroughStream {
pub url_route: &'static str,
/// A value of Python's `EndpointType`.
pub endpoint_type: &'static str,
}
/// What the Messages stream iterator keeps for its end-of-stream billing.
struct DeliveredStream {
chunks: Py<PyList>,
first_chunk: Option<Py<PyAny>>,
}
enum Pending {
@ -44,6 +65,8 @@ pub struct LegacyLogging {
error: Option<Py<PyBaseException>>,
body: Option<Py<PyDict>>,
headers: Option<Py<PyDict>>,
context: Option<RequestContext>,
stream: Option<DeliveredStream>,
asynchronous: bool,
internal: bool,
pending: Option<Pending>,
@ -77,6 +100,8 @@ impl LegacyLogging {
error: None,
body: None,
headers: None,
context: None,
stream: None,
asynchronous,
internal: false,
pending: None,
@ -85,8 +110,8 @@ impl LegacyLogging {
/// Deployment hooks are awaited, and Python's synchronous `@client` wrapper never
/// runs them.
fn deployment_hooks(&self, py: Python<'_>) -> PyResult<bool> {
Ok(self.asynchronous && DeploymentHooks::needed(py)?)
fn runs_deployment_hooks(&self) -> bool {
self.asynchronous
}
fn logger(&self) -> PyResult<&PythonLogger> {
@ -95,13 +120,13 @@ impl LegacyLogging {
})
}
fn prepare(&mut self, py: Python<'_>) -> PyResult<AdapterStep> {
fn prepare(&mut self, py: Python<'_>) -> PyResult<LifecycleStep> {
let prepared = prepare(py, self.call.kwargs().bind(py), self.logger()?)?.unbind();
self.call.set_kwargs(prepared);
Ok(AdapterStep::Arguments(self.call.kwargs().clone_ref(py)))
Ok(LifecycleStep::Arguments(self.call.kwargs().clone_ref(py)))
}
fn finalize(&mut self, py: Python<'_>) -> PyResult<AdapterStep> {
fn finalize(&mut self, py: Python<'_>) -> PyResult<LifecycleStep> {
finalize(
py,
&self.response,
@ -112,7 +137,7 @@ impl LegacyLogging {
)?;
self.response
.as_ref()
.map(|response| AdapterStep::Response(response.clone_ref(py)))
.map(|response| LifecycleStep::Response(response.clone_ref(py)))
.ok_or_else(missing_state)
}
@ -145,9 +170,7 @@ impl LegacyLogging {
.get_item("fallbacks")?
.is_none_or(|value| value.is_none())
{
if !logger.callbacks_needed(py, "async_success")? {
logger.success_bookkeeping(py, &self.response, &self.start, &self.end, true)?;
} else if logger.defers_async_logging(py) {
if logger.defers_async_logging(py) {
let pending = Py::new(
py,
PendingLogging {
@ -162,15 +185,72 @@ impl LegacyLogging {
logger.sync_success_for_async_call(py, &self.response, &self.start, &self.end)
}
fn stream_success(&self, py: Python<'_>, stream: &DeliveredStream) -> PyResult<()> {
let logger = self.logger()?;
let billing = self.surface.stream.ok_or_else(missing_state)?;
let billed = Streaming::Success.call(
py,
(
logger.object(py),
billing.url_route,
billing.endpoint_type,
&self.body,
&stream.chunks,
&self.start,
&self.end,
&stream.first_chunk,
),
);
match billed {
Err(error) if error.is_instance_of::<PyException>(py) => {
error.write_unraisable(py, Some(logger.object(py)));
Ok(())
}
result => result.map(|_| ()),
}
}
/// A failure after the stream reached the caller bills the delivered chunks as
/// partial usage. The sync path has no loop to schedule that on, so it falls back to
/// the plain failure handler.
fn stream_failure(&mut self, py: Python<'_>) -> PyResult<LifecycleStep> {
let (Some(logger), Some(error), Some(stream), Some(billing)) =
(&self.logger, &self.error, &self.stream, self.surface.stream)
else {
return Ok(LifecycleStep::Done);
};
if !self.asynchronous {
return self.dispatch_failure(py);
}
let scheduled = Streaming::Failure.call(
py,
(
logger.object(py),
billing.endpoint_type,
&self.body,
&stream.chunks,
error,
),
);
match scheduled {
Ok(awaitable) => {
self.pending = Some(Pending::AsyncFailure);
Ok(LifecycleStep::Await(awaitable.unbind()))
}
Err(failure) if is_cancellation(py, &failure) => Err(failure),
Err(_) => Ok(LifecycleStep::Done),
}
}
/// The sync failure handler, then the async one for async calls. Ordinary handler
/// errors never replace the selected failure or suppress the other family; a
/// cancellation does end the call.
fn dispatch_failure(&mut self, py: Python<'_>) -> PyResult<AdapterStep> {
fn dispatch_failure(&mut self, py: Python<'_>) -> PyResult<LifecycleStep> {
let (Some(logger), Some(error)) = (&self.logger, &self.error) else {
return Ok(AdapterStep::Done);
return Ok(LifecycleStep::Done);
};
if self.asynchronous && self.internal {
return Ok(AdapterStep::Done);
return Ok(LifecycleStep::Done);
}
if let Err(failure) = logger.failure(py, error, &self.start, &self.end, false)
&& is_cancellation(py, &failure)
@ -178,27 +258,27 @@ impl LegacyLogging {
return Err(failure);
}
if !self.asynchronous {
return Ok(AdapterStep::Done);
return Ok(LifecycleStep::Done);
}
match logger.failure(py, error, &self.start, &self.end, true) {
Ok(Some(awaitable)) => {
self.pending = Some(Pending::AsyncFailure);
Ok(AdapterStep::Await(awaitable))
Ok(LifecycleStep::Await(awaitable))
}
Ok(None) => Ok(AdapterStep::Done),
Ok(None) => Ok(LifecycleStep::Done),
Err(failure) if is_cancellation(py, &failure) => Err(failure),
Err(_) => Ok(AdapterStep::Done),
Err(_) => Ok(LifecycleStep::Done),
}
}
}
impl CallbackAdapter for LegacyLogging {
impl PythonLifecycle for LegacyLogging {
fn begin(
&mut self,
py: Python<'_>,
arguments: Py<PyDict>,
started_at: f64,
) -> PyResult<AdapterStep> {
) -> PyResult<LifecycleStep> {
self.call.set_kwargs(arguments);
self.start = datetime(py, started_at)?;
self.internal = is_internal_call(py)?;
@ -212,9 +292,9 @@ impl CallbackAdapter for LegacyLogging {
)?;
self.logger = Some(result.logger()?);
self.call.set_kwargs(result.kwargs()?);
if self.deployment_hooks(py)? {
if self.runs_deployment_hooks() {
self.pending = Some(Pending::DeploymentPreCall);
return Ok(AdapterStep::Await(DeploymentHooks::before_call(
return Ok(LifecycleStep::Await(DeploymentHooks::before_call(
py,
self.call.kwargs(),
self.surface.call_type,
@ -228,18 +308,16 @@ impl CallbackAdapter for LegacyLogging {
py: Python<'_>,
wire: Box<WireRequest>,
context: &RequestContext,
) -> PyResult<AdapterStep> {
) -> PyResult<LifecycleStep> {
let logger = self.logger()?;
logger.update_from_kwargs(py, self.call.kwargs(), &wire, context)?;
if !logger.callbacks_needed(py, "payload")? {
logger.record_api_call_start(py)?;
return Ok(AdapterStep::Wire(wire));
}
let body = to_py(py, &wire.body)?
.into_bound(py)
.cast_into::<PyDict>()?;
for name in context.passthrough_fields.iter() {
if let Some(value) = self.call.lookup(py, name)? {
for (name, sent) in wire.body.as_object().into_iter().flatten() {
if let Some(value) = self.call.lookup(py, name)?
&& from_py::<Value>(&value).is_ok_and(|caller| caller == *sent)
{
body.set_item(name, value)?;
}
}
@ -249,11 +327,11 @@ impl CallbackAdapter for LegacyLogging {
}
self.body = Some(body.clone().unbind());
self.headers = Some(headers.clone().unbind());
let api_key = self.call.lookup(py, "api_key")?;
self.context = Some(context.clone());
self.logger()?.pre_call(
py,
self.surface.input_description,
api_key.as_ref(),
context.api_key.as_ref().map(|api_key| api_key.expose()),
&body,
&headers,
&wire.url,
@ -262,7 +340,7 @@ impl CallbackAdapter for LegacyLogging {
.iter()
.map(|(name, value)| Ok((name.extract::<String>()?, value.extract::<String>()?)))
.collect::<PyResult<Vec<_>>>()?;
Ok(AdapterStep::Wire(Box::new(WireRequest {
Ok(LifecycleStep::Wire(Box::new(WireRequest {
body: from_py(&body)?,
headers,
..*wire
@ -274,12 +352,12 @@ impl CallbackAdapter for LegacyLogging {
py: Python<'_>,
response: Py<PyAny>,
timing: Timing,
) -> PyResult<AdapterStep> {
) -> PyResult<LifecycleStep> {
self.end = Some(datetime(py, timing.end_time)?);
self.response = Some(response);
if self.deployment_hooks(py)? {
if self.runs_deployment_hooks() {
self.pending = Some(Pending::DeploymentPostCall);
return Ok(AdapterStep::Await(DeploymentHooks::after_success(
return Ok(LifecycleStep::Await(DeploymentHooks::after_success(
py,
self.call.kwargs(),
&self.response,
@ -289,36 +367,50 @@ impl CallbackAdapter for LegacyLogging {
self.finalize(py)
}
fn emit(
&mut self,
py: Python<'_>,
event: &CallEvent,
public: Option<PublicValue<'_>>,
) -> PyResult<AdapterStep> {
match (event, public) {
(CallEvent::ResponseReceived { raw }, _) => {
let logger = self.logger()?;
if logger.callbacks_needed(py, "payload")? {
logger.post_call(py, &raw.body, self.body.as_ref(), self.headers.as_ref())?;
}
Ok(AdapterStep::Done)
fn emit(&mut self, py: Python<'_>, event: LifecycleEvent<'_>) -> PyResult<LifecycleStep> {
match event {
LifecycleEvent::Started { .. } => Ok(LifecycleStep::Done),
LifecycleEvent::Machine(MachineEvent::ResponseReceived { raw }) => {
let api_key = self
.context
.as_ref()
.and_then(|context| context.api_key.as_ref())
.map(|api_key| api_key.expose());
self.logger()?.post_call(
py,
&raw.body,
api_key,
self.body.as_ref(),
self.headers.as_ref(),
)?;
Ok(LifecycleStep::Done)
}
(CallEvent::Succeeded { timing }, Some(PublicValue::Response(response))) => {
LifecycleEvent::Succeeded { timing, response } => {
self.end = Some(datetime(py, timing.end_time)?);
self.response = Some(response.clone_ref(py));
self.dispatch_success(py)?;
Ok(AdapterStep::Done)
match &self.stream {
Some(stream) => self.stream_success(py, stream)?,
None => self.dispatch_success(py)?,
}
Ok(LifecycleStep::Done)
}
(CallEvent::Failed { timing, origin }, Some(PublicValue::Error(error))) => {
LifecycleEvent::Failed {
timing,
origin,
error,
} => {
self.end = Some(datetime(py, timing.end_time)?);
self.error = Some(error.clone_ref(py).into_value(py));
if *origin == FailureOrigin::Call
if self.stream.is_some() {
return self.stream_failure(py);
}
if origin == FailureOrigin::Call
&& self.logger.is_some()
&& self.deployment_hooks(py)?
&& self.runs_deployment_hooks()
{
let error = self.error.as_ref().ok_or_else(missing_state)?;
self.pending = Some(Pending::DeploymentFailure);
return Ok(AdapterStep::Await(DeploymentHooks::after_failure(
return Ok(LifecycleStep::Await(DeploymentHooks::after_failure(
py,
self.call.kwargs(),
error,
@ -327,11 +419,30 @@ impl CallbackAdapter for LegacyLogging {
}
self.dispatch_failure(py)
}
_ => Err(missing_state()),
}
}
fn resume(&mut self, py: Python<'_>, result: PyResult<Py<PyAny>>) -> PyResult<AdapterStep> {
fn opened(&mut self, py: Python<'_>) -> PyResult<()> {
if self.surface.stream.is_none() {
return Err(missing_state());
}
Streaming::Opened.call(py, (self.logger()?.object(py),))?;
self.stream = Some(DeliveredStream {
chunks: PyList::empty(py).unbind(),
first_chunk: None,
});
Ok(())
}
fn delivered(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()> {
let stream = self.stream.as_mut().ok_or_else(missing_state)?;
if stream.first_chunk.is_none() {
stream.first_chunk = Some(datetime(py, epoch_seconds())?);
}
stream.chunks.bind(py).append(chunk)
}
fn resume(&mut self, py: Python<'_>, result: PyResult<Py<PyAny>>) -> PyResult<LifecycleStep> {
match self.pending.take().ok_or_else(missing_state)? {
Pending::DeploymentPreCall => {
self.call
@ -345,7 +456,7 @@ impl CallbackAdapter for LegacyLogging {
Pending::DeploymentFailure => self.dispatch_failure(py),
Pending::AsyncFailure => match result {
Err(failure) if is_cancellation(py, &failure) => Err(failure),
_ => Ok(AdapterStep::Done),
_ => Ok(LifecycleStep::Done),
},
}
}
@ -357,7 +468,8 @@ impl CallbackAdapter for LegacyLogging {
error.write_unraisable(py, None);
}
self.body = None;
self.headers = None;
self.context = None;
self.stream = None;
}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
@ -369,8 +481,11 @@ impl CallbackAdapter for LegacyLogging {
visit.call(&self.end)?;
visit.call(&self.response)?;
visit.call(&self.error)?;
visit.call(&self.body)?;
visit.call(&self.headers)
if let Some(stream) = &self.stream {
visit.call(&stream.chunks)?;
visit.call(&stream.first_chunk)?;
}
visit.call(&self.body)
}
}

View file

@ -3,8 +3,8 @@
//! lifetime. No other callback host has that obligation, which is why nothing outside
//! this crate holds them.
use litellm_callbacks::{machine::Machine, route::Route};
use litellm_host_python::{RouteHost, run_call};
use litellm_host::{machine::Machine, route::Route};
use litellm_host_python::{RouteHost, lookup, run_call};
use pyo3::{
gc::{PyTraverseError, PyVisit},
prelude::*,
@ -63,21 +63,6 @@ impl PublicCall {
}
}
/// The caller's own object for a public argument, as every legacy reader resolves it: the
/// keyword if given, even an explicit `None`, else the bound request's attribute. A route
/// host projecting from the prepared keyword view uses the same rule, so the callbacks
/// and the provider see one object per argument.
pub fn lookup<'py>(
kwargs: &Bound<'py, PyDict>,
request: &Bound<'py, PyAny>,
name: &str,
) -> PyResult<Option<Bound<'py, PyAny>>> {
if let Some(value) = kwargs.get_item(name)? {
return Ok(Some(value));
}
request.getattr_opt(name)
}
/// Runs one native call under the legacy `Logging` contract: the route host projects from
/// the keyword view the contract prepares, and the contract observes the call.
pub fn run_legacy_call<H, M>(
@ -121,32 +106,6 @@ mod tests {
(call, locals)
}
#[test]
fn lookup_prefers_the_keyword_even_when_none_and_falls_back_to_the_request() {
Python::initialize();
Python::attach(|py| {
let (call, locals) = capture(
py,
c"
key = object()
document = {'type': 'document_url'}
class Request:
api_key = 'from-request'
api_base = 'from-request'
document = document
request = Request()
kwargs = {'api_key': key, 'api_base': None}
",
);
let key = locals.get_item("key").unwrap().unwrap();
let document = locals.get_item("document").unwrap().unwrap();
assert!(call.lookup(py, "api_key").unwrap().unwrap().is(&key));
assert!(call.lookup(py, "api_base").unwrap().unwrap().is_none());
assert!(call.lookup(py, "document").unwrap().unwrap().is(&document));
assert!(call.lookup(py, "model").unwrap().is_none());
});
}
#[test]
fn capture_copies_the_keyword_dict_without_copying_its_values() {
Python::initialize();

View file

@ -2,15 +2,14 @@
//! the deferred and worker-submitted success paths, and the sync-callbacks-for-async-calls
//! duplication. All of it expires with the legacy callback contract.
use litellm_callbacks::event::{RequestContext, WireRequest};
use litellm_host::event::{RequestContext, WireRequest};
use litellm_host_python::to_py;
use pyo3::{exceptions::PyBaseException, prelude::*, types::PyDict};
use crate::legacy_python::{Logging, Wrapper};
use crate::logger::PythonLogger;
pub trait LegacyCallbacks {
fn callbacks_needed(&self, py: Python<'_>, phase: &str) -> PyResult<bool>;
/// `Logging.update_from_kwargs`: what the logger is told about the request it is
/// about to see, with consumed credentials redacted.
fn update_from_kwargs(
@ -21,24 +20,23 @@ pub trait LegacyCallbacks {
context: &RequestContext,
) -> PyResult<()>;
fn record_api_call_start(&self, py: Python<'_>) -> PyResult<()>;
/// `Logging.pre_call`, or its payload-free shortcut when no input callback listens.
/// `Logging.pre_call`.
fn pre_call(
&self,
py: Python<'_>,
input: &str,
api_key: Option<&Bound<'_, PyAny>>,
api_key: Option<&str>,
body: &Bound<'_, PyDict>,
headers: &Bound<'_, PyDict>,
url: &str,
) -> PyResult<()>;
/// `Logging.post_call`, or its payload-free shortcut when no input callback listens.
/// `Logging.post_call`.
fn post_call(
&self,
py: Python<'_>,
original_response: &str,
api_key: Option<&str>,
body: Option<&Py<PyDict>>,
headers: Option<&Py<PyDict>>,
) -> PyResult<()>;
@ -82,16 +80,6 @@ pub trait LegacyCallbacks {
}
impl LegacyCallbacks for PythonLogger {
fn callbacks_needed(&self, py: Python<'_>, phase: &str) -> PyResult<bool> {
if !self.bridge_owned() {
return Ok(true);
}
py.import("litellm.rust_bridge.legacy_callbacks")?
.getattr("callbacks_needed")?
.call1((self.object(py), phase))?
.extract()
}
fn update_from_kwargs(
&self,
py: Python<'_>,
@ -100,18 +88,13 @@ impl LegacyCallbacks for PythonLogger {
context: &RequestContext,
) -> PyResult<()> {
let secret_fields: Vec<&str> = context.secret_fields.iter().map(String::as_str).collect();
let update = PyDict::new(py);
update.set_item("kwargs", redact(py, kwargs.bind(py), &secret_fields)?)?;
update.set_item("model", &context.model)?;
update.set_item(
"optional_params",
redact(
py,
&to_py(py, &context.optional_params)?
.into_bound(py)
.cast_into::<PyDict>()?,
&secret_fields,
)?,
let redacted_kwargs = redact(py, kwargs.bind(py), &secret_fields)?;
let optional_params = redact(
py,
&to_py(py, &context.optional_params)?
.into_bound(py)
.cast_into::<PyDict>()?,
&secret_fields,
)?;
let params = PyDict::new(py);
params.set_item(
@ -131,15 +114,17 @@ impl LegacyCallbacks for PythonLogger {
params.set_item(name, value)?;
}
}
update.set_item("litellm_params", params)?;
update.set_item("custom_llm_provider", &context.custom_llm_provider)?;
self.object(py)
.call_method("update_from_kwargs", (), Some(&update))?;
Ok(())
}
fn record_api_call_start(&self, py: Python<'_>) -> PyResult<()> {
self.object(py).call_method0("record_api_call_start_time")?;
Logging::Update.call(
py,
(
self.object(py),
redacted_kwargs,
&context.model,
optional_params,
params,
&context.custom_llm_provider,
),
)?;
Ok(())
}
@ -147,7 +132,7 @@ impl LegacyCallbacks for PythonLogger {
&self,
py: Python<'_>,
input: &str,
api_key: Option<&Bound<'_, PyAny>>,
api_key: Option<&str>,
body: &Bound<'_, PyDict>,
headers: &Bound<'_, PyDict>,
url: &str,
@ -156,17 +141,7 @@ impl LegacyCallbacks for PythonLogger {
additional.set_item("complete_input_dict", body)?;
additional.set_item("headers", headers)?;
additional.set_item("api_base", url)?;
let kwargs = PyDict::new(py);
kwargs.set_item("input", input)?;
kwargs.set_item("api_key", api_key)?;
kwargs.set_item("additional_args", &additional)?;
if self.callbacks_needed(py, "input")? {
self.object(py).call_method("pre_call", (), Some(&kwargs))?;
} else {
self.object(py)
.call_method("_pre_call", (), Some(&kwargs))?;
self.record_api_call_start(py)?;
}
Logging::PreCall.call(py, (self.object(py), input, api_key, &additional))?;
Ok(())
}
@ -174,37 +149,30 @@ impl LegacyCallbacks for PythonLogger {
&self,
py: Python<'_>,
original_response: &str,
api_key: Option<&str>,
body: Option<&Py<PyDict>>,
headers: Option<&Py<PyDict>>,
) -> PyResult<()> {
let additional = PyDict::new(py);
additional.set_item("complete_input_dict", body)?;
additional.set_item("headers", headers)?;
if self.callbacks_needed(py, "input")? {
let kwargs = PyDict::new(py);
kwargs.set_item("original_response", original_response)?;
kwargs.set_item("additional_args", &additional)?;
self.object(py)
.call_method("post_call", (), Some(&kwargs))?;
} else {
let response = py
.import("json")?
.call_method1("dumps", (original_response,))?;
self.object(py).call_method1(
"record_post_call",
(response, py.None(), py.None(), additional),
)?;
}
Logging::PostCall.call(
py,
(self.object(py), original_response, api_key, &additional),
)?;
Ok(())
}
fn defers_async_logging(&self, py: Python<'_>) -> bool {
self.object(py)
.getattr("_defer_async_logging")
.is_ok_and(|value| value.is_truthy().unwrap_or(false))
Logging::DefersAsync
.call(py, (self.object(py),))
.and_then(|value| value.extract())
.unwrap_or(false)
}
fn defer_success(&self, py: Python<'_>, pending: &Bound<'_, PyAny>) -> PyResult<()> {
self.object(py).setattr("_native_pending_logging", pending)
Logging::DeferSuccess.call(py, (self.object(py), pending))?;
Ok(())
}
fn sync_success_for_async_call(
@ -214,13 +182,7 @@ impl LegacyCallbacks for PythonLogger {
start: &Py<PyAny>,
end: &Option<Py<PyAny>>,
) -> PyResult<()> {
if !self.callbacks_needed(py, "sync_success_async")? {
return Ok(());
}
self.object(py).call_method1(
"handle_sync_success_callbacks_for_async_calls",
(response, start, end),
)?;
Logging::SyncSuccessForAsyncCall.call(py, (self.object(py), response, start, end))?;
Ok(())
}
@ -232,34 +194,11 @@ impl LegacyCallbacks for PythonLogger {
end: &Option<Py<PyAny>>,
asynchronous: bool,
) -> PyResult<Option<Py<PyAny>>> {
if !self.callbacks_needed(
py,
if asynchronous {
"async_failure"
} else {
"sync_failure"
},
)? {
py.import("litellm.rust_bridge.legacy_callbacks")?
.getattr("failure_bookkeeping")?
.call1((self.object(py), error, start, end, asynchronous))?;
return Ok(None);
}
let trace = py
.import("traceback")?
.getattr("format_exception")?
.call1((error,))?;
let trace = pyo3::types::PyString::new(py, "").call_method1("join", (trace,))?;
let value = self.object(py).call_method1(
if asynchronous {
"async_failure_handler"
} else {
"failure_handler"
},
(error, trace, start, end),
)?;
let value =
Logging::FailureHandler.call(py, (self.object(py), error, start, end, asynchronous))?;
Ok(asynchronous.then(|| value.unbind()))
}
fn submit_success(
&self,
py: Python<'_>,
@ -267,22 +206,7 @@ impl LegacyCallbacks for PythonLogger {
start: &Py<PyAny>,
end: &Option<Py<PyAny>>,
) -> PyResult<()> {
if !self.callbacks_needed(py, "sync_success")? {
return self.success_bookkeeping(py, response, start, end, false);
}
let context = py.import("contextvars")?.call_method0("copy_context")?;
py.import("litellm.litellm_core_utils.litellm_logging")?
.getattr("executor")?
.call_method1(
"submit",
(
context.getattr("run")?,
self.object(py).getattr("success_handler")?,
response,
start,
end,
),
)?;
Logging::SubmitSuccess.call(py, (self.object(py), response, start, end))?;
Ok(())
}
@ -293,18 +217,9 @@ impl LegacyCallbacks for PythonLogger {
start: &Py<PyAny>,
end: &Option<Py<PyAny>>,
) -> PyResult<()> {
if !self.callbacks_needed(py, "async_success")? {
return self.success_bookkeeping(py, response, start, end, true);
}
let context = py.import("contextvars")?.call_method0("copy_context")?;
let worker = py
.import("litellm.litellm_core_utils.logging_worker")?
.getattr("GLOBAL_LOGGING_WORKER")?
.getattr("ensure_initialized_and_enqueue")?;
let coroutine = self
.object(py)
.call_method1("async_success_handler", (response, start, end))?;
let enqueue = context.call_method1("run", (worker, &coroutine));
let coroutine =
Logging::AsyncSuccessHandler.call(py, (self.object(py), response, start, end))?;
let enqueue = Logging::Enqueue.call(py, (&coroutine,));
if enqueue.is_err()
&& let Err(error) = coroutine.call_method0("close")
{
@ -315,14 +230,7 @@ impl LegacyCallbacks for PythonLogger {
}
fn custom_pricing_fields(py: Python<'_>) -> PyResult<Vec<String>> {
py.import("litellm.types.utils")?
.getattr("CustomPricingLiteLLMParams")?
.getattr("model_fields")?
.cast_into::<PyDict>()?
.keys()
.iter()
.map(|name| name.extract::<String>())
.collect()
Logging::CustomPricingFields.call(py, ())?.extract()
}
fn redact(
@ -347,58 +255,5 @@ fn redact(
/// Proxy-internal calls skip the legacy success fan-out.
pub fn is_internal_call(py: Python<'_>) -> PyResult<bool> {
py.import("litellm._internal_context")?
.getattr("is_internal_call")?
.call_method0("get")?
.extract()
}
#[cfg(test)]
mod tests {
use pyo3::types::PyDict;
use super::*;
fn logger_whose_registries_need_no_input(py: Python<'_>, bridge_owned: bool) -> PythonLogger {
let locals = PyDict::new(py);
py.run(
c"
import sys
import types
for name in ('litellm', 'litellm.rust_bridge', 'litellm.rust_bridge.legacy_callbacks'):
sys.modules.setdefault(name, types.ModuleType(name))
legacy = sys.modules['litellm.rust_bridge.legacy_callbacks']
legacy.callbacks_needed = lambda logger, phase: logger.needed.get(phase, True)
class Logger:
needed = {'input': False}
logger = Logger()
",
Some(&locals),
Some(&locals),
)
.unwrap();
PythonLogger::new(
locals.get_item("logger").unwrap().unwrap().unbind(),
bridge_owned,
)
}
#[test]
fn a_caller_owned_logger_is_observed_in_full() {
Python::initialize();
Python::attach(|py| {
let logger = logger_whose_registries_need_no_input(py, false);
assert!(logger.callbacks_needed(py, "input").unwrap());
});
}
#[test]
fn a_bridge_owned_logger_is_elided_where_no_registry_needs_it() {
Python::initialize();
Python::attach(|py| {
let logger = logger_whose_registries_need_no_input(py, true);
assert!(!logger.callbacks_needed(py, "input").unwrap());
assert!(logger.callbacks_needed(py, "payload").unwrap());
});
}
Wrapper::IsInternalCall.call(py, ())?.extract()
}

View file

@ -0,0 +1,183 @@
use pyo3::prelude::*;
use strum::{IntoStaticStr, VariantArray};
const MODULE: &str = "litellm.rust_bridge.legacy_callbacks";
/// Every litellm Python internal the native call still borrows, grouped by the subsystem it
/// belongs to. Rust drives the call; these exist only so behaviour that Python owns today
/// (span tracking, the standard logging payload, spend, callback fan-out) keeps working.
/// A group is deleted once Rust owns that subsystem, so this enum only shrinks. Calling a
/// user's own callback is not borrowing and does not belong here.
///
/// `litellm/rust_bridge/legacy_callbacks.py` is the only Python module behind it, and
/// `python_contract.json` pins each function's parameters on both sides.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum LegacyPython {
Wrapper(Wrapper),
Logging(Logging),
DeploymentHooks(DeploymentHooks),
Streaming(Streaming),
}
/// The `@client` wrapper around the call: `function_setup`, limits, credentials,
/// response metadata and the correlation context.
#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, VariantArray)]
pub(crate) enum Wrapper {
#[strum(serialize = "setup")]
Setup,
#[strum(serialize = "check_limits")]
CheckLimits,
#[strum(serialize = "credential_list")]
CredentialList,
#[strum(serialize = "warn_unknown_credential")]
WarnUnknownCredential,
#[strum(serialize = "is_internal_call")]
IsInternalCall,
#[strum(serialize = "finalize")]
Finalize,
#[strum(serialize = "restore_context")]
RestoreContext,
}
/// litellm's `Logging` object and the sync and async callback fan-out behind it.
#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, VariantArray)]
pub(crate) enum Logging {
#[strum(serialize = "custom_pricing_fields")]
CustomPricingFields,
#[strum(serialize = "update_logging")]
Update,
#[strum(serialize = "pre_call")]
PreCall,
#[strum(serialize = "post_call")]
PostCall,
#[strum(serialize = "defers_async_logging")]
DefersAsync,
#[strum(serialize = "defer_success")]
DeferSuccess,
#[strum(serialize = "sync_success_for_async_call")]
SyncSuccessForAsyncCall,
#[strum(serialize = "submit_success")]
SubmitSuccess,
#[strum(serialize = "async_success_handler")]
AsyncSuccessHandler,
#[strum(serialize = "enqueue_logging")]
Enqueue,
#[strum(serialize = "failure_handler")]
FailureHandler,
}
/// The `litellm.utils` fan-outs that run every callback's deployment hook.
#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, VariantArray)]
pub(crate) enum DeploymentHooks {
#[strum(serialize = "before_deployment_call")]
BeforeDeploymentCall,
#[strum(serialize = "after_deployment_success")]
AfterDeploymentSuccess,
#[strum(serialize = "after_deployment_failure")]
AfterDeploymentFailure,
}
/// The Messages stream iterator's logging: the stream flag, the end-of-stream billing
/// from the delivered chunks, and the partial-usage failure path.
#[derive(Clone, Copy, Debug, IntoStaticStr, PartialEq, Eq, VariantArray)]
pub(crate) enum Streaming {
#[strum(serialize = "stream_opened")]
Opened,
#[strum(serialize = "stream_success")]
Success,
#[strum(serialize = "stream_failure")]
Failure,
}
impl LegacyPython {
fn name(self) -> &'static str {
match self {
Self::Wrapper(function) => function.into(),
Self::Logging(function) => function.into(),
Self::DeploymentHooks(function) => function.into(),
Self::Streaming(function) => function.into(),
}
}
pub(crate) fn call<'py, A>(self, py: Python<'py>, args: A) -> PyResult<Bound<'py, PyAny>>
where
A: pyo3::call::PyCallArgs<'py>,
{
py.import(MODULE)?.getattr(self.name())?.call1(args)
}
}
impl Wrapper {
pub(crate) fn call<'py, A>(self, py: Python<'py>, args: A) -> PyResult<Bound<'py, PyAny>>
where
A: pyo3::call::PyCallArgs<'py>,
{
LegacyPython::Wrapper(self).call(py, args)
}
}
impl Logging {
pub(crate) fn call<'py, A>(self, py: Python<'py>, args: A) -> PyResult<Bound<'py, PyAny>>
where
A: pyo3::call::PyCallArgs<'py>,
{
LegacyPython::Logging(self).call(py, args)
}
}
impl Streaming {
pub(crate) fn call<'py, A>(self, py: Python<'py>, args: A) -> PyResult<Bound<'py, PyAny>>
where
A: pyo3::call::PyCallArgs<'py>,
{
LegacyPython::Streaming(self).call(py, args)
}
}
impl DeploymentHooks {
pub(crate) fn call<'py, A>(self, py: Python<'py>, args: A) -> PyResult<Bound<'py, PyAny>>
where
A: pyo3::call::PyCallArgs<'py>,
{
LegacyPython::DeploymentHooks(self).call(py, args)
}
}
#[cfg(test)]
mod tests {
use std::collections::BTreeSet;
use strum::VariantArray;
use super::{DeploymentHooks, LegacyPython, Logging, Streaming, Wrapper};
use crate::test_support::PYTHON_CONTRACT;
#[test]
fn every_borrowed_function_is_in_the_python_contract() {
let contract: serde_json::Map<String, serde_json::Value> =
serde_json::from_str(PYTHON_CONTRACT).unwrap();
let declared: BTreeSet<&str> = contract.keys().map(String::as_str).collect();
let called: Vec<&str> = Wrapper::VARIANTS
.iter()
.map(|&function| LegacyPython::Wrapper(function))
.chain(
Logging::VARIANTS
.iter()
.map(|&function| LegacyPython::Logging(function)),
)
.chain(
DeploymentHooks::VARIANTS
.iter()
.map(|&function| LegacyPython::DeploymentHooks(function)),
)
.chain(
Streaming::VARIANTS
.iter()
.map(|&function| LegacyPython::Streaming(function)),
)
.map(LegacyPython::name)
.collect();
assert_eq!(called.len(), declared.len(), "a function is borrowed twice");
assert_eq!(called.into_iter().collect::<BTreeSet<_>>(), declared);
}
}

View file

@ -2,7 +2,7 @@
//! sync and async callback registries it fans out to, the deployment hooks, the deferred
//! proxy release, and the kwargs rewrites the wrapper makes on the way in (credential-name
//! inheritance, budget and retry-count limits). All of it sits behind one
//! [`CallbackAdapter`](litellm_host_python::CallbackAdapter), so the driver, the routes and
//! [`PythonLifecycle`](litellm_host_python::PythonLifecycle), so the driver, the routes and
//! core never learn which Python object is on the other end.
//!
//! Legacy callbacks receive the caller's own objects and may mutate them. [`PublicCall`]
@ -13,6 +13,7 @@ mod adapter;
mod call;
mod callbacks;
mod deferred;
mod legacy_python;
mod logger;
mod preparation;
#[cfg(test)]
@ -20,8 +21,8 @@ mod preparation;
mod test_support;
pub(crate) use adapter::LegacyLogging;
pub use adapter::LegacySurface;
pub use call::{PublicCall, lookup, run_legacy_call};
pub use adapter::{LegacySurface, PassThroughStream};
pub use call::{PublicCall, run_legacy_call};
pub(crate) use callbacks::{LegacyCallbacks, is_internal_call};
pub(crate) use logger::{DeploymentHooks, PythonLogger, finalize, setup};
pub(crate) use preparation::prepare;

View file

@ -5,34 +5,25 @@ use pyo3::{
types::{PyDict, PyTuple},
};
/// The `Logging` instance one call fans out through, and who owns it. A logger the caller
/// handed in is observed in full, because the caller reads it after the call; one this
/// crate built through `function_setup` is elided wherever no registry needs it.
use crate::legacy_python::{self, Wrapper};
/// The `Logging` instance one call fans out through.
pub struct PythonLogger {
object: Py<PyAny>,
bridge_owned: bool,
}
impl PythonLogger {
pub(crate) fn new(object: Py<PyAny>, bridge_owned: bool) -> Self {
Self {
object,
bridge_owned,
}
pub(crate) fn new(object: Py<PyAny>) -> Self {
Self { object }
}
pub(crate) fn object<'py>(&self, py: Python<'py>) -> &Bound<'py, PyAny> {
self.object.bind(py)
}
pub(crate) fn bridge_owned(&self) -> bool {
self.bridge_owned
}
pub fn clone_ref(&self, py: Python<'_>) -> Self {
Self {
object: self.object.clone_ref(py),
bridge_owned: self.bridge_owned,
}
}
@ -40,34 +31,17 @@ impl PythonLogger {
visit.call(&self.object)
}
pub fn success_bookkeeping(
&self,
py: Python<'_>,
response: &Option<Py<PyAny>>,
start: &Py<PyAny>,
end: &Option<Py<PyAny>>,
asynchronous: bool,
) -> PyResult<()> {
py.import("litellm.rust_bridge.legacy_callbacks")?
.getattr("success_bookkeeping")?
.call1((self.object(py), response, start, end, asynchronous))?;
Ok(())
}
pub fn restore_context(&self, py: Python<'_>) -> PyResult<()> {
py.import("litellm.utils")?
.getattr("_restore_correlation_context_if_supported")?
.call1((self.object(py),))?;
Wrapper::RestoreContext.call(py, (self.object(py),))?;
Ok(())
}
}
/// A bare Python object was not obtained from `setup`, so it is caller-owned.
impl FromPyObject<'_, '_> for PythonLogger {
type Error = PyErr;
fn extract(object: Borrowed<'_, '_, PyAny>) -> PyResult<Self> {
Ok(Self::new(object.to_owned().unbind(), false))
Ok(Self::new(object.to_owned().unbind()))
}
}
@ -75,9 +49,7 @@ pub struct SetupResult<'py>(Bound<'py, PyAny>);
impl SetupResult<'_> {
pub fn logger(&self) -> PyResult<PythonLogger> {
let object = self.0.getattr("logger")?.unbind();
let bridge_owned = self.0.getattr("bridge_owned")?.extract()?;
Ok(PythonLogger::new(object, bridge_owned))
Ok(PythonLogger::new(self.0.getattr("logger")?.unbind()))
}
pub fn kwargs(&self) -> PyResult<Py<PyDict>> {
@ -93,9 +65,8 @@ pub fn setup<'py>(
start: &Py<PyAny>,
asynchronous: bool,
) -> PyResult<SetupResult<'py>> {
py.import("litellm.rust_bridge.legacy_callbacks")?
.getattr("setup")?
.call1((call_type, args, kwargs, start, asynchronous))
Wrapper::Setup
.call(py, (call_type, args, kwargs, start, asynchronous))
.map(SetupResult)
}
@ -107,30 +78,20 @@ pub fn finalize(
start: &Py<PyAny>,
end: &Option<Py<PyAny>>,
) -> PyResult<()> {
py.import("litellm.rust_bridge.legacy_callbacks")?
.getattr("finalize")?
.call1((response, logger.object(py), kwargs, start, end))?;
Wrapper::Finalize.call(py, (response, logger.object(py), kwargs, start, end))?;
Ok(())
}
pub struct DeploymentHooks;
impl DeploymentHooks {
pub fn needed(py: Python<'_>) -> PyResult<bool> {
py.import("litellm.rust_bridge.legacy_callbacks")?
.getattr("deployment_callbacks_needed")?
.call0()?
.extract()
}
pub fn before_call(
py: Python<'_>,
kwargs: &Py<PyDict>,
call_type: &str,
) -> PyResult<Py<PyAny>> {
py.import("litellm.utils")?
.getattr("async_pre_call_deployment_hook")?
.call1((kwargs, call_type))
legacy_python::DeploymentHooks::BeforeDeploymentCall
.call(py, (kwargs, call_type))
.map(Bound::unbind)
}
@ -140,9 +101,8 @@ impl DeploymentHooks {
response: &Option<Py<PyAny>>,
call_type: &str,
) -> PyResult<Py<PyAny>> {
py.import("litellm.utils")?
.getattr("async_post_call_success_deployment_hook")?
.call1((kwargs, response, call_type))
legacy_python::DeploymentHooks::AfterDeploymentSuccess
.call(py, (kwargs, response, call_type))
.map(Bound::unbind)
}
@ -152,9 +112,8 @@ impl DeploymentHooks {
error: &Py<PyBaseException>,
call_type: &str,
) -> PyResult<Py<PyAny>> {
py.import("litellm.utils")?
.getattr("async_post_call_failure_deployment_hook")?
.call1((kwargs, error, call_type))
legacy_python::DeploymentHooks::AfterDeploymentFailure
.call(py, (kwargs, error, call_type))
.map(Bound::unbind)
}
}
@ -185,10 +144,6 @@ class Setup:
reads.append('logger')
return logger
@property
def bridge_owned(self):
reads.append('bridge_owned')
return True
@property
def kwargs(self):
reads.append('kwargs')
return []
@ -206,7 +161,6 @@ result = Setup()
.object(py)
.is(locals.get_item("logger").unwrap().unwrap())
);
assert!(logger.bridge_owned());
assert!(
result
.kwargs()
@ -220,17 +174,8 @@ result = Setup()
.unwrap()
.extract::<Vec<String>>()
.unwrap(),
["logger", "bridge_owned", "kwargs"]
["logger", "kwargs"]
);
});
}
#[test]
fn a_logger_extracted_from_a_bare_object_is_caller_owned() {
Python::initialize();
Python::attach(|py| {
let logger: PythonLogger = py.None().into_bound(py).extract().unwrap();
assert!(!logger.bridge_owned());
});
}
}

View file

@ -3,6 +3,8 @@ use pyo3::{
types::{PyDict, PyList},
};
use crate::legacy_python::Wrapper;
struct CredentialEntry<'py>(Bound<'py, PyAny>);
impl<'py> CredentialEntry<'py> {
@ -22,18 +24,19 @@ pub fn prepare<'py>(
) -> PyResult<Bound<'py, PyDict>> {
let arguments = kwargs.copy()?;
arguments.set_item("litellm_logging_obj", logger.object(py))?;
let litellm = py.import("litellm")?;
inherit_credentials(py, &litellm, &arguments)?;
py.import("litellm.rust_bridge.legacy_callbacks")?
.getattr("check_limits")?
.call1((&arguments,))?;
inherit_credentials(py, &arguments, || {
Ok(Wrapper::CredentialList
.call(py, ())?
.cast_into::<PyList>()?)
})?;
Wrapper::CheckLimits.call(py, (&arguments,))?;
Ok(arguments)
}
fn inherit_credentials(
py: Python<'_>,
litellm: &Bound<'_, PyModule>,
arguments: &Bound<'_, PyDict>,
fn inherit_credentials<'py>(
py: Python<'py>,
arguments: &Bound<'py, PyDict>,
credential_list: impl FnOnce() -> PyResult<Bound<'py, PyList>>,
) -> PyResult<()> {
let Some(requested) = arguments
.get_item("litellm_credential_name")?
@ -45,16 +48,13 @@ fn inherit_credentials(
return Ok(());
}
let requested: String = requested.extract()?;
let credentials = litellm.getattr("credential_list")?.cast_into::<PyList>()?;
let credentials = credential_list()?;
let names = credentials
.iter()
.map(|credential| CredentialEntry(credential).name())
.collect::<PyResult<Vec<_>>>()?;
let Some(index) = names.iter().position(|name| *name == requested) else {
py.import("litellm._logging")?.getattr("verbose_logger")?.call_method1(
"warning",
("litellm_credential_name=%s matched none of the %d loaded credentials; the request runs without it", requested, names.len()),
)?;
Wrapper::WarnUnknownCredential.call(py, (requested, names.len()))?;
return Ok(());
};
let selected = CredentialEntry(credentials.get_item(index)?);
@ -80,19 +80,19 @@ mod tests {
}
fn inherit(py: Python<'_>, locals: &Bound<'_, PyDict>) -> PyResult<()> {
let litellm = PyModule::new(py, "credential_host")?;
litellm.setattr(
"credential_list",
locals.get_item("credentials").unwrap().unwrap(),
)?;
inherit_credentials(
py,
&litellm,
&locals
.get_item("arguments")
.unwrap()
.unwrap()
.cast_into::<PyDict>()?,
|| {
Ok(locals
.get_item("credentials")?
.unwrap()
.cast_into::<PyList>()?)
},
)
}
@ -304,11 +304,11 @@ arguments = {'litellm_credential_name': 'ocr-test'}
fn falsy_credential_names_return_before_loading_credentials() {
Python::initialize();
Python::attach(|py| {
let litellm = PyModule::new(py, "credential_host").unwrap();
for name in [py.None(), py.eval(c"''", None, None).unwrap().unbind()] {
let arguments = PyDict::new(py);
arguments.set_item("litellm_credential_name", name).unwrap();
inherit_credentials(py, &litellm, &arguments).unwrap();
inherit_credentials(py, &arguments, || panic!("credentials must not be loaded"))
.unwrap();
}
});
}

View file

@ -16,7 +16,7 @@ fn defer<'py>(py: Python<'py>, script: &CStr) -> Bound<'py, PyDict> {
py,
PendingLogging {
pending: Some(PendingSuccess {
logger: PythonLogger::new(local(&locals, "logger").unbind(), true),
logger: PythonLogger::new(local(&locals, "logger").unbind()),
response: Some(local(&locals, "response").unbind()),
start: py.None(),
end: Some(py.None()),
@ -79,22 +79,6 @@ assert logger.calls == [], logger.calls
});
}
#[test]
fn a_release_after_the_async_callbacks_went_away_only_keeps_the_books() {
Python::initialize();
Python::attach(|py| {
let locals = defer(py, c"logger.needed = {'async_success': False}");
run(
py,
&locals,
c"
pending.release(True)
assert logger.calls == [('success_bookkeeping', True)], logger.calls
",
);
});
}
#[rstest]
#[case::ordinary_error(c"RuntimeError('queue full')", false)]
#[case::cancellation(c"asyncio.CancelledError()", true)]

View file

@ -1,7 +1,7 @@
use std::ffi::CStr;
use litellm_callbacks::event::{CallEvent, FailureOrigin, Timing};
use litellm_host_python::{AdapterStep, CallbackAdapter, PublicValue};
use litellm_host::event::{FailureOrigin, Timing};
use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle};
use pyo3::exceptions::asyncio::CancelledError;
use pyo3::prelude::*;
use pyo3::types::PyDict;
@ -24,7 +24,7 @@ fn begin<'py>(
py: Python<'py>,
locals: &Bound<'py, PyDict>,
asynchronous: bool,
) -> (LegacyLogging, AdapterStep) {
) -> (LegacyLogging, LifecycleStep) {
let mut logging = legacy_call(py, locals, asynchronous);
let kwargs = local(locals, "kwargs")
.cast_into::<PyDict>()
@ -34,15 +34,15 @@ fn begin<'py>(
(logging, step)
}
fn arguments<'py>(py: Python<'py>, step: AdapterStep) -> Bound<'py, PyDict> {
let AdapterStep::Arguments(arguments) = step else {
fn arguments<'py>(py: Python<'py>, step: LifecycleStep) -> Bound<'py, PyDict> {
let LifecycleStep::Arguments(arguments) = step else {
panic!("expected the prepared arguments");
};
arguments.into_bound(py)
}
fn awaits_deployment_hook(step: &AdapterStep) -> bool {
matches!(step, AdapterStep::Await(_))
fn awaits_deployment_hook(step: &LifecycleStep) -> bool {
matches!(step, LifecycleStep::Await(_))
}
#[rstest]
@ -97,6 +97,43 @@ assert checked is prepared
});
}
#[rstest]
#[case::synchronous(false)]
#[case::asynchronous(true)]
fn a_keyword_the_bridge_never_reads_reaches_every_reader_as_the_callers_object(
#[case] asynchronous: bool,
) {
Python::initialize();
Python::attach(|py| {
let locals = namespace(
py,
c"
opaque = object()
hooked = []
logger.hooks = {'pre': lambda kwargs: hooked.append(kwargs['vendor_extension']) or kwargs}
kwargs = {'logger': logger, 'vendor_extension': opaque}
",
);
let (mut logging, step) = begin(py, &locals, asynchronous);
let step = match step {
LifecycleStep::Await(hook_result) => logging.resume(py, Ok(hook_result)).unwrap(),
step => step,
};
locals.set_item("prepared", arguments(py, step)).unwrap();
locals.set_item("asynchronous", asynchronous).unwrap();
run(
py,
&locals,
c"
assert prepared['vendor_extension'] is opaque
[checked] = [value for name, value in logger.calls if name == 'check_limits']
assert checked['vendor_extension'] is opaque
assert hooked == ([opaque] if asynchronous else []), hooked
",
);
});
}
#[test]
fn response_returned_by_the_post_call_hook_is_finalized_and_returned() {
Python::initialize();
@ -121,7 +158,7 @@ logger.hooks = {'pre': lambda kwargs: kwargs}
let step = logging
.resume(py, Ok(local(&locals, "replacement").unbind()))
.unwrap();
let AdapterStep::Response(returned) = step else {
let LifecycleStep::Response(returned) = step else {
panic!("expected the finalized response");
};
assert!(returned.bind(py).is(local(&locals, "replacement")));
@ -180,13 +217,12 @@ fn failure_callbacks_run_after_the_failure_hook_however_it_ends(#[case] cancelle
.resume(py, Ok(local(&locals, "kwargs").unbind()))
.unwrap();
let failure = PyErr::from_value(local(&locals, "failure"));
let failed = CallEvent::Failed {
let failed = LifecycleEvent::Failed {
timing: TIMING,
origin: FailureOrigin::Call,
error: &failure,
};
let step = logging
.emit(py, &failed, Some(PublicValue::Error(&failure)))
.unwrap();
let step = logging.emit(py, failed).unwrap();
assert!(awaits_deployment_hook(&step));
let hook_result = if cancelled {
Err(CancelledError::new_err("cancelled"))
@ -195,7 +231,7 @@ fn failure_callbacks_run_after_the_failure_hook_however_it_ends(#[case] cancelle
};
assert!(matches!(
logging.resume(py, hook_result).unwrap(),
AdapterStep::Await(_)
LifecycleStep::Await(_)
));
run(
py,
@ -237,7 +273,7 @@ kwargs = {'logger': logger}
.unwrap()
.unbind();
let result = logging.begin(py, kwargs, 0.0).and_then(|step| match step {
AdapterStep::Await(_) => logging.resume(py, Ok(local(&locals, "kwargs").unbind())),
LifecycleStep::Await(_) => logging.resume(py, Ok(local(&locals, "kwargs").unbind())),
step => Ok(step),
});
let error = result.err().unwrap();

View file

@ -1,10 +1,12 @@
use std::ffi::CStr;
use litellm_callbacks::event::{CallEvent, Passthrough, RawResponse, RequestContext, WireRequest};
use litellm_host_python::{AdapterStep, CallbackAdapter};
use litellm_auth::SecretValue;
use litellm_host::event::{MachineEvent, RawResponse, RequestContext, WireRequest};
use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle, to_py};
use proptest::prelude::*;
use pyo3::prelude::*;
use rstest::rstest;
use serde_json::{Value, json};
use serde_json::{Map, Value, json};
use super::LegacyLogging;
use crate::PythonLogger;
@ -23,20 +25,12 @@ class PayloadLogger(StubLogger):
def pre_call(self, input, api_key, additional_args):
self.record('pre_call', None)
self.pre = additional_args
self.pre_api_key = api_key
on_pre_call(additional_args)
def _pre_call(self, input, api_key, additional_args):
self.record('_pre_call', None)
def record_api_call_start_time(self):
self.record('record_api_call_start_time', None)
def post_call(self, original_response, additional_args):
def post_call(self, original_response, api_key, additional_args):
self.record('post_call', None)
self.post = (original_response, additional_args)
def record_post_call(self, response, *rest):
self.record('record_post_call', response)
self.post = (original_response, api_key, additional_args)
request = Request()
kwargs = {}
@ -52,33 +46,47 @@ fn document(source: &str) -> Value {
json!({"type": "document_url", "document_url": source})
}
fn before_send(script: &CStr, caller: Value, body: Value) -> WireRequest {
before_send_with_secrets(script, caller, body, &[])
fn before_send(script: &CStr, body: Value) -> WireRequest {
before_send_with_secrets(script, json!({}), body, &[])
}
/// Runs `before_send` over `body` for a caller whose route-side view is `caller`, with the
/// Python objects `script` binds, then delivers the provider's raw response the way the
/// Runs `before_send` over `body` for a route whose parameters are `optional_params`, with
/// the Python objects `script` binds, then delivers the provider's raw response the way the
/// driver does and runs the script's `check()`.
fn before_send_with_secrets(
script: &CStr,
caller: Value,
optional_params: Value,
body: Value,
secret_fields: &[&str],
) -> WireRequest {
before_send_bound(&[], script, optional_params, body, secret_fields)
}
/// [`before_send_with_secrets`] with `bindings` placed in the namespace before `script` runs.
fn before_send_bound(
bindings: &[(&str, &Value)],
script: &CStr,
optional_params: Value,
body: Value,
secret_fields: &[&str],
) -> WireRequest {
Python::initialize();
Python::attach(|py| {
let locals = namespace(py, PAYLOAD_LOGGER);
for &(name, value) in bindings {
locals.set_item(name, to_py(py, value).unwrap()).unwrap();
}
run(py, &locals, script);
let mut logging = LegacyLogging {
logger: Some(PythonLogger::new(local(&locals, "logger").unbind(), true)),
logger: Some(PythonLogger::new(local(&locals, "logger").unbind())),
..legacy_call(py, &locals, false)
};
let context = RequestContext {
model: "model".into(),
custom_llm_provider: "provider".into(),
optional_params: caller.clone(),
passthrough_fields: Passthrough::unchanged(caller.as_object().unwrap(), &body),
optional_params,
secret_fields: secret_fields.iter().map(|name| name.to_string()).collect(),
api_key: Some(SecretValue::new("route-key")),
};
let wire = WireRequest {
url: "https://provider.invalid/ocr".into(),
@ -86,17 +94,17 @@ fn before_send_with_secrets(
body,
};
let step = logging.before_send(py, Box::new(wire), &context).unwrap();
let raw = CallEvent::ResponseReceived {
let raw = MachineEvent::ResponseReceived {
raw: RawResponse {
body: "raw response".into(),
},
};
assert!(matches!(
logging.emit(py, &raw, None).unwrap(),
AdapterStep::Done
logging.emit(py, LifecycleEvent::Machine(&raw)).unwrap(),
LifecycleStep::Done
));
run(py, &locals, c"check()");
let AdapterStep::Wire(wire) = step else {
let LifecycleStep::Wire(wire) = step else {
panic!("before_send did not hand back the wire request");
};
*wire
@ -129,11 +137,7 @@ def check():
")]
fn passthrough_keys_reach_pre_call_as_the_callers_own_objects(#[case] script: &CStr) {
let body = json!({"model": "model", "document": document(DOCUMENT), "pages": [0]});
let wire = before_send(
script,
json!({"document": document(DOCUMENT), "pages": [0]}),
body.clone(),
);
let wire = before_send(script, body.clone());
assert_eq!(wire.body, body);
}
@ -149,7 +153,6 @@ def check():
assert document['document_url'] == 'data:application/pdf;base64,ZWRpdGVk'
",
json!({"document": document(DOCUMENT)}),
json!({"document": document(DOCUMENT)}),
);
assert_eq!(wire.body["document"], document(EDITED));
}
@ -168,7 +171,6 @@ def check():
assert observed == [False], observed
assert document == {'type': 'document_url', 'document_url': 'https://example.invalid/scan.pdf'}
",
json!({"document": document("https://example.invalid/scan.pdf")}),
json!({"document": document(DOCUMENT)}),
);
assert_eq!(
@ -177,6 +179,23 @@ def check():
);
}
#[test]
fn a_caller_value_with_no_json_form_is_left_out_of_realiasing() {
let body = json!({"pages": [0]});
let wire = before_send(
c"
opaque = object()
kwargs = {'pages': opaque}
observed = []
on_pre_call = lambda args: observed.append(args['complete_input_dict']['pages'])
def check():
assert observed == [[0]], observed
",
body.clone(),
);
assert_eq!(wire.body, body);
}
#[rstest]
#[case::body(
c"
@ -192,7 +211,7 @@ def on_pre_call(args):
)]
fn rebinding_the_payload_envelope_does_not_reach_the_wire(#[case] script: &CStr) {
let body = json!({"document": document(DOCUMENT)});
let wire = before_send(script, json!({}), body.clone());
let wire = before_send(script, body.clone());
assert_eq!(wire.body, body);
assert_eq!(wire.headers, [("x-route".to_string(), "route".to_string())]);
}
@ -205,7 +224,6 @@ def on_pre_call(args):
args['headers']['x-callback'] = 'edited'
",
json!({}),
json!({}),
);
assert_eq!(
wire.headers,
@ -288,7 +306,7 @@ def on_pre_call(args):
)]
fn pre_call_body_edits_reach_the_wire(#[case] script: &CStr, #[case] expected: Value) {
let body = json!({"document": document(DOCUMENT)});
let wire = before_send(script, json!({"document": document(DOCUMENT)}), body);
let wire = before_send(script, body);
assert_eq!(wire.body, expected);
}
@ -302,7 +320,6 @@ def on_pre_call(args):
retained['x-retained'] = 'sent'
",
json!({}),
json!({}),
);
assert_eq!(
wire.headers,
@ -314,52 +331,193 @@ def on_pre_call(args):
}
#[test]
fn post_call_receives_the_raw_response_and_the_payload_dicts_pre_call_saw() {
fn post_call_receives_the_raw_response_the_route_key_and_the_body_and_headers_pre_call_saw() {
before_send(
c"
def check():
original_response, additional_args = logger.post
original_response, api_key, additional_args = logger.post
assert original_response == 'raw response', original_response
assert api_key == logger.pre_api_key == 'route-key', (api_key, logger.pre_api_key)
assert additional_args == {
'complete_input_dict': logger.pre['complete_input_dict'],
'headers': logger.pre['headers'],
}, additional_args
assert additional_args['complete_input_dict'] is logger.pre['complete_input_dict']
assert additional_args['headers'] is logger.pre['headers']
",
json!({}),
json!({"document": document(DOCUMENT)}),
);
}
#[rstest]
#[case::every_phase_listens(c"{}", &["pre_call", "post_call"])]
#[case::no_input_callback(
c"{'input': False}",
&["_pre_call", "record_api_call_start_time", "record_post_call"]
)]
#[case::no_payload_consumer(c"{'payload': False}", &["record_api_call_start_time"])]
fn payload_callbacks_run_only_for_the_phases_someone_listens_to(
#[case] needed: &CStr,
#[case] expected_calls: &[&str],
) {
let script = std::ffi::CString::new(format!(
"
logger.needed = {needed}
#[test]
fn every_request_runs_the_full_pre_call_and_post_call() {
let wire = before_send(
c"
def on_pre_call(args):
args['complete_input_dict']['include_image_base64'] = True
def check():
assert logger.names() == {expected_calls:?}, logger.calls
assert logger.names() == ['pre_call', 'post_call'], logger.calls
",
needed = needed.to_str().unwrap(),
expected_calls = expected_calls,
))
.unwrap();
let body = json!({"document": document(DOCUMENT)});
let wire = before_send(&script, json!({}), body.clone());
let edited = json!({"document": document(DOCUMENT), "include_image_base64": true});
json!({"document": document(DOCUMENT)}),
);
assert_eq!(
wire.body,
if expected_calls.contains(&"pre_call") {
edited
} else {
body
}
json!({"document": document(DOCUMENT), "include_image_base64": true})
);
}
/// What one pre-call callback does to the payload it is handed.
#[derive(Clone, Debug)]
enum Edit {
Nothing,
Set(String, Value),
Remove(String),
Rebind(Value),
RebindThenSetRetained(String, Value),
}
impl Edit {
fn script(&self) -> Value {
match self {
Self::Nothing => json!({"kind": "nothing"}),
Self::Set(key, value) => json!({"kind": "set", "key": key, "value": value}),
Self::Remove(key) => json!({"kind": "remove", "key": key}),
Self::Rebind(value) => json!({"kind": "rebind", "value": value}),
Self::RebindThenSetRetained(key, value) => {
json!({"kind": "rebind_then_set_retained", "key": key, "value": value})
}
}
}
/// The legacy contract: the provider is sent the body object `pre_call` received, as
/// the callback left it. Rebinding the envelope's key points the envelope elsewhere and
/// leaves that object alone.
fn sent(&self, body: &Map<String, Value>) -> Value {
let mut sent = body.clone();
match self {
Self::Nothing | Self::Rebind(_) => {}
Self::Set(key, value) | Self::RebindThenSetRetained(key, value) => {
sent.insert(key.clone(), value.clone());
}
Self::Remove(key) => {
sent.remove(key);
}
}
Value::Object(sent)
}
}
/// How the caller's keyword for a body key relates to what the route sends under it.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Caller {
PassedUnchanged,
RewrittenByTheRoute,
NotPassed,
}
const MODEL: &CStr = c"
aliased = {}
def on_pre_call(args):
body = args['complete_input_dict']
aliased.update({name: body[name] is kwargs[name] for name in unchanged})
kind = edit['kind']
if kind == 'set':
body[edit['key']] = edit['value']
elif kind == 'remove':
body.pop(edit['key'], None)
elif kind == 'rebind':
args['complete_input_dict'] = edit['value']
elif kind == 'rebind_then_set_retained':
args['complete_input_dict'] = {}
body[edit['key']] = edit['value']
def check():
assert aliased == {name: True for name in unchanged}, aliased
assert logger.names() == ['pre_call', 'post_call'], logger.calls
";
fn json_value() -> impl Strategy<Value = Value> {
let leaf = prop_oneof![
Just(Value::Null),
any::<bool>().prop_map(Value::from),
any::<i64>().prop_map(Value::from),
any::<f64>()
.prop_filter("JSON has no NaN or infinity", |number| number.is_finite())
.prop_map(Value::from),
".{0,8}".prop_map(Value::from),
];
leaf.prop_recursive(3, 24, 4, |inner| {
prop_oneof![
prop::collection::vec(inner.clone(), 0..4).prop_map(Value::from),
prop::collection::btree_map(key(), inner, 0..4)
.prop_map(|fields| Value::Object(fields.into_iter().collect())),
]
})
}
fn key() -> impl Strategy<Value = String> {
"[a-z]{1,6}"
}
fn caller() -> impl Strategy<Value = Caller> {
prop_oneof![
Just(Caller::PassedUnchanged),
Just(Caller::RewrittenByTheRoute),
Just(Caller::NotPassed),
]
}
fn edit() -> impl Strategy<Value = Edit> {
prop_oneof![
Just(Edit::Nothing),
(key(), json_value()).prop_map(|(key, value)| Edit::Set(key, value)),
key().prop_map(Edit::Remove),
json_value().prop_map(Edit::Rebind),
(key(), json_value()).prop_map(|(key, value)| Edit::RebindThenSetRetained(key, value)),
]
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(128))]
/// For any body, any caller keywords and any callback edit: every keyword the route
/// sends unchanged reaches `pre_call` as the caller's own object, and the provider is
/// sent exactly what the model says, so a callback that edits nothing changes nothing.
#[test]
fn the_wire_is_the_body_pre_call_received_as_the_callback_left_it(
fields in prop::collection::btree_map(key(), (json_value(), caller()), 0..5),
edit in edit(),
) {
let body: Map<String, Value> = fields
.iter()
.map(|(name, (value, _))| (name.clone(), value.clone()))
.collect();
let kwargs: Map<String, Value> = fields
.iter()
.filter_map(|(name, (value, caller))| match caller {
Caller::PassedUnchanged => Some((name.clone(), value.clone())),
Caller::RewrittenByTheRoute => Some((name.clone(), json!([value]))),
Caller::NotPassed => None,
})
.collect();
let unchanged: Value = fields
.iter()
.filter(|(_, (_, caller))| *caller == Caller::PassedUnchanged)
.map(|(name, _)| Value::from(name.clone()))
.collect();
let wire = before_send_bound(
&[
("kwargs", &Value::Object(kwargs)),
("unchanged", &unchanged),
("edit", &edit.script()),
],
MODEL,
json!({}),
Value::Object(body.clone()),
&[],
);
prop_assert_eq!(wire.body, edit.sent(&body));
prop_assert_eq!(wire.headers, [("x-route".to_string(), "route".to_string())]);
}
}

View file

@ -5,64 +5,94 @@ use pyo3::types::{PyDict, PyTuple};
use crate::{LegacyLogging, LegacySurface, PublicCall};
/// Stand-ins for every litellm function the legacy contract calls. Tests share one
/// interpreter and run concurrently, so each stub is installed idempotently and forwards to
/// the per-test `StubLogger` it is handed (directly, or as `kwargs['logger']`).
/// The parameters of every `legacy_callbacks` function, as the real module declares them.
/// `tests/test_litellm/rust_bridge/test_legacy_callbacks.py` pins this file to the Python
/// signatures, and [`namespace`] binds every fake call against it.
pub(crate) const PYTHON_CONTRACT: &str = include_str!("../python_contract.json");
/// Stand-ins for `legacy_callbacks`, the only Python module the crate calls. Tests
/// share one interpreter and run concurrently, so each fake is installed idempotently and
/// forwards to the per-test `StubLogger` it is handed (directly, or as `kwargs['logger']`).
/// Every fake is bound against the contract first, so a call the real module would reject
/// fails here too.
const STUBS: &CStr = c"
import contextvars
import inspect
import json
import sys
import traceback
import types
for name in (
'litellm',
'litellm.utils',
'litellm.types',
'litellm.types.utils',
'litellm._internal_context',
'litellm.litellm_core_utils',
'litellm.litellm_core_utils.logging_worker',
'litellm.litellm_core_utils.litellm_logging',
'litellm.rust_bridge',
'litellm.rust_bridge.legacy_callbacks',
):
for name in ('litellm', 'litellm.rust_bridge', 'litellm.rust_bridge.legacy_callbacks'):
sys.modules.setdefault(name, types.ModuleType(name))
legacy = sys.modules['litellm.rust_bridge.legacy_callbacks']
legacy.setup = lambda call_type, args, kwargs, start, asynchronous: types.SimpleNamespace(
logger=kwargs['logger_factory'](kwargs) if 'logger_factory' in kwargs else kwargs['logger'],
kwargs=kwargs,
bridge_owned=True,
)
legacy.deployment_callbacks_needed = lambda: True
legacy.check_limits = lambda arguments: arguments['logger'].check_limits(arguments)
legacy.callbacks_needed = lambda logger, phase: logger.needed.get(phase, True)
legacy.success_bookkeeping = lambda logger, response, start, end, asynchronous: logger.record(
'success_bookkeeping', asynchronous
)
legacy.failure_bookkeeping = lambda logger, error, start, end, asynchronous: logger.record(
'failure_bookkeeping', asynchronous
)
legacy.finalize = lambda response, logger, kwargs, start, end: logger.record('finalize', response)
CONTRACT = json.loads(python_contract)
utils = sys.modules['litellm.utils']
utils.async_pre_call_deployment_hook = lambda kwargs, call_type: kwargs['logger'].hook(
'pre', kwargs, call_type
)
utils.async_post_call_success_deployment_hook = lambda kwargs, response, call_type: kwargs[
'logger'
].hook('success', response, call_type)
utils.async_post_call_failure_deployment_hook = lambda kwargs, error, call_type: kwargs[
'logger'
].hook('failure', error, call_type)
utils._restore_correlation_context_if_supported = lambda logger: logger.record('restore', None)
internal = sys.modules['litellm._internal_context']
if not hasattr(internal, 'is_internal_call'):
internal.is_internal_call = contextvars.ContextVar('is_internal_call', default=False)
def contracted(name, fake):
signature = inspect.Signature(
[inspect.Parameter(parameter, inspect.Parameter.POSITIONAL_OR_KEYWORD) for parameter in CONTRACT[name]]
)
sys.modules['litellm.types.utils'].CustomPricingLiteLLMParams = type(
'CustomPricingLiteLLMParams', (), {'model_fields': {'ocr_cost_per_page': None}}
)
def checked(*args, **kwargs):
signature.bind(*args, **kwargs)
return fake(*args, **kwargs)
return checked
if not hasattr(legacy, 'is_internal'):
legacy.is_internal = contextvars.ContextVar('is_internal_call', default=False)
FAKES = {
'setup': lambda call_type, args, kwargs, start, asynchronous: types.SimpleNamespace(
logger=kwargs['logger_factory'](kwargs) if 'logger_factory' in kwargs else kwargs['logger'],
kwargs=kwargs,
),
'check_limits': lambda arguments: arguments['logger'].check_limits(arguments),
'finalize': lambda response, logger, kwargs, start, end: logger.record('finalize', response),
'update_logging': lambda logger, kwargs, model, optional_params, litellm_params, provider: logger.update_from_kwargs(
kwargs=kwargs,
model=model,
optional_params=optional_params,
litellm_params=litellm_params,
custom_llm_provider=provider,
),
'pre_call': lambda logger, input, api_key, additional_args: logger.pre_call(input, api_key, additional_args),
'post_call': lambda logger, original_response, api_key, additional_args: logger.post_call(
original_response, api_key, additional_args
),
'defers_async_logging': lambda logger: bool(getattr(logger, '_defer_async_logging', False)),
'defer_success': lambda logger, pending: setattr(logger, '_native_pending_logging', pending),
'sync_success_for_async_call': lambda logger, response, start, end: logger.handle_sync_success_callbacks_for_async_calls(
response, start, end
),
'failure_handler': lambda logger, error, start, end, asynchronous: (
logger.async_failure_handler if asynchronous else logger.failure_handler
)(error, ''.join(traceback.format_exception(error)), start, end),
'submit_success': lambda logger, response, start, end: logger.record('submit', (response, start, end)),
'async_success_handler': lambda logger, response, start, end: logger.async_success_handler(response, start, end),
'enqueue_logging': lambda coroutine: coroutine.enqueue(),
'restore_context': lambda logger: logger.record('restore', None),
'custom_pricing_fields': lambda: ('ocr_cost_per_page',),
'is_internal_call': lambda: legacy.is_internal.get(),
'credential_list': lambda: [],
'warn_unknown_credential': lambda name, loaded: None,
'before_deployment_call': lambda kwargs, call_type: kwargs['logger'].hook('pre', kwargs, call_type),
'after_deployment_success': lambda kwargs, response, call_type: kwargs['logger'].hook(
'success', response, call_type
),
'after_deployment_failure': lambda kwargs, error, call_type: kwargs['logger'].hook('failure', error, call_type),
'stream_opened': lambda logger: logger.record('stream_opened', None),
'stream_success': lambda logger, request_body, chunks, start, end, first_chunk: logger.record(
'stream_success', list(chunks)
),
'stream_failure': lambda logger, request_body, chunks, error: logger.record('stream_failure', error),
}
assert FAKES.keys() == CONTRACT.keys(), sorted(FAKES.keys() ^ CONTRACT.keys())
for name, fake in FAKES.items():
setattr(legacy, name, contracted(name, fake))
unraisable = sys.modules.setdefault(
@ -77,20 +107,6 @@ def unraisable_from(owner):
return [error for source, error in unraisable.events if source is owner]
class Worker:
def ensure_initialized_and_enqueue(self, coroutine):
return coroutine.enqueue()
class Executor:
def submit(self, run, handler, *args):
handler.__self__.record('submit', args)
sys.modules['litellm.litellm_core_utils.logging_worker'].GLOBAL_LOGGING_WORKER = Worker()
sys.modules['litellm.litellm_core_utils.litellm_logging'].executor = Executor()
class StubCoroutine:
def __init__(self, logger):
self.logger = logger
@ -106,7 +122,6 @@ class StubCoroutine:
class StubLogger:
def __init__(self):
self.calls = []
self.needed = {}
self.hooks = {}
self.on_enqueue = lambda coroutine: None
@ -147,6 +162,7 @@ logger = StubLogger()
/// A namespace with the stubs, `StubLogger` and a fresh `logger`, after `script` ran in it.
pub(crate) fn namespace<'py>(py: Python<'py>, script: &CStr) -> Bound<'py, PyDict> {
let locals = PyDict::new(py);
locals.set_item("python_contract", PYTHON_CONTRACT).unwrap();
py.run(STUBS, Some(&locals), Some(&locals)).unwrap();
py.run(script, Some(&locals), Some(&locals)).unwrap();
locals
@ -181,6 +197,7 @@ pub(crate) fn legacy_call(
LegacySurface {
call_type: "test",
input_description: "test input",
stream: None,
},
call,
asynchronous,

View file

@ -1,7 +1,7 @@
use std::ffi::CStr;
use litellm_callbacks::event::{CallEvent, FailureOrigin, Timing};
use litellm_host_python::{AdapterStep, CallbackAdapter, PublicValue};
use litellm_host::event::{FailureOrigin, Timing};
use litellm_host_python::{LifecycleEvent, LifecycleStep, PythonLifecycle};
use pyo3::exceptions::PyRuntimeError;
use pyo3::exceptions::asyncio::CancelledError;
use pyo3::prelude::*;
@ -19,52 +19,52 @@ const TIMING: Timing = Timing {
fn logged(py: Python<'_>, locals: &Bound<'_, PyDict>, asynchronous: bool) -> LegacyLogging {
LegacyLogging {
logger: Some(PythonLogger::new(local(locals, "logger").unbind(), true)),
logger: Some(PythonLogger::new(local(locals, "logger").unbind())),
..legacy_call(py, locals, asynchronous)
}
}
fn succeed(py: Python<'_>, locals: &Bound<'_, PyDict>, logging: &mut LegacyLogging) -> AdapterStep {
fn succeed(
py: Python<'_>,
locals: &Bound<'_, PyDict>,
logging: &mut LegacyLogging,
) -> LifecycleStep {
let response = local(locals, "response").unbind();
logging
.emit(
py,
&CallEvent::Succeeded { timing: TIMING },
Some(PublicValue::Response(&response)),
LifecycleEvent::Succeeded {
timing: TIMING,
response: &response,
},
)
.unwrap()
}
fn fail(py: Python<'_>, locals: &Bound<'_, PyDict>, logging: &mut LegacyLogging) -> AdapterStep {
fn fail(py: Python<'_>, locals: &Bound<'_, PyDict>, logging: &mut LegacyLogging) -> LifecycleStep {
let failure = PyErr::from_value(local(locals, "failure"));
logging
.emit(
py,
&CallEvent::Failed {
LifecycleEvent::Failed {
timing: TIMING,
origin: FailureOrigin::Host,
error: &failure,
},
Some(PublicValue::Error(&failure)),
)
.unwrap()
}
#[rstest]
#[case::sync_listened(false, c"", &["submit"])]
#[case::sync_unlistened(false, c"logger.needed = {'sync_success': False}", &["success_bookkeeping"])]
#[case::async_listened(
true,
c"",
&["async_success_handler", "enqueued", "sync_success_for_async_call"]
)]
#[case::async_unlistened(
true,
c"logger.needed = {'async_success': False, 'sync_success_async': False}",
&["success_bookkeeping"]
)]
#[case::async_deferred(true, c"logger._defer_async_logging = True", &["sync_success_for_async_call"])]
#[case::async_with_fallbacks(true, c"kwargs = {'fallbacks': ['other']}", &["sync_success_for_async_call"])]
fn success_reaches_only_the_callbacks_that_listen(
fn success_reaches_the_logging_handlers(
#[case] asynchronous: bool,
#[case] script: &CStr,
#[case] expected: &[&str],
@ -76,7 +76,7 @@ fn success_reaches_only_the_callbacks_that_listen(
let mut logging = logged(py, &locals, asynchronous);
assert!(matches!(
succeed(py, &locals, &mut logging),
AdapterStep::Done
LifecycleStep::Done
));
let names: Vec<String> = local(&locals, "logger")
.call_method0("names")
@ -109,7 +109,10 @@ fn internal_calls_skip_failure_callbacks_only_when_asynchronous(
internal: true,
..logged(py, &locals, asynchronous)
};
assert!(matches!(fail(py, &locals, &mut logging), AdapterStep::Done));
assert!(matches!(
fail(py, &locals, &mut logging),
LifecycleStep::Done
));
let names: Vec<String> = local(&locals, "logger")
.call_method0("names")
.unwrap()
@ -157,7 +160,7 @@ logger = FailingLogger()
let mut logging = logged(py, &locals, true);
assert!(matches!(
succeed(py, &locals, &mut logging),
AdapterStep::Done
LifecycleStep::Done
));
assert!(
logging
@ -173,14 +176,8 @@ logger = FailingLogger()
#[rstest]
#[case::sync_listened(false, c"", &["failure_handler"])]
#[case::sync_unlistened(false, c"logger.needed = {'sync_failure': False}", &["failure_bookkeeping"])]
#[case::async_listened(true, c"", &["failure_handler", "async_failure_handler"])]
#[case::async_unlistened(
true,
c"logger.needed = {'sync_failure': False, 'async_failure': False}",
&["failure_bookkeeping", "failure_bookkeeping"]
)]
fn failure_reaches_only_the_callbacks_that_listen(
fn failure_reaches_the_logging_handlers(
#[case] asynchronous: bool,
#[case] script: &CStr,
#[case] expected: &[&str],
@ -192,7 +189,10 @@ fn failure_reaches_only_the_callbacks_that_listen(
let mut logging = logged(py, &locals, asynchronous);
let step = fail(py, &locals, &mut logging);
let awaits_async_handler = expected.contains(&"async_failure_handler");
assert_eq!(matches!(step, AdapterStep::Await(_)), awaits_async_handler);
assert_eq!(
matches!(step, LifecycleStep::Await(_)),
awaits_async_handler
);
let names: Vec<String> = local(&locals, "logger")
.call_method0("names")
.unwrap()
@ -227,7 +227,7 @@ logger = FailingLogger()
let mut logging = logged(py, &locals, true);
assert!(matches!(
fail(py, &locals, &mut logging),
AdapterStep::Await(_)
LifecycleStep::Await(_)
));
assert!(
logging
@ -265,7 +265,7 @@ fn the_async_failure_handler_ends_the_call_unless_it_was_cancelled(
};
let expected = result.as_ref().err().map(|error| error.value(py).clone());
match logging.resume(py, result) {
Ok(step) => assert!(done && matches!(step, AdapterStep::Done)),
Ok(step) => assert!(done && matches!(step, LifecycleStep::Done)),
Err(propagated) => {
assert!(!done);
assert!(propagated.value(py).is(expected.unwrap()));

View file

@ -1,135 +0,0 @@
use std::time::{SystemTime, UNIX_EPOCH};
use serde_json::{Map, Value};
/// Seconds since the Unix epoch, on one clock for every host.
pub fn epoch_seconds() -> f64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|duration| duration.as_secs_f64())
.unwrap_or(0.0)
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct Timing {
pub start_time: f64,
pub end_time: f64,
}
/// The provider request as it is about to leave, offered to the host for rewriting.
#[derive(Clone, Debug, PartialEq)]
pub struct WireRequest {
pub url: String,
pub headers: Vec<(String, String)>,
pub body: Value,
}
/// What the route knows about the request it is sending, for a host that logs it. The
/// route owns these facts; a host reads them beside the wire request and never rewrites
/// them.
#[derive(Clone, Debug, PartialEq)]
pub struct RequestContext {
pub model: String,
pub custom_llm_provider: String,
/// The route's parameters before the provider transformation.
pub optional_params: Value,
pub passthrough_fields: Passthrough,
/// Optional-param names that carry credentials and must be redacted when logged.
pub secret_fields: Vec<String>,
}
/// Body keys whose values are the caller's inputs, unchanged by the route. The only way to
/// build one is to compare the two, so a route cannot name a key it rewrote.
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct Passthrough(Vec<String>);
impl Passthrough {
pub fn unchanged(caller: &Map<String, Value>, body: &Value) -> Self {
Self(
caller
.iter()
.filter(|(name, value)| body.get(name.as_str()) == Some(*value))
.map(|(name, _)| name.clone())
.collect(),
)
}
pub fn iter(&self) -> impl Iterator<Item = &str> {
self.0.iter().map(String::as_str)
}
pub fn contains(&self, name: &str) -> bool {
self.0.iter().any(|field| field == name)
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RawResponse {
pub body: String,
}
/// Whether a failure surfaced inside the call, including a host op the call asked for,
/// or in a host step around it (preparing the arguments, finalizing the response).
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum FailureOrigin {
Call,
Host,
}
#[derive(Clone, Debug, PartialEq)]
pub enum CallEvent {
ResponseReceived {
raw: RawResponse,
},
Succeeded {
timing: Timing,
},
Failed {
timing: Timing,
origin: FailureOrigin,
},
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use serde_json::json;
use super::*;
#[rstest]
#[case::unchanged_scalar(json!({"pages": [0]}), json!({"pages": [0]}), &["pages"])]
#[case::unchanged_explicit_null(json!({"pages": null}), json!({"pages": null}), &["pages"])]
#[case::unchanged_nested_object(
json!({"document": {"type": "document_url", "document_url": "https://a/b.pdf"}}),
json!({"document": {"type": "document_url", "document_url": "https://a/b.pdf"}, "model": "m"}),
&["document"]
)]
#[case::rewritten_value(
json!({"document": {"type": "document_url", "document_url": "https://a/b.pdf"}}),
json!({"document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}}),
&[]
)]
#[case::dropped_nested_field(
json!({"document": {"type": "image_url", "image_url": "https://a/b.png", "document_name": "b.png"}}),
json!({"document": {"type": "image_url", "image_url": "https://a/b.png"}}),
&[]
)]
#[case::added_nested_field(
json!({"document": {"type": "image_url", "image_url": "https://a/b.png"}}),
json!({"document": {"type": "image_url", "image_url": "https://a/b.png", "detail": "high"}}),
&[]
)]
#[case::reordered_array(json!({"pages": [0, 1]}), json!({"pages": [1, 0]}), &[])]
#[case::consumed_by_the_route(json!({"api_key": "k", "pages": [0]}), json!({"pages": [0]}), &["pages"])]
#[case::added_by_the_route(json!({}), json!({"model": "m"}), &[])]
#[case::non_object_body(json!({"pages": [0]}), json!([{"pages": [0]}]), &[])]
fn passthrough_is_exactly_the_callers_unchanged_keys(
#[case] caller: Value,
#[case] body: Value,
#[case] expected: &[&str],
) {
let passthrough = Passthrough::unchanged(caller.as_object().unwrap(), &body);
assert_eq!(passthrough.iter().collect::<Vec<_>>(), expected);
}
}

View file

@ -6,6 +6,7 @@ license.workspace = true
repository.workspace = true
[dependencies]
fancy-regex.workspace = true
litellm-types.workspace = true
serde.workspace = true
serde_json.workspace = true
@ -13,3 +14,6 @@ serde_path_to_error = "0.1"
serde_with.workspace = true
thiserror.workspace = true
url.workspace = true
[dev-dependencies]
rstest.workspace = true

View file

@ -0,0 +1,115 @@
use super::public::PublicError;
use super::rules::{Rule, contains_any};
/// The text branches of `_map_cohere_exception`, in its order.
pub(super) const RULES: &[Rule] = &[
Rule::new(
|mapping| {
contains_any(
&mapping.error_str,
&["invalid api token", "No API key provided."],
)
},
PublicError::Authentication,
),
Rule::new(
|mapping| mapping.error_str.contains("invalid type: parameter"),
PublicError::BadRequest,
),
Rule::new(
|mapping| mapping.error_str.contains("too many tokens"),
PublicError::ContextWindowExceeded,
),
Rule::new(
|mapping| {
mapping
.error_str
.to_lowercase()
.contains("internal server error")
},
PublicError::InternalServer,
),
Rule::new(
|mapping| mapping.status.is_none() && mapping.error_str.contains("invalid type:"),
PublicError::BadRequest,
),
Rule::new(
|mapping| mapping.status.is_none() && mapping.error_str.contains("Unexpected server error"),
PublicError::InternalServer,
),
];
#[cfg(test)]
mod tests {
use super::super::rules::first_match;
use super::super::testing::mapping;
use super::*;
fn classified(text: &str) -> Option<PublicError> {
classified_with(Some(400), text)
}
fn classified_with(status: Option<u16>, text: &str) -> Option<PublicError> {
first_match(RULES, &mapping(status, text)).map(|rule| rule.error)
}
#[rstest::rstest]
#[case::invalid_token("invalid api token", PublicError::Authentication)]
#[case::no_api_key("No API key provided.", PublicError::Authentication)]
#[case::invalid_parameter("invalid type: parameter x", PublicError::BadRequest)]
#[case::too_many_tokens("too many tokens", PublicError::ContextWindowExceeded)]
#[case::internal_server_text("Internal Server Error", PublicError::InternalServer)]
#[case::internal_server_any_case("INTERNAL server ERROR", PublicError::InternalServer)]
fn each_text_rule_claims_its_marker(#[case] text: &str, #[case] expected: PublicError) {
assert_eq!(classified(text), Some(expected));
}
#[rstest::rstest]
#[case::token_before_parameter(
"invalid api token invalid type: parameter",
PublicError::Authentication
)]
#[case::parameter_before_tokens(
"invalid type: parameter too many tokens",
PublicError::BadRequest
)]
#[case::tokens_before_internal(
"too many tokens Internal Server Error",
PublicError::ContextWindowExceeded
)]
fn the_earlier_rule_wins_when_two_apply(#[case] text: &str, #[case] expected: PublicError) {
assert_eq!(classified(text), Some(expected));
}
#[rstest::rstest]
#[case::invalid_type(None, "invalid type: x", Some(PublicError::BadRequest))]
#[case::unexpected_server_error(
None,
"Unexpected server error",
Some(PublicError::InternalServer)
)]
#[case::invalid_type_before_unexpected(
None,
"invalid type: x Unexpected server error",
Some(PublicError::BadRequest)
)]
#[case::internal_before_invalid_type(
None,
"internal server error invalid type: x",
Some(PublicError::InternalServer)
)]
#[case::invalid_type_with_a_status(Some(500), "invalid type: x", None)]
#[case::unexpected_with_a_status(Some(400), "Unexpected server error", None)]
fn the_trailing_rules_only_claim_failures_without_a_status(
#[case] status: Option<u16>,
#[case] text: &str,
#[case] expected: Option<PublicError>,
) {
assert_eq!(classified_with(status, text), expected);
}
#[test]
fn text_without_a_marker_is_left_to_the_status_table() {
assert_eq!(classified("rejected"), None);
}
}

View file

@ -0,0 +1,542 @@
//! A port of Python's `exception_type` for the routes that run in Rust. Rust decides the
//! public class, the message and the debug text; Python only builds the class.
//!
//! DIVERGENCES: where the Python mapper is inconsistent, the port follows one rule instead.
//! - The message is always `{Provider}Exception - {redacted text}`. Python's per-branch
//! labels (`RateLimitError: `, `litellm.RateLimitError: `, `Vertex_aiException BadRequestError`)
//! are dropped because every public class already prefixes `litellm.{Class}: `.
//! - The upstream response is always the real one. Python swaps in made-up `httpx.Response`
//! stubs on some Vertex branches, losing the body and `retry-after`.
//! - The debug text is always attached; Python passes it on some branches only.
//! - No family rule turns a status into a class; the shared status table owns that. So a
//! Vertex 502 is a `BadGatewayError` and an OpenAI-family 403 is a `PermissionDeniedError`.
//! Three rules read the status only to gate a text match, as Python does: the standalone
//! `429`, Vertex's wrapped 429 behind a 5xx, and Cohere's rules for failures with no status.
//! - A timeout text marker on an HTTP failure keeps the upstream response. Python's `Timeout`
//! carries none.
//! - Every family matches and reports the redacted text. Python's OpenAI mapper builds the
//! message from the unredacted text.
//! - A refused connection is an `APIConnectionError`, not the 500 Python's HTTP handler
//! synthesizes.
//! - Dropped Python rules: Vertex's bare `403` substring (it matches `4031 tokens`), Vertex's
//! `IndexError` quota marker (a Python client crash), the OpenAI SDK's missing-`api_key`
//! text and its `OPENAI` renaming, Cohere's `llm_provider="cohere"` override, and Cohere's
//! `CohereConnectionError` check (a Python SDK class name).
//!
//! KNOWN_GAPS: differences from the Python mapper that no Rust route can reach today. Each
//! one stops being acceptable at its trigger.
//! - The Vertex partner-model API base for "claude" models is not built into the debug text.
//! Trigger: a Vertex route whose models include Anthropic partner models.
//! - The debug text's `API Base` line is only the non-streaming Vertex URL. Python prefers an
//! explicit or provider-resolved `api_base`, uses `:streamGenerateContent` when streaming,
//! and has Gemini and OpenAI defaults. Trigger: the first route wired to this mapper, since
//! every route knows its `api_base`.
//! - The debug text has no `Messages:` line, which Python adds when
//! `redact_messages_in_exceptions` is off. Trigger: a wired route that carries messages.
//! - Python reports the provider `get_llm_provider` resolves for a stripped model name when
//! that name happens to be in the model cost map. Trigger: a route whose model names
//! overlap the cost map; that needs the provider resolution port, not a classifier change.
//! - `litellm_proxy` errors are not unwrapped into the proxied exception. Trigger: a Rust
//! route that calls a LiteLLM proxy.
//! - Only the OpenAI-compatible, Vertex AI and Cohere mappers are ported; every other
//! provider goes straight to the status table. Trigger: a Rust route for such a provider.
use super::secret_redaction::SecretRedactor;
mod cohere;
mod openai;
mod original;
mod public;
mod rules;
mod status;
mod vertex_ai;
pub use original::{ExceptionFamily, OriginalException};
pub use public::{MappedFailure, PublicError, UpstreamResponse};
use rules::{Rule, contains_any, first_match};
const TIMEOUT_MARKERS: &[&str] = &[
"Request Timeout Error",
"Request timed out",
"Timed out generating response",
"The read operation timed out",
];
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct ExceptionContext {
pub model: String,
pub custom_llm_provider: String,
pub asynchronous: bool,
pub vertex_project: Option<String>,
pub vertex_location: Option<String>,
pub model_group: Option<String>,
pub deployment: Option<String>,
pub user_api_key_alias: Option<String>,
pub user_api_key_team_alias: Option<String>,
}
/// What the rules read: the status of a provider response, if any, and the redacted text.
struct Mapping {
status: Option<u16>,
error_str: String,
}
pub fn exception_type(
context: &ExceptionContext,
redactor: Option<&SecretRedactor>,
original: &OriginalException,
) -> MappedFailure {
let (status, text, upstream) = match original {
OriginalException::Http {
status,
body,
headers,
} => (
Some(*status),
body.clone(),
Some(UpstreamResponse {
status: *status,
body: body.clone(),
headers: headers.clone(),
}),
),
OriginalException::Connection { message } | OriginalException::Plain { message } => {
(None, message.clone(), None)
}
OriginalException::Timeout {
timeout_seconds,
elapsed_seconds,
} => (
None,
timeout_message(context.asynchronous, *timeout_seconds, *elapsed_seconds),
None,
),
};
let mapping = Mapping {
status,
error_str: match redactor {
Some(redactor) => redactor.redact(&text),
None => text,
},
};
let family = ExceptionFamily::for_provider(&context.custom_llm_provider);
let (error, hint) = classify(family, original, &mapping);
MappedFailure {
error,
message: format!(
"{} - {}{hint}",
exception_provider(&context.custom_llm_provider),
mapping.error_str
),
upstream,
debug_info: extra_information(context, api_base(context).as_deref()),
}
}
fn classify(
family: ExceptionFamily,
original: &OriginalException,
mapping: &Mapping,
) -> (PublicError, &'static str) {
const TIMEOUT: PublicError = PublicError::Timeout { status: 408 };
if matches!(original, OriginalException::Timeout { .. })
|| contains_any(&mapping.error_str, TIMEOUT_MARKERS)
{
return (TIMEOUT, "");
}
if let Some(rule) = first_match(family_rules(family), mapping) {
return (rule.error, rule.hint);
}
let by_status = mapping.status.and_then(status::classify);
(by_status.unwrap_or(PublicError::ApiConnection), "")
}
fn family_rules(family: ExceptionFamily) -> &'static [Rule] {
match family {
ExceptionFamily::OpenAiCompatible => openai::RULES,
ExceptionFamily::VertexAi => vertex_ai::RULES,
ExceptionFamily::Cohere => cohere::RULES,
ExceptionFamily::Other => &[],
}
}
/// The text the Python HTTP handler's timeout carries: the sync and async handlers word it
/// differently.
fn timeout_message(
asynchronous: bool,
timeout_seconds: Option<f64>,
elapsed_seconds: Option<f64>,
) -> String {
let timeout = python_float(timeout_seconds);
if asynchronous {
let elapsed =
python_float(elapsed_seconds.map(|seconds| (seconds * 1000.0).round() / 1000.0));
format!("Connection timed out. Timeout passed={timeout}, time taken={elapsed} seconds")
} else {
format!("Connection timed out after {timeout} seconds.")
}
}
fn python_float(value: Option<f64>) -> String {
match value {
None => "None".to_string(),
Some(value) if value.fract() == 0.0 => format!("{value:.1}"),
Some(value) => value.to_string(),
}
}
fn exception_provider(provider: &str) -> String {
if provider == "openai" {
return "OpenAIException".to_string();
}
let mut characters = provider.chars();
match characters.next() {
Some(first) => format!("{}{}Exception", first.to_uppercase(), characters.as_str()),
None => String::new(),
}
}
fn api_base(context: &ExceptionContext) -> Option<String> {
match (&context.vertex_location, &context.vertex_project) {
(Some(location), Some(project)) => Some(format!(
"{location}-aiplatform.googleapis.com/v1/projects/{project}/locations/{location}/publishers/google/models/{}:generateContent",
context.model
)),
_ => None,
}
}
fn extra_information(context: &ExceptionContext, api_base: Option<&str>) -> String {
let lines = [
Some(format!("\nModel: {}", context.model)),
api_base.map(|api_base| format!("\nAPI Base: `{api_base}`")),
context
.model_group
.as_ref()
.map(|value| format!("\nmodel_group: `{value}`\n")),
context
.deployment
.as_ref()
.map(|value| format!("\ndeployment: `{value}`\n")),
context
.vertex_project
.as_ref()
.map(|value| format!("\nvertex_project: `{value}`\n")),
context
.vertex_location
.as_ref()
.map(|value| format!("\nvertex_location: `{value}`\n")),
];
let information: String = lines.into_iter().flatten().collect();
match &context.user_api_key_alias {
Some(alias) => format!(
"\n\nKey Name: `{alias}`\nTeam: `{}`{information}",
context.user_api_key_team_alias.as_deref().unwrap_or("None")
),
None => information,
}
}
#[cfg(test)]
mod testing {
use super::Mapping;
pub(super) fn mapping(status: Option<u16>, text: &str) -> Mapping {
Mapping {
status,
error_str: text.into(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
const DEBUG: &str = "\nModel: ocr-model";
fn context(provider: &str) -> ExceptionContext {
ExceptionContext {
model: "ocr-model".into(),
custom_llm_provider: provider.into(),
..ExceptionContext::default()
}
}
fn redactor() -> SecretRedactor {
SecretRedactor::new(16)
}
fn headers() -> Vec<(String, String)> {
vec![("retry-after".into(), "7".into())]
}
fn http(status: u16, body: &str) -> OriginalException {
OriginalException::Http {
status,
body: body.into(),
headers: headers(),
}
}
fn upstream(status: u16, body: &str) -> Option<UpstreamResponse> {
Some(UpstreamResponse {
status,
body: body.into(),
headers: headers(),
})
}
fn mapped(provider: &str, original: &OriginalException) -> MappedFailure {
exception_type(&context(provider), Some(&redactor()), original)
}
#[rstest::rstest]
#[case::openai_family("mistral", "rate limit reached", PublicError::RateLimit)]
#[case::vertex_family("vertex_ai", "Resource exhausted", PublicError::RateLimit)]
#[case::cohere_family("cohere", "too many tokens", PublicError::ContextWindowExceeded)]
fn a_family_text_rule_beats_the_status_and_keeps_the_real_response(
#[case] provider: &str,
#[case] body: &str,
#[case] expected: PublicError,
) {
let failure = mapped(provider, &http(401, body));
assert_eq!(failure.error, expected);
assert_eq!(failure.upstream, upstream(401, body));
}
#[test]
fn the_other_family_has_no_text_rules() {
assert_eq!(
mapped("reducto", &http(401, "rate limit reached")).error,
PublicError::Authentication
);
}
#[rstest::rstest]
#[case::openai_403_is_permission_denied("mistral", 403, PublicError::PermissionDenied)]
#[case::openai_409_is_bad_request("mistral", 409, PublicError::BadRequest)]
#[case::vertex_502_is_bad_gateway("vertex_ai", 502, PublicError::BadGateway)]
#[case::vertex_504_is_a_timeout("vertex_ai", 504, PublicError::Timeout { status: 504 })]
#[case::cohere_498_is_bad_request("cohere", 498, PublicError::BadRequest)]
#[case::other_503("reducto", 503, PublicError::ServiceUnavailable)]
fn without_a_text_rule_every_family_uses_the_status_table(
#[case] provider: &str,
#[case] status: u16,
#[case] expected: PublicError,
) {
assert_eq!(
mapped(provider, &http(status, "rejected")),
MappedFailure {
error: expected,
message: format!("{} - rejected", exception_provider(provider)),
upstream: upstream(status, "rejected"),
debug_info: DEBUG.into(),
}
);
}
#[rstest::rstest]
#[case::request_timeout_error("Request Timeout Error")]
#[case::request_timed_out("Request timed out")]
#[case::timed_out_generating("Timed out generating response")]
#[case::read_operation("The read operation timed out")]
fn timeout_markers_win_over_every_family(#[case] marker: &str) {
let body = format!("rate limit invalid api token {marker}");
for provider in ["mistral", "vertex_ai", "cohere", "reducto"] {
assert_eq!(
mapped(provider, &http(429, &body)).error,
PublicError::Timeout { status: 408 },
"{provider}"
);
}
}
#[test]
fn a_handler_timeout_is_a_408_without_a_response() {
let original = OriginalException::Timeout {
timeout_seconds: Some(0.5),
elapsed_seconds: Some(0.5031),
};
assert_eq!(
mapped("mistral", &original),
MappedFailure {
error: PublicError::Timeout { status: 408 },
message: "MistralException - Connection timed out after 0.5 seconds.".into(),
upstream: None,
debug_info: DEBUG.into(),
}
);
}
#[rstest::rstest]
#[case::refused_connection(OriginalException::Connection { message: "refused".into() })]
#[case::unparseable_response(OriginalException::Plain { message: "refused".into() })]
#[case::informational_status(OriginalException::Http { status: 399, body: "refused".into(), headers: Vec::new() })]
fn a_failure_no_rule_or_status_claims_is_a_connection_error(
#[case] original: OriginalException,
) {
let failure = mapped("reducto", &original);
assert_eq!(failure.error, PublicError::ApiConnection);
assert_eq!(failure.message, "ReductoException - refused");
}
#[test]
fn a_timeout_marker_on_a_response_keeps_the_response() {
let failure = mapped("reducto", &http(429, "Request timed out"));
assert_eq!(failure.error, PublicError::Timeout { status: 408 });
assert_eq!(failure.upstream, upstream(429, "Request timed out"));
}
#[test]
fn family_text_rules_also_classify_failures_without_a_response() {
let original = OriginalException::Plain {
message: "Request too large".into(),
};
assert_eq!(mapped("mistral", &original).error, PublicError::RateLimit);
}
#[rstest::rstest]
#[case::openai_family("mistral", "MistralException - rejected REDACTED")]
#[case::vertex_family("vertex_ai", "Vertex_aiException - rejected REDACTED")]
#[case::other_family("reducto", "ReductoException - rejected REDACTED")]
fn every_family_reports_the_redacted_text(#[case] provider: &str, #[case] message: &str) {
let failure = mapped(provider, &http(400, "rejected Bearer abcdefghijklmnop"));
assert_eq!(failure.message, message);
}
#[test]
fn redaction_runs_before_the_rules_see_the_text() {
let body = "db_password=rate_limit";
assert_eq!(
mapped("mistral", &http(400, body)).error,
PublicError::BadRequest
);
assert_eq!(
exception_type(&context("mistral"), None, &http(400, body)).error,
PublicError::RateLimit
);
}
#[test]
fn without_a_redactor_the_text_is_kept() {
let body = "rejected Bearer abcdefghijklmnop";
assert_eq!(
exception_type(&context("reducto"), None, &http(400, body)).message,
format!("ReductoException - {body}")
);
}
#[test]
fn a_rule_hint_follows_the_message() {
let failure = mapped("mistral", &http(400, "invalid_encrypted_content"));
assert_eq!(failure.error, PublicError::BadRequest);
assert!(
failure
.message
.starts_with("MistralException - invalid_encrypted_content\n\n This error occurs")
);
}
#[rstest::rstest]
#[case::sync(
false,
Some(0.5),
Some(0.5031),
"Connection timed out after 0.5 seconds."
)]
#[case::async_rounds_the_elapsed_time(
true,
Some(0.5),
Some(0.5031),
"Connection timed out. Timeout passed=0.5, time taken=0.503 seconds"
)]
#[case::whole_seconds_keep_a_decimal(
true,
Some(600.0),
Some(2.0),
"Connection timed out. Timeout passed=600.0, time taken=2.0 seconds"
)]
#[case::unknown_values_render_as_none(
true,
None,
None,
"Connection timed out. Timeout passed=None, time taken=None seconds"
)]
fn timeout_text_follows_the_delivery_mode(
#[case] asynchronous: bool,
#[case] timeout_seconds: Option<f64>,
#[case] elapsed_seconds: Option<f64>,
#[case] expected: &str,
) {
assert_eq!(
timeout_message(asynchronous, timeout_seconds, elapsed_seconds),
expected
);
}
#[test]
fn debug_information_follows_the_python_layout() {
let context = ExceptionContext {
vertex_project: Some("project".into()),
vertex_location: Some("region".into()),
model_group: Some("ocr".into()),
deployment: Some("deployment".into()),
user_api_key_alias: Some("key".into()),
..context("vertex_ai")
};
assert_eq!(
exception_type(&context, None, &http(400, "rejected")).debug_info,
concat!(
"\n\nKey Name: `key`\nTeam: `None`",
"\nModel: ocr-model",
"\nAPI Base: `region-aiplatform.googleapis.com/v1/projects/project/locations/region/publishers/google/models/ocr-model:generateContent`",
"\nmodel_group: `ocr`\n",
"\ndeployment: `deployment`\n",
"\nvertex_project: `project`\n",
"\nvertex_location: `region`\n",
)
);
}
#[rstest::rstest]
#[case::bare(ExceptionContext::default(), "\nModel: ")]
#[case::team_alias(
ExceptionContext { model: "m".into(), user_api_key_alias: Some("key".into()), user_api_key_team_alias: Some("team".into()), ..ExceptionContext::default() },
"\n\nKey Name: `key`\nTeam: `team`\nModel: m"
)]
#[case::team_alias_without_key_is_ignored(
ExceptionContext { model: "m".into(), user_api_key_team_alias: Some("team".into()), ..ExceptionContext::default() },
"\nModel: m"
)]
#[case::project_without_location_has_no_api_base(
ExceptionContext { model: "m".into(), vertex_project: Some("p".into()), ..ExceptionContext::default() },
"\nModel: m\nvertex_project: `p`\n"
)]
#[case::location_without_project_has_no_api_base(
ExceptionContext { model: "m".into(), vertex_location: Some("l".into()), ..ExceptionContext::default() },
"\nModel: m\nvertex_location: `l`\n"
)]
fn each_optional_context_field_adds_its_own_line(
#[case] context: ExceptionContext,
#[case] expected: &str,
) {
assert_eq!(
extra_information(&context, api_base(&context).as_deref()),
expected
);
}
#[rstest::rstest]
#[case::openai_keeps_its_brand("openai", "OpenAIException")]
#[case::lowercase("mistral", "MistralException")]
#[case::keeps_the_rest("azure_ai", "Azure_aiException")]
#[case::empty("", "")]
fn exception_provider_capitalizes_only_the_first_letter(
#[case] provider: &str,
#[case] expected: &str,
) {
assert_eq!(exception_provider(provider), expected);
}
}

View file

@ -0,0 +1,192 @@
use super::public::PublicError;
use super::rules::{Rule, contains_any, is_context_window_exceeded, is_rate_limit};
const ENCRYPTED_CONTENT_HELP: &str = "\n\n This error occurs when load balancing Responses API across deployments with different API keys.\n Encrypted content is tied to the organization that created it and cannot be decrypted by other organizations.\n\n Solution: Enable 'encrypted_content_affinity' to route follow-up requests to the correct deployment:\n\n router_settings:\n enable_pre_call_checks: true\n optional_pre_call_checks:\n - encrypted_content_affinity\n\n Learn more: https://docs.litellm.ai/docs/response_api#encrypted-content-affinity-multi-region-load-balancing";
/// The text branches of `_map_openai_exception`, in its order.
pub(super) const RULES: &[Rule] = &[
Rule::new(
|mapping| is_rate_limit(&mapping.error_str, mapping.status),
PublicError::RateLimit,
),
Rule::new(
|mapping| is_context_window_exceeded(&mapping.error_str),
PublicError::ContextWindowExceeded,
),
Rule::new(
|mapping| {
mapping.error_str.contains("invalid_request_error")
&& mapping.error_str.contains("model_not_found")
},
PublicError::NotFound,
),
Rule::new(
|mapping| mapping.error_str.contains("A timeout occurred"),
PublicError::Timeout { status: 408 },
),
Rule::new(
|mapping| {
let error_str = &mapping.error_str;
(error_str.contains("invalid_request_error")
&& error_str.contains("content_policy_violation"))
|| (error_str.contains("Invalid prompt")
&& error_str.contains("violating our usage policy"))
|| error_str
.to_lowercase()
.contains("request was rejected as a result of the safety system")
},
PublicError::ContentPolicyViolation,
),
Rule {
hint: ENCRYPTED_CONTENT_HELP,
..Rule::new(
|mapping| {
contains_any(
&mapping.error_str,
&["invalid_encrypted_content", "could not be verified"],
)
},
PublicError::BadRequest,
)
},
Rule::new(
|mapping| {
mapping.error_str.contains("invalid_request_error")
&& !mapping.error_str.contains("Incorrect API key provided")
},
PublicError::BadRequest,
),
Rule::new(
|mapping| {
contains_any(
&mapping.error_str,
&[
"Web server is returning an unknown error",
"The server had an error processing your request.",
],
)
},
PublicError::InternalServer,
),
Rule::new(
|mapping| mapping.error_str.contains("Request too large"),
PublicError::RateLimit,
),
Rule::new(
|mapping| {
mapping
.error_str
.contains("Mistral API raised a streaming error")
},
PublicError::Api { status: 500 },
),
];
#[cfg(test)]
mod tests {
use super::super::rules::first_match;
use super::super::testing::mapping;
use super::*;
fn classified(status: Option<u16>, text: &str) -> Option<PublicError> {
first_match(RULES, &mapping(status, text)).map(|rule| rule.error)
}
#[rstest::rstest]
#[case::rate_limit_phrase("rate limit reached", PublicError::RateLimit)]
#[case::context_window(
"This model's maximum context length is 10",
PublicError::ContextWindowExceeded
)]
#[case::model_not_found("invalid_request_error model_not_found", PublicError::NotFound)]
#[case::timeout_occurred("A timeout occurred", PublicError::Timeout { status: 408 })]
#[case::content_policy_error_code(
"invalid_request_error content_policy_violation",
PublicError::ContentPolicyViolation
)]
#[case::content_policy_usage_policy(
"Invalid prompt violating our usage policy",
PublicError::ContentPolicyViolation
)]
#[case::content_policy_safety_system(
"Request was rejected as a result of the safety system",
PublicError::ContentPolicyViolation
)]
#[case::encrypted_content("invalid_encrypted_content", PublicError::BadRequest)]
#[case::unverifiable_content("could not be verified", PublicError::BadRequest)]
#[case::invalid_request("invalid_request_error bad field", PublicError::BadRequest)]
#[case::unknown_server_error(
"Web server is returning an unknown error",
PublicError::InternalServer
)]
#[case::server_had_an_error(
"The server had an error processing your request.",
PublicError::InternalServer
)]
#[case::request_too_large("Request too large", PublicError::RateLimit)]
#[case::mistral_streaming_error(
"Mistral API raised a streaming error",
PublicError::Api { status: 500 }
)]
fn each_text_rule_claims_its_marker(#[case] text: &str, #[case] expected: PublicError) {
assert_eq!(classified(Some(400), text), Some(expected));
}
#[rstest::rstest]
#[case::rate_limit_before_context_window(
"rate limit and This model's maximum context length is 10",
PublicError::RateLimit
)]
#[case::context_window_before_content_policy(
"This model's maximum context length is 10 invalid_request_error content_policy_violation",
PublicError::ContextWindowExceeded
)]
#[case::model_not_found_before_invalid_request(
"invalid_request_error model_not_found",
PublicError::NotFound
)]
#[case::timeout_before_invalid_request(
"A timeout occurred invalid_request_error",
PublicError::Timeout { status: 408 }
)]
#[case::content_policy_before_invalid_request(
"invalid_request_error content_policy_violation",
PublicError::ContentPolicyViolation
)]
#[case::encrypted_content_before_invalid_request(
"invalid_request_error invalid_encrypted_content",
PublicError::BadRequest
)]
fn the_earlier_rule_wins_when_two_apply(#[case] text: &str, #[case] expected: PublicError) {
assert_eq!(classified(Some(400), text), Some(expected));
}
#[rstest::rstest]
#[case::encrypted_content("invalid_encrypted_content", ENCRYPTED_CONTENT_HELP)]
#[case::plain_invalid_request("invalid_request_error bad field", "")]
fn only_encrypted_content_failures_carry_the_affinity_help(
#[case] text: &str,
#[case] hint: &str,
) {
assert_eq!(
first_match(RULES, &mapping(Some(400), text)).map(|rule| rule.hint),
Some(hint)
);
}
#[rstest::rstest]
#[case::bad_key_is_left_to_the_status("invalid_request_error Incorrect API key provided")]
#[case::echoed_429_is_not_a_rate_limit("token 429 in the prompt")]
#[case::unmarked("rejected")]
fn text_without_a_marker_is_left_to_the_status_table(#[case] text: &str) {
assert_eq!(classified(Some(400), text), None);
}
#[test]
fn a_standalone_429_counts_with_a_429_status() {
assert_eq!(
classified(Some(429), "got 429 back"),
Some(PublicError::RateLimit)
);
}
}

View file

@ -0,0 +1,144 @@
/// A failure a Rust route produced, before any public class is chosen.
#[derive(Clone, Debug, PartialEq)]
pub enum OriginalException {
Http {
status: u16,
body: String,
headers: Vec<(String, String)>,
},
Connection {
message: String,
},
Timeout {
timeout_seconds: Option<f64>,
elapsed_seconds: Option<f64>,
},
/// A failure with no HTTP response behind it, such as an unparseable body or a local
/// file error.
Plain {
message: String,
},
}
/// Which provider-specific text rules apply before the shared status table.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum ExceptionFamily {
OpenAiCompatible,
VertexAi,
Cohere,
Other,
}
/// `openai_compatible_providers` in `litellm/constants.py`.
const OPENAI_COMPATIBLE_PROVIDERS: &[&str] = &[
"anyscale",
"groq",
"nvidia_nim",
"cerebras",
"baseten",
"sambanova",
"ai21_chat",
"ai21",
"volcengine",
"codestral",
"deepseek",
"tencent",
"deepinfra",
"perplexity",
"xinference",
"xai",
"zai",
"together_ai",
"fireworks_ai",
"empower",
"friendliai",
"azure_ai",
"github",
"litellm_proxy",
"hosted_vllm",
"llamafile",
"lm_studio",
"galadriel",
"github_copilot",
"chatgpt",
"novita",
"meta_llama",
"publicai",
"synthetic",
"tensormesh",
"apertis",
"nano-gpt",
"poe",
"chutes",
"parasail",
"libertai",
"featherless_ai",
"nscale",
"nebius",
"dashscope",
"qwencloud",
"qwen_ai_platform",
"modelscope",
"moonshot",
"v0",
"helicone",
"morph",
"lambda_ai",
"inception",
"hyperbolic",
"vercel_ai_gateway",
"aiml",
"wandb",
"cometapi",
"clarifai",
"docker_model_runner",
"ragflow",
"pinstripes",
"darkbloom",
"meta",
"cognition",
"scx-ai",
];
impl ExceptionFamily {
/// The provider dispatch at the top of Python's `exception_type`, in its order.
pub fn for_provider(provider: &str) -> Self {
match provider {
"openai" | "text-completion-openai" | "custom_openai" | "mistral" | "runwayml" => {
Self::OpenAiCompatible
}
provider if OPENAI_COMPATIBLE_PROVIDERS.contains(&provider) => Self::OpenAiCompatible,
"vertex_ai" | "vertex_ai_beta" | "gemini" => Self::VertexAi,
"cohere" | "cohere_chat" => Self::Cohere,
_ => Self::Other,
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[rstest::rstest]
#[case::openai("openai", ExceptionFamily::OpenAiCompatible)]
#[case::text_completion_openai("text-completion-openai", ExceptionFamily::OpenAiCompatible)]
#[case::custom_openai("custom_openai", ExceptionFamily::OpenAiCompatible)]
#[case::mistral("mistral", ExceptionFamily::OpenAiCompatible)]
#[case::runwayml("runwayml", ExceptionFamily::OpenAiCompatible)]
#[case::listed_compatible("azure_ai", ExceptionFamily::OpenAiCompatible)]
#[case::compatible_list_wins_over_its_own_mapper(
"together_ai",
ExceptionFamily::OpenAiCompatible
)]
#[case::vertex_ai("vertex_ai", ExceptionFamily::VertexAi)]
#[case::vertex_ai_beta("vertex_ai_beta", ExceptionFamily::VertexAi)]
#[case::gemini("gemini", ExceptionFamily::VertexAi)]
#[case::cohere("cohere", ExceptionFamily::Cohere)]
#[case::cohere_chat("cohere_chat", ExceptionFamily::Cohere)]
#[case::unported_mapper("anthropic", ExceptionFamily::Other)]
#[case::unknown("reducto", ExceptionFamily::Other)]
#[case::empty("", ExceptionFamily::Other)]
fn provider_selects_the_family(#[case] provider: &str, #[case] family: ExceptionFamily) {
assert_eq!(ExceptionFamily::for_provider(provider), family);
}
}

View file

@ -0,0 +1,77 @@
/// The public LiteLLM exception classes a Rust route failure can become. Python builds the
/// class; Rust decides which one.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum PublicError {
BadRequest,
ContextWindowExceeded,
ContentPolicyViolation,
Authentication,
PermissionDenied,
NotFound,
Timeout { status: u16 },
RateLimit,
InternalServer,
BadGateway,
ServiceUnavailable,
ApiConnection,
Api { status: u16 },
}
impl PublicError {
/// The `status_code` the Python class carries.
pub const fn status_code(self) -> u16 {
match self {
Self::BadRequest | Self::ContextWindowExceeded | Self::ContentPolicyViolation => 400,
Self::Authentication => 401,
Self::PermissionDenied => 403,
Self::NotFound => 404,
Self::RateLimit => 429,
Self::InternalServer | Self::ApiConnection => 500,
Self::BadGateway => 502,
Self::ServiceUnavailable => 503,
Self::Timeout { status } | Self::Api { status } => status,
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct UpstreamResponse {
pub status: u16,
pub body: String,
pub headers: Vec<(String, String)>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct MappedFailure {
pub error: PublicError,
pub message: String,
pub upstream: Option<UpstreamResponse>,
pub debug_info: String,
}
#[cfg(test)]
mod tests {
use super::*;
#[rstest::rstest]
#[case::bad_request(PublicError::BadRequest, 400)]
#[case::context_window(PublicError::ContextWindowExceeded, 400)]
#[case::content_policy(PublicError::ContentPolicyViolation, 400)]
#[case::authentication(PublicError::Authentication, 401)]
#[case::permission_denied(PublicError::PermissionDenied, 403)]
#[case::not_found(PublicError::NotFound, 404)]
#[case::request_timeout(PublicError::Timeout { status: 408 }, 408)]
#[case::gateway_timeout(PublicError::Timeout { status: 504 }, 504)]
#[case::rate_limit(PublicError::RateLimit, 429)]
#[case::internal_server(PublicError::InternalServer, 500)]
#[case::api_connection(PublicError::ApiConnection, 500)]
#[case::bad_gateway(PublicError::BadGateway, 502)]
#[case::service_unavailable(PublicError::ServiceUnavailable, 503)]
#[case::api(PublicError::Api { status: 501 }, 501)]
fn status_codes_are_the_ones_the_python_classes_set(
#[case] error: PublicError,
#[case] status: u16,
) {
assert_eq!(error.status_code(), status);
}
}

View file

@ -0,0 +1,176 @@
use std::sync::LazyLock;
use fancy_regex::Regex;
use serde_json::Value;
use super::Mapping;
use super::public::PublicError;
/// One text branch of a Python `_map_*_exception` function: when it applies, the class it
/// raises, and any help text appended to the message.
pub(super) struct Rule {
pub(super) when: fn(&Mapping) -> bool,
pub(super) error: PublicError,
pub(super) hint: &'static str,
}
impl Rule {
pub(super) const fn new(when: fn(&Mapping) -> bool, error: PublicError) -> Self {
Self {
when,
error,
hint: "",
}
}
}
/// The first rule that applies decides the class, as the `if`/`elif` chain does in Python.
pub(super) fn first_match<'r>(rules: &'r [Rule], mapping: &Mapping) -> Option<&'r Rule> {
rules.iter().find(|rule| (rule.when)(mapping))
}
pub(super) fn contains_any(text: &str, markers: &[&str]) -> bool {
markers.iter().any(|marker| text.contains(marker))
}
static STANDALONE_429: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"\b429\b").expect("valid regex"));
static RATE_LIMIT_PHRASE: LazyLock<Regex> =
LazyLock::new(|| Regex::new(r"rate[\s_\-]*limit").expect("valid regex"));
/// `ExceptionCheckers.is_error_str_rate_limit`.
pub(super) fn is_rate_limit(error_str: &str, status: Option<u16>) -> bool {
if STANDALONE_429.is_match(error_str).unwrap_or(false) && matches!(status, None | Some(429)) {
return true;
}
let lower = error_str.to_lowercase();
RATE_LIMIT_PHRASE.is_match(&lower).unwrap_or(false)
|| lower.contains("service tier capacity exceeded")
}
/// `ExceptionCheckers.is_error_str_context_window_exceeded`.
pub(super) fn is_context_window_exceeded(error_str: &str) -> bool {
let lower = error_str.to_lowercase();
if lower.contains("string_above_max_length") {
return false;
}
if lower.contains("invalid 'user'") && lower.contains("string too long") {
return false;
}
contains_any(
&lower,
&[
"exceed context limit",
"this model's maximum context length is",
"string too long. expected a string with maximum length",
"model's maximum context limit",
"is longer than the model's context length",
"input tokens exceed the configured limit",
"`inputs` tokens + `max_new_tokens` must be",
"exceeds the available context size",
"exceeds the maximum number of tokens allowed",
],
) || (lower.contains("current length is") && lower.contains("while limit is"))
|| (lower.contains("maximum input length is") && lower.contains("tokens"))
}
/// The integer `error.code` of a JSON error body, read the way Python's `int()` would.
pub(super) fn body_error_code(error_str: &str) -> Option<i64> {
let body: Value = serde_json::from_str(error_str).ok()?;
let Some(Value::Object(error)) = body.as_object()?.get("error") else {
return None;
};
match error.get("code")? {
Value::Number(number) => number
.as_i64()
.or_else(|| number.as_f64().map(|value| value.trunc() as i64)),
Value::String(code) => code.trim().replace('_', "").parse().ok(),
Value::Bool(flag) => Some(i64::from(*flag)),
_ => None,
}
}
#[cfg(test)]
mod tests {
use super::super::testing::mapping;
use super::*;
const ORDERED: &[Rule] = &[
Rule::new(
|mapping| mapping.error_str.contains("first"),
PublicError::NotFound,
),
Rule::new(|_| true, PublicError::ApiConnection),
];
#[rstest::rstest]
#[case::earlier_rule_wins("first and second", PublicError::NotFound)]
#[case::later_rule_when_the_earlier_does_not_apply("second", PublicError::ApiConnection)]
fn the_first_applicable_rule_decides(#[case] text: &str, #[case] expected: PublicError) {
let rule = first_match(ORDERED, &mapping(Some(400), text));
assert_eq!(rule.map(|rule| rule.error), Some(expected));
}
#[test]
fn no_applicable_rule_leaves_the_failure_to_the_caller() {
assert!(first_match(&ORDERED[..1], &mapping(Some(400), "second")).is_none());
}
#[rstest::rstest]
#[case::standalone_429_with_429_status("got 429 back", Some(429), true)]
#[case::standalone_429_with_other_status("got 429 back", Some(400), false)]
#[case::standalone_429_with_unknown_status("got 429 back", None, true)]
#[case::embedded_429("token4290", Some(429), false)]
#[case::phrase_spaced("Rate Limit reached", None, true)]
#[case::phrase_underscored("rate_limit", None, true)]
#[case::phrase_hyphenated("rate-limit", None, true)]
#[case::service_tier("Service tier capacity exceeded", None, true)]
#[case::unrelated("rejected", Some(429), false)]
fn rate_limit_detection(
#[case] text: &str,
#[case] status: Option<u16>,
#[case] expected: bool,
) {
assert_eq!(is_rate_limit(text, status), expected);
}
#[rstest::rstest]
#[case::exceed_context_limit("Exceed context limit", true)]
#[case::maximum_context_length("This model's maximum context length is 10", true)]
#[case::string_too_long("string too long. Expected a string with maximum length 5", true)]
#[case::maximum_context_limit("the model's maximum context limit", true)]
#[case::longer_than_context("prompt is longer than the model's context length", true)]
#[case::configured_limit("input tokens exceed the configured limit", true)]
#[case::max_new_tokens("`inputs` tokens + `max_new_tokens` must be <= 10", true)]
#[case::available_context("exceeds the available context size", true)]
#[case::maximum_tokens("exceeds the maximum number of tokens allowed", true)]
#[case::current_and_limit("current length is 9 while limit is 8", true)]
#[case::current_without_limit("current length is 9", false)]
#[case::maximum_input_tokens("maximum input length is 8 tokens", true)]
#[case::maximum_input_without_tokens("maximum input length is 8", false)]
#[case::string_above_max_length_wins("string_above_max_length exceed context limit", false)]
#[case::user_field_is_not_context(
"invalid 'user': string too long. expected a string with maximum length",
false
)]
#[case::unrelated("rejected", false)]
fn context_window_detection(#[case] text: &str, #[case] expected: bool) {
assert_eq!(is_context_window_exceeded(text), expected);
}
#[rstest::rstest]
#[case::integer(r#"{"error": {"code": 429}}"#, Some(429))]
#[case::float(r#"{"error": {"code": 429.9}}"#, Some(429))]
#[case::string(r#"{"error": {"code": " 4_29 "}}"#, Some(429))]
#[case::boolean(r#"{"error": {"code": true}}"#, Some(1))]
#[case::unparseable_string(r#"{"error": {"code": "slow"}}"#, None)]
#[case::null(r#"{"error": {"code": null}}"#, None)]
#[case::no_code(r#"{"error": {}}"#, None)]
#[case::error_not_an_object(r#"{"error": "429"}"#, None)]
#[case::no_error(r#"{"code": 429}"#, None)]
#[case::not_an_object("[429]", None)]
#[case::not_json("429", None)]
fn body_error_code_reads_the_nested_code(#[case] body: &str, #[case] expected: Option<i64>) {
assert_eq!(body_error_code(body), expected);
}
}

View file

@ -0,0 +1,49 @@
use super::public::PublicError;
/// `_map_exception_by_status`, the one place a provider status picks a class. Statuses
/// below 400 are not failures the table claims.
pub(super) fn classify(status: u16) -> Option<PublicError> {
let error = match status {
..400 => return None,
401 => PublicError::Authentication,
403 => PublicError::PermissionDenied,
404 => PublicError::NotFound,
408 | 504 => PublicError::Timeout { status },
429 => PublicError::RateLimit,
500 => PublicError::InternalServer,
502 => PublicError::BadGateway,
503 => PublicError::ServiceUnavailable,
400..500 => PublicError::BadRequest,
_ => PublicError::Api { status },
};
Some(error)
}
#[cfg(test)]
mod tests {
use super::*;
#[rstest::rstest]
#[case::below_client_errors(399, None)]
#[case::lowest_client_error(400, Some(PublicError::BadRequest))]
#[case::authentication(401, Some(PublicError::Authentication))]
#[case::permission_denied(403, Some(PublicError::PermissionDenied))]
#[case::not_found(404, Some(PublicError::NotFound))]
#[case::request_timeout(408, Some(PublicError::Timeout { status: 408 }))]
#[case::other_client_error(409, Some(PublicError::BadRequest))]
#[case::unprocessable(422, Some(PublicError::BadRequest))]
#[case::rate_limited(429, Some(PublicError::RateLimit))]
#[case::highest_client_error(499, Some(PublicError::BadRequest))]
#[case::internal_server(500, Some(PublicError::InternalServer))]
#[case::other_server_error(501, Some(PublicError::Api { status: 501 }))]
#[case::bad_gateway(502, Some(PublicError::BadGateway))]
#[case::service_unavailable(503, Some(PublicError::ServiceUnavailable))]
#[case::gateway_timeout(504, Some(PublicError::Timeout { status: 504 }))]
#[case::highest_server_error(599, Some(PublicError::Api { status: 599 }))]
fn every_mapped_status_and_the_fallback(
#[case] status: u16,
#[case] expected: Option<PublicError>,
) {
assert_eq!(classify(status), expected);
}
}

View file

@ -0,0 +1,177 @@
use super::public::PublicError;
use super::rules::{Rule, body_error_code, contains_any, is_context_window_exceeded};
const QUOTA_MARKERS: &[&str] = &[
"429 Quota exceeded",
"Quota exceeded for",
"Resource exhausted",
"429 Unable to submit request because the service is temporarily out of capacity.",
];
/// The text branches of `_map_vertex_exception`, in its order.
pub(super) const RULES: &[Rule] = &[
Rule::new(
|mapping| {
contains_any(
&mapping.error_str,
&[
"Vertex AI API has not been used in project",
"Unable to find your project",
],
)
},
PublicError::BadRequest,
),
Rule::new(
|mapping| {
mapping
.error_str
.contains("400 Request payload size exceeds")
|| is_context_window_exceeded(&mapping.error_str)
},
PublicError::ContextWindowExceeded,
),
Rule::new(
|mapping| {
contains_any(
&mapping.error_str,
&["None Unknown Error.", "Content has no parts."],
)
},
PublicError::InternalServer,
),
Rule::new(
|mapping| mapping.error_str.contains("API key not valid."),
PublicError::Authentication,
),
Rule::new(
|mapping| {
contains_any(
&mapping.error_str,
&[
"The response was blocked.",
"Output blocked by content filtering policy",
],
)
},
PublicError::ContentPolicyViolation,
),
Rule::new(
|mapping| {
contains_any(&mapping.error_str, QUOTA_MARKERS)
|| (mapping
.status
.is_some_and(|status| (500..600).contains(&status))
&& body_error_code(&mapping.error_str) == Some(429))
},
PublicError::RateLimit,
),
Rule::new(
|mapping| {
contains_any(
&mapping.error_str,
&["500 Internal Server Error", "The model is overloaded."],
)
},
PublicError::InternalServer,
),
];
#[cfg(test)]
mod tests {
use super::super::rules::first_match;
use super::super::testing::mapping;
use super::*;
fn classified(status: Option<u16>, text: &str) -> Option<PublicError> {
first_match(RULES, &mapping(status, text)).map(|rule| rule.error)
}
#[rstest::rstest]
#[case::api_not_enabled(
"Vertex AI API has not been used in project x",
PublicError::BadRequest
)]
#[case::project_not_found("Unable to find your project", PublicError::BadRequest)]
#[case::payload_too_large(
"400 Request payload size exceeds the limit",
PublicError::ContextWindowExceeded
)]
#[case::context_window(
"This model's maximum context length is 10",
PublicError::ContextWindowExceeded
)]
#[case::unknown_error("None Unknown Error.", PublicError::InternalServer)]
#[case::no_parts("Content has no parts.", PublicError::InternalServer)]
#[case::api_key_not_valid("API key not valid.", PublicError::Authentication)]
#[case::response_blocked("The response was blocked.", PublicError::ContentPolicyViolation)]
#[case::output_blocked(
"Output blocked by content filtering policy",
PublicError::ContentPolicyViolation
)]
#[case::quota_exceeded_429("429 Quota exceeded", PublicError::RateLimit)]
#[case::quota_exceeded_for("Quota exceeded for aiplatform", PublicError::RateLimit)]
#[case::resource_exhausted("Resource exhausted", PublicError::RateLimit)]
#[case::out_of_capacity(
"429 Unable to submit request because the service is temporarily out of capacity.",
PublicError::RateLimit
)]
#[case::internal_server_text("500 Internal Server Error", PublicError::InternalServer)]
#[case::overloaded("The model is overloaded.", PublicError::InternalServer)]
fn each_text_rule_claims_its_marker(#[case] text: &str, #[case] expected: PublicError) {
assert_eq!(classified(Some(400), text), Some(expected));
}
#[rstest::rstest]
#[case::server_error_wrapping_a_429(Some(503), Some(PublicError::RateLimit))]
#[case::lowest_server_error(Some(500), Some(PublicError::RateLimit))]
#[case::highest_server_error(Some(599), Some(PublicError::RateLimit))]
#[case::client_error(Some(400), None)]
#[case::no_status(None, None)]
fn a_wrapped_429_is_a_rate_limit_only_behind_a_server_error(
#[case] status: Option<u16>,
#[case] expected: Option<PublicError>,
) {
assert_eq!(
classified(status, r#"{"error": {"code": "429"}}"#),
expected
);
}
#[rstest::rstest]
#[case::project_before_payload_size(
"Unable to find your project 400 Request payload size exceeds",
PublicError::BadRequest
)]
#[case::context_window_before_unknown_error(
"This model's maximum context length is 10 None Unknown Error.",
PublicError::ContextWindowExceeded
)]
#[case::unknown_error_before_api_key(
"Content has no parts. API key not valid.",
PublicError::InternalServer
)]
#[case::api_key_before_blocked(
"API key not valid. The response was blocked.",
PublicError::Authentication
)]
#[case::blocked_before_quota(
"The response was blocked. Resource exhausted",
PublicError::ContentPolicyViolation
)]
#[case::quota_before_overloaded(
"Resource exhausted The model is overloaded.",
PublicError::RateLimit
)]
fn the_earlier_rule_wins_when_two_apply(#[case] text: &str, #[case] expected: PublicError) {
assert_eq!(classified(Some(400), text), Some(expected));
}
#[rstest::rstest]
#[case::a_403_in_the_text("got a 403 from 4031 tokens")]
#[case::python_client_crash("IndexError: list index out of range")]
#[case::unmarked("rejected")]
fn text_without_a_marker_is_left_to_the_status_table(#[case] text: &str) {
assert_eq!(classified(Some(400), text), None);
}
}

View file

@ -1,7 +1,9 @@
pub mod call_arguments;
pub mod core_helpers;
pub mod exception_mapping_utils;
pub mod get_llm_provider_logic;
pub mod params;
pub mod prompt_templates;
pub mod secret_redaction;
pub mod serde_compat;
pub mod url_utils;

View file

@ -0,0 +1,109 @@
use fancy_regex::Regex;
pub const REDACTED: &str = "REDACTED";
const DEFAULT_MINIMUM_CUSTOM_KEY_LENGTH: usize = 16;
fn minimum_custom_key_length() -> usize {
std::env::var("MINIMUM_CUSTOM_KEY_LENGTH")
.ok()
.and_then(|value| value.trim().parse().ok())
.unwrap_or(DEFAULT_MINIMUM_CUSTOM_KEY_LENGTH)
}
fn secret_patterns(minimum_custom_key_length: usize) -> String {
let sk_suffix_length = minimum_custom_key_length.saturating_sub("sk-".len());
[
r"-----BEGIN[A-Z \-]*PRIVATE KEY-----[\s\S]*?-----END[A-Z \-]*PRIVATE KEY-----",
r"\bya29\.[A-Za-z0-9_.~+/-]+",
r#"(?:client_secret|azure_password|azure_username)\s+[^\s,'"})\]{}>]+"#,
r"(?:AKIA|ASIA)[0-9A-Z]{16}",
r"Bearer\s+[A-Za-z0-9\-._~+/]{10,}=*",
r"Basic\s+[A-Za-z0-9+/]{10,}={0,2}",
&format!(r"sk-[A-Za-z0-9\-_]{{{sk_suffix_length},}}"),
r#"(?<=[?&])(?:api[_-]?key|\w*(?:token|password|passwd|client_secret|secret_key|_secret))=[^\s&'"]+"#,
r#"(?:api[_-]?key)['"]?\s*[:=]\s*['"]?[^\s,'"})\]{}>]{8,}"#,
r#"(?:x-api-key|api-key)['"]?\s*[:=]\s*['"]?[^\s,'"})\]{}>]+"#,
r"x-ak-[A-Za-z0-9\-_]{20,}",
r"AIza[0-9A-Za-z\-_]{35}",
r#"(?<=[?&])key=[^\s&'"]{8,}"#,
r#"(?:^|(?<=\W))\w*(?:password|passwd|client_secret|secret_key|_secret)['"]?\s*[:=]\s*['"]?[^\s,'"})\]{}>]+"#,
r#"(?<=://)[^\s'":]{0,4096}:[^\s'"]{1,4096}(?=@)"#,
r"dapi[0-9a-f]{32}",
r#"litellm\.[A-Za-z0-9_]*_key['"]?\s*[:=]\s*['"]?[^\s,'"})\]{}>]+"#,
r#"private_key['"]?\s*[:=]\s*['"]?(?:-----BEGIN[A-Z \-]*PRIVATE KEY-----[\s\S]*?-----END[A-Z \-]*PRIVATE KEY-----|[^\s,'"})\]{}>]+)"#,
concat!(
r"(?:master_key|xai_key|database_url|db_url|connection_string|",
r"aws_secret_access_key|aws_session_token|aws_access_key_id|",
r"signing_key|encryption_key|",
r"auth_token|access_token|refresh_token|",
r"slack_webhook_url|webhook_url|",
r"database_connection_string|",
r"huggingface_token|jwt_secret)",
r#"['"]?\s*[:=]\s*['"]?[^\s,'"})\]{}>]+"#,
),
r"\beyJ[A-Za-z0-9_-]{10,}\.[A-Za-z0-9_-]+\.[A-Za-z0-9_-]*",
r"(?<=[?&])sig=[A-Za-z0-9%+/=]+",
r#"\{[^{}]*"type"\s*:\s*"service_account"[^{}]*(?:\{[^{}]*\}[^{}]*)*\}"#,
]
.join("|")
}
/// Python's `_ENABLE_SECRET_REDACTION` pattern set, compiled once per configuration.
#[derive(Clone, Debug)]
pub struct SecretRedactor {
pattern: Regex,
}
impl SecretRedactor {
pub fn new(minimum_custom_key_length: usize) -> Self {
let pattern = Regex::new(&format!(
"(?i){}",
secret_patterns(minimum_custom_key_length)
))
.expect("secret redaction patterns compile");
Self { pattern }
}
/// `None` when `LITELLM_DISABLE_REDACT_SECRETS` turns redaction off.
pub fn from_env() -> Option<Self> {
let disabled = std::env::var("LITELLM_DISABLE_REDACT_SECRETS")
.is_ok_and(|value| value.eq_ignore_ascii_case("true"));
(!disabled).then(|| Self::new(minimum_custom_key_length()))
}
pub fn redact(&self, value: &str) -> String {
self.pattern.replace_all(value, REDACTED).into_owned()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[rstest::rstest]
#[case::bearer("auth failed: Bearer abcdefghijklmnop", "auth failed: REDACTED")]
#[case::sk_key("key sk-abcdefghijklmnopqrstuvwxyz rejected", "key REDACTED rejected")]
#[case::short_sk_key_is_kept("sk-abc", "sk-abc")]
#[case::query_param("GET /v1?api_key=secret123&x=1", "GET /v1?REDACTED&x=1")]
#[case::dict_repr("{'api_key': 'abcdefghij'}", "{'REDACTED'}")]
#[case::url_credentials("postgres://user:pass@host/db", "postgres://REDACTED@host/db")]
#[case::case_insensitive("BEARER ABCDEFGHIJKLMNOP", "REDACTED")]
#[case::aws_key("AKIAABCDEFGHIJKLMNOP", "REDACTED")]
#[case::sas_signature("https://x.blob/a?sv=1&sig=abc%2B=", "https://x.blob/a?sv=1&REDACTED")]
#[case::password_needs_word_boundary("db_password=hunter2", "REDACTED")]
#[case::plain_text_is_kept(r#"{"message": "rejected"}"#, r#"{"message": "rejected"}"#)]
fn redacts_the_same_spans_as_the_python_patterns(#[case] input: &str, #[case] expected: &str) {
assert_eq!(
SecretRedactor::new(DEFAULT_MINIMUM_CUSTOM_KEY_LENGTH).redact(input),
expected
);
}
#[test]
fn sk_threshold_follows_the_minimum_custom_key_length() {
let redactor = SecretRedactor::new(8);
assert_eq!(redactor.redact("sk-abcde"), REDACTED);
assert_eq!(redactor.redact("sk-abcd"), "sk-abcd");
}
}

View file

@ -9,7 +9,7 @@ autotests = false
[dependencies]
litellm-types.workspace = true
litellm-core-utils.workspace = true
litellm-callbacks.workspace = true
litellm-host.workspace = true
bytes.workspace = true
futures-util.workspace = true
base64.workspace = true

View file

@ -2,7 +2,6 @@ pub mod audio_transcription;
pub mod chat_completions;
pub mod constants;
pub mod error;
pub mod machine;
pub mod messages;
pub mod ocr;
pub mod responses;

View file

@ -1,88 +1,54 @@
use litellm_llms::custom_httpx::http_handler::http_request;
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
use std::time::Duration;
use super::{
Error, client::http_client, common_utils::truncate_error_body,
prepare::prepare_provider_request,
use litellm_llms::{
base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig,
custom_httpx::{http_handler::http_request, transport::Error as TransportError},
};
use crate::{constants::ANTHROPIC_MESSAGES_PROVIDER, messages::types::MessagesRequest};
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
use serde_json::Value;
pub(super) async fn execute_messages_provider_call(
request: MessagesRequest<'_>,
use super::{Error, client::http_client, common_utils::truncate_error_body};
pub(super) fn network(error: reqwest::Error) -> Error {
Error::Transport(TransportError::Network(error.to_string()))
}
pub(super) async fn send(
url: &str,
headers: &[(String, String)],
body: &Value,
timeout: Option<Duration>,
) -> Result<reqwest::Response, Error> {
let builder = headers.iter().fold(
http_client().post(url).json(body),
|builder, (key, value)| builder.header(key, value),
);
let builder = match timeout {
Some(duration) => builder.timeout(duration),
None => builder,
};
http_request(builder).await.map_err(network)
}
pub(super) async fn provider_error(response: reqwest::Response) -> Error {
let status = response.status().as_u16();
match response.text().await {
Ok(text) => Error::Transport(TransportError::Http {
status,
body: truncate_error_body(&text),
}),
Err(error) => network(error),
}
}
pub(super) fn decode_response(
config: &dyn BaseAnthropicMessagesConfig,
model: &str,
text: &str,
) -> Result<AnthropicMessagesResponse, Error> {
let request = prepare_provider_request(request)?;
let mut request_builder = http_client().post(&request.url).json(&request.body);
for (key, value) in &request.upstream_headers {
request_builder = request_builder.header(key, value);
}
if let Some(duration) = request.timeout {
request_builder = request_builder.timeout(duration);
}
let response = http_request(request_builder).await.map_err(|err| {
Error::Transport(litellm_llms::custom_httpx::transport::Error::Network(
err.to_string(),
))
})?;
let status = response.status();
let text = response.text().await.map_err(|err| {
Error::Transport(litellm_llms::custom_httpx::transport::Error::Network(
err.to_string(),
))
})?;
if !status.is_success() {
return Err(Error::Transport(
litellm_llms::custom_httpx::transport::Error::Http {
status: status.as_u16(),
body: truncate_error_body(&text),
},
));
}
let response = serde_json::from_str(&text)
let response = serde_json::from_str(text)
.map_err(|err| Error::InvalidResponse(format!("invalid messages response JSON: {err}")))?;
request
.config
.transform_anthropic_messages_response(&request.model, response)
config
.transform_anthropic_messages_response(model, response)
.map_err(Error::from)
}
pub(super) async fn execute_messages_provider_stream(
request: MessagesRequest<'_>,
) -> Result<reqwest::Response, Error> {
let request = prepare_provider_request(request)?;
if request.provider != ANTHROPIC_MESSAGES_PROVIDER {
return Err(Error::Unsupported("streaming messages for this provider"));
}
let mut request_builder = http_client().post(&request.url).json(&request.body);
for (key, value) in &request.upstream_headers {
request_builder = request_builder.header(key, value);
}
if let Some(duration) = request.timeout {
request_builder = request_builder.timeout(duration);
}
let response = http_request(request_builder).await.map_err(|err| {
Error::Transport(litellm_llms::custom_httpx::transport::Error::Network(
err.to_string(),
))
})?;
let status = response.status();
if !status.is_success() {
let text = response.text().await.map_err(|err| {
Error::Transport(litellm_llms::custom_httpx::transport::Error::Network(
err.to_string(),
))
})?;
return Err(Error::Transport(
litellm_llms::custom_httpx::transport::Error::Http {
status: status.as_u16(),
body: truncate_error_body(&text),
},
));
}
Ok(response)
}

View file

@ -1,11 +1,8 @@
//! The Anthropic Messages call, the Rust equivalent of Python's
//! `litellm.messages()`.
//!
//! [`messages`] is the top-level entrypoint: give it a model, a body, and
//! credentials, and it resolves the provider, transforms the request, calls the
//! provider, and returns a typed non-streaming response. [`messages_stream`]
//! is the streaming variant; it hands the raw upstream response back so a host
//! can splice the event stream to its own caller.
//! [`route`] is the call as a machine a host drives, streaming or not. [`messages`] runs
//! it in process for a caller that already holds the request and wants the message.
mod error;
pub mod types;
@ -14,17 +11,34 @@ mod client;
mod common_utils;
mod handler;
mod prepare;
use handler::{execute_messages_provider_call, execute_messages_provider_stream};
pub mod route;
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
use route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine};
use serde_json::Value;
use crate::messages::types::MessagesRequest;
pub async fn messages(request: MessagesRequest<'_>) -> Result<AnthropicMessagesResponse, Error> {
execute_messages_provider_call(request).await
}
pub async fn messages_stream(request: MessagesRequest<'_>) -> Result<reqwest::Response, Error> {
execute_messages_provider_stream(request).await
let Value::Object(body) = request.body else {
return Err(Error::InvalidRequest(
"messages body must be an object".into(),
));
};
let call = MessagesCall {
model: request.model.into(),
body,
api_key: request.api_key.map(Into::into),
api_base: request.api_base.map(Into::into),
custom_llm_provider: request.custom_llm_provider.map(Into::into),
extra_headers: request.extra_headers,
timeout: request.timeout,
};
match litellm_host::run::run(messages_machine(), &LocalMessagesHost::new(call)).await? {
MessagesOutput::Message(message) => Ok(*message),
MessagesOutput::Streamed => Err(Error::Unsupported(
"streamed responses need a streaming host",
)),
}
}
#[cfg(test)]

View file

@ -2,6 +2,7 @@ use litellm_core_utils::get_llm_provider_logic::{CustomLlmProvider, get_custom_l
use litellm_llms::base_llm::anthropic_messages::transformation::{
BaseAnthropicMessagesConfig, MessagesAuthStrategy,
};
use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest;
use serde_json::{Map, Value};
use super::{
@ -37,10 +38,14 @@ pub(super) fn prepare_provider_request(
let headers =
validate_environment(config, request.extra_headers, request.api_key, &env_lookup)?;
let typed_request = serde_json::from_value(request.body).map_err(|err| {
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
let typed_request: AnthropicMessagesRequest =
serde_json::from_value(request.body).map_err(|err| {
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}"))
})?;
let transformed = config.transform_anthropic_messages_request(AnthropicMessagesRequest {
model: model.clone(),
..typed_request
})?;
let transformed = config.transform_anthropic_messages_request(typed_request)?;
let body = serde_json::to_value(transformed).map_err(|err| {
Error::InvalidRequest(format!(
"failed to serialize Anthropic messages request: {err}"

View file

@ -0,0 +1,196 @@
use std::{sync::Mutex, time::Duration};
use bytes::Bytes;
use litellm_auth::SecretValue;
use litellm_core_utils::get_llm_provider_logic::get_custom_llm_provider;
use litellm_host::{
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
host::{Demand, Host},
machine::{HostChannel, MachineFault, RouteMachine},
route::Route,
};
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
use serde_json::{Map, Value};
use super::{
Error,
common_utils::messages_provider_config,
handler::{decode_response, network, provider_error, send},
prepare::prepare_provider_request,
types::MessagesRequest,
};
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MessagesOp {
ProjectRequest,
}
pub enum MessagesOpResult {
Request(Box<MessagesCall>),
}
/// The caller's request as the host projects it.
pub struct MessagesCall {
pub model: String,
pub body: Map<String, Value>,
pub api_key: Option<String>,
pub api_base: Option<String>,
pub custom_llm_provider: Option<String>,
pub extra_headers: Option<Map<String, Value>>,
pub timeout: Option<Duration>,
}
impl MessagesCall {
fn streams(&self) -> bool {
self.body.get("stream").and_then(Value::as_bool) == Some(true)
}
}
pub enum MessagesOutput {
Message(Box<AnthropicMessagesResponse>),
/// Every chunk already reached the host through `Deliver`.
Streamed,
}
pub struct Messages;
impl Route for Messages {
type Response = MessagesOutput;
type Error = Error;
type Op = MessagesOp;
type OpResult = MessagesOpResult;
type Chunk = Bytes;
type StreamHead = ();
}
impl From<MachineFault> for Error {
fn from(fault: MachineFault) -> Self {
Self::InvalidRequest(match fault {
MachineFault::Abandoned => "messages host driver was abandoned".into(),
MachineFault::Protocol(message) => format!("messages {message}"),
MachineFault::Mismatch => "invalid messages host operation result".into(),
})
}
}
pub type MessagesHost = HostChannel<Messages>;
pub type MessagesMachine = RouteMachine<Messages>;
/// Whether this route serves the request, decided before any callback runs so a host
/// can still run its own path.
pub fn supports(model: &str, custom_llm_provider: Option<&str>, stream: bool) -> bool {
let provider = get_custom_llm_provider(model, custom_llm_provider)
.map(|resolved| resolved.custom_llm_provider)
.or(custom_llm_provider);
match provider {
Some(ANTHROPIC_MESSAGES_PROVIDER) => true,
Some(provider) => !stream && messages_provider_config(provider).is_some(),
None => false,
}
}
/// The in-process host for a request already in hand. It answers projection once and
/// observes nothing.
pub struct LocalMessagesHost {
call: Mutex<Option<MessagesCall>>,
}
impl LocalMessagesHost {
pub fn new(call: MessagesCall) -> Self {
Self {
call: Mutex::new(Some(call)),
}
}
}
impl Host<Messages> for LocalMessagesHost {
async fn route(&self, op: MessagesOp) -> Result<MessagesOpResult, Error> {
match op {
MessagesOp::ProjectRequest => self
.call
.lock()
.unwrap_or_else(|error| error.into_inner())
.take()
.map(|call| MessagesOpResult::Request(Box::new(call)))
.ok_or_else(|| {
Error::InvalidRequest("messages request was already projected".into())
}),
}
}
}
pub fn messages_machine() -> MessagesMachine {
RouteMachine::new(|host| Box::pin(execute(host)))
}
async fn execute(host: MessagesHost) -> Result<MessagesOutput, Error> {
let MessagesOpResult::Request(call) = host.route(MessagesOp::ProjectRequest).await?;
let stream = call.streams();
let request = prepare_provider_request(MessagesRequest {
model: &call.model,
body: Value::Object(call.body.clone()),
api_key: call.api_key.as_deref(),
api_base: call.api_base.as_deref(),
custom_llm_provider: call.custom_llm_provider.as_deref(),
extra_headers: call.extra_headers.clone(),
timeout: call.timeout,
})?;
if stream && request.provider != ANTHROPIC_MESSAGES_PROVIDER {
return Err(Error::Unsupported("streaming messages for this provider"));
}
let context = RequestContext {
model: request.model.clone(),
custom_llm_provider: request.provider.clone(),
optional_params: Value::Object(
call.body
.iter()
.filter(|(name, _)| !matches!(name.as_str(), "model" | "messages"))
.map(|(name, value)| (name.clone(), value.clone()))
.collect(),
),
secret_fields: Vec::new(),
api_key: call.api_key.clone().map(SecretValue::new),
};
let wire = host
.before_send(
WireRequest {
url: request.url,
headers: request.upstream_headers,
body: request.body,
},
context,
)
.await?;
let response = send(&wire.url, &wire.headers, &wire.body, request.timeout).await?;
if !response.status().is_success() {
return Err(provider_error(response).await);
}
if stream {
return relay(&host, response).await;
}
let text = response.text().await.map_err(network)?;
host.emit(MachineEvent::ResponseReceived {
raw: RawResponse { body: text.clone() },
})
.await?;
decode_response(request.config, &request.model, &text)
.map(|message| MessagesOutput::Message(Box::new(message)))
}
/// Hands each upstream chunk to the caller as it arrives. A caller that stops reading
/// ends the upstream read, and the call completes with what it delivered.
async fn relay(
host: &MessagesHost,
mut response: reqwest::Response,
) -> Result<MessagesOutput, Error> {
if host.open(()).await? == Demand::Detached {
return Ok(MessagesOutput::Streamed);
}
while let Some(chunk) = response.chunk().await.map_err(network)? {
if host.deliver(chunk).await? == Demand::Detached {
break;
}
}
Ok(MessagesOutput::Streamed)
}

View file

@ -12,7 +12,7 @@ pub async fn perform(
client: &OcrClient,
request: LiteLLMOcrRequest,
) -> Result<LiteLLMOcrResponse, Error> {
litellm_callbacks::run::run(ocr_machine(client.clone()), &LocalOcrHost::new(request)).await
litellm_host::run::run(ocr_machine(client.clone()), &LocalOcrHost::new(request)).await
}
pub async fn ocr(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {

View file

@ -1,5 +1,6 @@
use futures_util::future::BoxFuture;
use litellm_callbacks::event::{CallEvent, Passthrough, RawResponse, RequestContext, WireRequest};
use litellm_auth::SecretValue;
use litellm_host::event::{MachineEvent, RawResponse, RequestContext, WireRequest};
use litellm_llms::{
base_llm::ocr::{
error::Error,
@ -36,6 +37,7 @@ pub(crate) struct OcrCallHooks {
custom_llm_provider: &'static str,
optional_params: Value,
secret_fields: Vec<String>,
api_key: Option<SecretValue>,
}
impl OcrCallHooks {
@ -51,28 +53,25 @@ impl OcrCallHooks {
.filter(|name| is_secret_param(name))
.cloned()
.collect(),
api_key: request.connection.api_key.clone(),
}
}
}
impl CallHooks<Error> for OcrCallHooks {
fn before_send(
&self,
wire: WireRequest,
passthrough_fields: Passthrough,
) -> BoxFuture<'_, Result<WireRequest, Error>> {
fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result<WireRequest, Error>> {
let context = RequestContext {
model: self.model.clone(),
custom_llm_provider: self.custom_llm_provider.into(),
optional_params: self.optional_params.clone(),
passthrough_fields,
secret_fields: self.secret_fields.clone(),
api_key: self.api_key.clone(),
};
Box::pin(self.host.before_send(wire, context))
}
fn response_received<'a>(&'a self, body: &'a [u8]) -> BoxFuture<'a, Result<(), Error>> {
Box::pin(self.host.emit(CallEvent::ResponseReceived {
Box::pin(self.host.emit(MachineEvent::ResponseReceived {
raw: RawResponse {
body: String::from_utf8_lossy(body).into_owned(),
},

View file

@ -21,8 +21,8 @@ mod cohere_tests;
#[path = "../../tests/deepseek_ocr.rs"]
mod deepseek_tests;
#[cfg(test)]
#[path = "../../tests/ocr/passthrough.rs"]
mod passthrough_tests;
#[path = "../../tests/ocr/document.rs"]
mod document_tests;
#[cfg(test)]
#[path = "../../tests/reducto_ocr.rs"]
mod reducto_tests;

View file

@ -1,4 +1,4 @@
use litellm_auth::{InputSource, Sourced};
use litellm_auth::{InputSource, SecretValue, Sourced};
use litellm_llms::base_llm::ocr::transformation::{
OcrConnection, OcrCredentialInputs, PreparedOcrRequest, credential_env,
};
@ -22,7 +22,7 @@ pub(crate) fn prepare_request(
.config
.get_api_key_env_var()
.and_then(credential_env)
.map(|value| Sourced::new(value, InputSource::Environment))
.map(|value| Sourced::new(SecretValue::new(value), InputSource::Environment))
})
});
let dynamic_api_base = credentials.dynamic_api_base.or_else(|| {

View file

@ -277,12 +277,18 @@ mod tests {
#[test]
fn connection_resolution_preserves_dynamic_precedence_and_input_sources() {
let connection = OcrConfigKind::Mistral.resolve_connection_params(OcrCredentialInputs {
api_key: Some(Sourced::new("explicit-key".into(), InputSource::Deployment)),
api_key: Some(Sourced::new(
litellm_auth::SecretValue::new("explicit-key"),
InputSource::Deployment,
)),
api_base: Some(Sourced::new(
"https://explicit.test".into(),
InputSource::Deployment,
)),
dynamic_api_key: Some(Sourced::new("dynamic-key".into(), InputSource::Environment)),
dynamic_api_key: Some(Sourced::new(
litellm_auth::SecretValue::new("dynamic-key"),
InputSource::Environment,
)),
dynamic_api_base: Some(Sourced::new(
"https://dynamic.test".into(),
InputSource::Request,
@ -292,7 +298,7 @@ mod tests {
connection
.api_key
.as_ref()
.map(|value| value.value().as_str()),
.map(|value| value.value().expose()),
Some("dynamic-key")
);
assert_eq!(
@ -318,22 +324,31 @@ mod tests {
fn empty_or_missing_dynamic_credentials_preserve_explicit_values(
#[case] dynamic_value: Option<&str>,
) {
let dynamic =
let dynamic_key = dynamic_value.map(|value| {
Sourced::new(
litellm_auth::SecretValue::new(value),
InputSource::Environment,
)
});
let dynamic_base =
dynamic_value.map(|value| Sourced::new(value.into(), InputSource::Environment));
let connection = OcrConfigKind::Mistral.resolve_connection_params(OcrCredentialInputs {
api_key: Some(Sourced::new("explicit-key".into(), InputSource::Deployment)),
api_key: Some(Sourced::new(
litellm_auth::SecretValue::new("explicit-key"),
InputSource::Deployment,
)),
api_base: Some(Sourced::new(
"https://explicit.test".into(),
InputSource::Deployment,
)),
dynamic_api_key: dynamic.clone(),
dynamic_api_base: dynamic,
dynamic_api_key: dynamic_key,
dynamic_api_base: dynamic_base,
});
assert_eq!(
connection
.api_key
.as_ref()
.map(|value| value.value().as_str()),
.map(|value| value.value().expose()),
Some("explicit-key")
);
assert_eq!(
@ -356,11 +371,18 @@ mod tests {
) {
let connection = OcrConfigKind::AzureDocumentIntelligence.resolve_connection_params(
OcrCredentialInputs {
api_key: explicit_key
.map(|value| Sourced::new(value.into(), InputSource::Deployment)),
api_key: explicit_key.map(|value| {
Sourced::new(
litellm_auth::SecretValue::new(value),
InputSource::Deployment,
)
}),
api_base: explicit_base
.map(|value| Sourced::new(value.into(), InputSource::Deployment)),
dynamic_api_key: Some(Sourced::new("dynamic-key".into(), InputSource::Environment)),
dynamic_api_key: Some(Sourced::new(
litellm_auth::SecretValue::new("dynamic-key"),
InputSource::Environment,
)),
dynamic_api_base: Some(Sourced::new(
"https://dynamic.test".into(),
InputSource::Deployment,
@ -371,7 +393,7 @@ mod tests {
connection
.api_key
.as_ref()
.map(|value| value.value().as_str()),
.map(|value| value.value().expose()),
explicit_key.map(|_| "dynamic-key")
);
assert_eq!(

View file

@ -1,8 +1,9 @@
use std::sync::{Arc, Mutex};
use litellm_auth::ResolvedCredential;
use litellm_callbacks::{
use litellm_host::{
event::{CallEvent, RequestContext, WireRequest},
machine::{HostChannel, HostTokenProvider, MachineFault, RouteMachine, TokenRoute},
route::Route,
};
use litellm_llms::{
@ -11,10 +12,7 @@ use litellm_llms::{
};
use super::handler::perform_ocr_request;
use crate::{
machine::{HostChannel, HostTokenProvider, MachineFault, RouteMachine, TokenRoute},
ocr::types::{LiteLLMOcrRequest, OcrDocumentInput, OcrFileContent, ResolvedOcrRequest},
};
use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput, OcrFileContent, ResolvedOcrRequest};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum OcrOp {
@ -39,6 +37,8 @@ impl Route for Ocr {
type Error = Error;
type Op = OcrOp;
type OpResult = OcrOpResult;
type Chunk = std::convert::Infallible;
type StreamHead = std::convert::Infallible;
}
impl TokenRoute for Ocr {
@ -54,16 +54,6 @@ impl TokenRoute for Ocr {
}
}
impl From<MachineFault> for Error {
fn from(fault: MachineFault) -> Self {
Self::InvalidRequest(match fault {
MachineFault::Abandoned => "OCR host driver was abandoned".into(),
MachineFault::Protocol(message) => format!("OCR {message}"),
MachineFault::Mismatch => "invalid OCR host operation result".into(),
})
}
}
pub type OcrHost = HostChannel<Ocr>;
pub type OcrMachine = RouteMachine<Ocr>;
@ -173,7 +163,7 @@ impl LocalOcrHost {
}
}
impl litellm_callbacks::host::Host<Ocr> for LocalOcrHost {
impl litellm_host::host::Host<Ocr> for LocalOcrHost {
async fn route(&self, op: OcrOp) -> Result<OcrOpResult, Error> {
match op {
OcrOp::ProjectRequest => self

View file

@ -1,7 +1,7 @@
use std::{collections::BTreeMap, path::PathBuf, time::Duration};
use bytes::Bytes;
use litellm_auth::{InputSource, TokenProviderHandle};
use litellm_auth::{InputSource, SecretValue, TokenProviderHandle};
use litellm_core_utils::call_arguments::CallArguments;
use litellm_llms::base_llm::ocr::{
error::Error,
@ -56,7 +56,7 @@ pub struct OcrFileContent {
/// credentials, and per-field provenance in `input_sources`.
#[derive(Clone, Debug, Default)]
pub struct OcrConnectionInputs {
pub api_key: Option<String>,
pub api_key: Option<SecretValue>,
pub api_base: Option<String>,
pub extra_headers: Map<String, Value>,
pub timeout: Option<Duration>,
@ -237,6 +237,16 @@ mod tests {
.unwrap()
}
#[test]
fn connection_inputs_debug_hides_the_api_key() {
let inputs = OcrConnectionInputs {
api_key: Some(SecretValue::new("caller-api-key")),
..OcrConnectionInputs::default()
};
assert!(!format!("{inputs:?}").contains("caller-api-key"));
}
#[test]
fn from_inputs_applies_connection_overrides_with_field_sources() {
let request = LiteLLMOcrRequest::from_inputs(
@ -245,7 +255,7 @@ mod tests {
None,
Default::default(),
OcrConnectionInputs {
api_key: Some(" key ".into()),
api_key: Some(SecretValue::new(" key ")),
api_base: Some("".into()),
extra_headers: json!({"x-a": "1"}).as_object().unwrap().clone(),
timeout: Some(Duration::from_secs(7)),
@ -259,7 +269,7 @@ mod tests {
.unwrap();
let api_key = request.credentials.api_key.as_ref().unwrap();
assert_eq!(api_key.clone().into_value(), "key");
assert_eq!(api_key.value().expose(), "key");
assert_eq!(api_key.source(), InputSource::Request);
assert!(request.credentials.api_base.is_none());
assert_eq!(

View file

@ -1,6 +1,6 @@
use std::{collections::BTreeMap, time::Duration};
use litellm_auth::InputSource;
use litellm_auth::{InputSource, SecretValue};
use litellm_llms::base_llm::ocr::{
error::Error,
transformation::{OcrDocument, decode_request_value},
@ -44,7 +44,7 @@ pub fn consumed_optional_param_names(
pub struct OcrWireRequest<D = Value> {
pub model: String,
pub document: D,
pub api_key: Option<String>,
pub api_key: Option<SecretValue>,
pub api_base: Option<String>,
pub custom_llm_provider: Option<String>,
pub extra_headers: Option<Map<String, Value>>,

View file

@ -1,4 +1,4 @@
use litellm_callbacks::event::CallEvent;
use litellm_host::event::{CallEvent, MachineEvent};
use litellm_llms::base_llm::ocr::error::Error;
use rstest::rstest;
use serde_json::{Value, json};
@ -69,7 +69,7 @@ async fn rejects_invalid_pages_features_and_format(
let result = decode_request(OcrWireRequest {
model: "azure_ai/doc-intelligence/prebuilt-read".into(),
document: json!({"type":"document_url","document_url":"https://example.com/a.pdf"}),
api_key: Some("key".into()),
api_key: Some(litellm_auth::SecretValue::new("key")),
api_base: Some(base),
custom_llm_provider: None,
extra_headers: None,
@ -263,7 +263,7 @@ async fn accepted_response_emits_response_received_before_polling() {
json!({}),
))
.with_observer(move |event| {
let CallEvent::ResponseReceived { raw } = event else {
let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event else {
return;
};
match request_count.lock().unwrap().len() {
@ -466,7 +466,7 @@ async fn model_id_is_encoded_and_dot_segments_are_rejected() {
mod transformation {
use std::sync::{Arc, Mutex};
use litellm_callbacks::event::CallEvent;
use litellm_host::event::{CallEvent, MachineEvent};
use litellm_llms::base_llm::ocr::transformation::OcrDocument;
use serde_json::{Value, json};
@ -646,7 +646,7 @@ mod transformation {
json!({}),
))
.with_observer(move |event| {
if let CallEvent::ResponseReceived { raw } = event {
if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event {
observed
.lock()
.unwrap()

View file

@ -1,7 +1,7 @@
use std::sync::{Arc, Mutex};
use litellm_callbacks::{
event::{CallEvent, WireRequest},
use litellm_host::{
event::{CallEvent, MachineEvent, WireRequest},
host::{Host, HostOp, HostResult},
machine::{HostFailure, Machine, MachineStep},
};
@ -81,7 +81,7 @@ fn request_boundary_selects_mistral_and_rejects_unknown_providers() {
let request = OcrWireRequest {
model: "mistral/model".into(),
document: json!({"type":"document_url","document_url":"https://example.com/doc.pdf"}),
api_key: Some("key".into()),
api_key: Some(litellm_auth::SecretValue::new("key")),
api_base: None,
custom_llm_provider: None,
extra_headers: None,
@ -97,7 +97,7 @@ fn request_boundary_selects_mistral_and_rejects_unknown_providers() {
decode_request(OcrWireRequest {
model: "model".into(),
document: json!({"type":"document_url","document_url":"https://example.com/doc.pdf"}),
api_key: Some("key".into()),
api_key: Some(litellm_auth::SecretValue::new("key")),
api_base: None,
custom_llm_provider: Some("unknown".into()),
extra_headers: None,
@ -194,7 +194,8 @@ async fn facade_uses_the_injected_http_client() {
fn event_name(event: &CallEvent) -> &'static str {
match event {
CallEvent::ResponseReceived { .. } => "response",
CallEvent::Started { .. } => "started",
CallEvent::Machine(MachineEvent::ResponseReceived { .. }) => "response",
CallEvent::Succeeded { .. } => "success",
CallEvent::Failed { .. } => "failure",
}
@ -235,7 +236,7 @@ async fn lifecycle_sends_headers_returned_by_the_before_send_operation() {
}
#[tokio::test]
async fn before_send_context_names_passthrough_fields_and_secrets() {
async fn before_send_context_names_the_route_and_its_secrets() {
let (base, _, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let observed = Arc::new(Mutex::new(None));
let captured = observed.clone();
@ -254,8 +255,6 @@ async fn before_send_context_names_passthrough_fields_and_secrets() {
assert_eq!(context.custom_llm_provider, "mistral");
assert_eq!(context.model, "model");
assert_eq!(wire.body["pages"], json!([0]));
assert!(context.passthrough_fields.contains("pages"));
assert!(context.passthrough_fields.contains("document"));
assert!(context.secret_fields.is_empty());
assert_eq!(context.optional_params["req_format"], "native");
@ -279,7 +278,6 @@ async fn before_send_context_names_passthrough_fields_and_secrets() {
perform_ocr_with(host).await.unwrap();
server.await.unwrap();
let context = observed.lock().unwrap().take().unwrap();
assert!(!context.passthrough_fields.contains("document"));
assert_eq!(context.secret_fields, ["client_secret"]);
}
@ -296,7 +294,7 @@ async fn lifecycle_orders_hooks_and_emits_one_success() {
server.await.unwrap();
assert_eq!(
*events.lock().unwrap(),
["before_send", "response", "success"]
["started", "before_send", "response", "success"]
);
assert_eq!(seen.lock().unwrap().len(), 1);
}
@ -311,7 +309,10 @@ async fn lifecycle_blocking_prevents_execution_and_emits_one_failure() {
);
let error = perform_ocr_with(host).await.unwrap_err();
assert!(matches!(error, OcrError::InvalidRequest(message) if message == "blocked"));
assert_eq!(*events.lock().unwrap(), ["before_send", "failure"]);
assert_eq!(
*events.lock().unwrap(),
["started", "before_send", "failure"]
);
}
#[tokio::test]
@ -330,7 +331,10 @@ async fn upstream_failure_emits_one_terminal_failure() {
);
assert!(perform_ocr_with(host).await.is_err());
server.await.unwrap();
assert_eq!(*events.lock().unwrap(), ["before_send", "failure"]);
assert_eq!(
*events.lock().unwrap(),
["started", "before_send", "failure"]
);
assert_eq!(seen.lock().unwrap().len(), 1);
}
@ -371,6 +375,7 @@ async fn drive_until(
intercept(*wire).map(|wire| HostResult::BeforeSend(Box::new(wire)))
}
HostOp::Emit(event) => {
let event = CallEvent::Machine(event);
ops.push(event_name(&event));
host.emit(&event)
.await
@ -414,7 +419,7 @@ async fn invalid_provider_response_emits_response_received_before_normalization_
let observed = responses_received.clone();
let host = LocalOcrHost::new(wire_request("mistral/model", &base, json!({}))).with_observer(
move |event| {
if let CallEvent::ResponseReceived { raw } = event {
if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event {
observed.lock().unwrap().push(raw.body.clone());
}
},
@ -815,7 +820,7 @@ impl Host<crate::ocr::route::Ocr> for CallerTokenHost {
async fn before_send(
&self,
wire: WireRequest,
_: &litellm_callbacks::event::RequestContext,
_: &litellm_host::event::RequestContext,
) -> Result<WireRequest, OcrError> {
let is_authorization = |name: &str| name.eq_ignore_ascii_case("authorization");
let authorization = wire
@ -850,7 +855,7 @@ async fn the_callers_azure_token_is_acquired_before_before_send_which_can_still_
trace: Mutex::new(Vec::new()),
};
litellm_callbacks::run::run(ocr_machine(ocr_client()), &host)
litellm_host::run::run(ocr_machine(ocr_client()), &host)
.await
.unwrap();
server.await.unwrap();

View file

@ -0,0 +1,152 @@
use litellm_host::event::WireRequest;
use litellm_llms::base_llm::ocr::error::Error;
use rstest::rstest;
use serde_json::{Value, json};
use super::test_support::{
MockResponse, SERVED_DOCUMENT, document_server, mock_server, perform_ocr_with, request_body,
wire_request_with_document,
};
use crate::ocr::route::LocalOcrHost;
#[derive(Clone, Copy, Debug)]
enum Route {
Mistral,
AzureAi,
VertexMistral,
AzureCohereParse,
Cohere,
}
impl Route {
fn model(self) -> &'static str {
match self {
Self::Mistral => "mistral/model",
Self::AzureAi => "azure_ai/model",
Self::VertexMistral => "vertex_ai/mistral-ocr-maas",
Self::AzureCohereParse => "azure_ai/cohere-parse",
Self::Cohere => "cohere/model",
}
}
fn document_type(self) -> &'static str {
match self {
Self::Mistral | Self::AzureAi | Self::VertexMistral => "document_url",
Self::AzureCohereParse | Self::Cohere => "image_url",
}
}
fn options(self) -> Value {
match self {
Self::Mistral | Self::AzureAi => json!({"pages": [0]}),
Self::VertexMistral => json!({"pages": [0], "vertex_project": "project-1"}),
Self::AzureCohereParse | Self::Cohere => json!({"output_format": "markdown"}),
}
}
}
/// What the host does to the wire request in `before_send`.
#[derive(Clone, Copy, Debug)]
enum Host {
Detached,
ReplacesDocument,
}
const REPLACED_DOCUMENT: &str = "data:image/png;base64,cmVwbGFjZWQ=";
impl Host {
fn before_send(self, wire: WireRequest) -> WireRequest {
let Value::Object(fields) = wire.body else {
return wire;
};
let body = fields
.into_iter()
.map(|(name, value)| match self {
Self::Detached => (name, value),
Self::ReplacesDocument if name == "document" => {
let document_type = value["type"].clone();
let key = document_type.as_str().unwrap_or_default().to_string();
(name, json!({"type": document_type, key: REPLACED_DOCUMENT}))
}
Self::ReplacesDocument => (name, value),
})
.collect();
WireRequest {
body: Value::Object(body),
..wire
}
}
}
struct Sent {
result: Result<(), Error>,
provider_body: Option<Value>,
}
async fn send(route: Route, host: Host, document_base: &str) -> Sent {
let (base, seen, provider) = mock_server(vec![MockResponse::json(json!({"pages": []}))]).await;
let document_type = route.document_type();
let document =
json!({"type": document_type, document_type: format!("{document_base}/scan.png")});
let request = wire_request_with_document(route.model(), &base, document, route.options());
let local =
LocalOcrHost::new(request).with_before_send(move |wire, _| Ok(host.before_send(wire)));
let result = perform_ocr_with(local).await.map(|_| ());
match result {
Ok(()) => provider.await.unwrap(),
Err(_) => provider.abort(),
}
let provider_body = seen
.lock()
.unwrap()
.first()
.map(|request| request_body(request));
Sent {
result,
provider_body,
}
}
fn served_document_uri() -> String {
use base64::Engine;
format!(
"data:image/png;base64,{}",
base64::engine::general_purpose::STANDARD.encode(SERVED_DOCUMENT)
)
}
#[rstest]
#[case::azure_ai(Route::AzureAi)]
#[case::vertex_mistral(Route::VertexMistral)]
#[case::azure_cohere_parse(Route::AzureCohereParse)]
#[tokio::test]
async fn inlining_routes_send_the_downloaded_document(#[case] route: Route) {
let (document_base, _documents) = document_server().await;
let sent = send(route, Host::Detached, &document_base).await;
sent.result.unwrap();
assert_eq!(
sent.provider_body.unwrap()["document"][route.document_type()],
json!(served_document_uri())
);
}
#[rstest]
#[tokio::test]
async fn document_replaced_by_the_host_reaches_the_provider(
#[values(
Route::Mistral,
Route::AzureAi,
Route::VertexMistral,
Route::AzureCohereParse,
Route::Cohere
)]
route: Route,
) {
let (document_base, _documents) = document_server().await;
let sent = send(route, Host::ReplacesDocument, &document_base).await;
sent.result.unwrap();
assert_eq!(
sent.provider_body.unwrap()["document"][route.document_type()],
json!(REPLACED_DOCUMENT)
);
}

View file

@ -1,282 +0,0 @@
use std::{
collections::BTreeSet,
sync::{Arc, Mutex},
};
use litellm_callbacks::event::{RequestContext, WireRequest};
use litellm_llms::base_llm::ocr::error::Error;
use rstest::rstest;
use rstest_reuse::{self, apply, template};
use serde_json::{Map, Value, json};
use super::test_support::{
MockResponse, SERVED_DOCUMENT, document_server, mock_server, perform_ocr_with, request_body,
wire_request_with_document,
};
use crate::ocr::route::LocalOcrHost;
#[derive(Clone, Copy, Debug)]
enum Route {
Mistral,
AzureAi,
VertexMistral,
AzureCohereParse,
Cohere,
}
impl Route {
fn model(self) -> &'static str {
match self {
Self::Mistral => "mistral/model",
Self::AzureAi => "azure_ai/model",
Self::VertexMistral => "vertex_ai/mistral-ocr-maas",
Self::AzureCohereParse => "azure_ai/cohere-parse",
Self::Cohere => "cohere/model",
}
}
fn document_type(self) -> &'static str {
match self {
Self::Mistral | Self::AzureAi | Self::VertexMistral => "document_url",
Self::AzureCohereParse | Self::Cohere => "image_url",
}
}
fn options(self) -> Value {
match self {
Self::Mistral | Self::AzureAi => json!({"pages": [0]}),
Self::VertexMistral => json!({"pages": [0], "vertex_project": "project-1"}),
Self::AzureCohereParse | Self::Cohere => json!({"output_format": "markdown"}),
}
}
}
#[derive(Clone, Copy, Debug)]
enum Source {
Inline,
Remote,
RemoteWithExtraField,
}
/// What the host does to the wire request in `before_send`.
#[derive(Clone, Copy, Debug)]
enum Host {
Detached,
/// What `litellm-callbacks-legacy` does before `pre_call`: every passthrough body key
/// is replaced by the caller's own value.
Realiasing,
ReplacesDocument,
}
const REPLACED_DOCUMENT: &str = "data:image/png;base64,cmVwbGFjZWQ=";
impl Host {
fn before_send(
self,
caller: &Map<String, Value>,
wire: WireRequest,
context: &RequestContext,
) -> WireRequest {
let Value::Object(fields) = wire.body else {
return wire;
};
let body = fields
.into_iter()
.map(|(name, value)| match self {
Self::Detached => (name, value),
Self::Realiasing => {
let aliased = context
.passthrough_fields
.contains(&name)
.then(|| caller.get(&name).cloned())
.flatten()
.unwrap_or(value);
(name, aliased)
}
Self::ReplacesDocument if name == "document" => {
let document_type = value["type"].clone();
let key = document_type.as_str().unwrap_or_default().to_string();
(name, json!({"type": document_type, key: REPLACED_DOCUMENT}))
}
Self::ReplacesDocument => (name, value),
})
.collect();
WireRequest {
body: Value::Object(body),
..wire
}
}
}
struct Sent {
caller: Map<String, Value>,
result: Result<(), Error>,
before_send: Option<(WireRequest, RequestContext)>,
provider_body: Option<Value>,
}
fn caller_document(route: Route, source: Source, document_base: &str) -> Value {
let document_type = route.document_type();
let remote = format!("{document_base}/scan.png");
match source {
Source::Inline => {
json!({"type": document_type, document_type: "data:image/png;base64,YWJj"})
}
Source::Remote => json!({"type": document_type, document_type: remote}),
Source::RemoteWithExtraField => {
json!({"type": document_type, document_type: remote, "document_name": "scan.png"})
}
}
}
async fn send(route: Route, source: Source, host: Host, document_base: &str) -> Sent {
let (base, seen, provider) = mock_server(vec![MockResponse::json(json!({"pages": []}))]).await;
let document = caller_document(route, source, document_base);
let caller: Map<String, Value> = route
.options()
.as_object()
.unwrap()
.clone()
.into_iter()
.chain([("document".to_string(), document.clone())])
.collect();
let observed = Arc::new(Mutex::new(None));
let captured = observed.clone();
let host_caller = caller.clone();
let request = wire_request_with_document(route.model(), &base, document, route.options());
let local = LocalOcrHost::new(request).with_before_send(move |wire, context| {
*captured.lock().unwrap() = Some((wire.clone(), context.clone()));
Ok(host.before_send(&host_caller, wire, context))
});
let result = perform_ocr_with(local).await.map(|_| ());
match result {
Ok(()) => provider.await.unwrap(),
Err(_) => provider.abort(),
}
let provider_body = seen
.lock()
.unwrap()
.first()
.map(|request| request_body(request));
let before_send = observed.lock().unwrap().take();
Sent {
caller,
result,
before_send,
provider_body,
}
}
fn served_document_uri() -> String {
use base64::Engine;
format!(
"data:image/png;base64,{}",
base64::engine::general_purpose::STANDARD.encode(SERVED_DOCUMENT)
)
}
#[template]
#[rstest]
fn every_route_and_source(
#[values(
Route::Mistral,
Route::AzureAi,
Route::VertexMistral,
Route::AzureCohereParse,
Route::Cohere
)]
route: Route,
#[values(Source::Inline, Source::Remote, Source::RemoteWithExtraField)] source: Source,
) {
}
#[template]
#[rstest]
fn every_route(
#[values(
Route::Mistral,
Route::AzureAi,
Route::VertexMistral,
Route::AzureCohereParse,
Route::Cohere
)]
route: Route,
) {
}
#[template]
#[rstest]
#[case::azure_ai(Route::AzureAi)]
#[case::vertex_mistral(Route::VertexMistral)]
#[case::azure_cohere_parse(Route::AzureCohereParse)]
fn inlining_routes(#[case] route: Route) {}
#[apply(every_route_and_source)]
#[tokio::test]
async fn passthrough_fields_are_exactly_the_caller_values_sent_unchanged(
route: Route,
source: Source,
) {
let (document_base, _documents) = document_server().await;
let sent = send(route, source, Host::Detached, &document_base).await;
sent.result.unwrap();
let (wire, context) = sent.before_send.unwrap();
let passthrough: BTreeSet<&str> = context.passthrough_fields.iter().collect();
let unchanged: BTreeSet<&str> = sent
.caller
.iter()
.filter(|(name, value)| wire.body.get(name.as_str()) == Some(*value))
.map(|(name, _)| name.as_str())
.collect();
assert_eq!(
passthrough,
unchanged,
"body: {:#}\ncaller: {:#}",
wire.body,
Value::Object(sent.caller.clone())
);
}
#[apply(every_route_and_source)]
#[tokio::test]
async fn realiasing_leaves_the_provider_request_unchanged(route: Route, source: Source) {
let (document_base, _documents) = document_server().await;
let detached = send(route, source, Host::Detached, &document_base).await;
let realiased = send(route, source, Host::Realiasing, &document_base).await;
detached.result.unwrap();
realiased.result.unwrap();
assert_eq!(realiased.provider_body, detached.provider_body);
}
#[apply(inlining_routes)]
#[tokio::test]
async fn inlining_routes_send_the_downloaded_document(
route: Route,
#[values(Host::Detached, Host::Realiasing)] host: Host,
) {
let (document_base, _documents) = document_server().await;
let sent = send(route, Source::Remote, host, &document_base).await;
sent.result.unwrap();
assert_eq!(
sent.provider_body.unwrap()["document"][route.document_type()],
json!(served_document_uri())
);
}
#[apply(every_route)]
#[tokio::test]
async fn document_replaced_by_the_host_reaches_the_provider(route: Route) {
let (document_base, _documents) = document_server().await;
let sent = send(
route,
Source::Remote,
Host::ReplacesDocument,
&document_base,
)
.await;
sent.result.unwrap();
assert_eq!(
sent.provider_body.unwrap()["document"][route.document_type()],
json!(REPLACED_DOCUMENT)
);
}

View file

@ -1,7 +1,7 @@
use std::sync::{Arc, Mutex};
use futures_util::future::BoxFuture;
use litellm_callbacks::event::{Passthrough, WireRequest};
use litellm_host::event::WireRequest;
use litellm_llms::{
base_llm::ocr::{error::Error, transformation::LiteLLMOcrResponse},
custom_httpx::llm_http_handler::{CallHooks, OcrClient},
@ -23,11 +23,7 @@ use crate::ocr::{
pub(crate) struct NoHooks;
impl CallHooks<Error> for NoHooks {
fn before_send(
&self,
wire: WireRequest,
_passthrough_fields: Passthrough,
) -> BoxFuture<'_, Result<WireRequest, Error>> {
fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result<WireRequest, Error>> {
Box::pin(async move { Ok(wire) })
}
@ -49,7 +45,7 @@ pub(crate) async fn perform_ocr(request: LiteLLMOcrRequest) -> Result<LiteLLMOcr
}
pub(crate) async fn perform_ocr_with(host: LocalOcrHost) -> Result<LiteLLMOcrResponse, Error> {
litellm_callbacks::run::run(ocr_machine(ocr_client()), &host).await
litellm_host::run::run(ocr_machine(ocr_client()), &host).await
}
pub(crate) fn wire_request(model: &str, base: &str, options: Value) -> LiteLLMOcrRequest {
@ -70,7 +66,7 @@ pub(crate) fn wire_request_with_document(
decode_request(OcrWireRequest {
model: model.into(),
document,
api_key: Some("test-key".into()),
api_key: Some(litellm_auth::SecretValue::new("test-key")),
api_base: Some(base.into()),
custom_llm_provider: None,
extra_headers: None,

View file

@ -1,4 +1,4 @@
use litellm_callbacks::event::{CallEvent, WireRequest};
use litellm_host::event::{CallEvent, MachineEvent, WireRequest};
use litellm_llms::base_llm::ocr::{error::Error, transformation::OcrDocument};
use rstest::rstest;
use serde_json::{Value, json};
@ -139,7 +139,7 @@ async fn response_received_stays_after_reducto_upload_and_parse() {
let request_count = seen.clone();
let host = LocalOcrHost::new(wire_request("reducto/parse-v3", &base, json!({}))).with_observer(
move |event| {
if let CallEvent::ResponseReceived { raw } = event {
if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event {
assert_eq!(request_count.lock().unwrap().len(), 2);
assert_eq!(raw.body, r#"{"result":{"chunks":[]}}"#);
}
@ -351,7 +351,7 @@ async fn guardrail_rewrites_document_before_upload() {
}
mod transformation {
use litellm_callbacks::event::{CallEvent, WireRequest};
use litellm_host::event::{CallEvent, MachineEvent, WireRequest};
use litellm_llms::{
base_llm::ocr::transformation::{BaseOcrConfig, OcrConnection, OcrRequestContext},
reducto::ocr::transformation::*,
@ -506,7 +506,7 @@ mod transformation {
let request_count = seen.clone();
let host = LocalOcrHost::new(wire_request("reducto/parse-v3", &base, json!({})))
.with_observer(move |event| {
if let CallEvent::ResponseReceived { raw } = event {
if let CallEvent::Machine(MachineEvent::ResponseReceived { raw }) = event {
assert_eq!(request_count.lock().unwrap().len(), 2);
assert_eq!(raw.body, r#"{"result":{"chunks":[]}}"#);
}

View file

@ -1,9 +1,10 @@
- Target invariants; implementation and runtime validation may lag these rules
- Keep this crate the CPython runtime adapter and nothing more: Serde marshalling, interpreter detachment, tokio/asyncio glue, the `Execution` handle, the call driver and the `CallbackAdapter`/`RouteHost` traits
- No LiteLLM domain dependencies beyond `litellm-callbacks`: no route types, no `Logging` policy, no public API registration, no cdylib build features
- Keep this crate the CPython runtime adapter and nothing more: Serde marshalling, interpreter detachment, tokio/asyncio glue, the `Execution` handle, the call driver and the `PythonLifecycle`/`RouteHost` traits
- No LiteLLM domain dependencies beyond `litellm-host`: no route types, no `Logging` policy, no public API registration, no cdylib build features
- The driver emits `Succeeded` or `Failed` exactly once and never dispatches after a cancellation; which Python objects consume those events is the adapter's business
- `RouteHost::invoke` receives the keyword view the adapter's `begin` returned, not the caller's dict; a route host that projects from it inherits that adapter's rewrites (for the legacy adapter: setup, deployment hooks, credential inheritance)
- A failure that surfaces inside the call, including a host op the call asked for, is mapped through the route's `map_failure`; a failure in `begin` or `after_success` is raised as is
- A native failure, including one a host op returns as `InvokeError::Native`, is classified exactly once through the route's `classify`; a Python exception raised inside the call, and a failure in `begin` or `after_success`, is raised as is
- A failing `classify` is raised with the native error's text as its `__context__`, never swallowed
- Use standard PyO3 ownership and conversion APIs
- Prefer `Bound<'py, T>` for attached operations/results, `Py<T>` for retention; binding/unbinding does not copy payloads
- Use `pythonize` for selected Serde data, never a JSON-text round trip; share conversion with `Pythonized<T>`

View file

@ -7,7 +7,7 @@ repository.workspace = true
[dependencies]
futures-util.workspace = true
litellm-callbacks.workspace = true
litellm-host.workspace = true
pyo3.workspace = true
pyo3-async-runtimes.workspace = true
pythonize.workspace = true

View file

@ -1,5 +1,5 @@
use litellm_callbacks::event::{CallEvent, RequestContext, Timing, WireRequest};
use litellm_callbacks::route::Route;
use litellm_host::event::{FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest};
use litellm_host::route::Route;
use pyo3::exceptions::PyRuntimeError;
use pyo3::gc::{PyTraverseError, PyVisit};
use pyo3::prelude::*;
@ -11,7 +11,7 @@ pub fn missing_state() -> PyErr {
/// What an adapter step produced: either the value the driver asked for, or a Python
/// awaitable the driver hands back to the caller's task before asking again.
pub enum AdapterStep {
pub enum LifecycleStep {
Await(Py<PyAny>),
Arguments(Py<PyDict>),
Wire(Box<WireRequest>),
@ -19,64 +19,97 @@ pub enum AdapterStep {
Done,
}
/// The host-typed value the driver attaches to a terminal event.
pub enum PublicValue<'a> {
Response(&'a Py<PyAny>),
Error(&'a PyErr),
/// What a lifecycle observes: the driver's start, the machine's own events, and one
/// terminal event carrying the public value the caller receives.
pub enum LifecycleEvent<'a> {
Started {
start_time: f64,
},
Machine(&'a MachineEvent),
Succeeded {
timing: Timing,
response: &'a Py<PyAny>,
},
Failed {
timing: Timing,
origin: FailureOrigin,
error: &'a PyErr,
},
}
/// One consumer of a call's lifecycle on the Python side. The driver calls the steps in
/// order: `begin` before the machine starts, `before_send` and `emit` while it runs,
/// `after_success` and one terminal `emit` after it completes. Whenever a step returns
/// [`AdapterStep::Await`], the driver awaits it in the caller's task and continues the
/// [`LifecycleStep::Await`], the driver awaits it in the caller's task and continues the
/// same step through `resume`.
///
/// A step that fails with an ordinary exception fails the call with that exception,
/// except on a terminal event, where the adapter is expected to report and swallow its
/// own errors. An exception that is not a `PyException`, such as a cancellation, ends
/// the call without further dispatch.
pub trait CallbackAdapter: Send + Sync {
pub trait PythonLifecycle: Send + Sync {
fn begin(
&mut self,
py: Python<'_>,
arguments: Py<PyDict>,
started_at: f64,
) -> PyResult<AdapterStep>;
) -> PyResult<LifecycleStep>;
fn before_send(
&mut self,
py: Python<'_>,
wire: Box<WireRequest>,
context: &RequestContext,
) -> PyResult<AdapterStep>;
) -> PyResult<LifecycleStep>;
fn after_success(
&mut self,
py: Python<'_>,
response: Py<PyAny>,
timing: Timing,
) -> PyResult<AdapterStep>;
) -> PyResult<LifecycleStep>;
fn emit(
&mut self,
py: Python<'_>,
event: &CallEvent,
public: Option<PublicValue<'_>>,
) -> PyResult<AdapterStep>;
fn emit(&mut self, py: Python<'_>, event: LifecycleEvent<'_>) -> PyResult<LifecycleStep>;
fn resume(&mut self, py: Python<'_>, result: PyResult<Py<PyAny>>) -> PyResult<AdapterStep>;
/// The call streams and its stream was handed to the caller. The caller is not
/// inside an await here, so this step and `delivered` cannot suspend.
fn opened(&mut self, py: Python<'_>) -> PyResult<()>;
/// One chunk of an open stream is about to reach the caller.
fn delivered(&mut self, py: Python<'_>, chunk: &Py<PyAny>) -> PyResult<()>;
fn resume(&mut self, py: Python<'_>, result: PyResult<Py<PyAny>>) -> PyResult<LifecycleStep>;
fn close(&mut self, py: Python<'_>);
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>;
}
/// The Python side of one route: answers the route's own operations, builds the public
/// response and maps failures to public exceptions.
pub trait RouteHost: Send + Sync {
type Route: Route;
/// Why a route operation the host answered did not produce a result: the route's own code
/// rejected it, which the route classifies like any other native failure, or Python code
/// raised, which reaches the caller as it was raised.
#[derive(Debug)]
pub enum InvokeError<E> {
Native(E),
Python(PyErr),
}
/// `arguments` is the keyword view the callback adapter's `begin` produced, not the
impl<E> From<PyErr> for InvokeError<E> {
fn from(error: PyErr) -> Self {
Self::Python(error)
}
}
/// The Python side of one route: answers the route's own operations, builds the public
/// response and classifies native failures into public exceptions.
pub trait RouteHost: Send + Sync {
type Route: Route<Error: std::fmt::Display>;
/// The public exception a native failure maps to, kept as a value until the driver
/// raises it.
type Failure: Into<PyErr>;
/// `arguments` is the keyword view the lifecycle's `begin` produced, not the
/// caller's own dict. A route host that projects from it inherits whatever that
/// adapter rewrote.
fn invoke(
@ -84,7 +117,7 @@ pub trait RouteHost: Send + Sync {
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
op: <Self::Route as Route>::Op,
) -> PyResult<<Self::Route as Route>::OpResult>;
) -> Result<<Self::Route as Route>::OpResult, InvokeError<<Self::Route as Route>::Error>>;
fn complete(
&mut self,
@ -92,12 +125,21 @@ pub trait RouteHost: Send + Sync {
response: <Self::Route as Route>::Response,
) -> PyResult<Py<PyAny>>;
fn native_error(error: <Self::Route as Route>::Error) -> PyErr;
/// One streamed chunk as the caller receives it.
fn chunk(
&mut self,
py: Python<'_>,
chunk: <Self::Route as Route>::Chunk,
) -> PyResult<Py<PyAny>>;
fn classify(
&self,
py: Python<'_>,
error: <Self::Route as Route>::Error,
) -> PyResult<Self::Failure>;
fn host_error(error: &PyErr) -> <Self::Route as Route>::Error;
fn map_failure(&self, py: Python<'_>, error: &PyErr) -> PyResult<PyErr>;
fn close(&mut self, py: Python<'_>);
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>;

View file

@ -0,0 +1,51 @@
use pyo3::{prelude::*, types::PyDict};
/// The caller's own object for a public argument: the keyword if given, even an explicit
/// `None`, else the bound request's attribute. Every reader of a public Python call uses
/// this rule, so the callbacks and the provider see one object per argument.
pub fn lookup<'py>(
kwargs: &Bound<'py, PyDict>,
request: &Bound<'py, PyAny>,
name: &str,
) -> PyResult<Option<Bound<'py, PyAny>>> {
if let Some(value) = kwargs.get_item(name)? {
return Ok(Some(value));
}
request.getattr_opt(name)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn lookup_prefers_the_keyword_even_when_none_and_falls_back_to_the_request() {
crate::initialize_python();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
c"
key = object()
document = {'type': 'document_url'}
class Request:
api_key = 'from-request'
api_base = 'from-request'
document = document
request = Request()
kwargs = {'api_key': key, 'api_base': None}
",
Some(&locals),
Some(&locals),
)
.unwrap();
let item = |name: &str| locals.get_item(name).unwrap().unwrap();
let kwargs = item("kwargs").cast_into::<PyDict>().unwrap();
let request = item("request");
let find = |name: &str| lookup(&kwargs, &request, name).unwrap();
assert!(find("api_key").unwrap().is(item("key")));
assert!(find("api_base").unwrap().is_none());
assert!(find("document").unwrap().is(item("document")));
assert!(find("model").is_none());
});
}
}

View file

@ -2,17 +2,19 @@ use std::sync::Arc;
use std::task::Poll;
use futures_util::future::{AbortHandle, Abortable};
use litellm_callbacks::event::{CallEvent, FailureOrigin, Timing, epoch_seconds};
use litellm_callbacks::host::{HostOp, HostResult, HostStep};
use litellm_callbacks::machine::{HostFailure, Machine, MachineStep};
use litellm_callbacks::route::Route;
use litellm_host::event::{FailureOrigin, Timing, epoch_seconds};
use litellm_host::host::{Demand, HostOp, HostResult, HostStep};
use litellm_host::machine::{HostFailure, Machine, MachineStep};
use litellm_host::route::Route;
use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError};
use pyo3::gc::{PyTraverseError, PyVisit};
use pyo3::prelude::*;
use pyo3::types::PyDict;
use tokio::sync::Mutex;
use crate::adapter::{AdapterStep, CallbackAdapter, PublicValue, RouteHost, missing_state};
use crate::adapter::{
InvokeError, LifecycleEvent, LifecycleStep, PythonLifecycle, RouteHost, missing_state,
};
use crate::execution::{poll_async_value, run_async_value, run_sync_value};
use crate::handle::{Execution, ExecutionBody, ExecutionStep};
@ -36,6 +38,7 @@ struct MachineState<M: Machine> {
enum Stage {
Begin,
Call,
Streaming,
AfterSuccess,
Succeeded(Py<PyAny>),
Failed(Py<PyBaseException>),
@ -43,6 +46,7 @@ enum Stage {
#[derive(Clone, Copy)]
enum Expect {
Started,
Arguments,
Wire,
Emitted,
@ -53,6 +57,8 @@ enum Expect {
enum Pending {
Native,
Adapter(Expect),
/// The stream handed to the caller waits for its next read or its close.
Consumer,
}
enum Next<H: RouteHost> {
@ -66,7 +72,7 @@ where
M: Machine<Route = H::Route, Complete = ResponseOf<H>> + 'static,
{
route: H,
adapter: Box<dyn CallbackAdapter>,
adapter: Box<dyn PythonLifecycle>,
machine: Option<Arc<Mutex<MachineState<M>>>>,
arguments: Option<Py<PyDict>>,
started_at: f64,
@ -84,7 +90,7 @@ pub fn run_call<H, M>(
py: Python<'_>,
machine: M,
route: H,
adapter: Box<dyn CallbackAdapter>,
adapter: Box<dyn PythonLifecycle>,
arguments: Py<PyDict>,
asynchronous: bool,
) -> PyResult<Py<PyAny>>
@ -118,7 +124,14 @@ where
}
match driver.resume(None)? {
ExecutionStep::Return(value) => Ok(value),
ExecutionStep::Await(_) => Err(PyRuntimeError::new_err("sync call suspended")),
ExecutionStep::Open => py
.import("litellm.rust_bridge.lifecycle")?
.getattr("SyncStream")?
.call1((Py::new(py, Execution::suspended(driver))?,))
.map(Bound::unbind),
ExecutionStep::Await(_) | ExecutionStep::Yield(_) => {
Err(PyRuntimeError::new_err("sync call suspended"))
}
}
}
@ -146,9 +159,11 @@ where
match (self.pending.take(), result) {
(None, None) => {
self.started_at = epoch_seconds();
let arguments = self.arguments.take().ok_or_else(missing_state)?;
match self.adapter.begin(py, arguments, self.started_at) {
Ok(step) => self.on_adapter(py, step, Expect::Arguments),
let started = LifecycleEvent::Started {
start_time: self.started_at,
};
match self.adapter.emit(py, started) {
Ok(step) => self.on_adapter(py, step, Expect::Started),
Err(error) => self.adapter_failed(py, error),
}
}
@ -157,6 +172,14 @@ where
self.run_steps(py, HostStep::Ready(result))
}
(Some(Pending::Native), Some(Err(error))) => self.interrupt(py, error),
(Some(Pending::Consumer), Some(read)) => {
let demand = if read.is_ok() {
Demand::More
} else {
Demand::Detached
};
self.resume_machine(py, Some(Ok(HostResult::Demand(demand))))
}
(Some(Pending::Adapter(expect)), Some(result)) => {
match self.adapter.resume(py, result) {
Ok(step) => self.on_adapter(py, step, expect),
@ -170,27 +193,28 @@ where
fn on_adapter(
&mut self,
py: Python<'_>,
step: AdapterStep,
step: LifecycleStep,
expect: Expect,
) -> PyResult<ExecutionStep> {
match (expect, step) {
(_, AdapterStep::Await(awaitable)) => {
(_, LifecycleStep::Await(awaitable)) => {
self.pending = Some(Pending::Adapter(expect));
Ok(ExecutionStep::Await(awaitable))
}
(Expect::Arguments, AdapterStep::Arguments(arguments)) => {
(Expect::Started, LifecycleStep::Done) => self.begin(py),
(Expect::Arguments, LifecycleStep::Arguments(arguments)) => {
self.arguments = Some(arguments);
self.stage = Stage::Call;
self.resume_machine(py, None)
}
(Expect::Wire, AdapterStep::Wire(wire)) => {
(Expect::Wire, LifecycleStep::Wire(wire)) => {
self.resume_machine(py, Some(Ok(HostResult::BeforeSend(wire))))
}
(Expect::Emitted, AdapterStep::Done) => {
(Expect::Emitted, LifecycleStep::Done) => {
self.resume_machine(py, Some(Ok(HostResult::Emitted)))
}
(Expect::Response, AdapterStep::Response(response)) => self.succeeded(py, response),
(Expect::Terminal, AdapterStep::Done) => match &self.stage {
(Expect::Response, LifecycleStep::Response(response)) => self.succeeded(py, response),
(Expect::Terminal, LifecycleStep::Done) => match &self.stage {
Stage::Succeeded(response) => Ok(ExecutionStep::Return(response.clone_ref(py))),
Stage::Failed(error) => Err(PyErr::from_value(error.bind(py).clone().into_any())),
_ => Err(missing_state()),
@ -199,10 +223,18 @@ where
}
}
fn begin(&mut self, py: Python<'_>) -> PyResult<ExecutionStep> {
let arguments = self.arguments.take().ok_or_else(missing_state)?;
match self.adapter.begin(py, arguments, self.started_at) {
Ok(step) => self.on_adapter(py, step, Expect::Arguments),
Err(error) => self.adapter_failed(py, error),
}
}
fn adapter_failed(&mut self, py: Python<'_>, error: PyErr) -> PyResult<ExecutionStep> {
match self.stage {
Stage::Begin | Stage::AfterSuccess => self.failure(py, error, FailureOrigin::Host),
Stage::Call => self.interrupt(py, error),
Stage::Call | Stage::Streaming => self.interrupt(py, error),
Stage::Succeeded(_) | Stage::Failed(_) => Err(error),
}
}
@ -248,14 +280,20 @@ where
let answer = match op {
HostOp::Route(op) => {
let arguments = self.arguments.as_ref().ok_or_else(missing_state)?;
self.route
.invoke(py, arguments.bind(py), op)
.map(HostResult::Route)
match self.route.invoke(py, arguments.bind(py), op) {
Ok(result) => Ok(HostResult::Route(result)),
Err(InvokeError::Native(error)) => {
return self
.resume_core(py, Some(Err(HostFailure::Error(error))))
.map(Next::Continue);
}
Err(InvokeError::Python(error)) => Err(error),
}
}
HostOp::BeforeSend { wire, context } => {
match self.adapter.before_send(py, wire, &context) {
Ok(AdapterStep::Wire(wire)) => Ok(HostResult::BeforeSend(wire)),
Ok(AdapterStep::Await(awaitable)) => {
Ok(LifecycleStep::Wire(wire)) => Ok(HostResult::BeforeSend(wire)),
Ok(LifecycleStep::Await(awaitable)) => {
self.pending = Some(Pending::Adapter(Expect::Wire));
return Ok(Next::Return(ExecutionStep::Await(awaitable)));
}
@ -263,9 +301,11 @@ where
Err(error) => Err(error),
}
}
HostOp::Emit(event) => match self.adapter.emit(py, &event, None) {
Ok(AdapterStep::Done) => Ok(HostResult::Emitted),
Ok(AdapterStep::Await(awaitable)) => {
HostOp::Open(_) => return self.opened(py).map(Next::Return),
HostOp::Deliver(chunk) => return self.delivered(py, chunk).map(Next::Return),
HostOp::Emit(event) => match self.adapter.emit(py, LifecycleEvent::Machine(&event)) {
Ok(LifecycleStep::Done) => Ok(HostResult::Emitted),
Ok(LifecycleStep::Await(awaitable)) => {
self.pending = Some(Pending::Adapter(Expect::Emitted));
return Ok(Next::Return(ExecutionStep::Await(awaitable)));
}
@ -279,6 +319,35 @@ where
}
}
fn opened(&mut self, py: Python<'_>) -> PyResult<ExecutionStep> {
self.stage = Stage::Streaming;
match self.adapter.opened(py) {
Ok(()) => {
self.pending = Some(Pending::Consumer);
Ok(ExecutionStep::Open)
}
Err(error) => self.interrupt(py, error),
}
}
fn delivered(
&mut self,
py: Python<'_>,
chunk: <RouteOf<H> as Route>::Chunk,
) -> PyResult<ExecutionStep> {
let chunk = match self.route.chunk(py, chunk) {
Ok(chunk) => chunk,
Err(error) => return self.interrupt(py, error),
};
match self.adapter.delivered(py, &chunk) {
Ok(()) => {
self.pending = Some(Pending::Consumer);
Ok(ExecutionStep::Yield(chunk))
}
Err(error) => self.interrupt(py, error),
}
}
fn interrupt(&mut self, py: Python<'_>, error: PyErr) -> PyResult<ExecutionStep> {
let cancelled = is_cancellation(py, &error);
let native = H::host_error(&error);
@ -349,6 +418,9 @@ where
Ok(public) => public,
Err(error) => return self.failure(py, error, FailureOrigin::Call),
};
if let Stage::Streaming = self.stage {
return self.succeeded(py, public);
}
self.stage = Stage::AfterSuccess;
match self.adapter.after_success(py, public, self.timing()) {
Ok(step) => self.on_adapter(py, step, Expect::Response),
@ -360,18 +432,35 @@ where
self.ended_at.get_or_insert_with(epoch_seconds);
let error = match self.interrupted.take() {
Some(retained) => PyErr::from_value(retained.into_bound(py).into_any()),
None => H::native_error(error),
None => self.classified(py, error),
};
self.failure(py, error, FailureOrigin::Call)
}
fn succeeded(&mut self, py: Python<'_>, response: Py<PyAny>) -> PyResult<ExecutionStep> {
let event = CallEvent::Succeeded {
timing: self.timing(),
/// The route's public exception for a native failure. When classification itself
/// fails, that failure is raised with the native error's text as its `__context__`.
fn classified(&self, py: Python<'_>, error: ErrorOf<H>) -> PyErr {
let native = error.to_string();
let classifier_error = match self.route.classify(py, error) {
Ok(failure) => return failure.into(),
Err(classifier_error) => classifier_error,
};
let step = self
.adapter
.emit(py, &event, Some(PublicValue::Response(&response)))?;
let attached = classifier_error.value(py).setattr(
"__context__",
PyRuntimeError::new_err(native).into_value(py),
);
match attached {
Ok(()) => classifier_error,
Err(error) => error,
}
}
fn succeeded(&mut self, py: Python<'_>, response: Py<PyAny>) -> PyResult<ExecutionStep> {
let event = LifecycleEvent::Succeeded {
timing: self.timing(),
response: &response,
};
let step = self.adapter.emit(py, event)?;
self.stage = Stage::Succeeded(response);
self.on_adapter(py, step, Expect::Terminal)
}
@ -386,18 +475,13 @@ where
if is_cancellation(py, &error) {
return Err(error);
}
let public = match origin {
FailureOrigin::Call => self.route.map_failure(py, &error).unwrap_or(error),
FailureOrigin::Host => error,
};
let event = CallEvent::Failed {
let event = LifecycleEvent::Failed {
timing: self.timing(),
origin,
error: &error,
};
let step = self
.adapter
.emit(py, &event, Some(PublicValue::Error(&public)))?;
self.stage = Stage::Failed(public.into_value(py));
let step = self.adapter.emit(py, event)?;
self.stage = Stage::Failed(error.into_value(py));
self.on_adapter(py, step, Expect::Terminal)
}
@ -450,8 +534,8 @@ where
mod tests {
use std::sync::{Arc, Mutex};
use litellm_callbacks::event::{RequestContext, WireRequest};
use litellm_callbacks::machine::{Interrupted, Step};
use litellm_host::event::{MachineEvent, RequestContext, WireRequest};
use litellm_host::machine::{Interrupted, Step};
use pyo3::exceptions::{PyBaseException, PyValueError};
use pyo3::types::PyDict;
@ -489,6 +573,12 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
#[derive(Clone, Debug, PartialEq, Eq)]
struct Error(String);
impl std::fmt::Display for Error {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str(&self.0)
}
}
struct Synthetic;
impl Route for Synthetic {
@ -496,6 +586,8 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
type Error = Error;
type Op = &'static str;
type OpResult = String;
type Chunk = std::convert::Infallible;
type StreamHead = std::convert::Infallible;
}
/// Yields the scripted ops in order, then completes or fails as scripted.
@ -518,8 +610,8 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
model: "model".into(),
custom_llm_provider: "provider".into(),
optional_params: serde_json::json!({}),
passthrough_fields: Default::default(),
secret_fields: Vec::new(),
api_key: None,
}
}
@ -534,6 +626,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
HostResult::Route(value) => value,
HostResult::BeforeSend(wire) => wire.url,
HostResult::Emitted => "emitted".into(),
HostResult::Demand(demand) => format!("{demand:?}"),
});
}
if !self.ops.is_empty() {
@ -566,25 +659,50 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
}
}
#[derive(Clone, Copy)]
enum OpScript {
Answer,
RaisePython,
RejectNatively,
}
struct SyntheticHost {
log: Log,
fail_op: bool,
op: OpScript,
classifier_fails: bool,
}
/// The fake route's public exception, kept as a value so a test sees what `classify`
/// produced before the driver raises it.
#[derive(Debug, PartialEq, Eq)]
struct Classified(String);
impl From<Classified> for PyErr {
fn from(classified: Classified) -> Self {
PyValueError::new_err(format!("classified: {}", classified.0))
}
}
impl RouteHost for SyntheticHost {
type Route = Synthetic;
type Failure = Classified;
fn invoke(
&mut self,
_: Python<'_>,
arguments: &Bound<'_, PyDict>,
op: &'static str,
) -> PyResult<String> {
) -> Result<String, InvokeError<Error>> {
self.log.push(format!("route:{op}"));
if self.fail_op {
return Err(PyValueError::new_err("op failed"));
match self.op {
OpScript::Answer => Ok(format!("{op}:{}", arguments.len())),
OpScript::RaisePython => Err(PyValueError::new_err("op failed").into()),
OpScript::RejectNatively => Err(InvokeError::Native(Error("op rejected".into()))),
}
Ok(format!("{op}:{}", arguments.len()))
}
fn chunk(&mut self, _: Python<'_>, chunk: std::convert::Infallible) -> PyResult<Py<PyAny>> {
match chunk {}
}
fn complete(&mut self, py: Python<'_>, response: String) -> PyResult<Py<PyAny>> {
@ -594,22 +712,18 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
.unbind())
}
fn native_error(error: Error) -> PyErr {
PyValueError::new_err(error.0)
fn classify(&self, _: Python<'_>, error: Error) -> PyResult<Classified> {
self.log.push(format!("classify:{error}"));
if self.classifier_fails {
return Err(pyo3::exceptions::PyTypeError::new_err("classifier failed"));
}
Ok(Classified(error.0))
}
fn host_error(error: &PyErr) -> Error {
Error(error.to_string())
}
fn map_failure(&self, py: Python<'_>, error: &PyErr) -> PyResult<PyErr> {
self.log.push("map_failure");
Ok(PyValueError::new_err(format!(
"mapped: {}",
error.value(py)
)))
}
fn close(&mut self, _: Python<'_>) {
self.log.push("route.close");
}
@ -632,13 +746,18 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
script: AdapterScript,
}
impl CallbackAdapter for SyntheticAdapter {
fn begin(&mut self, _: Python<'_>, arguments: Py<PyDict>, _: f64) -> PyResult<AdapterStep> {
impl PythonLifecycle for SyntheticAdapter {
fn begin(
&mut self,
_: Python<'_>,
arguments: Py<PyDict>,
_: f64,
) -> PyResult<LifecycleStep> {
self.log.push("begin");
if matches!(self.script, AdapterScript::FailBegin) {
return Err(PyValueError::new_err("begin failed"));
}
Ok(AdapterStep::Arguments(arguments))
Ok(LifecycleStep::Arguments(arguments))
}
fn before_send(
@ -646,9 +765,9 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
_: Python<'_>,
wire: Box<WireRequest>,
_: &RequestContext,
) -> PyResult<AdapterStep> {
) -> PyResult<LifecycleStep> {
self.log.push("before_send");
Ok(AdapterStep::Wire(Box::new(WireRequest {
Ok(LifecycleStep::Wire(Box::new(WireRequest {
url: "rewritten".into(),
..*wire
})))
@ -659,41 +778,48 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
py: Python<'_>,
response: Py<PyAny>,
_: Timing,
) -> PyResult<AdapterStep> {
) -> PyResult<LifecycleStep> {
self.log.push("after_success");
match self.script {
AdapterScript::ReplaceResponse => Ok(AdapterStep::Response(
AdapterScript::ReplaceResponse => Ok(LifecycleStep::Response(
"replaced".into_pyobject(py)?.into_any().unbind(),
)),
AdapterScript::FailAfterSuccess => {
Err(PyValueError::new_err("after_success failed"))
}
AdapterScript::Plain | AdapterScript::FailBegin => {
Ok(AdapterStep::Response(response))
Ok(LifecycleStep::Response(response))
}
}
}
fn emit(
&mut self,
py: Python<'_>,
event: &CallEvent,
public: Option<PublicValue<'_>>,
) -> PyResult<AdapterStep> {
self.log.push(match (event, public) {
(CallEvent::ResponseReceived { raw }, None) => format!("response:{}", raw.body),
(CallEvent::Succeeded { .. }, Some(PublicValue::Response(value))) => {
format!("succeeded:{}", value.bind(py))
fn emit(&mut self, py: Python<'_>, event: LifecycleEvent<'_>) -> PyResult<LifecycleStep> {
self.log.push(match event {
LifecycleEvent::Started { .. } => "started".into(),
LifecycleEvent::Machine(MachineEvent::ResponseReceived { raw }) => {
format!("response:{}", raw.body)
}
(CallEvent::Failed { origin, .. }, Some(PublicValue::Error(error))) => {
LifecycleEvent::Succeeded { response, .. } => {
format!("succeeded:{}", response.bind(py))
}
LifecycleEvent::Failed { origin, error, .. } => {
format!("failed:{origin:?}:{}", error.value(py))
}
_ => "unexpected".into(),
});
Ok(AdapterStep::Done)
Ok(LifecycleStep::Done)
}
fn resume(&mut self, _: Python<'_>, _: PyResult<Py<PyAny>>) -> PyResult<AdapterStep> {
fn opened(&mut self, _: Python<'_>) -> PyResult<()> {
self.log.push("opened");
Ok(())
}
fn delivered(&mut self, _: Python<'_>, _: &Py<PyAny>) -> PyResult<()> {
self.log.push("delivered");
Ok(())
}
fn resume(&mut self, _: Python<'_>, _: PyResult<Py<PyAny>>) -> PyResult<LifecycleStep> {
Err(missing_state())
}
@ -709,15 +835,31 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
fn run_scripted(
py: Python<'_>,
machine: ScriptedMachine,
fail_op: bool,
op: OpScript,
script: AdapterScript,
asynchronous: bool,
) -> (PyResult<Py<PyAny>>, Vec<String>) {
let log = Log::default();
let route = SyntheticHost {
log: Log(log.0.clone()),
fail_op,
};
run_hosted(
py,
machine,
SyntheticHost {
log: Log::default(),
op,
classifier_fails: false,
},
script,
asynchronous,
)
}
fn run_hosted(
py: Python<'_>,
machine: ScriptedMachine,
route: SyntheticHost,
script: AdapterScript,
asynchronous: bool,
) -> (PyResult<Py<PyAny>>, Vec<String>) {
let log = Log(route.log.0.clone());
let adapter = SyntheticAdapter {
log: Log(log.0.clone()),
script,
@ -756,8 +898,8 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
wire: Box::new(wire()),
context: Box::new(context()),
},
HostOp::Emit(CallEvent::ResponseReceived {
raw: litellm_callbacks::event::RawResponse { body: "raw".into() },
HostOp::Emit(MachineEvent::ResponseReceived {
raw: litellm_host::event::RawResponse { body: "raw".into() },
}),
],
outcome: Some(Ok("done".into())),
@ -777,7 +919,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
let (result, log) = run_scripted(
py,
success_machine(),
false,
OpScript::Answer,
AdapterScript::Plain,
asynchronous,
);
@ -785,6 +927,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
assert_eq!(
log,
[
"started",
"begin",
"route:project",
"before_send",
@ -800,28 +943,75 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
});
}
fn failing_machine() -> ScriptedMachine {
ScriptedMachine {
ops: vec![HostOp::Route("project")],
outcome: Some(Err(Error("provider exploded".into()))),
answers: Vec::new(),
}
}
#[test]
fn machine_failures_are_mapped_and_dispatched_once_as_call_failures() {
fn a_native_failure_is_classified_once_and_reported_classified() {
let _guard = PYTHON_GLOBALS
.lock()
.unwrap_or_else(|error| error.into_inner());
crate::initialize_python();
Python::attach(|py| {
let machine = ScriptedMachine {
ops: vec![HostOp::Route("project")],
outcome: Some(Err(Error("provider exploded".into()))),
answers: Vec::new(),
};
let (result, log) = run_scripted(py, machine, false, AdapterScript::Plain, false);
let error = result.unwrap_err();
assert_eq!(error.value(py).to_string(), "mapped: provider exploded");
install_lifecycle_module(py);
for asynchronous in [false, true] {
let (result, log) = run_scripted(
py,
failing_machine(),
OpScript::Answer,
AdapterScript::Plain,
asynchronous,
);
let error = result.unwrap_err();
assert!(error.is_instance_of::<PyValueError>(py));
assert_eq!(error.value(py).to_string(), "classified: provider exploded");
assert_eq!(
log,
[
"started",
"begin",
"route:project",
"classify:provider exploded",
"failed:Call:classified: provider exploded",
"adapter.close",
"route.close",
]
);
}
});
}
#[test]
fn a_native_rejection_from_a_host_operation_is_classified_once() {
let _guard = PYTHON_GLOBALS
.lock()
.unwrap_or_else(|error| error.into_inner());
crate::initialize_python();
Python::attach(|py| {
let (result, log) = run_scripted(
py,
success_machine(),
OpScript::RejectNatively,
AdapterScript::Plain,
false,
);
assert_eq!(
result.unwrap_err().value(py).to_string(),
"classified: op rejected"
);
assert_eq!(
log,
[
"started",
"begin",
"route:project",
"map_failure",
"failed:Call:mapped: provider exploded",
"classify:op rejected",
"failed:Call:classified: op rejected",
"adapter.close",
"route.close",
]
@ -830,18 +1020,72 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
}
#[test]
fn host_operation_failures_interrupt_the_call_and_keep_the_python_exception() {
fn a_python_exception_from_a_host_operation_is_reported_as_raised() {
let _guard = PYTHON_GLOBALS
.lock()
.unwrap_or_else(|error| error.into_inner());
crate::initialize_python();
Python::attach(|py| {
let (result, log) =
run_scripted(py, success_machine(), true, AdapterScript::Plain, false);
let (result, log) = run_scripted(
py,
success_machine(),
OpScript::RaisePython,
AdapterScript::Plain,
false,
);
let error = result.unwrap_err();
assert_eq!(error.value(py).to_string(), "mapped: op failed");
assert!(!log.contains(&"before_send".to_string()));
assert!(log.contains(&"failed:Call:mapped: op failed".to_string()));
assert!(error.is_instance_of::<PyValueError>(py));
assert_eq!(error.value(py).to_string(), "op failed");
assert_eq!(
log,
[
"started",
"begin",
"route:project",
"failed:Call:op failed",
"adapter.close",
"route.close",
]
);
});
}
#[test]
fn a_failing_classifier_surfaces_with_the_native_error_as_context() {
let _guard = PYTHON_GLOBALS
.lock()
.unwrap_or_else(|error| error.into_inner());
crate::initialize_python();
Python::attach(|py| {
let (result, log) = run_hosted(
py,
failing_machine(),
SyntheticHost {
log: Log::default(),
op: OpScript::Answer,
classifier_fails: true,
},
AdapterScript::Plain,
false,
);
let error = result.unwrap_err();
assert!(error.is_instance_of::<pyo3::exceptions::PyTypeError>(py));
assert_eq!(error.value(py).to_string(), "classifier failed");
let context = error.value(py).getattr("__context__").unwrap();
assert!(context.is_instance_of::<PyRuntimeError>());
assert_eq!(context.str().unwrap().to_string(), "provider exploded");
assert_eq!(
log,
[
"started",
"begin",
"route:project",
"classify:provider exploded",
"failed:Call:classifier failed",
"adapter.close",
"route.close",
]
);
});
}
@ -855,7 +1099,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
let (result, log) = run_scripted(
py,
success_machine(),
false,
OpScript::Answer,
AdapterScript::FailBegin,
false,
);
@ -864,6 +1108,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
assert_eq!(
log,
[
"started",
"begin",
"failed:Host:begin failed",
"adapter.close",
@ -885,7 +1130,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
let (result, log) = run_scripted(
py,
success_machine(),
false,
OpScript::Answer,
AdapterScript::ReplaceResponse,
asynchronous,
);
@ -908,7 +1153,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
let (result, log) = run_scripted(
py,
success_machine(),
false,
OpScript::Answer,
AdapterScript::FailAfterSuccess,
asynchronous,
);
@ -938,12 +1183,13 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
struct Cancelling(Log);
impl RouteHost for Cancelling {
type Route = Synthetic;
type Failure = Classified;
fn invoke(
&mut self,
py: Python<'_>,
_: &Bound<'_, PyDict>,
_: &'static str,
) -> PyResult<String> {
) -> Result<String, InvokeError<Error>> {
self.0.push("route");
Err(PyErr::from_value(
py.import("asyncio")
@ -952,21 +1198,26 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
.unwrap()
.call0()
.unwrap(),
))
)
.into())
}
fn chunk(
&mut self,
_: Python<'_>,
chunk: std::convert::Infallible,
) -> PyResult<Py<PyAny>> {
match chunk {}
}
fn complete(&mut self, _: Python<'_>, _: String) -> PyResult<Py<PyAny>> {
Err(missing_state())
}
fn native_error(error: Error) -> PyErr {
PyValueError::new_err(error.0)
fn classify(&self, _: Python<'_>, error: Error) -> PyResult<Classified> {
self.0.push("classify");
Ok(Classified(error.0))
}
fn host_error(error: &PyErr) -> Error {
Error(error.to_string())
}
fn map_failure(&self, _: Python<'_>, _: &PyErr) -> PyResult<PyErr> {
self.0.push("map_failure");
Err(missing_state())
}
fn close(&mut self, _: Python<'_>) {}
fn traverse(&self, _: &PyVisit<'_>) -> Result<(), PyTraverseError> {
Ok(())
@ -988,7 +1239,10 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
)
.unwrap_err();
assert!(!error.is_instance_of::<pyo3::exceptions::PyException>(py));
assert_eq!(log.entries(), ["begin", "route", "adapter.close"]);
assert_eq!(
log.entries(),
["started", "begin", "route", "adapter.close"]
);
});
}

View file

@ -8,6 +8,10 @@ use pyo3::prelude::*;
pub enum ExecutionStep {
Return(Py<PyAny>),
Await(Py<PyAny>),
/// The call streams: the caller gets a stream over this execution, which stays
/// suspended until the stream asks for a chunk.
Open,
Yield(Py<PyAny>),
}
pub trait ExecutionBody: Send + Sync {
@ -34,6 +38,13 @@ impl Execution {
}
}
/// An execution already started elsewhere and now waiting for its next input.
pub fn suspended(body: impl ExecutionBody + 'static) -> Self {
Self {
state: ExecutionState::Suspended(Box::new(body)),
}
}
fn advance(
slf: &Bound<'_, Self>,
py: Python<'_>,
@ -64,6 +75,8 @@ impl Execution {
let step = body.resume(result)?;
let (tag, value, suspended) = match step {
ExecutionStep::Await(value) => ("Await", value, true),
ExecutionStep::Open => ("Open", py.None(), true),
ExecutionStep::Yield(value) => ("Yield", value, true),
ExecutionStep::Return(value) => ("Complete", value, false),
};
let step = py

View file

@ -1,9 +1,10 @@
//! The CPython runtime adapter: value marshalling, interpreter detachment, the tokio and
//! asyncio glue, and the driver that runs a native [`Machine`](litellm_callbacks::machine::Machine)
//! against a Python route host and a callback adapter. Everything here is Python-specific by
//! asyncio glue, and the driver that runs a native [`Machine`](litellm_host::machine::Machine)
//! against a Python route host and a Python lifecycle. Everything here is Python-specific by
//! construction; another host language gets its own crate of the same shape.
mod adapter;
mod argument;
mod callable;
mod driver;
mod execution;
@ -11,7 +12,10 @@ mod gil;
mod handle;
mod marshal;
pub use adapter::{AdapterStep, CallbackAdapter, PublicValue, RouteHost, missing_state};
pub use adapter::{
InvokeError, LifecycleEvent, LifecycleStep, PythonLifecycle, RouteHost, missing_state,
};
pub use argument::lookup;
pub use callable::wrap_failure;
pub use driver::run_call;
pub use execution::{poll_async_value, run_async, run_async_value, run_sync, run_sync_value};

View file

@ -1,13 +1,14 @@
[package]
name = "litellm-callbacks"
name = "litellm-host"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
[dependencies]
litellm-auth.workspace = true
serde_json.workspace = true
tokio = { workspace = true, features = ["sync"] }
[dev-dependencies]
rstest.workspace = true
tokio = { workspace = true, features = ["macros"] }

View file

@ -0,0 +1,76 @@
use std::time::{SystemTime, UNIX_EPOCH};
use serde_json::Value;
/// Seconds since the Unix epoch, on one clock for every host.
pub fn epoch_seconds() -> f64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|duration| duration.as_secs_f64())
.unwrap_or(0.0)
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct Timing {
pub start_time: f64,
pub end_time: f64,
}
/// The provider request as it is about to leave, offered to the host for rewriting.
#[derive(Clone, Debug, PartialEq)]
pub struct WireRequest {
pub url: String,
pub headers: Vec<(String, String)>,
pub body: Value,
}
/// What the route knows about the request it is sending, for a host that logs it. The
/// route owns these facts; a host reads them beside the wire request and never rewrites
/// them.
#[derive(Clone, Debug, PartialEq)]
pub struct RequestContext {
pub model: String,
pub custom_llm_provider: String,
/// The route's parameters before the provider transformation.
pub optional_params: Value,
/// Optional-param names that carry credentials and must be redacted when logged.
pub secret_fields: Vec<String>,
/// The credential the route resolved for the provider call.
pub api_key: Option<litellm_auth::SecretValue>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct RawResponse {
pub body: String,
}
/// Whether a failure surfaced inside the call, including a host op the call asked for,
/// or in a host step around it (preparing the arguments, finalizing the response).
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum FailureOrigin {
Call,
Host,
}
/// What a machine reports while it runs.
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum MachineEvent {
ResponseReceived { raw: RawResponse },
}
/// What an in-process host observes: the machine's own events between the driver's
/// start and terminal ones.
#[derive(Clone, Debug, PartialEq)]
pub enum CallEvent {
Started {
start_time: f64,
},
Machine(MachineEvent),
Succeeded {
timing: Timing,
},
Failed {
timing: Timing,
origin: FailureOrigin,
},
}

View file

@ -1,6 +1,6 @@
use std::future::Future;
use crate::event::{CallEvent, RequestContext, WireRequest};
use crate::event::{CallEvent, MachineEvent, RequestContext, WireRequest};
use crate::route::Route;
/// One suspension point of a native call, performed by the host.
@ -10,13 +10,26 @@ pub enum HostOp<R: Route> {
wire: Box<WireRequest>,
context: Box<RequestContext>,
},
Emit(CallEvent),
Emit(MachineEvent),
/// The response streams: the host hands the caller a stream and answers once the
/// caller asks for the first chunk or goes away.
Open(R::StreamHead),
/// The next chunk of an open stream, answered once the caller asks for the one after.
Deliver(R::Chunk),
}
pub enum HostResult<R: Route> {
Route(R::OpResult),
BeforeSend(Box<WireRequest>),
Emitted,
Demand(Demand),
}
/// Whether the caller of a streamed call still reads it.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Demand {
More,
Detached,
}
/// A host answer that is either available now or arrives once the host's own
@ -42,4 +55,12 @@ pub trait Host<R: Route>: Send + Sync {
fn emit(&self, _event: &CallEvent) -> impl Future<Output = Result<(), R::Error>> + Send {
async { Ok(()) }
}
fn open(&self, _head: R::StreamHead) -> impl Future<Output = Result<Demand, R::Error>> + Send {
async { Ok(Demand::More) }
}
fn deliver(&self, _chunk: R::Chunk) -> impl Future<Output = Result<Demand, R::Error>> + Send {
async { Ok(Demand::More) }
}
}

View file

@ -1,7 +1,7 @@
//! The contract between a native call and the host runtime that drives it.
//!
//! A host is whatever sits on the far side of the language boundary: CPython today,
//! another runtime later. Core implements [`machine::Machine`] per route and never learns
//! another runtime later. Core runs each route on a [`machine::RouteMachine`] and never learns
//! which host is on the other end. The machine yields [`host::HostOp`]s; a driver answers
//! them, observes [`event::CallEvent`]s and may rewrite the wire request before it is sent.

View file

@ -1,9 +1,8 @@
use std::sync::Arc;
use litellm_auth::{Error, ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle};
use litellm_callbacks::route::Route;
use super::{HostChannel, MachineFault};
use crate::route::Route;
use litellm_auth::{Error, ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle};
/// A route whose host can mint credentials on the call's behalf.
pub trait TokenRoute: Route {

View file

@ -1,6 +1,12 @@
mod auth;
mod route_machine;
use std::future::Future;
use std::pin::Pin;
pub use auth::{HostTokenProvider, TokenRoute};
pub use route_machine::{ExecuteFuture, HostChannel, MachineFault, RouteMachine};
use crate::host::{HostOp, HostResult};
use crate::route::Route;

View file

@ -2,18 +2,16 @@
//! place, and turns the host operations that future requests into [`Machine`] steps. No
//! task is spawned; dropping the machine drops the in-flight call.
mod auth;
use std::{future::Future, pin::Pin};
pub use auth::{HostTokenProvider, TokenRoute};
use litellm_callbacks::{
event::{CallEvent, RequestContext, WireRequest},
host::{HostOp, HostResult},
machine::{HostFailure, Interrupted, Machine, MachineStep, Step},
use tokio::sync::{mpsc, oneshot};
use super::{HostFailure, Interrupted, Machine, MachineStep, Step};
use crate::{
event::{MachineEvent, RequestContext, WireRequest},
host::{Demand, HostOp, HostResult},
route::Route,
};
use tokio::sync::{mpsc, oneshot};
/// The machine's own failures, distinct from anything the provider call reports.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
@ -82,12 +80,27 @@ where
}
}
pub async fn emit(&self, event: CallEvent) -> Result<(), R::Error> {
pub async fn emit(&self, event: MachineEvent) -> Result<(), R::Error> {
match self.invoke(HostOp::Emit(event)).await? {
HostResult::Emitted => Ok(()),
_ => Err(MachineFault::Mismatch.into()),
}
}
pub async fn open(&self, head: R::StreamHead) -> Result<Demand, R::Error> {
self.demand(HostOp::Open(head)).await
}
pub async fn deliver(&self, chunk: R::Chunk) -> Result<Demand, R::Error> {
self.demand(HostOp::Deliver(chunk)).await
}
async fn demand(&self, op: HostOp<R>) -> Result<Demand, R::Error> {
match self.invoke(op).await? {
HostResult::Demand(demand) => Ok(demand),
_ => Err(MachineFault::Mismatch.into()),
}
}
}
enum Execution<R: Route> {

View file

@ -6,4 +6,9 @@ pub trait Route: Send + Sync + 'static {
type Error: Clone + Send + Sync + 'static;
type Op: Send + 'static;
type OpResult: Send + 'static;
/// One piece of a streamed response, handed to the caller as it arrives. A route
/// that never streams uses `Infallible`.
type Chunk: Send + 'static;
/// What the route knows once a streamed response starts, before its first chunk.
type StreamHead: Send + 'static;
}

View file

@ -11,6 +11,7 @@ where
H: Host<M::Route>,
{
let start_time = epoch_seconds();
let _ = host.emit(&CallEvent::Started { start_time }).await;
let mut result = None;
let outcome = loop {
let step = match machine.resume(result.take()).await {
@ -24,7 +25,12 @@ where
.before_send(*wire, &context)
.await
.map(|wire| HostResult::BeforeSend(Box::new(wire))),
HostOp::Emit(event) => host.emit(&event).await.map(|()| HostResult::Emitted),
HostOp::Emit(event) => host
.emit(&CallEvent::Machine(event))
.await
.map(|()| HostResult::Emitted),
HostOp::Open(head) => host.open(head).await.map(HostResult::Demand),
HostOp::Deliver(chunk) => host.deliver(chunk).await.map(HostResult::Demand),
};
match answer {
Ok(answer) => result = Some(answer),
@ -60,6 +66,8 @@ mod tests {
type Error = &'static str;
type Op = &'static str;
type OpResult = ();
type Chunk = std::convert::Infallible;
type StreamHead = std::convert::Infallible;
}
struct Scripted {
@ -102,6 +110,7 @@ mod tests {
async fn emit(&self, event: &CallEvent) -> Result<(), &'static str> {
self.seen.lock().unwrap().push(match event {
CallEvent::Started { .. } => "started".into(),
CallEvent::Succeeded { .. } => "succeeded".into(),
CallEvent::Failed { .. } => "failed".into(),
other => format!("{other:?}"),
@ -124,7 +133,7 @@ mod tests {
assert_eq!(outcome, Ok(()));
assert_eq!(
*host.seen.lock().unwrap(),
["route:project", "route:send", "succeeded"]
["started", "route:project", "route:send", "succeeded"]
);
}
@ -133,7 +142,7 @@ mod tests {
let host = Recording::default();
let outcome = run(scripted(&[], Err("boom")), &host).await;
assert_eq!(outcome, Err("boom"));
assert_eq!(*host.seen.lock().unwrap(), ["failed"]);
assert_eq!(*host.seen.lock().unwrap(), ["started", "failed"]);
let host = Recording {
fail: Some("send"),
@ -143,7 +152,35 @@ mod tests {
assert_eq!(outcome, Err("host failed"));
assert_eq!(
*host.seen.lock().unwrap(),
["route:project", "route:send", "failed"]
["started", "route:project", "route:send", "failed"]
);
}
struct StartTimes(Mutex<Vec<f64>>);
impl Host<Unit> for StartTimes {
async fn route(&self, _: &'static str) -> Result<(), &'static str> {
Ok(())
}
async fn emit(&self, event: &CallEvent) -> Result<(), &'static str> {
if let CallEvent::Started { start_time }
| CallEvent::Succeeded {
timing: Timing { start_time, .. },
} = event
{
self.0.lock().unwrap().push(*start_time);
}
Err("observer failed")
}
}
#[tokio::test]
async fn started_opens_the_call_at_the_terminal_start_time_and_cannot_fail_it() {
let host = StartTimes(Mutex::default());
assert_eq!(run(scripted(&["project"], Ok(())), &host).await, Ok(()));
let times = host.0.lock().unwrap();
assert_eq!(times.len(), 2);
assert_eq!(times[0], times[1]);
}
}

View file

@ -15,7 +15,7 @@ litellm-auth.workspace = true
litellm-auth-aws.workspace = true
litellm-auth-azure.workspace = true
litellm-auth-gcp.workspace = true
litellm-callbacks.workspace = true
litellm-host.workspace = true
litellm-framing.workspace = true
base64.workspace = true
bytes.workspace = true

View file

@ -150,7 +150,7 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOcrConfig {
api_key: inputs.api_key.and_then(|key| {
inputs
.dynamic_api_key
.filter(|value| !value.value().is_empty())
.filter(|value| !value.value().expose().is_empty())
.or(Some(key))
}),
api_base: inputs.api_base.and_then(|base| {
@ -592,12 +592,17 @@ impl AzureDocumentIntelligenceOcrConfig {
)?;
return Ok(connection.extra_headers.clone());
}
let key = nonblank(connection.api_key.clone())
.map(|value| Sourced::new(value, connection.api_key_source))
.or_else(|| {
nonblank(self.get_api_key_env_var().and_then(env_lookup))
.map(|value| Sourced::new(value, InputSource::Environment))
});
let key = nonblank(
connection
.api_key
.as_ref()
.map(|key| key.expose().to_string()),
)
.map(|value| Sourced::new(value, connection.api_key_source))
.or_else(|| {
nonblank(self.get_api_key_env_var().and_then(env_lookup))
.map(|value| Sourced::new(value, InputSource::Environment))
});
if let Some(key) = key {
super::super::common_utils::validate_destination(connection, key.source())?;
return Ok(
@ -796,7 +801,7 @@ mod tests {
#[tokio::test]
async fn request_endpoint_accepts_request_owned_key() {
let connection = OcrConnection {
api_key: Some("request-key".into()),
api_key: Some(litellm_auth::SecretValue::new("request-key")),
api_key_source: InputSource::Request,
api_base: Some("https://request.example".into()),
api_base_source: InputSource::Request,

View file

@ -142,12 +142,17 @@ impl AzureAiOcrConfig {
super::common_utils::validate_destination(connection, connection.extra_headers_source)?;
return Ok(connection.extra_headers.clone());
}
let key = nonblank(connection.api_key.clone())
.map(|value| Sourced::new(value, connection.api_key_source))
.or_else(|| {
nonblank(self.get_api_key_env_var().and_then(env_lookup))
.map(|value| Sourced::new(value, InputSource::Environment))
});
let key = nonblank(
connection
.api_key
.as_ref()
.map(|key| key.expose().to_string()),
)
.map(|value| Sourced::new(value, connection.api_key_source))
.or_else(|| {
nonblank(self.get_api_key_env_var().and_then(env_lookup))
.map(|value| Sourced::new(value, InputSource::Environment))
});
if let Some(key) = key {
super::common_utils::validate_destination(connection, key.source())?;
return Ok(bearer_headers(connection, key.value()));
@ -196,7 +201,7 @@ mod tests {
#[fixture]
fn connection() -> OcrConnection {
OcrConnection {
api_key: Some("request-key".into()),
api_key: Some(litellm_auth::SecretValue::new("request-key")),
api_base: Some("https://example.com".into()),
..Default::default()
}
@ -288,7 +293,7 @@ mod tests {
#[tokio::test]
async fn request_endpoint_accepts_request_owned_key() {
let connection = OcrConnection {
api_key: Some("request-key".into()),
api_key: Some(litellm_auth::SecretValue::new("request-key")),
api_key_source: InputSource::Request,
api_base: Some("https://request.example".into()),
api_base_source: InputSource::Request,

View file

@ -102,6 +102,17 @@ pub enum Error {
Headers(#[from] crate::custom_httpx::http_handler::HeaderError),
}
impl From<litellm_host::machine::MachineFault> for Error {
fn from(fault: litellm_host::machine::MachineFault) -> Self {
use litellm_host::machine::MachineFault;
Self::InvalidRequest(match fault {
MachineFault::Abandoned => "OCR host driver was abandoned".into(),
MachineFault::Protocol(message) => format!("OCR {message}"),
MachineFault::Mismatch => "invalid OCR host operation result".into(),
})
}
}
impl From<litellm_core_utils::call_arguments::ArgumentError> for Error {
fn from(error: litellm_core_utils::call_arguments::ArgumentError) -> Self {
Self::RequestField {

View file

@ -1,6 +1,6 @@
use std::{collections::BTreeMap, future::Future, time::Duration};
use litellm_auth::{InputSource, Sourced, TokenProviderHandle};
use litellm_auth::{InputSource, SecretValue, Sourced, TokenProviderHandle};
use litellm_core_utils::{
call_arguments::CallArguments,
serde_compat::{FiniteF64, LaxI64},
@ -90,21 +90,22 @@ pub enum OcrResponseFormat {
#[derive(Clone, Default)]
pub struct OcrCredentialInputs {
pub api_key: Option<Sourced<String>>,
pub dynamic_api_key: Option<Sourced<String>>,
pub api_key: Option<Sourced<SecretValue>>,
pub dynamic_api_key: Option<Sourced<SecretValue>>,
pub api_base: Option<Sourced<String>>,
pub dynamic_api_base: Option<Sourced<String>>,
}
impl OcrCredentialInputs {
pub fn new(
api_key: Option<String>,
api_key: Option<SecretValue>,
api_key_source: InputSource,
api_base: Option<String>,
api_base_source: InputSource,
) -> Self {
Self {
api_key: nonblank(api_key).map(|value| Sourced::new(value, api_key_source)),
api_key: nonblank(api_key.as_ref().map(|key| key.expose().to_string()))
.map(|value| Sourced::new(SecretValue::new(value), api_key_source)),
dynamic_api_key: None,
api_base: nonblank(api_base).map(|value| Sourced::new(value, api_base_source)),
dynamic_api_base: None,
@ -159,7 +160,7 @@ fn nonblank(value: Option<String>) -> Option<String> {
#[derive(Clone)]
pub struct OcrConnection {
pub api_key: Option<String>,
pub api_key: Option<SecretValue>,
pub api_key_source: InputSource,
pub api_base: Option<String>,
pub api_base_source: InputSource,
@ -209,7 +210,7 @@ impl Default for OcrConnection {
#[derive(Clone, Default)]
pub struct ResolvedOcrCredentials {
pub api_key: Option<Sourced<String>>,
pub api_key: Option<Sourced<SecretValue>>,
pub api_base: Option<Sourced<String>>,
}
@ -428,7 +429,7 @@ pub trait BaseOcrConfig: Send + Sync + Sized + 'static {
ResolvedOcrCredentials {
api_key: inputs
.dynamic_api_key
.filter(|value| !value.value().is_empty())
.filter(|value| !value.value().expose().is_empty())
.or(inputs.api_key),
api_base: inputs
.dynamic_api_base

View file

@ -179,8 +179,8 @@ impl CohereParseConfig {
}
let key = connection
.api_key
.as_deref()
.map(str::trim)
.as_ref()
.map(|key| key.expose().trim())
.filter(|key| !key.is_empty())
.map(str::to_string)
.or_else(|| {
@ -718,7 +718,7 @@ mod tests {
assert!(matches!(
CohereParseConfig.resolve_headers(
&OcrConnection {
api_key: Some(" ".into()),
api_key: Some(litellm_auth::SecretValue::new(" ")),
..Default::default()
},
&|_| None,

View file

@ -3,9 +3,9 @@ use std::{sync::OnceLock, time::Duration};
use bytes::{Bytes, BytesMut};
use futures_util::future::BoxFuture;
use litellm_auth_gcp::VertexAuth;
use litellm_callbacks::event::{Passthrough, WireRequest};
use litellm_host::event::WireRequest;
use serde::{Serialize, de::DeserializeOwned};
use serde_json::{Map, Value};
use serde_json::Value;
use crate::{
base_llm::ocr::{
@ -26,11 +26,7 @@ use crate::{
/// The route's view of one call, handed to provider code that has to reach the
/// caller's hooks mid-flight (guardrails on the outgoing body, raw response events).
pub trait CallHooks<E>: Send + Sync {
fn before_send(
&self,
wire: WireRequest,
passthrough_fields: Passthrough,
) -> BoxFuture<'_, Result<WireRequest, E>>;
fn before_send(&self, wire: WireRequest) -> BoxFuture<'_, Result<WireRequest, E>>;
fn response_received<'a>(&'a self, body: &'a [u8]) -> BoxFuture<'a, Result<(), E>>;
}
@ -231,9 +227,8 @@ pub async fn transform_request_body<C: BaseOcrConfig, B: Serialize>(
config.get_supported_ocr_params(&request.model),
)?;
config.validate_request_body(&composed)?;
let passthrough_fields = Passthrough::unchanged(&caller_inputs(request)?, &composed);
let changed = hooks
.before_send(wire_request(url, headers, composed), passthrough_fields)
.before_send(wire_request(url, headers, composed))
.await?;
if !changed.body.is_object() {
return Err(Error::RequestField {
@ -252,21 +247,6 @@ fn wire_request(url: &str, headers: &[(String, String)], body: Value) -> WireReq
}
}
fn caller_inputs(request: &PreparedOcrRequest) -> Result<Map<String, Value>, Error> {
let document = request
.caller_document
.then(|| serde_json::to_value(&request.document))
.transpose()
.map_err(|_| Error::RequestField {
path: "document".into(),
})?;
let params: Map<String, Value> = request.optional_params.clone().into();
Ok(params
.into_iter()
.chain(document.map(|document| ("document".to_string(), document)))
.collect())
}
pub fn build_http_request<B: Serialize>(
client: &OcrClient,
request: &PreparedOcrRequest,
@ -294,9 +274,7 @@ pub async fn guardrail_document(
let body = serde_json::to_value(&request.document).map_err(|_| Error::RequestField {
path: "document".into(),
})?;
let changed = hooks
.before_send(wire_request(url, headers, body), Passthrough::default())
.await?;
let changed = hooks.before_send(wire_request(url, headers, body)).await?;
let document = decode_request_value(changed.body, "guardrail.document")?;
Ok((document, changed.headers))
}

View file

@ -135,8 +135,8 @@ impl MistralOcrConfig {
}
let api_key = connection
.api_key
.as_deref()
.map(str::trim)
.as_ref()
.map(|key| key.expose().trim())
.filter(|key| !key.is_empty())
.map(str::to_string)
.or_else(|| {
@ -212,7 +212,7 @@ mod tests {
#[default(vec![])] extra_headers: Vec<(String, String)>,
) -> OcrConnection {
OcrConnection {
api_key: api_key.map(str::to_string),
api_key: api_key.map(litellm_auth::SecretValue::new),
extra_headers,
..OcrConnection::default()
}

View file

@ -442,8 +442,8 @@ fn resolve_headers(
}
let api_key = connection
.api_key
.as_deref()
.map(str::trim)
.as_ref()
.map(|key| key.expose().trim())
.filter(|key| !key.is_empty())
.map(str::to_string)
.or_else(|| {
@ -629,7 +629,7 @@ mod tests {
#[test]
fn explicit_key_precedes_environment_key() {
let connection = OcrConnection {
api_key: Some("passed-key".into()),
api_key: Some(litellm_auth::SecretValue::new("passed-key")),
..Default::default()
};
let headers = resolve_headers(&connection, &|_| Some("env-key".into())).unwrap();
@ -639,7 +639,7 @@ mod tests {
#[test]
fn blank_explicit_key_uses_environment_key() {
let connection = OcrConnection {
api_key: Some(" ".into()),
api_key: Some(litellm_auth::SecretValue::new(" ")),
..Default::default()
};
let headers = resolve_headers(&connection, &|_| Some(" env-key ".into())).unwrap();

View file

@ -134,7 +134,10 @@ impl VertexAiOcrConfig {
.vertex_auth()
.validate_environment(
connection.extra_headers.clone(),
connection.api_key.as_deref(),
connection
.api_key
.as_ref()
.map(litellm_auth::SecretValue::expose),
config,
&credential_env,
)

View file

@ -1,7 +1,7 @@
- Target invariants, not completion claims; these supersede older conflicting bridge guidance
- Keep this crate the product-specific PyO3 consumer of `litellm-host-python`
- Own registration, input projection, the route host and the caller callables it answers operations with (file readers, token providers), public response/error construction and the per-call composition of machine, route host and callback contract
- Legacy callback sharing (the caller's args, kwargs and request object, body/header roots, `passthrough_fields` re-aliasing) lives in `litellm-callbacks-legacy` behind `PublicCall` and `run_legacy_call`; the bridge hands the public call over and keeps no copy
- Legacy callback sharing (the caller's args, kwargs and request object, body/header roots, re-aliasing unchanged body keys) lives in `litellm-callbacks-legacy` behind `PublicCall` and `run_legacy_call`; the bridge hands the public call over and keeps no copy
- Value-oriented execution, sync waiting, nested-runtime checks, signal polling and panic containment live in `litellm-host-python`; native async work uses `pyo3-async-runtimes`, Serde output uses `Pythonized<T>`
- Core owns typed native state, the route machine, provider preparation/I/O and normalization; the host driver owns terminal events; the legacy adapter in `litellm-callbacks-legacy` owns `Logging` dispatch policy
- Python, Rust SDK and gateway use one lifecycle-bearing core route entrypoint; provider helpers stay private, never bridge-accessible transport drivers

View file

@ -18,10 +18,6 @@ pub(crate) struct RouteOptions {
pub(crate) timeout: Option<Duration>,
}
pub(crate) fn body_argument(value: &Bound<'_, PyAny>) -> PyResult<Map<String, Value>> {
required_object("body", from_py_argument(value)?)
}
pub(crate) fn messages_argument(value: &Bound<'_, PyAny>) -> PyResult<Vec<Value>> {
match from_py_argument(value)? {
Value::Array(values) => Ok(values),
@ -192,18 +188,6 @@ mod tests {
json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}])
);
let body = py
.eval(
c"{'model': 'claude', 'metadata': {'user': '1'}}",
None,
None,
)
.unwrap();
assert_eq!(
Value::Object(body_argument(&body).unwrap()),
json!({"model": "claude", "metadata": {"user": "1"}})
);
let params = py.eval(c"{'temperature': 0.2}", None, None).unwrap();
assert_eq!(
optional_params_argument(&params).unwrap(),

View file

@ -1,88 +0,0 @@
use litellm_core::messages::{Error, messages as run_messages, types::MessagesRequest};
use litellm_host_python::{run_async, run_sync};
use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse;
use pyo3::prelude::*;
use serde_json::{Map, Value};
use crate::{
errors::messages_error_to_pyerr,
marshal::{RouteOptions, body_argument, extra_headers_argument, optional_timeout},
};
async fn execute(
body: Map<String, Value>,
options: RouteOptions,
) -> Result<AnthropicMessagesResponse, Error> {
let RouteOptions {
model,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout,
} = options;
run_messages(MessagesRequest {
model: &model,
body: Value::Object(body),
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
timeout,
})
.await
}
#[pyfunction]
#[pyo3(signature = (model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None))]
#[expect(
clippy::too_many_arguments,
reason = "one parameter per Python keyword"
)]
pub(crate) fn messages(
py: Python<'_>,
model: String,
#[pyo3(from_py_with = body_argument)] body: Map<String, Value>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
#[pyo3(from_py_with = extra_headers_argument)] extra_headers: Option<Map<String, Value>>,
timeout_seconds: Option<f64>,
) -> PyResult<Py<PyAny>> {
let options = RouteOptions {
model,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout: optional_timeout(timeout_seconds),
};
run_sync(py, execute(body, options), messages_error_to_pyerr)
}
#[pyfunction]
#[pyo3(signature = (model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None))]
#[expect(
clippy::too_many_arguments,
reason = "one parameter per Python keyword"
)]
pub(crate) fn amessages<'py>(
py: Python<'py>,
model: String,
#[pyo3(from_py_with = body_argument)] body: Map<String, Value>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
#[pyo3(from_py_with = extra_headers_argument)] extra_headers: Option<Map<String, Value>>,
timeout_seconds: Option<f64>,
) -> PyResult<Bound<'py, PyAny>> {
let options = RouteOptions {
model,
api_key,
api_base,
custom_llm_provider,
extra_headers,
timeout: optional_timeout(timeout_seconds),
};
run_async(py, execute(body, options), messages_error_to_pyerr)
}

View file

@ -0,0 +1,186 @@
use bytes::Bytes;
use litellm_core::messages::{
Error,
route::{Messages, MessagesCall, MessagesOp, MessagesOpResult, MessagesOutput},
};
use litellm_host_python::{InvokeError, RouteHost, from_py, lookup, to_py};
use litellm_llms::custom_httpx::transport::Error as TransportError;
use pyo3::{
exceptions::{PyException, PyValueError},
gc::{PyTraverseError, PyVisit},
prelude::*,
types::{PyBytes, PyDict},
};
use serde_json::{Map, Value};
use crate::{
errors::{RustUpstreamError, messages_error_to_pyerr},
marshal::{optional_timeout, python_timeout_seconds},
};
/// The Anthropic Messages body fields a caller may pass besides `model` and `messages`,
/// as `AnthropicMessagesRequestOptionalParams` declares them.
const BODY_FIELDS: [&str; 20] = [
"max_tokens",
"metadata",
"stop_sequences",
"stream",
"system",
"temperature",
"thinking",
"tool_choice",
"tools",
"top_k",
"inference_geo",
"top_p",
"mcp_servers",
"context_management",
"container",
"output_format",
"speed",
"output_config",
"cache_control",
"reasoning_effort",
];
/// The Python side of the Messages route: projects the prepared arguments and builds the
/// public response, chunks and exceptions.
pub(super) struct MessagesRouteHost {
request: Py<PyAny>,
}
impl MessagesRouteHost {
pub(super) fn new(request: Py<PyAny>) -> Self {
Self { request }
}
fn project(&self, py: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult<MessagesCall> {
let request = self.request.bind(py);
let argument = |name: &str| -> PyResult<Option<Bound<'_, PyAny>>> {
Ok(lookup(arguments, request, name)?.filter(|value| !value.is_none()))
};
let string = |name: &str| -> PyResult<Option<String>> {
argument(name)?.map(|value| value.extract()).transpose()
};
let model = string("model")?.ok_or_else(|| PyValueError::new_err("model is required"))?;
let messages =
argument("messages")?.ok_or_else(|| PyValueError::new_err("messages is required"))?;
let fields = BODY_FIELDS
.iter()
.filter_map(|name| match argument(name) {
Ok(Some(value)) => Some(from_py(&value).map(|value| ((*name).to_string(), value))),
Ok(None) => None,
Err(error) => Some(Err(error)),
})
.collect::<PyResult<Vec<(String, Value)>>>()?;
let body = [
("model".to_string(), Value::String(model.clone())),
("messages".to_string(), from_py(&messages)?),
]
.into_iter()
.chain(fields)
.collect::<Map<String, Value>>();
let timeout = argument("timeout")?
.map(|value| python_timeout_seconds(py, value.unbind()))
.transpose()?
.flatten();
Ok(MessagesCall {
model,
body,
api_key: string("api_key")?,
api_base: string("api_base")?,
custom_llm_provider: string("custom_llm_provider")?,
extra_headers: argument("extra_headers")?
.map(|value| from_py(&value))
.transpose()?,
timeout: optional_timeout(timeout),
})
}
fn provider(&self, py: Python<'_>) -> String {
self.request
.bind(py)
.getattr("custom_llm_provider")
.and_then(|value| value.extract::<Option<String>>())
.ok()
.flatten()
.unwrap_or_else(|| "anthropic".into())
}
fn map_failure(&self, py: Python<'_>, error: PyErr) -> PyErr {
if !error.is_instance_of::<PyException>(py) {
return error;
}
let mapped = py
.import("litellm.rust_bridge.messages.route_host")
.and_then(|module| module.getattr("map_failure"))
.and_then(|map| map.call1((error.value(py), self.request.bind(py), self.provider(py))))
.and_then(|mapped| {
mapped
.extract::<Py<pyo3::exceptions::PyBaseException>>()
.map_err(PyErr::from)
});
match mapped {
Ok(mapped) => PyErr::from_value(mapped.into_bound(py).into_any()),
Err(_) => error,
}
}
}
impl RouteHost for MessagesRouteHost {
type Route = Messages;
type Failure = PyErr;
fn invoke(
&mut self,
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
op: MessagesOp,
) -> Result<MessagesOpResult, InvokeError<Error>> {
match op {
MessagesOp::ProjectRequest => self
.project(py, arguments)
.map(|call| MessagesOpResult::Request(Box::new(call)))
.map_err(|error| InvokeError::Python(self.map_failure(py, error))),
}
}
fn complete(&mut self, py: Python<'_>, response: MessagesOutput) -> PyResult<Py<PyAny>> {
match response {
MessagesOutput::Message(message) => py
.import("litellm.rust_bridge.messages.route_host")?
.getattr("response")?
.call1((to_py(py, message.as_ref())?,))
.map(Bound::unbind),
MessagesOutput::Streamed => Ok(py.None()),
}
}
fn chunk(&mut self, py: Python<'_>, chunk: Bytes) -> PyResult<Py<PyAny>> {
Ok(PyBytes::new(py, &chunk).into_any().unbind())
}
fn classify(&self, py: Python<'_>, error: Error) -> PyResult<PyErr> {
let native = match error {
Error::Transport(TransportError::Http { status, body }) => {
let error = RustUpstreamError::new_err((status, body));
error
.value(py)
.setattr("headers", Vec::<(String, String)>::new())?;
error
}
other => messages_error_to_pyerr(other),
};
Ok(self.map_failure(py, native))
}
fn host_error(error: &PyErr) -> Error {
Error::InvalidRequest(error.to_string())
}
fn close(&mut self, _: Python<'_>) {}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.request)
}
}

View file

@ -0,0 +1,68 @@
mod host;
use host::MessagesRouteHost;
use litellm_callbacks_legacy::{LegacySurface, PassThroughStream, PublicCall, run_legacy_call};
use litellm_core::messages::route::{messages_machine, supports};
use pyo3::{
prelude::*,
types::{PyDict, PyTuple},
};
use crate::errors::RustBridgeDeclined;
const SURFACE: LegacySurface = LegacySurface {
call_type: "anthropic_messages",
input_description: "Messages",
stream: Some(PassThroughStream {
url_route: "/v1/messages",
endpoint_type: "anthropic",
}),
};
fn run_messages(
py: Python<'_>,
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
asynchronous: bool,
) -> PyResult<Py<PyAny>> {
let model: String = request.getattr("model")?.extract()?;
let provider: Option<String> = request.getattr("custom_llm_provider")?.extract()?;
let stream = request
.getattr("stream")?
.extract::<Option<bool>>()?
.unwrap_or(false);
if !supports(&model, provider.as_deref(), stream) {
return Err(RustBridgeDeclined::new_err(
"the Rust Messages route does not serve this provider",
));
}
run_legacy_call(
py,
SURFACE,
PublicCall::capture(&request, &args, &kwargs)?,
messages_machine(),
MessagesRouteHost::new(request.unbind()),
asynchronous,
)
}
#[pyfunction]
pub(crate) fn messages(
py: Python<'_>,
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
run_messages(py, request, args, kwargs, false)
}
#[pyfunction]
pub(crate) fn amessages(
py: Python<'_>,
request: Bound<'_, PyAny>,
args: Bound<'_, PyTuple>,
kwargs: Bound<'_, PyDict>,
) -> PyResult<Py<PyAny>> {
run_messages(py, request, args, kwargs, true)
}

View file

@ -22,11 +22,6 @@ mod tests {
"atranscription",
"(model, audio, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None)",
),
(
"messages",
"amessages",
"(model, body, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None)",
),
(
"chat_completions",
"achat_completions",
@ -113,25 +108,6 @@ value = Broken()
);
assert_eq!(async_chat_error.to_string(), sync_chat_error.to_string());
let invalid_body = PyList::empty(py);
let sync_messages_error = module
.getattr("messages")
.and_then(|function| function.call1(("model", &invalid_body)))
.expect_err("sync Messages should reject a non-dict body");
let async_messages_error = module
.getattr("amessages")
.and_then(|function| function.call1(("model", &invalid_body)))
.expect_err("async Messages should reject a non-dict body");
assert_eq!(
sync_messages_error.to_string(),
"ValueError: body must be a dict"
);
assert_eq!(
async_messages_error.to_string(),
sync_messages_error.to_string()
);
let invalid_headers = PyList::empty(py);
let kwargs = PyDict::new(py);
kwargs
@ -193,13 +169,6 @@ value = Broken()
headers_kwargs
.set_item("extra_headers", &invalid)
.expect("kwargs should accept extra_headers");
let invalid_body = PyList::empty(py);
let error = module
.getattr("messages")
.and_then(|function| function.call(("model", &invalid_body), Some(&headers_kwargs)))
.expect_err("body should be validated before headers");
assert_eq!(error.to_string(), "ValueError: body must be a dict");
let invalid_payload =
PyModule::new(py, "invalid_payload").expect("invalid payload should be created");
let error = module

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