mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-21 00:21:49 +00:00
Merge remote-tracking branch 'origin/main' into litellm_vertex_chirp3_streaming_stt
This commit is contained in:
commit
dc11c34e3a
152 changed files with 5504 additions and 9075 deletions
50
.github/scripts/_agent_shin_actions.py
vendored
50
.github/scripts/_agent_shin_actions.py
vendored
|
|
@ -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)
|
||||
211
.github/scripts/agent_shin_shared.py
vendored
211
.github/scripts/agent_shin_shared.py
vendored
|
|
@ -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()
|
||||
573
.github/scripts/close_low_quality_prs.py
vendored
573
.github/scripts/close_low_quality_prs.py
vendored
|
|
@ -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())
|
||||
282
.github/scripts/triage-requirements.txt
vendored
282
.github/scripts/triage-requirements.txt
vendored
|
|
@ -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
|
||||
1797
.github/scripts/triage_with_llm.py
vendored
1797
.github/scripts/triage_with_llm.py
vendored
File diff suppressed because it is too large
Load diff
92
.github/workflows/close_low_quality_prs.yml
vendored
92
.github/workflows/close_low_quality_prs.yml
vendored
|
|
@ -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[@]}"
|
||||
|
|
@ -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"
|
||||
172
.github/workflows/triage_reconsider.yml
vendored
172
.github/workflows/triage_reconsider.yml
vendored
|
|
@ -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)"
|
||||
91
litellm-rust/Cargo.lock
generated
91
litellm-rust/Cargo.lock
generated
|
|
@ -2027,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]]
|
||||
|
|
@ -2056,8 +2050,8 @@ dependencies = [
|
|||
"futures-util",
|
||||
"litellm-auth",
|
||||
"litellm-auth-aws",
|
||||
"litellm-callbacks",
|
||||
"litellm-core-utils",
|
||||
"litellm-host",
|
||||
"litellm-llms",
|
||||
"litellm-types",
|
||||
"mime_guess",
|
||||
|
|
@ -2110,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",
|
||||
|
|
@ -2139,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",
|
||||
|
|
@ -2598,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"
|
||||
|
|
@ -2679,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"
|
||||
|
|
@ -2842,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"
|
||||
|
|
@ -3241,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"
|
||||
|
|
@ -4096,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"
|
||||
|
|
@ -4209,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"
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
118
litellm-rust/crates/callbacks-legacy/python_contract.json
Normal file
118
litellm-rust/crates/callbacks-legacy/python_contract.json
Normal 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"
|
||||
]
|
||||
}
|
||||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
}
|
||||
|
|
|
|||
183
litellm-rust/crates/callbacks-legacy/src/legacy_python.rs
Normal file
183
litellm-rust/crates/callbacks-legacy/src/legacy_python.rs
Normal 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);
|
||||
}
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
}
|
||||
});
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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())]);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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()));
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
196
litellm-rust/crates/core/src/messages/route.rs
Normal file
196
litellm-rust/crates/core/src/messages/route.rs
Normal 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)
|
||||
}
|
||||
|
|
@ -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> {
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
},
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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(|| {
|
||||
|
|
|
|||
|
|
@ -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!(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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!(
|
||||
|
|
|
|||
|
|
@ -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>>,
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
152
litellm-rust/crates/core/tests/ocr/document.rs
Normal file
152
litellm-rust/crates/core/tests/ocr/document.rs
Normal 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)
|
||||
);
|
||||
}
|
||||
|
|
@ -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)
|
||||
);
|
||||
}
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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":[]}}"#);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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>`
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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>;
|
||||
|
|
|
|||
51
litellm-rust/crates/host-python/src/argument.rs
Normal file
51
litellm-rust/crates/host-python/src/argument.rs
Normal 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());
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
@ -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"]
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
|
|
@ -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"] }
|
||||
76
litellm-rust/crates/host/src/event.rs
Normal file
76
litellm-rust/crates/host/src/event.rs
Normal 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,
|
||||
},
|
||||
}
|
||||
|
|
@ -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) }
|
||||
}
|
||||
}
|
||||
|
|
@ -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.
|
||||
|
||||
|
|
@ -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 {
|
||||
|
|
@ -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;
|
||||
|
||||
|
|
@ -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> {
|
||||
|
|
@ -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;
|
||||
}
|
||||
|
|
@ -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]);
|
||||
}
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(¶ms).unwrap(),
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
186
litellm-rust/crates/python-bridge/src/routes/messages/host.rs
Normal file
186
litellm-rust/crates/python-bridge/src/routes/messages/host.rs
Normal 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)
|
||||
}
|
||||
}
|
||||
68
litellm-rust/crates/python-bridge/src/routes/messages/mod.rs
Normal file
68
litellm-rust/crates/python-bridge/src/routes/messages/mod.rs
Normal 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)
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
use litellm_auth::ResolvedCredential;
|
||||
use litellm_core::ocr::route::{Ocr, OcrOp, OcrOpResult};
|
||||
use litellm_host_python::{RouteHost, missing_state, to_py};
|
||||
use litellm_host_python::{InvokeError, RouteHost, missing_state, to_py};
|
||||
use litellm_llms::base_llm::ocr::{error::Error, transformation::LiteLLMOcrResponse};
|
||||
use pyo3::{
|
||||
exceptions::PyBaseException,
|
||||
exceptions::{PyBaseException, PyException},
|
||||
gc::{PyTraverseError, PyVisit},
|
||||
prelude::*,
|
||||
types::PyDict,
|
||||
|
|
@ -57,12 +57,8 @@ impl OcrRouteHost {
|
|||
.ok_or_else(missing_state)?
|
||||
.acquire(py)
|
||||
}
|
||||
}
|
||||
|
||||
impl RouteHost for OcrRouteHost {
|
||||
type Route = Ocr;
|
||||
|
||||
fn invoke(
|
||||
fn answer(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
arguments: &Bound<'_, PyDict>,
|
||||
|
|
@ -88,6 +84,40 @@ impl RouteHost for OcrRouteHost {
|
|||
}
|
||||
}
|
||||
|
||||
fn map_failure(&self, py: Python<'_>, error: PyErr) -> PyErr {
|
||||
if !error.is_instance_of::<PyException>(py) {
|
||||
return error;
|
||||
}
|
||||
let provider = match &self.data {
|
||||
OcrHostData::Projected(handles) => handles.provider,
|
||||
_ => "",
|
||||
};
|
||||
let mapped = py
|
||||
.import("litellm.rust_bridge.ocr.route_host")
|
||||
.and_then(|module| module.getattr("map_failure"))
|
||||
.and_then(|map| map.call1((error.value(py), self.request.bind(py), provider)))
|
||||
.and_then(|mapped| mapped.extract::<Py<PyBaseException>>().map_err(PyErr::from));
|
||||
match mapped {
|
||||
Ok(mapped) => PyErr::from_value(mapped.into_bound(py).into_any()),
|
||||
Err(_) => error,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl RouteHost for OcrRouteHost {
|
||||
type Route = Ocr;
|
||||
type Failure = PyErr;
|
||||
|
||||
fn invoke(
|
||||
&mut self,
|
||||
py: Python<'_>,
|
||||
arguments: &Bound<'_, PyDict>,
|
||||
op: OcrOp,
|
||||
) -> Result<OcrOpResult, InvokeError<Error>> {
|
||||
self.answer(py, arguments, op)
|
||||
.map_err(|error| InvokeError::Python(self.map_failure(py, error)))
|
||||
}
|
||||
|
||||
fn complete(&mut self, py: Python<'_>, response: LiteLLMOcrResponse) -> PyResult<Py<PyAny>> {
|
||||
py.import("litellm.rust_bridge.ocr.route_host")?
|
||||
.getattr("response")?
|
||||
|
|
@ -95,27 +125,18 @@ impl RouteHost for OcrRouteHost {
|
|||
.map(Bound::unbind)
|
||||
}
|
||||
|
||||
fn native_error(error: Error) -> PyErr {
|
||||
ocr_error_to_pyerr(error)
|
||||
fn chunk(&mut self, _: Python<'_>, chunk: std::convert::Infallible) -> PyResult<Py<PyAny>> {
|
||||
match chunk {}
|
||||
}
|
||||
|
||||
fn classify(&self, py: Python<'_>, error: Error) -> PyResult<PyErr> {
|
||||
Ok(self.map_failure(py, ocr_error_to_pyerr(error)))
|
||||
}
|
||||
|
||||
fn host_error(error: &PyErr) -> Error {
|
||||
Error::InvalidRequest(error.to_string())
|
||||
}
|
||||
|
||||
fn map_failure(&self, py: Python<'_>, error: &PyErr) -> PyResult<PyErr> {
|
||||
let provider = match &self.data {
|
||||
OcrHostData::Projected(handles) => handles.provider,
|
||||
_ => "",
|
||||
};
|
||||
let mapped: Py<PyBaseException> = py
|
||||
.import("litellm.rust_bridge.ocr.route_host")?
|
||||
.getattr("map_failure")?
|
||||
.call1((error.value(py), self.request.bind(py), provider))?
|
||||
.extract()?;
|
||||
Ok(PyErr::from_value(mapped.into_bound(py).into_any()))
|
||||
}
|
||||
|
||||
fn close(&mut self, _: Python<'_>) {
|
||||
self.data = OcrHostData::Released;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -15,6 +15,7 @@ use pyo3::{
|
|||
const SURFACE: LegacySurface = LegacySurface {
|
||||
call_type: "ocr",
|
||||
input_description: "OCR document processing",
|
||||
stream: None,
|
||||
};
|
||||
|
||||
const ASYNC_SURFACE: LegacySurface = LegacySurface {
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
use litellm_auth::SecretValue;
|
||||
use litellm_core::ocr::{
|
||||
types::{LiteLLMOcrRequest, OcrDocumentInput},
|
||||
wire::{OcrWireRequest, consumed_optional_params, decode_document, decode_request_input},
|
||||
|
|
@ -31,7 +32,7 @@ struct OcrArguments<'a, 'py> {
|
|||
|
||||
impl<'py> OcrArguments<'_, 'py> {
|
||||
fn lookup(&self, name: &str) -> PyResult<Bound<'py, PyAny>> {
|
||||
litellm_callbacks_legacy::lookup(self.kwargs, self.request, name)?
|
||||
litellm_host_python::lookup(self.kwargs, self.request, name)?
|
||||
.ok_or_else(|| PyValueError::new_err(format!("missing argument: {name}")))
|
||||
}
|
||||
|
||||
|
|
@ -47,8 +48,11 @@ impl<'py> OcrArguments<'_, 'py> {
|
|||
self.lookup("document")
|
||||
}
|
||||
|
||||
fn api_key(&self) -> PyResult<Option<String>> {
|
||||
self.lookup("api_key")?.extract()
|
||||
fn api_key(&self) -> PyResult<Option<SecretValue>> {
|
||||
Ok(self
|
||||
.lookup("api_key")?
|
||||
.extract::<Option<String>>()?
|
||||
.map(SecretValue::new))
|
||||
}
|
||||
|
||||
fn api_base(&self) -> PyResult<Option<String>> {
|
||||
|
|
|
|||
|
|
@ -325,12 +325,27 @@ def _outgoing_trace_context(parent_span: object) -> Context | None:
|
|||
return None
|
||||
|
||||
|
||||
def _propagated_context(headers: Mapping[str, str], request_context: Context) -> Context:
|
||||
"""``request_context`` when it continues the trace ``headers`` already name, else the
|
||||
caller's own context, so an explicit upstream ``traceparent`` (``x-pass-traceparent``)
|
||||
is never swapped for an unrelated trace and its ``tracestate`` survives."""
|
||||
caller: Final = extract_traceparent(headers)
|
||||
if caller is None:
|
||||
return request_context
|
||||
caller_span: Final = get_current_span(caller).get_span_context()
|
||||
request_span: Final = get_current_span(request_context).get_span_context()
|
||||
if not caller_span.is_valid or caller_span.trace_id == request_span.trace_id:
|
||||
return request_context
|
||||
return caller
|
||||
|
||||
|
||||
def inject_trace_context(headers: Mapping[str, str], parent_span: object = None) -> dict[str, str]:
|
||||
"""``headers`` plus W3C ``traceparent``/``tracestate`` for this request's span.
|
||||
|
||||
Parent preference: ``parent_span`` (the request span auth stashed on the key), then
|
||||
the anchored request root span, then the ambient active span. Only trace context is
|
||||
injected, never Baggage. Unchanged when no valid span exists anywhere.
|
||||
injected, never Baggage. Unchanged when no valid span exists anywhere. A ``traceparent``
|
||||
already in ``headers`` from a different trace is forwarded as-is instead of replaced.
|
||||
"""
|
||||
context: Final = _outgoing_trace_context(parent_span)
|
||||
if context is None:
|
||||
|
|
@ -338,7 +353,7 @@ def inject_trace_context(headers: Mapping[str, str], parent_span: object = None)
|
|||
carrier: Final = { # mutable-ok: OpenTelemetry propagator requires a mutable carrier
|
||||
key: value for key, value in headers.items() if key.lower() not in _W3C_TRACE_HEADERS
|
||||
}
|
||||
_PROPAGATOR.inject(carrier, context=context)
|
||||
_PROPAGATOR.inject(carrier, context=_propagated_context(headers, context))
|
||||
return carrier
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -224,6 +224,12 @@ _FINISH_REASON_MAP: Final[dict[str, OpenAIChatCompletionFinishReason]] = {
|
|||
"IMAGE_PROHIBITED_CONTENT": "content_filter",
|
||||
"TOO_MANY_TOOL_CALLS": "stop",
|
||||
"MALFORMED_RESPONSE": "stop",
|
||||
"NO_IMAGE": "content_filter",
|
||||
"IMAGE_RECITATION": "content_filter",
|
||||
"IMAGE_OTHER": "content_filter",
|
||||
"ESCALATION": "content_filter",
|
||||
"UNEXPECTED_TOOL_CALL": "stop",
|
||||
"MISSING_THOUGHT_SIGNATURE": "stop",
|
||||
# Zhipu GLM
|
||||
"network_error": "stop",
|
||||
"sensitive": "content_filter",
|
||||
|
|
|
|||
|
|
@ -1265,11 +1265,6 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
additional_args.get("api_base", "")
|
||||
)
|
||||
|
||||
def record_api_call_start_time(self) -> None:
|
||||
self.model_call_details["api_call_start_time"] = datetime.datetime.now()
|
||||
if self.model_call_details.get("first_api_call_start_time") is None:
|
||||
self.model_call_details["first_api_call_start_time"] = self.model_call_details["api_call_start_time"]
|
||||
|
||||
def pre_call(self, input, api_key, model=None, additional_args={}):
|
||||
# Log the exact input to the LLM API
|
||||
try:
|
||||
|
|
@ -1334,7 +1329,15 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
"LiteLLM.LoggingError: [Non-Blocking] Exception occurred while logging %s", e
|
||||
)
|
||||
|
||||
self.record_api_call_start_time()
|
||||
self.model_call_details["api_call_start_time"] = datetime.datetime.now()
|
||||
# Set-once first provider-handoff instant. api_call_start_time
|
||||
# is overwritten on every retry, so it can't measure one-time
|
||||
# preprocessing; pinning the first attempt excludes retry loops
|
||||
# + backoff. Logging object only — must NOT go into
|
||||
# litellm_params["metadata"] (caller request metadata, typed
|
||||
# Dict[str, str], echoed downstream; a datetime breaks it).
|
||||
if self.model_call_details.get("first_api_call_start_time") is None:
|
||||
self.model_call_details["first_api_call_start_time"] = self.model_call_details["api_call_start_time"]
|
||||
# Input Integration Logging -> If you want to log the fact that an attempt to call the model was made
|
||||
callbacks: Final = litellm.input_callback + (self.dynamic_input_callbacks or [])
|
||||
for callback in callbacks:
|
||||
|
|
@ -1468,21 +1471,16 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
"""
|
||||
return _get_masked_values(headers, ignore_sensitive_values=ignore_sensitive_headers)
|
||||
|
||||
def record_post_call(
|
||||
self, original_response: object, input: object, api_key: object, additional_args: dict[str, object]
|
||||
) -> None:
|
||||
self.model_call_details["input"] = input
|
||||
self.model_call_details["api_key"] = api_key
|
||||
self.model_call_details["original_response"] = original_response
|
||||
self.model_call_details["additional_args"] = additional_args
|
||||
self.model_call_details["log_event_type"] = "post_api_call"
|
||||
|
||||
def post_call(self, original_response, input=None, api_key=None, additional_args={}):
|
||||
# Log the exact result from the LLM API, for streaming - log the type of response received
|
||||
if isinstance(original_response, dict):
|
||||
original_response = json.dumps(original_response, default=str)
|
||||
try:
|
||||
self.record_post_call(original_response, input, api_key, additional_args)
|
||||
self.model_call_details["input"] = input
|
||||
self.model_call_details["api_key"] = api_key
|
||||
self.model_call_details["original_response"] = original_response
|
||||
self.model_call_details["additional_args"] = additional_args
|
||||
self.model_call_details["log_event_type"] = "post_api_call"
|
||||
|
||||
attr: Literal["warning", "debug"]
|
||||
if self.litellm_request_debug:
|
||||
|
|
@ -2177,7 +2175,6 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
logging_result,
|
||||
start_time,
|
||||
end_time,
|
||||
build_logging_payload: bool = True,
|
||||
):
|
||||
"""Resolve hidden params, compute response cost, and emit the standard logging payload."""
|
||||
hidden_params: Final = getattr(logging_result, "_hidden_params", {})
|
||||
|
|
@ -2202,9 +2199,6 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
else:
|
||||
self.model_call_details["response_cost"] = self._response_cost_calculator(result=logging_result)
|
||||
|
||||
if not build_logging_payload:
|
||||
return
|
||||
|
||||
self.model_call_details["standard_logging_object"] = self._build_standard_logging_payload(
|
||||
logging_result, start_time, end_time
|
||||
)
|
||||
|
|
@ -2266,7 +2260,6 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
end_time=None,
|
||||
cache_hit=None,
|
||||
standard_logging_object: StandardLoggingPayload | None = None,
|
||||
build_logging_payload: bool = True,
|
||||
):
|
||||
try:
|
||||
if start_time is None:
|
||||
|
|
@ -2304,7 +2297,6 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
logging_result=logging_result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
build_logging_payload=build_logging_payload,
|
||||
)
|
||||
elif standard_logging_object is not None:
|
||||
self.model_call_details["standard_logging_object"] = standard_logging_object
|
||||
|
|
@ -3328,9 +3320,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
except Exception as e:
|
||||
verbose_logger.debug("Error in _handle_callback_failure: %s", e)
|
||||
|
||||
def _failure_handler_helper_fn(
|
||||
self, exception, traceback_exception, start_time=None, end_time=None, build_logging_payload: bool = True
|
||||
):
|
||||
def _failure_handler_helper_fn(self, exception, traceback_exception, start_time=None, end_time=None):
|
||||
if start_time is None:
|
||||
start_time = self.start_time
|
||||
if end_time is None:
|
||||
|
|
@ -3365,9 +3355,6 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
metadata: Final = self.model_call_details["litellm_params"].get("metadata", {}) or {}
|
||||
metadata.update(exception.headers)
|
||||
|
||||
if not build_logging_payload:
|
||||
return start_time, end_time
|
||||
|
||||
## STANDARDIZED LOGGING PAYLOAD
|
||||
|
||||
self.model_call_details["standard_logging_object"] = get_standard_logging_object_payload(
|
||||
|
|
|
|||
|
|
@ -94,6 +94,7 @@ from litellm.utils import (
|
|||
from ..common_utils import (
|
||||
AnthropicError,
|
||||
AnthropicModelInfo,
|
||||
eager_input_streaming_flag,
|
||||
process_anthropic_headers,
|
||||
strip_advisor_blocks_from_messages,
|
||||
)
|
||||
|
|
@ -732,10 +733,20 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
|
||||
input_anthropic_schema: Final = sanitize_input_schema_for_anthropic(_input_schema)
|
||||
|
||||
_tool: Final = AnthropicMessagesTool(
|
||||
name=tool["function"]["name"],
|
||||
input_schema=input_anthropic_schema,
|
||||
type="custom",
|
||||
_eager_input_streaming: Final = eager_input_streaming_flag(tool)
|
||||
_tool: Final = (
|
||||
AnthropicMessagesTool(
|
||||
name=tool["function"]["name"],
|
||||
input_schema=input_anthropic_schema,
|
||||
type="custom",
|
||||
)
|
||||
if _eager_input_streaming is None
|
||||
else AnthropicMessagesTool(
|
||||
name=tool["function"]["name"],
|
||||
input_schema=input_anthropic_schema,
|
||||
type="custom",
|
||||
eager_input_streaming=_eager_input_streaming,
|
||||
)
|
||||
)
|
||||
|
||||
_description: Final = tool["function"].get("description")
|
||||
|
|
|
|||
|
|
@ -10,7 +10,7 @@ from types import MappingProxyType
|
|||
from typing import Any, Final, Literal
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
|
||||
from pydantic import BaseModel, ConfigDict, StrictBool, TypeAdapter, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm.constants import (
|
||||
|
|
@ -19,6 +19,7 @@ from litellm.constants import (
|
|||
DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET,
|
||||
DEFAULT_REASONING_EFFORT_XHIGH_THINKING_BUDGET,
|
||||
)
|
||||
from litellm.exceptions import UnsupportedParamsError
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
get_file_ids_from_messages,
|
||||
is_encrypted_reasoning_block,
|
||||
|
|
@ -231,6 +232,27 @@ def optionally_handle_anthropic_oauth(headers: dict, api_key: str | None) -> tup
|
|||
return headers, api_key
|
||||
|
||||
|
||||
class _EagerInputStreamingFunction(BaseModel):
|
||||
eager_input_streaming: StrictBool | None = None
|
||||
|
||||
|
||||
class _EagerInputStreamingTool(BaseModel):
|
||||
eager_input_streaming: StrictBool | None = None
|
||||
function: _EagerInputStreamingFunction | None = None
|
||||
|
||||
|
||||
def eager_input_streaming_flag(tool: object) -> bool | None:
|
||||
try:
|
||||
parsed: Final = _EagerInputStreamingTool.model_validate(tool)
|
||||
except ValidationError as error:
|
||||
if isinstance(tool, Mapping):
|
||||
raise UnsupportedParamsError(message="eager_input_streaming must be a boolean") from error
|
||||
return None
|
||||
if parsed.eager_input_streaming is not None:
|
||||
return parsed.eager_input_streaming
|
||||
return parsed.function.eager_input_streaming if parsed.function is not None else None
|
||||
|
||||
|
||||
class AnthropicError(BaseLLMException):
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -373,6 +395,9 @@ class AnthropicModelInfo(BaseLLMModelInfo):
|
|||
|
||||
return False
|
||||
|
||||
def is_eager_input_streaming_used(self, tools: Sequence[object] | None) -> bool:
|
||||
return any(eager_input_streaming_flag(tool) is True for tool in tools or ())
|
||||
|
||||
@staticmethod
|
||||
def _supports_sampling_params(model: str) -> bool:
|
||||
"""Claude 4.7+ (Opus 4.7/4.8, Fable 5) removed sampling params: the API
|
||||
|
|
|
|||
|
|
@ -111,6 +111,7 @@ from litellm.litellm_core_utils.reasoning_effort_utils import (
|
|||
reasoning_effort_from_thinking_budget,
|
||||
)
|
||||
from litellm.llms.anthropic.common_utils import (
|
||||
eager_input_streaming_flag,
|
||||
is_empty_unsigned_thinking_block,
|
||||
normalize_anthropic_tool_use_id,
|
||||
strip_encrypted_reasoning_blocks_from_anthropic_messages,
|
||||
|
|
@ -197,6 +198,15 @@ def target_supports_mid_conversation_system(model: str | None, custom_llm_provid
|
|||
return supports_mid_conversation_system(model=model, custom_llm_provider=custom_llm_provider)
|
||||
|
||||
|
||||
def _chat_tool_param(function_chunk: ChatCompletionToolParamFunctionChunk, tool: object) -> ChatCompletionToolParam:
|
||||
eager_input_streaming: Final = eager_input_streaming_flag(tool)
|
||||
if eager_input_streaming is None:
|
||||
return ChatCompletionToolParam(type="function", function=function_chunk)
|
||||
return ChatCompletionToolParam(
|
||||
type="function", function=function_chunk, eager_input_streaming=eager_input_streaming
|
||||
)
|
||||
|
||||
|
||||
class AnthropicAdapter:
|
||||
def __init__(self) -> None:
|
||||
pass
|
||||
|
|
@ -770,6 +780,7 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
"cache_control",
|
||||
"strict",
|
||||
"type",
|
||||
"eager_input_streaming",
|
||||
]
|
||||
|
||||
for idx, tool in enumerate(tools):
|
||||
|
|
@ -808,7 +819,7 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
for k, v in tool.items():
|
||||
if k not in mapped_tool_params: # pass additional computer kwargs
|
||||
function_chunk.setdefault("parameters", {}).update({k: v})
|
||||
tool_param = ChatCompletionToolParam(type="function", function=function_chunk)
|
||||
tool_param = _chat_tool_param(function_chunk, tool)
|
||||
self._add_cache_control_if_applicable(tool, tool_param, model)
|
||||
new_tools.append(tool_param)
|
||||
|
||||
|
|
@ -1399,6 +1410,8 @@ class LiteLLMAnthropicMessagesAdapter:
|
|||
return "max_tokens"
|
||||
elif openai_finish_reason == "tool_calls":
|
||||
return "tool_use"
|
||||
elif openai_finish_reason in ["content_filter", "refusal"]:
|
||||
return "refusal"
|
||||
return "end_turn"
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -48,6 +48,7 @@ from litellm.llms.bedrock.request_metadata import (
|
|||
merge_bedrock_invoke_headers,
|
||||
resolve_bedrock_request_metadata,
|
||||
)
|
||||
from litellm.types.llms.anthropic import ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER
|
||||
from litellm.types.llms.bedrock import *
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
|
|
@ -1517,12 +1518,6 @@ class AmazonConverseConfig(BaseConfig):
|
|||
"""Process tools and collect anthropic_beta values."""
|
||||
bedrock_tools: list[ToolBlock] = []
|
||||
|
||||
# Collect anthropic_beta values from user headers
|
||||
anthropic_beta_list: Final = []
|
||||
if headers:
|
||||
user_betas: Final = get_anthropic_beta_from_headers(headers)
|
||||
anthropic_beta_list.extend(user_betas)
|
||||
|
||||
# Separate pre-formatted Bedrock tools (e.g. systemTool from web_search_options)
|
||||
# from OpenAI-format tools that need transformation via _bedrock_tools_pt
|
||||
filtered_tools: Final = []
|
||||
|
|
@ -1542,6 +1537,17 @@ class AmazonConverseConfig(BaseConfig):
|
|||
continue
|
||||
filtered_tools.append(tool)
|
||||
|
||||
base_model: Final = BedrockModelInfo.get_base_model(model)
|
||||
client_beta_list: Final = get_anthropic_beta_from_headers(headers or {})
|
||||
eager_beta: Final = (
|
||||
(ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER,)
|
||||
if base_model.startswith("anthropic")
|
||||
and AnthropicModelInfo().is_eager_input_streaming_used(filtered_tools)
|
||||
and ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER not in client_beta_list
|
||||
else ()
|
||||
)
|
||||
anthropic_beta_list: Final = [*client_beta_list, *eager_beta]
|
||||
|
||||
# Only separate tools if computer use tools are actually present
|
||||
if filtered_tools and self.is_computer_use_tool_used(filtered_tools, model):
|
||||
# Separate computer use tools from regular function tools
|
||||
|
|
@ -1619,7 +1625,6 @@ class AmazonConverseConfig(BaseConfig):
|
|||
|
||||
# Opus 4.5 gates ``output_config.effort`` behind a beta header;
|
||||
# Claude 4.6/4.7 accept it without one.
|
||||
base_model: Final = BedrockModelInfo.get_base_model(model)
|
||||
if base_model.startswith("anthropic"):
|
||||
output_config: Final = additional_request_params.get("output_config")
|
||||
if (
|
||||
|
|
|
|||
|
|
@ -24,8 +24,12 @@ from litellm.llms.bedrock.common_utils import (
|
|||
normalize_custom_field_on_tools,
|
||||
normalize_tool_input_schema_types_for_bedrock_invoke,
|
||||
strip_unsupported_bedrock_invoke_output_config_keys,
|
||||
tools_without_eager_input_streaming,
|
||||
)
|
||||
from litellm.types.llms.anthropic import (
|
||||
ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER,
|
||||
ANTHROPIC_TOOL_SEARCH_BETA_HEADER,
|
||||
)
|
||||
from litellm.types.llms.anthropic import ANTHROPIC_TOOL_SEARCH_BETA_HEADER
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
|
@ -237,6 +241,9 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
|
|||
# Hoist `custom.defer_loading` then drop `custom` (Bedrock doesn't support it)
|
||||
normalize_custom_field_on_tools(anthropic_request)
|
||||
normalize_tool_input_schema_types_for_bedrock_invoke(anthropic_request)
|
||||
outbound_tools: Final = tools_without_eager_input_streaming(anthropic_request)
|
||||
if outbound_tools is not None:
|
||||
anthropic_request["tools"] = outbound_tools
|
||||
return anthropic_request
|
||||
|
||||
def _compute_bedrock_invoke_beta_headers(
|
||||
|
|
@ -269,6 +276,9 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
|
|||
if bedrock_supports_tool_search(model):
|
||||
beta_set.add("tool-search-tool-2025-10-19")
|
||||
|
||||
if self.is_eager_input_streaming_used(tools):
|
||||
beta_set.add(ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER)
|
||||
|
||||
auto_beta_list: Final = filter_and_transform_beta_headers(
|
||||
beta_headers=list(beta_set - user_beta_set),
|
||||
provider="bedrock",
|
||||
|
|
|
|||
|
|
@ -9,7 +9,7 @@ import functools
|
|||
import json
|
||||
import os
|
||||
import re
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -18,6 +18,7 @@ if TYPE_CHECKING:
|
|||
from litellm.types.llms.bedrock import BedrockCreateBatchRequest
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm import verbose_logger
|
||||
|
|
@ -330,6 +331,17 @@ def normalize_custom_field_on_tools(request_body: dict) -> None:
|
|||
tool["defer_loading"] = deferred
|
||||
|
||||
|
||||
_TOOL_DICTS_ADAPTER: Final = TypeAdapter(tuple[Mapping[str, object], ...])
|
||||
|
||||
|
||||
def tools_without_eager_input_streaming(request_body: Mapping[str, object]) -> Sequence[object] | None:
|
||||
try:
|
||||
tools: Final = _TOOL_DICTS_ADAPTER.validate_python(request_body.get("tools"))
|
||||
except ValidationError:
|
||||
return None
|
||||
return [{key: value for key, value in tool.items() if key != "eager_input_streaming"} for tool in tools]
|
||||
|
||||
|
||||
def normalize_json_schema_custom_types_to_object(schema: dict) -> None:
|
||||
"""
|
||||
In-place: replace JSON Schema ``type: \"custom\"`` with ``\"object\"`` (iterative walk).
|
||||
|
|
|
|||
|
|
@ -39,6 +39,7 @@ from litellm.llms.bedrock.common_utils import (
|
|||
normalize_custom_field_on_tools,
|
||||
normalize_tool_input_schema_types_for_bedrock_invoke,
|
||||
strip_unsupported_bedrock_invoke_output_config_keys,
|
||||
tools_without_eager_input_streaming,
|
||||
)
|
||||
from litellm.llms.bedrock.request_metadata import (
|
||||
bedrock_request_metadata_headers,
|
||||
|
|
@ -46,6 +47,7 @@ from litellm.llms.bedrock.request_metadata import (
|
|||
)
|
||||
from litellm.types.llms.anthropic import (
|
||||
ANTHROPIC_BETA_HEADER_VALUES,
|
||||
ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER,
|
||||
ANTHROPIC_TOOL_SEARCH_BETA_HEADER,
|
||||
)
|
||||
from litellm.types.llms.bedrock import BedrockInvokeAnthropicMessagesRequest
|
||||
|
|
@ -525,6 +527,9 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
if injected_thinking_for_clear_thinking:
|
||||
beta_set.add("interleaved-thinking-2025-05-14")
|
||||
|
||||
if anthropic_model_info.is_eager_input_streaming_used(tools):
|
||||
beta_set.add(ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER)
|
||||
|
||||
self._filter_context_management_for_bedrock_invoke(
|
||||
anthropic_messages_request=anthropic_messages_request,
|
||||
beta_set=beta_set,
|
||||
|
|
@ -719,6 +724,10 @@ class AmazonAnthropicClaudeMessagesConfig(
|
|||
if filtered_betas:
|
||||
anthropic_messages_request["anthropic_beta"] = filtered_betas
|
||||
|
||||
outbound_tools: Final = tools_without_eager_input_streaming(anthropic_messages_request)
|
||||
if outbound_tools is not None:
|
||||
anthropic_messages_request["tools"] = outbound_tools
|
||||
|
||||
remaining_output_config: Final = anthropic_messages_request.get("output_config")
|
||||
if (
|
||||
litellm.drop_params is True
|
||||
|
|
|
|||
|
|
@ -144,6 +144,7 @@ from litellm.types.llms.openai import (
|
|||
from litellm.types.realtime import RealtimeQueryParams
|
||||
from litellm.types.rerank import RerankResponse
|
||||
from litellm.types.responses.main import DeleteResponseResult
|
||||
from litellm.types.responses.streaming_websocket import ResponsesWebSocketRequestDefaults
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.types.utils import (
|
||||
CallTypes,
|
||||
|
|
@ -6593,6 +6594,7 @@ class BaseLLMHTTPHandler:
|
|||
litellm_metadata: dict[str, object] | None = None,
|
||||
custom_llm_provider: str | None = None,
|
||||
first_message: str | None = None,
|
||||
request_defaults: ResponsesWebSocketRequestDefaults | None = None,
|
||||
**kwargs: Any,
|
||||
):
|
||||
"""
|
||||
|
|
@ -6747,6 +6749,7 @@ class BaseLLMHTTPHandler:
|
|||
output_guardrail_callbacks=_ws_output_guardrail_callbacks,
|
||||
quota_callbacks=_ws_quota_callbacks,
|
||||
authorized_model=model,
|
||||
request_defaults=request_defaults,
|
||||
)
|
||||
await streaming.bidirectional_forward()
|
||||
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ import time
|
|||
from collections.abc import Callable, Mapping, Sequence
|
||||
from copy import deepcopy
|
||||
from functools import partial
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast, get_args
|
||||
|
||||
import httpx
|
||||
|
||||
|
|
@ -57,6 +57,7 @@ from litellm.types.llms.vertex_ai import (
|
|||
ContentType,
|
||||
FunctionCallingConfig,
|
||||
FunctionDeclaration,
|
||||
GeminiFinishReason,
|
||||
GeminiThinkingConfig,
|
||||
GenerateContentResponseBody,
|
||||
HttpxPartType,
|
||||
|
|
@ -1330,25 +1331,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
"IMAGE_PROHIBITED_CONTENT": "The token generation was stopped as the response was flagged for prohibited image content.",
|
||||
}
|
||||
|
||||
_GEMINI_FINISH_REASON_KEYS = frozenset(
|
||||
{
|
||||
"STOP",
|
||||
"MAX_TOKENS",
|
||||
"SAFETY",
|
||||
"RECITATION",
|
||||
"FINISH_REASON_UNSPECIFIED",
|
||||
"MALFORMED_FUNCTION_CALL",
|
||||
"LANGUAGE",
|
||||
"OTHER",
|
||||
"BLOCKLIST",
|
||||
"PROHIBITED_CONTENT",
|
||||
"SPII",
|
||||
"IMAGE_SAFETY",
|
||||
"IMAGE_PROHIBITED_CONTENT",
|
||||
"TOO_MANY_TOOL_CALLS",
|
||||
"MALFORMED_RESPONSE",
|
||||
}
|
||||
)
|
||||
_GEMINI_FINISH_REASON_KEYS: Final[frozenset[str]] = frozenset(get_args(GeminiFinishReason))
|
||||
|
||||
@staticmethod
|
||||
def get_finish_reason_mapping() -> dict[str, OpenAIChatCompletionFinishReason]:
|
||||
|
|
@ -2232,22 +2215,23 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
|
||||
grounding_metadata: Final[list[dict]] = []
|
||||
url_context_metadata: Final[list[dict]] = []
|
||||
image_response: list[ImageURLListItem] | None = None
|
||||
safety_ratings: Final[list] = []
|
||||
citation_metadata: Final[list] = []
|
||||
chat_completion_message: Final[ChatCompletionResponseMessage] = {"role": "assistant"}
|
||||
chat_completion_logprobs: ChoiceLogprobs | None = None
|
||||
tools: list[ChatCompletionToolCallChunk] | None = []
|
||||
functions: ChatCompletionToolCallFunctionChunk | None = None
|
||||
thinking_blocks: list[ChatCompletionThinkingBlock] | None = None
|
||||
reasoning_content: str | None = None
|
||||
thought_signatures: Sequence[str] | None = None
|
||||
server_side_tool_invocations: list[dict[str, object]] | None = None
|
||||
|
||||
for idx, candidate in enumerate(_candidates):
|
||||
if "content" not in candidate:
|
||||
if "content" not in candidate and "finishReason" not in candidate:
|
||||
continue
|
||||
|
||||
image_response: list[ImageURLListItem] | None = None
|
||||
chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"}
|
||||
chat_completion_logprobs: ChoiceLogprobs | None = None
|
||||
tools: list[ChatCompletionToolCallChunk] | None = None
|
||||
functions: ChatCompletionToolCallFunctionChunk | None = None
|
||||
thinking_blocks: list[ChatCompletionThinkingBlock] | None = None
|
||||
reasoning_content: str | None = None
|
||||
thought_signatures: Sequence[str] | None = None
|
||||
server_side_tool_invocations: list[dict[str, object]] | None = None
|
||||
|
||||
# Extract metadata using helper function
|
||||
(
|
||||
candidate_grounding_metadata,
|
||||
|
|
@ -2261,7 +2245,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
safety_ratings.extend(candidate_safety_ratings)
|
||||
citation_metadata.extend(candidate_citation_metadata)
|
||||
|
||||
if "parts" in candidate["content"]:
|
||||
if "content" in candidate and candidate["content"] and "parts" in candidate["content"]:
|
||||
(
|
||||
content,
|
||||
reasoning_content,
|
||||
|
|
@ -2368,14 +2352,18 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
)
|
||||
model_response.choices.append(choice)
|
||||
elif isinstance(model_response, ModelResponse):
|
||||
native_finish_reason = candidate.get("finishReason")
|
||||
choice = litellm.Choices(
|
||||
finish_reason=VertexGeminiConfig._check_finish_reason(
|
||||
chat_completion_message, candidate.get("finishReason")
|
||||
chat_completion_message, native_finish_reason
|
||||
),
|
||||
index=candidate.get("index", idx),
|
||||
message=chat_completion_message,
|
||||
logprobs=chat_completion_logprobs,
|
||||
enhancements=None,
|
||||
provider_specific_fields=(
|
||||
{"native_finish_reason": native_finish_reason} if native_finish_reason is not None else None
|
||||
),
|
||||
)
|
||||
model_response.choices.append(choice)
|
||||
|
||||
|
|
@ -3173,12 +3161,10 @@ class ModelResponseIterator:
|
|||
self.has_seen_tool_calls = True
|
||||
break
|
||||
|
||||
# _process_candidates skips candidates without a "content" part, so a
|
||||
# content-less chunk leaves choices empty and the downstream streaming
|
||||
# handler hits IndexError on choices[0]. This covers the final chunk
|
||||
# (finishReason, no content) and mid-stream metadata-only chunks
|
||||
# (grounding/web-search/thought, no content and no finishReason — seen
|
||||
# with web_search + reasoning) by emitting an empty-delta choice.
|
||||
# _process_candidates skips candidates with neither "content" nor
|
||||
# "finishReason", so a metadata-only chunk (grounding/web-search/thought,
|
||||
# seen with web_search + reasoning) leaves choices empty and the downstream
|
||||
# streaming handler hits IndexError on choices[0]. Emit an empty-delta choice.
|
||||
if not model_response.choices and _candidates:
|
||||
from litellm.types.utils import Delta, StreamingChoices
|
||||
|
||||
|
|
|
|||
|
|
@ -1146,37 +1146,35 @@ def responses_api_bridge_check(
|
|||
return model_info, model
|
||||
|
||||
|
||||
def _should_allow_input_examples(custom_llm_provider: str | None, model: str) -> bool:
|
||||
_ANTHROPIC_ONLY_TOOL_KEYS: Final = frozenset({"input_examples", "eager_input_streaming"})
|
||||
|
||||
|
||||
def _is_claude_tool_target(custom_llm_provider: str | None, model: str) -> bool:
|
||||
if custom_llm_provider == "anthropic":
|
||||
return True
|
||||
if custom_llm_provider == "azure_ai" or custom_llm_provider == "bedrock" or custom_llm_provider == "vertex_ai":
|
||||
return "claude" in model.lower()
|
||||
model_lower: Final = model.lower()
|
||||
if custom_llm_provider == "bedrock":
|
||||
return "claude" in model_lower or ("arn:" in model_lower and ":bedrock:" in model_lower)
|
||||
if custom_llm_provider == "azure_ai" or custom_llm_provider == "vertex_ai":
|
||||
return "claude" in model_lower
|
||||
return False
|
||||
|
||||
|
||||
def _drop_input_examples_from_tool(tool: dict) -> dict:
|
||||
tool_copy: Final = tool.copy()
|
||||
tool_copy.pop("input_examples", None)
|
||||
function = tool_copy.get("function")
|
||||
if isinstance(function, dict):
|
||||
function = function.copy()
|
||||
function.pop("input_examples", None)
|
||||
tool_copy["function"] = function
|
||||
return tool_copy
|
||||
def _without_anthropic_only_tool_keys(tool: dict) -> dict:
|
||||
kept: Final = {key: value for key, value in tool.items() if key not in _ANTHROPIC_ONLY_TOOL_KEYS}
|
||||
function: Final = tool.get("function")
|
||||
if not isinstance(function, dict):
|
||||
return kept
|
||||
return {
|
||||
**kept,
|
||||
"function": {key: value for key, value in function.items() if key not in _ANTHROPIC_ONLY_TOOL_KEYS},
|
||||
}
|
||||
|
||||
|
||||
def _drop_input_examples_from_tools(
|
||||
tools: list[dict] | None,
|
||||
) -> list[dict] | None:
|
||||
def _drop_anthropic_only_tool_keys(tools: list[dict] | None) -> list[dict] | None:
|
||||
if tools is None:
|
||||
return None
|
||||
cleaned_tools: Final[list[dict]] = []
|
||||
for tool in tools:
|
||||
if isinstance(tool, dict):
|
||||
cleaned_tools.append(_drop_input_examples_from_tool(tool))
|
||||
else:
|
||||
cleaned_tools.append(tool)
|
||||
return cleaned_tools
|
||||
return [_without_anthropic_only_tool_keys(tool) if isinstance(tool, dict) else tool for tool in tools]
|
||||
|
||||
|
||||
class _ProxyAuthHeadersProvider(Protocol):
|
||||
|
|
@ -5360,8 +5358,8 @@ def completion(
|
|||
api_base=api_base,
|
||||
)
|
||||
|
||||
if not _should_allow_input_examples(custom_llm_provider=custom_llm_provider, model=model):
|
||||
tools = _drop_input_examples_from_tools(tools=tools)
|
||||
if not _is_claude_tool_target(custom_llm_provider=custom_llm_provider, model=model):
|
||||
tools = _drop_anthropic_only_tool_keys(tools=tools)
|
||||
|
||||
if provider_specific_header is not None:
|
||||
headers.update(
|
||||
|
|
|
|||
|
|
@ -50869,7 +50869,7 @@
|
|||
"vertex_ai/google/gemma-4-26b-a4b-it-maas": {
|
||||
"input_cost_per_token": 1.5e-07,
|
||||
"litellm_provider": "vertex_ai-openai_models",
|
||||
"max_input_tokens": 256000,
|
||||
"max_input_tokens": 262144,
|
||||
"max_output_tokens": 128000,
|
||||
"max_tokens": 128000,
|
||||
"mode": "chat",
|
||||
|
|
|
|||
|
|
@ -33378,6 +33378,10 @@
|
|||
"cache_control": {
|
||||
"$ref": "#/components/schemas/ChatCompletionCachedContent"
|
||||
},
|
||||
"eager_input_streaming": {
|
||||
"title": "Eager Input Streaming",
|
||||
"type": "boolean"
|
||||
},
|
||||
"function": {
|
||||
"$ref": "#/components/schemas/ChatCompletionToolParamFunctionChunk"
|
||||
},
|
||||
|
|
@ -33407,6 +33411,10 @@
|
|||
"title": "Description",
|
||||
"type": "string"
|
||||
},
|
||||
"eager_input_streaming": {
|
||||
"title": "Eager Input Streaming",
|
||||
"type": "boolean"
|
||||
},
|
||||
"name": {
|
||||
"title": "Name",
|
||||
"type": "string"
|
||||
|
|
|
|||
|
|
@ -2593,10 +2593,12 @@ class ProxyBaseLLMRequestProcessing:
|
|||
|
||||
async def refresh_stream_headers() -> Mapping[str, str]:
|
||||
"""`custom_headers` rebuilt for whichever deployment served the stream."""
|
||||
if not getattr(response, "fallback_headers_adopted", False):
|
||||
return custom_headers
|
||||
return self._stream_response_headers(
|
||||
hidden_params=get_hidden_params_dict(response),
|
||||
hidden_params=(
|
||||
get_hidden_params_dict(response)
|
||||
if getattr(response, "fallback_headers_adopted", False)
|
||||
else hidden_params
|
||||
),
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
logging_obj=logging_obj,
|
||||
version=version,
|
||||
|
|
|
|||
|
|
@ -12,13 +12,36 @@ import sys
|
|||
import threading
|
||||
import time
|
||||
import traceback
|
||||
from collections.abc import AsyncGenerator, Awaitable, Callable, Coroutine, Mapping, Sequence
|
||||
from collections.abc import (
|
||||
AsyncGenerator,
|
||||
AsyncIterable,
|
||||
AsyncIterator,
|
||||
Awaitable,
|
||||
Callable,
|
||||
Coroutine,
|
||||
Mapping,
|
||||
Sequence,
|
||||
)
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import date, datetime, timedelta, timezone
|
||||
from email.mime.multipart import MIMEMultipart
|
||||
from email.mime.text import MIMEText
|
||||
from functools import partial
|
||||
from types import MappingProxyType
|
||||
from typing import TYPE_CHECKING, Any, ClassVar, Final, Literal, Optional, Protocol, TypeVar, Union, cast, overload
|
||||
from typing import (
|
||||
TYPE_CHECKING,
|
||||
Any,
|
||||
ClassVar,
|
||||
Final,
|
||||
Generic,
|
||||
Literal,
|
||||
Optional,
|
||||
Protocol,
|
||||
TypeVar,
|
||||
Union,
|
||||
cast,
|
||||
overload,
|
||||
)
|
||||
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
|
|
@ -123,6 +146,7 @@ from litellm.proxy._types import (
|
|||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.common_utils.callback_utils import add_guardrail_to_applied_guardrails_header
|
||||
from litellm.proxy.common_utils.config_sync_pubsub import publish_config_param_change
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.db.create_views import (
|
||||
|
|
@ -438,6 +462,36 @@ def _enrich_http_exception_with_guardrail_context(exc: BaseException, callback:
|
|||
detail.setdefault("guardrail_mode", event_hook)
|
||||
|
||||
|
||||
def _record_raising_guardrail(request_data: Mapping[str, object], callback: object) -> None:
|
||||
guardrail_name: Final[object] = getattr(callback, "guardrail_name", None)
|
||||
if isinstance(request_data, dict) and isinstance(guardrail_name, str):
|
||||
add_guardrail_to_applied_guardrails_header(request_data=request_data, guardrail_name=guardrail_name)
|
||||
|
||||
|
||||
class _UpstreamStreamBoundary(Generic[_T]):
|
||||
__slots__ = ("_upstream", "failure")
|
||||
|
||||
def __init__(self, upstream: AsyncIterable[_T]) -> None:
|
||||
self._upstream: Final = upstream.__aiter__()
|
||||
self.failure: BaseException | None = None
|
||||
|
||||
def __aiter__(self) -> "_UpstreamStreamBoundary[_T]":
|
||||
return self
|
||||
|
||||
async def __anext__(self) -> _T:
|
||||
try:
|
||||
return await self._upstream.__anext__()
|
||||
except StopAsyncIteration:
|
||||
raise
|
||||
except Exception as e:
|
||||
self.failure = e
|
||||
raise
|
||||
|
||||
|
||||
class _StreamIteratorHook(Protocol[_T]):
|
||||
def __call__(self, *, response: AsyncIterator[_T]) -> AsyncGenerator[_T, None]: ...
|
||||
|
||||
|
||||
def _is_client_error_exception(exc: Exception) -> bool:
|
||||
if isinstance(exc, HTTPException):
|
||||
return exc.status_code < 500
|
||||
|
|
@ -1816,13 +1870,19 @@ class ProxyLogging:
|
|||
)
|
||||
if expected_if_unmutated is not None:
|
||||
callback.mark_pre_call_hook_ran(expected_if_unmutated)
|
||||
result: Final = await self._process_guardrail_callback(
|
||||
callback=callback,
|
||||
data=input_data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=call_type,
|
||||
event_type=GuardrailEventHooks.pre_call,
|
||||
)
|
||||
try:
|
||||
result: Final = await self._process_guardrail_callback(
|
||||
callback=callback,
|
||||
data=input_data,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
call_type=call_type,
|
||||
event_type=GuardrailEventHooks.pre_call,
|
||||
)
|
||||
except SensitiveDataRouteException:
|
||||
raise
|
||||
except Exception:
|
||||
_record_raising_guardrail(data, callback)
|
||||
raise
|
||||
if (
|
||||
scans_raw_request
|
||||
and expected_if_unmutated is not None
|
||||
|
|
@ -2045,13 +2105,18 @@ class ProxyLogging:
|
|||
_merge_pipeline_metadata_writes(data, result.modified_data)
|
||||
|
||||
if result.terminal_action == "block":
|
||||
blocking_step: Final = result.step_results[-1] if result.step_results else None
|
||||
callback: Final = (
|
||||
PipelineExecutor.find_guardrail_callback(blocking_step.guardrail_name)
|
||||
if blocking_step is not None
|
||||
else None
|
||||
)
|
||||
if callback is not None:
|
||||
_record_raising_guardrail(data, callback)
|
||||
original_exception: Final = result.original_exception
|
||||
if original_exception is not None and not _exception_changes_request_flow(original_exception):
|
||||
blocking_step: Final = result.step_results[-1] if result.step_results else None
|
||||
if blocking_step is not None:
|
||||
callback: Final = PipelineExecutor.find_guardrail_callback(blocking_step.guardrail_name)
|
||||
if callback is not None:
|
||||
_enrich_http_exception_with_guardrail_context(original_exception, callback)
|
||||
if callback is not None:
|
||||
_enrich_http_exception_with_guardrail_context(original_exception, callback)
|
||||
raise original_exception
|
||||
|
||||
step_results_serializable: Final = [
|
||||
|
|
@ -2317,8 +2382,10 @@ class ProxyLogging:
|
|||
if data is not None:
|
||||
self._process_guardrail_metadata(data)
|
||||
return data
|
||||
except Exception as e:
|
||||
raise e
|
||||
except Exception:
|
||||
if data is not None:
|
||||
self._process_guardrail_metadata(data)
|
||||
raise
|
||||
|
||||
async def _run_parallel_pre_call_guardrails(
|
||||
self,
|
||||
|
|
@ -2376,6 +2443,8 @@ class ProxyLogging:
|
|||
# live kwargs.
|
||||
if callback.scan_raw_request and not isinstance(result, BaseException) and result is not None:
|
||||
callback.mark_pre_call_hook_ran(data)
|
||||
if isinstance(result, BaseException) and not isinstance(result, SensitiveDataRouteException):
|
||||
_record_raising_guardrail(data, callback)
|
||||
raised: Final = tuple(result for result in results if isinstance(result, BaseException))
|
||||
blocking: Final = next((exc for exc in raised if not _exception_changes_request_flow(exc)), None)
|
||||
if blocking is not None:
|
||||
|
|
@ -2454,7 +2523,12 @@ class ProxyLogging:
|
|||
break
|
||||
|
||||
@staticmethod
|
||||
async def _run_guardrail_with_metrics(callback: object, coro: Awaitable[_T], hook_type: str) -> _T:
|
||||
async def _run_guardrail_with_metrics(
|
||||
callback: object,
|
||||
coro: Awaitable[_T],
|
||||
hook_type: str,
|
||||
request_data: Mapping[str, object],
|
||||
) -> _T:
|
||||
"""
|
||||
Await `coro`, recording its latency and status to the
|
||||
`litellm_guardrail_latency_seconds` metric under `hook_type`, and
|
||||
|
|
@ -2474,6 +2548,7 @@ class ProxyLogging:
|
|||
status = "error"
|
||||
error_type = type(e).__name__
|
||||
_enrich_http_exception_with_guardrail_context(e, callback)
|
||||
_record_raising_guardrail(request_data, callback)
|
||||
raise
|
||||
finally:
|
||||
ProxyLogging._emit_guardrail_metrics(
|
||||
|
|
@ -2486,21 +2561,19 @@ class ProxyLogging:
|
|||
|
||||
@staticmethod
|
||||
async def _wrap_streaming_iterator_with_enrichment(
|
||||
callback: object, gen: AsyncGenerator[_T, None]
|
||||
callback: object,
|
||||
response: AsyncIterable[_T],
|
||||
hook: _StreamIteratorHook[_T],
|
||||
request_data: Mapping[str, object],
|
||||
) -> AsyncGenerator[_T, None]:
|
||||
"""
|
||||
Yield from `gen`; if iteration raises an HTTPException with dict detail,
|
||||
enrich the detail with the originating callback's `guardrail_name` and
|
||||
`guardrail_mode` before re-raising. Used to wrap each layer of the
|
||||
async_post_call_streaming_iterator_hook chain so the enrichment is
|
||||
attributed to the callback that produced the chunk pipeline at that
|
||||
point in the chain.
|
||||
"""
|
||||
upstream: Final = _UpstreamStreamBoundary(response)
|
||||
try:
|
||||
async for chunk in gen:
|
||||
async for chunk in hook(response=upstream):
|
||||
yield chunk
|
||||
except Exception as e:
|
||||
_enrich_http_exception_with_guardrail_context(e, callback)
|
||||
if e is not upstream.failure:
|
||||
_enrich_http_exception_with_guardrail_context(e, callback)
|
||||
_record_raising_guardrail(request_data, callback)
|
||||
raise
|
||||
|
||||
# Cache for callback-capability detection. Keyed on a signature of
|
||||
|
|
@ -2735,6 +2808,7 @@ class ProxyLogging:
|
|||
call_type=call_type,
|
||||
),
|
||||
"during_call",
|
||||
request_data=data,
|
||||
)
|
||||
return
|
||||
await self._run_guardrail_with_metrics(
|
||||
|
|
@ -2745,6 +2819,7 @@ class ProxyLogging:
|
|||
call_type=call_type,
|
||||
),
|
||||
"during_call",
|
||||
request_data=data,
|
||||
)
|
||||
|
||||
async def failed_tracking_alert(
|
||||
|
|
@ -3263,6 +3338,7 @@ class ProxyLogging:
|
|||
response=response,
|
||||
),
|
||||
"post_call",
|
||||
request_data=data,
|
||||
)
|
||||
else:
|
||||
guardrail_response = await self._run_guardrail_with_metrics(
|
||||
|
|
@ -3273,6 +3349,7 @@ class ProxyLogging:
|
|||
response=response,
|
||||
),
|
||||
"post_call",
|
||||
request_data=data,
|
||||
)
|
||||
|
||||
if guardrail_response is not None:
|
||||
|
|
@ -3336,6 +3413,7 @@ class ProxyLogging:
|
|||
response=response,
|
||||
),
|
||||
"post_call",
|
||||
request_data=data,
|
||||
)
|
||||
else:
|
||||
await self._run_guardrail_with_metrics(
|
||||
|
|
@ -3346,6 +3424,7 @@ class ProxyLogging:
|
|||
response=response,
|
||||
),
|
||||
"post_call",
|
||||
request_data=data,
|
||||
)
|
||||
|
||||
results: Final = await asyncio.gather(
|
||||
|
|
@ -3409,6 +3488,7 @@ class ProxyLogging:
|
|||
request_data=request_data,
|
||||
),
|
||||
"post_mcp_call",
|
||||
request_data=request_data,
|
||||
)
|
||||
return response
|
||||
|
||||
|
|
@ -3650,27 +3730,27 @@ class ProxyLogging:
|
|||
)
|
||||
else kind
|
||||
)
|
||||
if effective_kind == "override":
|
||||
current_response = self._wrap_streaming_iterator_with_enrichment(
|
||||
resolved_callback,
|
||||
resolved_callback.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=current_response,
|
||||
request_data=request_data,
|
||||
),
|
||||
hook: _StreamIteratorHook[object] = (
|
||||
partial(
|
||||
resolved_callback.async_post_call_streaming_iterator_hook,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=request_data,
|
||||
)
|
||||
else:
|
||||
# kind == "apply_guardrail": route through unified_guardrail
|
||||
current_response = self._wrap_streaming_iterator_with_enrichment(
|
||||
resolved_callback,
|
||||
unified_guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=request_data,
|
||||
response=current_response,
|
||||
guardrail_to_apply=resolved_callback,
|
||||
buffer_until_moderated_default=(kind == "override"),
|
||||
),
|
||||
if effective_kind == "override"
|
||||
else partial(
|
||||
unified_guardrail.async_post_call_streaming_iterator_hook,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=request_data,
|
||||
guardrail_to_apply=resolved_callback,
|
||||
buffer_until_moderated_default=(kind == "override"),
|
||||
)
|
||||
)
|
||||
current_response = self._wrap_streaming_iterator_with_enrichment(
|
||||
resolved_callback,
|
||||
current_response,
|
||||
hook,
|
||||
request_data=request_data,
|
||||
)
|
||||
|
||||
pipeline_translation: Final = (
|
||||
resolve_endpoint_translation(user_api_key_dict, None) if post_call_pipelines else None
|
||||
|
|
|
|||
|
|
@ -61,6 +61,7 @@ from litellm.types.llms.openai import (
|
|||
ChatCompletionToolParamFunctionChunk,
|
||||
ChatCompletionUserMessage,
|
||||
GenericChatCompletionMessage,
|
||||
IncompleteDetails,
|
||||
InputTokensDetails,
|
||||
OpenAIChatCompletionTextObject,
|
||||
OpenAIMcpServerTool,
|
||||
|
|
@ -111,6 +112,9 @@ ResponseTools: TypeAlias = Sequence[Mapping[str, object]] | None
|
|||
ChatToolParam: TypeAlias = ChatCompletionToolParam | OpenAIMcpServerTool
|
||||
NAMESPACE_DESCRIPTION_SEPARATOR: Final = "\n\n"
|
||||
NAMESPACE_MEMBER_TYPES_WITH_CHAT_TOOLS: Final = frozenset({"function", "custom"})
|
||||
_INCOMPLETE_REASON_BY_FINISH_REASON: Final[Mapping[str, Literal["max_output_tokens", "content_filter"]]] = (
|
||||
MappingProxyType({"length": "max_output_tokens", "content_filter": "content_filter", "refusal": "content_filter"})
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
@ -2020,6 +2024,8 @@ class LiteLLMCompletionResponsesConfig:
|
|||
chat_completion_tool["allowed_callers"] = tool.get("allowed_callers")
|
||||
if tool.get("input_examples"):
|
||||
chat_completion_tool["input_examples"] = tool.get("input_examples")
|
||||
if tool.get("eager_input_streaming") is not None:
|
||||
chat_completion_tool["eager_input_streaming"] = tool.get("eager_input_streaming")
|
||||
return ResponsesToolChatForm(
|
||||
chat_tools=(cast(ChatCompletionToolParam, chat_completion_tool),), web_search_options=None
|
||||
)
|
||||
|
|
@ -2096,6 +2102,8 @@ class LiteLLMCompletionResponsesConfig:
|
|||
responses_tool["allowed_callers"] = tool.get("allowed_callers")
|
||||
if tool.get("input_examples") is not None:
|
||||
responses_tool["input_examples"] = tool.get("input_examples")
|
||||
if tool.get("eager_input_streaming") is not None:
|
||||
responses_tool["eager_input_streaming"] = tool.get("eager_input_streaming")
|
||||
result.append(responses_tool)
|
||||
else:
|
||||
# mcp or other: pass through unchanged
|
||||
|
|
@ -2295,6 +2303,18 @@ class LiteLLMCompletionResponsesConfig:
|
|||
# Default to completed for unknown finish reasons
|
||||
return "completed"
|
||||
|
||||
@staticmethod
|
||||
def _incomplete_details_for_finish_reason(
|
||||
finish_reason: str | None,
|
||||
existing: IncompleteDetails | None,
|
||||
) -> IncompleteDetails | None:
|
||||
if existing is not None:
|
||||
return existing
|
||||
if finish_reason is None:
|
||||
return None
|
||||
reason: Final = _INCOMPLETE_REASON_BY_FINISH_REASON.get(finish_reason)
|
||||
return IncompleteDetails(reason=reason) if reason is not None else None
|
||||
|
||||
@staticmethod
|
||||
def _tool_call_id_from_responses_item(item_id: str | None, call_id: str | None) -> str:
|
||||
"""Bedrock Mantle returns a non-unique, index-based ``call_id`` (``call_0``,
|
||||
|
|
@ -2411,13 +2431,18 @@ class LiteLLMCompletionResponsesConfig:
|
|||
if choices and len(choices) > 0:
|
||||
finish_reason = choices[0].finish_reason
|
||||
|
||||
incomplete_details: Final = LiteLLMCompletionResponsesConfig._incomplete_details_for_finish_reason(
|
||||
finish_reason=finish_reason,
|
||||
existing=getattr(chat_completion_response, "incomplete_details", None),
|
||||
)
|
||||
|
||||
responses_api_response: Final[ResponsesAPIResponse] = ResponsesAPIResponse(
|
||||
id=chat_completion_response.id,
|
||||
created_at=chat_completion_response.created,
|
||||
model=chat_completion_response.model,
|
||||
object="response",
|
||||
error=getattr(chat_completion_response, "error", None),
|
||||
incomplete_details=getattr(chat_completion_response, "incomplete_details", None),
|
||||
incomplete_details=incomplete_details,
|
||||
instructions=getattr(chat_completion_response, "instructions", None),
|
||||
metadata=getattr(chat_completion_response, "metadata", {}),
|
||||
output=LiteLLMCompletionResponsesConfig._transform_chat_completion_choices_to_responses_output(
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue