diff --git a/.circleci/config.yml b/.circleci/config.yml index dbeb412506f..abcdbf45187 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -133,6 +133,26 @@ commands: done echo "record/replay proxy did not become ready" >&2 exit 1 + start_fake_openai_endpoint: + description: "Start the canned OpenAI mock (tests/_fake_openai_endpoint_server.py) on host port 8190 and wait until healthy. Models whose api_base points here (via FAKE_OPENAI_API_BASE) get well-formed chat/text/embedding responses with realistic usage, so the E2E run neither pays for nor depends on the live provider. A request whose model is '429' returns HTTP 429 for rate-limit/cooldown tests. Run after uv deps are synced." + steps: + - run: + name: Start fake OpenAI endpoint + background: true + command: | + uv run --no-sync python tests/_fake_openai_endpoint_server.py --host 0.0.0.0 --port 8190 + - run: + name: Wait for fake OpenAI endpoint + command: | + for i in $(seq 1 30); do + if curl -sf http://localhost:8190/health >/dev/null 2>&1; then + echo "fake OpenAI endpoint is up" + exit 0 + fi + sleep 1 + done + echo "fake OpenAI endpoint did not become ready" >&2 + exit 1 setup_litellm_enterprise_pip: steps: - run: @@ -168,6 +188,8 @@ jobs: name: win/default shell: powershell.exe working_directory: ~/project + environment: + UV_PYTHON: "3.11" steps: - checkout - run: @@ -200,7 +222,7 @@ jobs: if (-not (Select-String -Path $PROFILE -SimpleMatch $uvBin -Quiet)) { Add-Content -Path $PROFILE -Value "`$env:Path = `"$uvBin;`$env:Path`"" } - uv sync --frozen --group dev --python (Get-Command python).Source + uv sync --frozen --group dev --python 3.11 - run: name: Run Windows-specific test command: | @@ -594,6 +616,8 @@ jobs: working_directory: ~/project resource_class: large parallelism: 4 + environment: + FAKE_OPENAI_API_BASE: http://127.0.0.1:8190 steps: - checkout - setup_google_dns @@ -609,6 +633,7 @@ jobs: paths: - ~/.cache/uv key: v1-uv-cache-{{ checksum "uv.lock" }} + - start_fake_openai_endpoint # Run pytest and generate JUnit XML report - setup_litellm_enterprise_pip - run: @@ -1549,6 +1574,7 @@ jobs: name: Install Dependencies command: | uv sync --frozen --all-groups --all-extras --python 3.12 + - start_fake_openai_endpoint - start_postgres: db_name: litellm_test - attach_workspace: @@ -1586,6 +1612,7 @@ jobs: -e DATABASE_URL="postgresql://postgres:postgres@host.docker.internal:5432/litellm_test" \ -e DEFAULT_NUM_WORKERS_LITELLM_PROXY=1 \ -e DISABLE_SCHEMA_UPDATE="True" \ + -e FAKE_OPENAI_API_BASE=http://host.docker.internal:8190 \ --name my-app \ --add-host=host.docker.internal:host-gateway \ -v $(pwd)/litellm/proxy/example_config_yaml/bad_schema.prisma:/app/schema.prisma \ @@ -1648,6 +1675,7 @@ jobs: zstd -d litellm-docker-database.tar.zst --stdout | docker load docker tag litellm-docker-database:ci my-app:latest - start_openai_record_replay_proxy + - start_fake_openai_endpoint - run: name: Run Docker container command: | @@ -1655,6 +1683,7 @@ jobs: -p 4000:4000 \ -e DATABASE_URL=postgresql://postgres:postgres@host.docker.internal:5432/circle_test \ -e USE_PRISMA_MIGRATE=True \ + -e FAKE_OPENAI_API_BASE=http://host.docker.internal:8190 \ -e AZURE_API_KEY=$AZURE_API_KEY \ -e REDIS_HOST=$REDIS_HOST \ -e REDIS_PASSWORD=$REDIS_PASSWORD \ @@ -1817,6 +1846,7 @@ jobs: zstd -d litellm-docker-database.tar.zst --stdout | docker load docker images | grep litellm-docker-database - start_openai_record_replay_proxy + - start_fake_openai_endpoint - run: name: Run Docker container # intentionally give bad redis credentials here @@ -1830,6 +1860,7 @@ jobs: -e REDIS_PORT=$REDIS_PORT \ -e LITELLM_MASTER_KEY="sk-1234" \ -e OPENAI_API_KEY=$OPENAI_API_KEY \ + -e FAKE_OPENAI_API_BASE=http://host.docker.internal:8190 \ -e LITELLM_LICENSE=$LITELLM_LICENSE \ -e OTEL_EXPORTER="in_memory" \ -e AWS_ACCESS_KEY_ID=$AWS_ACCESS_KEY_ID \ @@ -1889,6 +1920,7 @@ jobs: -e REDIS_PORT=$REDIS_PORT \ -e LITELLM_MASTER_KEY="sk-1234" \ -e OPENAI_API_KEY=$OPENAI_API_KEY \ + -e FAKE_OPENAI_API_BASE=http://host.docker.internal:8190 \ -e LITELLM_LICENSE="bad-license" \ --add-host host.docker.internal:host-gateway \ --name my-app-3 \ @@ -1938,6 +1970,7 @@ jobs: uv sync --frozen --all-groups --all-extras --python 3.12 - start_postgres - start_redis + - start_fake_openai_endpoint - attach_workspace: at: ~/project - run: @@ -1961,6 +1994,7 @@ jobs: -e REDIS_PORT=6379 \ -e LITELLM_MASTER_KEY="sk-1234" \ -e OPENAI_API_KEY=$OPENAI_API_KEY \ + -e FAKE_OPENAI_API_BASE=http://host.docker.internal:8190 \ -e LITELLM_LICENSE=$LITELLM_LICENSE \ -e AWS_ACCESS_KEY_ID=$AWS_ACCESS_KEY_ID \ -e AWS_SECRET_ACCESS_KEY=$AWS_SECRET_ACCESS_KEY \ @@ -2020,6 +2054,7 @@ jobs: command: | uv sync --frozen --all-groups --all-extras --python 3.12 - start_postgres + - start_fake_openai_endpoint - attach_workspace: at: ~/project - run: @@ -2039,6 +2074,7 @@ jobs: -e REDIS_PASSWORD=$REDIS_PASSWORD \ -e REDIS_PORT=$REDIS_PORT \ -e LITELLM_MASTER_KEY="sk-1234" \ + -e FAKE_OPENAI_API_BASE=http://host.docker.internal:8190 \ -e LITELLM_LICENSE=$LITELLM_LICENSE \ -e USE_DDTRACE=True \ -e DD_API_KEY=$DD_API_KEY \ @@ -2060,6 +2096,7 @@ jobs: -e REDIS_PASSWORD=$REDIS_PASSWORD \ -e REDIS_PORT=$REDIS_PORT \ -e LITELLM_MASTER_KEY="sk-1234" \ + -e FAKE_OPENAI_API_BASE=http://host.docker.internal:8190 \ -e LITELLM_LICENSE=$LITELLM_LICENSE \ -e USE_DDTRACE=True \ -e DD_API_KEY=$DD_API_KEY \ @@ -2112,6 +2149,7 @@ jobs: command: | uv sync --frozen --all-groups --all-extras --python 3.12 - start_postgres + - start_fake_openai_endpoint - attach_workspace: at: ~/project - run: @@ -2129,6 +2167,7 @@ jobs: -e DATABASE_URL=postgresql://postgres:postgres@host.docker.internal:5432/circle_test \ -e STORE_MODEL_IN_DB="True" \ -e LITELLM_MASTER_KEY="sk-1234" \ + -e FAKE_OPENAI_API_BASE=http://host.docker.internal:8190 \ -e LITELLM_LICENSE=$LITELLM_LICENSE \ --add-host host.docker.internal:host-gateway \ --name my-app \ @@ -2187,6 +2226,7 @@ jobs: command: | docker build -t my-app:latest -f docker/build_from_pip/Dockerfile.build_from_pip . - start_postgres + - start_fake_openai_endpoint - run: name: Run Docker container # intentionally give bad redis credentials here @@ -2200,6 +2240,7 @@ jobs: -e REDIS_PORT=$REDIS_PORT \ -e LITELLM_MASTER_KEY="sk-1234" \ -e OPENAI_API_KEY=$OPENAI_API_KEY \ + -e FAKE_OPENAI_API_BASE=http://host.docker.internal:8190 \ -e LITELLM_LICENSE=$LITELLM_LICENSE \ -e OTEL_EXPORTER="in_memory" \ -e AWS_ACCESS_KEY_ID=$AWS_ACCESS_KEY_ID \ diff --git a/.github/scripts/_agent_shin_actions.py b/.github/scripts/_agent_shin_actions.py new file mode 100644 index 00000000000..b3d1ff055b3 --- /dev/null +++ b/.github/scripts/_agent_shin_actions.py @@ -0,0 +1,50 @@ +"""Dry-run wrapper(s) around Agent Shin GitHub mutations. + +The rollout scripts currently need only one mutation wrapped, so this module +exposes a single ``maybe_post_comment`` helper. It takes a ``dry_run: bool`` +keyword argument and the body is intentionally trivial: + + if dry_run: + print(...) # log what we would do, return + return + real_mutation(...) # otherwise, actually do it + +That shape means a dry-run preview differs from the real run in exactly one +line per side effect: the call site. So when you `python3 script.py` locally +without ``--close``, you can be confident the actions printed are the ones the +GitHub Action would have performed (modulo ordering on retry/error paths, +which are deliberately simple). Any further mutation a rollout script needs +should get the same ``maybe_*`` treatment instead of calling the raw +``triage_with_llm`` mutation directly. + +Importing from this module pulls in the real mutation from ``triage_with_llm`` +— call sites in the rollout scripts should NEVER import ``post_comment`` +directly; that would skip the dry-run gate and is the bug class this module +exists to prevent. +""" + +from __future__ import annotations + +import sys +import textwrap + +# Import the module itself rather than the bare names so monkeypatching +# `triage_with_llm.post_comment` (or any of the other mutations) in tests is +# reflected here — `from triage_with_llm import post_comment` would bind the +# original function to a local name and bypass the patch, defeating the whole +# point of these wrappers. +import triage_with_llm + + +def _log(line: str) -> None: + """Print a single dry-run line to stdout (one log statement per side effect).""" + print(line, file=sys.stdout, flush=True) + + +def maybe_post_comment(repo: str, number: int, body: str, *, dry_run: bool) -> None: + """Post a comment on ``repo#number`` — or, in dry-run, log what we would post.""" + if dry_run: + _log(f"[DRY RUN] comment {repo}#{number}:") + _log(textwrap.indent(body, " ")) + return + triage_with_llm.post_comment(repo, number, body) diff --git a/.github/scripts/agent_shin_shared.py b/.github/scripts/agent_shin_shared.py new file mode 100644 index 00000000000..8f3dc3c2322 --- /dev/null +++ b/.github/scripts/agent_shin_shared.py @@ -0,0 +1,211 @@ +"""Constants and helpers shared by Agent Shin's triage scripts. + +Both `triage_with_llm.py` (the LLM-judge entrypoint) and +`close_low_quality_prs.py` (the daily Greptile-score sweep) need to +agree on the same notions of: + + * What counts as a Greptile-authored review comment + (``GREPTILE_BOT_LOGINS``) and how to extract a confidence score from + its body (``SCORE_PATTERN`` / :func:`extract_greptile_score`). + * How long the 2-hour grace window is (``GRACE_PERIOD_SECONDS``) and + the HTML marker stamped into a grace-warning comment so the *other* + script can see "Agent Shin already warned" and behave accordingly + (``GRACE_COMMENT_MARKER``). + * Who Agent Shin is on GitHub (``AGENT_SHIN_DEFAULT_BOT_LOGIN``). + * How GitHub-style ISO-8601 timestamps round-trip into timezone-aware + :class:`datetime.datetime` (:func:`parse_iso8601`). + +Keeping these in one module means a future change (new Greptile output +format, a longer grace window, a new allowlisted account) is a single edit +instead of two — the original split version had to call out in comments +that the two copies "must stay in sync" precisely because nothing +enforced it. +""" + +from __future__ import annotations + +import datetime as dt +import json +import os +import re +import subprocess +from typing import Iterable + +GREPTILE_BOT_LOGINS = frozenset({"greptile-apps", "greptile-apps[bot]"}) + +SCORE_PATTERN = re.compile( + r"confidence\s*score\s*[:\-]?\s*(\d+)\s*/\s*5", + re.IGNORECASE, +) + +GRACE_COMMENT_MARKER = "" + +# 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 = "" + +# 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 ``, 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() diff --git a/.github/scripts/close_low_quality_prs.py b/.github/scripts/close_low_quality_prs.py new file mode 100644 index 00000000000..7b9bbb579e3 --- /dev/null +++ b/.github/scripts/close_low_quality_prs.py @@ -0,0 +1,573 @@ +#!/usr/bin/env python3 +""" +Auto-close low-quality pull requests. + +Closes open PRs (including drafts, regardless of age) that satisfy ALL of: + 1. Have a Greptile (`greptile-apps`) review comment whose latest + "Confidence Score: X/5" is below the configured threshold (default: 4). + 2. Are authored by an external OSS contributor (internal BerriAI + contributors are exempt). + 3. Do not carry an opt-out label (default: "do not close"). + +`--min-age-days` is retained as an opt-in safety net for one-off backfill +runs (default: 0). The team's intent is that the count of open PRs equals +the count of PRs internal collaborators need to action on, so neither age +nor draft status acts as a free pass. + +For each match, the script posts an explanatory comment and closes the PR. +Because OSS contributors *cannot* reopen a PR closed by the bot/maintainer +(GitHub limitation), the close-comment instructs them to push their fixes +and **open a fresh PR**, or to comment `@agent-shin reconsider` on the +closed PR to have the LLM judge re-evaluate (and reopen on pass). + +Requires the `gh` CLI to be authenticated. + +Usage examples: + # Dry run (default) - prints what would be closed + python3 close_low_quality_prs.py + + # Actually close matching PRs + python3 close_low_quality_prs.py --close + + # Restrict to PRs at least N days old (one-off backfill safety net) + python3 close_low_quality_prs.py --min-age-days 7 --min-score 4 --close +""" + +from __future__ import annotations + +import argparse +import datetime as dt +import json +import os +import subprocess +import sys +from typing import Iterable + +# Add this script's directory to `sys.path` so the sibling +# `agent_shin_shared` module is importable when the script is invoked +# directly (e.g. `python3 .github/scripts/close_low_quality_prs.py ...`). +sys.path.insert(0, os.path.dirname(os.path.abspath(__file__))) + +from agent_shin_shared import ( # noqa: E402 -- sys.path adjusted above + AGENT_SHIN_CLOSE_MARKER, + ALLOWLIST_LOGINS, + GRACE_COMMENT_MARKER, + GRACE_PERIOD_SECONDS, + GREPTILE_BOT_LOGINS, + SCORE_PATTERN, + extract_greptile_score, + gh, + list_open_items, + parse_iso8601, + seconds_since_latest_marker_comment, +) + +# `GREPTILE_BOT_LOGINS` and `SCORE_PATTERN` (Greptile's GitHub App login +# variants and the "Confidence Score: X/5" regex) are imported from +# `agent_shin_shared` so the LLM judge in `triage_with_llm.py` and this +# daily Greptile sweep read the score through the same set of logins +# and the same regex. + +# `author_association` values for internal BerriAI contributors who should be +# exempt from auto-triage. +INTERNAL_AUTHOR_ASSOCIATIONS = frozenset({"OWNER", "MEMBER", "COLLABORATOR"}) + +# Default labels that exempt a PR from auto-close. Defined at module scope (not +# as a mutable argparse default) so that `--optout-label foo` REPLACES the +# defaults instead of appending to them — the argparse `action="append"` + +# `default=[...]` combination silently mutates the shared default list. +DEFAULT_OPTOUT_LABELS = ("do not close", "keep open", "wip") + +# `GRACE_COMMENT_MARKER` (HTML marker appended to grace-period warning +# comments — used by either script to recognize that a warning was +# already posted) and `GRACE_PERIOD_SECONDS` (length of the grace +# period between the warning and the actual auto-close, 2 hours) are +# imported from `agent_shin_shared` so the Agent Shin LLM judge and +# this daily Greptile sweep agree on the same marker and duration. + + +def fetch_open_prs(repo: str | None) -> list[dict]: + """Fetch all open PRs (number, createdAt, isDraft, labels, author). + + Includes drafts: `gh pr list --state open` returns both ready-for-review + and draft PRs by default. This is the desired behavior — drafts are not + a free pass; the internal-collaborator open-PR queue should reflect every + PR that needs human attention regardless of draft status. + """ + fields = "number,title,createdAt,isDraft,labels,author,url" + return list_open_items("pr", repo=repo, fields=fields) + + +def fetch_pr_author_association(pr_number: int, repo: str | None) -> str: + """Return the GitHub `author_association` for a PR, uppercase. + + Values: OWNER, MEMBER, COLLABORATOR, CONTRIBUTOR, FIRST_TIME_CONTRIBUTOR, + FIRST_TIMER, MANNEQUIN, NONE. Returns "" on lookup failure. + """ + endpoint = ( + f"repos/{repo}/pulls/{pr_number}" + if repo + else f"repos/{{owner}}/{{repo}}/pulls/{pr_number}" + ) + try: + data = json.loads(gh("api", endpoint)) + except subprocess.CalledProcessError: + return "" + return (data.get("author_association") or "").upper() + + +def is_external_pr_author(pr: dict, repo: str | None) -> bool: + """Return True if the PR author is an external OSS contributor. + + Internal = `OWNER` / `MEMBER` / `COLLABORATOR` association, or a bot login. + """ + login = ((pr.get("author") or {}).get("login") or "").lower() + if login.endswith("[bot]") or login in {"dependabot", "github-actions"}: + return False + association = fetch_pr_author_association(pr["number"], repo) + # Fail-safe: if the API lookup failed (empty string), treat the author as + # internal so we don't auto-close their PR. Auto-close is destructive, so + # an unknown association should never make a PR eligible for closing. + if not association or association in INTERNAL_AUTHOR_ASSOCIATIONS: + return False + return True + + +def fetch_pr_comments(pr_number: int, repo: str | None) -> list[dict]: + """Fetch issue-level comments on a PR (where Greptile posts its summary).""" + endpoint = ( + f"repos/{repo}/issues/{pr_number}/comments?per_page=100" + if repo + else f"repos/{{owner}}/{{repo}}/issues/{pr_number}/comments?per_page=100" + ) + raw = gh("api", "--paginate", endpoint) + comments: list[dict] = [] + for line in raw.strip().splitlines(): + line = line.strip() + if not line: + continue + try: + parsed = json.loads(line) + except json.JSONDecodeError: + # A malformed line should not blow up the whole sweep. Skip and + # carry on so the remaining PRs in this run still get evaluated. + continue + if isinstance(parsed, list): + comments.extend(parsed) + else: + comments.append(parsed) + return comments + + +def has_optout_label(pr: dict, optout_labels: set[str]) -> bool: + labels = {label.get("name", "").lower() for label in pr.get("labels", [])} + return bool(labels & {lbl.lower() for lbl in optout_labels}) + + +def seconds_since_last_grace_warning( + comments: Iterable[dict], + *, + bot_login: str | None = None, + now: dt.datetime | None = None, +) -> float | None: + """Return seconds since the bot's most recent grace-period warning, or + None if no such warning has ever been posted on this PR. + + Thin wrapper over + `agent_shin_shared.seconds_since_latest_marker_comment` — the + centralized helper handles the bot-author filter, marker match, + timestamp parsing, and `now` injection. Keeping this wrapper + preserves the closer's "already-fetched comments + injectable now" + interface so callers (and tests) don't need to change. + """ + return seconds_since_latest_marker_comment( + comments, + marker=GRACE_COMMENT_MARKER, + bot_login=bot_login, + now=now, + ) + + +def format_grace_warning_comment(score: int, threshold: int) -> str: + """Comment posted on the FIRST low-Greptile-score detection — gives + the contributor a 2-hour grace window before the auto-close fires on + the next daily cron run. + + Mirrors `format_grace_warning_pr_comment` in + `triage_with_llm.py` in spirit (2-hour grace + escape hatches), but + framed around Greptile's confidence score instead of the LLM judge's + rubric since the close trigger here is the Greptile signal. + """ + return ( + "🚅 Hi, thanks for the PR! I'm **Agent Shin**, the automated triage bot for this " + "repository.\n" + "\n" + "Heads up: Greptile's most recent review scored this PR " + f"**{score}/5**, below our merge bar of **{threshold}/5**.\n" + "\n" + "If the score isn't lifted in the next **2 hours**, I'll auto-close this PR. That's " + "**not** us saying the change isn't worthwhile. We want the open-PR list to mirror " + "what a maintainer can act on *right now*, so contributors like you don't get lost in " + "a backlog. Take your time; everything below still works after the close.\n" + "\n" + "**During the grace period:** push fixes that address Greptile's feedback, then comment " + "`@greptileai` to request a fresh review. If " + f"the new score is **{threshold}/5 or higher**, the PR stays open and no further " + "action is needed on your side.\n" + "\n" + "**If the PR does get auto-closed in 2 hours, you still have an easy recovery path:**\n" + "\n" + "- Comment `@greptileai` to request a fresh review. **This still works even after " + f"the PR is closed**, and a score of {threshold}/5 or higher is one of the signals " + "that lifts the PR back into the review queue. A low Greptile score isn't a blocker.\n" + "- Comment `@agent-shin reconsider` after pushing fixes; I'll re-run the rubric and " + "reopen the PR if both gates (description rubric + Greptile score) now pass.\n" + "\n" + f"{GRACE_COMMENT_MARKER}" + ) + + +def post_grace_warning( + pr: dict, + score: int, + threshold: int, + repo: str | None, + dry_run: bool, +) -> None: + """Post the 2-hour grace-period warning comment on `pr`. + + The warning carries `GRACE_COMMENT_MARKER` so subsequent runs can + detect that the contributor has already been told about the + pending close. Does NOT close the PR — the close happens on the + next eligible run after `GRACE_PERIOD_SECONDS` elapses (handled + by `close_pr`). + """ + pr_number = pr["number"] + repo_args = ["--repo", repo] if repo else [] + + if dry_run: + print( + f" [DRY RUN] Would post grace warning to PR #{pr_number} " + f"(greptile={score}/5): {pr['title']}" + ) + return + + comment_body = format_grace_warning_comment(score, threshold) + gh("pr", "comment", str(pr_number), "--body", comment_body, *repo_args) + print(f" Posted grace warning on PR #{pr_number} (greptile={score}/5)") + + +def format_close_comment(score: int, threshold: int) -> str: + """Comment posted when a low-Greptile-score PR is auto-closed. + + Carries `AGENT_SHIN_CLOSE_MARKER` so the `@agent-shin reconsider` path + (guarded by `was_closed_by_agent_shin`) recognizes this as an Agent Shin + close and is allowed to reopen the PR once it passes again; without the + marker that recovery path the comment advertises silently rejects the + contributor. + """ + score_sentence = ( + f"Greptile's most recent review scored this PR **{score}/5**, below " + f"our merge bar of **{threshold}/5**, and the 2-hour grace period since " + "the warning has elapsed.\n\n" + ) + return ( + f"Closing as part of automated PR triage.\n\n" + f"{score_sentence}" + "We close low-confidence PRs aggressively to keep the review queue " + "manageable for maintainers and contributors alike. **This is not a " + "rejection of the idea.** To bring this back:\n\n" + "1. Push the fixes that address Greptile's feedback (continue using " + "your existing branch is fine).\n" + "2. **Open a new PR** with the updated branch. Greptile will review " + "it again, and if it scores " + f"**{threshold}/5 or higher** a maintainer will take another look.\n\n" + "_Why open a new PR instead of reopening this one?_ GitHub does not " + "let external contributors reopen a PR that was closed by a bot or " + "maintainer, so a fresh PR is the most reliable path forward. If you " + "would prefer this exact PR re-evaluated, comment " + "`@agent-shin reconsider` once you've pushed the fixes; Agent Shin " + "will re-run triage and reopen this PR if it now meets the bar. " + "You can also comment `@greptileai` to request a fresh Greptile " + "review; that works **even after the PR is closed**.\n\n" + "Thanks for contributing to LiteLLM. We know auto-closures can sting; " + "the goal is to keep the project healthy, not to dismiss your work." + f"\n\n{AGENT_SHIN_CLOSE_MARKER}" + ) + + +def close_pr( + pr: dict, + score: int, + threshold: int, + age_days: int, + repo: str | None, + dry_run: bool, + label: str | None, +) -> None: + """Post the explanatory comment and close the PR.""" + pr_number = pr["number"] + repo_args = ["--repo", repo] if repo else [] + + if dry_run: + print( + f" [DRY RUN] Would close PR #{pr_number} " + f"(age={age_days}d, greptile={score}/5): {pr['title']}" + ) + return + + comment_body = format_close_comment(score, threshold) + gh("pr", "comment", str(pr_number), "--body", comment_body, *repo_args) + + if label: + try: + gh("pr", "edit", str(pr_number), "--add-label", label, *repo_args) + except subprocess.CalledProcessError as exc: + stderr = (exc.stderr or "").strip() + print(f" warn: failed to add label '{label}' to #{pr_number}: {stderr}") + + gh("pr", "close", str(pr_number), *repo_args) + print(f" Closed PR #{pr_number} (greptile={score}/5, age={age_days}d)") + + +def evaluate_pr( + pr: dict, + now: dt.datetime, + min_age_days: int, + min_score: int, + repo: str | None, + optout_labels: set[str], + allowlist: frozenset[str] = ALLOWLIST_LOGINS, +) -> tuple[str, int | None, int | None]: + """Decide what to do with `pr` on this triage run. + + Returns (action, score_or_none, age_days_or_none) where action is one of: + "skip-too-young", "skip-optout-label", "skip-not-allowlisted", + "skip-internal", "skip-no-greptile-score", "skip-score-ok", + "warn-grace", "skip-in-grace-period", or "close". + + Drafts are NOT skipped — the goal is "open PR count == PRs internal + collaborators need to action on", and a draft that Greptile scored <4/5 + is still in that queue. Authors can opt out via the `wip` label (see + `DEFAULT_OPTOUT_LABELS`) if they need to keep a long-lived draft open. + + Grace-period semantics: the first time a PR fails the rubric, the + action is `warn-grace` — the caller should post a warning comment but + NOT close the PR. On a subsequent run, if the warning is still less + than `GRACE_PERIOD_SECONDS` old AND the PR still fails, the action is + `skip-in-grace-period`. Once the warning ages out and the rubric is + still failing, the action is `close`. + """ + if has_optout_label(pr, optout_labels): + return ("skip-optout-label", None, None) + + created = parse_iso8601(pr["createdAt"]) + age_days = (now - created).days + # `min_age_days` defaults to 0 (close as soon as Greptile scores low). + # Set a positive value via --min-age-days for one-off backfill runs that + # want to skip very-young PRs. + if min_age_days > 0 and age_days < min_age_days: + return ("skip-too-young", None, age_days) + + # While the allowlist is active it is the sole author gate: only those + # logins are acted on and the external-only restriction is bypassed for + # them. Otherwise auto-close only external OSS contributors — internal + # contributors (BerriAI org members) handle their own backlog. + login = ((pr.get("author") or {}).get("login") or "").lower() + if allowlist: + if login not in allowlist: + return ("skip-not-allowlisted", None, age_days) + elif not is_external_pr_author(pr, repo): + return ("skip-internal", None, age_days) + + comments = fetch_pr_comments(pr["number"], repo) + extraction = extract_greptile_score(comments) + if extraction is None: + return ("skip-no-greptile-score", None, age_days) + + score, _ = extraction + if score >= min_score: + return ("skip-score-ok", score, age_days) + + grace_age = seconds_since_last_grace_warning(comments, now=now) + if grace_age is None: + return ("warn-grace", score, age_days) + if grace_age < GRACE_PERIOD_SECONDS: + return ("skip-in-grace-period", score, age_days) + + return ("close", score, age_days) + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--repo", + type=str, + default=None, + help="Repository (owner/repo). Auto-detected if omitted.", + ) + parser.add_argument( + "--min-age-days", + type=int, + default=0, + help=( + "Minimum age (in days) before a PR is eligible. Default 0 = " + "close as soon as Greptile flags it. Set a positive value for " + "one-off backfill runs that want to spare very-young PRs." + ), + ) + parser.add_argument( + "--min-score", + type=int, + default=4, + choices=range(1, 6), + help="Greptile score below which a PR is closed (default: 4 -> closes <4/5).", + ) + parser.add_argument( + "--optout-label", + action="append", + default=None, + help=( + "Label(s) that exempt a PR from auto-close. Repeat to add more. " + "Case-insensitive. When omitted, defaults to " + f"{list(DEFAULT_OPTOUT_LABELS)!r}; passing this flag REPLACES the " + "defaults (argparse `append` with a mutable default would append " + "instead, which we explicitly avoid)." + ), + ) + parser.add_argument( + "--close-label", + type=str, + default=None, + help=( + "Optional label to add to PRs that get auto-closed " + "(e.g. 'auto-closed-low-quality'). Must already exist on the repo." + ), + ) + parser.add_argument( + "--close", + action="store_true", + help="Actually close matching PRs (default is dry-run).", + ) + parser.add_argument( + "--limit", + type=int, + default=None, + help="Maximum number of PRs to close in one run (safety net).", + ) + args = parser.parse_args() + + dry_run = not args.close + if dry_run: + print("=== DRY RUN MODE (pass --close to actually close PRs) ===\n") + + print("Fetching open PRs...") + prs = fetch_open_prs(args.repo) + print(f"Found {len(prs)} open PRs.\n") + + now = dt.datetime.now(dt.timezone.utc) + optout_labels = set(args.optout_label or DEFAULT_OPTOUT_LABELS) + + closed = 0 + summary = { + "close": 0, + "warn-grace": 0, + "skip-in-grace-period": 0, + "skip-too-young": 0, + "skip-optout-label": 0, + "skip-not-allowlisted": 0, + "skip-internal": 0, + "skip-no-greptile-score": 0, + "skip-score-ok": 0, + } + + # `warned` tracks grace-warning comments posted in this run so the + # `--limit` safety net bounds *all* destructive write actions, not + # just closures. Without this cap, a backlog of PRs failing the + # threshold simultaneously could flood contributors with comments. + warned = 0 + for pr in sorted(prs, key=lambda p: p["createdAt"]): + try: + action, score, age_days = evaluate_pr( + pr, + now, + args.min_age_days, + args.min_score, + args.repo, + optout_labels, + ) + summary[action] = summary.get(action, 0) + 1 + + if action == "warn-grace": + assert score is not None + print( + f"#{pr['number']}: \"{pr['title']}\" " + f"(age={age_days}d, greptile={score}/5) -> warn-grace" + ) + post_grace_warning( + pr, + score=score, + threshold=args.min_score, + repo=args.repo, + dry_run=dry_run, + ) + if not dry_run: + warned += 1 + if args.limit is not None and (warned + closed) >= args.limit: + print( + f"\nReached --limit={args.limit} " + f"(closed={closed}, warned={warned}); stopping." + ) + break + continue + + if action != "close": + continue + + assert score is not None and age_days is not None + print( + f"#{pr['number']}: \"{pr['title']}\" " + f"(age={age_days}d, greptile={score}/5) -> close" + ) + close_pr( + pr, + score=score, + threshold=args.min_score, + age_days=age_days, + repo=args.repo, + dry_run=dry_run, + label=args.close_label, + ) + + if not dry_run: + closed += 1 + if args.limit is not None and (warned + closed) >= args.limit: + print( + f"\nReached --limit={args.limit} " + f"(closed={closed}, warned={warned}); stopping." + ) + break + except Exception as exc: # noqa: BLE001 - per-PR errors don't abort the sweep + summary["error"] = summary.get("error", 0) + 1 + print( + f"!! PR #{pr.get('number')}: {exc}", + file=sys.stderr, + ) + continue + + print("\n=== Summary ===") + for key, value in summary.items(): + print(f" {key:28s} {value}") + if dry_run: + print(f"\nTotal would close: {summary['close']}") + else: + print(f"\nTotal closed: {closed}") + print( + f"Total {'would warn (grace)' if dry_run else 'warned (grace)'}: " + f"{summary['warn-grace']}" + ) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/.github/scripts/triage-requirements.txt b/.github/scripts/triage-requirements.txt new file mode 100644 index 00000000000..a18f05fbb95 --- /dev/null +++ b/.github/scripts/triage-requirements.txt @@ -0,0 +1,282 @@ +# Hash-pinned dependency set for the Agent Shin triage scripts. +# Installed in privileged triage workflows, so every package is pinned to an +# exact version with SHA-256 hashes and installed with pip --require-hashes. +# +# Regenerate after bumping openai: +# echo 'openai==' \ +# | 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 diff --git a/.github/scripts/triage_rollout_heads_up.py b/.github/scripts/triage_rollout_heads_up.py new file mode 100644 index 00000000000..a5dedb1c9e7 --- /dev/null +++ b/.github/scripts/triage_rollout_heads_up.py @@ -0,0 +1,557 @@ +#!/usr/bin/env python3 +"""One-shot 7-day heads-up sweep for the Agent Shin rollout. + +Posts a friendly "the OSS triage bot kicks in next Monday" comment on every +open external PR/issue that currently *would* fail the new rubric — i.e., +every PR/issue Agent Shin would close once the rollout completes. The point +is to give contributors a full week to fix their description before the bot +ever takes a destructive action, so nobody is surprised by an auto-close. + +The script is designed to run **exactly once** at rollout, fired by a manual +``workflow_dispatch`` (``dry_run=false``) on the heads-up workflow. Re-runs +are safe: every comment is stamped with the hidden ``HEADS_UP_MARKER`` and +PRs/issues that already carry the marker are skipped. + +Dry-run vs. real run +-------------------- +Defaults to dry-run. Passing ``--close`` flips into real mode. Every GitHub +mutation goes through ``_agent_shin_actions``, which has a one-line +``if dry_run: log else: do_it`` per call, so the only difference between a +dry-run preview and the real run is the call site that actually hits the +GitHub API. + +Local preview:: + + python3 .github/scripts/triage_rollout_heads_up.py --repo BerriAI/litellm + +Real run (the manual rollout dispatch uses this):: + + python3 .github/scripts/triage_rollout_heads_up.py --repo BerriAI/litellm --close +""" + +from __future__ import annotations + +import argparse +import datetime as dt +import json +import os +import sys +from pathlib import Path +from typing import Any + +# Make the sibling triage_with_llm + _agent_shin_actions importable when this +# script is invoked directly (the GitHub workflow does `python3 .github/scripts/...`). +_SCRIPTS_DIR = Path(__file__).resolve().parent +if str(_SCRIPTS_DIR) not in sys.path: + sys.path.insert(0, str(_SCRIPTS_DIR)) + +from _agent_shin_actions import maybe_post_comment # noqa: E402 +from agent_shin_shared import ( # noqa: E402 + AGENT_SHIN_DEFAULT_BOT_LOGIN, + ALLOWLIST_LOGINS, + list_open_items, +) +from triage_with_llm import ( # noqa: E402 + DEFAULT_MODEL, + call_llm_judge, + fetch_issue, + fetch_pr, + gh, + is_internal_contributor, + review_gate, + triage, +) + +# Hidden marker so re-runs skip PRs/issues we've already notified. Distinct from +# the within-grace / ready / regressed markers so it can't be confused with the +# steady-state lifecycle comments. +HEADS_UP_MARKER = "" + +# Placeholder until the litellm-docs PR ships. The rollout blog post explains +# the new rubric, the 7-day grace, and how to recover after an auto-close. +# TODO(docs): replace with the canonical URL once the litellm-docs PR merges. +ROLLOUT_BLOG_URL = "https://docs.litellm.ai/docs/agent_shin_triage_rollout" + +# Default cutoff is one week from "now". Computed at runtime so the wording +# stays correct even if the rollout is merged later than planned. The user can +# override with --close-on YYYY-MM-DD when running the script manually. +DEFAULT_GRACE_DAYS = 7 + +# The daily auto-close sweeps (close_low_quality_prs.yml at 09:00 UTC and +# review_gate.yml at 09:30 UTC) are what actually close a still-failing item, +# so the deadline we promise contributors has to name that wall-clock moment. +ACTIVATION_TIME_UTC = "09:00 UTC" + + +def _format_cutoff(cutoff: dt.date) -> str: + """Human-readable, timezone-explicit cutoff, e.g. ``Monday, June 1, 2026 + (09:00 UTC)`` — the moment a still-failing PR/issue gets closed.""" + return ( + f"{cutoff.strftime('%A, %B')} {cutoff.day}, {cutoff.year} " + f"({ACTIVATION_TIME_UTC})" + ) + + +def _rubric_section_pr() -> str: + return ( + "**Going forward, every external PR needs ONE of:**\n" + "\n" + "- A linked GitHub issue using a closing keyword: " + "`Fixes #1234`, `Closes #1234`, or `Resolves #1234`, OR\n" + "- All three of: a clear **problem description**, **expected vs. " + "actual behavior**, and **end-to-end QA proof** (at least one of a " + "short screen recording / video, before/after screenshots, or the " + "exact commands you ran with their real output; mocked or stubbed " + "runs don't count).\n" + "\n" + "PRs also need a **Greptile confidence score of 4/5 or higher** before " + "the bot will tag them `ready for review`. You can `@greptileai` to " + "request a fresh review at any time, including after the PR is closed." + ) + + +def _rubric_section_issue() -> str: + return ( + "**Going forward, every external issue needs:**\n" + "\n" + "- For **bug reports**: end-to-end evidence of the bug (at least one " + "of a screen recording / video, a screenshot, or the exact commands " + "you ran with their real output / traceback) plus expected vs. actual " + "behavior. Written steps with no run output don't count, and mocked " + "or stubbed runs don't count.\n" + "- For **feature requests**: a clear description of the proposed " + "feature plus a use case + concrete example (config, API call, UI " + "flow, or scenario showing what's blocked today)." + ) + + +def _description_only_note(kind: str) -> str: + noun = "PR" if kind == "pr" else "issue" + return ( + f"⚠️ **The requirements must live in the {noun} *description*, not in " + "comments.** Some PRs/issues collect 100+ comments from humans and " + "bots; reading the entire thread on every triage run would balloon " + "GitHub API usage (we'd start getting 429'd) and blow out the LLM " + "judge's context. The bot only reads the description, so anything " + "you add as a comment will be invisible to it." + ) + + +def _missing_section(verdict: dict, greptile_score: int | None) -> str: + """Bullet list of what's currently missing on this PR/issue. + + Combines the LLM judge's `missing` list (rubric items) with a Greptile + shortfall (for PRs) so the contributor sees one list of things to fix. + """ + missing = list(verdict.get("missing") or []) + if greptile_score is not None and greptile_score < 4: + missing.insert( + 0, + f"Greptile's most recent review scored this PR {greptile_score}/5 " + "(below the 4/5 bar Agent Shin will require).", + ) + if not missing: + return ( + "_The bot couldn't articulate a specific missing piece; see the " + "rubric link above and double-check the description includes all " + "of it before the rollout._" + ) + bullets = "\n".join(f"- {m}" for m in missing) + return f"**What this one is currently missing:**\n\n{bullets}" + + +def _recovery_section(kind: str) -> str: + if kind == "pr": + return ( + "**If the bot closes this PR after the rollout:** update the " + "description with the missing pieces, then either open a fresh " + "PR or comment `@agent-shin reconsider` on the closed PR. If " + "Greptile re-scores you at 4/5 or higher I'll reopen and tag " + "the PR `ready for review`. (`@greptileai` works on closed PRs " + "too; a fresh review is one of the signals that lifts you back " + "into the queue.) This is **not** us losing interest in your " + "change; far from it. We just need open PRs to be a list of " + "things a maintainer can act on, so we can get to yours faster." + ) + return ( + "**If the bot closes this issue after the rollout:** edit the issue " + "description to add the missing pieces, then comment `@agent-shin " + "reconsider` on the closed issue. I'll re-evaluate and, if the rubric " + "is met, reopen it. (GitHub doesn't let external authors reopen an " + "issue a maintainer or bot closed, so the comment is the reliable " + "path.) This is **not** us saying the bug isn't real or the request " + "isn't useful; it's so the remaining open issues are a list of things " + "a maintainer can act on." + ) + + +def format_heads_up_comment( + *, kind: str, verdict: dict, greptile_score: int | None, cutoff: dt.date +) -> str: + """Compose the friendly 7-day heads-up comment posted on a failing PR/issue.""" + noun = "PR" if kind == "pr" else "issue" + rubric = _rubric_section_pr() if kind == "pr" else _rubric_section_issue() + cutoff_str = _format_cutoff(cutoff) + explanation = (verdict.get("explanation") or "").strip() + explanation_block = ( + f"> _(The judge's note for this one: {explanation})_\n\n" if explanation else "" + ) + + return ( + "🚅 **Heads-up: we're turning on the OSS triage bot in " + f"{DEFAULT_GRACE_DAYS} days, on {cutoff_str}.**\n" + "\n" + "We're rolling out **Agent Shin**, an LLM-as-judge triage bot for " + f"external {noun}s. Once it's live, the bot reads each open " + f"{noun}'s description, scores it against a small rubric, and " + f"auto-closes any {noun} that's missing the basics, with a single " + f"comment explaining what's missing and how to recover. Full " + f"context: [Agent Shin rollout blog post]({ROLLOUT_BLOG_URL}).\n" + "\n" + f"{rubric}\n" + "\n" + f"{_description_only_note(kind)}\n" + "\n" + f"{_missing_section(verdict, greptile_score)}\n" + "\n" + f"{explanation_block}" + "**Timeline (you have a week):**\n" + "\n" + f"- We turn the bot on in {DEFAULT_GRACE_DAYS} days, on " + f"**{cutoff_str}**. You have until then to update this {noun}'s " + "description with the missing pieces above.\n" + f"- If this {noun} still fails the rubric at **{cutoff_str}**, " + "we'll close it.\n" + f"- From then on the bot runs daily, and every {noun} that fails " + "the rubric gets a **2-hour lifetime**: one warning comment, then " + "auto-close 2 hours later.\n" + "\n" + f"{_recovery_section(kind)}\n" + "\n" + f"{HEADS_UP_MARKER}" + ) + + +def _list_open_numbers(repo: str, kind: str) -> list[int]: + """Return every open PR or issue number in ``repo``. + + Delegates to ``list_open_items`` so the full backlog is fetched (no cap) + and the `gh {pr,issue} list` invocation stays in one shared place. ``gh + issue list`` would include PRs, but ``list_open_items`` uses the dedicated + command per kind, so the two never mix. + """ + return [ + item["number"] for item in list_open_items(kind, repo=repo, fields="number") + ] + + +def _has_heads_up_marker(item: dict) -> bool: + """Cheap fast-path: check the PR/issue body itself for the marker. + + The marker is appended to the *comment* we post, not the body, so this + will only fire if the body literally contains the marker text. We still + do the comment-marker check separately below; this body check just lets + us short-circuit for PRs/issues that quote the marker for any reason. + """ + body = item.get("body") or "" + return HEADS_UP_MARKER in body + + +def _comments_have_marker(repo: str, number: int) -> bool: + """True if the bot already posted a comment carrying the marker. + + Used for idempotency: a re-run skips items the previous run notified. + Filters by author (matching the sibling marker-checks in + ``triage_with_llm._has_marker`` and + ``agent_shin_shared.seconds_since_latest_marker_comment``) so a + contributor who quotes the heads-up via GitHub's "Quote reply" — which + preserves HTML comments in the raw markdown — can't trick the + idempotency check into silently skipping a real heads-up. + + Comments live on the unified issues endpoint regardless of whether the + item is a PR or an issue, so no ``kind`` argument is required here. + """ + expected_login = ( + os.environ.get("AGENT_SHIN_BOT_LOGIN") or AGENT_SHIN_DEFAULT_BOT_LOGIN + ).lower() + raw = gh( + "api", + "--paginate", + f"repos/{repo}/issues/{number}/comments?per_page=100", + ) + for line in raw.splitlines(): + line = line.strip() + if not line: + continue + try: + payload = json.loads(line) + except json.JSONDecodeError: + continue + comments = payload if isinstance(payload, list) else [payload] + for comment in comments: + author = ((comment.get("user") or {}).get("login") or "").lower() + if author != expected_login: + continue + if HEADS_UP_MARKER in (comment.get("body") or ""): + return True + return False + + +def _evaluate_pr(*, repo: str, number: int, model: str, judge: Any = None) -> dict: + """Run the future PR rubric (review_gate) in dry-run and return the result.""" + return review_gate( + repo=repo, + number=number, + close=False, # we only want the verdict, never act here + model=model, + judge=judge, + ) + + +def _evaluate_issue(*, repo: str, number: int, model: str, judge: Any = None) -> dict: + """Run the future issue rubric (triage kind='issue') in dry-run.""" + return triage( + repo=repo, + kind="issue", + number=number, + close=False, + model=model, + judge=judge, + ) + + +def _would_be_closed(kind: str, result: dict) -> bool: + """True if the future triage would auto-close this PR/issue based on the + rubric (regardless of grace-period gating). + + For PRs we trust ``review_gate``'s ``passing`` field — it combines the LLM + verdict and the Greptile score. For issues we read the LLM verdict + directly. Both fields are ``None``/missing on skip paths + (skip-internal-author, skip-llm-error, etc.) where the future bot would + NOT close the item — those return False. + """ + if kind == "pr": + passing = result.get("passing") + if passing is None: + return False # skipped — nothing for the heads-up to warn about + return passing is False + verdict = result.get("verdict") or {} + return (verdict.get("verdict") or "").lower() == "fail" + + +def _process_one( + *, + repo: str, + kind: str, + number: int, + model: str, + cutoff: dt.date, + dry_run: bool, + judge: Any = None, + skip_marker_check: bool = False, + allowlist: frozenset[str] = ALLOWLIST_LOGINS, +) -> dict: + """Evaluate one PR/issue and post a heads-up if it would be auto-closed. + + Returns a per-item dict for the summary table. + """ + base = {"kind": kind, "number": number} + fetcher = fetch_pr if kind == "pr" else fetch_issue + item = fetcher(repo, number) + + if (item.get("state") or "") != "open": + return {**base, "action": "skip-not-open"} + if allowlist: + login = (item.get("user") or {}).get("login") or "" + if login.lower() not in allowlist: + return {**base, "action": "skip-not-allowlisted"} + elif is_internal_contributor(item): + return {**base, "action": "skip-internal-author"} + if not skip_marker_check and _has_heads_up_marker(item): + return {**base, "action": "skip-already-marked-in-body"} + if not skip_marker_check and _comments_have_marker(repo, number): + return {**base, "action": "skip-already-notified"} + + if kind == "pr": + result = _evaluate_pr(repo=repo, number=number, model=model, judge=judge) + else: + result = _evaluate_issue(repo=repo, number=number, model=model, judge=judge) + + if not _would_be_closed(kind, result): + return {**base, "action": "skip-passing", "evaluator": result.get("action")} + + verdict = result.get("verdict") or {} + greptile_score = result.get("greptile_score") if kind == "pr" else None + comment = format_heads_up_comment( + kind=kind, verdict=verdict, greptile_score=greptile_score, cutoff=cutoff + ) + maybe_post_comment(repo, number, comment, dry_run=dry_run) + return { + **base, + "action": "heads-up-posted" if not dry_run else "would-post-heads-up", + "verdict": (verdict.get("verdict") or "").lower(), + "greptile_score": greptile_score, + } + + +def _print_summary(results: list[dict]) -> None: + """Tally per-action counts so a dry-run preview tells you at a glance how + many comments the real run would post.""" + counts: dict[str, int] = {} + for r in results: + counts[r["action"]] = counts.get(r["action"], 0) + 1 + print("\n=== rollout heads-up summary ===") + for action in sorted(counts): + print(f" {action:35s} {counts[action]}") + print(f" total {len(results)}") + + +def run( + *, + repo: str, + close: bool, + cutoff: dt.date, + model: str, + kinds: tuple[str, ...] = ("pr", "issue"), + judge: Any = None, + only_numbers: dict[str, list[int]] | None = None, + skip_marker_check: bool = False, +) -> list[dict]: + """Sweep ``repo`` and post heads-up comments. Returns the per-item results.""" + dry_run = not close + if dry_run: + print( + f"[DRY RUN] sweeping {repo}; --close not passed, no comments will be posted." + ) + else: + print(f"[REAL RUN] sweeping {repo}; comments WILL be posted.") + print(f"Cutoff date in comment body: {cutoff.isoformat()}") + + results: list[dict] = [] + for kind in kinds: + if only_numbers and kind in only_numbers: + numbers = list(only_numbers[kind]) + else: + numbers = _list_open_numbers(repo, kind) + print(f"\n--- {kind}s: {len(numbers)} open ---") + for n in numbers: + try: + result = _process_one( + repo=repo, + kind=kind, + number=n, + model=model, + cutoff=cutoff, + dry_run=dry_run, + judge=judge, + skip_marker_check=skip_marker_check, + ) + except ( + Exception + ) as exc: # noqa: BLE001 - per-item errors don't abort the sweep + result = { + "kind": kind, + "number": n, + "action": "error", + "error": str(exc), + } + print(f"!! {kind}#{n}: {exc}", file=sys.stderr) + print(f" {kind}#{n}: {result['action']}") + results.append(result) + _print_summary(results) + return results + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--repo", required=True, help="owner/repo") + parser.add_argument( + "--close", + action="store_true", + help=( + "Actually post comments. Without this flag the script is in " + "dry-run mode and only logs what it would do." + ), + ) + parser.add_argument( + "--close-on", + type=dt.date.fromisoformat, + default=None, + help=( + "Cutoff date shown in the heads-up comment as the rollout date " + f"(default: today + {DEFAULT_GRACE_DAYS} days)." + ), + ) + parser.add_argument( + "--model", + default=os.environ.get("TRIAGE_MODEL") or DEFAULT_MODEL, + help=f"Model for the rubric LLM judge (default: {DEFAULT_MODEL}).", + ) + parser.add_argument( + "--kind", + choices=("pr", "issue", "both"), + default="both", + help="Restrict the sweep to PRs or issues only (default: both).", + ) + parser.add_argument( + "--only-pr", + type=int, + action="append", + default=[], + help="Limit the PR sweep to these PR numbers (repeat for several).", + ) + parser.add_argument( + "--only-issue", + type=int, + action="append", + default=[], + help="Limit the issue sweep to these issue numbers (repeat for several).", + ) + parser.add_argument( + "--ignore-existing-marker", + action="store_true", + help=( + "Re-post on PRs/issues that already carry the heads-up marker. " + "Useful for testing the comment wording on a known PR." + ), + ) + args = parser.parse_args() + + cutoff = args.close_on or ( + dt.datetime.now(dt.timezone.utc).date() + dt.timedelta(days=DEFAULT_GRACE_DAYS) + ) + + kinds: tuple[str, ...] + if args.kind == "pr": + kinds = ("pr",) + elif args.kind == "issue": + kinds = ("issue",) + else: + kinds = ("pr", "issue") + + only: dict[str, list[int]] = {} + if args.only_pr: + only["pr"] = args.only_pr + if args.only_issue: + only["issue"] = args.only_issue + + # The script must NOT hit the LLM in dry-run if no key is set — we still + # want a useful preview that says "skip-no-llm-key" for items that would + # have been judged. Production runs require OPENAI_API_KEY. + if args.close and not os.environ.get("OPENAI_API_KEY"): + parser.error("OPENAI_API_KEY must be set for --close (real-run) mode.") + + run( + repo=args.repo, + close=args.close, + cutoff=cutoff, + model=args.model, + kinds=kinds, + only_numbers=only or None, + skip_marker_check=args.ignore_existing_marker, + ) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/.github/scripts/triage_with_llm.py b/.github/scripts/triage_with_llm.py new file mode 100644 index 00000000000..d2536058e01 --- /dev/null +++ b/.github/scripts/triage_with_llm.py @@ -0,0 +1,1778 @@ +#!/usr/bin/env python3 +""" +Agent Shin — LLM-as-judge triage for external OSS pull requests and issues. + +Evaluates a single PR or issue against the contribution rubric and, when the +LLM judge marks it as failing, posts an explanatory comment + closes the +PR/issue. Re-triggers on `reopened` so contributors can iterate back in by +filling in the missing pieces and reopening. + +Internal BerriAI contributors (`author_association` in {OWNER, MEMBER, +COLLABORATOR}) and bot accounts are skipped entirely. + +Usage: + triage_with_llm.py --repo owner/repo --pr 1234 + triage_with_llm.py --repo owner/repo --issue 5678 + triage_with_llm.py --repo owner/repo --pr 1234 --close # actually close + triage_with_llm.py --repo owner/repo --pr 1234 --print-prompt # show prompt + +Defaults are SAFE: without `--close` the script writes a verdict to stdout (and, +when running in GitHub Actions, to $GITHUB_STEP_SUMMARY) but takes no GitHub +write actions. + +Environment: + GH_TOKEN / GITHUB_TOKEN - for `gh` CLI auth (auto-set in Actions) + OPENAI_API_KEY - required when --close is passed + OPENAI_BASE_URL - optional (route to any OpenAI-compatible API) + TRIAGE_MODEL - optional model override (default: gpt-5.4-mini) +""" + +from __future__ import annotations + +import argparse +import datetime as dt +import json +import os +import re +import subprocess +import sys +import textwrap +import urllib.parse +from typing import Any, 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/triage_with_llm.py ...`) and +# also when the tests load this script via +# `importlib.util.spec_from_file_location`. +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, + AGENT_SHIN_DEFAULT_BOT_LOGIN, + ALLOWLIST_LOGINS, + GRACE_COMMENT_MARKER, + GRACE_PERIOD_SECONDS, + GREPTILE_BOT_LOGINS, + SCORE_PATTERN, + extract_greptile_score, + gh, + parse_iso8601, + seconds_since_latest_marker_comment, +) + +DEFAULT_MODEL = "gpt-5.4-mini" + +INTERNAL_ASSOCIATIONS = frozenset({"OWNER", "MEMBER", "COLLABORATOR"}) + +# `AGENT_SHIN_DEFAULT_BOT_LOGIN` is imported from `agent_shin_shared`. +# When the workflow uses the default `secrets.GITHUB_TOKEN`, the +# closure / reopen event's `actor.login` is `github-actions[bot]`. The +# env override `AGENT_SHIN_BOT_LOGIN` exists for local debugging and for +# repos that wire Agent Shin to a PAT. + +# HTML marker appended to every reconsider verdict comment. We grep for this +# on subsequent reconsider triggers to enforce a short cooldown so that +# repeated `@agent-shin reconsider` comments don't burn CI/LLM budget. +# Using a unique HTML comment keeps the marker invisible to humans while +# being trivially greppable from a comments-list API response. +RECONSIDER_COMMENT_MARKER = "" + +# Minimum gap between two reconsider verdicts on the same PR/issue. Set to +# 10 minutes — long enough that a contributor can't trivially spam the +# trigger, short enough that a genuine "I just pushed a fix and reupdated +# the body" iteration loop isn't punished. +RECONSIDER_RATE_LIMIT_SECONDS = 600 + +# `GRACE_COMMENT_MARKER` (HTML marker on the grace-period warning comment +# posted on the first low-quality detection — used on subsequent triage +# runs to detect that a warning was already posted and measure how long +# ago it was 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 daily Greptile sweep and the +# LLM judge agree on the same marker and duration. + +# --- Review-gate ("ready for review" label lifecycle) configuration ---------- +# The review gate keeps a single label in sync with whether a PR currently +# clears BOTH quality bars: the LLM rubric (clear problem + expected/actual + +# QA proof, or a linked issue) AND Greptile's most recent confidence score. +READY_FOR_REVIEW_LABEL = "ready for review" +DEFAULT_GRACE_DAYS = 1 # 24h before an un-passing, un-tagged PR is auto-closed +DEFAULT_MIN_GREPTILE_SCORE = 4 # Greptile < 4/5 counts as "not passing" + +# Hidden HTML-comment markers stamped into review-gate comments. They never +# render in the GitHub UI but let the gate detect its own prior actions so it +# (a) posts the within-grace "what's missing" notice at most once and (b) can +# tell a first-time pass ("ready for review") from a recovery after a +# regression ("all clear again"). +READY_MARKER = "" +REGRESSED_MARKER = "" +WITHIN_GRACE_MARKER = "" + +# `GREPTILE_BOT_LOGINS` (Greptile's GitHub App login variants — +# `greptile-apps[bot]` in REST API comments, `greptile-apps` in +# `gh pr view --json` output) and `SCORE_PATTERN` (regex matching lines +# like `Confidence Score: 3/5`) are imported from `agent_shin_shared` +# so the daily sweep and the review gate read the score through the +# same set of logins / patterns. + +# `AGENT_SHIN_CLOSE_MARKER` is imported from `agent_shin_shared` so this LLM +# judge and the daily Greptile sweep stamp the same marker on their close +# comments — `was_closed_by_agent_shin` keys the reconsider reopen path off it. + +# Model families that require `reasoning_effort` to be set, and that reject +# `temperature != 1` unless `reasoning_effort` is "none". For these models we +# pass `reasoning_effort="none"` so a `temperature=0` deterministic judgment +# is still accepted. See litellm/llms/openai/chat/gpt_5_transformation.py for +# the full set of constraints LiteLLM applies to these models. +GPT5_FAMILY_PREFIX = "gpt-5" + +# Regexes for picking off "obvious passes" without burning LLM tokens. +# +# Keep this list to GitHub's documented PR-closing keywords only +# (https://docs.github.com/issues/tracking-your-work-with-issues/linking-a-pull-request-to-an-issue). +# Casual mentions like "see #1234" or "ref #1234" are intentionally NOT +# auto-passed — they should fall through to the LLM judge, which has the +# stricter rubric "a bare issue number without a closing keyword counts only +# if it's clearly the related issue (not a passing mention)". +LINKED_ISSUE_PATTERN = re.compile( + r"\b(?:fixes|fix|fixed|closes|close|closed|resolves|resolve|resolved)\s+" + r"(?:#\d+|https?://github\.com/[\w.-]+/[\w.-]+/issues/\d+)", + re.IGNORECASE, +) +HTML_COMMENT_PATTERN = re.compile(r"", re.DOTALL) + + +# --------------------------------------------------------------------------- +# gh helpers +# +# `gh` is imported from `agent_shin_shared` so a future change (timeout, +# logging, retry) only needs to be made once. + + +def fetch_pr(repo: str, number: int) -> dict: + """Return the full GitHub REST representation of a PR.""" + return json.loads(gh("api", f"repos/{repo}/pulls/{number}")) + + +def fetch_issue(repo: str, number: int) -> dict: + """Return the full GitHub REST representation of an issue.""" + return json.loads(gh("api", f"repos/{repo}/issues/{number}")) + + +def post_comment(repo: str, number: int, body: str) -> None: + """Post an issue-style comment (works for both issues and PRs).""" + gh( + "api", + f"repos/{repo}/issues/{number}/comments", + "-X", + "POST", + "-f", + f"body={body}", + ) + + +def close_pr(repo: str, number: int) -> None: + """Close a pull request (state=closed).""" + gh( + "api", + f"repos/{repo}/pulls/{number}", + "-X", + "PATCH", + "-f", + "state=closed", + ) + + +def reopen_pr(repo: str, number: int) -> None: + """Reopen a previously-closed pull request (state=open). + + Used by the `@agent-shin reconsider` comment-trigger flow: the bot has + write access via GH_TOKEN, so it can reopen on the contributor's behalf + even though GitHub doesn't let the OSS author do it themselves. + """ + gh( + "api", + f"repos/{repo}/pulls/{number}", + "-X", + "PATCH", + "-f", + "state=open", + ) + + +def close_issue(repo: str, number: int, *, not_planned: bool = True) -> None: + """Close an issue, marking state_reason=not_planned by default.""" + args = [ + "api", + f"repos/{repo}/issues/{number}", + "-X", + "PATCH", + "-f", + "state=closed", + ] + if not_planned: + args.extend(["-f", "state_reason=not_planned"]) + gh(*args) + + +def reopen_issue(repo: str, number: int) -> None: + """Reopen a previously-closed issue (state=open, state_reason=reopened).""" + gh( + "api", + f"repos/{repo}/issues/{number}", + "-X", + "PATCH", + "-f", + "state=open", + "-f", + "state_reason=reopened", + ) + + +def add_label(repo: str, number: int, label: str) -> None: + """Add a label to a PR/issue (GitHub creates the label if it's missing).""" + gh( + "api", + f"repos/{repo}/issues/{number}/labels", + "-X", + "POST", + "-f", + f"labels[]={label}", + ) + + +def remove_label(repo: str, number: int, label: str) -> None: + """Remove a label from a PR/issue. A missing label (404) is not an error.""" + encoded = urllib.parse.quote(label, safe="") + try: + gh( + "api", + f"repos/{repo}/issues/{number}/labels/{encoded}", + "-X", + "DELETE", + ) + except subprocess.CalledProcessError as exc: + stderr = (exc.stderr or "").lower() + if "404" in stderr or "not found" in stderr: + return + raise + + +def _iter_paginated_json(*api_args: str) -> Any: + """Yield JSON objects from `gh api --paginate ... -q '.[]'`. + + `gh api --paginate` on a JSON-array endpoint concatenates pages into + one stream; `-q '.[]'` flattens that stream into newline-delimited + objects (jq-style). This keeps memory bounded for chatty endpoints + like issue events/comments on long-lived PRs. + """ + raw = gh("api", "--paginate", *api_args, "-q", ".[]") + for line in raw.splitlines(): + line = line.strip() + if not line: + continue + try: + yield json.loads(line) + except json.JSONDecodeError: + # A malformed line should not blow up the whole guard. Skip and + # carry on — at worst the guard fail-closes (returns False / + # None) and the caller treats it as "unknown". + continue + + +def fetch_last_close_event( + repo: str, number: int +) -> tuple[str | None, dt.datetime | None]: + """Return the actor login and timestamp of the most recent `closed` event. + + Either field may be None: actor when the events API returns nothing + (unusual for a closed item, but possible on transient errors), and + timestamp when the event lacks `created_at` or the value can't be + parsed. `was_closed_by_agent_shin` fail-closes on either. + """ + actor: str | None = None + closed_at: dt.datetime | None = None + for event in _iter_paginated_json(f"repos/{repo}/issues/{number}/events"): + if event.get("event") != "closed": + continue + actor = (event.get("actor") or {}).get("login") + created = event.get("created_at") + if not created: + closed_at = None + continue + try: + closed_at = parse_iso8601(created) + except ValueError: + closed_at = None + return actor, closed_at + + +# How much older than the latest `closed` event the Agent Shin marker +# comment is allowed to be while still counting as "this close was Agent +# Shin's". Agent Shin posts the close comment immediately before closing, +# so the marker timestamp is normally at most a few seconds before the +# close event; the buffer just absorbs clock skew between the comments +# API and the events API. +AGENT_SHIN_CLOSE_MARKER_SKEW_SECONDS = 300 + + +def was_closed_by_agent_shin( + repo: str, number: int, *, bot_login: str | None = None +) -> bool: + """Return True iff Agent Shin itself most-recently closed this PR/issue. + + This is the guard that stops `@agent-shin reconsider` from reopening an + item Agent Shin did not close — a maintainer closing for non-rubric + reasons (security, duplicate, design rejection), or a different workflow + (stale/duplicate sweeps) closing under the shared `github-actions[bot]` + identity. Three independent signals must all hold, because that identity + is not unique to Agent Shin and a marker comment from a prior + closed/reopened cycle would otherwise vouch for an unrelated close: + + 1. The most recent `closed` event's actor is the bot identity. + 2. Agent Shin left one of its auto-close comments, detected via + `AGENT_SHIN_CLOSE_MARKER`. The actor check alone can't tell an + Agent Shin close from any other `github-actions[bot]` close. + 3. That marker comment was posted at (or just before) the latest + close event, not on a previous close in an + Agent-Shin-close -> reconsider-reopen -> other-bot-reclose cycle. + + The check is intentionally fail-closed: any uncertainty about who closed + the item is treated as "not Agent Shin" so the destructive reopen path + stays gated. + """ + expected = ( + bot_login + or os.environ.get("AGENT_SHIN_BOT_LOGIN") + or AGENT_SHIN_DEFAULT_BOT_LOGIN + ).lower() + actor, closed_at = fetch_last_close_event(repo, number) + if not actor or actor.lower() != expected or closed_at is None: + return False + marker_seconds = seconds_since_last_agent_shin_close( + repo, number, bot_login=bot_login + ) + if marker_seconds is None: + return False + close_age_seconds = (dt.datetime.now(dt.timezone.utc) - closed_at).total_seconds() + return marker_seconds <= close_age_seconds + AGENT_SHIN_CLOSE_MARKER_SKEW_SECONDS + + +def _seconds_since_latest_marker_comment( + repo: str, + number: int, + *, + marker: str, + bot_login: str | None = None, +) -> float | None: + """Return seconds since the bot's most recent comment with ``marker``. + + Fetches comments via `_iter_paginated_json` and delegates the + iteration / author-filter / timestamp logic to + `agent_shin_shared.seconds_since_latest_marker_comment` so the daily + Greptile sweep and the LLM judge use one source of truth for the + "bot already posted X" detection. The wall-clock `now` is resolved + against this module's `dt` so tests that freeze time via + `monkeypatch.setattr(triage_module, "dt", ...)` still apply. + """ + return seconds_since_latest_marker_comment( + _iter_paginated_json(f"repos/{repo}/issues/{number}/comments"), + marker=marker, + bot_login=bot_login, + now=dt.datetime.now(dt.timezone.utc), + ) + + +def seconds_since_last_reconsider_verdict( + repo: str, number: int, *, bot_login: str | None = None +) -> float | None: + """Return seconds since the bot's most recent reconsider verdict comment. + + Detects comments by matching the HTML marker `RECONSIDER_COMMENT_MARKER` + appended by `format_reopen_comment` and + `format_reconsider_still_failing_comment`. Returns None when the bot + has never posted a reconsider verdict on this PR/issue (or when the + only matching comments are missing a `created_at` timestamp, which + shouldn't happen on a real GitHub response). + """ + return _seconds_since_latest_marker_comment( + repo, number, marker=RECONSIDER_COMMENT_MARKER, bot_login=bot_login + ) + + +def seconds_since_last_grace_warning( + repo: str, number: int, *, bot_login: str | None = None +) -> float | None: + """Return seconds since the bot's most recent grace-period warning. + + Detects warning comments by matching the HTML marker + `GRACE_COMMENT_MARKER` appended by `format_grace_warning_pr_comment` + and `format_grace_warning_issue_comment`. Returns None when no + grace warning has ever been posted on this PR/issue — that's the + "first low-quality detection" signal that drives the warning path. + """ + return _seconds_since_latest_marker_comment( + repo, number, marker=GRACE_COMMENT_MARKER, bot_login=bot_login + ) + + +def seconds_since_last_agent_shin_close( + repo: str, number: int, *, bot_login: str | None = None +) -> float | None: + """Return seconds since Agent Shin's most recent auto-close comment. + + Detects close comments by matching `AGENT_SHIN_CLOSE_MARKER` (stamped by + `format_pr_close_comment` / `format_issue_close_comment`). Returns None + when Agent Shin has never closed this PR/issue — the signal + `was_closed_by_agent_shin` uses to keep the reconsider reopen path gated + against closures performed by other workflows sharing the bot identity. + """ + return _seconds_since_latest_marker_comment( + repo, number, marker=AGENT_SHIN_CLOSE_MARKER, bot_login=bot_login + ) + + +# --------------------------------------------------------------------------- +# Author classification + + +def is_internal_contributor(item: dict) -> bool: + """Return True if the PR/issue author should be exempted from triage. + + Fail-safe: if `author_association` is missing or empty (which should never + happen on a successful GitHub REST response but is possible on schema + changes or partial responses), treat the author as INTERNAL so the + destructive close path never fires on an unknown contributor. This matches + the sibling `is_external_pr_author` in `close_low_quality_prs.py`. + """ + login = ((item.get("user") or {}).get("login") or "").lower() + if login.endswith("[bot]") or login in {"dependabot", "github-actions"}: + return True + association = (item.get("author_association") or "").upper() + if not association or association in INTERNAL_ASSOCIATIONS: + return True + return False + + +# --------------------------------------------------------------------------- +# Greptile score + age helpers (`extract_greptile_score`, `parse_iso8601`) +# live in `agent_shin_shared` — they're imported at the top of this module +# so both `triage_with_llm.py` and `close_low_quality_prs.py` share a +# single source of truth for the Confidence-Score regex and ISO-8601 +# parsing. + + +# --------------------------------------------------------------------------- +# Prompt construction + + +def strip_html_comments(text: str) -> str: + """Remove HTML comments — template placeholder text shouldn't fool the judge.""" + return HTML_COMMENT_PATTERN.sub("", text or "") + + +def has_linked_issue(text: str) -> bool: + """Heuristic: does this body link to an open issue (Fixes #123 etc.)?""" + return bool(LINKED_ISSUE_PATTERN.search(strip_html_comments(text or ""))) + + +def build_pr_prompt(*, title: str, body: str) -> str: + cleaned_body = strip_html_comments(body or "").strip() or "(empty)" + # Dedent the static template *before* interpolating dynamic fields so that + # multi-line bodies (whose 2nd+ lines start at column 0) don't defeat the + # common-indent computation in textwrap.dedent. + template = textwrap.dedent(""" + You are "Agent Shin", the OSS triage bot for the LiteLLM open-source + repository (BerriAI/litellm). Decide whether this external pull request + meets the project's contribution standards. + + A PR PASSES triage only if BOTH (1) AND (2) are satisfied. A linked + issue alone is NOT enough — it covers context, not proof. + + (1) CONTEXT — the PR provides AT LEAST ONE of: + (a) A link to a related GitHub issue. Acceptable forms: + "Fixes #1234", "Closes #1234", "Resolves #1234", + "Refs https://github.com/BerriAI/litellm/issues/1234". A + bare "#1234" without a closing keyword counts only if it + is clearly the related issue (not a passing mention). + (b) A clear problem description in the body (what bug or + missing feature this addresses, beyond the title) AND + expected vs. actual behavior (or, for features, "what's + possible now vs. with this PR"). + + (2) END-TO-END QA PROOF: the PR body contains AT LEAST ONE of: + (a) A screen recording / video showing the behavior before + and after the change (the bug reproducing, then the fix + working). For a brand-new feature with no meaningful + "before", a recording of it working end-to-end is fine. + (b) A screenshot (or before/after screenshots) showing the + fix or feature working. + (c) Specific commands that were actually run (curl, python, + a CLI invocation, etc.) PAIRED WITH their real + output, demonstrating the change works end-to-end against + the real system. Commands whose external dependencies + (LLM provider, DB, network) are mocked or stubbed do NOT + satisfy (2c); they are not end-to-end. + + `has_qa_proof` must be set to `true` only when (2a), (2b), + or a non-mocked (2c) is actually present in the body. If the + only "proof" is mocked tests, `has_qa_proof` is `false` and + the verdict is "fail". + + The following do NOT count as QA proof: + - Generic claims like "I tested it", "works locally", "all + tests pass", or a checked "I added tests" checkbox with no + output shown. + - A description of what tests exist or were added, without + their actual output in the PR body. + - `pytest` (or any test runner) executed against the + repository's own unit tests. Those mock the LLM provider, + DB, and network, so they are NOT end-to-end and never + satisfy (2), no matter how much passing output is pasted. + - A linked issue. The linked issue is context (1a), never + proof (2). + + FAIL the PR if EITHER (1) or (2) is missing. Do not bias toward PASS: + if QA proof is absent, the verdict is "fail" even when the rest of + the PR is well-written. + + Respond with a single JSON object, no prose: + + {{ + "verdict": "pass" | "fail", + "linked_issue": boolean, + "has_problem_description": boolean, + "has_expected_vs_actual": boolean, + "has_qa_proof": boolean, + "qa_proof_type": "video" | "screenshot" | "commands_with_output" | "none", + "missing": ["plain-english strings naming what is missing"], + "explanation": "1-2 sentence reasoning for the team to skim" + }} + + --- + PR title: {title} + + PR body: + --- + {cleaned_body} + --- + """).strip() + return template.format(title=title, cleaned_body=cleaned_body) + + +def build_issue_prompt(*, title: str, body: str) -> str: + cleaned_body = strip_html_comments(body or "").strip() or "(empty)" + # Dedent the static template *before* interpolating dynamic fields so that + # multi-line bodies (whose 2nd+ lines start at column 0) don't defeat the + # common-indent computation in textwrap.dedent. + template = textwrap.dedent(""" + You are "Agent Shin", the OSS triage bot for the LiteLLM open-source + repository (BerriAI/litellm). Decide whether this GitHub issue meets + the project's reporting standards. + + For a BUG REPORT the issue PASSES triage only when it contains BOTH: + (1) END-TO-END EVIDENCE OF THE BUG (the "before"; set + `has_repro=true` only when this is present): AT LEAST ONE of: + (a) A screen recording / video of the bug happening. + (b) A screenshot of the bug. + (c) The exact command(s) actually run (curl, python, a CLI + invocation, etc.) PAIRED WITH their real output, traceback, + or logs showing the failure against the real system. + Commands whose external dependencies (LLM provider, DB, + network) are mocked or stubbed do NOT count. + Prose-only "steps to reproduce" with no run output, video, or + screenshot do NOT satisfy (1). + (2) Expected vs. actual behavior (`has_expected_vs_actual`). + + FAIL the bug report if either (1) or (2) is missing. Do not bias + toward PASS: if the bug isn't demonstrated end-to-end, the verdict is + "fail" even when the report is well-written. + + For a FEATURE REQUEST the issue PASSES triage only when it contains + ALL of: + - A clear description of the proposed feature (what should LiteLLM do + that it does not today). + - Motivation / use case with a concrete example (config, API call, + UI flow, or scenario showing what's blocked today). + + For an issue that is neither a bug report nor a feature request (a + question, support request, or discussion), PASS as long as it has a + clear, specific ask and is not empty or template placeholder text. + + Respond with a single JSON object, no prose: + + {{ + "verdict": "pass" | "fail", + "kind": "bug" | "feature" | "other", + "has_repro": boolean, + "has_expected_vs_actual": boolean, + "has_motivation_example": boolean, + "missing": ["plain-english strings naming what is missing"], + "explanation": "1-2 sentence reasoning for the team to skim" + }} + + --- + Issue title: {title} + + Issue body: + --- + {cleaned_body} + --- + """).strip() + return template.format(title=title, cleaned_body=cleaned_body) + + +# --------------------------------------------------------------------------- +# LLM call + verdict parsing + + +def call_llm_judge( + prompt: str, *, model: str, api_key: str, base_url: str | None +) -> str: + """Call an OpenAI-compatible chat completions endpoint. Returns raw text.""" + # Import inside the function so unit tests that monkey-patch this never + # need the openai package installed. + from openai import OpenAI + + client = ( + OpenAI(api_key=api_key, base_url=base_url) + if base_url + else OpenAI(api_key=api_key) + ) + kwargs: dict[str, Any] = { + "model": model, + "messages": [{"role": "user", "content": prompt}], + "temperature": 0, + "response_format": {"type": "json_object"}, + } + # gpt-5.x reasoning models reject `temperature != 1` unless + # `reasoning_effort` is explicitly "none". Set it via `extra_body` so this + # works across openai SDK versions regardless of whether the SDK natively + # types `reasoning_effort` as a top-level chat-completions param yet. + if model.lower().startswith(GPT5_FAMILY_PREFIX): + kwargs["extra_body"] = {"reasoning_effort": "none"} + response = client.chat.completions.create(**kwargs) + return response.choices[0].message.content or "" + + +def parse_verdict(raw: str) -> dict: + """Parse the LLM's JSON response. Tolerates ```json fences and stray text.""" + if not raw: + raise ValueError("empty LLM response") + text = raw.strip() + if text.startswith("```"): + text = re.sub(r"^```(?:json)?\s*", "", text) + text = re.sub(r"\s*```$", "", text) + try: + return json.loads(text) + except json.JSONDecodeError: + match = re.search(r"\{.*\}", text, re.DOTALL) + if not match: + raise ValueError(f"could not extract JSON from LLM response: {raw[:200]}") + return json.loads(match.group(0)) + + +# --------------------------------------------------------------------------- +# Comment composition + + +def _format_missing(missing: list[str]) -> str: + if not missing: + return "- (see explanation below)" + return "\n".join(f"- {m}" for m in missing) + + +# Rubric items the judge can mark present. The first element of each tuple is +# the verdict-JSON boolean field, the second is the human-readable label we +# render in the "what you got right" section of close / grace-warning comments. +_PR_PRESENT_LABELS: tuple[tuple[str, str], ...] = ( + ("linked_issue", "Linked a related GitHub issue"), + ("has_problem_description", "Clear problem description"), + ("has_expected_vs_actual", "Expected vs. actual behavior"), + ("has_qa_proof", "End-to-end QA proof"), +) + +# Issue rubric labels grouped by `kind`. The judge sets `kind` to one of +# {"bug", "feature", "other"}; when "other" we render both groups so we don't +# silently drop a present-flag the judge actually set to True. +_ISSUE_BUG_LABELS: tuple[tuple[str, str], ...] = ( + ( + "has_repro", + "End-to-end evidence of the bug (video, screenshot, or command + real output)", + ), + ("has_expected_vs_actual", "Expected vs. actual behavior"), +) +_ISSUE_FEATURE_LABELS: tuple[tuple[str, str], ...] = ( + ("has_motivation_example", "Motivation and concrete example"), +) + + +def _format_present_for_pr(verdict: dict) -> list[str]: + """Human-readable rubric items the judge confirmed are present on a PR. + + Drives the "what you got right" section in close / grace-warning comments. + The user gave explicit feedback: contributors should see what they nailed + *before* the list of gaps, so the comment doesn't read as pure rejection. + """ + return [label for field, label in _PR_PRESENT_LABELS if verdict.get(field)] + + +def _format_present_for_issue(verdict: dict) -> list[str]: + """Human-readable rubric items the judge confirmed are present on an issue. + + Branches on the judge's `kind` field. For `"other"` (or missing kind) we + render the union so a present-flag isn't dropped just because the judge + couldn't classify the issue cleanly. + """ + kind = (verdict.get("kind") or "").lower() + groups: list[tuple[tuple[str, str], ...]] = [] + if kind in ("bug", "other", ""): + groups.append(_ISSUE_BUG_LABELS) + if kind in ("feature", "other", ""): + groups.append(_ISSUE_FEATURE_LABELS) + out: list[str] = [] + for group in groups: + for field, label in group: + if verdict.get(field) and label not in out: + out.append(label) + return out + + +def _format_present_block(items: list[str]) -> str: + """Render the optional "what you got right" block. Empty string when the + judge didn't confirm anything as present — better to omit the section + entirely than to show "What you got right: (nothing)". + """ + if not items: + return "" + bullets = "\n".join(f"- ✅ {item}" for item in items) + return f"**What you got right:**\n\n{bullets}\n\n" + + +def format_pr_close_comment(verdict: dict) -> str: + missing_lines = _format_missing(verdict.get("missing") or []) + present_block = _format_present_block(_format_present_for_pr(verdict)) + explanation = verdict.get("explanation") or "" + return ( + "🚅 Hi, thanks for the PR! I'm **Agent Shin**, the automated triage bot for this " + "repository. " + "[What's this and why am I getting it?](https://docs.litellm.ai/blog/agent-shin-triage)\n" + "\n" + "I read the description against our " + "[contribution rubric](https://github.com/BerriAI/litellm/blob/main/.github/pull_request_template.md). " + "Here's how it lined up:\n" + "\n" + f"{present_block}" + "**What's still missing:**\n" + "\n" + f"{missing_lines}\n" + "\n" + f"> {explanation}\n" + "\n" + "**Closing this PR isn't a rejection of the change.** We want the open-PR list to " + "mirror what a maintainer can act on *right now*, so contributors don't get lost in a " + 'backlog. A closed PR is a soft "park this for later"; your work is still here, ' + "the diff is still here, and getting it reopened is one comment away. Take your time.\n" + "\n" + "**To bring this PR back:**\n" + "\n" + "- Update the description with the missing pieces, then comment `@agent-shin reconsider` " + "on this PR. I'll re-evaluate and reopen if it now passes.\n" + "- Or **Open a new PR** with the same fix and the updated description. GitHub doesn't " + "always let external contributors reopen a bot-closed PR, so a fresh PR is the most " + "reliable path back into the review queue.\n" + "- If Greptile's most recent score on this PR was below 4/5, comment `@greptileai` to " + "request a fresh review; that **still works even after the PR is closed**, and a " + "stronger score is one of the signals that lifts the PR back into the queue. A low " + "Greptile score isn't a blocker.\n" + "\n" + '**What "end-to-end QA proof" means**, since it\'s the most common gap: at least one ' + "of a short before/after screen recording / video (the bug reproducing, then the fix " + "working; for a brand-new feature, a recording of it working end-to-end), a screenshot " + "(or before/after screenshots) of it working, or the exact commands you ran paired " + "with their **real output** against the real system. Running `pytest` on the repo's " + "unit tests doesn't count; those mock the LLM provider, DB, and network, so they " + "aren't end-to-end. Output from a real, no-mocks integration run is what we look " + "for. A linked issue alone isn't enough either: it covers context, not proof. See " + "[the full rubric](https://docs.litellm.ai/blog/agent-shin-triage#the-rubric-for-pull-requests).\n" + "\n" + "Internal BerriAI contributors: this rubric doesn't apply to you; ping a maintainer.\n" + "\n" + "_(I'm an LLM, so I'm not infallible. If you think I got this wrong, comment " + "`@agent-shin reconsider` or ping a maintainer; they'll override me.)_" + f"\n\n{AGENT_SHIN_CLOSE_MARKER}" + ) + + +def format_issue_close_comment(verdict: dict) -> str: + missing_lines = _format_missing(verdict.get("missing") or []) + present_block = _format_present_block(_format_present_for_issue(verdict)) + explanation = verdict.get("explanation") or "" + return ( + "🚅 Hi, thanks for filing this! I'm **Agent Shin**, the automated triage bot for this " + "repository. " + "[What's this and why am I getting it?](https://docs.litellm.ai/blog/agent-shin-triage)\n" + "\n" + "I read the issue against our reporting checklist. Here's how it lined up:\n" + "\n" + f"{present_block}" + "**What's still missing:**\n" + "\n" + f"{missing_lines}\n" + "\n" + f"> {explanation}\n" + "\n" + "**Closing this isn't us saying the bug isn't real or the request isn't useful.** We " + "want the open-issue list to mirror what a maintainer can act on *right now*, so " + "reports like yours don't get buried in a backlog. A closed issue is a soft \"park " + 'this for later"; your report is still here, and getting it reopened is one comment ' + "away. Take your time.\n" + "\n" + "**To bring this issue back:**\n" + "\n" + "1. Edit the issue description to add the missing pieces:\n" + " - For **bug reports**: end-to-end evidence of the bug (a screen recording / " + "video, a screenshot, or the exact commands you ran with their real output / " + "traceback) plus expected vs. actual behavior. Written steps with no run output, " + "video, or screenshot don't count, and mocked or stubbed runs don't count.\n" + " - For **feature requests**: a concrete description of what should change, plus a " + "use case and example (config / API call / UI flow).\n" + "2. Comment `@agent-shin reconsider`. I'll re-run triage and reopen the issue if it " + "now meets the bar. (GitHub doesn't let external authors reopen an issue a maintainer " + "or bot closed, so the comment-based reconsider is the reliable path.)\n" + "\n" + "Internal BerriAI contributors: this rubric doesn't apply to you; ping a maintainer.\n" + "\n" + "_(I'm an LLM, so I'm not infallible. If you think I got this wrong, comment " + "`@agent-shin reconsider` or ping a maintainer; they'll override me.)_" + f"\n\n{AGENT_SHIN_CLOSE_MARKER}" + ) + + +def format_grace_warning_pr_comment(verdict: dict) -> str: + """Comment posted on the FIRST low-quality detection — gives the + contributor a 2-hour grace window to fix the PR before the next + triage run actually closes it. + + This is the "before-close" warning. On the second triage run, if the + grace marker is older than `GRACE_PERIOD_SECONDS` AND the PR still + fails the rubric, the close path runs (which posts + `format_pr_close_comment` and closes the PR). + """ + missing_lines = _format_missing(verdict.get("missing") or []) + present_block = _format_present_block(_format_present_for_pr(verdict)) + explanation = verdict.get("explanation") or "" + return ( + "🚅 Hi, thanks for the PR! I'm **Agent Shin**, the automated triage bot for this " + "repository. " + "[What's this and why am I getting it?](https://docs.litellm.ai/blog/agent-shin-triage)\n" + "\n" + "I read the description against our " + "[contribution rubric](https://github.com/BerriAI/litellm/blob/main/.github/pull_request_template.md). " + "Here's how it lined up:\n" + "\n" + f"{present_block}" + "**What's still missing:**\n" + "\n" + f"{missing_lines}\n" + "\n" + f"> {explanation}\n" + "\n" + "If the description isn't updated in the next **2 hours**, I'll auto-close this PR. " + "That's **not** us saying we don't care about the change; we want the open-PR list to " + "mirror what a maintainer can act on *right now*, so contributors don't get lost in a " + 'backlog. A closed PR is a soft "park this for later," not a rejection. Take your ' + "time; everything below still works after the close.\n" + "\n" + "**During the grace period:** just update the PR description with the missing pieces. " + "No need to ping me; I'll re-check on the next sweep and skip the auto-close if it " + "now passes. See " + "[what counts as QA proof](https://docs.litellm.ai/blog/agent-shin-triage#the-rubric-for-pull-requests) " + "for the full rubric (a linked issue alone isn't enough; it covers context, not proof).\n" + "\n" + "**If the PR does get auto-closed in 2 hours, you still have easy recovery paths:**\n" + "\n" + "- Comment `@agent-shin reconsider` after updating the description. I'll re-evaluate " + "and reopen the PR if it now passes.\n" + "- Comment `@greptileai` to request a fresh Greptile review; that **still works even " + "after the PR is closed**, and a stronger score is one of the signals that lifts the " + "PR back into the queue. So a low Greptile score isn't a blocker either.\n" + "\n" + "Internal BerriAI contributors: this rubric doesn't apply to you; ping a maintainer.\n" + "\n" + "_(I'm an LLM, so I'm not infallible. If you think I got this wrong, ping a " + "maintainer; they'll override me.)_\n" + "\n" + f"{GRACE_COMMENT_MARKER}" + ) + + +def format_grace_warning_issue_comment(verdict: dict) -> str: + """Issue analogue of `format_grace_warning_pr_comment`.""" + missing_lines = _format_missing(verdict.get("missing") or []) + present_block = _format_present_block(_format_present_for_issue(verdict)) + explanation = verdict.get("explanation") or "" + return ( + "🚅 Hi, thanks for filing this! I'm **Agent Shin**, the automated triage bot for this " + "repository. " + "[What's this and why am I getting it?](https://docs.litellm.ai/blog/agent-shin-triage)\n" + "\n" + "I read the issue against our reporting checklist. Here's how it lined up:\n" + "\n" + f"{present_block}" + "**What's still missing:**\n" + "\n" + f"{missing_lines}\n" + "\n" + f"> {explanation}\n" + "\n" + "If the issue isn't updated in the next **2 hours**, I'll auto-close it. That's **not** us " + "saying the bug isn't real or the request isn't useful; we want the open-issue list " + "to mirror what a maintainer can act on *right now*, so reports like yours don't get " + 'buried in a backlog. A closed issue is a soft "park this for later," not a ' + "rejection. Take your time; reopening is one comment away.\n" + "\n" + "**During the grace period:** just edit the issue description with the missing " + "pieces. No need to ping me; I'll re-check on the next sweep and skip the auto-close " + "if it now passes.\n" + "\n" + "Missing pieces, depending on what this is:\n" + "\n" + "- For **bug reports**: end-to-end evidence of the bug (a screen recording / video, a " + "screenshot, or the exact commands you ran with their real output / traceback) plus " + "expected vs. actual behavior. Written steps with no run output don't count, and " + "mocked or stubbed runs don't count.\n" + "- For **feature requests**: a concrete description of what should change, plus a use " + "case and example (config / API call / UI flow).\n" + "\n" + "**If the issue does get auto-closed in 2 hours**, comment `@agent-shin reconsider` " + "and I'll re-evaluate. If it now meets the bar, I'll reopen the issue.\n" + "\n" + "Internal BerriAI contributors: this rubric doesn't apply to you; ping a maintainer.\n" + "\n" + "_(I'm an LLM, so I'm not infallible. If you think I got this wrong, ping a " + "maintainer; they'll override me.)_\n" + "\n" + f"{GRACE_COMMENT_MARKER}" + ) + + +# --------------------------------------------------------------------------- +# Step-summary helpers + + +def write_step_summary(content: str) -> None: + """When running inside GitHub Actions, append to the step summary file.""" + path = os.environ.get("GITHUB_STEP_SUMMARY") + if not path: + return + try: + with open(path, "a", encoding="utf-8") as handle: + handle.write(content) + if not content.endswith("\n"): + handle.write("\n") + except OSError as exc: + print(f"warn: failed to write step summary: {exc}", file=sys.stderr) + + +# --------------------------------------------------------------------------- +# Core orchestration + + +def format_reopen_comment(kind: str) -> str: + """Comment posted when Agent Shin reopens after a successful reconsider.""" + noun = "PR" if kind == "pr" else "issue" + # The trailing HTML marker is used by `seconds_since_last_reconsider_verdict` + # to enforce a cooldown between repeated `@agent-shin reconsider` triggers. + # Keep the marker on its own line so it doesn't disturb the rendered text. + return ( + f"♻️ **Re-evaluated and reopened.** Thanks for updating the {noun}!\n" + "\n" + "Agent Shin re-ran triage on the latest description and it now meets " + "the bar. A maintainer will take another look soon; please don't " + f"close this {noun} again unless asked to.\n" + "\n" + "_(If a maintainer ends up closing this for non-rubric reasons, that " + "decision stands; comment `@agent-shin reconsider` again only if you " + "have substantively new information.)_\n" + "\n" + f"{RECONSIDER_COMMENT_MARKER}" + ) + + +def format_reconsider_still_failing_comment(kind: str, verdict: dict) -> str: + """Comment posted when reconsider re-runs triage but the verdict is still fail.""" + missing_lines = _format_missing(verdict.get("missing") or []) + explanation = verdict.get("explanation") or "" + noun = "PR" if kind == "pr" else "issue" + # The trailing HTML marker is used by `seconds_since_last_reconsider_verdict` + # to enforce a cooldown between repeated `@agent-shin reconsider` triggers. + return ( + f"⏸️ **Re-evaluated; this {noun} still doesn't meet the rubric.**\n" + "\n" + "Agent Shin re-ran triage on the current description but is still " + "missing:\n" + "\n" + f"{missing_lines}\n" + "\n" + f"> {explanation}\n" + "\n" + "Update the description with the missing pieces and comment " + "`@agent-shin reconsider` again, or ping a maintainer if you think " + "I got this wrong.\n" + "\n" + "_(I'm an LLM and I'm not infallible.)_\n" + "\n" + f"{RECONSIDER_COMMENT_MARKER}" + ) + + +# --------------------------------------------------------------------------- +# Review gate — "ready for review" label lifecycle + +_UNSET = object() + + +def _combine_missing( + verdict: dict, greptile_score: int | None, min_score: int +) -> list[str]: + """Merge the LLM rubric's `missing` list with a Greptile-score shortfall.""" + missing = list(verdict.get("missing") or []) + if greptile_score is not None and greptile_score < min_score: + missing.insert( + 0, + f"Greptile's most recent review scored this PR {greptile_score}/5 " + f"(below the {min_score}/5 bar)", + ) + return missing or ["(see explanation below)"] + + +def _has_marker( + comments: Iterable[dict], marker: str, *, bot_login: str | None = None +) -> bool: + """Return True iff the bot itself posted a comment containing ``marker``. + + Filters by author so a contributor who quotes the marker (e.g. via + GitHub's "Quote reply" feature, which preserves HTML comments in + raw markdown) is not mistaken for a bot action — that would + silently suppress notifications or change which "recovered" wording + is selected. Matches the author-filter pattern used by the sibling + `_seconds_since_latest_marker_comment` helper. + """ + expected_login = ( + bot_login + or os.environ.get("AGENT_SHIN_BOT_LOGIN") + or AGENT_SHIN_DEFAULT_BOT_LOGIN + ).lower() + for comment in comments: + author = ((comment.get("user") or {}).get("login") or "").lower() + if author != expected_login: + continue + if marker in (comment.get("body") or ""): + return True + return False + + +def format_ready_for_review_comment( + verdict: dict, + greptile_score: int | None, + min_greptile_score: int = DEFAULT_MIN_GREPTILE_SCORE, +) -> str: + """Posted the first time a PR clears the bar (label added).""" + score_line = ( + f" Greptile scored it **{greptile_score}/5**." + if greptile_score is not None + else "" + ) + explanation = verdict.get("explanation") or "" + return ( + "✅ **Triage passed, tagging `ready for review`.**\n" + "\n" + "Agent Shin checked this PR against the " + "[contribution rubric](https://github.com/BerriAI/litellm/blob/main/.github/pull_request_template.md) " + "and it clears the bar (a linked issue, or a clear problem description " + f"+ expected vs. actual + QA proof).{score_line}\n" + "\n" + f"> {explanation}\n" + "\n" + "A maintainer will take it from here. If a later re-check finds the PR " + f"has regressed (Greptile drops below {min_greptile_score}/5, " + "the QA proof is removed, etc.) I'll pull the tag and comment with " + "what's missing; fix it and the tag comes back automatically.\n" + f"{READY_MARKER}" + ) + + +def format_all_clear_comment(verdict: dict, greptile_score: int | None) -> str: + """Posted when a PR recovers after a regression (label re-added).""" + score_line = ( + f" Greptile is back to **{greptile_score}/5**." + if greptile_score is not None + else "" + ) + explanation = verdict.get("explanation") or "" + return ( + "✅ **All clear again, re-adding `ready for review`.**\n" + "\n" + "Thanks for addressing the earlier feedback. On re-check this PR meets " + f"the contribution bar once more.{score_line}\n" + "\n" + f"> {explanation}\n" + "\n" + "A maintainer will take another look.\n" + f"{READY_MARKER}" + ) + + +def format_regression_comment( + missing: list[str], explanation: str, grace_days: int +) -> str: + """Posted when a previously-tagged PR regresses (label removed, PR stays open). + + Discloses the same ``grace_days`` deadline the state machine enforces: + once that window elapses with the PR still failing, the close path fires. + Hiding the deadline behind a bare "stays open" would surprise contributors + with an auto-close they were never warned about. + """ + window = "24 hours" if grace_days == 1 else f"{grace_days} days" + return ( + "⚠️ **Removing the `ready for review` tag.**\n" + "\n" + "On a re-check this PR no longer meets the contribution bar. What's " + "missing now:\n" + "\n" + f"{_format_missing(missing)}\n" + "\n" + f"> {explanation}\n" + "\n" + f"The PR stays open for ~{window}; address the points above and Agent " + 'Shin will post an "all clear" comment and re-add the tag ' + "automatically. If the points still aren't addressed after that " + "window, the PR is auto-closed; that's not a rejection, and you can " + "comment `@agent-shin reconsider` to have it re-evaluated and reopened " + "once it passes.\n" + f"{REGRESSED_MARKER}" + ) + + +def format_within_grace_comment( + missing: list[str], explanation: str, grace_days: int +) -> str: + """Posted once while a failing PR is still inside its grace window.""" + window = "24 hours" if grace_days == 1 else f"{grace_days} days" + return ( + "🚅 Hi, thanks for the PR! This is **Agent Shin**, the automated triage " + "bot. This PR doesn't quite meet the contribution bar yet:\n" + "\n" + f"{_format_missing(missing)}\n" + "\n" + f"> {explanation}\n" + "\n" + f"You have ~{window} from when this PR was opened to add the missing " + "pieces; just update the description and I'll re-check on the next " + "sweep. Once it passes I'll tag it `ready for review`. If it does get " + "auto-closed, that's not a rejection; comment `@agent-shin reconsider` " + "and I'll re-evaluate and reopen if it now passes.\n" + f"{WITHIN_GRACE_MARKER}" + ) + + +def review_gate( + *, + repo: str, + number: int, + close: bool, + model: str, + judge: Any = None, + greptile_score: Any = _UNSET, + comments: Any = _UNSET, + now: dt.datetime | None = None, + grace_days: int = DEFAULT_GRACE_DAYS, + min_greptile_score: int = DEFAULT_MIN_GREPTILE_SCORE, + label: str = READY_FOR_REVIEW_LABEL, + allowlist: frozenset[str] = ALLOWLIST_LOGINS, +) -> dict: + """Reconcile the `ready for review` label with a PR's current quality. + + A PR is *passing* when it clears BOTH gates: the LLM rubric (linked issue, + or problem description + expected/actual + QA proof) AND Greptile's most + recent confidence score (>= ``min_greptile_score``; absence of a score is + not held against the PR). The gate then drives a small state machine, using + the label itself as the persisted state so comments fire only on + transitions (never on every scheduled run): + + passing, untagged -> add label + "ready for review" / "all clear" + passing, tagged -> noop-passing + not passing, tagged -> remove label + regression comment (stays open) + not passing, untagged, old -> close + comment (past the grace window) + not passing, untagged, new -> one-time "what's missing" notice (within grace) + + ``close`` gates every destructive side effect: with ``close=False`` the + function returns a ``would-*`` preview and touches nothing, mirroring the + dry-run contract of :func:`triage`. ``judge``/``greptile_score``/ + ``comments``/``now`` are injectable for tests; in production they are + resolved from the OpenAI judge, the PR's Greptile comment, the live comment + list, and the wall clock respectively. + """ + item = fetch_pr(repo, number) + + title = item.get("title") or "" + body = item.get("body") or "" + login = (item.get("user") or {}).get("login") or "" + association = item.get("author_association") or "" + state = item.get("state") or "" + # GitHub label names are case-insensitive; compare lowercased so a repo + # that already has e.g. "Ready for Review" is recognized as the same + # label as our READY_FOR_REVIEW_LABEL constant ("ready for review"). + labels_now = {(lbl.get("name") or "").lower() for lbl in (item.get("labels") or [])} + label_key = label.lower() + created_raw = item.get("created_at") or "" + + base_result = { + "kind": "pr", + "number": number, + "title": title, + "author": login, + "author_association": association, + "state": state, + "labeled": label_key in labels_now, + "review_gate": True, + } + + if state != "open": + return {**base_result, "action": "skip-not-open"} + + if allowlist: + if login.lower() not in allowlist: + return {**base_result, "action": "skip-not-allowlisted"} + elif is_internal_contributor(item): + return {**base_result, "action": "skip-internal-author"} + + # Resolve the comment list once — used for both the Greptile score and the + # marker-based dedup below. + if comments is _UNSET: + comments = list(_iter_paginated_json(f"repos/{repo}/issues/{number}/comments")) + + # --- rubric verdict: linked-issue short-circuit, else the LLM judge ------- + if has_linked_issue(body): + verdict = { + "verdict": "pass", + "linked_issue": True, + "missing": [], + "explanation": "Linked-issue regex matched; LLM was not called.", + } + rubric_pass = True + else: + prompt = build_pr_prompt(title=title, body=body) + if judge is None: + api_key = os.environ.get("OPENAI_API_KEY") + if not api_key: + return {**base_result, "action": "skip-no-llm-key"} + base_url = os.environ.get("OPENAI_BASE_URL") or None + + def judge(p: str) -> str: + return call_llm_judge( + p, model=model, api_key=api_key, base_url=base_url + ) + + try: + verdict = parse_verdict(judge(prompt)) + except Exception as exc: # noqa: BLE001 - judge errors must never act + return {**base_result, "action": "skip-llm-error", "error": str(exc)} + rubric_pass = (verdict.get("verdict") or "").lower() == "pass" + + # --- Greptile score ------------------------------------------------------- + if greptile_score is _UNSET: + extraction = extract_greptile_score(comments) + greptile_score = extraction[0] if extraction else None + greptile_ok = greptile_score is None or greptile_score >= min_greptile_score + passing = rubric_pass and greptile_ok + + # --- age ------------------------------------------------------------------ + age_days = None + if created_raw: + reference = now or dt.datetime.now(dt.timezone.utc) + age_days = (reference - parse_iso8601(created_raw)).days + + label_present = label_key in labels_now + explanation = verdict.get("explanation") or "" + # When the rubric short-circuited to pass (linked-issue regex) but + # Greptile dragged the PR below the bar, the synthetic verdict's + # explanation ("LLM was not called") would mislead a contributor reading + # the regression / close comment. Surface the real reason instead. + if rubric_pass and not greptile_ok: + explanation = ( + f"Greptile's most recent review scored this PR " + f"{greptile_score}/5 (below the {min_greptile_score}/5 bar)." + ) + verdict = {**verdict, "explanation": explanation} + base_result = { + **base_result, + "verdict": verdict, + "greptile_score": greptile_score, + "passing": passing, + "age_days": age_days, + } + + if passing: + if label_present: + return {**base_result, "action": "noop-passing"} + recovered = _has_marker(comments, REGRESSED_MARKER) + comment = ( + format_all_clear_comment(verdict, greptile_score) + if recovered + else format_ready_for_review_comment( + verdict, greptile_score, min_greptile_score + ) + ) + if not close: + return {**base_result, "action": "would-label-ready", "comment": comment} + post_comment(repo, number, comment) + add_label(repo, number, label) + return {**base_result, "action": "labeled-ready", "comment": comment} + + missing = _combine_missing(verdict, greptile_score, min_greptile_score) + + if label_present: + comment = format_regression_comment(missing, explanation, grace_days) + if not close: + return {**base_result, "action": "would-remove-label", "comment": comment} + remove_label(repo, number, label) + post_comment(repo, number, comment) + return {**base_result, "action": "label-removed-regressed", "comment": comment} + + # Not passing and not tagged. If the PR was previously tagged and then + # regressed (we removed the label and posted REGRESSED_MARKER), honor the + # "PR stays open — fix it and the tag comes back" promise from + # `format_regression_comment` and skip the close path. Without this guard, + # any PR older than `grace_days` would be closed on the next evaluation, + # giving the contributor no realistic window to address the regression. + # + # The promise has a deliberate expiration: once `grace_days` have elapsed + # since the regression notice, fall through to the close path so a PR that + # was abandoned post-regression doesn't sit open forever. + if _has_marker(comments, REGRESSED_MARKER): + reference = now or dt.datetime.now(dt.timezone.utc) + seconds_since_regression = seconds_since_latest_marker_comment( + comments, marker=REGRESSED_MARKER, now=reference + ) + grace_seconds = grace_days * 86400 + if seconds_since_regression is None or seconds_since_regression < grace_seconds: + return {**base_result, "action": "regressed-already-notified"} + + # Not passing and not tagged: close if past the grace window, else notify once. + if age_days is not None and age_days >= grace_days: + comment = format_pr_close_comment({**verdict, "missing": missing}) + if not close: + return {**base_result, "action": "would-close", "comment": comment} + post_comment(repo, number, comment) + close_pr(repo, number) + return {**base_result, "action": "closed", "comment": comment} + + if _has_marker(comments, WITHIN_GRACE_MARKER): + return {**base_result, "action": "within-grace-already-notified"} + comment = format_within_grace_comment(missing, explanation, grace_days) + if not close: + return { + **base_result, + "action": "would-notify-within-grace", + "comment": comment, + } + post_comment(repo, number, comment) + return {**base_result, "action": "within-grace-notified", "comment": comment} + + +def triage( + *, + repo: str, + kind: str, + number: int, + close: bool, + model: str, + judge: Any = None, + print_prompt: bool = False, + reconsider: bool = False, + allowlist: frozenset[str] = ALLOWLIST_LOGINS, +) -> dict: + """Triage a single PR or issue. Returns a result dict for logging/tests. + + `judge` is an optional callable `(prompt) -> str` for tests / dry-run with + a stub. In production, leave it None and the script uses `call_llm_judge`. + + When `reconsider=True`, the closed-state guard is skipped and a + fail-but-no-comment is replaced with a "still failing" comment + leave + closed; a pass triggers `reopen_pr`/`reopen_issue` plus a reopen comment. + Reconsider mode is intended for the `@agent-shin reconsider` comment + trigger. Like regular triage, `close=False` keeps reconsider in dry-run + (returns `would-reopen` / `would-reconsider-still-failing` so a local + operator can preview without write side effects); the workflow only + passes `--close` when `AGENT_SHIN_ENABLED=true`. + + Reconsider mode adds two extra safety guards on top of the regular + triage skip-internal-author check: + + 1. **Bot-closed guard.** Only reopens if the most recent close was + performed by the bot identity (default `github-actions[bot]`). + This stops a contributor from using `@agent-shin reconsider` to + override a maintainer's close for non-rubric reasons. + 2. **Rate-limit guard.** If the bot has already posted a reconsider + verdict on this PR/issue within `RECONSIDER_RATE_LIMIT_SECONDS`, + skip — repeated triggers from the same contributor shouldn't burn + CI minutes or LLM budget. + """ + fetcher = {"pr": fetch_pr, "issue": fetch_issue}[kind] + item = fetcher(repo, number) + + title = item.get("title") or "" + body = item.get("body") or "" + login = (item.get("user") or {}).get("login") or "" + association = item.get("author_association") or "" + state = item.get("state") or "" + + base_result = { + "kind": kind, + "number": number, + "title": title, + "author": login, + "author_association": association, + "state": state, + "reconsider": reconsider, + } + + # Reconsider only makes sense on a closed PR/issue. A "reconsider on an + # open PR" is a no-op (the regular triage flow already evaluates open + # PRs); return a clear skip so the workflow can short-circuit. + if reconsider: + if state != "closed": + return {**base_result, "action": "skip-not-closed"} + else: + if state != "open": + return {**base_result, "action": "skip-not-open"} + + if allowlist: + if login.lower() not in allowlist: + return {**base_result, "action": "skip-not-allowlisted"} + elif is_internal_contributor(item): + return {**base_result, "action": "skip-internal-author"} + + # Reconsider-only guards — these run BEFORE the LLM call so a + # maintainer-closed PR / rate-limited trigger never spends LLM budget. + if reconsider: + if not was_closed_by_agent_shin(repo, number): + return {**base_result, "action": "skip-not-bot-closed"} + age = seconds_since_last_reconsider_verdict(repo, number) + if age is not None and age < RECONSIDER_RATE_LIMIT_SECONDS: + return { + **base_result, + "action": "skip-rate-limited", + "rate_limit_age_seconds": age, + "rate_limit_window_seconds": RECONSIDER_RATE_LIMIT_SECONDS, + } + + if kind == "pr": + # Short-circuit: if body very clearly links a related issue, just pass. + if has_linked_issue(body): + base = { + **base_result, + "action": "pass-linked-issue", + "verdict": { + "verdict": "pass", + "linked_issue": True, + "explanation": "Linked-issue regex matched; LLM was not called.", + }, + } + if reconsider: + # Pass-on-reconsider -> reopen the PR with a friendly comment. + reopen_body = format_reopen_comment(kind) + if not close: + return { + **base, + "action": "would-reopen", + "comment": reopen_body, + } + post_comment(repo, number, reopen_body) + reopen_pr(repo, number) + return { + **base, + "action": "reopened", + "comment": reopen_body, + } + return base + prompt = build_pr_prompt(title=title, body=body) + else: + prompt = build_issue_prompt(title=title, body=body) + + if print_prompt: + return {**base_result, "action": "print-prompt", "prompt": prompt} + + if judge is None: + api_key = os.environ.get("OPENAI_API_KEY") + if not api_key: + # No key configured — never take a destructive action. Report skip. + return { + **base_result, + "action": "skip-no-llm-key", + "prompt_preview": prompt[:200], + } + base_url = os.environ.get("OPENAI_BASE_URL") or None + + def judge(p: str) -> str: + return call_llm_judge(p, model=model, api_key=api_key, base_url=base_url) + + try: + raw = judge(prompt) + verdict = parse_verdict(raw) + except Exception as exc: # noqa: BLE001 - judge errors must never close PRs + return {**base_result, "action": "skip-llm-error", "error": str(exc)} + + decision = (verdict.get("verdict") or "").lower() + + if reconsider: + # Reconsider: an explicit `pass` -> reopen + post reopen comment; + # anything else (fail, missing/malformed verdict, typo) -> leave + # closed + post a "still failing" comment so the contributor can + # iterate again. Reopen is destructive, so a flaky/empty verdict + # must not satisfy the gate. + # In dry-run (`close=False`) we return `would-*` actions instead + # of touching GitHub state, mirroring the regular triage flow's + # `would-close`. This lets a local operator preview the outcome + # of `python triage_with_llm.py --reconsider --pr N` without + # risking accidental comments or reopens. + if decision == "pass": + reopen_body = format_reopen_comment(kind) + if not close: + return { + **base_result, + "action": "would-reopen", + "verdict": verdict, + "comment": reopen_body, + } + post_comment(repo, number, reopen_body) + if kind == "pr": + reopen_pr(repo, number) + else: + reopen_issue(repo, number) + return { + **base_result, + "action": "reopened", + "verdict": verdict, + "comment": reopen_body, + } + still_failing = format_reconsider_still_failing_comment(kind, verdict) + if not close: + return { + **base_result, + "action": "would-reconsider-still-failing", + "verdict": verdict, + "comment": still_failing, + } + post_comment(repo, number, still_failing) + return { + **base_result, + "action": "reconsider-still-failing", + "verdict": verdict, + "comment": still_failing, + } + + if decision != "fail": + return {**base_result, "action": "pass-llm", "verdict": verdict} + + # Grace-period flow: on the first low-quality detection, post a warning + # comment instead of closing immediately. On a subsequent triage run + # (manual re-trigger, or the daily `close_low_quality_prs.py` cron + # finding the same PR in its own pass), if `GRACE_PERIOD_SECONDS` has + # elapsed since the warning AND the PR still fails the rubric, close. + grace_age = seconds_since_last_grace_warning(repo, number) + if grace_age is None: + warning_body = ( + format_grace_warning_pr_comment(verdict) + if kind == "pr" + else format_grace_warning_issue_comment(verdict) + ) + if not close: + return { + **base_result, + "action": "would-warn-grace", + "verdict": verdict, + "comment": warning_body, + } + post_comment(repo, number, warning_body) + return { + **base_result, + "action": "warned-grace", + "verdict": verdict, + "comment": warning_body, + } + if grace_age < GRACE_PERIOD_SECONDS: + return { + **base_result, + "action": "skip-in-grace-period", + "verdict": verdict, + "grace_age_seconds": grace_age, + "grace_period_seconds": GRACE_PERIOD_SECONDS, + } + + # The grace window has elapsed. `--close` still gates the destructive + # write so a dry-run preview never posts or closes — the workflow only + # passes `--close` when `AGENT_SHIN_ENABLED=true`, which keeps the bot + # inert by default. + if not close: + return {**base_result, "action": "would-close", "verdict": verdict} + + comment_body = ( + format_pr_close_comment(verdict) + if kind == "pr" + else format_issue_close_comment(verdict) + ) + post_comment(repo, number, comment_body) + if kind == "pr": + close_pr(repo, number) + else: + close_issue(repo, number) + + return { + **base_result, + "action": "closed", + "verdict": verdict, + "comment": comment_body, + } + + +# --------------------------------------------------------------------------- +# CLI + + +def render_summary(result: dict) -> str: + """Render a human-readable summary block (used for stdout + step summary).""" + lines = ["## Agent Shin verdict", ""] + lines.append( + f"- **{result['kind'].upper()} #{result['number']}**: {result.get('title', '')}" + ) + lines.append( + f"- **Author**: `{result.get('author', '')}` ({result.get('author_association', '')})" + ) + lines.append(f"- **State**: {result.get('state', '')}") + lines.append(f"- **Action**: `{result['action']}`") + verdict = result.get("verdict") + if verdict: + lines.append("") + lines.append("```json") + lines.append(json.dumps(verdict, indent=2)) + lines.append("```") + error = result.get("error") + if error: + lines.append("") + lines.append(f"_LLM error: {error}_") + comment = result.get("comment") + if comment: + lines.append("") + lines.append("### Posted comment:") + lines.append("") + lines.append("> " + comment.replace("\n", "\n> ")) + return "\n".join(lines) + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--repo", required=True, help="Repository (owner/repo).") + target = parser.add_mutually_exclusive_group(required=True) + target.add_argument("--pr", type=int, help="Pull request number to triage.") + target.add_argument("--issue", type=int, help="Issue number to triage.") + parser.add_argument( + "--close", + action="store_true", + help="Actually post comment + close on fail (default: dry run).", + ) + parser.add_argument( + "--model", + # `os.environ.get("TRIAGE_MODEL", DEFAULT_MODEL)` would return "" when + # GitHub Actions exposes an unset repo variable as an empty-string env + # var, silently bypassing DEFAULT_MODEL and causing every call to fail + # as `skip-llm-error`. The `or` guard collapses empty -> default. + default=os.environ.get("TRIAGE_MODEL") or DEFAULT_MODEL, + help=f"OpenAI-compatible model name (default: {DEFAULT_MODEL}).", + ) + parser.add_argument( + "--print-prompt", + action="store_true", + help="Print the prompt that would be sent to the judge and exit.", + ) + parser.add_argument( + "--reconsider", + action="store_true", + help=( + "Re-run triage on a CLOSED PR/issue and reopen it on pass. " + "Used by the `@agent-shin reconsider` comment-trigger workflow. " + "Only invoke this from a workflow that has already gated on " + "AGENT_SHIN_ENABLED=true and verified the commenter is the " + "PR/issue author or an internal collaborator." + ), + ) + parser.add_argument( + "--review-gate", + action="store_true", + help=( + "Reconcile the `ready for review` label for an OPEN PR: tag on " + "pass, remove the tag + comment on regression, close after the " + "grace window if it never passed. PR-only." + ), + ) + parser.add_argument( + "--grace-days", + type=int, + default=DEFAULT_GRACE_DAYS, + help=( + "Review-gate only: hours/24 a failing, un-tagged PR may stay open " + f"before auto-close (default: {DEFAULT_GRACE_DAYS} = 24h)." + ), + ) + parser.add_argument( + "--min-greptile-score", + type=int, + default=DEFAULT_MIN_GREPTILE_SCORE, + choices=range(1, 6), + help=( + "Review-gate only: Greptile score below which a PR counts as not " + f"passing (default: {DEFAULT_MIN_GREPTILE_SCORE} -> <4/5 regresses)." + ), + ) + args = parser.parse_args() + + kind = "pr" if args.pr is not None else "issue" + number = args.pr if args.pr is not None else args.issue + + if args.review_gate: + if kind != "pr": + parser.error("--review-gate applies to pull requests only (use --pr).") + result = review_gate( + repo=args.repo, + number=number, + close=args.close, + model=args.model, + grace_days=args.grace_days, + min_greptile_score=args.min_greptile_score, + ) + else: + result = triage( + repo=args.repo, + kind=kind, + number=number, + close=args.close, + model=args.model, + print_prompt=args.print_prompt, + reconsider=args.reconsider, + ) + + if result.get("action") == "print-prompt": + print(result["prompt"]) + return 0 + + summary = render_summary(result) + print(summary) + write_step_summary(summary + "\n") + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/.github/workflows/close_low_quality_prs.yml b/.github/workflows/close_low_quality_prs.yml new file mode 100644 index 00000000000..2401be84000 --- /dev/null +++ b/.github/workflows/close_low_quality_prs.yml @@ -0,0 +1,92 @@ +name: Close Low-Quality PRs + +# Auto-close any open PR (including drafts, regardless of age) authored by an +# external OSS contributor that Greptile reviewed with a confidence score +# below 4/5. Closures are explained in a comment that tells the contributor +# to push fixes and open a fresh PR (since OSS authors cannot reopen a PR +# closed by a bot/maintainer) or comment `@agent-shin reconsider` to have +# Agent Shin re-evaluate. +# +# Manual one-off run: +# gh workflow run "Close Low-Quality PRs" -f close=true +# +# Dry-run preview (no PRs are touched): +# gh workflow run "Close Low-Quality PRs" -f close=false + +on: + schedule: + # Daily at 09:00 UTC. Pairs well with the stale-issue workflow at midnight. + - cron: "0 9 * * *" + workflow_dispatch: + inputs: + close: + description: "Actually close matching PRs (false = dry run)." + required: false + default: "false" + type: choice + options: + - "true" + - "false" + min_age_days: + description: "Minimum PR age in days (default 0 = no age filter)." + required: false + default: "0" + min_score: + description: "Greptile score below which a PR is closed (1-5)." + required: false + default: "4" + limit: + description: "Maximum number of PRs to close in a single run." + required: false + default: "25" + +permissions: + contents: read + pull-requests: write + issues: write + +jobs: + close-low-quality-prs: + if: github.repository == 'BerriAI/litellm' + runs-on: ubuntu-latest + steps: + - name: Checkout triage script + uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + with: + sparse-checkout: .github/scripts + persist-credentials: false + + - name: Set up Python + uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 + with: + python-version: "3.12" + + - name: Run low-quality PR closer + env: + GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} + # Scheduled runs are ALWAYS dry-run, even when AGENT_SHIN_ENABLED is + # "true", so the team can QA the closer's verdicts in step summaries + # before any contributor sees a PR closed. Real closures only happen + # on manual workflow_dispatch with close=true (and the variable set). + CLOSE_FLAG: ${{ github.event.inputs.close || 'false' }} + AGENT_SHIN_ENABLED: ${{ vars.AGENT_SHIN_ENABLED }} + MIN_AGE_DAYS: ${{ github.event.inputs.min_age_days || '0' }} + MIN_SCORE: ${{ github.event.inputs.min_score || '4' }} + LIMIT: ${{ github.event.inputs.limit || '25' }} + run: | + set -euo pipefail + ARGS=( + --repo "${{ github.repository }}" + --min-age-days "${MIN_AGE_DAYS}" + --min-score "${MIN_SCORE}" + --limit "${LIMIT}" + ) + if [ "${AGENT_SHIN_ENABLED:-false}" != "true" ]; then + echo "::notice::AGENT_SHIN_ENABLED is not 'true' -> forcing dry-run regardless of close input." + elif [ "${GITHUB_EVENT_NAME:-}" = "workflow_dispatch" ] && [ "${CLOSE_FLAG}" = "true" ]; then + ARGS+=(--close) + echo "::notice::Running in close-on-fail mode." + else + echo "::notice::AGENT_SHIN_ENABLED is true but this trigger is dry-run (scheduled event or close=false)." + fi + python3 .github/scripts/close_low_quality_prs.py "${ARGS[@]}" diff --git a/.github/workflows/review_gate.yml b/.github/workflows/review_gate.yml new file mode 100644 index 00000000000..ba4b488b79d --- /dev/null +++ b/.github/workflows/review_gate.yml @@ -0,0 +1,131 @@ +name: Agent Shin — review gate + +# Keeps the `ready for review` label in sync with whether an external PR +# currently clears BOTH the LLM rubric AND Greptile's confidence score. +# +# pass -> add `ready for review` + a "passed / all clear" comment +# regress -> remove the label + a "what's missing" comment (PR stays open) +# fail, <24h old -> a one-time "what's missing" notice (grace window) +# fail, >24h old -> close + a comment (reopen via `@agent-shin reconsider`) +# +# DRY-RUN BY DEFAULT. Every side effect (label add/remove, comment, close) is +# gated behind `--close`, which is only added when the repo variable +# `AGENT_SHIN_ENABLED == "true"`. Until then runs only write the verdict to the +# workflow step summary. +# +# Manual single PR: gh workflow run "Agent Shin — review gate" -f pr_number=NNN +# Manual dry-run: gh workflow run "Agent Shin — review gate" -f close=false +# +# We use `pull_request_target` so the workflow can read repo secrets and run +# against fork PRs. Fork code is never checked out — only PR metadata is read +# via `gh api`. + +on: + pull_request_target: + types: [opened, reopened, synchronize, ready_for_review] + schedule: + # Daily at 09:30 UTC — re-reconciles labels as Greptile re-reviews land. + - cron: "30 9 * * *" + workflow_dispatch: + inputs: + pr_number: + description: "Single PR to reconcile (omit to sweep all open PRs)." + required: false + close: + description: "If AGENT_SHIN_ENABLED=true, actually act (false = dry run)." + required: false + default: "false" + type: choice + options: + - "true" + - "false" + grace_days: + description: "Hours/24 a failing, un-tagged PR may stay open before close." + required: false + default: "1" + min_greptile_score: + description: "Greptile score below which a PR counts as not passing (1-5)." + required: false + default: "4" + +permissions: + contents: read + issues: write + pull-requests: write + +jobs: + review-gate: + if: github.repository == 'BerriAI/litellm' + runs-on: ubuntu-latest + steps: + - name: Checkout triage script + uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + with: + sparse-checkout: .github/scripts + persist-credentials: false + + - name: Set up Python + uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 + with: + python-version: "3.12" + + - name: Install LLM client + run: pip install --no-cache-dir --require-hashes -r .github/scripts/triage-requirements.txt + + - name: Run review gate + env: + GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} + # Mirror the triage workflow: only expose the LLM key when the bot is + # enabled or a collaborator triggers it manually, so an external user + # can't force paid LLM calls by churning a fork PR while the bot is + # still in dry-run. + OPENAI_API_KEY: ${{ (vars.AGENT_SHIN_ENABLED == 'true' || github.event_name == 'workflow_dispatch') && secrets.OPENAI_API_KEY || '' }} + OPENAI_BASE_URL: ${{ vars.OPENAI_BASE_URL }} + TRIAGE_MODEL: ${{ vars.TRIAGE_MODEL }} + AGENT_SHIN_ENABLED: ${{ vars.AGENT_SHIN_ENABLED }} + CLOSE_FLAG: ${{ github.event.inputs.close || 'false' }} + GRACE_DAYS: ${{ github.event.inputs.grace_days || '1' }} + MIN_GREPTILE_SCORE: ${{ github.event.inputs.min_greptile_score || '4' }} + EVENT_PR: ${{ github.event.pull_request.number }} + INPUT_PR: ${{ github.event.inputs.pr_number }} + run: | + set -euo pipefail + COMMON=(--review-gate --grace-days "${GRACE_DAYS}" --min-greptile-score "${MIN_GREPTILE_SCORE}") + + # Fail-safe gating, identical philosophy to the Greptile closer: + # - AGENT_SHIN_ENABLED must be the EXACT string "true" to act at all. + # - A manual dispatch can still preview with close=false. + # - Automatic triggers (PR events, schedule) act once enabled — that + # is the whole point of the gate (re-tag / un-tag automatically). + DO_CLOSE="false" + if [ "${AGENT_SHIN_ENABLED:-false}" != "true" ]; then + echo "::notice::AGENT_SHIN_ENABLED is not 'true' -> dry-run (no labels/comments/closes)." + elif [ "${GITHUB_EVENT_NAME:-}" = "workflow_dispatch" ] && [ "${CLOSE_FLAG:-false}" = "true" ]; then + DO_CLOSE="true" + echo "::notice::Manual run -> acting for real." + elif [ "${GITHUB_EVENT_NAME:-}" != "workflow_dispatch" ]; then + DO_CLOSE="true" + echo "::notice::Enabled automatic trigger (${GITHUB_EVENT_NAME:-}) -> acting for real." + else + echo "::notice::Manual dispatch with close=false -> dry-run." + fi + if [ "${DO_CLOSE}" = "true" ]; then + COMMON+=(--close) + fi + + # Single PR (PR event or explicit input) vs. sweep over all open PRs. + TARGET_PR="${EVENT_PR:-${INPUT_PR:-}}" + if [ -n "${TARGET_PR}" ]; then + python3 .github/scripts/triage_with_llm.py --repo "${{ github.repository }}" --pr "${TARGET_PR}" "${COMMON[@]}" + else + echo "::notice::Sweeping all open PRs." + # Match GH_LIST_ALL_LIMIT in agent_shin_shared.py: gh lists newest-first, + # so any cap below the real backlog silently drops the *oldest* PRs — + # exactly the stale ones this daily sweep is meant to reconcile. + mapfile -t NUMBERS < <(gh pr list --repo "${{ github.repository }}" --state open --limit 100000 --json number --jq '.[].number') + for n in "${NUMBERS[@]}"; do + echo "::group::PR #${n}" + python3 .github/scripts/triage_with_llm.py --repo "${{ github.repository }}" --pr "${n}" "${COMMON[@]}" || echo "::warning::review gate errored on #${n}" + echo "::endgroup::" + done + fi diff --git a/.github/workflows/test-linting.yml b/.github/workflows/test-linting.yml index 2e967f3ed3f..de7e1b68346 100644 --- a/.github/workflows/test-linting.yml +++ b/.github/workflows/test-linting.yml @@ -77,18 +77,19 @@ jobs: run: | uv run --no-sync python scripts/ruff_strict_gate.py --base "$BASE_SHA" + - name: Check type-discipline budget (mutable collections / casts / type guards / kwargs / unexplained suppressions, delta vs base) + env: + BASE_SHA: ${{ github.event.pull_request.base.sha }} + run: | + uv run --no-sync python scripts/type_discipline_gate.py --base "$BASE_SHA" + - name: Print OpenAI version run: | uv run --no-sync python -c "import openai; print(f'OpenAI version: {openai.__version__}')" - - name: Run MyPy type checking - run: | - cd litellm - (uv run --no-sync mypy . || true) | uv run --no-sync python ../scripts/type_check_gate.py --tool mypy - - name: Run basedpyright type checking run: | - (uv run --no-sync basedpyright --outputjson || true) | uv run --no-sync python scripts/type_check_gate.py --tool basedpyright + (uv run --no-sync basedpyright --outputjson || true) | uv run --no-sync python scripts/type_check_gate.py - name: Check for circular imports run: | @@ -100,22 +101,20 @@ jobs: run: | uv run --no-sync python -c "from litellm import *" || (echo '🚨 import failed, this means you introduced unprotected imports! 🚨'; exit 1) - any-discipline: - # Separate job: the first run cold-builds litellm's type cache (~2 min, ~3 GB), - # so keep it off the main lint job's time budget. Subsequent runs reuse the - # cached .mypy_cache_any and only re-type-check the changed files. + # Intentionally NON-GATING. This job turns red when a *-budget.json ceiling is + # raised (or a rule/budget is dropped) so a loosening is obvious in review, but it + # must be kept OUT of the branch-protection required-checks list so a justified + # bump can still be merged by a human who has seen and accepted the red. + budget-ratchet: runs-on: ubuntu-latest - timeout-minutes: 10 + timeout-minutes: 5 + permissions: + contents: read steps: - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 - # Check out the PR head, not the default refs/pull/N/merge: the merge ref - # folds in newer base commits, which the diff-based gates (ruff delta, - # Any-discipline) would otherwise blame on this branch. with: - ref: ${{ github.event.pull_request.head.sha }} fetch-depth: 0 - clean: true persist-credentials: false - name: Set up Python @@ -123,32 +122,11 @@ jobs: with: python-version: "3.12" - - name: Set up uv - uses: astral-sh/setup-uv@37802adc94f370d6bfd71619e3f0bf239e1f3b78 # v7 - with: - version: "0.10.9" - - - name: Install dependencies - run: | - uv sync --frozen - - # Keyed on deps + mypy config (which fix the type cache's validity), not on - # source content, so changed files always differ from the restored cache. - # The gate also defensively invalidates each target's cache entry, so - # correctness never depends on cache freshness -- this is purely for speed. - - name: Restore Any-gate type cache - uses: actions/cache@0057852bfaa89a56745cba8c7296529d2fc39830 # v4.3.0 - with: - path: .mypy_cache_any - key: any-mypy-cache-${{ runner.os }}-py3.12-${{ hashFiles('uv.lock', 'litellm/mypy.ini') }} - restore-keys: | - any-mypy-cache-${{ runner.os }}-py3.12- - - - name: Check Any discipline on changed lines + - name: Ratchet check (budgets may only decrease; non-gating) env: BASE_SHA: ${{ github.event.pull_request.base.sha }} run: | - uv run --no-sync python scripts/check_any_discipline.py --changed --base "$BASE_SHA" + python scripts/budget_ratchet_check.py --base "$BASE_SHA" secret-scan: runs-on: ubuntu-latest diff --git a/.github/workflows/triage_issue_with_llm.yml b/.github/workflows/triage_issue_with_llm.yml new file mode 100644 index 00000000000..765453cf2c6 --- /dev/null +++ b/.github/workflows/triage_issue_with_llm.yml @@ -0,0 +1,96 @@ +name: Agent Shin — Issue triage + +# LLM-as-judge triage for external GitHub issues. +# +# DRY-RUN BY DEFAULT. See .github/workflows/triage_pr_with_llm.yml for the +# enablement procedure — same repo variable (`AGENT_SHIN_ENABLED=true`) +# unlocks the PR and issue triage flows together. + +on: + issues: + types: [opened, reopened] + workflow_dispatch: + inputs: + issue_number: + description: "Issue number to triage manually." + required: true + close: + description: "If true and AGENT_SHIN_ENABLED=true, actually close on fail." + required: false + default: "false" + type: choice + options: + - "true" + - "false" + +permissions: + contents: read + issues: write + +jobs: + triage: + if: github.repository == 'BerriAI/litellm' + runs-on: ubuntu-latest + steps: + - name: Checkout triage script + uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + with: + sparse-checkout: .github/scripts + persist-credentials: false + + - name: Set up Python + uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 + with: + python-version: "3.12" + + - name: Install LLM client + run: pip install --no-cache-dir --require-hashes -r .github/scripts/triage-requirements.txt + + - name: Run Agent Shin + env: + GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} + # Only expose the LLM key when the bot is enabled or a collaborator + # triggers it manually, so an external user can't force paid LLM + # calls by churning issues while the bot is still in dry-run. + # The Python script calls the LLM whenever this var is set + # (regardless of `--close`); stripping `--close` doesn't suppress + # the API call, only the destructive side effects. + OPENAI_API_KEY: ${{ (vars.AGENT_SHIN_ENABLED == 'true' || github.event_name == 'workflow_dispatch') && secrets.OPENAI_API_KEY || '' }} + OPENAI_BASE_URL: ${{ vars.OPENAI_BASE_URL }} + TRIAGE_MODEL: ${{ vars.TRIAGE_MODEL }} + AGENT_SHIN_ENABLED: ${{ vars.AGENT_SHIN_ENABLED }} + DISPATCH_CLOSE: ${{ github.event.inputs.close }} + ISSUE_NUMBER: ${{ github.event.issue.number || github.event.inputs.issue_number }} + run: | + set -euo pipefail + ARGS=(--repo "${{ github.repository }}" --issue "${ISSUE_NUMBER}") + # Fail-safe gating: only the EXACT string "true" enables the + # destructive --close path. The workflow_dispatch input is a + # `choice` dropdown of "true"/"false" so the UI is constrained, + # but the API (`gh workflow run -f close=...`) accepts any + # string, and a `!= "false"` check would treat "True", "yes", + # "1", "TRUE", typos, and accidental whitespace as enabling + # closure. Mirror the Greptile closer's `= "true"` pattern. + if [ "${AGENT_SHIN_ENABLED:-false}" = "true" ] && [ "${DISPATCH_CLOSE:-false}" = "true" ]; then + ARGS+=(--close) + echo "::notice::Agent Shin is ENABLED and running in close-on-fail mode." + elif [ "${AGENT_SHIN_ENABLED:-false}" = "true" ]; then + echo "::notice::Agent Shin is ENABLED but this trigger is dry-run (workflow_dispatch close != 'true')." + else + echo "::notice::Agent Shin is in DRY-RUN mode (AGENT_SHIN_ENABLED is not 'true'). No comments will be posted; no issues will be closed." + fi + # Automatic `issues` events stay dry-run regardless until the team + # explicitly invokes workflow_dispatch with close=true. + if [ "${GITHUB_EVENT_NAME:-}" = "issues" ]; then + # filter out --close rather than substituting to "" (which would + # leave an empty positional arg that argparse rejects) + FILTERED=() + for arg in "${ARGS[@]}"; do + if [ "${arg}" != "--close" ]; then + FILTERED+=("${arg}") + fi + done + ARGS=("${FILTERED[@]}") + echo "::notice::issues trigger -> forcing dry-run." + fi + python3 .github/scripts/triage_with_llm.py "${ARGS[@]}" diff --git a/.github/workflows/triage_pr_with_llm.yml b/.github/workflows/triage_pr_with_llm.yml new file mode 100644 index 00000000000..936547598fb --- /dev/null +++ b/.github/workflows/triage_pr_with_llm.yml @@ -0,0 +1,110 @@ +name: Agent Shin — PR triage + +# LLM-as-judge triage for external pull requests. +# +# DRY-RUN BY DEFAULT. Closures and public comments are gated on the repo +# variable `AGENT_SHIN_ENABLED` being set to the string `"true"`. Until then, +# every run only writes its verdict to the workflow step summary so the team +# can QA the judge's decisions before flipping it on. +# +# To enable for real: +# 1. Add a repo secret `OPENAI_API_KEY` (or compatible). +# 2. Set repo variable `AGENT_SHIN_ENABLED` to `true` +# (Settings > Secrets and variables > Actions > Variables). +# +# We use `pull_request_target` so the workflow has access to repo secrets +# and runs against PRs from forks. We never check out fork code — only read +# PR metadata via `gh api`, so this is safe. + +on: + pull_request_target: + types: [opened, reopened] + workflow_dispatch: + inputs: + pr_number: + description: "PR number to triage manually." + required: true + close: + description: "If true and AGENT_SHIN_ENABLED=true, actually close on fail." + required: false + default: "false" + type: choice + options: + - "true" + - "false" + +permissions: + contents: read + issues: write + pull-requests: write + +jobs: + triage: + if: github.repository == 'BerriAI/litellm' + runs-on: ubuntu-latest + steps: + - name: Checkout triage script + uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + with: + sparse-checkout: .github/scripts + persist-credentials: false + + - name: Set up Python + uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 + with: + python-version: "3.12" + + - name: Install LLM client + run: pip install --no-cache-dir --require-hashes -r .github/scripts/triage-requirements.txt + + - name: Run Agent Shin + env: + GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} + # Only expose the LLM key when the bot is enabled or a collaborator + # triggers it manually, so an external user can't force paid LLM + # calls by churning a fork PR while the bot is still in dry-run. + # The Python script calls the LLM whenever this var is set + # (regardless of `--close`); stripping `--close` doesn't suppress + # the API call, only the destructive side effects. + OPENAI_API_KEY: ${{ (vars.AGENT_SHIN_ENABLED == 'true' || github.event_name == 'workflow_dispatch') && secrets.OPENAI_API_KEY || '' }} + OPENAI_BASE_URL: ${{ vars.OPENAI_BASE_URL }} + TRIAGE_MODEL: ${{ vars.TRIAGE_MODEL }} + AGENT_SHIN_ENABLED: ${{ vars.AGENT_SHIN_ENABLED }} + DISPATCH_CLOSE: ${{ github.event.inputs.close }} + PR_NUMBER: ${{ github.event.pull_request.number || github.event.inputs.pr_number }} + run: | + set -euo pipefail + ARGS=(--repo "${{ github.repository }}" --pr "${PR_NUMBER}") + # Fail-safe gating: only the EXACT string "true" enables the + # destructive --close path. The workflow_dispatch input is a + # `choice` dropdown of "true"/"false" so the UI is constrained, + # but the API (`gh workflow run -f close=...`) accepts any + # string, and a `!= "false"` check would treat "True", "yes", + # "1", "TRUE", typos, and accidental whitespace as enabling + # closure. Mirror the Greptile closer's `= "true"` pattern. + if [ "${AGENT_SHIN_ENABLED:-false}" = "true" ] && [ "${DISPATCH_CLOSE:-false}" = "true" ]; then + ARGS+=(--close) + echo "::notice::Agent Shin is ENABLED and running in close-on-fail mode." + elif [ "${AGENT_SHIN_ENABLED:-false}" = "true" ]; then + echo "::notice::Agent Shin is ENABLED but this trigger is dry-run (workflow_dispatch close != 'true' or scheduled event)." + else + echo "::notice::Agent Shin is in DRY-RUN mode (AGENT_SHIN_ENABLED is not 'true'). No comments will be posted; no PRs will be closed." + fi + # On the scheduled/automatic pull_request_target trigger we default to + # dry-run regardless, so the team can review verdicts in the step + # summary before any contributor sees a comment. Only the manual + # workflow_dispatch path (with close=true) closes PRs. + if [ "${GITHUB_EVENT_NAME:-}" = "pull_request_target" ]; then + # strip any --close added above (filter out, don't substitute + # to empty string — that would leave a stray "" positional arg + # that argparse rejects) + FILTERED=() + for arg in "${ARGS[@]}"; do + if [ "${arg}" != "--close" ]; then + FILTERED+=("${arg}") + fi + done + ARGS=("${FILTERED[@]}") + echo "::notice::pull_request_target trigger -> forcing dry-run." + fi + python3 .github/scripts/triage_with_llm.py "${ARGS[@]}" diff --git a/.github/workflows/triage_reconsider.yml b/.github/workflows/triage_reconsider.yml new file mode 100644 index 00000000000..f35f681d09a --- /dev/null +++ b/.github/workflows/triage_reconsider.yml @@ -0,0 +1,172 @@ +name: Agent Shin — reconsider + +# Comment-trigger workflow: when the PR/issue author (or an internal +# collaborator) comments `@agent-shin reconsider` on a CLOSED PR/issue, +# Agent Shin re-runs LLM-judge triage on the current title+body and: +# +# - on PASS: posts a "re-evaluated and reopened" comment + reopens. +# - on FAIL: posts a "still missing X" comment and leaves it closed, +# so the contributor can iterate again. +# +# This exists because GitHub does NOT let an external (non-write-access) +# OSS contributor reopen a PR/issue closed by a bot or maintainer. Without +# this comment trigger, a contributor whose PR Agent Shin auto-closed +# would have no path back into the review queue except opening a fresh PR +# (which loses the original PR's history). The bot, on the other hand, +# has write access via GH_TOKEN and can reopen on their behalf. +# +# DRY-RUN BY DEFAULT — gated on `vars.AGENT_SHIN_ENABLED == 'true'` just +# like the other Agent Shin workflows. The workflow also gates on the +# commenter being either the PR/issue author or an internal collaborator +# (OWNER/MEMBER/COLLABORATOR) so random commenters cannot DOS the LLM +# judge or force a reopen. + +on: + issue_comment: + types: [created] + +permissions: + contents: read + issues: write + pull-requests: write + +jobs: + reconsider: + if: | + github.repository == 'BerriAI/litellm' + && contains(github.event.comment.body, '@agent-shin reconsider') + runs-on: ubuntu-latest + steps: + - name: Authorize commenter + # Only the PR/issue author OR an internal collaborator may trigger + # a reconsider. Outside random commenters could otherwise spam the + # phrase to burn LLM budget or, if a fail-open bug were ever + # introduced, force a reopen on someone else's behalf. + # + # We expose the authorization decision as a step output and gate + # every subsequent (potentially destructive) step on it. A `run:` + # step with `exit 0` would NOT stop the job — only `if:` gating + # on a known-true output is safe here. + id: auth + env: + COMMENTER: ${{ github.event.comment.user.login }} + AUTHOR: ${{ github.event.issue.user.login }} + ASSOCIATION: ${{ github.event.comment.author_association }} + run: | + set -euo pipefail + if [ "${COMMENTER}" = "${AUTHOR}" ]; then + echo "::notice::Authorized: commenter is the PR/issue author." + echo "authorized=true" >> "$GITHUB_OUTPUT" + exit 0 + fi + case "${ASSOCIATION}" in + OWNER|MEMBER|COLLABORATOR) + echo "::notice::Authorized: commenter is an internal collaborator (${ASSOCIATION})." + echo "authorized=true" >> "$GITHUB_OUTPUT" + ;; + *) + echo "::notice::Commenter '${COMMENTER}' (${ASSOCIATION}) is not authorized to trigger reconsider; skipping subsequent steps." + echo "authorized=false" >> "$GITHUB_OUTPUT" + ;; + esac + + - name: React 👀 to acknowledge the reconsider + # Add an eyes reaction to the triggering comment the moment we accept + # it, so the contributor gets instant feedback that the bot saw their + # `@agent-shin reconsider` before the slower triage steps run. Gated on + # AGENT_SHIN_ENABLED so dry-run leaves no visible trace. Best-effort: + # a reactions API hiccup must never fail the actual reconsider. + if: steps.auth.outputs.authorized == 'true' && vars.AGENT_SHIN_ENABLED == 'true' + env: + GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} + COMMENT_ID: ${{ github.event.comment.id }} + run: | + set -euo pipefail + gh api --method POST \ + -H "Accept: application/vnd.github+json" \ + "repos/${{ github.repository }}/issues/comments/${COMMENT_ID}/reactions" \ + -f content=eyes \ + || echo "::warning::failed to add 👀 reaction (non-fatal)" + + - name: Checkout triage script + if: steps.auth.outputs.authorized == 'true' + uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + with: + sparse-checkout: .github/scripts + persist-credentials: false + + - name: Set up Python + if: steps.auth.outputs.authorized == 'true' + uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 + with: + python-version: "3.12" + + - name: Install LLM client + if: steps.auth.outputs.authorized == 'true' + run: pip install --no-cache-dir --require-hashes -r .github/scripts/triage-requirements.txt + + - name: Run Agent Shin reconsider + if: steps.auth.outputs.authorized == 'true' + env: + GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} + # Only expose the LLM key when the bot is enabled, so a PR/issue + # author can't force paid LLM calls by spamming `@agent-shin + # reconsider` while the bot is still in dry-run. The Python script + # calls the LLM whenever this var is set (regardless of `--close`); + # stripping `--close` doesn't suppress the API call, only the + # destructive side effects. Mirror the gating used by every other + # Agent Shin workflow (triage_pr_with_llm.yml, review_gate.yml, ...). + OPENAI_API_KEY: ${{ vars.AGENT_SHIN_ENABLED == 'true' && secrets.OPENAI_API_KEY || '' }} + OPENAI_BASE_URL: ${{ vars.OPENAI_BASE_URL }} + TRIAGE_MODEL: ${{ vars.TRIAGE_MODEL }} + AGENT_SHIN_ENABLED: ${{ vars.AGENT_SHIN_ENABLED }} + # `issue_comment` events fire for both issues and PR comments. + # `issue.pull_request` is set iff this is a PR comment, so we use + # its presence to decide whether to invoke `--pr N` or `--issue N`. + IS_PR: ${{ github.event.issue.pull_request != null }} + NUMBER: ${{ github.event.issue.number }} + run: | + set -euo pipefail + if [ "${IS_PR}" = "true" ]; then + ARGS=(--repo "${{ github.repository }}" --pr "${NUMBER}" --reconsider) + else + ARGS=(--repo "${{ github.repository }}" --issue "${NUMBER}" --reconsider) + fi + # Reconsider's destructive actions (post comment + reopen) are + # gated on `--close`, mirroring the regular triage workflows. + # When AGENT_SHIN_ENABLED is not the EXACT string "true", we + # still run the script so its verdict + would-X action lands in + # the step summary for QA — but without `--close`, the script + # returns `would-reopen` / `would-reconsider-still-failing` + # instead of touching GitHub state. + # + # Use the positive `= "true"` gate (not `!= "true" -> exit`) so + # the workflow guardrails in + # tests/test_litellm/test_github_triage_workflows.py see the + # canonical fail-safe enable pattern. Unknown values like + # "True", "yes", "1", or typos fall through to the dry-run + # branch, which is the safe default. + if [ "${AGENT_SHIN_ENABLED:-false}" = "true" ]; then + ARGS+=(--close) + echo "::notice::Agent Shin reconsider ENABLED — running real triage (close=true)." + else + echo "::notice::AGENT_SHIN_ENABLED is not 'true' -> reconsider stays in dry-run (no comment, no reopen)." + fi + python3 .github/scripts/triage_with_llm.py "${ARGS[@]}" + + - name: React 👍 when the reconsider finishes + # Once the reconsider run has completed successfully, add a thumbs-up so + # the contributor sees the bot is done (the 👀 stays, signalling + # seen -> handled). `success()` keeps this from firing if the run + # errored, and the AGENT_SHIN_ENABLED gate keeps dry-run inert. + if: success() && steps.auth.outputs.authorized == 'true' && vars.AGENT_SHIN_ENABLED == 'true' + env: + GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} + COMMENT_ID: ${{ github.event.comment.id }} + run: | + set -euo pipefail + gh api --method POST \ + -H "Accept: application/vnd.github+json" \ + "repos/${{ github.repository }}/issues/comments/${COMMENT_ID}/reactions" \ + -f content=+1 \ + || echo "::warning::failed to add 👍 reaction (non-fatal)" diff --git a/.github/workflows/triage_rollout_heads_up.yml b/.github/workflows/triage_rollout_heads_up.yml new file mode 100644 index 00000000000..903960151e2 --- /dev/null +++ b/.github/workflows/triage_rollout_heads_up.yml @@ -0,0 +1,92 @@ +name: Agent Shin — rollout heads-up (one-shot) + +# Fires the 7-day heads-up comment on every open external PR/issue that the +# new triage bot would auto-close. The real sweep is a deliberate one-shot: +# trigger it at rollout via a manual `workflow_dispatch` with `dry_run=false`. +# The script is idempotent (skips items that already carry the +# `` marker), so a re-run is harmless. +# +# The automatic push trigger runs DRY-RUN only, so merging the script to +# `litellm_internal_staging` never posts a comment; it just confirms the +# workflow is wired up. Posting real comments requires the manual dispatch, +# which is also the only trigger that exposes `OPENAI_API_KEY`. The heads-up +# is intentionally NOT gated on `AGENT_SHIN_ENABLED`: it has to warn +# contributors while that flag is still off, ahead of the flip that turns on +# auto-closing. +# +# The workflow is a thin shell over `.github/scripts/triage_rollout_heads_up.py`. +# Dry-run vs. real run differ in EXACTLY one CLI flag (`--close`), added only +# on a manual dispatch with `dry_run=false`. + +on: + push: + branches: + - litellm_internal_staging + paths: + # The presence of this script on staging IS the rollout merge marker. + # Editing the file later would re-fire the workflow; that's safe because + # the script skips PRs/issues that already have the heads-up marker. + - ".github/scripts/triage_rollout_heads_up.py" + workflow_dispatch: + inputs: + dry_run: + description: "Dry run (true = preview only, false = actually post comments)." + required: false + default: "true" + type: choice + options: + - "true" + - "false" + +permissions: + contents: read + issues: write + pull-requests: write + +jobs: + heads-up: + if: github.repository == 'BerriAI/litellm' + runs-on: ubuntu-latest + steps: + - name: Checkout triage scripts + uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + with: + sparse-checkout: .github/scripts + persist-credentials: false + + - name: Set up Python + uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 + with: + python-version: "3.12" + + - name: Install LLM client + run: pip install --no-cache-dir --require-hashes -r .github/scripts/triage-requirements.txt + + - name: Run heads-up sweep + env: + GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} + # Only the manual dispatch (the real-run trigger) needs the LLM key. + # The automatic push trigger runs dry-run and never posts, so it gets + # no key. Mirrors the sibling triage workflows, which expose the key + # only on an enabled/dispatched run rather than unconditionally. + OPENAI_API_KEY: ${{ github.event_name == 'workflow_dispatch' && secrets.OPENAI_API_KEY || '' }} + OPENAI_BASE_URL: ${{ vars.OPENAI_BASE_URL }} + TRIAGE_MODEL: ${{ vars.TRIAGE_MODEL }} + # The real run is a deliberate manual dispatch with dry_run=false. + # Use the EXACT "false" comparison so any unexpected input value + # fail-closes to dry-run (mirrors the AGENT_SHIN_ENABLED pattern in + # the sibling workflows). The automatic push trigger always stays + # dry-run, so merging the script never posts. + DRY_RUN_INPUT: ${{ github.event.inputs.dry_run }} + run: | + set -euo pipefail + ARGS=(--repo "${{ github.repository }}") + if [ "${GITHUB_EVENT_NAME:-}" = "workflow_dispatch" ] && [ "${DRY_RUN_INPUT:-true}" = "false" ]; then + ARGS+=(--close) + echo "::notice::Manual rollout dispatch with dry_run=false -> heads-up comments WILL be posted." + elif [ "${GITHUB_EVENT_NAME:-}" = "workflow_dispatch" ]; then + echo "::notice::Manual dispatch in dry-run mode -> previewing only, no comments will be posted." + else + echo "::notice::Automatic push trigger -> dry-run preview only. Fire the real rollout sweep with a manual workflow_dispatch (dry_run=false)." + fi + python3 .github/scripts/triage_rollout_heads_up.py "${ARGS[@]}" diff --git a/.gitignore b/.gitignore index 54ae53bb2c9..fda3311fe02 100644 --- a/.gitignore +++ b/.gitignore @@ -74,8 +74,6 @@ tests/local_testing/log.txt .codegpt litellm/proxy/_new_new_secret_config.yaml litellm/proxy/custom_guardrail.py -**/.mypy_cache/ -**/.mypy_cache_any/ litellm/proxy/application.log tests/llm_translation/vertex_test_account.json tests/llm_translation/test_vertex_key.json diff --git a/CLAUDE.md b/CLAUDE.md index 48dc3d81d94..2070b6fcdd6 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -36,11 +36,11 @@ Don't hesitate to use values in .env to get needed API keys and other secrets, a Run tests, format your code, and lint your code before each commit -When you fix violations gated by `ruff-strict-budget.json`, `mypy-code-budget.json`, or `basedpyright-code-budget.json`, run `make lint-budget-update` and commit the lowered baselines so the ceilings ratchet down instead of leaving stale headroom +When you fix violations gated by `ruff-strict-budget.json` or `basedpyright-code-budget.json`, run `make lint-budget-update` and commit the lowered baselines so the ceilings ratchet down instead of leaving stale headroom -If you're trying to create a new function that relies on untyped stuff, instead of adding more Any's and bringing it closer to the max, just validate it in the caller with Pydantic (a model or `TypeAdapter` that returns the typed thing or raises will do) and then pass the now typed variable in +If you're trying to create a new function that relies on untyped stuff, instead of adding more Any's and pushing `reportAny` / `reportExplicitAny` closer to their basedpyright ceilings, just validate it in the caller with Pydantic (a model or `TypeAdapter` that returns the typed thing or raises will do) and then pass the now typed variable in -The Any-discipline gate (`make lint-any`, also a CI job) fails when a line you changed under `litellm/` holds a value typed `Any`, including the `X | Any`. Ideally `# any-ok: ` is never used; treat it as a last resort for a genuine typed/untyped boundary that Pydantic truly can't model +If you get an LIT001 or LIT002 fail, refactor the code to follow functional programming best practices rather than introducing mutable data structures. For example, build values in one shot with comprehensions or generators wrapped in `tuple()` / `frozenset()` instead of seeding an empty `list`/`dict`/`set` and mutating it over time. Ideally `# mutable-ok` is never used; reach for it only as a genuine last resort when an immutable rewrite is truly impossible, and always pair it with a real reason Ask to commit and push your work when you're done (or if you're confident that your code is good and works, just do it) diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 97a8d53f831..1080579d0fa 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -154,8 +154,7 @@ Individual linting commands: ```bash make format-check # Check Black formatting make lint-ruff # Run Ruff linting -make lint-mypy # Run MyPy type checking -make lint-any # Fail on Any-typed values on changed lines +make lint-basedpyright # Run basedpyright type checking make check-circular-imports # Check for circular imports make check-import-safety # Check import safety ``` @@ -217,7 +216,7 @@ LiteLLM follows the [Google Python Style Guide](https://google.github.io/stylegu Our automated quality checks include: - **Black** for consistent code formatting - **Ruff** for linting and code quality -- **MyPy** for static type checking +- **basedpyright** for static type checking - **Circular import detection** - **Import safety validation** @@ -231,7 +230,7 @@ If `make lint` fails: 1. **Formatting issues**: Run `make format` to auto-fix 2. **Ruff issues**: Check the output and fix manually -3. **MyPy issues**: Add proper type hints +3. **basedpyright issues**: Add proper type hints 4. **Circular imports**: Refactor import dependencies 5. **Import safety**: Fix any unprotected imports @@ -246,7 +245,7 @@ If `make test-unit` fails: ### 3. Common Development Tips -- **Use type hints**: MyPy requires proper type annotations +- **Use type hints**: basedpyright requires proper type annotations - **Write descriptive commit messages**: Help reviewers understand your changes - **Keep PRs focused**: One feature/fix per PR - **Test edge cases**: Don't just test the happy path diff --git a/Makefile b/Makefile index f0563b273c2..6183dff1556 100644 --- a/Makefile +++ b/Makefile @@ -5,8 +5,8 @@ test-unit-integrations test-unit-core-utils test-unit-other test-unit-root \ test-proxy-unit-a test-proxy-unit-b test-integration test-unit-helm \ info lint lint-dev format \ - lint-mypy lint-mypy-budget-update lint-basedpyright lint-basedpyright-budget-update \ - lint-ruff-budget lint-ruff-budget-update lint-budget-update lint-any \ + lint-basedpyright lint-basedpyright-budget-update \ + lint-ruff-budget lint-ruff-budget-update lint-budget-update \ install-dev install-proxy-dev install-test-deps install-hooks \ install-helm-unittest check-circular-imports check-import-safety @@ -22,17 +22,14 @@ help: @echo " make install-hooks - Install git hooks (Conventional Commits + Branches)" @echo " make format - Apply Black code formatting" @echo " make format-check - Check Black code formatting (matches CI)" - @echo " make lint - Run all linting (Ruff, MyPy, Black check, circular imports, import safety)" + @echo " make lint - Run all linting (Ruff, basedpyright, Black check, circular imports, import safety)" @echo " make lint-ruff - Run Ruff linting only" - @echo " make lint-mypy - Run MyPy (disallow_untyped_defs), gated by per-rule error counts" - @echo " make lint-mypy-budget-update - Re-capture the MyPy per-rule budget (ratchet)" @echo " make lint-basedpyright - Run basedpyright strict, gated by per-rule error counts" @echo " make lint-basedpyright-budget-update - Re-capture the basedpyright per-rule budget (ratchet)" @echo " make lint-black - Check Black formatting (matches CI)" @echo " make lint-ruff-budget - Gate the codebase total of each strict ruff rule against its ceiling" @echo " make lint-ruff-budget-update - Re-capture per-rule baselines in ruff-strict-budget.json (ratchet)" - @echo " make lint-budget-update - Re-capture all three ratchet budgets (ruff + mypy + basedpyright)" - @echo " make lint-any - Fail if changed lines under litellm/ hold an Any-typed value" + @echo " make lint-budget-update - Re-capture all ratchet budgets (ruff + basedpyright)" @echo " make check-circular-imports - Check for circular imports" @echo " make check-import-safety - Check import safety" @echo " make test - Run all tests" @@ -126,17 +123,11 @@ lint-ruff-FULL-dev: install-dev if [ -n "$$files" ]; then echo "$$files" | xargs $(UV_RUN) ruff check; \ else echo "No changed .py files to check."; fi -lint-mypy: install-dev - cd litellm && ($(UV_RUN) mypy . || true) | $(UV_RUN) python ../scripts/type_check_gate.py --tool mypy - -lint-mypy-budget-update: install-dev - cd litellm && ($(UV_RUN) mypy . || true) | $(UV_RUN) python ../scripts/type_check_gate.py --tool mypy --update - lint-basedpyright: install-dev - ($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --tool basedpyright + ($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py lint-basedpyright-budget-update: install-dev - ($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --tool basedpyright --update + ($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --update lint-black: format-check @@ -146,11 +137,8 @@ lint-ruff-budget: install-dev lint-ruff-budget-update: install-dev $(UV_RUN) python scripts/ruff_strict_gate.py --update -# Ratchet all three budgets in one shot (ruff strict + mypy + basedpyright) -lint-budget-update: lint-ruff-budget-update lint-mypy-budget-update lint-basedpyright-budget-update - -lint-any: install-dev - $(UV_RUN) python scripts/check_any_discipline.py --changed +# Ratchet all budgets in one shot (ruff strict + basedpyright) +lint-budget-update: lint-ruff-budget-update lint-basedpyright-budget-update check-circular-imports: install-dev cd litellm && $(UV_RUN) python ../tests/documentation_tests/test_circular_imports.py && cd .. @@ -159,10 +147,10 @@ check-import-safety: install-dev @$(UV_RUN) python -c "from litellm import *; print('[from litellm import *] OK! no issues!');" || (echo '🚨 import failed, this means you introduced unprotected imports! 🚨'; exit 1) # Combined linting (matches test-linting.yml workflow) -lint: format-check lint-ruff lint-mypy lint-basedpyright check-circular-imports check-import-safety lint-ruff-budget lint-any +lint: format-check lint-ruff lint-basedpyright check-circular-imports check-import-safety lint-ruff-budget # Faster linting for local development (only checks changed code) -lint-dev: lint-format-changed lint-mypy lint-any check-circular-imports check-import-safety +lint-dev: lint-format-changed check-circular-imports check-import-safety # Testing targets test: install-test-deps diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index b531e0e17df..73bc5c47703 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -1,11 +1,11 @@ { "reportAny": { "baseline": 24954, - "slack": 10 + "slack": 2500 }, "reportArgumentType": { "baseline": 1863, - "slack": 3 + "slack": 180 }, "reportAssignmentType": { "baseline": 220, @@ -33,7 +33,7 @@ }, "reportExplicitAny": { "baseline": 6931, - "slack": 10 + "slack": 700 }, "reportFunctionMemberAccess": { "baseline": 7, @@ -77,7 +77,7 @@ }, "reportMissingTypeArgument": { "baseline": 10612, - "slack": 10 + "slack": 1000 }, "reportMissingTypeStubs": { "baseline": 27, @@ -113,7 +113,7 @@ }, "reportPrivateUsage": { "baseline": 1625, - "slack": 10 + "slack": 160 }, "reportRedeclaration": { "baseline": 8, @@ -133,7 +133,7 @@ }, "reportUnknownArgumentType": { "baseline": 30603, - "slack": 10 + "slack": 3000 }, "reportUnknownLambdaType": { "baseline": 76, @@ -141,15 +141,15 @@ }, "reportUnknownMemberType": { "baseline": 27322, - "slack": 10 + "slack": 2500 }, "reportUnknownParameterType": { "baseline": 13636, - "slack": 10 + "slack": 1000 }, "reportUnknownVariableType": { "baseline": 21776, - "slack": 10 + "slack": 2000 }, "reportUnnecessaryCast": { "baseline": 118, diff --git a/db_scripts/create_views.py b/db_scripts/create_views.py index 3027b38958d..2b34664452d 100644 --- a/db_scripts/create_views.py +++ b/db_scripts/create_views.py @@ -15,7 +15,7 @@ db = Prisma( ) -async def check_view_exists(): # noqa: PLR0915 +async def check_view_exists(): """ Checks if the LiteLLM_VerificationTokenView and MonthlyGlobalSpend exists in the user's db. @@ -34,8 +34,7 @@ async def check_view_exists(): # noqa: PLR0915 print("LiteLLM_VerificationTokenView Exists!") # noqa except Exception: # If an error occurs, the view does not exist, so create it - await db.execute_raw( - """ + await db.execute_raw(""" CREATE VIEW "LiteLLM_VerificationTokenView" AS SELECT v.*, @@ -45,8 +44,7 @@ async def check_view_exists(): # noqa: PLR0915 t.rpm_limit AS team_rpm_limit FROM "LiteLLM_VerificationToken" v LEFT JOIN "LiteLLM_TeamTable" t ON v.team_id = t.team_id; - """ - ) + """) print("LiteLLM_VerificationTokenView Created!") # noqa diff --git a/docker/build_from_pip/litellm_config.yaml b/docker/build_from_pip/litellm_config.yaml index 51223026170..f54647853ef 100644 --- a/docker/build_from_pip/litellm_config.yaml +++ b/docker/build_from_pip/litellm_config.yaml @@ -3,7 +3,7 @@ model_list: litellm_params: model: openai/fake api_key: fake-key - api_base: https://exampleopenaiendpoint-production.up.railway.app/ + api_base: os.environ/FAKE_OPENAI_API_BASE general_settings: alerting: ["slack"] \ No newline at end of file diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index a1f63f388b4..6830147116d 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -412,7 +412,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): detail=f"User {user_api_key_dict.user_id} does not have access to the file {file_id}", ) - async def async_pre_call_hook( # noqa: PLR0915 + async def async_pre_call_hook( self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache, diff --git a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py index 75229bacc8f..a057df65500 100644 --- a/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py +++ b/enterprise/litellm_enterprise/proxy/management_endpoints/project_endpoints.py @@ -483,7 +483,7 @@ async def new_project( response_model=LiteLLM_ProjectTable, ) @management_endpoint_wrapper -async def update_project( # noqa: PLR0915 +async def update_project( data: UpdateProjectRequest, http_request: Request, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), diff --git a/litellm/__init__.py b/litellm/__init__.py index 0d6a788e368..cffdbacf597 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -213,6 +213,15 @@ standard_logging_payload_excluded_fields: Optional[List[str]] = ( log_raw_request_response: bool = False redact_messages_in_exceptions: Optional[bool] = False redact_user_api_key_info: Optional[bool] = False +# When True (default — preserves historical behavior), the Router appends +# internal config names (model_group, fallback model groups, deployment +# timeouts, fallback failure details) onto exception messages and surfaces +# them to clients via ProxyException.message. Set to False if you do NOT +# want the proxy's internal model_name / fallback wiring visible to clients. +# Deprecation: planned to flip to False (redact by default) in a future +# major release; opt in early with `litellm.expose_router_debug_in_errors +# = False`. +expose_router_debug_in_errors: bool = True filter_invalid_headers: Optional[bool] = False add_user_information_to_llm_headers: Optional[bool] = ( None # adds user_id, team_id, token hash (params from StandardLoggingMetadata) to request headers @@ -235,6 +244,17 @@ modify_params = bool(os.getenv("LITELLM_MODIFY_PARAMS", False)) use_chat_completions_url_for_anthropic_messages: bool = bool( os.getenv("LITELLM_USE_CHAT_COMPLETIONS_URL_FOR_ANTHROPIC_MESSAGES", False) ) # When True, routes OpenAI /v1/messages requests to chat/completions instead of the Responses API +# When True, strip the OpenAI-flavored `usage.total_tokens` field that +# LiteLLM injects into non-streaming /v1/messages responses, bringing the +# wire response into line with the Anthropic spec (matches the streaming +# SSE path, which already omits total_tokens). Default False to preserve +# backward compatibility for clients that read the LiteLLM-shaped +# `usage.total_tokens` today. Planned to flip to True in a future major +# release; opt in early via Python: +# `litellm.strip_anthropic_total_tokens = True` +# Or via `litellm_settings.strip_anthropic_total_tokens: true` in +# config.yaml. +strip_anthropic_total_tokens: bool = False route_all_chat_openai_to_responses: bool = ( os.getenv("LITELLM_ROUTE_ALL_CHAT_OPENAI_TO_RESPONSES", "false").lower() == "true" ) # When True, routes all OpenAI /chat/completions requests through the Responses API bridge @@ -413,7 +433,7 @@ anthropic_beta_headers_url: str = os.getenv( "LITELLM_ANTHROPIC_BETA_HEADERS_URL", "https://raw.githubusercontent.com/BerriAI/litellm/main/litellm/anthropic_beta_headers_config.json", ) -suppress_debug_info = False +suppress_debug_info: bool = False dynamodb_table_name: Optional[str] = None s3_callback_params: Optional[Dict] = None s3_audit_callback_params: Optional[Dict] = None diff --git a/litellm/_redis.py b/litellm/_redis.py index e2b04f795cb..1b6e1a5e4b0 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -311,7 +311,7 @@ def get_redis_url_from_environment(): return f"{redis_protocol}://{auth_part}{os.environ['REDIS_HOST']}:{os.environ['REDIS_PORT']}" -def _get_redis_client_logic(**env_overrides): # noqa: PLR0915 +def _get_redis_client_logic(**env_overrides): """ Common functionality across sync + async redis client implementations """ diff --git a/litellm/a2a_protocol/main.py b/litellm/a2a_protocol/main.py index dcb5cb74ec4..2b6f2cd12b4 100644 --- a/litellm/a2a_protocol/main.py +++ b/litellm/a2a_protocol/main.py @@ -436,7 +436,7 @@ def _build_streaming_logging_obj( return logging_obj -async def asend_message_streaming( # noqa: PLR0915 +async def asend_message_streaming( a2a_client: Optional["A2AClientType"] = None, request: Optional["SendStreamingMessageRequest"] = None, api_base: Optional[str] = None, diff --git a/litellm/batches/main.py b/litellm/batches/main.py index 15ee9303969..f124882b5a4 100644 --- a/litellm/batches/main.py +++ b/litellm/batches/main.py @@ -157,7 +157,7 @@ async def acreate_batch( @client -def create_batch( # noqa: PLR0915 +def create_batch( completion_window: Literal["24h"], endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions"], input_file_id: str, diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index 48691335b40..2a8bd856040 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -394,7 +394,7 @@ class LLMCachingHandler: return cr["model"] return None - def _process_async_embedding_cached_response( # noqa: PLR0915 + def _process_async_embedding_cached_response( self, final_embedding_cached_response: Optional[EmbeddingResponse], cached_result: List[Optional[CachedEmbedding]], diff --git a/litellm/caching/qdrant_semantic_cache.py b/litellm/caching/qdrant_semantic_cache.py index cb521efca05..68d3b8c20b3 100644 --- a/litellm/caching/qdrant_semantic_cache.py +++ b/litellm/caching/qdrant_semantic_cache.py @@ -28,7 +28,7 @@ from .base_cache import BaseCache class QdrantSemanticCache(BaseCache): CACHE_KEY_FIELD_NAME = "litellm_cache_key" - def __init__( # noqa: PLR0915 + def __init__( self, qdrant_api_base=None, qdrant_api_key=None, diff --git a/litellm/caching/redis_cache.py b/litellm/caching/redis_cache.py index 7239bea7853..263e1df2ee7 100644 --- a/litellm/caching/redis_cache.py +++ b/litellm/caching/redis_cache.py @@ -369,6 +369,8 @@ class RedisCache(BaseCache): """ Make sure each key starts with the given namespace """ + if key is None: + return key # type: ignore[return-value] if self.namespace is not None and not key.startswith(self.namespace): key = self.namespace + ":" + key diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index dabf09f8b2a..3fa6b983e5f 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -1211,7 +1211,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): return self.chunk_parser(json.loads(str_line)) @staticmethod - def translate_responses_chunk_to_openai_stream( # noqa: PLR0915 + def translate_responses_chunk_to_openai_stream( parsed_chunk: Union[dict, BaseModel], ) -> "ModelResponseStream": """ diff --git a/litellm/constants.py b/litellm/constants.py index b51d15b6d25..a3ea68c7949 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -510,6 +510,8 @@ DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE = os.getenv( "DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE", "streaming.chunk.yield" ) +LITELLM_HTTP_STATUS_CLIENT_DISCONNECTED = 499 + EMAIL_BUDGET_ALERT_TTL = int( os.getenv("EMAIL_BUDGET_ALERT_TTL", 24 * 60 * 60) ) # 24 hours in seconds diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index e934c6a6f83..27a146df7bf 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -94,6 +94,7 @@ from litellm.types.utils import ( LlmProviders, LlmProvidersSet, ModelInfo, + ServiceTier, StandardBuiltInToolsParams, TranscriptionUsageDurationObject, TranscriptionUsageTokensObject, @@ -288,7 +289,7 @@ def _transcription_usage_has_token_details( return (prompt_tokens_val > 0) or (completion_tokens_val > 0) -def cost_per_token( # noqa: PLR0915 +def cost_per_token( model: str = "", prompt_tokens: int = 0, completion_tokens: int = 0, @@ -614,7 +615,9 @@ def cost_per_token( # noqa: PLR0915 service_tier=service_tier, ) elif custom_llm_provider == "anthropic": - return anthropic_cost_per_token(model=model, usage=usage_block) + return anthropic_cost_per_token( + model=model, usage=usage_block, service_tier=service_tier + ) elif custom_llm_provider == "bedrock": return bedrock_cost_per_token( model=model, usage=usage_block, service_tier=service_tier @@ -885,6 +888,23 @@ def _map_traffic_type_to_service_tier(traffic_type: Optional[str]) -> Optional[s return service_tier +def _normalize_service_tier(service_tier: object) -> str | None: + """ + Reduce a service_tier value to a concrete billable tier string or None. + + "auto" is a routing preference and any non-string value is not a billable + tier, so both defer to standard pricing (or to the tier the provider reports + on the response usage) instead of crashing the downstream cost-key lookup, + which calls service_tier.lower() + """ + if ( + not isinstance(service_tier, str) + or service_tier.lower() == ServiceTier.AUTO.value + ): + return None + return service_tier + + def _get_usage_object( completion_response: Any, ) -> Optional[Usage]: @@ -1136,7 +1156,7 @@ def _store_cost_breakdown_in_logging_obj( pass -def completion_cost( # noqa: PLR0915 +def completion_cost( completion_response=None, model: Optional[str] = None, prompt="", @@ -1224,6 +1244,8 @@ def completion_cost( # noqa: PLR0915 if service_tier is None and optional_params is not None: service_tier = optional_params.get("service_tier") + service_tier = _normalize_service_tier(service_tier) + # Extract service_tier from completion_response if not provided if service_tier is None and completion_response is not None: if isinstance(completion_response, BaseModel): @@ -1231,6 +1253,8 @@ def completion_cost( # noqa: PLR0915 elif isinstance(completion_response, dict): service_tier = completion_response.get("service_tier") + service_tier = _normalize_service_tier(service_tier) + # Extract service_tier from usage object if not provided if service_tier is None and cost_per_token_usage_object is not None: if isinstance(cost_per_token_usage_object, BaseModel): @@ -1240,6 +1264,8 @@ def completion_cost( # noqa: PLR0915 elif isinstance(cost_per_token_usage_object, dict): service_tier = cost_per_token_usage_object.get("service_tier") + service_tier = _normalize_service_tier(service_tier) + selected_model = _select_model_name_for_cost_calc( model=model, completion_response=completion_response, diff --git a/litellm/images/main.py b/litellm/images/main.py index d95b7287d20..8b108ded4c9 100644 --- a/litellm/images/main.py +++ b/litellm/images/main.py @@ -195,7 +195,7 @@ def image_generation( @client -def image_generation( # noqa: PLR0915 +def image_generation( prompt: str, model: Optional[str] = None, n: Optional[int] = None, @@ -738,7 +738,7 @@ def image_variation( @client -def image_edit( # noqa: PLR0915 +def image_edit( image: Optional[Union[FileTypes, List[FileTypes]]] = None, prompt: Optional[str] = None, model: Optional[str] = None, diff --git a/litellm/integrations/SlackAlerting/slack_alerting.py b/litellm/integrations/SlackAlerting/slack_alerting.py index e7be004e62e..2108ebae312 100644 --- a/litellm/integrations/SlackAlerting/slack_alerting.py +++ b/litellm/integrations/SlackAlerting/slack_alerting.py @@ -351,7 +351,7 @@ class SlackAlerting(CustomBatchLogger): except Exception: return 0 - async def send_daily_reports(self, router) -> bool: # noqa: PLR0915 + async def send_daily_reports(self, router) -> bool: """ Send a daily report on: - Top 5 deployments with most failed requests @@ -1373,7 +1373,7 @@ Model Info: return False - async def send_alert( # noqa: PLR0915 + async def send_alert( self, message: str, level: Literal["Low", "Medium", "High"], diff --git a/litellm/integrations/braintrust_logging.py b/litellm/integrations/braintrust_logging.py index 9b1c5077882..6a6313f72e1 100644 --- a/litellm/integrations/braintrust_logging.py +++ b/litellm/integrations/braintrust_logging.py @@ -133,9 +133,7 @@ class BraintrustLogger(CustomLogger): self.default_project_id = project_dict["id"] - def log_success_event( # noqa: PLR0915 - self, kwargs, response_obj, start_time, end_time - ): + def log_success_event(self, kwargs, response_obj, start_time, end_time): verbose_logger.debug("REACHES BRAINTRUST SUCCESS") try: litellm_call_id = kwargs.get("litellm_call_id") @@ -271,9 +269,7 @@ class BraintrustLogger(CustomLogger): except Exception as e: raise e # don't use verbose_logger.exception, if exception is raised - async def async_log_success_event( # noqa: PLR0915 - self, kwargs, response_obj, start_time, end_time - ): + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): verbose_logger.debug("REACHES BRAINTRUST SUCCESS") try: litellm_call_id = kwargs.get("litellm_call_id") diff --git a/litellm/integrations/langfuse/langfuse.py b/litellm/integrations/langfuse/langfuse.py index 0efc7d66876..b1c6956a16c 100644 --- a/litellm/integrations/langfuse/langfuse.py +++ b/litellm/integrations/langfuse/langfuse.py @@ -549,7 +549,7 @@ class LangFuseLogger: ) ) - def _log_langfuse_v2( # noqa: PLR0915 + def _log_langfuse_v2( self, user_id: Optional[str], metadata: dict, diff --git a/litellm/integrations/mock_client_factory.py b/litellm/integrations/mock_client_factory.py index 02a927fe64f..9b912ce70c8 100644 --- a/litellm/integrations/mock_client_factory.py +++ b/litellm/integrations/mock_client_factory.py @@ -107,7 +107,7 @@ def _is_url_match(url, matchers: List[str]) -> bool: return False -def create_mock_client_factory(config: MockClientConfig): # noqa: PLR0915 +def create_mock_client_factory(config: MockClientConfig): """ Factory function that creates mock client functions based on configuration. diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index fc37b6a34d8..6b50ef49b49 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -2198,9 +2198,7 @@ class OpenTelemetry(OTELGenAISemconvMixin, CustomLogger): return kv_pairs - def set_attributes( # noqa: PLR0915 - self, span: Span, kwargs, response_obj: Optional[Any] - ): + def set_attributes(self, span: Span, kwargs, response_obj: Optional[Any]): try: if self.callback_name == "langtrace": from litellm.integrations.langtrace import LangtraceAttributes diff --git a/litellm/integrations/otel/model/config.py b/litellm/integrations/otel/model/config.py index ca46182bc66..4f7c3277ebb 100644 --- a/litellm/integrations/otel/model/config.py +++ b/litellm/integrations/otel/model/config.py @@ -184,6 +184,20 @@ class OpenTelemetryV2Config(BaseSettings): ), ) + @field_validator("capture_message_content", mode="before") + @classmethod + def _normalize_capture_message_content(cls, value: object) -> object: + """Fold the capture mode to its canonical lower_snake_case form. + + V1 read this env var case-insensitively, so operators set the + UPPER_SNAKE_CASE form (e.g. ``SPAN_AND_EVENT``). The canonical values + here are lower_snake_case; normalizing at the boundary keeps both + spellings working and lets every downstream comparison stay exact. + """ + if isinstance(value, str): + return value.lower() + return value + @field_validator( "baggage_promoted_keys", "baggage_metadata_keys", diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 2119527a8e5..c63f114514a 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -75,7 +75,7 @@ class PrometheusLogger(CustomLogger): return cb return None - def __init__( # noqa: PLR0915 + def __init__( self, **kwargs, ): @@ -2255,7 +2255,7 @@ class PrometheusLogger(CustomLogger): or _litellm_params_metadata.get("user_agent"), } - def set_llm_deployment_failure_metrics(self, request_kwargs: dict): # noqa: PLR0915 + def set_llm_deployment_failure_metrics(self, request_kwargs: dict): """ Sets Failure metrics when an LLM API call fails diff --git a/litellm/litellm_core_utils/exception_mapping_utils.py b/litellm/litellm_core_utils/exception_mapping_utils.py index 0d35da9fa1a..edb97b310d7 100644 --- a/litellm/litellm_core_utils/exception_mapping_utils.py +++ b/litellm/litellm_core_utils/exception_mapping_utils.py @@ -170,6 +170,16 @@ def get_error_message(error_obj) -> Optional[str]: ####### EXCEPTION MAPPING ################ +def _get_body_error_code(error_str: str) -> int | None: + """Return error.code from a JSON error body, or None if not parseable.""" + try: + body = json.loads(error_str) + code = body.get("error", {}).get("code") + return int(code) if code is not None else None + except Exception: + return None + + def _get_response_headers(original_exception: Exception) -> Optional[httpx.Headers]: """ Extract and return the response headers from an exception, if present. @@ -234,7 +244,7 @@ def extract_and_raise_litellm_exception( ) -def exception_type( # type: ignore # noqa: PLR0915 +def exception_type( # type: ignore model, original_exception, custom_llm_provider, @@ -1415,6 +1425,29 @@ def exception_type( # type: ignore # noqa: PLR0915 ), ), ) + elif ( + isinstance(getattr(original_exception, "status_code", None), int) + and 500 <= original_exception.status_code < 600 + and _get_body_error_code(error_str) == 429 + ): + # upstream gateway wraps a 429 inside a 5xx envelope + # e.g. HTTP 500/503 with {"error":{"code":429,...}}. + # Scoped to 5xx so HTTP 400/401 with body code:429 + # still maps to BadRequestError / AuthenticationError. + exception_mapping_worked = True + raise RateLimitError( + message=f"litellm.RateLimitError: {custom_llm_provider}Exception - {error_str}", + model=model, + llm_provider=custom_llm_provider, + litellm_debug_info=extra_information, + response=httpx.Response( + status_code=429, + request=httpx.Request( + method="POST", + url=" https://cloud.google.com/vertex-ai/", + ), + ), + ) elif ( "500 Internal Server Error" in error_str or "The model is overloaded." in error_str diff --git a/litellm/litellm_core_utils/get_llm_provider_logic.py b/litellm/litellm_core_utils/get_llm_provider_logic.py index b3a48769d6a..625f8416517 100644 --- a/litellm/litellm_core_utils/get_llm_provider_logic.py +++ b/litellm/litellm_core_utils/get_llm_provider_logic.py @@ -154,7 +154,7 @@ def handle_anthropic_text_model_custom_llm_provider( return model, custom_llm_provider -def get_llm_provider( # noqa: PLR0915 +def get_llm_provider( model: str, custom_llm_provider: Optional[str] = None, api_base: Optional[str] = None, @@ -573,7 +573,7 @@ def get_llm_provider( # noqa: PLR0915 ) -def _get_openai_compatible_provider_info( # noqa: PLR0915 +def _get_openai_compatible_provider_info( model: str, api_base: Optional[str], api_key: Optional[str], diff --git a/litellm/litellm_core_utils/get_supported_openai_params.py b/litellm/litellm_core_utils/get_supported_openai_params.py index 65c238344e9..e87042b9101 100644 --- a/litellm/litellm_core_utils/get_supported_openai_params.py +++ b/litellm/litellm_core_utils/get_supported_openai_params.py @@ -5,7 +5,7 @@ from litellm.exceptions import BadRequestError from litellm.types.utils import LlmProviders, LlmProvidersSet -def get_supported_openai_params( # noqa: PLR0915 +def get_supported_openai_params( model: str, custom_llm_provider: Optional[str] = None, request_type: Literal[ diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 8af3bfa9c0a..347efbcbc97 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -986,7 +986,7 @@ class Logging(LiteLLMLoggingBaseClass): self._get_masked_api_base(additional_args.get("api_base", "")) ) - def pre_call(self, input, api_key, model=None, additional_args={}): # noqa: PLR0915 + def pre_call(self, input, api_key, model=None, additional_args={}): # Log the exact input to the LLM API litellm.error_logs["PRE_CALL"] = locals() try: @@ -2124,7 +2124,7 @@ class Logging(LiteLLMLoggingBaseClass): await self.async_success_handler(result=complete_streaming_response) return - def success_handler( # noqa: PLR0915 + def success_handler( self, result=None, start_time=None, end_time=None, cache_hit=None, **kwargs ): verbose_logger.debug( @@ -2589,7 +2589,7 @@ class Logging(LiteLLMLoggingBaseClass): ), ) - async def async_success_handler( # noqa: PLR0915 + async def async_success_handler( self, result=None, start_time=None, end_time=None, cache_hit=None, **kwargs ): """ @@ -3041,7 +3041,7 @@ class Logging(LiteLLMLoggingBaseClass): kwargs=self.model_call_details, ) # type: ignore - def failure_handler( # noqa: PLR0915 + def failure_handler( self, exception, traceback_exception, start_time=None, end_time=None ): verbose_logger.debug( @@ -3758,7 +3758,7 @@ def _get_masked_values( } -def set_callbacks(callback_list, function_id=None): # noqa: PLR0915 +def set_callbacks(callback_list, function_id=None): """ Globally sets the callback client """ @@ -3859,7 +3859,7 @@ def set_callbacks(callback_list, function_id=None): # noqa: PLR0915 return None -def _init_custom_logger_compatible_class( # noqa: PLR0915 +def _init_custom_logger_compatible_class( logging_integration: _custom_logger_compatible_callbacks_literal, internal_usage_cache: Optional[DualCache], llm_router: Optional[ @@ -4616,7 +4616,7 @@ def _maybe_auto_initialize_arize_phoenix(_in_memory_loggers: list) -> None: ) -def get_custom_logger_compatible_class( # noqa: PLR0915 +def get_custom_logger_compatible_class( logging_integration: _custom_logger_compatible_callbacks_literal, ) -> Optional[CustomLogger]: try: @@ -5421,19 +5421,20 @@ class StandardLoggingPayloadSetup: tb_lines[:MAXIMUM_TRACEBACK_LINES_TO_LOG] ) # Limit to first 100 lines + # Prefer the `.message` attribute (set by ProxyException and every + # litellm.exceptions.* class) over str(exc); ProxyException does not + # call super().__init__() nor define __str__, so str() on it returns + # an empty string, which used to silently strip the human-readable + # message from spend_logs.metadata.error_information. + # Use isinstance, not truthiness: an explicit empty string on + # `.message` is a deliberate value and must not be replaced by + # `str(exc)`. explicit_message = getattr(original_exception, "message", None) - error_message = ( - explicit_message - if isinstance(explicit_message, str) and explicit_message - else str(original_exception) - ) + if isinstance(explicit_message, str): + error_message = explicit_message + else: + error_message = str(original_exception) if original_exception else "" - # Duck-typed read so bare-Exception subclasses like - # `litellm.BudgetExceededError` can participate without joining the - # RateLimitError hierarchy (which would break `except BudgetExceededError`). - # Validated against the enum value sets so a third-party exception that - # happens to declare a `.category` or `.rate_limit_type` string attribute - # can't leak garbage into the payload or Prometheus label cardinality. rate_limit_category = validate_rate_limit_category( getattr(original_exception, "category", None) ) @@ -5446,11 +5447,42 @@ class StandardLoggingPayloadSetup: error_class=error_class, llm_provider=_llm_provider_in_exception, traceback=traceback_info, - error_message=error_message if original_exception else "", + error_message=error_message, error_rate_limit_category=rate_limit_category, error_rate_limit_type=rate_limit_type, ) + @staticmethod + def get_error_information_for_logging_payload( + metadata: dict, + original_exception: Exception | None, + error_str: str | None, + ) -> tuple[StandardLoggingPayloadErrorInformation, str | None]: + error_information = StandardLoggingPayloadSetup.get_error_information( + original_exception=original_exception, + ) + if not metadata.get("client_disconnected"): + return error_information, error_str + + client_disconnect_error = metadata.get("error_information") + if isinstance(client_disconnect_error, dict): + error_information = cast( + StandardLoggingPayloadErrorInformation, + client_disconnect_error, + ) + else: + error_information = cast( + StandardLoggingPayloadErrorInformation, + { + "error_code": "499", + "error_message": "Client disconnected the request", + "error_class": "ClientDisconnected", + }, + ) + if not error_str: + error_str = "Client disconnected the request" + return error_information, error_str + @staticmethod def get_response_time( start_time_float: float, @@ -5778,8 +5810,12 @@ def get_standard_logging_object_payload( api_base=litellm_params.get("api_base"), ) - error_information = StandardLoggingPayloadSetup.get_error_information( - original_exception=original_exception, + error_information, error_str = ( + StandardLoggingPayloadSetup.get_error_information_for_logging_payload( + metadata=metadata, + original_exception=original_exception, + error_str=error_str, + ) ) ## get final response object ## diff --git a/litellm/litellm_core_utils/llm_cost_calc/utils.py b/litellm/litellm_core_utils/llm_cost_calc/utils.py index d75850984a9..7a7fde3087e 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/utils.py +++ b/litellm/litellm_core_utils/llm_cost_calc/utils.py @@ -191,6 +191,11 @@ def _get_service_tier_cost_key(base_key: str, service_tier: Optional[str]) -> st return base_key +def _parse_above_token_threshold(key: str) -> float: + threshold_str = key.split("_above_")[1].split("_tokens")[0] + return float(threshold_str.replace("k", "")) * (1000 if "k" in threshold_str else 1) + + def _get_token_base_cost( model_info: ModelInfo, usage: Usage, service_tier: Optional[str] = None ) -> Tuple[float, float, float, float, float]: @@ -256,15 +261,13 @@ def _get_token_base_cost( # Only sort the threshold keys (typically 1-2 keys instead of 66+) threshold: Optional[float] = None - for key in sorted(threshold_keys, reverse=True): + for key in sorted(threshold_keys, key=_parse_above_token_threshold, reverse=True): value = model_info.get(key) if value is not None: try: # Handle both formats: _above_128k_tokens and _above_128_tokens threshold_str = key.split("_above_")[1].split("_tokens")[0] - threshold = float(threshold_str.replace("k", "")) * ( - 1000 if "k" in threshold_str else 1 - ) + threshold = _parse_above_token_threshold(key) if usage.prompt_tokens > threshold: # Prefer a service_tier-specific above-threshold key when available, # e.g. input_cost_per_token_priority_above_200k_tokens for Gemini @@ -303,40 +306,54 @@ def _get_token_base_cost( # Apply tiered pricing to cache costs cache_creation_tiered_key = ( - f"cache_creation_input_token_cost_above_{threshold_str}_tokens" + _get_service_tier_cost_key( + f"cache_creation_input_token_cost_above_{threshold_str}_tokens", + service_tier, + ) + if service_tier + else f"cache_creation_input_token_cost_above_{threshold_str}_tokens" + ) + cache_creation_1hr_tiered_key = ( + _get_service_tier_cost_key( + f"cache_creation_input_token_cost_above_1hr_above_{threshold_str}_tokens", + service_tier, + ) + if service_tier + else f"cache_creation_input_token_cost_above_1hr_above_{threshold_str}_tokens" ) - cache_creation_1hr_tiered_key = f"cache_creation_input_token_cost_above_1hr_above_{threshold_str}_tokens" cache_read_tiered_key = ( - f"cache_read_input_token_cost_above_{threshold_str}_tokens" + _get_service_tier_cost_key( + f"cache_read_input_token_cost_above_{threshold_str}_tokens", + service_tier, + ) + if service_tier + else f"cache_read_input_token_cost_above_{threshold_str}_tokens" ) - if cache_creation_tiered_key in model_info: - cache_creation_cost = cast( - float, - _get_cost_per_unit( - model_info, - cache_creation_tiered_key, - cache_creation_cost, - ), - ) + cache_creation_cost = cast( + float, + _get_cost_per_unit( + model_info, + cache_creation_tiered_key, + cache_creation_cost, + ), + ) - if cache_creation_1hr_tiered_key in model_info: - cache_creation_cost_above_1hr = cast( - float, - _get_cost_per_unit( - model_info, - cache_creation_1hr_tiered_key, - cache_creation_cost_above_1hr, - ), - ) + cache_creation_cost_above_1hr = cast( + float, + _get_cost_per_unit( + model_info, + cache_creation_1hr_tiered_key, + cache_creation_cost_above_1hr, + ), + ) - if cache_read_tiered_key in model_info: - cache_read_cost = cast( - float, - _get_cost_per_unit( - model_info, cache_read_tiered_key, cache_read_cost - ), - ) + cache_read_cost = cast( + float, + _get_cost_per_unit( + model_info, cache_read_tiered_key, cache_read_cost + ), + ) break except (IndexError, ValueError): @@ -683,7 +700,7 @@ def _get_regional_uplift_multiplier( return 1.0 -def generic_cost_per_token( # noqa: PLR0915 +def generic_cost_per_token( model: str, usage: Usage, custom_llm_provider: str, diff --git a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py index 4e5b53a13d7..016bb6b1e22 100644 --- a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py +++ b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py @@ -471,7 +471,7 @@ def _should_convert_tool_call_to_json_mode( return False -def convert_to_model_response_object( # noqa: PLR0915 +def convert_to_model_response_object( response_object: Optional[dict] = None, model_response_object: Optional[ Union[ diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 5059e612f2f..b95b73398ac 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -1475,7 +1475,7 @@ def convert_to_gemini_tool_call_invoke( ) -def convert_to_gemini_tool_call_result( # noqa: PLR0915 +def convert_to_gemini_tool_call_result( message: Union[ChatCompletionToolMessage, ChatCompletionFunctionMessage], last_message_with_tool_calls: Optional[dict], model: Optional[str] = None, @@ -2227,7 +2227,7 @@ def _sanitize_empty_text_content( return message -def _add_missing_tool_results( # noqa: PLR0915 +def _add_missing_tool_results( current_message: AllMessageValues, messages: List[AllMessageValues], current_index: int, @@ -2484,7 +2484,7 @@ def sanitize_messages_for_tool_calling( return sanitized_messages -def anthropic_messages_pt( # noqa: PLR0915 +def anthropic_messages_pt( messages: List[AllMessageValues], model: str, llm_provider: str, @@ -3278,7 +3278,7 @@ def convert_to_cohere_tool_invoke(tool_calls: list) -> List[ToolCallObject]: return cohere_tool_invoke -def cohere_messages_pt_v2( # noqa: PLR0915 +def cohere_messages_pt_v2( messages: List, model: str, llm_provider: str, @@ -4703,7 +4703,7 @@ class BedrockConverseMessagesProcessor: return messages @staticmethod - async def _bedrock_converse_messages_pt_async( # noqa: PLR0915 + async def _bedrock_converse_messages_pt_async( messages: List, model: str, llm_provider: str, @@ -5133,7 +5133,7 @@ class BedrockConverseMessagesProcessor: return assistant_parts -def _bedrock_converse_messages_pt( # noqa: PLR0915 +def _bedrock_converse_messages_pt( messages: List, model: str, llm_provider: str, diff --git a/litellm/litellm_core_utils/realtime_streaming.py b/litellm/litellm_core_utils/realtime_streaming.py index c8f87d96e2f..c56a70177bf 100644 --- a/litellm/litellm_core_utils/realtime_streaming.py +++ b/litellm/litellm_core_utils/realtime_streaming.py @@ -1198,7 +1198,7 @@ class RealTimeStreaming: item["content"] = new_content return item - async def client_ack_messages(self): # noqa: PLR0915 + async def client_ack_messages(self): try: while True: message = await self.websocket.receive_text() diff --git a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py index d51b937d434..04f6b1241c3 100644 --- a/litellm/litellm_core_utils/streaming_chunk_builder_utils.py +++ b/litellm/litellm_core_utils/streaming_chunk_builder_utils.py @@ -209,7 +209,7 @@ class ChunkProcessor: ) return response - def get_combined_tool_content( # noqa: PLR0915 + def get_combined_tool_content( self, tool_call_chunks: List[Dict[str, Any]] ) -> List[ChatCompletionMessageToolCall]: tool_calls_list: List[ChatCompletionMessageToolCall] = [] diff --git a/litellm/litellm_core_utils/streaming_handler.py b/litellm/litellm_core_utils/streaming_handler.py index 7e4bf895a79..888a9658396 100644 --- a/litellm/litellm_core_utils/streaming_handler.py +++ b/litellm/litellm_core_utils/streaming_handler.py @@ -967,7 +967,7 @@ class CustomStreamWrapper: delta, model_response.choices[0].delta, attribute ) - def return_processed_chunk_logic( # noqa: PLR0915, C901 + def return_processed_chunk_logic( # noqa: C901 self, completion_obj: Dict[str, Any], model_response: ModelResponseStream, @@ -1145,7 +1145,7 @@ class CustomStreamWrapper: del model_response.choices[0].delta.reasoning_content return - def chunk_creator(self, chunk: Any): # type: ignore # noqa: PLR0915 + def chunk_creator(self, chunk: Any): # type: ignore if hasattr(chunk, "id"): self.response_id = chunk.id model_response = self.model_response_creator() @@ -1887,7 +1887,7 @@ class CustomStreamWrapper: model_response.choices[0].finish_reason = "tool_calls" return model_response - def __next__(self) -> "ModelResponseStream": # noqa: PLR0915 + def __next__(self) -> "ModelResponseStream": cache_hit = False if ( self.custom_llm_provider is not None @@ -2077,7 +2077,7 @@ class CustomStreamWrapper: return self.completion_stream - async def __anext__(self) -> "ModelResponseStream": # noqa: PLR0915 + async def __anext__(self) -> "ModelResponseStream": cache_hit = False if ( self.custom_llm_provider is not None diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index 74b41062174..c766c6edec1 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -744,6 +744,17 @@ def _count_content_list( thinking_text = str(c.get("thinking", "")) if thinking_text: num_tokens += count_function(thinking_text) + elif c["type"] == "tool_reference": + # Anthropic tool-search reference block: a lightweight pointer to + # a deferred tool, e.g. {"type": "tool_reference", "tool_name": ...}. + # The full tool definition is counted via the `tools` param, so we + # only count the referenced name here. Without this branch, + # token_counter raises on tool-search traffic; on the streaming + # anthropic_messages path that nulls response_cost and causes the + # proxy to drop the SpendLogs row entirely (silent cost undercount). + tool_name = str(c.get("tool_name") or "") + if tool_name: + num_tokens += count_function(tool_name) else: content_type = ( c.get("type", type(c).__name__) @@ -752,7 +763,7 @@ def _count_content_list( ) raise ValueError( f"Invalid content item type: {content_type}. " - f"Expected str or dict with 'type' field (text, image_url, tool_use, tool_result, thinking)." + f"Expected str or dict with 'type' field (text, image_url, tool_use, tool_result, thinking, tool_reference)." ) return num_tokens except Exception as e: diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index 2fb29b32a61..5d14f3cc4ae 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -772,7 +772,7 @@ class ModelResponseIterator: ) return results - def chunk_parser(self, chunk: dict) -> ModelResponseStream: # noqa: PLR0915 + def chunk_parser(self, chunk: dict) -> ModelResponseStream: try: type_chunk = chunk.get("type", "") or "" diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index e8c1e659e9f..c24c990f356 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -605,7 +605,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ) return _tool_choice - def _map_tool_helper( # noqa: PLR0915 + def _map_tool_helper( self, tool: ChatCompletionToolParam, ) -> Tuple[Optional[AllAnthropicToolsValues], Optional[AnthropicMcpServerTool]]: @@ -1399,7 +1399,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): return None - def map_openai_params( # noqa: PLR0915 + def map_openai_params( self, non_default_params: dict, optional_params: dict, @@ -2213,6 +2213,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): inference_geo: Optional[str] = None if "inference_geo" in _usage and _usage["inference_geo"] is not None: inference_geo = _usage["inference_geo"] + service_tier = cast( + str | None, + _usage.get("service_tier"), + ) iterations: Optional[List[Any]] = _usage.get("iterations") if iterations: @@ -2324,6 +2328,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): ), inference_geo=inference_geo, speed=speed, + service_tier=service_tier, ) return usage diff --git a/litellm/llms/anthropic/cost_calculation.py b/litellm/llms/anthropic/cost_calculation.py index 6a031498dae..44081ea9e79 100644 --- a/litellm/llms/anthropic/cost_calculation.py +++ b/litellm/llms/anthropic/cost_calculation.py @@ -18,7 +18,9 @@ if TYPE_CHECKING: import litellm -def _compute_cache_only_cost(model_info: "ModelInfo", usage: "Usage") -> float: +def _compute_cache_only_cost( + model_info: "ModelInfo", usage: "Usage", service_tier: str | None = None +) -> float: """ Return only the cache-related portion of the prompt cost (cache read + cache write). @@ -36,7 +38,9 @@ def _compute_cache_only_cost(model_info: "ModelInfo", usage: "Usage") -> float: cache_creation_cost, cache_creation_cost_above_1hr, cache_read_cost, - ) = _get_token_base_cost(model_info=model_info, usage=usage) + ) = _get_token_base_cost( + model_info=model_info, usage=usage, service_tier=service_tier + ) cache_cost = float(prompt_tokens_details["cache_hit_tokens"]) * cache_read_cost @@ -56,19 +60,26 @@ def _compute_cache_only_cost(model_info: "ModelInfo", usage: "Usage") -> float: return cache_cost -def cost_per_token(model: str, usage: "Usage") -> Tuple[float, float]: +def cost_per_token( + model: str, usage: "Usage", service_tier: str | None = None +) -> Tuple[float, float]: """ Calculates the cost per token for a given model, prompt tokens, and completion tokens. Input: - model: str, the model name without provider prefix - usage: LiteLLM Usage block, containing anthropic caching information + - service_tier: the service tier the request was served at (e.g. "priority"), + read from the Anthropic response usage and used to select tier-specific pricing Returns: Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd """ prompt_cost, completion_cost = generic_cost_per_token( - model=model, usage=usage, custom_llm_provider="anthropic" + model=model, + usage=usage, + custom_llm_provider="anthropic", + service_tier=service_tier, ) # Apply provider_specific_entry multipliers for geo/speed routing @@ -89,7 +100,9 @@ def cost_per_token(model: str, usage: "Usage") -> Tuple[float, float]: multiplier *= provider_specific_entry.get("fast", 1.0) if multiplier != 1.0: - cache_cost = _compute_cache_only_cost(model_info=model_info, usage=usage) + cache_cost = _compute_cache_only_cost( + model_info=model_info, usage=usage, service_tier=service_tier + ) prompt_cost = (prompt_cost - cache_cost) * multiplier + cache_cost completion_cost *= multiplier except Exception: diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py index f049abcf47f..a8e2fceb4ee 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py @@ -372,7 +372,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): cache_read_input_tokens=0, ) - def __next__(self): # noqa: PLR0915 + def __next__(self): from .transformation import LiteLLMAnthropicMessagesAdapter try: @@ -618,7 +618,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): ) raise StopIteration - async def __anext__(self): # noqa: PLR0915 + async def __anext__(self): from .transformation import LiteLLMAnthropicMessagesAdapter try: diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index 150f056dc81..bf425637b56 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -332,7 +332,14 @@ class LiteLLMAnthropicMessagesAdapter: if isinstance(source, dict) else getattr(source, "cache_control", None) ) - if cache_control and model and self.is_anthropic_claude_model(model): + if ( + cache_control + and model + and ( + self.is_anthropic_claude_model(model) + or self.is_bedrock_arn_model(model) + ) + ): # TypedDict objects support dict operations at runtime # Use type ignore consistent with codebase pattern (see anthropic/chat/transformation.py:432) if isinstance(target, dict): @@ -376,7 +383,7 @@ class LiteLLMAnthropicMessagesAdapter: isinstance(tool_type, str) and tool_type.startswith("web_search") ) or tool_name == "web_search" - def translate_anthropic_messages_to_openai( # noqa: PLR0915 + def translate_anthropic_messages_to_openai( self, messages: List[ Union[ @@ -752,6 +759,20 @@ class LiteLLMAnthropicMessagesAdapter: model_lower = model.lower() return "anthropic" in model_lower or "claude" in model_lower + @staticmethod + def is_bedrock_arn_model(model: str) -> bool: + """ + Check if the model string is a Bedrock ARN, such as an Application + Inference Profile (e.g. arn:aws:bedrock:us-east-1:123:application-inference-profile/id). + + These ARNs contain neither "anthropic" nor "claude", so is_anthropic_claude_model + cannot identify them even though, on the /v1/messages endpoint, they point at Claude. + Match ":bedrock:" in the ARN service field so another service's ARN that merely names + bedrock in a resource (arn:aws:sagemaker:.../my-bedrock-endpoint) is not matched. + """ + model_lower = model.lower() + return "arn:" in model_lower and ":bedrock:" in model_lower + @staticmethod def translate_thinking_for_model( thinking: Dict[str, Any], diff --git a/litellm/llms/anthropic/experimental_pass_through/context_management/editors/compact.py b/litellm/llms/anthropic/experimental_pass_through/context_management/editors/compact.py index 4aae85b17fe..6479ee999b0 100644 --- a/litellm/llms/anthropic/experimental_pass_through/context_management/editors/compact.py +++ b/litellm/llms/anthropic/experimental_pass_through/context_management/editors/compact.py @@ -97,7 +97,7 @@ def _read_summary_max_tokens_setting() -> int: return COMPACT_SUMMARY_MAX_TOKENS -async def _check_summary_model_access( # noqa: PLR0915 +async def _check_summary_model_access( user_api_key_auth: Any, summary_model: str, llm_router: Any, @@ -970,7 +970,7 @@ def apply_client_compaction_block_history( ) -async def apply_compact_20260112( # noqa: PLR0915 +async def apply_compact_20260112( *, model: str, messages: List[Dict[str, Any]], diff --git a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py index 5f1362e259f..04819a416a2 100644 --- a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/streaming_iterator.py @@ -66,7 +66,7 @@ class AnthropicResponsesStreamWrapper: self._current_block_index += 1 return self._current_block_index - def _process_event(self, event: Any) -> None: # noqa: PLR0915 + def _process_event(self, event: Any) -> None: """Convert one Responses API event into zero or more Anthropic chunks queued for emission.""" event_type = getattr(event, "type", None) if event_type is None and isinstance(event, dict): diff --git a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py index 2badc2a3276..4fb1ddf5c46 100644 --- a/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/responses_adapters/transformation.py @@ -51,7 +51,7 @@ class LiteLLMAnthropicToResponsesAPIAdapter: return source.get("url") return None - def translate_messages_to_responses_input( # noqa: PLR0915 + def translate_messages_to_responses_input( self, messages: List[ Union[ diff --git a/litellm/llms/bedrock/chat/agentcore/transformation.py b/litellm/llms/bedrock/chat/agentcore/transformation.py index 44ba1ce3c86..3cd3a249c33 100644 --- a/litellm/llms/bedrock/chat/agentcore/transformation.py +++ b/litellm/llms/bedrock/chat/agentcore/transformation.py @@ -218,8 +218,20 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): - Qualifier goes as query parameter - Only the payload goes in the request body + Payload shape: + - ``prompt`` is always present and contains the text-only flatten of the + last message's content (existing behavior). + - ``content`` is added ONLY when the ``forward_multimodal_content`` litellm + param is truthy AND the last message's ``content`` is a list containing a + non-text block (e.g. ``image_url``, ``file``, ``input_audio``). The list is + forwarded verbatim so the agent's ``@app.entrypoint`` handler can parse the + OpenAI-shaped multimodal blocks. This is opt-in because an AgentCore agent + must be explicitly written to read ``payload["content"]``; by default the + payload stays byte-identical to the legacy ``{"prompt": "..."}`` shape. + Returns: - dict: Payload dict containing the prompt + dict: Payload dict containing the prompt and (optionally) the OpenAI + content list. """ verbose_logger.debug( f"AgentCore transform_request - optional_params keys: {list(optional_params.keys())}" @@ -231,6 +243,20 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): # Create the payload - this is what goes in the body (raw JSON) payload: dict = {"prompt": prompt} + # Opt-in: when forward_multimodal_content is set, forward the OpenAI content + # list verbatim under "content" so an attachment-aware agent can read the raw + # blocks (image_url, file, etc.). Default off keeps the payload byte-identical + # to the legacy {"prompt": "..."} shape for agents that only read the prompt. + if self._should_forward_multimodal_content(optional_params, litellm_params): + last_content = messages[-1].get("content") + if isinstance(last_content, list) and any( + isinstance(block, dict) and block.get("type") not in (None, "text") + for block in last_content + ): + # Copy so the payload never aliases messages[-1]["content"]; shallow, + # not deep, to avoid cloning large base64 media on the request path. + payload["content"] = list(last_content) + # Get or generate session ID - this goes in the header runtime_session_id = self._get_runtime_session_id(optional_params) headers["X-Amzn-Bedrock-AgentCore-Runtime-Session-Id"] = runtime_session_id @@ -246,6 +272,29 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): verbose_logger.debug(f"PAYLOAD: {payload}") return payload + @staticmethod + def _should_forward_multimodal_content( + optional_params: dict, litellm_params: dict + ) -> bool: + """Whether to forward raw OpenAI content blocks under ``payload["content"]``. + + Opt-in via the ``forward_multimodal_content`` litellm param (default ``False``) + because AgentCore agents must be explicitly written to read the field. The + value may arrive as a bool or a config/env string ("true", "1", ...). Checks + ``optional_params`` first (where other AgentCore params land), then + ``litellm_params``. + """ + for source in (optional_params, litellm_params): + if not isinstance(source, dict): + continue + value = source.get("forward_multimodal_content") + if value is None: + continue + if isinstance(value, str): + return value.strip().lower() in ("1", "true", "yes", "on") + return bool(value) + return False + def _extract_sse_json(self, line: str) -> Optional[Dict]: """Extract and parse JSON from an SSE data line.""" if not line.startswith("data:"): diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 7e1020000f4..7b1064ccef9 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -248,7 +248,7 @@ class BedrockConverseLLM(BaseAWSLLM): encoding=encoding, ) - def completion( # noqa: PLR0915 + def completion( self, model: str, messages: list, diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index b5e5e4de6fc..bb261ec85b2 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -2189,7 +2189,7 @@ class AmazonConverseConfig(BaseConfig): real_tools = [t for i, t in enumerate(tools) if i not in json_tool_indices] return real_tools if real_tools else None - def _transform_response( # noqa: PLR0915 + def _transform_response( self, model: str, response: httpx.Response, diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index 0a1322a751e..75b560b4d6d 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -473,7 +473,7 @@ class BedrockLLM(BaseAWSLLM): prompt += f"{message['content']}" return prompt, chat_history # type: ignore - def process_response( # noqa: PLR0915 + def process_response( self, model: str, response: httpx.Response, @@ -765,7 +765,7 @@ class BedrockLLM(BaseAWSLLM): return model_response - def completion( # noqa: PLR0915 + def completion( self, model: str, messages: list, diff --git a/litellm/llms/bedrock/embed/embedding.py b/litellm/llms/bedrock/embed/embedding.py index 27dc785bf57..b6aa99842d7 100644 --- a/litellm/llms/bedrock/embed/embedding.py +++ b/litellm/llms/bedrock/embed/embedding.py @@ -388,7 +388,7 @@ class BedrockEmbedding(BaseAWSLLM): batch_data=batch_data, ) - def embeddings( # noqa: PLR0915 + def embeddings( self, model: str, input: List[str], diff --git a/litellm/llms/bedrock/image_edit/stability_transformation.py b/litellm/llms/bedrock/image_edit/stability_transformation.py index 2d73e47003d..d00d62a8530 100644 --- a/litellm/llms/bedrock/image_edit/stability_transformation.py +++ b/litellm/llms/bedrock/image_edit/stability_transformation.py @@ -149,7 +149,7 @@ class BedrockStabilityImageEditConfig(BaseImageEditConfig): return mapped_params - def transform_image_edit_request( # noqa: PLR0915 + def transform_image_edit_request( self, model: str, prompt: Optional[str], diff --git a/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py b/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py index 2d6bdb5298a..0522bb249e1 100644 --- a/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py +++ b/litellm/llms/bedrock/passthrough/guardrail_translation/handler.py @@ -226,7 +226,7 @@ class BedrockPassthroughGuardrailHandler(BaseTranslation): return _is_converse_endpoint(endpoint) @staticmethod - async def de_anonymize_event_stream( # noqa: PLR0915 + async def de_anonymize_event_stream( body_bytes: bytes, proxy_logging_obj: "ProxyLogging", user_api_key_dict: "UserAPIKeyAuth", diff --git a/litellm/llms/bedrock_mantle/chat/transformation.py b/litellm/llms/bedrock_mantle/chat/transformation.py index 18f051f8524..1504e89c58e 100644 --- a/litellm/llms/bedrock_mantle/chat/transformation.py +++ b/litellm/llms/bedrock_mantle/chat/transformation.py @@ -4,8 +4,10 @@ Amazon Bedrock Mantle - OpenAI-compatible inference engine in Amazon Bedrock. API docs: https://docs.aws.amazon.com/bedrock/latest/userguide/bedrock-mantle.html Base URL: https://bedrock-mantle.{region}.api.aws/v1 -Auth: AWS Bedrock API key as Bearer token (set via BEDROCK_MANTLE_API_KEY env var) - or region-aware key via BEDROCK_MANTLE_{REGION}_API_KEY. +Auth: Bearer token (litellm_params.api_key, BEDROCK_MANTLE_API_KEY, or the + standard AWS_BEARER_TOKEN_BEDROCK) when present; otherwise AWS SigV4 + (service "bedrock") over the standard credential chain. See + BedrockMantleAuthMixin in common_utils. """ from typing import Iterator, AsyncIterator, Any, List, Optional, Tuple, Union @@ -13,20 +15,26 @@ from typing import Iterator, AsyncIterator, Any, List, Optional, Tuple, Union import litellm from litellm._logging import verbose_logger from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM +from litellm.llms.bedrock_mantle.common_utils import ( + BEDROCK_MANTLE_DEFAULT_REGION, + BedrockMantleAuthMixin, +) from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import AllMessageValues from litellm.types.router import GenericLiteLLMParams from ...openai_like.chat.transformation import OpenAILikeChatConfig -BEDROCK_MANTLE_DEFAULT_REGION = "us-east-1" - -class BedrockMantleChatConfig(OpenAILikeChatConfig): +class BedrockMantleChatConfig(BedrockMantleAuthMixin, OpenAILikeChatConfig): """ Transformation config for Amazon Bedrock Mantle OpenAI-compatible API. """ + def __init__(self, aws_signer: BaseAWSLLM | None = None): + super().__init__() + self._aws_signer = aws_signer or BaseAWSLLM() + @property def custom_llm_provider(self) -> Optional[str]: return "bedrock_mantle" @@ -54,7 +62,7 @@ class BedrockMantleChatConfig(OpenAILikeChatConfig): or get_secret_str("BEDROCK_MANTLE_API_BASE") or f"https://bedrock-mantle.{region}.api.aws/v1" ) - dynamic_api_key = api_key or get_secret_str("BEDROCK_MANTLE_API_KEY") + dynamic_api_key = self._resolve_bearer_token(api_key) return api_base, dynamic_api_key def validate_environment( diff --git a/litellm/llms/bedrock_mantle/common_utils.py b/litellm/llms/bedrock_mantle/common_utils.py new file mode 100644 index 00000000000..8c092f345d9 --- /dev/null +++ b/litellm/llms/bedrock_mantle/common_utils.py @@ -0,0 +1,115 @@ +""" +Shared auth and region resolution for the Amazon Bedrock Mantle backends. + +Mantle authenticates with a Bearer token when one is available +(litellm_params.api_key, BEDROCK_MANTLE_API_KEY, or the standard +AWS_BEARER_TOKEN_BEDROCK); otherwise it falls back to AWS SigV4 (service +"bedrock") over the standard credential chain (IAM role / access key / profile / +web identity). The Chat Completions and Responses backends share this behaviour +through BedrockMantleAuthMixin so the two paths can never drift apart. +""" + +import re +from typing import Tuple + +from botocore.exceptions import ( + CredentialRetrievalError, + NoCredentialsError, + PartialCredentialsError, + ProfileNotFound, +) + +from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM +from litellm.secret_managers.main import get_secret_str + +BEDROCK_MANTLE_DEFAULT_REGION = "us-east-1" + +# Standard Mantle host: https://bedrock-mantle..api.aws (group 1 = region). +MANTLE_HOST_RE = re.compile( + r"^https?://bedrock-mantle\.([^/.]+)\.api\.aws", re.IGNORECASE +) + + +class BedrockMantleAuthMixin: + _aws_signer: BaseAWSLLM + + @staticmethod + def _resolve_bearer_token(api_key: str | None) -> str | None: + return ( + api_key + or get_secret_str("BEDROCK_MANTLE_API_KEY") + or get_secret_str("AWS_BEARER_TOKEN_BEDROCK") + ) + + @staticmethod + def _resolve_region(params: dict) -> str: + region = params.get("aws_region_name") + if region: + BaseAWSLLM._validate_aws_region_name(region) + return region + base = params.get("api_base") or get_secret_str("BEDROCK_MANTLE_API_BASE") + if base: + match = MANTLE_HOST_RE.match(base.rstrip("/")) + if match: + return match.group(1) + return ( + get_secret_str("BEDROCK_MANTLE_REGION") + or get_secret_str("AWS_REGION_NAME") + or get_secret_str("AWS_REGION") + or BEDROCK_MANTLE_DEFAULT_REGION + ) + + def sign_request( + self, + headers: dict, + optional_params: dict, + request_data: dict, + api_base: str, + api_key: str | None = None, + model: str | None = None, + stream: bool | None = None, + fake_stream: bool | None = None, + ) -> Tuple[dict, bytes | None]: + bearer = self._resolve_bearer_token(api_key) + if not bearer: + # SigV4 path. Pin the credential-scope region to the region of the actual + # signing URL so the SigV4 scope and the URL host can never disagree, even + # when a stale api_base and aws_region_name point at different regions. + # Fall back to _resolve_region only for custom proxy hosts that do not + # match the standard Mantle URL pattern. Also drop any caller Authorization + # so _sign_request's restore-original-Authorization step cannot override + # the SigV4 header. + host_match = MANTLE_HOST_RE.match(api_base.rstrip("/")) + optional_params = { + **optional_params, + "aws_region_name": ( + host_match.group(1) + if host_match + else self._resolve_region({**optional_params, "api_base": api_base}) + ), + } + headers = {k: v for k, v in headers.items() if k.lower() != "authorization"} + try: + return self._aws_signer._sign_request( + service_name="bedrock", + headers=headers, + optional_params=optional_params, + request_data=request_data, + api_base=api_base, + api_key=bearer, + model=model, + stream=stream, + fake_stream=fake_stream, + ) + except ( + NoCredentialsError, + PartialCredentialsError, + ProfileNotFound, + CredentialRetrievalError, + ) as e: + raise ValueError( + "Bedrock Mantle auth failed: no Bearer token and no usable AWS " + "credentials. Set BEDROCK_MANTLE_API_KEY (or AWS_BEARER_TOKEN_BEDROCK) " + "or pass api_key for Bearer auth, or provide AWS credentials " + "(IAM role / access key / profile / web identity) for SigV4." + ) from e diff --git a/litellm/llms/bedrock_mantle/responses/transformation.py b/litellm/llms/bedrock_mantle/responses/transformation.py index b409666a967..2e30f85fd0e 100644 --- a/litellm/llms/bedrock_mantle/responses/transformation.py +++ b/litellm/llms/bedrock_mantle/responses/transformation.py @@ -15,26 +15,20 @@ role / access key / profile / web identity), signed via the shared BaseAWSLLM._sign_request after the request body is finalized. """ -import re -from typing import Any, Dict, List, Optional, Tuple - -from botocore.exceptions import ( - CredentialRetrievalError, - NoCredentialsError, - PartialCredentialsError, - ProfileNotFound, -) +from typing import Any, Dict, List, Optional from litellm._logging import verbose_logger from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM +from litellm.llms.bedrock_mantle.common_utils import ( + MANTLE_HOST_RE, + BedrockMantleAuthMixin, +) from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import LlmProviders -BEDROCK_MANTLE_DEFAULT_REGION = "us-east-1" - # Checked longest/most-specific first so a full endpoint URL collapses to host # in one pass and the appended path never doubles. _BASE_SUFFIXES_TO_STRIP = ( @@ -45,18 +39,13 @@ _BASE_SUFFIXES_TO_STRIP = ( "/v1", ) -# Standard Mantle host: https://bedrock-mantle..api.aws (group 1 = region). -_MANTLE_HOST_RE = re.compile( - r"^https?://bedrock-mantle\.([^/.]+)\.api\.aws", re.IGNORECASE -) - # Per Bedrock Mantle Responses API validation errors. _BEDROCK_MANTLE_SUPPORTED_RESPONSE_TOOL_TYPES = frozenset( {"function", "mcp", "custom", "namespace", "tool_search"} ) -class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig): +class BedrockMantleResponsesAPIConfig(BedrockMantleAuthMixin, OpenAIResponsesAPIConfig): def __init__( self, aws_signer: Optional[BaseAWSLLM] = None, @@ -70,24 +59,6 @@ class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig): def custom_llm_provider(self) -> LlmProviders: return LlmProviders.BEDROCK_MANTLE - @staticmethod - def _resolve_region(params: dict) -> str: - region = params.get("aws_region_name") - if region: - BaseAWSLLM._validate_aws_region_name(region) - return region - base = params.get("api_base") or get_secret_str("BEDROCK_MANTLE_API_BASE") - if base: - match = _MANTLE_HOST_RE.match(base.rstrip("/")) - if match: - return match.group(1) - return ( - get_secret_str("BEDROCK_MANTLE_REGION") - or get_secret_str("AWS_REGION_NAME") - or get_secret_str("AWS_REGION") - or BEDROCK_MANTLE_DEFAULT_REGION - ) - def get_complete_url( self, api_base: Optional[str], @@ -107,7 +78,7 @@ class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig): # For the standard Mantle host (including the default-region base that # responses/main.py auto-injects into litellm_params.api_base), pin to the # single resolved region so aws_region_name wins; preserve custom proxy hosts. - if _MANTLE_HOST_RE.match(base): + if MANTLE_HOST_RE.match(base): base = f"https://bedrock-mantle.{region}.api.aws" path = "/openai/v1/responses" if self.use_openai_path else "/v1/responses" return f"{base}{path}" @@ -116,13 +87,9 @@ class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig): self, headers: dict, model: str, litellm_params: Optional[GenericLiteLLMParams] ) -> dict: litellm_params = litellm_params or GenericLiteLLMParams() - api_key = ( - litellm_params.api_key - or get_secret_str("BEDROCK_MANTLE_API_KEY") - or get_secret_str("AWS_BEARER_TOKEN_BEDROCK") - ) - if api_key: - headers["Authorization"] = f"Bearer {api_key}" + bearer = self._resolve_bearer_token(litellm_params.api_key) + if bearer: + headers["Authorization"] = f"Bearer {bearer}" if litellm_params.aws_bedrock_project_id: headers["OpenAI-Project"] = litellm_params.aws_bedrock_project_id return headers @@ -182,58 +149,3 @@ class BedrockMantleResponsesAPIConfig(OpenAIResponsesAPIConfig): params.pop("tools", None) return params - - def sign_request( - self, - headers: dict, - optional_params: dict, - request_data: dict, - api_base: str, - api_key: Optional[str] = None, - model: Optional[str] = None, - stream: Optional[bool] = None, - fake_stream: Optional[bool] = None, - ) -> Tuple[dict, Optional[bytes]]: - bearer = ( - api_key - or get_secret_str("BEDROCK_MANTLE_API_KEY") - or get_secret_str("AWS_BEARER_TOKEN_BEDROCK") - ) - if not bearer: - # SigV4 path. Pin the credential-scope region to the region of the actual - # signing URL (api_base, already region-resolved by get_complete_url) so the - # SigV4 scope and the URL host can never disagree. Resolve from api_base first, - # then fall back to the regular precedence. Also drop any caller Authorization - # so _sign_request's restore-original-Authorization step cannot override the - # SigV4 header. - optional_params = { - **optional_params, - "aws_region_name": self._resolve_region( - {**optional_params, "api_base": api_base} - ), - } - headers = {k: v for k, v in headers.items() if k.lower() != "authorization"} - try: - return self._aws_signer._sign_request( - service_name="bedrock", - headers=headers, - optional_params=optional_params, - request_data=request_data, - api_base=api_base, - api_key=bearer, - model=model, - stream=stream, - fake_stream=fake_stream, - ) - except ( - NoCredentialsError, - PartialCredentialsError, - ProfileNotFound, - CredentialRetrievalError, - ) as e: - raise ValueError( - "Bedrock Mantle auth failed: no Bearer token and no usable AWS " - "credentials. Set BEDROCK_MANTLE_API_KEY (or AWS_BEARER_TOKEN_BEDROCK) " - "or pass api_key for Bearer auth, or provide AWS credentials " - "(IAM role / access key / profile / web identity) for SigV4." - ) from e diff --git a/litellm/llms/custom_httpx/aiohttp_transport.py b/litellm/llms/custom_httpx/aiohttp_transport.py index 62f707b3622..b97a59a93a6 100644 --- a/litellm/llms/custom_httpx/aiohttp_transport.py +++ b/litellm/llms/custom_httpx/aiohttp_transport.py @@ -116,6 +116,16 @@ class AiohttpResponseStream(httpx.AsyncByteStream): # For other exceptions, use the normal mapping with map_aiohttp_exceptions(): raise + finally: + # Release the aiohttp connection when iteration ends for any + # reason (read timeout, cancellation from a client disconnect, + # GeneratorExit). Without this, abnormally terminated streams + # permanently hold a slot in the TCPConnector pool; once the + # pool is exhausted every request to that host times out (408) + # until the proxy is restarted, even after the backend recovers. + # On a fully-read response the connection was already released + # at EOF and close() is a no-op. + self._aiohttp_response.close() async def aclose(self) -> None: with map_aiohttp_exceptions(): diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 5575385fb28..8ac5b47c6e7 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -5676,7 +5676,7 @@ class BaseLLMHTTPHandler: ) raise - async def async_responses_websocket( # noqa: PLR0915 + async def async_responses_websocket( self, model: str, websocket: Any, diff --git a/litellm/llms/fireworks_ai/chat/transformation.py b/litellm/llms/fireworks_ai/chat/transformation.py index cca3b3da37a..341c2fc7350 100644 --- a/litellm/llms/fireworks_ai/chat/transformation.py +++ b/litellm/llms/fireworks_ai/chat/transformation.py @@ -392,7 +392,10 @@ class FireworksAIConfig(OpenAIGPTConfig): headers: dict, ) -> dict: if not model.startswith("accounts/") and "#" not in model: - model = f"accounts/fireworks/models/{model}" + if model.endswith("-fast"): + model = f"accounts/fireworks/routers/{model}" + else: + model = f"accounts/fireworks/models/{model}" messages = self._transform_messages_helper( messages=messages, model=model, litellm_params=litellm_params ) diff --git a/litellm/llms/gemini/realtime/transformation.py b/litellm/llms/gemini/realtime/transformation.py index 51fa395d899..74f6cd4d831 100644 --- a/litellm/llms/gemini/realtime/transformation.py +++ b/litellm/llms/gemini/realtime/transformation.py @@ -1378,7 +1378,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): raise ValueError(f"Unknown openai event: {key}, value: {value}") return openai_event - def transform_realtime_response( # noqa: PLR0915 + def transform_realtime_response( self, message: Union[str, bytes], model: str, diff --git a/litellm/llms/hosted_vllm/chat/transformation.py b/litellm/llms/hosted_vllm/chat/transformation.py index 1824314865c..7c42e6a9a00 100644 --- a/litellm/llms/hosted_vllm/chat/transformation.py +++ b/litellm/llms/hosted_vllm/chat/transformation.py @@ -2,6 +2,7 @@ Translate from OpenAI's `/v1/chat/completions` to VLLM's `/v1/chat/completions` """ +import json from typing import ( Any, Coroutine, @@ -22,7 +23,9 @@ from litellm.litellm_core_utils.prompt_templates.factory import _parse_mime_type from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import ( AllMessageValues, + ChatCompletionAssistantToolCall, ChatCompletionFileObject, + ChatCompletionToolCallFunctionChunk, ChatCompletionVideoObject, ChatCompletionVideoUrlObject, ) @@ -101,26 +104,18 @@ class HostedVLLMChatConfig(OpenAIGPTConfig): ) -> dict: _tools = non_default_params.pop("tools", None) if _tools is not None: - # remove 'additionalProperties' from tools _tools = _remove_additional_properties(_tools) - # remove 'strict' from tools _tools = _remove_strict_from_schema(_tools) if isinstance(_tools, list): _tools = self._convert_custom_tools_to_function_tools(_tools) if _tools is not None: non_default_params["tools"] = _tools - # Handle thinking parameter - convert Anthropic-style to OpenAI-style reasoning_effort - # vLLM is OpenAI-compatible, so it understands reasoning_effort, not thinking - # Reference: https://github.com/BerriAI/litellm/issues/19761 thinking = non_default_params.pop("thinking", None) if thinking is not None and isinstance(thinking, dict): if thinking.get("type") == "enabled": - # Only convert if reasoning_effort not already set if "reasoning_effort" not in non_default_params: budget_tokens = thinking.get("budget_tokens", 0) - # Map budget_tokens to reasoning_effort level - # Same logic as Anthropic adapter (translate_anthropic_thinking_to_reasoning_effort) if budget_tokens >= 10000: non_default_params["reasoning_effort"] = "high" elif budget_tokens >= 5000: @@ -137,20 +132,13 @@ class HostedVLLMChatConfig(OpenAIGPTConfig): def _get_openai_compatible_provider_info( self, api_base: Optional[str], api_key: Optional[str] ) -> Tuple[Optional[str], Optional[str]]: - api_base = api_base or get_secret_str("HOSTED_VLLM_API_BASE") # type: ignore + api_base = api_base or get_secret_str("HOSTED_VLLM_API_BASE") dynamic_api_key = ( api_key or get_secret_str("HOSTED_VLLM_API_KEY") or "fake-api-key" - ) # vllm does not require an api key + ) return api_base, dynamic_api_key def _is_video_file(self, content_item: ChatCompletionFileObject) -> bool: - """ - Check if the file is a video - - - format: video/ - - file_data: base64 encoded video data - - file_id: infer mp4 from extension - """ file = content_item.get("file", {}) format = file.get("format") file_data = file.get("file_data") @@ -205,29 +193,69 @@ class HostedVLLMChatConfig(OpenAIGPTConfig): """ Support translating: - video files from file_id or file_data to video_url - - thinking_blocks on assistant messages to content blocks + - thinking_blocks on assistant messages are removed, and content lists + are converted to strings for vLLM compatibility """ for message in messages: if message["role"] == "assistant": - thinking_blocks = message.pop("thinking_blocks", None) # type: ignore - if thinking_blocks: - new_content: list = [ - ( - { - "type": block["type"], - "thinking": block.get("thinking", ""), + message.pop("thinking_blocks", None) + existing_content = message.get("content") + if isinstance(existing_content, list): + text_parts = [] + tool_calls: list[ChatCompletionAssistantToolCall] = [] + content_blocks: list[object] = [] + has_structured_content = False + for c in existing_content: + if isinstance(c, dict) and c.get("type") == "text": + text_parts.append(c.get("text", "")) + content_blocks.append(c) + elif isinstance(c, dict) and c.get("type") == "tool_use": + tool_input = c.get("input", {}) + tool_calls.append( + ChatCompletionAssistantToolCall( + id=c.get("id"), + type="function", + function=ChatCompletionToolCallFunctionChunk( + name=c.get("name"), + arguments=( + tool_input + if isinstance( + tool_input, + str, + ) + else json.dumps(tool_input) + ), + ), + ) + ) + else: + content_blocks.append(c) + has_structured_content = True + if tool_calls: + existing_tool_calls = message.get("tool_calls") + if isinstance(existing_tool_calls, list): + existing_tool_call_ids = { + tool_call.get("id") + for tool_call in existing_tool_calls + if isinstance(tool_call, dict) + and tool_call.get("id") is not None } - if block.get("type") == "thinking" - else {"type": block["type"], "data": block.get("data", "")} - ) - for block in thinking_blocks - ] - existing_content = message.get("content") - if isinstance(existing_content, str): - new_content.append({"type": "text", "text": existing_content}) - elif isinstance(existing_content, list): - new_content.extend(existing_content) - message["content"] = new_content # type: ignore + new_tool_calls = [ + tool_call + for tool_call in tool_calls + if tool_call.get("id") not in existing_tool_call_ids + ] + if new_tool_calls: + message["tool_calls"] = ( + existing_tool_calls + new_tool_calls + ) + else: + message["tool_calls"] = tool_calls + content_str = "\n".join(text_parts) + new_content = ( + content_blocks if has_structured_content else content_str + ) + message["content"] = new_content # type: ignore[typeddict-item] elif message["role"] == "user": message_content = message.get("content") if message_content and isinstance(message_content, list): @@ -243,6 +271,7 @@ class HostedVLLMChatConfig(OpenAIGPTConfig): message_content[idx] = self._convert_file_to_video_url( content_item ) + if is_async: return super()._transform_messages( messages, model, is_async=cast(Literal[True], True) diff --git a/litellm/llms/huggingface/embedding/transformation.py b/litellm/llms/huggingface/embedding/transformation.py index 88d42cfcdcc..7cddda617a9 100644 --- a/litellm/llms/huggingface/embedding/transformation.py +++ b/litellm/llms/huggingface/embedding/transformation.py @@ -404,7 +404,7 @@ class HuggingFaceEmbeddingConfig(BaseConfig): ) return completion_response - def convert_to_model_response_object( # noqa: PLR0915 + def convert_to_model_response_object( self, completion_response: Union[List[Dict[str, Any]], Dict[str, Any]], model_response: ModelResponse, diff --git a/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py index 5464b5bb7ee..b8b750b8c12 100644 --- a/litellm/llms/openai/chat/gpt_transformation.py +++ b/litellm/llms/openai/chat/gpt_transformation.py @@ -18,6 +18,9 @@ from typing import ( overload, ) +import os +from urllib.parse import urlparse + import httpx import litellm @@ -426,6 +429,32 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): ) return messages, tools + def _should_preserve_cache_control_for_endpoint( + self, + custom_llm_provider: str | None, + api_base: str | None, + ) -> bool: + """ + The generic `openai` provider also reaches OpenAI-compatible endpoints + (a LiteLLM proxy, vLLM, an Anthropic-compatible gateway) via a custom + api_base. Those can understand cache_control, so it must survive there. + Real OpenAI cannot, so it is still stripped for an openai.com host. + """ + if custom_llm_provider != "openai": + return False + resolved_api_base = ( + api_base + or litellm.api_base + or os.getenv("OPENAI_BASE_URL") + or os.getenv("OPENAI_API_BASE") + ) + if not resolved_api_base: + return False + hostname = urlparse(resolved_api_base).hostname + if hostname is None: + return False + return hostname != "openai.com" and not hostname.endswith(".openai.com") + def transform_request( self, model: str, @@ -441,11 +470,14 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): dict: The transformed request. Sent as the body of the API call. """ messages = self._transform_messages(messages=messages, model=model) - messages, tools = self.remove_cache_control_flag_from_messages_and_tools( - model=model, messages=messages, tools=optional_params.get("tools", []) - ) - if tools is not None and len(tools) > 0: - optional_params["tools"] = tools + if not self._should_preserve_cache_control_for_endpoint( + litellm_params.get("custom_llm_provider"), litellm_params.get("api_base") + ): + messages, tools = self.remove_cache_control_flag_from_messages_and_tools( + model=model, messages=messages, tools=optional_params.get("tools", []) + ) + if tools is not None and len(tools) > 0: + optional_params["tools"] = tools optional_params.pop("max_retries", None) @@ -466,16 +498,19 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): transformed_messages = await self._transform_messages( messages=messages, model=model, is_async=True ) - ( - transformed_messages, - tools, - ) = self.remove_cache_control_flag_from_messages_and_tools( - model=model, - messages=transformed_messages, - tools=optional_params.get("tools", []), - ) - if tools is not None and len(tools) > 0: - optional_params["tools"] = tools + if not self._should_preserve_cache_control_for_endpoint( + litellm_params.get("custom_llm_provider"), litellm_params.get("api_base") + ): + ( + transformed_messages, + tools, + ) = self.remove_cache_control_flag_from_messages_and_tools( + model=model, + messages=transformed_messages, + tools=optional_params.get("tools", []), + ) + if tools is not None and len(tools) > 0: + optional_params["tools"] = tools if self.__class__._is_base_class: return { "model": model, diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index 194f29648c4..ea905d8ebca 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -608,7 +608,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): return streaming_response - def completion( # type: ignore # noqa: PLR0915 + def completion( # type: ignore self, model_response: ModelResponse, timeout: Union[float, httpx.Timeout], diff --git a/litellm/llms/openai/transcriptions/whisper_transformation.py b/litellm/llms/openai/transcriptions/whisper_transformation.py index fa507e1bc26..2c01156fe05 100644 --- a/litellm/llms/openai/transcriptions/whisper_transformation.py +++ b/litellm/llms/openai/transcriptions/whisper_transformation.py @@ -1,3 +1,4 @@ +import json from typing import List, Optional, Union from httpx import Headers, Response @@ -107,9 +108,7 @@ class OpenAIWhisperAudioTranscriptionConfig(BaseAudioTranscriptionConfig): """ data = {"model": model, "file": audio_file, **optional_params} - if "response_format" not in data or ( - data["response_format"] == "text" or data["response_format"] == "json" - ): + if "response_format" not in data: data["response_format"] = ( "verbose_json" # ensures 'duration' is received - used for cost calculation ) @@ -133,10 +132,11 @@ class OpenAIWhisperAudioTranscriptionConfig(BaseAudioTranscriptionConfig): ) -> TranscriptionResponse: try: raw_response_json = raw_response.json() - except Exception as e: - raise ValueError( - f"Error transforming response to json: {str(e)}\nResponse: {raw_response.text}" - ) + except json.JSONDecodeError: + content_type = raw_response.headers.get("content-type", "").lower() + if "application/json" in content_type: + raise + return TranscriptionResponse(text=raw_response.text) if any( key in raw_response_json diff --git a/litellm/llms/openrouter/chat/transformation.py b/litellm/llms/openrouter/chat/transformation.py index 0d7850e8c74..107d5c25e6d 100644 --- a/litellm/llms/openrouter/chat/transformation.py +++ b/litellm/llms/openrouter/chat/transformation.py @@ -50,11 +50,15 @@ class OpenrouterConfig(OpenAIGPTConfig): def map_openai_params( self, - non_default_params: dict, + non_default_params: dict[str, object], optional_params: dict, model: str, drop_params: bool, ) -> dict: + # OpenRouter expects "xhigh" instead of "max" for reasoning_effort. + if non_default_params.get("reasoning_effort") == "max": + non_default_params = {**non_default_params, "reasoning_effort": "xhigh"} + mapped_openai_params = super().map_openai_params( non_default_params, optional_params, model, drop_params ) diff --git a/litellm/llms/perplexity/cost_calculator.py b/litellm/llms/perplexity/cost_calculator.py index 0f9c3cad841..bf055f91aa0 100644 --- a/litellm/llms/perplexity/cost_calculator.py +++ b/litellm/llms/perplexity/cost_calculator.py @@ -58,11 +58,8 @@ def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]: ## CALCULATE OUTPUT COST output_cost_per_token = _safe_float_cast(model_info.get("output_cost_per_token")) - completion_cost: float = (usage.completion_tokens or 0) * output_cost_per_token - ## ADD REASONING TOKENS COST (if present) reasoning_tokens = getattr(usage, "reasoning_tokens", 0) or 0 - # Also check completion_tokens_details if reasoning_tokens is not directly available if ( reasoning_tokens == 0 and hasattr(usage, "completion_tokens_details") @@ -73,9 +70,19 @@ def cost_per_token(model: str, usage: Usage) -> Tuple[float, float]: ) reasoning_cost_value = model_info.get("output_cost_per_reasoning_token") + + # `completion_tokens` includes `reasoning_tokens` per the OpenAI/Perplexity usage + # convention (codified for the central path in PR #18607). When a reasoning rate is + # configured we subtract before the output-rate multiplication so the reasoning + # tokens are not billed twice. if reasoning_tokens > 0 and reasoning_cost_value is not None: - reasoning_cost_per_token = _safe_float_cast(reasoning_cost_value) - completion_cost += reasoning_tokens * reasoning_cost_per_token + non_reasoning_completion_tokens = max( + 0, (usage.completion_tokens or 0) - reasoning_tokens + ) + completion_cost: float = non_reasoning_completion_tokens * output_cost_per_token + completion_cost += reasoning_tokens * _safe_float_cast(reasoning_cost_value) + else: + completion_cost = (usage.completion_tokens or 0) * output_cost_per_token ## ADD SEARCH QUERIES COST (if present) num_search_queries = 0 diff --git a/litellm/llms/predibase/chat/transformation.py b/litellm/llms/predibase/chat/transformation.py index 3d251d24b0d..ce004f60bfc 100644 --- a/litellm/llms/predibase/chat/transformation.py +++ b/litellm/llms/predibase/chat/transformation.py @@ -129,7 +129,7 @@ class PredibaseConfig(BaseConfig): optional_params["response_format"] = value return optional_params - def transform_response( # noqa: PLR0915 + def transform_response( self, model: str, raw_response: Response, diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index c578d6cd28b..f5a2b268263 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -678,7 +678,7 @@ def check_if_part_exists_in_parts( return False -def _gemini_convert_messages_with_history( # noqa: PLR0915 +def _gemini_convert_messages_with_history( messages: List[AllMessageValues], model: Optional[str] = None, litellm_params: Optional[dict] = None, @@ -1176,7 +1176,7 @@ def _rewrite_google_maps_response_format(data: RequestBody) -> None: _rewrite_mime_type_to_response_format(generation_config) -def _transform_request_body( # noqa: PLR0915 +def _transform_request_body( messages: List[AllMessageValues], model: str, optional_params: dict, diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 3ec7b0814dd..c171538b9c0 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -614,9 +614,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return googleSearch, googleSearchRetrieval, enterpriseWebSearch, urlContext - def _map_function( # noqa: PLR0915 - self, value: List[dict], optional_params: dict - ) -> List[Tools]: + def _map_function(self, value: List[dict], optional_params: dict) -> List[Tools]: """ Map OpenAI-style tools/functions to Vertex AI format. @@ -1173,7 +1171,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): optional_params["include_server_side_tool_invocations"] = True return - def map_openai_params( # noqa: PLR0915 + def map_openai_params( self, non_default_params: Dict, optional_params: Dict, @@ -1904,7 +1902,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return False @staticmethod - def _calculate_usage( # noqa: PLR0915 + def _calculate_usage( completion_response: Union[ GenerateContentResponseBody, BidiGenerateContentServerMessage ], @@ -2380,7 +2378,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return annotations @staticmethod - def _process_candidates( # noqa: PLR0915 + def _process_candidates( _candidates: List[Candidates], model_response: Union[ModelResponse, "ModelResponseStream"], standard_optional_params: dict, @@ -2846,6 +2844,7 @@ async def make_call( sync_stream=False, logging_obj=logging_obj, response_headers=response.headers, + response=response, ) # LOGGING logging_obj.post_call( @@ -2889,6 +2888,7 @@ def make_sync_call( sync_stream=True, logging_obj=logging_obj, response_headers=response.headers, + response=response, ) # LOGGING @@ -3350,12 +3350,14 @@ class ModelResponseIterator: sync_stream: bool, logging_obj: LoggingClass, response_headers: Optional[Dict[str, str]] = None, + response: httpx.Response | None = None, ): from litellm.litellm_core_utils.prompt_templates.common_utils import ( check_is_function_call, ) self.streaming_response = streaming_response + self.response = response self.chunk_type: Literal["valid_json", "accumulated_json"] = "valid_json" self.accumulated_json = "" self.sent_first_chunk = False @@ -3657,3 +3659,41 @@ class ModelResponseIterator: raise StopAsyncIteration except ValueError as e: raise RuntimeError(f"Error parsing chunk: {e},\nReceived chunk: {chunk}") + + async def aclose(self) -> None: + iterator = getattr( + self, + "async_response_iterator", + self.streaming_response, + ) + if iterator is not None and hasattr(iterator, "aclose"): + try: + await iterator.aclose() + except Exception as e: # noqa: BLE001 + verbose_logger.debug( + "ModelResponseIterator.aclose: error closing iterator: %s", e + ) + if self.response is not None: + try: + await self.response.aclose() + except Exception as e: # noqa: BLE001 + verbose_logger.debug( + "ModelResponseIterator.aclose: error closing response: %s", e + ) + + def close(self) -> None: + iterator = getattr(self, "response_iterator", self.streaming_response) + if iterator is not None and hasattr(iterator, "close"): + try: + iterator.close() + except Exception as e: # noqa: BLE001 + verbose_logger.debug( + "ModelResponseIterator.close: error closing iterator: %s", e + ) + if self.response is not None: + try: + self.response.close() + except Exception as e: # noqa: BLE001 + verbose_logger.debug( + "ModelResponseIterator.close: error closing response: %s", e + ) diff --git a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py index 99165c37c93..165dac24903 100644 --- a/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py +++ b/litellm/llms/vertex_ai/gemini_embeddings/batch_embed_content_handler.py @@ -124,7 +124,7 @@ class GoogleBatchEmbeddings(VertexLLM): return resolved_files - def batch_embeddings( # noqa: PLR0915 + def batch_embeddings( self, model: str, input: GeminiEmbeddingInput, diff --git a/litellm/llms/vertex_ai/vertex_ai_non_gemini.py b/litellm/llms/vertex_ai/vertex_ai_non_gemini.py index 222820d7ee5..c134dee7ad4 100644 --- a/litellm/llms/vertex_ai/vertex_ai_non_gemini.py +++ b/litellm/llms/vertex_ai/vertex_ai_non_gemini.py @@ -77,7 +77,7 @@ def _set_client_in_cache(client_cache_key: str, vertex_llm_model: Any): ) -def completion( # noqa: PLR0915 +def completion( model: str, messages: list, model_response: ModelResponse, @@ -485,7 +485,7 @@ def completion( # noqa: PLR0915 ) -async def async_completion( # noqa: PLR0915 +async def async_completion( llm_model, mode: str, prompt: str, diff --git a/litellm/main.py b/litellm/main.py index ab83cbbac5f..5becb807092 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -1086,7 +1086,7 @@ def _build_custom_pricing_entry( @tracer.wrap() @client -def completion( # type: ignore # noqa: PLR0915 +def completion( # type: ignore model: str, # Optional OpenAI params: see https://platform.openai.com/docs/api-reference/chat/create messages: List = [], @@ -4881,7 +4881,7 @@ def embedding( @client -def embedding( # noqa: PLR0915 +def embedding( model, input=[], # Optional params @@ -6128,7 +6128,7 @@ async def atext_completion( @client -def text_completion( # noqa: PLR0915 +def text_completion( prompt: Union[ str, List[Union[str, List[Union[str, List[int]]]]] ], # Required: The prompt(s) to generate completions for. @@ -6667,7 +6667,7 @@ async def atranscription(*args, **kwargs) -> TranscriptionResponse: @client -def transcription( # noqa: PLR0915 +def transcription( model: str, file: FileTypes, ## OPTIONAL OPENAI PARAMS ## @@ -6974,7 +6974,7 @@ async def aspeech(*args, **kwargs) -> HttpxBinaryResponseContent: @client -def speech( # noqa: PLR0915 +def speech( model: str, input: str, voice: Optional[Union[str, dict]] = None, @@ -7437,22 +7437,7 @@ def speech( # noqa: PLR0915 async def ahealth_check( model_params: dict, - mode: Optional[ - Literal[ - "chat", - "completion", - "embedding", - "audio_speech", - "audio_transcription", - "image_generation", - "video_generation", - "batch", - "rerank", - "realtime", - "responses", - "ocr", - ] - ] = "chat", + mode: str | None = "chat", prompt: Optional[str] = None, input: Optional[List] = None, ): @@ -7665,7 +7650,7 @@ def stream_chunk_builder_text_completion( return TextCompletionResponse(**response) -def stream_chunk_builder( # noqa: PLR0915 +def stream_chunk_builder( chunks: list, messages: Optional[list] = None, start_time=None, diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index f563ad0c5b5..39d612f252d 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -2528,6 +2528,100 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "azure_ai/gpt-5.5": { + "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_above_272k_tokens": 1e-06, + "cache_read_input_token_cost_priority": 1e-06, + "cache_read_input_token_cost_above_272k_tokens_priority": 2e-06, + "input_cost_per_token": 5e-06, + "input_cost_per_token_above_272k_tokens": 1e-05, + "input_cost_per_token_priority": 1e-05, + "input_cost_per_token_above_272k_tokens_priority": 2e-05, + "litellm_provider": "azure_ai", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 3e-05, + "output_cost_per_token_above_272k_tokens": 4.5e-05, + "output_cost_per_token_priority": 6e-05, + "output_cost_per_token_above_272k_tokens_priority": 9e-05, + "source": "https://ai.azure.com/catalog/models/gpt-5.5", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_service_tier": true, + "supports_vision": true, + "supports_web_search": true, + "supports_none_reasoning_effort": true, + "supports_xhigh_reasoning_effort": true, + "supports_minimal_reasoning_effort": false + }, + "azure_ai/gpt-5.5-2026-04-23": { + "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_above_272k_tokens": 1e-06, + "cache_read_input_token_cost_priority": 1e-06, + "cache_read_input_token_cost_above_272k_tokens_priority": 2e-06, + "input_cost_per_token": 5e-06, + "input_cost_per_token_above_272k_tokens": 1e-05, + "input_cost_per_token_priority": 1e-05, + "input_cost_per_token_above_272k_tokens_priority": 2e-05, + "litellm_provider": "azure_ai", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 3e-05, + "output_cost_per_token_above_272k_tokens": 4.5e-05, + "output_cost_per_token_priority": 6e-05, + "output_cost_per_token_above_272k_tokens_priority": 9e-05, + "source": "https://ai.azure.com/catalog/models/gpt-5.5", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_service_tier": true, + "supports_vision": true, + "supports_web_search": true, + "supports_none_reasoning_effort": true, + "supports_xhigh_reasoning_effort": true, + "supports_minimal_reasoning_effort": false + }, "azure_ai/gpt-5.4": { "cache_read_input_token_cost": 2.5e-07, "cache_read_input_token_cost_above_272k_tokens": 5e-07, @@ -10068,6 +10162,8 @@ }, "claude-sonnet-4-5": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, + "cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, "input_cost_per_token_above_200k_tokens": 6e-06, @@ -10097,6 +10193,8 @@ }, "claude-sonnet-4-5-20250929": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, + "cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, "input_cost_per_token_above_200k_tokens": 6e-06, @@ -10127,6 +10225,7 @@ }, "claude-sonnet-4-6": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, "litellm_provider": "anthropic", @@ -10155,6 +10254,8 @@ }, "claude-sonnet-4-5-20250929-v1:0": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, + "cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, "input_cost_per_token_above_200k_tokens": 6e-06, @@ -10811,13 +10912,13 @@ "supports_tool_choice": true }, "command-r7b-12-2024": { - "input_cost_per_token": 1.5e-07, + "input_cost_per_token": 3.75e-08, "litellm_provider": "cohere_chat", "max_input_tokens": 128000, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_token": 3.75e-08, + "output_cost_per_token": 1.5e-07, "source": "https://docs.cohere.com/v2/docs/command-r7b", "supports_function_calling": true, "supports_tool_choice": true @@ -14511,6 +14612,38 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "fireworks_ai/accounts/fireworks/models/deepseek-v4-flash": { + "cache_read_input_token_cost": 2.8e-08, + "input_cost_per_token": 1.4e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 384000, + "max_tokens": 384000, + "mode": "chat", + "output_cost_per_token": 2.8e-07, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/accounts/fireworks/models/deepseek-v4-pro": { + "cache_read_input_token_cost": 1.45e-07, + "input_cost_per_token": 1.74e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 384000, + "max_tokens": 384000, + "mode": "chat", + "output_cost_per_token": 3.48e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, "fireworks_ai/accounts/fireworks/models/firefunction-v2": { "input_cost_per_token": 9e-07, "litellm_provider": "fireworks_ai", @@ -14586,43 +14719,64 @@ "input_cost_per_token": 1.4e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 202800, - "max_output_tokens": 202800, - "max_tokens": 202800, + "max_output_tokens": 131072, + "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 4.4e-06, - "source": "https://fireworks.ai/models/fireworks/glm-5p1", + "source": "https://docs.fireworks.ai/serverless/pricing", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/accounts/fireworks/models/glm-5p2": { + "cache_read_input_token_cost": 2.6e-07, + "input_cost_per_token": 1.4e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 4.4e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false }, "fireworks_ai/accounts/fireworks/models/gpt-oss-120b": { + "cache_read_input_token_cost": 1.5e-08, "input_cost_per_token": 1.5e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 131072, - "max_output_tokens": 131072, - "max_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 6e-07, - "source": "https://fireworks.ai/pricing", + "source": "https://docs.fireworks.ai/serverless/pricing", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "supports_vision": false }, "fireworks_ai/accounts/fireworks/models/gpt-oss-20b": { - "input_cost_per_token": 5e-08, + "cache_read_input_token_cost": 3.5e-08, + "input_cost_per_token": 7e-08, "litellm_provider": "fireworks_ai", "max_input_tokens": 131072, - "max_output_tokens": 131072, - "max_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, "mode": "chat", - "output_cost_per_token": 2e-07, - "source": "https://fireworks.ai/pricing", + "output_cost_per_token": 3e-07, + "source": "https://docs.fireworks.ai/serverless/pricing", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "supports_vision": false }, "fireworks_ai/accounts/fireworks/models/kimi-k2-instruct": { "input_cost_per_token": 6e-07, @@ -14678,6 +14832,38 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "fireworks_ai/accounts/fireworks/models/kimi-k2p6": { + "cache_read_input_token_cost": 1.6e-07, + "input_cost_per_token": 9.5e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 4e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "fireworks_ai/accounts/fireworks/models/kimi-k2p7-code": { + "cache_read_input_token_cost": 1.9e-07, + "input_cost_per_token": 9.5e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 4e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "fireworks_ai/accounts/fireworks/models/llama-v3p1-405b-instruct": { "input_cost_per_token": 3e-06, "litellm_provider": "fireworks_ai", @@ -14795,6 +14981,38 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "fireworks_ai/accounts/fireworks/models/minimax-m2p7": { + "cache_read_input_token_cost": 6e-08, + "input_cost_per_token": 3e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 196608, + "max_output_tokens": 196608, + "max_tokens": 196608, + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/accounts/fireworks/models/minimax-m3": { + "cache_read_input_token_cost": 6e-08, + "input_cost_per_token": 3e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 512000, + "max_output_tokens": 512000, + "max_tokens": 512000, + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, "fireworks_ai/accounts/fireworks/models/mixtral-8x22b-instruct-hf": { "input_cost_per_token": 1.2e-06, "litellm_provider": "fireworks_ai", @@ -14847,6 +15065,38 @@ "supports_response_schema": true, "supports_tool_choice": false }, + "fireworks_ai/deepseek-v4-flash": { + "cache_read_input_token_cost": 2.8e-08, + "input_cost_per_token": 1.4e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 384000, + "max_tokens": 384000, + "mode": "chat", + "output_cost_per_token": 2.8e-07, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/deepseek-v4-pro": { + "cache_read_input_token_cost": 1.45e-07, + "input_cost_per_token": 1.74e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 384000, + "max_tokens": 384000, + "mode": "chat", + "output_cost_per_token": 3.48e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, "fireworks_ai/glm-4p7": { "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 6e-07, @@ -14867,15 +15117,80 @@ "input_cost_per_token": 1.4e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 202800, - "max_output_tokens": 202800, - "max_tokens": 202800, + "max_output_tokens": 131072, + "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 4.4e-06, - "source": "https://fireworks.ai/models/fireworks/glm-5p1", + "source": "https://docs.fireworks.ai/serverless/pricing", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/glm-5p1-fast": { + "cache_read_input_token_cost": 5.2e-07, + "input_cost_per_token": 2.8e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 202800, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 8.8e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/glm-5p2": { + "cache_read_input_token_cost": 2.6e-07, + "input_cost_per_token": 1.4e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 4.4e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/gpt-oss-120b": { + "cache_read_input_token_cost": 1.5e-08, + "input_cost_per_token": 1.5e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 6e-07, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/gpt-oss-20b": { + "cache_read_input_token_cost": 3.5e-08, + "input_cost_per_token": 7e-08, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 3e-07, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false }, "fireworks_ai/kimi-k2p5": { "cache_read_input_token_cost": 1e-07, @@ -14891,6 +15206,70 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "fireworks_ai/kimi-k2p6": { + "cache_read_input_token_cost": 1.6e-07, + "input_cost_per_token": 9.5e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 4e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "fireworks_ai/kimi-k2p6-fast": { + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 2e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 8e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "fireworks_ai/kimi-k2p7-code": { + "cache_read_input_token_cost": 1.9e-07, + "input_cost_per_token": 9.5e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 4e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "fireworks_ai/kimi-k2p7-code-fast": { + "cache_read_input_token_cost": 3.8e-07, + "input_cost_per_token": 1.9e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 8e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "fireworks_ai/minimax-m2p1": { "cache_read_input_token_cost": 3e-08, "input_cost_per_token": 3e-07, @@ -14905,6 +15284,54 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "fireworks_ai/minimax-m2p7": { + "cache_read_input_token_cost": 6e-08, + "input_cost_per_token": 3e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 196608, + "max_output_tokens": 196608, + "max_tokens": 196608, + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/minimax-m3": { + "cache_read_input_token_cost": 6e-08, + "input_cost_per_token": 3e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 512000, + "max_output_tokens": 512000, + "max_tokens": 512000, + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/qwen3p7-plus": { + "cache_read_input_token_cost": 8e-08, + "input_cost_per_token": 4e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 262144, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_token": 1.6e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "fireworks_ai/nomic-ai/nomic-embed-text-v1": { "input_cost_per_token": 8e-09, "litellm_provider": "fireworks_ai-embedding-models", @@ -25103,6 +25530,21 @@ "supports_tool_choice": true, "supports_vision": true }, + "mistral/mistral-medium-3-5": { + "input_cost_per_token": 1.5e-06, + "litellm_provider": "mistral", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 7.5e-06, + "source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "mistral/mistral-small": { "input_cost_per_token": 1e-07, "litellm_provider": "mistral", @@ -39351,6 +39793,22 @@ "litellm_provider": "fireworks_ai", "mode": "chat" }, + "fireworks_ai/accounts/fireworks/models/qwen3p7-plus": { + "cache_read_input_token_cost": 8e-08, + "input_cost_per_token": 4e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 262144, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_token": 1.6e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "fireworks_ai/accounts/fireworks/models/qwq-32b": { "max_tokens": 131072, "max_input_tokens": 131072, @@ -39513,6 +39971,54 @@ "litellm_provider": "fireworks_ai", "mode": "chat" }, + "fireworks_ai/accounts/fireworks/routers/glm-5p1-fast": { + "cache_read_input_token_cost": 5.2e-07, + "input_cost_per_token": 2.8e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 202800, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 8.8e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/accounts/fireworks/routers/kimi-k2p6-fast": { + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 2e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 8e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "fireworks_ai/accounts/fireworks/routers/kimi-k2p7-code-fast": { + "cache_read_input_token_cost": 3.8e-07, + "input_cost_per_token": 1.9e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 8e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "novita/deepseek/deepseek-v3.2": { "litellm_provider": "novita", "mode": "chat", @@ -42456,4 +42962,105 @@ "supports_reasoning": true, "source": "https://serverless.tensormesh.ai/v1/models/openrouter" } -} \ No newline at end of file +, + "deepseek-v4-flash": { + "cache_creation_input_token_cost": 0.0, + "cache_read_input_token_cost": 2.8e-09, + "input_cost_per_token": 1.4e-07, + "input_cost_per_token_cache_hit": 2.8e-09, + "litellm_provider": "deepseek", + "max_input_tokens": 1000000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 2.8e-07, + "source": "https://api-docs.deepseek.com/quick_start/pricing", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "deepseek-v4-pro": { + "cache_creation_input_token_cost": 0.0, + "cache_read_input_token_cost": 3.625e-09, + "input_cost_per_token": 4.35e-07, + "input_cost_per_token_cache_hit": 3.625e-09, + "litellm_provider": "deepseek", + "max_input_tokens": 1000000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 8.7e-07, + "source": "https://api-docs.deepseek.com/quick_start/pricing", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "deepseek/deepseek-v4-flash": { + "cache_creation_input_token_cost": 0.0, + "cache_read_input_token_cost": 2.8e-09, + "input_cost_per_token": 1.4e-07, + "input_cost_per_token_cache_hit": 2.8e-09, + "litellm_provider": "deepseek", + "max_input_tokens": 1000000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 2.8e-07, + "source": "https://api-docs.deepseek.com/quick_start/pricing", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "deepseek/deepseek-v4-pro": { + "cache_creation_input_token_cost": 0.0, + "cache_read_input_token_cost": 3.625e-09, + "input_cost_per_token": 4.35e-07, + "input_cost_per_token_cache_hit": 3.625e-09, + "litellm_provider": "deepseek", + "max_input_tokens": 1000000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 8.7e-07, + "source": "https://api-docs.deepseek.com/quick_start/pricing", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": false + } +} diff --git a/litellm/mypy.ini b/litellm/mypy.ini deleted file mode 100644 index b65e11bab42..00000000000 --- a/litellm/mypy.ini +++ /dev/null @@ -1,22 +0,0 @@ -[mypy] -warn_return_any = True -ignore_missing_imports = True -disallow_untyped_defs = True -mypy_path = litellm/stubs -namespace_packages = True -disable_error_code = - annotation-unchecked, - import-untyped - -[mypy-litellm.*] -ignore_missing_imports = False - -[mypy-google.*] -ignore_missing_imports = True - -[mypy-cryptography.hazmat.bindings._rust.x509] -ignore_errors = True - -[mypy-fastuuid.*] -ignore_missing_imports = True -ignore_errors = True \ No newline at end of file diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 1535daeb01d..e47fc84b533 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -125,7 +125,7 @@ class MCPRequestHandler: LITELLM_MCP_ACCESS_GROUPS_HEADER_NAME = SpecialHeaders.mcp_access_groups.value @staticmethod - async def process_mcp_request( # noqa: PLR0915 + async def process_mcp_request( scope: Scope, ) -> Tuple[ UserAPIKeyAuth, diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 5e419b5c0a3..afec884cd96 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -1421,8 +1421,11 @@ class MCPServerManager: "No allowed MCP Servers found for user api key auth." ) return list(combined_servers) - except Exception as e: - verbose_logger.warning(f"Failed to get allowed MCP servers: {str(e)}.") + except Exception: # noqa: BLE001 + verbose_logger.exception( + "Failed to get allowed MCP servers; team-level object_permission " + "grants may be dropped. Falling back to global servers only." + ) return allow_all_server_ids async def resolve_toolset_tool_permissions( @@ -3355,7 +3358,7 @@ class MCPServerManager: ) ) - async def _call_regular_mcp_tool( # noqa: PLR0915 + async def _call_regular_mcp_tool( self, mcp_server: MCPServer, original_tool_name: str, diff --git a/litellm/proxy/_experimental/mcp_server/sampling_handler.py b/litellm/proxy/_experimental/mcp_server/sampling_handler.py index 1637c9eb0b9..b659ba6f813 100644 --- a/litellm/proxy/_experimental/mcp_server/sampling_handler.py +++ b/litellm/proxy/_experimental/mcp_server/sampling_handler.py @@ -661,7 +661,7 @@ def _convert_openai_response_to_mcp_result( ) -async def _check_model_access( # noqa: PLR0915 +async def _check_model_access( model: str, user_api_key_auth: Any ) -> Optional["ErrorData"]: """Enforce model-permission checks for MCP sampling requests. diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 746fc4e7d3f..08e42e918e9 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -617,7 +617,7 @@ if MCP_AVAILABLE: active_mcp_session_var.reset(_session_reset_token) @server.call_tool() - async def mcp_server_tool_call( # noqa: PLR0915 + async def mcp_server_tool_call( name: str, arguments: Dict[str, Any] | None ) -> CallToolResult: """ @@ -1036,7 +1036,14 @@ if MCP_AVAILABLE: allowed_mcp_servers: List[MCPServer], ) -> List[MCPServer]: """ - Get the filtered MCP servers from the MCP server names + Get the filtered MCP servers from the MCP server names. + + Fails closed when ``mcp_servers`` is explicitly provided (path- or + header-derived) but none of the names resolve to a server alias or + access group the caller can access. The previous behavior returned + the full ``allowed_mcp_servers`` set, which silently widened scope + when a client targeted ``/mcp//`` and made URL/header + namespacing appear to work when it did not. """ filtered_server: dict[str, MCPServer] = {} @@ -1076,6 +1083,17 @@ if MCP_AVAILABLE: if filtered_server: return list(filtered_server.values()) + if mcp_servers is not None: + # Caller asked for a specific scope but nothing resolved. Fail + # closed so URL/header namespacing cannot silently fall back to + # the caller's full allowed-server set. + verbose_logger.debug( + "MCP scope filter resolved to no servers for requested names %s; " + "returning empty list (fail-closed).", + mcp_servers, + ) + return [] + return allowed_mcp_servers def _tool_name_matches(tool_name: str, filter_list: List[str]) -> bool: @@ -1591,7 +1609,7 @@ if MCP_AVAILABLE: _mcp_gateway_initialize_instructions.reset(instructions_token) _mcp_gateway_server_name.reset(server_name_token) - async def _get_tools_from_mcp_servers( # noqa: PLR0915 + async def _get_tools_from_mcp_servers( user_api_key_auth: Optional[UserAPIKeyAuth], mcp_auth_header: Optional[str], mcp_servers: Optional[List[str]], @@ -2435,7 +2453,7 @@ if MCP_AVAILABLE: }, ) - async def execute_mcp_tool( # noqa: PLR0915 + async def execute_mcp_tool( name: str, arguments: Dict[str, Any], allowed_mcp_servers: List[MCPServer], @@ -3642,7 +3660,7 @@ if MCP_AVAILABLE: detail="Forbidden", ) - async def handle_streamable_http_mcp( # noqa: PLR0915 + async def handle_streamable_http_mcp( scope: Scope, receive: Receive, send: Send ) -> None: """Handle MCP requests through StreamableHTTP.""" diff --git a/litellm/proxy/_experimental/mcp_server/utils.py b/litellm/proxy/_experimental/mcp_server/utils.py index 97cfa74ea45..b0141d3207c 100644 --- a/litellm/proxy/_experimental/mcp_server/utils.py +++ b/litellm/proxy/_experimental/mcp_server/utils.py @@ -23,9 +23,20 @@ import os from urllib.parse import quote # Constants -LITELLM_MCP_SERVER_NAME = "litellm-mcp-server" +# +# NOTE: The environment-backed values below are read once, when this module is +# first imported, and cached for the lifetime of the process. Changing the +# corresponding environment variables after import has no effect unless the +# module is reloaded (e.g. ``importlib.reload``). Tests that override these +# variables must reload this module — see +# ``tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_identity_env.py``. +LITELLM_MCP_SERVER_NAME = os.environ.get( + "LITELLM_MCP_SERVER_NAME", "litellm-mcp-server" +) LITELLM_MCP_SERVER_VERSION = "1.0.0" -LITELLM_MCP_SERVER_DESCRIPTION = "MCP Server for LiteLLM" +LITELLM_MCP_SERVER_DESCRIPTION = os.environ.get( + "LITELLM_MCP_SERVER_DESCRIPTION", "MCP Server for LiteLLM" +) MCP_TOOL_PREFIX_SEPARATOR = os.environ.get("MCP_TOOL_PREFIX_SEPARATOR", "-") MCP_TOOL_PREFIX_FORMAT = "{server_name}{separator}{tool_name}" diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 291bf0a1372..f9c23475951 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -361,6 +361,16 @@ class LiteLLMRoutes(enum.Enum): "/realtime?{model}", "/v1/realtime?{model}", "/openai/v1/realtime?{model}", + # realtime (GA WebRTC HTTP routes) + "/realtime/client_secrets", + "/v1/realtime/client_secrets", + "/openai/v1/realtime/client_secrets", + "/realtime/calls", + "/v1/realtime/calls", + "/openai/v1/realtime/calls", + "/realtime/transcription_sessions", + "/v1/realtime/transcription_sessions", + "/openai/v1/realtime/transcription_sessions", # responses API "/responses", "/v1/responses", @@ -2153,6 +2163,10 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): master_key: Optional[str] = Field( None, description="require a key for all calls to proxy" ) + allow_cli_sso_verification_uri_complete: bool | None = Field( + None, + description="opt-in to RFC 8628 verification_uri_complete for the CLI SSO device flow, pre-filling the user_code in the browser. Off by default; intended for same-host clients where the device that starts the flow and the browser run on the same machine", + ) database_url: Optional[str] = Field( None, description="connect to a postgres db - needed for generating temporary keys + tracking spend / key", diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index 7b2f75e1cff..7446f61ad1c 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -509,7 +509,7 @@ async def get_agent_card( tags=["[beta] A2A Agents"], dependencies=[Depends(user_api_key_auth)], ) -async def invoke_agent_a2a( # noqa: PLR0915 +async def invoke_agent_a2a( agent_id: str, request: Request, fastapi_response: Response, diff --git a/litellm/proxy/anthropic_endpoints/endpoints.py b/litellm/proxy/anthropic_endpoints/endpoints.py index 1995ff275c9..856b788b54b 100644 --- a/litellm/proxy/anthropic_endpoints/endpoints.py +++ b/litellm/proxy/anthropic_endpoints/endpoints.py @@ -5,6 +5,7 @@ Unified /v1/messages endpoint - (Anthropic Spec) from fastapi import APIRouter, Depends, HTTPException, Request, Response from fastapi.responses import JSONResponse +import litellm from litellm._logging import verbose_proxy_logger from litellm.anthropic_interface.exceptions import AnthropicExceptionMapping from litellm.integrations.custom_guardrail import ModifyResponseException @@ -23,6 +24,40 @@ from litellm.types.utils import TokenCountResponse router = APIRouter() +def _strip_total_tokens_from_anthropic_response(response: Any) -> None: + """Remove the OpenAI-flavored `usage.total_tokens` field that LiteLLM + injects into Anthropic /v1/messages responses. + + The Anthropic /v1/messages spec only defines: + input_tokens, output_tokens, cache_creation_input_tokens, + cache_read_input_tokens, cache_creation.{ephemeral_5m,ephemeral_1h} + The streaming SSE path (message_delta.usage) already does not include + total_tokens; this brings the non-streaming path into the same shape. + + Handles both shapes returned by `base_process_llm_request`: + - plain `dict` (most common — `AnthropicMessagesResponse` is a TypedDict + and is `dict` at runtime) + - Pydantic model whose `usage` attribute is dict-shaped (e.g. a + BaseModel that holds raw Anthropic usage as a `dict[str, int]`) + + Streaming results (StreamingResponse, AsyncIterator, etc.) and Pydantic + models with strongly-typed Usage sub-models are left untouched — + those paths either have separate serialization handling or impose + type constraints the helper does not try to subvert. + """ + if response is None: + return + if isinstance(response, dict): + usage = response.get("usage") + if isinstance(usage, dict) and "total_tokens" in usage: + usage.pop("total_tokens", None) + return + # Pydantic-model fallback: only mutate if `usage` is a dict. + usage = getattr(response, "usage", None) + if isinstance(usage, dict) and "total_tokens" in usage: + usage.pop("total_tokens", None) + + @router.post( "/v1/messages", tags=["[beta] Anthropic `/v1/messages`"], @@ -72,6 +107,18 @@ async def anthropic_response( user_api_base=user_api_base, version=version, ) + # Optionally strip the non-Anthropic `usage.total_tokens` field + # LiteLLM adds internally. Anthropic's official /v1/messages spec + # only defines input_tokens / output_tokens / cache_*_input_tokens; + # total_tokens is an OpenAI convention. Default off + # (`litellm.strip_anthropic_total_tokens = False`) to preserve + # backward compatibility for clients that currently read it; set + # to True to align the wire response with the spec (and with the + # streaming SSE path, which already omits total_tokens). + # spend_logs / Prometheus still compute total internally — this + # only affects the wire response. + if litellm.strip_anthropic_total_tokens: + _strip_total_tokens_from_anthropic_response(result) return result except ModifyResponseException as e: # Guardrail flagged content in passthrough mode - return 200 with violation message diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index aa967732a90..814346eddf8 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -519,7 +519,7 @@ MODEL_DISCOVERY_ROUTES = frozenset( ) -async def common_checks( # noqa: PLR0915 +async def common_checks( request_body: dict, team_object: Optional[LiteLLM_TeamTable], user_object: Optional[LiteLLM_UserTable], diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index c868d3d22b2..94b2ed84f20 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -2,6 +2,7 @@ import os import re import sys from functools import lru_cache +from logging import Logger from typing import Any, Dict, FrozenSet, List, Mapping, Optional, Tuple, Union from fastapi import HTTPException, Request, status @@ -995,6 +996,45 @@ def get_project_model_tpm_limit( return None +def custom_auth_common_checks_warning( + *, + custom_auth_configured: bool, + run_common_checks: bool, +) -> str | None: + if not custom_auth_configured or run_common_checks: + return None + return ( + "custom_auth is configured but 'custom_auth_run_common_checks' is not set. " + "Problem: budgets, model-access allowlists, and per-model rate limits configured " + "on your DB team/project records will NOT be enforced for custom-auth requests " + "(rate limits set directly on the returned UserAPIKeyAuth still apply). " + "Fix: set 'general_settings.custom_auth_run_common_checks: true'. " + "Docs: https://docs.litellm.ai/docs/proxy/custom_auth" + ) + + +_custom_auth_common_checks_warning_emitted = False + + +def warn_once_if_custom_auth_skips_common_checks( + *, + custom_auth_configured: bool, + run_common_checks: bool, + logger: Logger = verbose_proxy_logger, +) -> None: + global _custom_auth_common_checks_warning_emitted + if _custom_auth_common_checks_warning_emitted: + return + message = custom_auth_common_checks_warning( + custom_auth_configured=custom_auth_configured, + run_common_checks=run_common_checks, + ) + if message is None: + return + logger.warning(message) + _custom_auth_common_checks_warning_emitted = True + + def is_pass_through_provider_route(route: str) -> bool: PROVIDER_SPECIFIC_PASS_THROUGH_ROUTES = [ "vertex-ai", @@ -1227,6 +1267,14 @@ _MODEL_ROUTING_BODY_TARGET_MODEL_ROUTE_MARKERS = ( "/vector_stores", ) _MODEL_ROUTING_COMPLETION_MODEL_ROUTE_MARKERS = ("/evals",) +# Realtime WebRTC routes carry the effective model inside the nested +# ``session.model`` field (see realtime_endpoints.endpoints), so the model the +# request will actually use is not present at the top level. Extract it here so +# can_key_call_model() validates the real target model. +_MODEL_ROUTING_SESSION_MODEL_ROUTE_MARKERS = ( + "/realtime/client_secrets", + "/realtime/calls", +) _MODEL_ROUTING_ID_FIELDS = ( "file_id", "input_file_id", @@ -1409,6 +1457,12 @@ def _extract_model_candidates_from_request( _append_model_candidates(candidates, body_model) if uses_body_target_model_sources or not body_model: _append_model_candidates(candidates, request_data.get("target_model_names")) + if _route_matches_any_marker( + route=route, markers=_MODEL_ROUTING_SESSION_MODEL_ROUTE_MARKERS + ): + session = request_data.get("session") + if isinstance(session, dict): + _append_model_candidates(candidates, session.get("model")) if uses_completion_model_sources and isinstance( request_data.get("completion"), dict ): diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index fd6ff2ada7f..90845dfd824 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -1954,7 +1954,7 @@ class JWTAuthManager: return None, None, None @staticmethod - async def auth_builder( # noqa: PLR0915 + async def auth_builder( api_key: str, jwt_handler: JWTHandler, request_data: dict, diff --git a/litellm/proxy/auth/model_checks.py b/litellm/proxy/auth/model_checks.py index 00f276dc970..b89db51c6f1 100644 --- a/litellm/proxy/auth/model_checks.py +++ b/litellm/proxy/auth/model_checks.py @@ -10,6 +10,7 @@ from litellm.repositories.object_permission_repository import ObjectPermissionRe from litellm.router import Router from litellm.router_utils.fallback_event_handlers import get_fallback_model_group from litellm.types.router import CredentialLiteLLMParams, LiteLLM_Params +from litellm.types.utils import LlmProviders from litellm.utils import get_valid_models _CREDENTIAL_LITELLM_PARAM_FIELDS = set(CredentialLiteLLMParams.model_fields) @@ -308,10 +309,21 @@ def get_known_models_from_wildcard( # add model prefix to wildcard models wildcard_models = [f"{model_prefix}{model}" for model in wildcard_models] + known_providers = {provider.value for provider in LlmProviders} suffix_appended_wildcard_models = [] for model in wildcard_models: if not model.startswith(wildcard_provider_prefix): - model = f"{wildcard_provider_prefix}/{model}" + # `get_provider_models` returns provider-prefixed ids (e.g. "ollama/gemma3:1b"). + # When the wildcard uses a custom prefix (e.g. "ollama_server1/*" to distinguish + # multiple instances), replace that existing provider prefix instead of stacking + # both, which would otherwise yield an uncallable "ollama_server1/ollama/gemma3:1b". + # Only strip the leading segment when it is a known provider, so ids whose first + # segment is an org rather than a provider (e.g. "meta-llama/Llama-3-8B") keep it. + leading, sep, model_suffix = model.partition("/") + if sep and leading in known_providers: + model = f"{wildcard_provider_prefix}/{model_suffix}" + else: + model = f"{wildcard_provider_prefix}/{model}" suffix_appended_wildcard_models.append(model) return suffix_appended_wildcard_models or [] diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index 666c01562b5..6f359e52eeb 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -979,7 +979,7 @@ def _ensure_parent_otel_span_on_request_state(request: Request) -> None: request.state.parent_otel_span = parent_otel_span -async def _user_api_key_auth_builder( # noqa: PLR0915 +async def _user_api_key_auth_builder( request: Request, api_key: str, azure_api_key_header: str, @@ -2126,7 +2126,7 @@ def _team_obj_from_token(valid_token: UserAPIKeyAuth) -> LiteLLM_TeamTableCached @tracer.wrap() -async def _run_centralized_common_checks( # noqa: PLR0915 +async def _run_centralized_common_checks( user_api_key_auth_obj: UserAPIKeyAuth, request: Request, request_data: dict, diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index ea479a5721b..344f90aa144 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -58,7 +58,7 @@ router = APIRouter() dependencies=[Depends(user_api_key_auth)], tags=["batch"], ) -async def create_batch( # noqa: PLR0915 +async def create_batch( request: Request, fastapi_response: Response, provider: Optional[str] = None, @@ -343,7 +343,7 @@ async def create_batch( # noqa: PLR0915 dependencies=[Depends(user_api_key_auth)], tags=["batch"], ) -async def retrieve_batch( # noqa: PLR0915 +async def retrieve_batch( request: Request, fastapi_response: Response, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 21ecc08f44a..bfeaf84930d 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -31,6 +31,7 @@ from litellm.constants import ( DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE, DEFAULT_MAX_RECURSE_DEPTH, LITELLM_DETAILED_TIMING, + LITELLM_HTTP_STATUS_CLIENT_DISCONNECTED, MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG, STREAM_SSE_DATA_PREFIX, ) @@ -67,7 +68,12 @@ if TYPE_CHECKING: else: ProxyConfig = Any from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request -from litellm.types.utils import ModelResponse, ModelResponseStream, Usage +from litellm.types.utils import ( + ModelResponse, + ModelResponseStream, + StandardLoggingPayloadErrorInformation, + Usage, +) # Datadog streaming spans are a no-op when ddtrace is not enabled, but the # ``with tracer.trace(...)`` context manager still allocates a NullSpan and @@ -77,6 +83,63 @@ from litellm.types.utils import ModelResponse, ModelResponseStream, Usage _DD_STREAMING_TRACE_ENABLED = not isinstance(tracer, NullTracer) +_CLIENT_DISCONNECTED_ERROR_INFORMATION: StandardLoggingPayloadErrorInformation = { + "error_code": str(LITELLM_HTTP_STATUS_CLIENT_DISCONNECTED), + "error_message": "Client disconnected the request", + "error_class": "ClientDisconnected", +} + + +def _apply_client_disconnect_metadata(target_metadata: dict[str, object]) -> None: + target_metadata["client_disconnected"] = True + target_metadata["error_information"] = dict(_CLIENT_DISCONNECTED_ERROR_INFORMATION) + + +async def _record_streaming_client_disconnect_if_needed( + request: Request | None, + request_data: dict, + client_disconnected: bool = False, +) -> bool: + if not client_disconnected: + if request is None: + return False + try: + disconnected = await request.is_disconnected() + except Exception: # noqa: BLE001 + return False + if not disconnected: + return False + + logging_obj = request_data.get("litellm_logging_obj") + if logging_obj is not None: + litellm_params = logging_obj.model_call_details.setdefault("litellm_params", {}) + _apply_client_disconnect_metadata(litellm_params.setdefault("metadata", {})) + _apply_client_disconnect_metadata( + logging_obj.model_call_details.setdefault("metadata", {}) + ) + + _apply_client_disconnect_metadata(request_data.setdefault("metadata", {})) + litellm_params = request_data.setdefault("litellm_params", {}) + _apply_client_disconnect_metadata(litellm_params.setdefault("metadata", {})) + + verbose_proxy_logger.debug( + "Recorded streaming client disconnect with error_code=499 for litellm_call_id=%s", + request_data.get("litellm_call_id"), + ) + return True + + +async def _cancel_pending_gather_tasks(tasks: list["asyncio.Task[Any]"]) -> None: + pending_tasks = [task for task in tasks if not task.done()] + for task in pending_tasks: + task.cancel() + for task in pending_tasks: + try: + await task + except (asyncio.CancelledError, Exception): # noqa: BLE001 + pass + + def _serialize_http_exception_detail( detail: Any, ) -> Tuple[str, Optional[dict]]: @@ -242,20 +305,6 @@ def _extract_error_from_sse_chunk(event_line: Union[str, bytes]) -> dict: return default_error -async def _aclose_upstream_response(response: Any) -> None: - """Release the upstream HTTP connection when a stream ends for any - reason, including client disconnect. Mirrors the finally block of - async_data_generator in proxy_server.py.""" - with anyio.CancelScope(shield=True): - if hasattr(response, "aclose"): - try: - await response.aclose() - except BaseException as e: - verbose_proxy_logger.debug( - "error closing upstream response stream: %s", e - ) - - class _UpstreamClosingStreamingResponse(StreamingResponse): """StreamingResponse that always closes its body iterator and the wrapped upstream generator. @@ -300,7 +349,7 @@ class _UpstreamClosingStreamingResponse(StreamingResponse): ) -async def create_response( # noqa: PLR0915 +async def create_response( generator: AsyncGenerator[str, None], media_type: str, headers: dict, @@ -1168,7 +1217,7 @@ class ProxyBaseLLMRequestProcessing: _payload_str, ) - async def base_process_llm_request( # noqa: PLR0915 + async def base_process_llm_request( self, request: Request, fastapi_response: Response, @@ -1358,19 +1407,22 @@ class ProxyBaseLLMRequestProcessing: user_model=user_model, user_api_key_dict=user_api_key_dict, ) - tasks.append(llm_call) + llm_call_task = asyncio.create_task(llm_call) + tasks.append(llm_call_task) - # wait for call to end llm_responses = asyncio.gather( *tasks ) # run the moderation check in parallel to the actual llm api call - if general_settings.get("cancel_on_disconnect", False): - responses = await _await_llm_call_cancelling_on_disconnect( - request, llm_responses - ) - else: - responses = await llm_responses + try: + if general_settings.get("cancel_on_disconnect", False): + responses = await _await_llm_call_cancelling_on_disconnect( + request, llm_responses + ) + else: + responses = await llm_responses + finally: + await _cancel_pending_gather_tasks(tasks) response = responses[1] @@ -1555,6 +1607,7 @@ class ProxyBaseLLMRequestProcessing: user_api_key_dict=user_api_key_dict, request_data=self.data, proxy_logging_obj=proxy_logging_obj, + request=request, ) ) return await create_response( @@ -1568,6 +1621,7 @@ class ProxyBaseLLMRequestProcessing: response=response, user_api_key_dict=user_api_key_dict, request_data=self.data, + request=request, ) if route_type == "aresponses": # Streaming /v1/responses returns here without @@ -2324,6 +2378,13 @@ class ProxyBaseLLMRequestProcessing: self._apply_router_cooldown_retry_after(headers, e) + if isinstance(e, ProxyException): + e.headers = { + **e.headers, + **{k: v if isinstance(v, str) else str(v) for k, v in headers.items()}, + } + raise e + if isinstance(e, HTTPException): raw_detail = getattr(e, "detail", str(e)) message, structured_fields = _serialize_http_exception_detail(raw_detail) @@ -2412,6 +2473,39 @@ class ProxyBaseLLMRequestProcessing: else: return chunk + @staticmethod + async def _finalize_streaming_generator_cleanup( + request: Request | None, + request_data: dict, + response: Any, + stream_completed: bool = False, + client_disconnected: bool = False, + ) -> None: + with anyio.CancelScope(shield=True): + should_record_client_disconnect = client_disconnected or ( + not stream_completed + ) + recorded_client_disconnect = False + if should_record_client_disconnect: + recorded_client_disconnect = ( + await _record_streaming_client_disconnect_if_needed( + request, + request_data, + client_disconnected, + ) + ) + if recorded_client_disconnect: + ProxyLogging._fire_deferred_stream_logging(request_data) + + if hasattr(response, "aclose"): + try: + await response.aclose() + except BaseException as e: # noqa: BLE001 + verbose_proxy_logger.debug( + "async_streaming_data_generator: error closing response stream: %s", + e, + ) + @staticmethod async def async_streaming_data_generator( response: Any, @@ -2421,6 +2515,7 @@ class ProxyBaseLLMRequestProcessing: *, serialize_chunk: StreamChunkSerializer, serialize_error: StreamErrorSerializer, + request: Request | None = None, ) -> AsyncGenerator[str, None]: """ Shared streaming data generator: runs proxy iterator hook, per-chunk hook, @@ -2445,6 +2540,8 @@ class ProxyBaseLLMRequestProcessing: and not cost_injection_enabled ) debug_enabled = verbose_proxy_logger.isEnabledFor(logging.DEBUG) + stream_completed = False + client_disconnected = False try: str_so_far = "" async for ( @@ -2492,6 +2589,7 @@ class ProxyBaseLLMRequestProcessing: ) ) yield serialize_chunk(chunk) + stream_completed = True except (asyncio.CancelledError, GeneratorExit): # Client disconnected mid-stream. CancelledError / GeneratorExit # are BaseException and bypass the success/failure logging @@ -2499,9 +2597,11 @@ class ProxyBaseLLMRequestProcessing: # release it here. This is the outermost generator Starlette closes # on disconnect, so the nested iterator hook (which only sees # GeneratorExit on GC) cannot own the refund. - proxy_logging_obj._release_max_parallel_requests_on_disconnect( - user_api_key_dict - ) + if not stream_completed: + proxy_logging_obj._release_max_parallel_requests_on_disconnect( + user_api_key_dict + ) + client_disconnected = True raise except Exception as e: verbose_proxy_logger.exception( @@ -2530,9 +2630,16 @@ class ProxyBaseLLMRequestProcessing: param=getattr(e, "param", "None"), code=getattr(e, "status_code", 500), ) + stream_completed = True yield serialize_error(proxy_exception) finally: - await _aclose_upstream_response(response) + await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup( + request=request, + request_data=request_data, + response=response, + stream_completed=stream_completed, + client_disconnected=client_disconnected, + ) @staticmethod def async_sse_data_generator( @@ -2540,6 +2647,7 @@ class ProxyBaseLLMRequestProcessing: user_api_key_dict: UserAPIKeyAuth, request_data: dict, proxy_logging_obj: ProxyLogging, + request: Request | None = None, ) -> AsyncGenerator[str, None]: """ Anthropic /messages and Google /generateContent streaming data generator require SSE events. @@ -2558,6 +2666,7 @@ class ProxyBaseLLMRequestProcessing: serialize_error=lambda proxy_exc: ( f"{STREAM_SSE_DATA_PREFIX}{json.dumps({'error': proxy_exc.to_dict()})}\n\n" ), + request=request, ) @staticmethod diff --git a/litellm/proxy/common_utils/callback_utils.py b/litellm/proxy/common_utils/callback_utils.py index ab1eeaf1646..71dce163b78 100644 --- a/litellm/proxy/common_utils/callback_utils.py +++ b/litellm/proxy/common_utils/callback_utils.py @@ -36,7 +36,7 @@ if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging -def initialize_callbacks_on_proxy( # noqa: PLR0915 +def initialize_callbacks_on_proxy( value: Any, premium_user: bool, config_file_path: str, diff --git a/litellm/proxy/common_utils/model_listing_utils.py b/litellm/proxy/common_utils/model_listing_utils.py new file mode 100644 index 00000000000..3a70377037d --- /dev/null +++ b/litellm/proxy/common_utils/model_listing_utils.py @@ -0,0 +1,167 @@ +"""Team-scoped (BYOK) model-name translation for the model listing endpoints. + +`/v1/models`, `/models`, and `GET /v1/models/{id}` should surface the public +`team_public_model_name` rather than the internal routing key +`model_name_{team_id}_{uuid}`, consistent with `/v1/model/info`. The internal +key still routes regardless; this is a presentation-layer swap only and does not +touch access-group or auth semantics (see issue #28382). Operators can pin the +legacy internal names with `general_settings.use_team_public_model_name: false`. +""" + +from __future__ import annotations + +from collections.abc import Mapping +from typing import TYPE_CHECKING, cast + +if TYPE_CHECKING: + from litellm.router import Router + + +class TeamModelNameTranslator: + """Translates internal team routing keys to their public names for the model + listing/retrieve responses. Stateless; the live router and general_settings + are injected per call so the unit tests can drive it without globals. + """ + + @staticmethod + def _internal_public_pair(model: object) -> tuple[str, str] | None: + """`(internal_routing_key, public_name)` for a team-scoped row, else None.""" + if not isinstance(model, dict): + return None + model_dict = cast(dict[str, object], model) # any-ok: checked + model_info_raw: object = model_dict.get("model_info") + if not isinstance(model_info_raw, Mapping): + return None + model_info = cast(Mapping[str, object], model_info_raw) # any-ok: checked + team_id = model_info.get("team_id") + team_public = model_info.get("team_public_model_name") + name = model_dict.get("model_name") + if ( + isinstance(team_id, str) + and isinstance(team_public, str) + and isinstance(name, str) + and team_id + and team_public + and name.startswith(f"model_name_{team_id}_") + ): + return name, team_public + return None + + @staticmethod + def _is_enabled(general_settings: Mapping[str, object]) -> bool: + return general_settings.get("use_team_public_model_name", True) is not False + + @staticmethod + def build_internal_to_public_map( + llm_router: "Router | None", + general_settings: Mapping[str, object], + ) -> dict[str, str]: + """Internal team routing key -> public `team_public_model_name`. + + Empty when disabled via the legacy flag, the router is absent, or the + router model list is malformed. + """ + if llm_router is None or not TeamModelNameTranslator._is_enabled( + general_settings + ): + return {} + router_model_list = llm_router.get_model_list() + if not isinstance(router_model_list, list): + return {} + return dict( + pair + for pair in ( + TeamModelNameTranslator._internal_public_pair(model) + for model in router_model_list + ) + if pair is not None + ) + + @staticmethod + def _response_to_lookup_map( + model_names: list[str], + internal_to_public: dict[str, str], + ) -> dict[str, str]: + """Map each public response id to the first internal lookup id seen in + `model_names`, preserving first-occurrence order. First-wins keeps list + and retrieve in agreement on which accessible deployment a shared public + id resolves to: a global iterated before a colliding team alias stays + the listed entry, and sibling team rows collapse to their first + occurrence. + """ + result: dict[str, str] = {} + for name in model_names: + result.setdefault(internal_to_public.get(name, name), name) + return result + + @staticmethod + def listing_entries( + model_names: list[str], + llm_router: "Router | None", + general_settings: Mapping[str, object], + ) -> list[tuple[str, str]]: + """`(response_id, metadata_lookup_id)` for each listed model, de-duplicated + by response_id while preserving order. + + For team-scoped rows `response_id` is the public name shown to the client, + while `metadata_lookup_id` stays the internal routing key so downstream + metadata/fallback lookups (keyed by the routing name) still resolve. The + lookup id is always one of `model_names` (the caller's accessible set), so + a public name shared across teams never resolves to another team's + internal key. Both ids are identical for unmapped names (globals, + access-group keys). + """ + internal_to_public = TeamModelNameTranslator.build_internal_to_public_map( + llm_router, general_settings + ) + if not internal_to_public: + return [(name, name) for name in model_names] + return list( + TeamModelNameTranslator._response_to_lookup_map( + model_names, internal_to_public + ).items() + ) + + @staticmethod + def translate_listing( + model_names: list[str], + llm_router: "Router | None", + general_settings: Mapping[str, object], + ) -> list[str]: + """Public-name view of `model_names` (the `response_id` of each listing + entry). Sibling deployments sharing a public name collapse to one entry + while preserving order; unmapped names pass through. + """ + return [ + entry[0] + for entry in TeamModelNameTranslator.listing_entries( + model_names, llm_router, general_settings + ) + ] + + @staticmethod + def resolve_public_name( + model_id: str, + available_models: list[str], + llm_router: "Router | None", + general_settings: Mapping[str, object], + ) -> str: + """Resolve a public team name back to the internal routing key the router + indexes by, so `GET /v1/models/{id}` accepts the name the listing returns. + + Resolution is restricted to `available_models` (the caller's accessible + set) so colliding public names across teams never resolve across an access + boundary. Uses the same first-occurrence dedup as `listing_entries` so a + public id advertised by `/v1/models` resolves to the same internal + deployment that the listing's metadata was built from. Returns `model_id` + unchanged when it is not an accessible public team name (already-internal + names and globals pass through). + """ + internal_to_public = TeamModelNameTranslator.build_internal_to_public_map( + llm_router, general_settings + ) + if not internal_to_public: + return model_id + return TeamModelNameTranslator._response_to_lookup_map( + available_models, internal_to_public + ).get(model_id, model_id) diff --git a/litellm/proxy/common_utils/timezone_utils.py b/litellm/proxy/common_utils/timezone_utils.py index 700a9197f6f..32f9f47d519 100644 --- a/litellm/proxy/common_utils/timezone_utils.py +++ b/litellm/proxy/common_utils/timezone_utils.py @@ -15,7 +15,7 @@ def get_budget_reset_timezone(): return getattr(litellm, "timezone", None) or "UTC" -def get_budget_reset_time(budget_duration: str): +def get_budget_reset_time(budget_duration: str) -> datetime: """ Get the budget reset time based on the configured timezone. Falls back to UTC if not specified. diff --git a/litellm/proxy/db/create_views.py b/litellm/proxy/db/create_views.py index 97525a528d0..d9e21fc5d2a 100644 --- a/litellm/proxy/db/create_views.py +++ b/litellm/proxy/db/create_views.py @@ -11,7 +11,7 @@ _db = Any _VIEW_NOT_FOUND_MARKERS = ("does not exist", "no such table", "undefined table") -async def create_missing_views(db: _db): # noqa: PLR0915 +async def create_missing_views(db: _db): """ -------------------------------------------------- NOTE: Copy of `litellm/db_scripts/create_views.py`. diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index e7f14df5294..aab92a54577 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -1128,7 +1128,7 @@ class DBSpendUpdateWriter: "_flush_tool_discovery_queue error (non-blocking): %s", e ) - async def _commit_spend_updates_to_db( # noqa: PLR0915 + async def _commit_spend_updates_to_db( self, prisma_client: PrismaClient, n_retry_times: int, @@ -1613,199 +1613,215 @@ class DBSpendUpdateWriter: start_time = time.time() try: - for i in range(n_retry_times + 1): - try: - # Sort the transactions to minimize the probability of deadlocks by reducing the chance of concurrent - # trasactions locking the same rows/ranges in different orders. - transactions_to_process = dict( - sorted( - daily_spend_transactions.items(), - # Normally to avoid deadlocks we would sort by the index, but since we have sprinkled indexes - # on our schema like we're discount Salt Bae, we just sort by all fields that have an index, - # in an ad-hoc (but hopefully sensible) order of indexes. The actual ordering matters less than - # ensuring that all concurrent transactions sort in the same order. - # We could in theory use the dict key, as it contains basically the same fields, but this is more - # robust to future changes in the key format. - # If _update_daily_spend ever gets the ability to write to multiple tables at once, the sorting - # should sort by the table first. - key=lambda x: ( - x[1].get("date") or "", - x[1].get(entity_id_field) or "", - x[1].get("api_key") or "", - x[1].get("model") or "", - x[1].get("custom_llm_provider") or "", - ), - )[:BATCH_SIZE] - ) - - if len(transactions_to_process) == 0: - verbose_proxy_logger.debug( - f"No new transactions to process for daily {entity_type} spend update" - ) - break - + while daily_spend_transactions: + for i in range(n_retry_times + 1): try: - async with prisma_client.db.batch_() as batcher: - for _, transaction in transactions_to_process.items(): - entity_id = transaction.get(entity_id_field) + # Sort the transactions to minimize the probability of deadlocks by reducing the chance of concurrent + # trasactions locking the same rows/ranges in different orders. + transactions_to_process = dict( + sorted( + daily_spend_transactions.items(), + # Normally to avoid deadlocks we would sort by the index, but since we have sprinkled indexes + # on our schema like we're discount Salt Bae, we just sort by all fields that have an index, + # in an ad-hoc (but hopefully sensible) order of indexes. The actual ordering matters less than + # ensuring that all concurrent transactions sort in the same order. + # We could in theory use the dict key, as it contains basically the same fields, but this is more + # robust to future changes in the key format. + # If _update_daily_spend ever gets the ability to write to multiple tables at once, the sorting + # should sort by the table first. + key=lambda x: ( + x[1].get("date") or "", + x[1].get(entity_id_field) or "", + x[1].get("api_key") or "", + x[1].get("model") or "", + x[1].get("custom_llm_provider") or "", + ), + )[:BATCH_SIZE] + ) - # Construct the where clause dynamically - where_clause = { - unique_constraint_name: { + if len(transactions_to_process) == 0: + verbose_proxy_logger.debug( + f"No new transactions to process for daily {entity_type} spend update" + ) + return + + try: + async with prisma_client.db.batch_() as batcher: + for _, transaction in transactions_to_process.items(): + entity_id = transaction.get(entity_id_field) + + # Construct the where clause dynamically + where_clause = { + unique_constraint_name: { + entity_id_field: entity_id, + "date": transaction["date"], + "api_key": transaction["api_key"], + "model": transaction["model"], + "custom_llm_provider": transaction.get( + "custom_llm_provider" + ) + or "", + "mcp_namespaced_tool_name": transaction.get( + "mcp_namespaced_tool_name" + ) + or "", + "endpoint": transaction.get("endpoint") + or "", + } + } + + # Get the table dynamically + table = getattr(batcher, table_name) + + # Common data structure for both create and update + common_data = { entity_id_field: entity_id, "date": transaction["date"], "api_key": transaction["api_key"], - "model": transaction["model"], - "custom_llm_provider": transaction.get( - "custom_llm_provider" - ) - or "", + "model": transaction.get("model"), + "model_group": transaction.get("model_group"), "mcp_namespaced_tool_name": transaction.get( "mcp_namespaced_tool_name" ) or "", + "custom_llm_provider": transaction.get( + "custom_llm_provider" + ), "endpoint": transaction.get("endpoint") or "", - } - } - - # Get the table dynamically - table = getattr(batcher, table_name) - - # Common data structure for both create and update - common_data = { - entity_id_field: entity_id, - "date": transaction["date"], - "api_key": transaction["api_key"], - "model": transaction.get("model"), - "model_group": transaction.get("model_group"), - "mcp_namespaced_tool_name": transaction.get( - "mcp_namespaced_tool_name" - ) - or "", - "custom_llm_provider": transaction.get( - "custom_llm_provider" - ), - "endpoint": transaction.get("endpoint") or "", - "prompt_tokens": transaction["prompt_tokens"], - "completion_tokens": transaction[ - "completion_tokens" - ], - "spend": transaction["spend"], - "api_requests": transaction["api_requests"], - "successful_requests": transaction[ - "successful_requests" - ], - "failed_requests": transaction["failed_requests"], - } - - # Add cache-related fields if they exist - if "cache_read_input_tokens" in transaction: - common_data["cache_read_input_tokens"] = ( - transaction.get("cache_read_input_tokens", 0) - ) - if "cache_creation_input_tokens" in transaction: - common_data["cache_creation_input_tokens"] = ( - transaction.get( - "cache_creation_input_tokens", 0 - ) - ) - - if entity_type == "tag" and "request_id" in transaction: - common_data["request_id"] = transaction.get( - "request_id" - ) - - # Create update data structure - update_data = { - "prompt_tokens": { - "increment": transaction["prompt_tokens"] - }, - "completion_tokens": { - "increment": transaction["completion_tokens"] - }, - "spend": {"increment": transaction["spend"]}, - "api_requests": { - "increment": transaction["api_requests"] - }, - "successful_requests": { - "increment": transaction["successful_requests"] - }, - "failed_requests": { - "increment": transaction["failed_requests"] - }, - } - - # Add cache-related fields to update if they exist - if "cache_read_input_tokens" in transaction: - update_data["cache_read_input_tokens"] = { - "increment": transaction.get( - "cache_read_input_tokens", 0 - ) - } - if "cache_creation_input_tokens" in transaction: - update_data["cache_creation_input_tokens"] = { - "increment": transaction.get( - "cache_creation_input_tokens", 0 - ) + "prompt_tokens": transaction["prompt_tokens"], + "completion_tokens": transaction[ + "completion_tokens" + ], + "spend": transaction["spend"], + "api_requests": transaction["api_requests"], + "successful_requests": transaction[ + "successful_requests" + ], + "failed_requests": transaction[ + "failed_requests" + ], } - if entity_type == "tag" and "request_id" in transaction: - update_data["request_id"] = transaction.get( - "request_id" + # Add cache-related fields if they exist + if "cache_read_input_tokens" in transaction: + common_data["cache_read_input_tokens"] = ( + transaction.get( + "cache_read_input_tokens", 0 + ) + ) + if "cache_creation_input_tokens" in transaction: + common_data["cache_creation_input_tokens"] = ( + transaction.get( + "cache_creation_input_tokens", 0 + ) + ) + + if ( + entity_type == "tag" + and "request_id" in transaction + ): + common_data["request_id"] = transaction.get( + "request_id" + ) + + # Create update data structure + update_data = { + "prompt_tokens": { + "increment": transaction["prompt_tokens"] + }, + "completion_tokens": { + "increment": transaction[ + "completion_tokens" + ] + }, + "spend": {"increment": transaction["spend"]}, + "api_requests": { + "increment": transaction["api_requests"] + }, + "successful_requests": { + "increment": transaction[ + "successful_requests" + ] + }, + "failed_requests": { + "increment": transaction["failed_requests"] + }, + } + + # Add cache-related fields to update if they exist + if "cache_read_input_tokens" in transaction: + update_data["cache_read_input_tokens"] = { + "increment": transaction.get( + "cache_read_input_tokens", 0 + ) + } + if "cache_creation_input_tokens" in transaction: + update_data["cache_creation_input_tokens"] = { + "increment": transaction.get( + "cache_creation_input_tokens", 0 + ) + } + + if ( + entity_type == "tag" + and "request_id" in transaction + ): + update_data["request_id"] = transaction.get( + "request_id" + ) + + # Add endpoint to update_data so existing rows get their endpoint field updated + update_data["endpoint"] = ( + transaction.get("endpoint") or "" ) - # Add endpoint to update_data so existing rows get their endpoint field updated - update_data["endpoint"] = ( - transaction.get("endpoint") or "" - ) + table.upsert( + where=where_clause, + data={ + "create": common_data, + "update": update_data, + }, + ) + except Exception as batch_error: + # Log detailed error information for debugging batch upsert failures + # This helps diagnose issues like unique constraint violations + spend_log_error( + "Daily %s spend batch upsert failed. " + "Table: %s, Constraint: %s, Batch size: %d, Error: %s", + entity_type, + table_name, + unique_constraint_name, + len(transactions_to_process), + str(batch_error), + exc=batch_error, + ) + raise - table.upsert( - where=where_clause, - data={ - "create": common_data, - "update": update_data, - }, - ) - except Exception as batch_error: - # Log detailed error information for debugging batch upsert failures - # This helps diagnose issues like unique constraint violations - spend_log_error( - "Daily %s spend batch upsert failed. " - "Table: %s, Constraint: %s, Batch size: %d, Error: %s", - entity_type, - table_name, - unique_constraint_name, - len(transactions_to_process), - str(batch_error), - exc=batch_error, + verbose_proxy_logger.debug( + f"Processed {len(transactions_to_process)} daily {entity_type} transactions in {time.time() - start_time:.2f}s" ) - raise - verbose_proxy_logger.debug( - f"Processed {len(transactions_to_process)} daily {entity_type} transactions in {time.time() - start_time:.2f}s" - ) + # Remove processed transactions + for key in transactions_to_process.keys(): + daily_spend_transactions.pop(key, None) - # Remove processed transactions - for key in transactions_to_process.keys(): - daily_spend_transactions.pop(key, None) + break - break - - except DB_CONNECTION_ERROR_TYPES as e: - if i >= n_retry_times: - _raise_failed_update_spend_exception( - e=e, - start_time=start_time, - proxy_logging_obj=proxy_logging_obj, + except DB_CONNECTION_ERROR_TYPES as e: + if i >= n_retry_times: + _raise_failed_update_spend_exception( + e=e, + start_time=start_time, + proxy_logging_obj=proxy_logging_obj, + ) + await asyncio.sleep( + # Sleep a random amount to avoid retrying and deadlocking again: when two transactions deadlock they are + # cancelled basically at the same time, so if they wait the same time they will also retry at the same time + # and thus they are more likely to deadlock again. + # Instead, we sleep a random amount so that they retry at slightly different times, lowering the chance of + # repeated deadlocks, and therefore of exceeding the retry limit. + random.uniform(2**i, 2 ** (i + 1)) ) - await asyncio.sleep( - # Sleep a random amount to avoid retrying and deadlocking again: when two transactions deadlock they are - # cancelled basically at the same time, so if they wait the same time they will also retry at the same time - # and thus they are more likely to deadlock again. - # Instead, we sleep a random amount so that they retry at slightly different times, lowering the chance of - # repeated deadlocks, and therefore of exceeding the retry limit. - random.uniform(2**i, 2 ** (i + 1)) - ) except Exception as e: if "transactions_to_process" in locals(): diff --git a/litellm/proxy/db/db_transaction_queue/pod_lock_manager.py b/litellm/proxy/db/db_transaction_queue/pod_lock_manager.py index 5e0ddef9eaa..2cbc0646567 100644 --- a/litellm/proxy/db/db_transaction_queue/pod_lock_manager.py +++ b/litellm/proxy/db/db_transaction_queue/pod_lock_manager.py @@ -1,4 +1,5 @@ import asyncio +import json from litellm._uuid import uuid from typing import TYPE_CHECKING, Any, Optional @@ -167,8 +168,11 @@ end self._release_lock_script = script_register( self._COMPARE_AND_DELETE_LOCK_SCRIPT ) + # acquire_lock stores the pod_id via async_set_cache, which + # JSON-encodes the value; compare against the same encoding so + # the Lua equality check matches and the lock is released result = await self._release_lock_script( - keys=[lock_key], args=[self.pod_id] + keys=[lock_key], args=[json.dumps(self.pod_id)] ) return int(result or 0) except Exception: diff --git a/litellm/proxy/dev_config.yaml b/litellm/proxy/dev_config.yaml new file mode 100644 index 00000000000..e437ed7a118 --- /dev/null +++ b/litellm/proxy/dev_config.yaml @@ -0,0 +1,191 @@ +model_list: + # ---------- Anthropic native ---------- + - model_name: anthropic-haiku-4-5 + litellm_params: + model: anthropic/claude-haiku-4-5 + api_key: os.environ/ANTHROPIC_API_KEY + - model_name: anthropic-sonnet-4-5 + litellm_params: + model: anthropic/claude-sonnet-4-5 + api_key: os.environ/ANTHROPIC_API_KEY + - model_name: anthropic-opus-4-5 + litellm_params: + model: anthropic/claude-opus-4-5 + api_key: os.environ/ANTHROPIC_API_KEY + - model_name: anthropic-sonnet-4-6 + litellm_params: + model: anthropic/claude-sonnet-4-6 + api_key: os.environ/ANTHROPIC_API_KEY + - model_name: anthropic-opus-4-6 + litellm_params: + model: anthropic/claude-opus-4-6 + api_key: os.environ/ANTHROPIC_API_KEY + - model_name: anthropic-opus-4-7 + litellm_params: + model: anthropic/claude-opus-4-7 + api_key: os.environ/ANTHROPIC_API_KEY + - model_name: anthropic-opus-4-8 + litellm_params: + model: anthropic/claude-opus-4-8 + api_key: os.environ/ANTHROPIC_API_KEY + + # ---------- Bedrock Invoke ---------- + - model_name: bedrock-invoke-haiku-4-5 + litellm_params: + model: bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0 + aws_region_name: us-east-1 + - model_name: bedrock-invoke-sonnet-4-5 + litellm_params: + model: bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0 + aws_region_name: us-east-1 + - model_name: bedrock-invoke-opus-4-5 + litellm_params: + model: bedrock/us.anthropic.claude-opus-4-5-20251101-v1:0 + aws_region_name: us-east-1 + - model_name: bedrock-invoke-sonnet-4-6 + litellm_params: + model: bedrock/us.anthropic.claude-sonnet-4-6 + aws_region_name: us-east-1 + - model_name: bedrock-invoke-opus-4-6 + litellm_params: + model: bedrock/us.anthropic.claude-opus-4-6-v1 + aws_region_name: us-east-1 + - model_name: bedrock-invoke-opus-4-7 + litellm_params: + model: bedrock/global.anthropic.claude-opus-4-7 + aws_region_name: us-east-1 + - model_name: bedrock-invoke-opus-4-8 + litellm_params: + model: bedrock/global.anthropic.claude-opus-4-8 + aws_region_name: us-east-1 + + # ---------- Bedrock Converse ---------- + - model_name: bedrock-converse-haiku-4-5 + litellm_params: + model: bedrock/converse/us.anthropic.claude-haiku-4-5-20251001-v1:0 + aws_region_name: us-east-1 + - model_name: bedrock-converse-sonnet-4-5 + litellm_params: + model: bedrock/converse/us.anthropic.claude-sonnet-4-5-20250929-v1:0 + aws_region_name: us-east-1 + - model_name: bedrock-converse-opus-4-5 + litellm_params: + model: bedrock/converse/us.anthropic.claude-opus-4-5-20251101-v1:0 + aws_region_name: us-east-1 + - model_name: bedrock-converse-sonnet-4-6 + litellm_params: + model: bedrock/converse/us.anthropic.claude-sonnet-4-6 + aws_region_name: us-east-1 + - model_name: bedrock-converse-opus-4-6 + litellm_params: + model: bedrock/converse/us.anthropic.claude-opus-4-6-v1 + aws_region_name: us-east-1 + - model_name: bedrock-converse-opus-4-7 + litellm_params: + model: bedrock/converse/global.anthropic.claude-opus-4-7 + aws_region_name: us-east-1 + - model_name: bedrock-converse-opus-4-8 + litellm_params: + model: bedrock/converse/global.anthropic.claude-opus-4-8 + aws_region_name: us-east-1 + + # ---------- Vertex AI (Anthropic on Vertex) ---------- + - model_name: vertex-haiku-4-5 + litellm_params: + model: vertex_ai/claude-haiku-4-5@20251001 + vertex_project: os.environ/VERTEXAI_PROJECT + vertex_location: global + - model_name: vertex-sonnet-4-5 + litellm_params: + model: vertex_ai/claude-sonnet-4-5@20250929 + vertex_project: os.environ/VERTEXAI_PROJECT + vertex_location: global + - model_name: vertex-opus-4-5 + litellm_params: + model: vertex_ai/claude-opus-4-5@20251101 + vertex_project: os.environ/VERTEXAI_PROJECT + vertex_location: global + - model_name: vertex-sonnet-4-6 + litellm_params: + model: vertex_ai/claude-sonnet-4-6 + vertex_project: os.environ/VERTEXAI_PROJECT + vertex_location: global + - model_name: vertex-opus-4-6 + litellm_params: + model: vertex_ai/claude-opus-4-6 + vertex_project: os.environ/VERTEXAI_PROJECT + vertex_location: global + - model_name: vertex-opus-4-7 + litellm_params: + model: vertex_ai/claude-opus-4-7 + vertex_project: os.environ/VERTEXAI_PROJECT + vertex_location: global + - model_name: vertex-opus-4-8 + litellm_params: + model: vertex_ai/claude-opus-4-8 + vertex_project: os.environ/VERTEXAI_PROJECT + vertex_location: global + + # ---------- Gemini Enterprise Agent Platform ---------- + - model_name: gemini-claude-code + litellm_params: + model: vertex_ai/gemini-2.5-pro + vertex_project: os.environ/VERTEXAI_PROJECT + vertex_location: global + vertex_credentials: os.environ/GEMINI_CLAUDE_CODE_VERTEX_CREDENTIALS + extra_body: + labels: + workload: claude-code + source: litellm + environment: internal + reconciliation_group: claude-code-gemini + + # ---------- Azure AI Foundry (Anthropic on Azure) ---------- + - model_name: azure-haiku-4-5 + litellm_params: + model: azure_ai/claude-haiku-4-5 + api_base: os.environ/AZURE_AI_API_BASE + api_key: os.environ/AZURE_AI_API_KEY + - model_name: azure-sonnet-4-5 + litellm_params: + model: azure_ai/claude-sonnet-4-5 + api_base: os.environ/AZURE_AI_API_BASE + api_key: os.environ/AZURE_AI_API_KEY + - model_name: azure-opus-4-5 + litellm_params: + model: azure_ai/claude-opus-4-5 + api_base: os.environ/AZURE_AI_API_BASE + api_key: os.environ/AZURE_AI_API_KEY + - model_name: azure-sonnet-4-6 + litellm_params: + model: azure_ai/claude-sonnet-4-6 + api_base: os.environ/AZURE_AI_API_BASE + api_key: os.environ/AZURE_AI_API_KEY + - model_name: azure-opus-4-6 + litellm_params: + model: azure_ai/claude-opus-4-6 + api_base: os.environ/AZURE_AI_API_BASE + api_key: os.environ/AZURE_AI_API_KEY + - model_name: azure-opus-4-7 + litellm_params: + model: azure_ai/claude-opus-4-7 + api_base: os.environ/AZURE_AI_API_BASE + api_key: os.environ/AZURE_AI_API_KEY + - model_name: azure-opus-4-8 + litellm_params: + model: azure_ai/claude-opus-4-8 + api_base: os.environ/AZURE_AI_API_BASE + api_key: os.environ/AZURE_AI_API_KEY + + # ---------- OpenAI ---------- + - model_name: gpt-5.5 + litellm_params: + model: openai/gpt-5.5 + api_key: os.environ/OPENAI_API_KEY + +general_settings: + master_key: sk-1234 + +litellm_settings: + drop_params: True + telemetry: False diff --git a/litellm/proxy/discovery_endpoints/ui_discovery_endpoints.py b/litellm/proxy/discovery_endpoints/ui_discovery_endpoints.py index 233df5c6c57..0e7e67aa37f 100644 --- a/litellm/proxy/discovery_endpoints/ui_discovery_endpoints.py +++ b/litellm/proxy/discovery_endpoints/ui_discovery_endpoints.py @@ -25,6 +25,10 @@ async def get_ui_config(): or general_settings.get("auto_redirect_ui_login_to_sso", False) is True ) admin_ui_disabled = os.getenv("DISABLE_ADMIN_UI", "false").lower() == "true" + hide_default_credentials_hint = bool( + os.getenv("LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT", "false").lower() == "true" + or general_settings.get("hide_default_credentials_hint", False) is True + ) sso_configured = _has_user_setup_sso() @@ -38,6 +42,7 @@ async def get_ui_config(): auto_redirect_to_sso=sso_configured and auto_redirect_ui_login_to_sso, admin_ui_disabled=admin_ui_disabled, sso_configured=sso_configured, + hide_default_credentials_hint=hide_default_credentials_hint, is_control_plane=is_control_plane, workers=proxy_config.worker_registry if is_control_plane else [], ) diff --git a/litellm/proxy/example_config_yaml/disable_schema_update.yaml b/litellm/proxy/example_config_yaml/disable_schema_update.yaml index 5dcbd0dbd57..6c1f535f8f3 100644 --- a/litellm/proxy/example_config_yaml/disable_schema_update.yaml +++ b/litellm/proxy/example_config_yaml/disable_schema_update.yaml @@ -3,12 +3,12 @@ model_list: litellm_params: model: openai/fake api_key: fake-key - api_base: https://exampleopenaiendpoint-production.up.railway.app/ + api_base: os.environ/FAKE_OPENAI_API_BASE - model_name: gpt-4 litellm_params: model: openai/gpt-4 api_key: fake-key - api_base: https://exampleopenaiendpoint-production.up.railway.app/ + api_base: os.environ/FAKE_OPENAI_API_BASE litellm_settings: callbacks: ["gcs_bucket"] diff --git a/litellm/proxy/example_config_yaml/enterprise_config.yaml b/litellm/proxy/example_config_yaml/enterprise_config.yaml index 337e85177e5..037d31009ea 100644 --- a/litellm/proxy/example_config_yaml/enterprise_config.yaml +++ b/litellm/proxy/example_config_yaml/enterprise_config.yaml @@ -3,7 +3,7 @@ model_list: litellm_params: model: openai/fake api_key: fake-key - api_base: https://exampleopenaiendpoint-production.up.railway.app/ + api_base: os.environ/FAKE_OPENAI_API_BASE tags: ["teamA"] model_info: id: "team-a-model" diff --git a/litellm/proxy/example_config_yaml/multi_instance_simple_config.yaml b/litellm/proxy/example_config_yaml/multi_instance_simple_config.yaml index f83160a7a22..b353924f000 100644 --- a/litellm/proxy/example_config_yaml/multi_instance_simple_config.yaml +++ b/litellm/proxy/example_config_yaml/multi_instance_simple_config.yaml @@ -3,7 +3,7 @@ model_list: litellm_params: model: openai/my-fake-model api_key: my-fake-key - api_base: https://exampleopenaiendpoint-production.up.railway.app/ + api_base: os.environ/FAKE_OPENAI_API_BASE litellm_settings: cache: True diff --git a/litellm/proxy/example_config_yaml/otel_test_config.yaml b/litellm/proxy/example_config_yaml/otel_test_config.yaml index 7f18e513437..2ebadcbc167 100644 --- a/litellm/proxy/example_config_yaml/otel_test_config.yaml +++ b/litellm/proxy/example_config_yaml/otel_test_config.yaml @@ -3,7 +3,7 @@ model_list: litellm_params: model: openai/gpt-5-mini api_key: fake-key - api_base: https://exampleopenaiendpoint-production.up.railway.app/ + api_base: os.environ/FAKE_OPENAI_API_BASE tags: ["teamA"] model_info: id: "team-a-model" @@ -11,7 +11,7 @@ model_list: litellm_params: model: openai/gpt-5-mini api_key: fake-key - api_base: https://exampleopenaiendpoint-production.up.railway.app/ + api_base: os.environ/FAKE_OPENAI_API_BASE tags: ["teamB"] model_info: id: "team-b-model" @@ -24,7 +24,7 @@ model_list: litellm_params: model: openai/429 api_key: fake-key - api_base: https://exampleopenaiendpoint-production.up.railway.app + api_base: os.environ/FAKE_OPENAI_API_BASE - model_name: llava-hf litellm_params: model: openai/llava-hf/llava-v1.6-vicuna-7b-hf @@ -35,12 +35,12 @@ model_list: - model_name: bedrock/* litellm_params: model: bedrock/* - api_base: https://exampleopenaiendpoint-production.up.railway.app/ + api_base: os.environ/FAKE_OPENAI_API_BASE - model_name: openai/* litellm_params: model: openai/* api_key: os.environ/OPENAI_API_KEY - api_base: https://exampleopenaiendpoint-production.up.railway.app/ + api_base: os.environ/FAKE_OPENAI_API_BASE litellm_settings: diff --git a/litellm/proxy/example_config_yaml/spend_tracking_config.yaml b/litellm/proxy/example_config_yaml/spend_tracking_config.yaml index dfed2194b58..60adadbd8d4 100644 --- a/litellm/proxy/example_config_yaml/spend_tracking_config.yaml +++ b/litellm/proxy/example_config_yaml/spend_tracking_config.yaml @@ -3,7 +3,7 @@ model_list: litellm_params: model: openai/gpt-5-mini api_key: fake-key - api_base: https://exampleopenaiendpoint-production.up.railway.app/ + api_base: os.environ/FAKE_OPENAI_API_BASE general_settings: use_redis_transaction_buffer: true diff --git a/litellm/proxy/example_config_yaml/store_model_db_config.yaml b/litellm/proxy/example_config_yaml/store_model_db_config.yaml index b9cd2302046..5b77a53b4b7 100644 --- a/litellm/proxy/example_config_yaml/store_model_db_config.yaml +++ b/litellm/proxy/example_config_yaml/store_model_db_config.yaml @@ -3,7 +3,7 @@ model_list: litellm_params: model: openai/my-fake-model api_key: my-fake-key - api_base: https://exampleopenaiendpoint-production.up.railway.app/ + api_base: os.environ/FAKE_OPENAI_API_BASE general_settings: store_model_in_db: true diff --git a/litellm/proxy/google_endpoints/endpoints.py b/litellm/proxy/google_endpoints/endpoints.py index cc20f0cf3b3..6427835c250 100644 --- a/litellm/proxy/google_endpoints/endpoints.py +++ b/litellm/proxy/google_endpoints/endpoints.py @@ -107,6 +107,7 @@ async def google_stream_generate_content( data["stream"] = True # google-genai SDK (?alt=sse) must not receive OpenAI's data: [DONE] terminator. data["_litellm_skip_openai_stream_done"] = True + data["_litellm_raw_sse_stream"] = True processor = ProxyBaseLLMRequestProcessing(data=data) try: diff --git a/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py b/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py index 5b5f91195e7..d70c8e4f310 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py +++ b/litellm/proxy/guardrails/guardrail_hooks/aim/aim.py @@ -9,7 +9,6 @@ import json import os from typing import TYPE_CHECKING, Any, AsyncGenerator, Optional, Type, Union -from fastapi import HTTPException from pydantic import BaseModel from websockets.asyncio.client import ClientConnection, connect @@ -21,7 +20,7 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, ) -from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy._types import ProxyException, UserAPIKeyAuth from litellm.proxy.guardrails._content_utils import ( apply_redacted_messages_back, build_inspection_messages, @@ -129,6 +128,16 @@ class AimGuardrail(CustomGuardrail): verbose_proxy_logger.error(f"Aim: {action_type} action") return data + @staticmethod + def _rejection(message: str, *, openai_code: str | None = None) -> ProxyException: + return ProxyException( + message=message, + type="invalid_request_error", + param=None, + code=400, + openai_code=openai_code, + ) + def _handle_block_action(self, analysis_result: Any, required_action: Any) -> None: detection_message = required_action.get("detection_message", None) verbose_proxy_logger.info( @@ -136,7 +145,7 @@ class AimGuardrail(CustomGuardrail): policies=list(analysis_result["policy_drill_down"].keys()), ), ) - raise HTTPException(status_code=400, detail=detection_message) + raise self._rejection(detection_message, openai_code="content_policy_violation") def _anonymize_request(self, res: Any, data: dict) -> dict: verbose_proxy_logger.info("Aim: anonymize action") @@ -148,14 +157,11 @@ class AimGuardrail(CustomGuardrail): # parts from a multimodal request — degrade to block so the # multimodal payload is never silently rewritten. if has_non_string_content(data): - raise HTTPException( - status_code=400, - detail=( - "Aim: anonymize action requested for multimodal input " - "but mask-in-place would drop non-text parts. Send the " - "request with plain string content to use anonymize, " - "or rely on block-mode policies." - ), + raise self._rejection( + "Aim: anonymize action requested for multimodal input " + "but mask-in-place would drop non-text parts. Send the " + "request with plain string content to use anonymize, " + "or rely on block-mode policies." ) redacted_messages = [ { @@ -287,9 +293,9 @@ class AimGuardrail(CustomGuardrail): if aim_output_guardrail_result and aim_output_guardrail_result.get( "detection_message" ): - raise HTTPException( - status_code=400, - detail=aim_output_guardrail_result.get("detection_message"), + raise self._rejection( + aim_output_guardrail_result.get("detection_message"), + openai_code="content_policy_violation", ) if aim_output_guardrail_result and aim_output_guardrail_result.get( "redacted_output" diff --git a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py index ff802223f21..72b9b7dc3c1 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py +++ b/litellm/proxy/guardrails/guardrail_hooks/lakera_ai.py @@ -121,7 +121,7 @@ class lakeraAI_Moderation(CustomGuardrail): return None - async def _check( # noqa: PLR0915 + async def _check( self, data: dict, user_api_key_dict: UserAPIKeyAuth, diff --git a/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py b/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py index 7e6f3dac008..b3b8fbdb2a5 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py +++ b/litellm/proxy/guardrails/guardrail_hooks/openai/moderations.py @@ -25,7 +25,11 @@ from litellm.llms.custom_httpx.http_handler import ( httpxSpecialProvider, ) from litellm.types.guardrails import GuardrailEventHooks -from litellm.types.utils import GenericGuardrailAPIInputs, GuardrailStatus +from litellm.types.utils import ( + GenericGuardrailAPIInputs, + GuardrailStatus, + GuardrailTracingDetail, +) from .base import OpenAIGuardrailBase @@ -287,6 +291,7 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail): start_time=start_time, end_time=end_time, event_type=event_type, + tracing_detail=self._build_tracing_detail(guardrail_response), ) return response @@ -328,9 +333,36 @@ class OpenAIModerationGuardrail(OpenAIGuardrailBase, CustomGuardrail): start_time=start_time, end_time=end_time, event_type=event_type, + tracing_detail=self._build_tracing_detail(guardrail_response), ) raise e + @staticmethod + def _build_tracing_detail( + guardrail_response: Union[dict, str, Exception], + ) -> Optional[GuardrailTracingDetail]: + """ + Pull the flagged category names out of the moderation response so trace + backends can index a short, queryable ``guardrail_violation_categories`` + attribute instead of the full ``guardrail_response`` blob, whose + ``category_scores`` map (one float per category) blows past indexed-field + length limits on backends like ELK (1024 chars). + """ + if not isinstance(guardrail_response, dict): + return None + + results = guardrail_response.get("results") or [] + violation_categories = [ + category + for result in results + if isinstance(result, dict) + for category, is_flagged in (result.get("categories") or {}).items() + if is_flagged + ] + if not violation_categories: + return None + return GuardrailTracingDetail(violation_categories=violation_categories) + @staticmethod def get_config_model() -> Optional[Type["GuardrailConfigModel"]]: """ diff --git a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py index e5200394b55..e8887fa712a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py +++ b/litellm/proxy/guardrails/guardrail_hooks/panw_prisma_airs/panw_prisma_airs.py @@ -261,7 +261,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): ) return "" - async def _call_panw_api( # noqa: PLR0915 + async def _call_panw_api( self, content: str = "", is_response: bool = False, @@ -1762,7 +1762,7 @@ class PanwPrismaAirsHandler(CustomGuardrail): return rd.get("name") if ("arguments" in rd or "mcp_arguments" in rd) else None @log_guardrail_information - async def apply_guardrail( # noqa: PLR0915 + async def apply_guardrail( self, inputs: GenericGuardrailAPIInputs, request_data: dict, diff --git a/litellm/proxy/guardrails/guardrail_hooks/presidio.py b/litellm/proxy/guardrails/guardrail_hooks/presidio.py index 7d6d1adb05e..a8afe4efe2a 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/presidio.py +++ b/litellm/proxy/guardrails/guardrail_hooks/presidio.py @@ -741,6 +741,18 @@ class _OPTIONAL_PresidioPIIMasking(CustomGuardrail): For multiple messages in /chat/completions, we'll need to call them in parallel. """ + # Respect the configured event hook. In `logging_only` mode (and any config that + # excludes pre_call) the live request must not be masked - masking is applied to a + # copy at logging time via `async_logging_hook`. Without this gate the request sent + # to the model would carry anonymization tokens and the response would echo them. + if ( + self.should_run_guardrail( + data=data, + event_type=GuardrailEventHooks.pre_call, + ) + is not True + ): + return data try: content_safety = data.get("content_safety", None) diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 2a2c758fa8a..09fff71062b 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -288,7 +288,7 @@ class UnifiedLLMGuardrails(CustomLogger): return response - async def async_post_call_streaming_iterator_hook( # noqa: PLR0915 + async def async_post_call_streaming_iterator_hook( self, user_api_key_dict: UserAPIKeyAuth, response: Any, diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index 4a28143e617..be51234e7bc 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -6,6 +6,7 @@ import random import sys import threading import time +from collections.abc import Mapping from typing import List, Optional import litellm @@ -42,23 +43,50 @@ MINIMAL_DISPLAY_PARAMS = ["model", "mode_error"] # endpoints that reject unknown fields with 400 "Unknown parameter: # 'max_tokens'". Allow-list so new modes are safe by default. # Per-deployment override: `model_info.health_check_supports_max_tokens`. -_MAX_TOKEN_SUPPORT_MODES: frozenset = frozenset({"chat", "completion", "responses"}) +_MAX_TOKEN_SUPPORT_MODES: frozenset[str] = frozenset( + {"chat", "completion", "responses"} +) -def _should_inject_health_check_max_tokens(model_info: dict) -> bool: +def _resolve_health_check_mode( + model_info: Mapping[str, object], litellm_params: Mapping[str, object] +) -> str | None: + """ + Effective mode for a deployment's health-check probe. + + Prefers operator-set `model_info.mode`; otherwise resolves it from the model + cost map, which understands `bedrock/` and cross-region inference-profile + prefixes (`us.`, `eu.`, `apac.`). Without this, non-chat Bedrock deployments + (e.g. embeddings) are probed as chat, so `max_tokens` is injected and the + request 400s on "extraneous key [max_tokens]". + """ + explicit_mode = model_info.get("mode") + if isinstance(explicit_mode, str): + return explicit_mode + model = litellm_params.get("model") + if not isinstance(model, str): + return None + try: + return litellm.get_model_info(model=model).get("mode") + except Exception: + return None + + +def _should_inject_health_check_max_tokens( + model_info: Mapping[str, object], mode: str | None +) -> bool: """ Whether the health-check probe should include `max_tokens`. Order: 1. `model_info.health_check_supports_max_tokens` (operator override). - 2. `_MAX_TOKEN_SUPPORT_MODES`. Missing `mode` is treated as `chat` + 2. `_MAX_TOKEN_SUPPORT_MODES`. An unresolvable mode is treated as `chat` for backward compatibility. """ explicit = model_info.get("health_check_supports_max_tokens") if explicit is not None: return bool(explicit) - mode = model_info.get("mode") or "chat" - return mode in _MAX_TOKEN_SUPPORT_MODES + return (mode or "chat") in _MAX_TOKEN_SUPPORT_MODES # Health-check modes that forward `reasoning_effort` to the provider (chat-style calls). @@ -165,7 +193,9 @@ async def run_with_timeout(task, timeout): async def _run_model_health_check(model: dict): litellm_params = model["litellm_params"] model_info = model.get("model_info", {}) - mode = model_info.get("mode", None) + mode = _resolve_health_check_mode( + model_info, litellm_params # any-ok: untyped router config dict + ) litellm_params = _update_litellm_params_for_health_check(model_info, litellm_params) timeout = model_info.get("health_check_timeout") or HEALTH_CHECK_TIMEOUT_SECONDS @@ -421,10 +451,15 @@ def _update_litellm_params_for_health_check( reject unknown fields with 400 "Unknown parameter: 'max_tokens'". - updates the `model` param with the `health_check_model` if it exists Doc: https://docs.litellm.ai/docs/proxy/health#wildcard-routes - updates the `voice` param with the `health_check_voice` for `audio_speech` mode if it exists Doc: https://docs.litellm.ai/docs/proxy/health#text-to-speech-models - - for Bedrock models with region routing (bedrock/region/model), strips the litellm routing prefix but preserves the model ID + - for Bedrock models with region routing (bedrock/region/model), strips the litellm routing prefix but preserves the model ID, and pins `custom_llm_provider` to `bedrock` (only when the deployment hasn't already set one, so an explicit `bedrock_converse` survives) so the bare model id still resolves to the provider (e.g. cross-region ids like `us.cohere.embed-v4:0`) """ + mode = _resolve_health_check_mode( + model_info, litellm_params # any-ok: untyped router config dict + ) litellm_params["messages"] = _get_random_llm_message() - if _should_inject_health_check_max_tokens(model_info): + if _should_inject_health_check_max_tokens( + model_info, mode # any-ok: untyped router config dict + ): _resolved_max_tokens = _resolve_health_check_max_tokens( model_info, litellm_params ) @@ -432,7 +467,7 @@ def _update_litellm_params_for_health_check( litellm_params["max_tokens"] = _resolved_max_tokens # Per-model reasoning effort for health checks only (e.g. reasoning_effort=none). - if model_info.get("mode", None) in _HEALTH_CHECK_MODES_SUPPORTING_REASONING_EFFORT: + if mode in _HEALTH_CHECK_MODES_SUPPORTING_REASONING_EFFORT: _hc_reasoning_effort = model_info.get("health_check_reasoning_effort", None) if _hc_reasoning_effort is not None: litellm_params["reasoning_effort"] = _hc_reasoning_effort @@ -440,7 +475,7 @@ def _update_litellm_params_for_health_check( _health_check_model = model_info.get("health_check_model", None) if _health_check_model is not None: litellm_params["model"] = _health_check_model - if model_info.get("mode", None) == "audio_speech": + if mode == "audio_speech": litellm_params["voice"] = model_info.get("health_check_voice", "alloy") # Handle Bedrock region routing format: bedrock/region/model @@ -477,6 +512,10 @@ def _update_litellm_params_for_health_check( model = "/".join(filtered_parts) litellm_params["model"] = model + if not litellm_params.get("custom_llm_provider"): # any-ok: untyped router dict + litellm_params["custom_llm_provider"] = ( # any-ok: untyped router dict + "bedrock" + ) return litellm_params diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index e0d018d4344..8a432eb2f42 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -167,7 +167,7 @@ async def test_endpoint(request: Request): tags=["health"], dependencies=[Depends(user_api_key_auth)], ) -async def health_services_endpoint( # noqa: PLR0915 +async def health_services_endpoint( user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), service: services = fastapi.Query(description="Specify the service being hit."), ): diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index 23af23e78bd..d36e9858b5a 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -239,7 +239,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): request_count_end_user_id=results[5], ) - async def async_pre_call_hook( # noqa: PLR0915 + async def async_pre_call_hook( self, user_api_key_dict: UserAPIKeyAuth, cache: DualCache, @@ -506,9 +506,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger): return - async def async_log_success_event( # noqa: PLR0915 - self, kwargs, response_obj, start_time, end_time - ): + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): from litellm.proxy.common_utils.callback_utils import ( get_model_group_from_litellm_kwargs, ) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 0587ce1cc29..2c5937d9506 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -162,6 +162,8 @@ _UNTRUSTED_METADATA_CONTROL_FIELDS = ( "secret_fields", "_guardrail_pipelines", "_pipeline_managed_guardrails", + "client_disconnected", + "error_information", PRE_CALL_EXECUTED_GUARDRAILS_KEY, ) @@ -1317,7 +1319,7 @@ class LiteLLMProxyRequestSetup: ) -async def add_litellm_data_to_request( # noqa: PLR0915 +async def add_litellm_data_to_request( data: dict, request: Request, user_api_key_dict: UserAPIKeyAuth, diff --git a/litellm/proxy/management_endpoints/budget_management_endpoints.py b/litellm/proxy/management_endpoints/budget_management_endpoints.py index 698155a5c26..e35ec2933d0 100644 --- a/litellm/proxy/management_endpoints/budget_management_endpoints.py +++ b/litellm/proxy/management_endpoints/budget_management_endpoints.py @@ -183,10 +183,23 @@ async def update_budget( except ValueError as e: raise HTTPException(status_code=400, detail={"error": str(e)}) + # recompute budget_reset_at when the duration changes, unless the caller pinned a reset time explicitly + recomputed_reset_at = ( + { + "budget_reset_at": get_budget_reset_time( + budget_duration=budget_obj.budget_duration + ) + } + if budget_obj.budget_duration is not None + and "budget_reset_at" not in budget_obj.model_fields_set + else {} + ) + response = await BudgetRepository(prisma_client).table.update( where={"budget_id": budget_obj.budget_id}, data={ **budget_obj.model_dump(exclude_unset=True), # type: ignore + **recomputed_reset_at, "updated_by": user_api_key_dict.user_id or litellm_proxy_admin_name, }, # type: ignore ) diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index 13107b68864..341a8767db0 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -1,5 +1,5 @@ import asyncio -from datetime import datetime, timedelta +from datetime import datetime from types import SimpleNamespace from typing import Any, Callable, Dict, List, Optional, Set, Tuple, Union @@ -35,16 +35,25 @@ _PRISMA_TO_PG_TABLE: Dict[str, str] = { def update_metrics(existing_metrics: SpendMetrics, record: Any) -> SpendMetrics: - """Update metrics with new record data.""" - existing_metrics.spend += record.spend - existing_metrics.prompt_tokens += record.prompt_tokens - existing_metrics.completion_tokens += record.completion_tokens - existing_metrics.total_tokens += record.prompt_tokens + record.completion_tokens - existing_metrics.cache_read_input_tokens += record.cache_read_input_tokens - existing_metrics.cache_creation_input_tokens += record.cache_creation_input_tokens - existing_metrics.api_requests += record.api_requests - existing_metrics.successful_requests += record.successful_requests - existing_metrics.failed_requests += record.failed_requests + """Update metrics with new record data. + + Rollup rows can carry None for numeric fields when SUM() spans zero rows + (e.g. a key with no spend), so coalesce to 0 before accumulating to avoid + a TypeError. Mirrors the handling in ``_record_to_spend_metrics``. + """ + prompt_tokens = record.prompt_tokens or 0 + completion_tokens = record.completion_tokens or 0 + existing_metrics.spend += record.spend or 0.0 + existing_metrics.prompt_tokens += prompt_tokens + existing_metrics.completion_tokens += completion_tokens + existing_metrics.total_tokens += prompt_tokens + completion_tokens + existing_metrics.cache_read_input_tokens += record.cache_read_input_tokens or 0 + existing_metrics.cache_creation_input_tokens += ( + record.cache_creation_input_tokens or 0 + ) + existing_metrics.api_requests += record.api_requests or 0 + existing_metrics.successful_requests += record.successful_requests or 0 + existing_metrics.failed_requests += record.failed_requests or 0 return existing_metrics @@ -390,38 +399,24 @@ def _adjust_dates_for_timezone( timezone_offset_minutes: Optional[int], ) -> Tuple[str, str]: """ - Adjust date range to account for timezone differences. + Pass-through for the local date range; the timezone offset is intentionally ignored here. - The database stores dates in UTC. When a user in a different timezone - selects a local date range, we need to expand the UTC query range to - capture all records that fall within their local date range. + The aggregation table (e.g. LiteLLM_DailyUserSpend) stores spend in whole-UTC-day + buckets keyed on date as YYYY-MM-DD. Any conversion from a local date range to a + UTC date range using only date arithmetic must round to whole UTC days, allowing up + to 24h of slop at each boundary. The previous implementation expanded the SQL range + by an extra full UTC day on whichever side the offset pointed, which pulled in 24h + of unrelated bucket data per boundary and produced approximately 100% over-counting + on single-day queries (e.g. IST May 29 returning UTC May 28 + UTC May 29 in full). + Sums of single-day queries then exceeded the equivalent multi-day aggregate, which + is mathematically impossible. - Args: - start_date: Start date in YYYY-MM-DD format (user's local date) - end_date: End date in YYYY-MM-DD format (user's local date) - timezone_offset_minutes: Minutes behind UTC (positive = west of UTC) - This matches JavaScript's Date.getTimezoneOffset() convention. - For example: PST = +480 (8 hours * 60 = 480 minutes behind UTC) - - Returns: - Tuple of (adjusted_start_date, adjusted_end_date) in YYYY-MM-DD format + Treating the local date as the UTC date trades a small one-time boundary slop for + correct, monotonic, additive results across single-day and multi-day queries. A + later fix can introduce hour-level buckets or pro-rata weighting on adjacent UTC + days; both require data the current schema does not store. """ - if timezone_offset_minutes is None or timezone_offset_minutes == 0: - return start_date, end_date - - start = datetime.strptime(start_date, "%Y-%m-%d") - end = datetime.strptime(end_date, "%Y-%m-%d") - - if timezone_offset_minutes > 0: - # West of UTC (Americas): local evening extends into next UTC day - # e.g., Feb 4 23:59 PST = Feb 5 07:59 UTC - end = end + timedelta(days=1) - else: - # East of UTC (Asia/Europe): local morning starts in previous UTC day - # e.g., Feb 4 00:00 IST = Feb 3 18:30 UTC - start = start - timedelta(days=1) - - return start.strftime("%Y-%m-%d"), end.strftime("%Y-%m-%d") + return start_date, end_date def _build_where_conditions( @@ -726,7 +721,7 @@ def _key_metadata( return KeyMetadata(key_alias=meta.get("key_alias"), team_id=meta.get("team_id")) -def _aggregate_grouping_sets_records_sync( # noqa: PLR0915 +def _aggregate_grouping_sets_records_sync( *, records: List[Any], api_key_metadata: Dict[str, Dict[str, Any]], diff --git a/litellm/proxy/management_endpoints/common_utils.py b/litellm/proxy/management_endpoints/common_utils.py index 458cba686e6..f28bcc2bcd4 100644 --- a/litellm/proxy/management_endpoints/common_utils.py +++ b/litellm/proxy/management_endpoints/common_utils.py @@ -397,7 +397,7 @@ def _set_object_metadata_field( field_name: Name of the metadata field to set value: Value to set for the field """ - if field_name in LiteLLM_ManagementEndpoint_MetadataFields_Premium: + if field_name in LiteLLM_ManagementEndpoint_MetadataFields_Premium and value: _premium_user_check(field_name) object_data.metadata = object_data.metadata or {} @@ -563,13 +563,11 @@ def _update_metadata_field(updated_kv: dict, field_name: str) -> None: field_name: Name of the metadata field being updated """ if field_name in LiteLLM_ManagementEndpoint_MetadataFields_Premium: - value = updated_kv.get(field_name) - # Skip the premium check for empty collections ([] or {}). - # The UI sends these as defaults even when the user hasn't configured - # any enterprise features (see issue #20304). However, we still - # proceed with the update so that users can intentionally clear a - # previously-set field by sending an empty list/dict. - if value is not None and value != [] and value != {}: + # The UI sends falsy defaults (False, [], {}) even when the user has not + # enabled any enterprise feature (see #20304, #30285); require a license + # only for a truthy value. The falsy value is still persisted below so a + # previously-set field can be cleared. + if updated_kv.get(field_name): _premium_user_check() if field_name in updated_kv and updated_kv[field_name] is not None: diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index b0210e1123f..143d61a0b3a 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -18,6 +18,7 @@ import os import re import secrets import traceback +from collections.abc import Mapping from datetime import datetime, timedelta, timezone from typing import Any, Callable, Dict, List, Literal, Optional, Tuple, cast @@ -59,6 +60,9 @@ from litellm.proxy.common_utils.rbac_utils import check_org_admin_can_generate_k from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.hooks.key_management_event_hooks import KeyManagementEventHooks +from litellm.proxy.hooks.model_max_budget_limiter import ( + VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX, +) from litellm.proxy.management_endpoints.common_utils import ( _check_passthrough_routes_caller_permission, _is_user_org_admin_for_team, @@ -675,7 +679,7 @@ def _enforce_upperbound_key_params( ) -async def _common_key_generation_helper( # noqa: PLR0915 +async def _common_key_generation_helper( data: GenerateKeyRequest, user_api_key_dict: UserAPIKeyAuth, litellm_changed_by: Optional[str], @@ -1789,7 +1793,8 @@ def prepare_metadata_fields( if k in LiteLLM_ManagementEndpoint_MetadataFields_Premium: from litellm.proxy.utils import _premium_user_check - _premium_user_check(k) + if v: + _premium_user_check(k) casted_metadata[k] = v except Exception as e: @@ -3225,6 +3230,65 @@ async def delete_key_fn( raise handle_exception_on_proxy(e) +async def _get_model_max_budget_current_spend( + api_key_hash: str, + model: str, + budget_config: BudgetConfig, + user_api_key_cache: UserApiKeyCache, +) -> float: + virtual_key_model_spend_cache_key = ( + f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:" + f"{api_key_hash}:{model}:{budget_config.budget_duration}" + ) + current_spend: float | None = await user_api_key_cache.async_get_cache( + key=virtual_key_model_spend_cache_key, + ) + if current_spend is None: + model_without_prefix = model.split("/")[-1] if "/" in model else model + virtual_key_model_spend_cache_key = ( + f"{VIRTUAL_KEY_SPEND_CACHE_KEY_PREFIX}:" + f"{api_key_hash}:{model_without_prefix}:{budget_config.budget_duration}" + ) + current_spend = await user_api_key_cache.async_get_cache( + key=virtual_key_model_spend_cache_key, + ) + try: + return float(current_spend or 0.0) + except (TypeError, ValueError): + return 0.0 + + +async def _build_model_max_budget_usage( + api_key_hash: str, + model_max_budget: Mapping[str, Mapping[str, object]], + user_api_key_cache: UserApiKeyCache | None, +) -> dict[str, dict[str, object]]: + if user_api_key_cache is None or not model_max_budget: + return {} + + result: dict[str, dict[str, object]] = {} + for model, budget_info in model_max_budget.items(): + try: + budget_config = BudgetConfig.model_validate(budget_info) + if budget_config.budget_duration is None: + continue + duration_in_seconds(budget_config.budget_duration) + except Exception: # noqa: BLE001 + continue + spend = await _get_model_max_budget_current_spend( + api_key_hash=api_key_hash, + model=model, + budget_config=budget_config, + user_api_key_cache=user_api_key_cache, + ) + result[model] = { + "current_spend": round(spend, 4), + "budget_limit": budget_config.max_budget, + "time_period": budget_config.budget_duration, + } + return result + + @router.post( "/v2/key/info", tags=["key management"], @@ -3252,7 +3316,7 @@ async def info_key_fn_v2( -d {"keys": ["sk-1", "sk-2", "sk-3"]} ``` """ - from litellm.proxy.proxy_server import prisma_client + from litellm.proxy.proxy_server import prisma_client, user_api_key_cache try: if prisma_client is None: @@ -3298,7 +3362,19 @@ async def info_key_fn_v2( k_dict = k.model_dump() except Exception: k_dict = k.dict() - k_dict.pop("token", None) + k_token_hash = k_dict.pop("token", None) + + model_max_budget = k_dict.get("model_max_budget") or {} + budget_table = k_dict.get("litellm_budget_table") or {} + if not model_max_budget and isinstance(budget_table, dict): + model_max_budget = budget_table.get("model_max_budget") or {} + if model_max_budget and k_token_hash: + k_dict["model_max_budget_usage"] = await _build_model_max_budget_usage( + api_key_hash=k_token_hash, + model_max_budget=model_max_budget, + user_api_key_cache=user_api_key_cache, + ) + filtered_key_info.append(k_dict) return {"key": data.keys, "info": filtered_key_info} @@ -3336,7 +3412,7 @@ async def info_key_fn( -H "Authorization: Bearer sk-test-example-key-123" ``` """ - from litellm.proxy.proxy_server import prisma_client + from litellm.proxy.proxy_server import prisma_client, user_api_key_cache try: if prisma_client is None: @@ -3381,7 +3457,18 @@ async def info_key_fn( except Exception: # if using pydantic v1 key_info = key_info.dict() - key_info.pop("token") + key_token_hash = key_info.pop("token") + + model_max_budget = key_info.get("model_max_budget") or {} + budget_table = key_info.get("litellm_budget_table") or {} + if not model_max_budget and isinstance(budget_table, dict): + model_max_budget = budget_table.get("model_max_budget") or {} + if model_max_budget and key_token_hash: + key_info["model_max_budget_usage"] = await _build_model_max_budget_usage( + api_key_hash=key_token_hash, + model_max_budget=model_max_budget, + user_api_key_cache=user_api_key_cache, + ) # Attach object_permission if object_permission_id is set key_info = await attach_object_permission_to_dict(key_info, prisma_client) @@ -3419,7 +3506,7 @@ def _check_model_access_group( return True -async def generate_key_helper_fn( # noqa: PLR0915 +async def generate_key_helper_fn( request_type: Literal[ "user", "key" ], # identifies if this request is from /user/new or /key/generate @@ -4070,7 +4157,7 @@ async def delete_key_aliases( ) -async def _rotate_master_key( # noqa: PLR0915 +async def _rotate_master_key( prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, current_master_key: str, diff --git a/litellm/proxy/management_endpoints/mcp_management_endpoints.py b/litellm/proxy/management_endpoints/mcp_management_endpoints.py index 0df4675b67f..e86982307e7 100644 --- a/litellm/proxy/management_endpoints/mcp_management_endpoints.py +++ b/litellm/proxy/management_endpoints/mcp_management_endpoints.py @@ -49,6 +49,8 @@ from litellm.constants import LITELLM_PROXY_ADMIN_NAME from litellm.proxy._experimental.mcp_server.utils import ( build_env_var_setup_url, collect_env_var_references, + LITELLM_MCP_SERVER_DESCRIPTION, + LITELLM_MCP_SERVER_NAME, get_server_prefix, parse_admin_env_vars, ) @@ -89,8 +91,6 @@ def does_mcp_server_exist( DEFAULT_MCP_REGISTRY_VERSION = "1.0.0" -LITELLM_MCP_SERVER_NAME = "litellm-mcp-server" -LITELLM_MCP_SERVER_DESCRIPTION = "MCP Server for LiteLLM" try: importlib.import_module("mcp") diff --git a/litellm/proxy/management_endpoints/team_endpoints.py b/litellm/proxy/management_endpoints/team_endpoints.py index 85b640d6b6c..3d90e7b5ab9 100644 --- a/litellm/proxy/management_endpoints/team_endpoints.py +++ b/litellm/proxy/management_endpoints/team_endpoints.py @@ -933,7 +933,7 @@ def _check_team_budget_update_authority( response_model=LiteLLM_TeamTable, ) @management_endpoint_wrapper -async def new_team( # noqa: PLR0915 +async def new_team( data: NewTeamRequest, http_request: Request, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), @@ -1637,7 +1637,7 @@ def validate_team_org_change( "/team/update", tags=["team management"], dependencies=[Depends(user_api_key_auth)] ) @management_endpoint_wrapper -async def update_team( # noqa: PLR0915 +async def update_team( data: UpdateTeamRequest, http_request: Request, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), @@ -2961,6 +2961,7 @@ async def team_member_update( returned_team_info: TeamInfoResponseObject = await team_info( http_request=http_request, team_id=data.team_id, + key_limit=None, user_api_key_dict=user_api_key_dict, ) @@ -3577,6 +3578,9 @@ async def team_info( team_id: str = fastapi.Query( default=None, description="Team ID in the request parameters" ), + key_limit: int | None = fastapi.Query( + default=None, description="Limit the number of keys returned", gt=0 + ), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): """ @@ -3632,6 +3636,7 @@ async def team_info( table_name="key", query_type="find_all", expires=datetime.now(), + limit=key_limit, ) if keys is None: diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 5af1dd321ed..427c87e0f44 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -145,6 +145,9 @@ _CLI_SSO_START_RATE_LIMIT_WINDOW_SECONDS = 60 _CLI_SSO_START_RATE_LIMIT_MAX_ATTEMPTS = 30 _CLI_SSO_USER_CODE_ALPHABET = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789" _CLI_SSO_LOGIN_ID_RE = re.compile(r"^cli-[A-Za-z0-9_-]{12,124}$") +_CLI_SSO_USER_CODE_RE = re.compile( + rf"^[{_CLI_SSO_USER_CODE_ALPHABET}]{{4}}-[{_CLI_SSO_USER_CODE_ALPHABET}]{{4}}$" +) _CLI_SSO_SCALAR_TYPES = (str, int, float, bool) _CLI_SSO_DEST_KEY_RE = re.compile(r"^[A-Za-z0-9_.-]+$") _CLI_SSO_SECRET_KEY_FRAGMENTS = frozenset( @@ -182,6 +185,41 @@ def _is_valid_cli_sso_login_id(login_id: Optional[str]) -> bool: return isinstance(login_id, str) and bool(_CLI_SSO_LOGIN_ID_RE.fullmatch(login_id)) +def _is_valid_cli_sso_user_code(user_code: str | None) -> bool: + return isinstance(user_code, str) and bool( + _CLI_SSO_USER_CODE_RE.fullmatch(user_code) + ) + + +def _cli_sso_verification_uri_complete_enabled() -> bool: + from litellm.proxy.proxy_server import general_settings + + return bool(general_settings.get("allow_cli_sso_verification_uri_complete", False)) + + +def _cli_sso_start_response_body( + *, + login_id: str, + poll_secret: str, + user_code: str, + verification_uri_complete: str | None, +) -> dict[str, str | int]: + if verification_uri_complete is None: + return { + "login_id": login_id, + "poll_secret": poll_secret, + "user_code": user_code, + "expires_in": CLI_SSO_SESSION_TTL_SECONDS, + } + return { + "login_id": login_id, + "poll_secret": poll_secret, + "user_code": user_code, + "verification_uri_complete": verification_uri_complete, + "expires_in": CLI_SSO_SESSION_TTL_SECONDS, + } + + def _get_cli_sso_start_rate_limit_cache_key( request: Request, use_x_forwarded_for: Optional[bool] = False ) -> str: @@ -478,10 +516,20 @@ def _cli_poll_attribution_metadata_from_session( def _render_cli_sso_verification_page( - verify_url: str, browser_complete_token: str + verify_url: str, + browser_complete_token: str, + prefill_user_code: str | None = None, ) -> str: escaped_verify_url = escape(verify_url, quote=True) escaped_browser_complete_token = escape(browser_complete_token, quote=True) + user_code_value_attr = ( + f' value="{escape(prefill_user_code, quote=True)}"' if prefill_user_code else "" + ) + instructions = ( + "Confirm the verification code below to finish this login." + if prefill_user_code + else "Enter the verification code shown in your terminal to finish this login." + ) return f""" @@ -535,11 +583,11 @@ def _render_cli_sso_verification_page(

Complete CLI Login

-

Enter the verification code shown in your terminal to finish this login.

+

{instructions}

- +
@@ -573,12 +621,29 @@ async def cli_sso_start(request: Request): } _set_cli_sso_flow(login_id=login_id, cache=user_api_key_cache, flow=flow) - return { - "login_id": login_id, - "poll_secret": poll_secret, - "user_code": user_code, - "expires_in": CLI_SSO_SESSION_TTL_SECONDS, - } + verification_uri_complete: str | None = ( + ( + get_custom_url( + request_base_url=str(request.base_url), route="sso/key/generate" + ) + + "?" + + urlencode( + { + "source": LITELLM_CLI_SOURCE_IDENTIFIER, + "key": login_id, + "user_code": user_code, + } + ) + ) + if _cli_sso_verification_uri_complete_enabled() + else None + ) + return _cli_sso_start_response_body( + login_id=login_id, + poll_secret=poll_secret, + user_code=user_code, + verification_uri_complete=verification_uri_complete, + ) @router.post( @@ -829,6 +894,7 @@ async def google_login( key: Optional[str] = None, existing_key: Optional[str] = None, return_to: Optional[str] = None, + user_code: str | None = None, ): """ Create Proxy API Keys using Google Workspace SSO. Requires setting PROXY_BASE_URL in .env @@ -897,6 +963,7 @@ async def google_login( cli_state: Optional[str] = SSOAuthenticationHandler._get_cli_state( source=source, key=key, + user_code=(user_code if _cli_sso_verification_uri_complete_enabled() else None), ) # check if user defined a custom auth sso sign in handler, if yes, use it @@ -1921,14 +1988,16 @@ async def auth_callback(request: Request, state: Optional[str] = None): ) if state and state.startswith(f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:"): - # State format: {PREFIX}:{login_id} - state_parts = state.split(":", 1) + # State format: {PREFIX}:{login_id}[:{user_code}] + state_parts = state.split(":", 2) key_id = state_parts[1] if len(state_parts) > 1 else None + prefill_user_code = state_parts[2] if len(state_parts) > 2 else None verbose_proxy_logger.info("CLI SSO callback detected") return await cli_sso_callback( request=request, key=key_id, + prefill_user_code=prefill_user_code, result=result, received_response=received_response, ) @@ -2008,6 +2077,7 @@ async def _complete_cli_sso_callback_session( prisma_client: PrismaClient, user_api_key_cache: UserApiKeyCache, proxy_logging_obj: ProxyLogging, + prefill_user_code: str | None = None, ): from fastapi.responses import HTMLResponse @@ -2071,6 +2141,7 @@ async def _complete_cli_sso_callback_session( content=_render_cli_sso_verification_page( verify_url=verify_url, browser_complete_token=browser_complete_token, + prefill_user_code=prefill_user_code, ), status_code=200, ) @@ -2081,6 +2152,7 @@ async def cli_sso_callback( key: Optional[str] = None, result: Optional[Union[OpenID, dict]] = None, received_response: Optional[dict] = None, + prefill_user_code: str | None = None, ): """CLI SSO callback - stores session info for JWT generation on polling""" verbose_proxy_logger.info("CLI SSO callback") @@ -2137,6 +2209,7 @@ async def cli_sso_callback( prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, proxy_logging_obj=proxy_logging_obj, + prefill_user_code=prefill_user_code, ) except ProxyException: raise @@ -3053,21 +3126,27 @@ class SSOAuthenticationHandler: @staticmethod def _get_cli_state( - source: Optional[str], key: Optional[str], existing_key: Optional[str] = None + source: str | None, + key: str | None, + existing_key: str | None = None, + user_code: str | None = None, ) -> Optional[str]: """ Checks the request 'source' if a cli state token was passed in This is used to authenticate through the CLI login flow. - The state parameter format is: {PREFIX}:{login_id} + The state parameter format is: {PREFIX}:{login_id}[:{user_code}] - The state parameter is used to pass data through the OAuth flow without changing the callback URL + - user_code is appended only for the opt-in verification_uri_complete flow so the verify page can pre-fill it """ from litellm.constants import ( LITELLM_CLI_SESSION_TOKEN_PREFIX, ) if source == LITELLM_CLI_SOURCE_IDENTIFIER and key: + if _is_valid_cli_sso_user_code(user_code): + return f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:{key}:{user_code}" return f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:{key}" else: return None @@ -3145,7 +3224,7 @@ class SSOAuthenticationHandler: ) @staticmethod - async def get_redirect_response_from_openid( # noqa: PLR0915 + async def get_redirect_response_from_openid( result: Union[OpenID, dict, CustomOpenID], request: Request, received_response: Optional[dict] = None, diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index 2ba1d937c04..bb3033e2a6c 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -14,6 +14,8 @@ from litellm.types.utils import SpecialEnums if TYPE_CHECKING: from fastapi import Request + from litellm.router import Router + def _is_base64_encoded_unified_file_id(b64_uid: str) -> Union[str, Literal[False]]: # Ensure b64_uid is a string and not a mock object @@ -300,6 +302,92 @@ def get_credentials_for_model( return credentials +def get_team_provider_credentials( + llm_router: Optional["Router"], + team_models: List[str], + custom_llm_provider: str, + team_id: Optional[str] = None, +) -> Optional[dict]: + """ + Resolve upstream credentials for a provider-scoped file operation + (e.g. GET /v1/files), which doesn't pin a model. + + Priority: + 1. The team's own (BYOK) deployment for this provider — a deployment whose + ``model_info.team_id`` matches ``team_id``. This keeps team-scoped listings + on the team's own provider account/key instead of a shared global one. + 2. Fallback: any deployment the team is granted access to for this provider, + expanding wildcard routes and the all-proxy-models sentinel. + + Credential lookup is always scoped to the team's allowlist, so a team can + never resolve a provider key for a deployment it isn't authorized to use. + Returns None when the router is unavailable or no authorized deployment + matches, so the caller can fall back to default credential resolution. + """ + if llm_router is None: + return None + + def _provider_credentials(model_id: str) -> Optional[dict]: + credentials = llm_router.get_deployment_credentials_with_provider( + model_id=model_id + ) + if ( + credentials is not None + and credentials.get("custom_llm_provider") == custom_llm_provider + ): + return credentials + return None + + # 1. Prefer the team's own BYOK deployment, matched by model_info.team_id. + if team_id is not None: + for deployment in llm_router.model_list or []: + model_info = deployment.get("model_info") or {} + if model_info.get("team_id") != team_id: + continue + deployment_id = model_info.get("id") + if deployment_id is None: + continue + credentials = _provider_credentials(deployment_id) + if credentials is not None: + return credentials + + # 2. Fall back to deployments the team is allowed to access. The + # all-proxy-models sentinel isn't expanded by get_complete_model_list, so + # normalize it to an empty allowlist, which defers to the team-scoped + # proxy model list. A team with a restricted allowlist (e.g. anthropic + # only) therefore never resolves another provider's key. + from litellm.proxy._types import SpecialModelNames + from litellm.proxy.auth.model_checks import get_complete_model_list + + grants_all_models = SpecialModelNames.all_proxy_models.value in team_models + effective_team_models = [] if grants_all_models else team_models + + proxy_model_list = llm_router.get_model_names(team_id=team_id) + model_access_groups = llm_router.get_model_access_groups() + models_to_try = list( + dict.fromkeys( + get_complete_model_list( + key_models=[], + team_models=effective_team_models, + proxy_model_list=proxy_model_list, + user_model=None, + infer_model_from_keys=False, + return_wildcard_routes=True, + llm_router=llm_router, + model_access_groups=model_access_groups, + include_model_access_groups=True, + team_id=team_id, + ) + ) + ) + for model_name in models_to_try: + credentials = _provider_credentials(model_name) + if credentials is not None: + return credentials + + return None + + def prepare_data_with_credentials( data: dict, credentials: dict, diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index 3e5873c2655..d7dab350154 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -43,6 +43,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( encode_file_id_with_model, extract_file_creation_params, get_credentials_for_model, + get_team_provider_credentials, handle_model_based_routing, prepare_data_with_credentials, validate_managed_files_requirement, @@ -284,7 +285,7 @@ async def route_create_file( dependencies=[Depends(user_api_key_auth)], tags=["files"], ) -async def create_file( # noqa: PLR0915 +async def create_file( request: Request, fastapi_response: Response, purpose: str = Form(...), @@ -589,7 +590,7 @@ async def create_file( # noqa: PLR0915 dependencies=[Depends(user_api_key_auth)], tags=["files"], ) -async def get_file_content( # noqa: PLR0915 +async def get_file_content( request: Request, fastapi_response: Response, file_id: str, @@ -1351,14 +1352,20 @@ async def list_files( status_code=400, detail="target_model_names on list files must be a list of one model name. Example: ['gpt-4o']", ) - ## Use router to list fine-tuning jobs for that model if llm_router is None: raise HTTPException( status_code=500, detail="LLM Router not initialized. Ensure models added to proxy.", ) - data["model"] = target_model_names_list[0] - response = await llm_router.afile_list( + credentials = get_credentials_for_model( + llm_router=llm_router, + model_id=target_model_names_list[0], + operation_context="file list", + ) + prepare_data_with_credentials(data=data, credentials=credentials) + response = await litellm.afile_list( + custom_llm_provider=credentials["custom_llm_provider"], + purpose=purpose, **data, ) else: @@ -1370,6 +1377,18 @@ async def list_files( or "openai" ) + # No model/target_model_names pinned: resolve upstream credentials from + # the team's deployment for this provider so the call is authenticated + # against the team's own account (e.g. the team's openai deployment). + team_credentials = get_team_provider_credentials( + llm_router=llm_router, + team_models=user_api_key_dict.team_models or [], + custom_llm_provider=custom_llm_provider, + team_id=user_api_key_dict.team_id, + ) + if team_credentials is not None: + prepare_data_with_credentials(data=data, credentials=team_credentials) + response = await litellm.afile_list( custom_llm_provider=custom_llm_provider, purpose=purpose, **data # type: ignore ) diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py index a912a88a993..6feb4e36bf9 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/anthropic_passthrough_logging_handler.py @@ -366,7 +366,7 @@ class AnthropicPassthroughLoggingHandler: ) @staticmethod - def _collapse_pure_text_chunks( # noqa: PLR0915 + def _collapse_pure_text_chunks( all_chunks: Sequence[Union[str, bytes]], ) -> Optional[List[str]]: """ @@ -551,7 +551,7 @@ class AnthropicPassthroughLoggingHandler: return complete_streaming_response @staticmethod - def batch_creation_handler( # noqa: PLR0915 + def batch_creation_handler( httpx_response: httpx.Response, logging_obj: LiteLLMLoggingObj, url_route: str, diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py index 9f353226dd0..b77c6e2f655 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py @@ -275,7 +275,7 @@ class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler): return litellm_model_response, response_cost @staticmethod - def openai_passthrough_handler( # noqa: PLR0915 + def openai_passthrough_handler( httpx_response: httpx.Response, response_body: dict, logging_obj: LiteLLMLoggingObj, diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py index 6a138532617..73d4245670a 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py @@ -645,7 +645,7 @@ class VertexPassthroughLoggingHandler: return kwargs @staticmethod - def batch_prediction_jobs_handler( # noqa: PLR0915 + def batch_prediction_jobs_handler( httpx_response: httpx.Response, logging_obj: LiteLLMLoggingObj, url_route: str, diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index e3cb9dec884..b84746758fb 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -143,7 +143,7 @@ async def set_env_variables_in_header(custom_headers: Optional[dict]) -> Optiona return headers -async def chat_completion_pass_through_endpoint( # noqa: PLR0915 +async def chat_completion_pass_through_endpoint( fastapi_response: Response, request: Request, adapter_id: str, @@ -701,7 +701,7 @@ from litellm.passthrough.timeout_utils import ( ) -async def pass_through_request( # noqa: PLR0915 +async def pass_through_request( request: Request, target: str, custom_headers: dict, @@ -1540,7 +1540,7 @@ async def _parse_request_data_by_content_type( return query_params_data, custom_body_data, file_data, stream -def create_pass_through_route( # noqa: PLR0915 +def create_pass_through_route( endpoint, target: str, custom_headers: Optional[Mapping[str, Any]] = None, @@ -1776,7 +1776,7 @@ def create_websocket_passthrough_route( return websocket_endpoint_func -async def websocket_passthrough_request( # noqa: PLR0915 +async def websocket_passthrough_request( websocket: WebSocket, target: str, custom_headers: dict, diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index bd9746d2a47..9c4d7b1bb5d 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -205,6 +205,37 @@ class ProxyInitializationHelpers: ) return uvicorn_args + @staticmethod + def _apply_uvicorn_max_requests_jitter( + uvicorn_args: dict, + max_requests_before_restart: Optional[int], + jitter: int, + ) -> None: + """ + Stagger uvicorn worker restarts via limit_max_requests_jitter (uvicorn>=0.41.0). + """ + import inspect + + import uvicorn + + if max_requests_before_restart is None: + print( + "\033[1;33mLiteLLM Proxy: --max_requests_before_restart_jitter " + "has no effect without --max_requests_before_restart\033[0m\n" + ) + return + if ( + "limit_max_requests_jitter" + in inspect.signature(uvicorn.Config.__init__).parameters + ): + uvicorn_args["limit_max_requests_jitter"] = jitter + else: + print( + f"\033[1;33mLiteLLM Proxy: --max_requests_before_restart_jitter " + f"requires uvicorn>=0.41.0, but installed uvicorn=={uvicorn.__version__}. " + f"Ignoring the flag.\033[0m" + ) + @staticmethod def _get_reload_options(config_path: Optional[str]) -> dict: """Build uvicorn reload kwargs so --reload also reacts to .env and YAML edits.""" @@ -387,6 +418,7 @@ class ProxyInitializationHelpers: ssl_certfile_path: str, ssl_keyfile_path: str, max_requests_before_restart: Optional[int] = None, + max_requests_before_restart_jitter: Optional[int] = None, ): """ Run litellm with `gunicorn` @@ -467,6 +499,16 @@ class ProxyInitializationHelpers: # Optional: recycle workers after N requests to mitigate memory growth if max_requests_before_restart is not None: gunicorn_options["max_requests"] = max_requests_before_restart + if max_requests_before_restart_jitter is not None: + if max_requests_before_restart is None: + print( + "\033[1;33mLiteLLM Proxy: --max_requests_before_restart_jitter " + "has no effect without --max_requests_before_restart\033[0m\n" + ) + else: + gunicorn_options["max_requests_jitter"] = ( + max_requests_before_restart_jitter + ) # Clean up prometheus .db files when a worker exits (prevents ghost gauge values) if os.environ.get("PROMETHEUS_MULTIPROC_DIR"): @@ -791,6 +833,18 @@ class ProxyInitializationHelpers: help="Restart worker after this many requests (uvicorn: limit_max_requests, gunicorn: max_requests)", envvar="MAX_REQUESTS_BEFORE_RESTART", ) +@click.option( + "--max_requests_before_restart_jitter", + default=None, + type=int, + help=( + "Stagger worker restarts by adding a random amount in [0, jitter] to " + "--max_requests_before_restart so workers do not recycle at the same time " + "(uvicorn: limit_max_requests_jitter, requires uvicorn>=0.41.0; gunicorn: max_requests_jitter). " + "Has no effect without --max_requests_before_restart." + ), + envvar="MAX_REQUESTS_BEFORE_RESTART_JITTER", +) @click.option( "--enforce_prisma_migration_check", is_flag=True, @@ -814,7 +868,7 @@ class ProxyInitializationHelpers: default=False, help="Enable uvicorn hot reload (dev only). Also reloads when the --config YAML file changes. Incompatible with --num_workers>1, --run_gunicorn, and --run_hypercorn.", ) -def run_server( # noqa: PLR0915 +def run_server( cli_args, host, port, @@ -858,6 +912,7 @@ def run_server( # noqa: PLR0915 keepalive_timeout, timeout_worker_healthcheck, max_requests_before_restart, + max_requests_before_restart_jitter: Optional[int], enforce_prisma_migration_check: bool, use_v2_migration_resolver: bool, reload: bool, @@ -1260,6 +1315,12 @@ def run_server( # noqa: PLR0915 if max_requests_before_restart is not None: uvicorn_args["limit_max_requests"] = max_requests_before_restart if run_gunicorn is False and run_hypercorn is False and run_granian is False: + if max_requests_before_restart_jitter is not None: + ProxyInitializationHelpers._apply_uvicorn_max_requests_jitter( + uvicorn_args=uvicorn_args, + max_requests_before_restart=max_requests_before_restart, + jitter=max_requests_before_restart_jitter, + ) if ssl_certfile_path is not None and ssl_keyfile_path is not None: print( f"\033[1;32mLiteLLM Proxy: Using SSL with certfile: {ssl_certfile_path} and keyfile: {ssl_keyfile_path}\033[0m\n" @@ -1287,6 +1348,7 @@ def run_server( # noqa: PLR0915 ssl_certfile_path=ssl_certfile_path, ssl_keyfile_path=ssl_keyfile_path, max_requests_before_restart=max_requests_before_restart, + max_requests_before_restart_jitter=max_requests_before_restart_jitter, ) elif run_hypercorn is True: ProxyInitializationHelpers._init_hypercorn_server( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 1f765aa8d63..e6ce92344ff 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -15,6 +15,7 @@ import threading import time import traceback import warnings +from collections.abc import Mapping from datetime import datetime, timedelta, timezone from typing import ( TYPE_CHECKING, @@ -259,6 +260,7 @@ from litellm.proxy.auth.auth_checks import ( from litellm.proxy.auth.auth_utils import ( check_response_size_is_safe, is_request_body_safe, + warn_once_if_custom_auth_skips_common_checks, ) from litellm.proxy.auth.handle_jwt import JWTHandler from litellm.proxy.auth.litellm_license import LicenseCheck @@ -301,6 +303,7 @@ from litellm.proxy.common_utils.load_config_utils import ( get_config_file_contents_from_gcs, get_file_contents_from_s3, ) +from litellm.proxy.common_utils.model_listing_utils import TeamModelNameTranslator from litellm.proxy.common_utils.openai_endpoint_utils import ( remove_sensitive_info_from_deployment, ) @@ -745,7 +748,7 @@ async def _initialize_shared_aiohttp_session(): @asynccontextmanager -async def proxy_startup_event(app: FastAPI): # noqa: PLR0915 +async def proxy_startup_event(app: FastAPI): global prisma_client, master_key, use_background_health_checks, llm_router, llm_model_list, general_settings, proxy_budget_rescheduler_min_time, proxy_budget_rescheduler_max_time, litellm_proxy_admin_name, db_writer_client, store_model_in_db, premium_user, _license_check, proxy_batch_polling_interval, shared_aiohttp_session import json @@ -888,16 +891,26 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915 asyncio.create_task(_run_pw_migration()) + ## use_redis_transaction_buffer: fall back to a standalone Redis (REDIS_* env) + ## when the proxy cache backend is not Redis ## + transaction_buffer_redis_cache = redis_usage_cache + if transaction_buffer_redis_cache is None: + transaction_buffer_redis_cache = ( + ProxyStartupEvent._get_transaction_buffer_redis_cache( + general_settings=general_settings + ) + ) + ProxyStartupEvent._initialize_startup_logging( llm_router=llm_router, proxy_logging_obj=proxy_logging_obj, - redis_usage_cache=redis_usage_cache, + redis_usage_cache=transaction_buffer_redis_cache, ) ## Validate use_redis_transaction_buffer requires Redis cache ## ProxyStartupEvent._validate_redis_transaction_buffer_config( general_settings=general_settings, - redis_usage_cache=redis_usage_cache, + redis_usage_cache=transaction_buffer_redis_cache, ) ## SEMANTIC TOOL FILTER ## @@ -2496,7 +2509,7 @@ async def _invalidate_spend_counter(counter_key: str): ) -async def update_cache( # noqa: PLR0915 +async def update_cache( token: Optional[str], user_id: Optional[str], end_user_id: Optional[str], @@ -3900,7 +3913,7 @@ class ProxyConfig: premium_user = _license_check.is_premium() return - async def load_config( # noqa: PLR0915 + async def load_config( self, router: Optional[litellm.Router], config_file_path: str ): """ @@ -4361,6 +4374,12 @@ class ProxyConfig: user_custom_auth = get_instance_fn( value=custom_auth, config_file_path=config_file_path ) + warn_once_if_custom_auth_skips_common_checks( + custom_auth_configured=custom_auth is not None, + run_common_checks=bool( + general_settings.get("custom_auth_run_common_checks", False) + ), + ) custom_key_generate = general_settings.get("custom_key_generate", None) if custom_key_generate is not None: @@ -6631,7 +6650,7 @@ def save_worker_config(**data): os.environ["WORKER_CONFIG"] = json.dumps(data) -async def initialize( # noqa: PLR0915 +async def initialize( model=None, alias=None, api_base=None, @@ -7022,10 +7041,33 @@ def _format_streaming_sse_chunk(chunk: Union[str, bytes]) -> Union[str, bytes]: return f"data: {chunk}\n\n" -async def async_data_generator( # noqa: PLR0915 - response, user_api_key_dict: UserAPIKeyAuth, request_data: dict +_SSE_FRAME_DELIMITERS = ("\r\n\r\n", "\n\n", "\r\r") +_MAX_RAW_SSE_BUFFER_CHARS = 8 * 1024 * 1024 + + +def _pop_complete_sse_frame(buffer: str) -> tuple[str | None, str]: + delimiter_positions = [ + (position, delimiter) + for delimiter in _SSE_FRAME_DELIMITERS + if (position := buffer.find(delimiter)) != -1 + ] + if not delimiter_positions: + return None, buffer + + position, delimiter = min(delimiter_positions, key=lambda item: item[0]) + frame_end = position + len(delimiter) + return buffer[:frame_end], buffer[frame_end:] + + +async def async_data_generator( + response, + user_api_key_dict: UserAPIKeyAuth, + request_data: dict, + request: Request | None = None, ): verbose_proxy_logger.debug("inside generator") + stream_completed = False + client_disconnected = False try: error_message: Optional[str] = None requested_model_from_client = _get_client_requested_model_for_streaming( @@ -7047,6 +7089,8 @@ async def async_data_generator( # noqa: PLR0915 # happened to ship a streaming-iterator override (the default). needs_iterator_wrap = proxy_logging_obj.needs_iterator_wrap() needs_per_chunk_hook = proxy_logging_obj.needs_per_chunk_streaming_hook() + is_raw_sse_stream = bool(request_data.get("_litellm_raw_sse_stream")) + raw_sse_buffer = "" if needs_iterator_wrap: stream_iterator = proxy_logging_obj.async_post_call_streaming_iterator_hook( @@ -7077,14 +7121,38 @@ async def async_data_generator( # noqa: PLR0915 if isinstance(chunk, BaseModel): chunk = _serialize_streaming_chunk(chunk) elif isinstance(chunk, bytes): - # Some upstream streaming iterators (e.g. AsyncGoogleGenAIGenerateContentStreamingIterator - # for /v1beta/.../streamGenerateContent) yield raw SSE bytes from Gemini. - # Decode to str so the f-string below does not emit a Python b'...' literal, - # and pass already-formatted SSE through unchanged to avoid double "data:" prefix. chunk = chunk.decode("utf-8", errors="replace") - if chunk.startswith(("data:", "event:", ":")): - yield chunk if chunk.endswith("\n\n") else chunk + "\n\n" + if is_raw_sse_stream: + raw_sse_buffer += chunk + while True: + frame, raw_sse_buffer = _pop_complete_sse_frame(raw_sse_buffer) + if frame is None: + break + yield frame + if len(raw_sse_buffer) > _MAX_RAW_SSE_BUFFER_CHARS: + raise ValueError( + "Raw SSE stream exceeded maximum buffered size without a frame delimiter" + ) continue + if chunk.startswith(("data:", "event:", ":")): + yield ( + chunk + if chunk.endswith(_SSE_FRAME_DELIMITERS) + else chunk + "\n\n" + ) + continue + elif isinstance(chunk, str) and is_raw_sse_stream: + raw_sse_buffer += chunk + while True: + frame, raw_sse_buffer = _pop_complete_sse_frame(raw_sse_buffer) + if frame is None: + break + yield frame + if len(raw_sse_buffer) > _MAX_RAW_SSE_BUFFER_CHARS: + raise ValueError( + "Raw SSE stream exceeded maximum buffered size without a frame delimiter" + ) + continue elif isinstance(chunk, str) and chunk.startswith("data: "): error_message = chunk break @@ -7094,12 +7162,20 @@ async def async_data_generator( # noqa: PLR0915 except Exception as e: yield f"data: {str(e)}\n\n" + stream_completed = True if not needs_iterator_wrap: # The iterator-wrap path fires deferred logging itself; fire it # here for the no-wrap fast path so non-callback deployments # still flush their post-stream logging. ProxyLogging._fire_deferred_stream_logging(request_data) + if raw_sse_buffer: + yield ( + raw_sse_buffer + if raw_sse_buffer.endswith(_SSE_FRAME_DELIMITERS) + else raw_sse_buffer + "\n\n" + ) + if error_message is not None: yield error_message # OpenAI-compatible streams terminate with data: [DONE]; Google GenAI (?alt=sse) does not. @@ -7113,9 +7189,11 @@ async def async_data_generator( # noqa: PLR0915 # it here. This is the outermost generator Starlette closes on # disconnect, so it fires reliably regardless of needs_iterator_wrap # (a nested iterator hook would only see GeneratorExit on GC). - proxy_logging_obj._release_max_parallel_requests_on_disconnect( - user_api_key_dict - ) + if not stream_completed: + proxy_logging_obj._release_max_parallel_requests_on_disconnect( + user_api_key_dict + ) + client_disconnected = True raise except Exception as e: verbose_proxy_logger.exception( @@ -7149,30 +7227,33 @@ async def async_data_generator( # noqa: PLR0915 code=getattr(e, "status_code", 500), ) error_returned = json.dumps({"error": proxy_exception.to_dict()}) + stream_completed = True yield f"data: {error_returned}\n\n" finally: - # Close the response stream to release the underlying HTTP connection - # back to the connection pool. This prevents pool exhaustion when - # clients disconnect mid-stream. - # Shield from cancellation so the close awaits can complete. - with anyio.CancelScope(shield=True): - if hasattr(response, "aclose"): - try: - await response.aclose() - except BaseException as e: - verbose_proxy_logger.debug( - "async_data_generator: error closing response stream: %s", - e, - ) + from litellm.proxy.common_request_processing import ( + ProxyBaseLLMRequestProcessing, + ) + + await ProxyBaseLLMRequestProcessing._finalize_streaming_generator_cleanup( + request=request, + request_data=request_data, + response=response, + stream_completed=stream_completed, + client_disconnected=client_disconnected, + ) def select_data_generator( - response, user_api_key_dict: UserAPIKeyAuth, request_data: dict + response, + user_api_key_dict: UserAPIKeyAuth, + request_data: dict, + request: Request | None = None, ): return async_data_generator( response=response, user_api_key_dict=user_api_key_dict, request_data=request_data, + request=request, ) @@ -7250,15 +7331,48 @@ class ProxyStartupEvent: if _use_redis_transaction_buffer and redis_usage_cache is None: raise ValueError( "`use_redis_transaction_buffer` is enabled in general_settings " - "but no Redis cache is configured. This will cause spend updates " + "but no Redis is configured. This will cause spend updates " "to not be tracked. Add a Redis cache in litellm_settings:\n\n" "litellm_settings:\n" " cache: true\n" " cache_params:\n" " type: redis\n" - " url: os.environ/REDIS_URL\n" + " url: os.environ/REDIS_URL\n\n" + "or set REDIS_* environment variables (e.g. REDIS_HOST, " + "REDIS_PORT, REDIS_PASSWORD, or REDIS_URL) to use a standalone " + "Redis for the transaction buffer." ) + @staticmethod + def _get_transaction_buffer_redis_cache( + general_settings: dict, + ) -> RedisCache | None: + """ + Builds a standalone Redis cache from REDIS_* environment variables so + use_redis_transaction_buffer can run when the proxy cache backend is not + Redis (e.g. disk, s3). + + Returns None when the buffer is disabled, or when no Redis host or url + is set in the environment. + """ + from litellm._redis import _redis_kwargs_from_environment + from litellm.secret_managers.main import str_to_bool + + _use_redis_transaction_buffer: bool | str | None = general_settings.get( + "use_redis_transaction_buffer", False + ) + if isinstance(_use_redis_transaction_buffer, str): + _use_redis_transaction_buffer = str_to_bool(_use_redis_transaction_buffer) + + if not _use_redis_transaction_buffer: + return None + + redis_env_kwargs = _redis_kwargs_from_environment() + if "host" not in redis_env_kwargs and "url" not in redis_env_kwargs: + return None + + return RedisCache(**redis_env_kwargs) + @classmethod async def _initialize_semantic_tool_filter( cls, @@ -7470,7 +7584,7 @@ class ProxyStartupEvent: ) @classmethod - async def initialize_scheduled_background_jobs( # noqa: PLR0915 + async def initialize_scheduled_background_jobs( cls, general_settings: dict, prisma_client: PrismaClient, @@ -8264,6 +8378,8 @@ async def model_list( """ global llm_model_list, general_settings, llm_router, prisma_client, user_api_key_cache, proxy_logging_obj + settings = cast(dict[str, object], general_settings) # any-ok: legacy settings + from litellm.proxy.management_endpoints.common_utils import ( _user_has_admin_privileges, ) @@ -8343,16 +8459,21 @@ async def model_list( if hidden_names: all_models = [m for m in all_models if m not in hidden_names] - # Build response data with all proxy models + # Surface the public team name by default; legacy internal keys via flag. + # The internal routing key drives the metadata/fallback lookup, while the + # public name is what the client sees as the model id. model_data = [] - for model in all_models: + for response_id, lookup_id in TeamModelNameTranslator.listing_entries( + all_models, llm_router, settings + ): model_info = create_model_info_response( - model_id=model, + model_id=lookup_id, provider="openai", include_metadata=include_metadata or False, fallback_type=fallback_type, llm_router=llm_router, ) + model_info["id"] = response_id model_data.append(model_info) return dict( @@ -8380,16 +8501,21 @@ async def model_list( if hidden_names: all_models = [m for m in all_models if m not in hidden_names] - # Build response data + # Surface the public team name by default; legacy internal keys via flag. + # The internal routing key drives the metadata/fallback lookup, while the + # public name is what the client sees as the model id. model_data = [] - for model in all_models: + for response_id, lookup_id in TeamModelNameTranslator.listing_entries( + all_models, llm_router, settings + ): model_info = create_model_info_response( - model_id=model, + model_id=lookup_id, provider="openai", include_metadata=include_metadata or False, fallback_type=fallback_type, llm_router=llm_router, ) + model_info["id"] = response_id model_data.append(model_info) return dict( @@ -8411,6 +8537,8 @@ async def model_list( async def model_info( model_id: str, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), + team_id: Optional[str] = None, + healthy_only: Optional[bool] = False, ): """ Retrieve information about a specific model accessible to your API key. @@ -8420,16 +8548,21 @@ async def model_info( Follows OpenAI API specification for individual model retrieval. https://platform.openai.com/docs/api-reference/models/retrieve + + Query parameters mirror `/v1/models` so the same caller context (team + scoping, health filtering, paused deployments) drives both endpoints; the + listing's public id must resolve to the same internal deployment here. """ global llm_model_list, general_settings, llm_router, prisma_client, user_api_key_cache, proxy_logging_obj + settings = cast(dict[str, object], general_settings) # any-ok: legacy settings + from litellm.proxy.utils import ( create_model_info_response, get_available_models_for_user, validate_model_access, ) - # Get available models for the user all_models = await get_available_models_for_user( user_api_key_dict=user_api_key_dict, llm_router=llm_router, @@ -8437,21 +8570,43 @@ async def model_info( user_model=user_model, prisma_client=prisma_client, proxy_logging_obj=proxy_logging_obj, - team_id=None, + team_id=team_id, include_model_access_groups=False, only_model_access_groups=False, return_wildcard_routes=False, user_api_key_cache=user_api_key_cache, ) + # Mirror /v1/models' visibility filter so first-occurrence resolution + # cannot land on a deployment the listing had hidden. + blocked_names = ( + llm_router.get_fully_blocked_model_names() if llm_router is not None else set() + ) + unhealthy_names: set[str] = set() + if healthy_only and llm_router is not None: + unhealthy_names = await llm_router.async_get_fully_unhealthy_model_names() + hidden_names = blocked_names | unhealthy_names + if hidden_names: + all_models = [m for m in all_models if m not in hidden_names] + + internal_to_public = TeamModelNameTranslator.build_internal_to_public_map( + llm_router, settings + ) + resolved_model_id = TeamModelNameTranslator.resolve_public_name( + model_id=model_id, + available_models=all_models, + llm_router=llm_router, + general_settings=settings, + ) + # Validate that the requested model is accessible - validate_model_access(model_id=model_id, available_models=all_models) + validate_model_access(model_id=resolved_model_id, available_models=all_models) # Get provider information from the router deployment if llm_router is None: raise HTTPException(status_code=500, detail="Router not initialized") - deployment = llm_router.get_deployment_by_model_group_name(model_id) + deployment = llm_router.get_deployment_by_model_group_name(resolved_model_id) if deployment is None: raise HTTPException( status_code=404, @@ -8461,9 +8616,9 @@ async def model_info( # Use the actual litellm model from the deployment to get provider info _, provider, _, _ = litellm.get_llm_provider(model=deployment.litellm_params.model) - # Return the model information in the same format as the list endpoint + response_id = internal_to_public.get(resolved_model_id, model_id) return create_model_info_response( - model_id=model_id, + model_id=response_id, provider=provider, include_metadata=False, fallback_type=None, @@ -8609,6 +8764,7 @@ async def chat_completion( response=_streaming_response, user_api_key_dict=user_api_key_dict, request_data=_data, + request=request, ) return StreamingResponse( @@ -8643,6 +8799,7 @@ async def chat_completion( response=_streaming_response, user_api_key_dict=user_api_key_dict, request_data=_data, + request=request, ) return StreamingResponse( @@ -8681,7 +8838,7 @@ async def chat_completion( dependencies=[Depends(user_api_key_auth)], tags=["completions"], ) -async def completion( # noqa: PLR0915 +async def completion( request: Request, fastapi_response: Response, model: Optional[str] = None, @@ -8791,6 +8948,7 @@ async def completion( # noqa: PLR0915 response=_streaming_response, user_api_key_dict=user_api_key_dict, request_data=_data, + request=request, ) return StreamingResponse( @@ -8837,6 +8995,7 @@ async def completion( # noqa: PLR0915 response=_streaming_response, user_api_key_dict=user_api_key_dict, request_data=data, + request=request, ) return StreamingResponse( @@ -13309,6 +13468,7 @@ async def async_queue_request( user_api_key_dict=user_api_key_dict, response=response, request_data=data, + request=request, ), media_type="text/event-stream", ) @@ -14411,7 +14571,7 @@ async def invitation_delete( dependencies=[Depends(user_api_key_auth)], include_in_schema=False, ) -async def update_config( # noqa: PLR0915 +async def update_config( config_info: ConfigYAML, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), ): @@ -15028,7 +15188,7 @@ async def delete_callback( include_in_schema=False, dependencies=[Depends(user_api_key_auth)], ) -async def get_config(): # noqa: PLR0915 +async def get_config(): """ For Admin UI - allows admin to view config via UI # return the callbacks and the env variables for the callback diff --git a/litellm/proxy/public_endpoints/provider_create_fields.json b/litellm/proxy/public_endpoints/provider_create_fields.json index d1137dec0c0..e55e65161b0 100644 --- a/litellm/proxy/public_endpoints/provider_create_fields.json +++ b/litellm/proxy/public_endpoints/provider_create_fields.json @@ -1305,6 +1305,16 @@ "provider_display_name": "Google AI Studio", "litellm_provider": "gemini", "credential_fields": [ + { + "key": "api_base", + "label": "API Base", + "placeholder": "https://generativelanguage.googleapis.com/v1beta", + "tooltip": "Leave blank to let LiteLLM pick the right Gemini API version automatically (v1alpha for Gemini 3+ models, v1beta otherwise). Override only when fronting Gemini through a custom gateway; if you do, include the version prefix (e.g. /v1beta) but not the trailing slash. LiteLLM appends '/models/{model}:generateContent'.", + "required": false, + "field_type": "text", + "options": null, + "default_value": null + }, { "key": "api_key", "label": "API Key", diff --git a/litellm/proxy/realtime_endpoints/endpoints.py b/litellm/proxy/realtime_endpoints/endpoints.py index a953dbec6b7..409ce6f50f2 100644 --- a/litellm/proxy/realtime_endpoints/endpoints.py +++ b/litellm/proxy/realtime_endpoints/endpoints.py @@ -130,6 +130,12 @@ async def _prepare_client_secret_session( session_model = req.session.model if req.session else None model: str = session_model or req.model or _DEFAULT_REALTIME_MODEL if session_type != "transcription": + await can_key_call_resolved_model( + model=model, + valid_token=user_api_key_dict, + llm_model_list=llm_model_list, + llm_router=llm_router, + ) return model, session_data, session_type transcription_model_candidates = _transcription_model_candidates_from_session( diff --git a/litellm/proxy/response_polling/background_streaming.py b/litellm/proxy/response_polling/background_streaming.py index 03039d4f441..a69e6734d71 100644 --- a/litellm/proxy/response_polling/background_streaming.py +++ b/litellm/proxy/response_polling/background_streaming.py @@ -21,7 +21,7 @@ from litellm.proxy.response_polling.polling_handler import ResponsePollingHandle from litellm.types.llms.openai import ResponsesAPIStatus -async def background_streaming_task( # noqa: PLR0915 +async def background_streaming_task( polling_id: str, data: dict, polling_handler: ResponsePollingHandler, diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 3626a21516d..bbd8b75fdd6 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -264,7 +264,7 @@ async def add_shared_session_to_data(data: dict) -> None: pass -async def route_request( # noqa: PLR0915 - Complex routing function, refactoring tracked separately +async def route_request( data: dict, llm_router: Optional[LitellmRouter], user_model: Optional[str], diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index ef06adb27fc..0ba77dcd2f0 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -1734,7 +1734,7 @@ async def calculate_spend(request: SpendCalculateRequest): 200: {"model": List[LiteLLM_SpendLogs]}, }, ) -async def ui_view_spend_logs( # noqa: PLR0915 +async def ui_view_spend_logs( request: Request, api_key: Optional[str] = fastapi.Query( default=None, @@ -2273,7 +2273,7 @@ async def ui_view_request_response_for_request_id( 200: {"model": List[LiteLLM_SpendLogs]}, }, ) -async def view_spend_logs( # noqa: PLR0915 +async def view_spend_logs( api_key: Optional[str] = fastapi.Query( default=None, description="Get spend logs based on api key", diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index d215294fd04..aef06a3c668 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -228,9 +228,7 @@ def _extract_usage_for_ocr_call(response_obj: Any, response_obj_dict: dict) -> d return {} -def get_logging_payload( # noqa: PLR0915 - kwargs, response_obj, start_time, end_time -) -> SpendLogsPayload: +def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogsPayload: if kwargs is None: kwargs = {} diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 74f12a1eeb1..a7bc94f7430 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -46,6 +46,7 @@ from litellm.proxy._types import ( ) from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.proxy.model_listing import ModelInfoResponse from litellm.types.utils import CallTypes, CallTypesLiteral try: @@ -2088,7 +2089,7 @@ class ProxyLogging: litellm_call_id=request_data.get("litellm_call_id", ""), status="fail" ) if AlertType.llm_exceptions in self.alert_types and not isinstance( - original_exception, HTTPException + original_exception, (HTTPException, ProxyException) ): """ Just alert on LLM API exceptions. Do not alert on user errors @@ -2192,6 +2193,7 @@ class ProxyLogging: e.g should only return True for: - Authentication Errors from user_api_key_auth - HTTP HTTPException (rate limit errors) + - ProxyException (guardrail blocks, budget / rate-limit errors) """ ######################################################### @@ -2208,7 +2210,7 @@ class ProxyLogging: ): return False - return isinstance(original_exception, HTTPException) or ( + return isinstance(original_exception, (HTTPException, ProxyException)) or ( error_type == ProxyErrorTypes.auth_error ) @@ -3430,6 +3432,7 @@ class PrismaClient: r.expires = r.expires.isoformat() elif query_type == "find_all" and team_id is not None: response = await VerificationTokenRepository(self).table.find_many( + take=limit, where={"team_id": team_id}, include={"litellm_budget_table": True}, ) @@ -6310,56 +6313,61 @@ def create_model_info_response( include_metadata: bool = False, fallback_type: Optional[str] = None, llm_router: Optional["Router"] = None, -) -> dict: +) -> ModelInfoResponse: """ - Create a standardized model info response. + Create a standardized OpenAI-compatible model object. - Args: - model_id: The model ID - provider: The model provider - include_metadata: Whether to include metadata - fallback_type: Type of fallbacks to include - llm_router: LiteLLM router instance - - Returns: - Dictionary containing model information + When include_metadata is true, attaches the model's configured fallbacks + (resolved via the router under fallback_type, defaulting to "general"). + Raises HTTPException(400) for an unknown fallback_type. """ from litellm.proxy.auth.model_checks import get_all_fallbacks - model_info = { + base: ModelInfoResponse = { "id": model_id, "object": "model", "created": DEFAULT_MODEL_CREATED_AT_TIME, "owned_by": provider, } - # Add metadata if requested - if include_metadata: - metadata = {} - - # Default fallback_type to "general" if include_metadata is true - effective_fallback_type = ( - fallback_type if fallback_type is not None else "general" - ) - - # Validate fallback_type - valid_fallback_types = ["general", "context_window", "content_policy"] - if effective_fallback_type not in valid_fallback_types: - raise HTTPException( - status_code=400, - detail=f"Invalid fallback_type. Must be one of: {valid_fallback_types}", + # Surface context-window limits for OpenAI-compatible discovery clients. + # Only emitted when known, so wildcard routes and limitless backends stay clean. + # Limits are best-effort enrichment, so a single malformed deployment degrades + # to the base response rather than 500-ing the whole listing. + if llm_router is not None: + try: + model_group_info = llm_router.get_model_group_info(model_id) + except Exception as e: + verbose_proxy_logger.debug( + "create_model_info_response: get_model_group_info failed for %s: %s", + model_id, + e, ) + model_group_info = None + if model_group_info is not None: + if model_group_info.max_input_tokens is not None: + base["max_input_tokens"] = int(model_group_info.max_input_tokens) + if model_group_info.max_output_tokens is not None: + base["max_output_tokens"] = int(model_group_info.max_output_tokens) - fallbacks = get_all_fallbacks( - model=model_id, - llm_router=llm_router, - fallback_type=effective_fallback_type, + if not include_metadata: + return base + + effective_fallback_type = fallback_type if fallback_type is not None else "general" + + valid_fallback_types = ["general", "context_window", "content_policy"] + if effective_fallback_type not in valid_fallback_types: + raise HTTPException( + status_code=400, + detail=f"Invalid fallback_type. Must be one of: {valid_fallback_types}", ) - metadata["fallbacks"] = fallbacks - model_info["metadata"] = metadata - - return model_info + fallbacks = get_all_fallbacks( + model=model_id, + llm_router=llm_router, + fallback_type=effective_fallback_type, + ) + return {**base, "metadata": {"fallbacks": fallbacks}} def validate_model_access( diff --git a/litellm/proxy/wildcard_config.yaml b/litellm/proxy/wildcard_config.yaml new file mode 100644 index 00000000000..7c178690836 --- /dev/null +++ b/litellm/proxy/wildcard_config.yaml @@ -0,0 +1,52 @@ +model_list: + # ---------- Anthropic native ---------- + - model_name: "anthropic/*" + litellm_params: + model: "anthropic/*" + api_key: os.environ/ANTHROPIC_API_KEY + + # ---------- Bedrock ---------- + - model_name: "bedrock/*" + litellm_params: + model: "bedrock/*" + aws_region_name: us-east-1 + + # ---------- Vertex AI ---------- + - model_name: "vertex_ai/*" + litellm_params: + model: "vertex_ai/*" + vertex_project: os.environ/VERTEXAI_PROJECT + vertex_location: global + + # ---------- Azure AI Foundry ---------- + - model_name: "azure_ai/*" + litellm_params: + model: "azure_ai/*" + api_base: os.environ/AZURE_AI_API_BASE + api_key: os.environ/AZURE_AI_API_KEY + + # ---------- Azure OpenAI ---------- + - model_name: "azure/*" + litellm_params: + model: "azure/*" + api_base: os.environ/AZURE_API_BASE + api_key: os.environ/AZURE_API_KEY + + # ---------- Gemini ---------- + - model_name: "gemini/*" + litellm_params: + model: "gemini/*" + api_key: os.environ/GEMINI_API_KEY + + # ---------- OpenAI ---------- + - model_name: "openai/*" + litellm_params: + model: "openai/*" + api_key: os.environ/OPENAI_API_KEY + +general_settings: + master_key: sk-1234 + +litellm_settings: + drop_params: True + telemetry: False diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 7031ecaa1a0..f6f0a92def0 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -284,7 +284,7 @@ async def arealtime_calls( @wrapper_client -async def _arealtime( # noqa: PLR0915 +async def _arealtime( model: str, websocket: Any, # fastapi websocket api_base: Optional[str] = None, diff --git a/litellm/rerank_api/main.py b/litellm/rerank_api/main.py index e27585116ce..e40e12e9197 100644 --- a/litellm/rerank_api/main.py +++ b/litellm/rerank_api/main.py @@ -75,7 +75,7 @@ async def arerank( @client -def rerank( # noqa: PLR0915 +def rerank( model: str, query: str, documents: List[Union[str, Dict[str, Any]]], diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py index 24b5db28571..acb7487f430 100644 --- a/litellm/responses/mcp/chat_completions_handler.py +++ b/litellm/responses/mcp/chat_completions_handler.py @@ -77,7 +77,7 @@ def _add_mcp_metadata_to_response( setattr(message, "provider_specific_fields", provider_fields) -async def acompletion_with_mcp( # noqa: PLR0915 +async def acompletion_with_mcp( model: str, messages: List, tools: Optional[List] = None, diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 94cff6922b5..df5de205d45 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -644,7 +644,7 @@ class LiteLLM_Proxy_MCP_Handler: return result_text or "Tool executed successfully" @staticmethod - async def _execute_tool_calls( # noqa: PLR0915 + async def _execute_tool_calls( tool_server_map: dict[str, str], tool_calls: List[Any], user_api_key_auth: Any, diff --git a/litellm/router.py b/litellm/router.py index d256fbd003f..79f591425a7 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -241,7 +241,7 @@ class Router: lowesttpm_logger: Optional[LowestTPMLoggingHandler] = None optional_callbacks: Optional[List[Union[CustomLogger, Callable, str]]] = None - def __init__( # noqa: PLR0915 + def __init__( self, model_list: Optional[ Union[List[DeploymentTypedDict], List[Dict[str, Any]]] @@ -2887,7 +2887,7 @@ class Router: f"Silent experiment failed for model {silent_model}: {str(e)}" ) - async def _acompletion( # noqa: PLR0915 + async def _acompletion( self, model: str, messages: List[Dict[str, str]], **kwargs ) -> Union[ ModelResponse, @@ -3045,7 +3045,8 @@ class Router: deployment_timeout_param = _timeout_debug_deployment_dict.get( "litellm_params", {} ).get("timeout", None) - e.message += f"\n\nDeployment Info: request_timeout: {deployment_request_timeout_param}\ntimeout: {deployment_timeout_param}" + if litellm.expose_router_debug_in_errors: + e.message += f"\n\nDeployment Info: request_timeout: {deployment_request_timeout_param}\ntimeout: {deployment_timeout_param}" # Set per-deployment num_retries on exception for retry logic if deployment is not None: self._set_deployment_num_retries_on_exception(e, deployment) @@ -5158,7 +5159,7 @@ class Router: ) raise e - async def _acreate_file( # noqa: PLR0915 + async def _acreate_file( self, model: str, **kwargs, @@ -6469,7 +6470,7 @@ class Router: # propagate so they remain visible. return None - async def async_function_with_fallbacks_common_utils( # noqa: PLR0915 + async def async_function_with_fallbacks_common_utils( self, e: Exception, disable_fallbacks: Optional[bool], @@ -6646,7 +6647,8 @@ class Router: ) ) - e.message += "\n{}".format(error_message) + if litellm.expose_router_debug_in_errors: + e.message += "\n{}".format(error_message) elif isinstance(e, litellm.ContentPolicyViolationError): if content_policy_fallbacks is not None: content_policy_fallback_model_group: Optional[List[str]] = ( @@ -6681,7 +6683,8 @@ class Router: ) ) - e.message += "\n{}".format(error_message) + if litellm.expose_router_debug_in_errors: + e.message += "\n{}".format(error_message) if fallbacks is not None and model_group is not None: verbose_router_logger.debug(f"inside model fallbacks: {fallbacks}") ( @@ -6699,7 +6702,10 @@ class Router: verbose_router_logger.info( f"No fallback model group found for original model_group={model_group}. Fallbacks={fallbacks}" ) - if hasattr(original_exception, "message"): + if ( + hasattr(original_exception, "message") + and litellm.expose_router_debug_in_errors + ): original_exception.message += f"No fallback model group found for original model_group={model_group}. Fallbacks={fallbacks}" # type: ignore raise original_exception @@ -6730,7 +6736,10 @@ class Router: ) fallback_failure_exception_str = str(new_exception) - if hasattr(original_exception, "message"): + if ( + hasattr(original_exception, "message") + and litellm.expose_router_debug_in_errors + ): # add the available fallbacks to the exception original_exception.message += ". Received Model Group={}\nAvailable Model Group Fallbacks={}".format( # type: ignore model_group, @@ -6845,7 +6854,7 @@ class Router: ) @tracer.wrap() - async def async_function_with_retries(self, *args, **kwargs): # noqa: PLR0915 + async def async_function_with_retries(self, *args, **kwargs): verbose_router_logger.debug("Inside async function with retries.") original_function = kwargs.pop("original_function") fallbacks = kwargs.pop("fallbacks", self.fallbacks) @@ -9326,7 +9335,7 @@ class Router: return model_info - def _set_model_group_info( # noqa: PLR0915 + def _set_model_group_info( self, model_group: str, user_facing_model_group_name: str ) -> Optional[ModelGroupInfo]: """ @@ -10568,7 +10577,7 @@ class Router: ) return client - def _pre_call_checks( # noqa: PLR0915 + def _pre_call_checks( self, model: str, healthy_deployments: List, diff --git a/litellm/router_strategy/lowest_cost.py b/litellm/router_strategy/lowest_cost.py index 54498363f51..3f641d4f0fb 100644 --- a/litellm/router_strategy/lowest_cost.py +++ b/litellm/router_strategy/lowest_cost.py @@ -190,7 +190,7 @@ class LowestCostLoggingHandler(CustomLogger): ) pass - async def async_get_available_deployments( # noqa: PLR0915 + async def async_get_available_deployments( self, model_group: str, healthy_deployments: list, diff --git a/litellm/router_strategy/lowest_latency.py b/litellm/router_strategy/lowest_latency.py index 870b3f29d48..3adb8d43920 100644 --- a/litellm/router_strategy/lowest_latency.py +++ b/litellm/router_strategy/lowest_latency.py @@ -35,9 +35,7 @@ class LowestLatencyLoggingHandler(CustomLogger): self.router_cache = router_cache self.routing_args = RoutingArgs(**routing_args) - def log_success_event( # noqa: PLR0915 - self, kwargs, response_obj, start_time, end_time - ): + def log_success_event(self, kwargs, response_obj, start_time, end_time): try: """ Update latency usage on success @@ -259,9 +257,7 @@ class LowestLatencyLoggingHandler(CustomLogger): ) pass - async def async_log_success_event( # noqa: PLR0915 - self, kwargs, response_obj, start_time, end_time - ): + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): try: """ Update latency usage on success @@ -413,7 +409,7 @@ class LowestLatencyLoggingHandler(CustomLogger): ) pass - def _get_available_deployments( # noqa: PLR0915 + def _get_available_deployments( self, model_group: str, healthy_deployments: list, diff --git a/litellm/router_strategy/lowest_tpm_rpm.py b/litellm/router_strategy/lowest_tpm_rpm.py index 488f8450941..f807ba7232a 100644 --- a/litellm/router_strategy/lowest_tpm_rpm.py +++ b/litellm/router_strategy/lowest_tpm_rpm.py @@ -158,7 +158,7 @@ class LowestTPMLoggingHandler(CustomLogger): verbose_router_logger.debug(traceback.format_exc()) pass - def get_available_deployments( # noqa: PLR0915 + def get_available_deployments( self, model_group: str, healthy_deployments: list, diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index 62e706a0cf5..eb756e3cf8b 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -244,9 +244,13 @@ def _check_non_standard_fallback_format(fallbacks: Optional[List[Any]]) -> bool: if all(isinstance(item, str) for item in fallbacks): return True elif all(isinstance(item, dict) for item in fallbacks): - for key in LiteLLMParamsTypedDict.__annotations__.keys(): - if key in fallbacks[0].keys(): - return True + for item in fallbacks: + for key in LiteLLMParamsTypedDict.__annotations__.keys(): + if key in item: + # If the value is a list, it's likely a standard fallback model group mapping + # (e.g. {"model": ["backup"]}) rather than a parameter override. + if not isinstance(item[key], list): + return True return False diff --git a/litellm/secret_managers/aws_secret_manager_v2.py b/litellm/secret_managers/aws_secret_manager_v2.py index 4461e34396e..ef3c821caf1 100644 --- a/litellm/secret_managers/aws_secret_manager_v2.py +++ b/litellm/secret_managers/aws_secret_manager_v2.py @@ -29,6 +29,7 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.proxy._types import KeyManagementSystem from litellm.types.llms.custom_http import httpxSpecialProvider +from litellm.types.secret_managers.main import KeyManagementSettings from .base_secret_manager import BaseSecretManager @@ -43,6 +44,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): aws_profile_name: Optional[str] = None, aws_web_identity_token: Optional[str] = None, aws_sts_endpoint: Optional[str] = None, + replica_regions: list[str] | None = None, **kwargs, ): BaseSecretManager.__init__(self, **kwargs) @@ -56,6 +58,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): self.aws_profile_name = aws_profile_name self.aws_web_identity_token = aws_web_identity_token self.aws_sts_endpoint = aws_sts_endpoint + self.replica_regions: list[str] = replica_regions or [] @classmethod def validate_environment(cls): @@ -75,7 +78,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): def load_aws_secret_manager( cls, use_aws_secret_manager: Optional[bool], - key_management_settings: Optional[Any] = None, + key_management_settings: KeyManagementSettings | None = None, ): """ Initialize AWSSecretsManagerV2 with settings from key_management_settings @@ -110,6 +113,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): "aws_sts_endpoint": getattr( key_management_settings, "aws_sts_endpoint", None ), + "replica_regions": key_management_settings.replica_regions, } # Remove None values aws_kwargs = {k: v for k, v in aws_kwargs.items() if v is not None} @@ -316,6 +320,90 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): params={"timeout": timeout}, ) + try: + response = await async_client.post( + url=endpoint_url, + headers=headers, + data=body.decode("utf-8"), + ) + response.raise_for_status() + create_response = response.json() + except httpx.HTTPStatusError as err: + raise ValueError(f"HTTP error occurred: {err.response.text}") + except httpx.TimeoutException: + raise ValueError("Timeout error occurred") + + if self.replica_regions: + try: + await self.async_replicate_secret( + secret_name=secret_name, + replica_regions=self.replica_regions, + optional_params=optional_params, + timeout=timeout, + ) + verbose_logger.debug( + "Replicated secret '%s' to regions: %s", + secret_name, + self.replica_regions, + ) + except Exception as replication_err: # noqa: BLE001 + verbose_logger.warning( + "Failed to replicate secret '%s' to regions %s: %s — key was created successfully.", + secret_name, + self.replica_regions, + str(replication_err), + ) + + return create_response + + async def async_replicate_secret( + self, + secret_name: str, + replica_regions: list[str], + optional_params: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> dict[str, object]: + """ + Replicate a secret to additional AWS regions using ReplicateSecretToRegions. + + Called after a successful CreateSecret when replica_regions is configured. + Replication is best-effort — callers should not depend on this for correctness. + + Args: + secret_name: Name or ARN of the secret to replicate + replica_regions: List of target AWS region names, e.g. ["us-west-2"] + optional_params: Additional AWS parameters + timeout: Request timeout + + Returns: + dict: AWS response, or {} if replica_regions is empty + """ + if not replica_regions: + return {} + + verbose_logger.info( + "ReplicateSecretToRegions called for secret '%s' in regions %s", + secret_name, + replica_regions, + ) + + data: dict[str, object] = { + "SecretId": secret_name, + "AddReplicaRegions": [{"Region": r} for r in replica_regions], + } + + endpoint_url, headers, body = self._prepare_request( + action="ReplicateSecretToRegions", + secret_name=secret_name, + optional_params=optional_params, + request_data=data, + ) + + async_client = get_async_httpx_client( + llm_provider=httpxSpecialProvider.SecretManager, + params={"timeout": timeout}, + ) + try: response = await async_client.post( url=endpoint_url, headers=headers, data=body.decode("utf-8") diff --git a/litellm/secret_managers/main.py b/litellm/secret_managers/main.py index 5c31d81f04c..f4b1d4a1b69 100644 --- a/litellm/secret_managers/main.py +++ b/litellm/secret_managers/main.py @@ -156,7 +156,7 @@ def get_secret_bool( return str_to_bool(_secret_value) -def get_secret( # noqa: PLR0915 +def get_secret( secret_name: str, default_value: Optional[Union[str, bool]] = None, ): diff --git a/litellm/secret_managers/secret_manager_handler.py b/litellm/secret_managers/secret_manager_handler.py index 4ff94d18eff..3a3cf6272dc 100644 --- a/litellm/secret_managers/secret_manager_handler.py +++ b/litellm/secret_managers/secret_manager_handler.py @@ -23,7 +23,7 @@ def _is_base64(s): return False -def get_secret_from_manager( # noqa: PLR0915 +def get_secret_from_manager( client: Any, key_manager: str, secret_name: str, diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 55216caa941..6eb65d7be02 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -211,6 +211,9 @@ class PiiEntityType(str, Enum): # UK UK_NHS = "UK_NHS" UK_NINO = "UK_NINO" + UK_PASSPORT = "UK_PASSPORT" + UK_POSTCODE = "UK_POSTCODE" + UK_VEHICLE_REGISTRATION = "UK_VEHICLE_REGISTRATION" # Spain ES_NIF = "ES_NIF" ES_NIE = "ES_NIE" @@ -265,7 +268,13 @@ PII_ENTITY_CATEGORIES_MAP = { PiiEntityType.US_PASSPORT, PiiEntityType.US_SSN, ], - PiiEntityCategory.UK: [PiiEntityType.UK_NHS, PiiEntityType.UK_NINO], + PiiEntityCategory.UK: [ + PiiEntityType.UK_NHS, + PiiEntityType.UK_NINO, + PiiEntityType.UK_PASSPORT, + PiiEntityType.UK_POSTCODE, + PiiEntityType.UK_VEHICLE_REGISTRATION, + ], PiiEntityCategory.SPAIN: [PiiEntityType.ES_NIF, PiiEntityType.ES_NIE], PiiEntityCategory.ITALY: [ PiiEntityType.IT_FISCAL_CODE, @@ -319,8 +328,7 @@ class PresidioPresidioConfigModelUserInterface(BaseModel): presidio_filter_scope: Optional[Literal["input", "output", "both"]] = Field( default=None, description=( - "Where to apply Presidio checks: 'input' (user -> model), " - "'output' (model -> user), or 'both' (default)." + "Where to apply Presidio checks: 'input' (user -> model), 'output' (model -> user), or 'both' (default)." ), ) output_parse_pii: Optional[bool] = Field( diff --git a/litellm/types/proxy/discovery_endpoints/ui_discovery_endpoints.py b/litellm/types/proxy/discovery_endpoints/ui_discovery_endpoints.py index 46cd3f49f1a..5498474ea9f 100644 --- a/litellm/types/proxy/discovery_endpoints/ui_discovery_endpoints.py +++ b/litellm/types/proxy/discovery_endpoints/ui_discovery_endpoints.py @@ -11,5 +11,6 @@ class UiDiscoveryEndpoints(BaseModel): auto_redirect_to_sso: bool admin_ui_disabled: bool sso_configured: bool + hide_default_credentials_hint: bool = False is_control_plane: bool = False workers: List[WorkerRegistryEntry] = [] diff --git a/litellm/types/proxy/model_listing.py b/litellm/types/proxy/model_listing.py new file mode 100644 index 00000000000..c3330da0d66 --- /dev/null +++ b/litellm/types/proxy/model_listing.py @@ -0,0 +1,21 @@ +"""Response types for the model listing/retrieve endpoints (/v1/models, /models).""" + +from typing import Literal + +from typing_extensions import NotRequired, TypedDict + + +class ModelInfoMetadata(TypedDict): + fallbacks: list[str] + + +class ModelInfoResponse(TypedDict): + """OpenAI-compatible model object. `metadata` is present only when the + endpoint is called with include_metadata=true. + """ + + id: str + object: Literal["model"] + created: int + owned_by: str + metadata: NotRequired[ModelInfoMetadata] diff --git a/litellm/types/secret_managers/main.py b/litellm/types/secret_managers/main.py index b0a294188cd..00a092a3c93 100644 --- a/litellm/types/secret_managers/main.py +++ b/litellm/types/secret_managers/main.py @@ -72,3 +72,12 @@ class KeyManagementSettings(LiteLLMPydanticObjectBase): aws_sts_endpoint: Optional[str] = None """Custom STS endpoint URL (useful for VPC endpoints or testing)""" + + replica_regions: Optional[List[str]] = None + """ + Optional list of additional AWS regions to replicate secrets to after CreateSecret. + Uses the AWS Secrets Manager ReplicateSecretToRegions API. Replication is + best-effort — failure to replicate does not fail key creation. + Example: ["us-west-2", "eu-west-1"] + Only applies when key_management_system is "aws_secret_manager". + """ diff --git a/litellm/types/utils.py b/litellm/types/utils.py index a5032942011..80034e50393 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -196,7 +196,9 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False): float ] # OpenAI priority service tier pricing cache_read_input_token_cost_above_200k_tokens: Optional[float] + cache_read_input_token_cost_above_200k_tokens_priority: Optional[float] cache_read_input_token_cost_above_272k_tokens: Optional[float] + cache_read_input_token_cost_above_272k_tokens_priority: Optional[float] cache_read_input_token_cost_above_512k_tokens: Optional[float] input_cost_per_character: Optional[float] # only for vertex ai models input_cost_per_audio_token: Optional[float] @@ -204,9 +206,11 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False): input_cost_per_token_above_200k_tokens: Optional[ float ] # only for vertex ai gemini-2.5-pro models + input_cost_per_token_above_200k_tokens_priority: Optional[float] input_cost_per_token_above_272k_tokens: Optional[ float ] # GPT-5.4/5.4-pro: prompts >272K priced at 2x input + input_cost_per_token_above_272k_tokens_priority: Optional[float] input_cost_per_token_above_512k_tokens: Optional[ float ] # MiniMax-M3: prompts >512K priced at 2x input @@ -240,9 +244,11 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False): output_cost_per_token_above_200k_tokens: Optional[ float ] # only for vertex ai gemini-2.5-pro models + output_cost_per_token_above_200k_tokens_priority: Optional[float] output_cost_per_token_above_272k_tokens: Optional[ float ] # GPT-5.4/5.4-pro: prompts >272K priced at 1.5x output + output_cost_per_token_above_272k_tokens_priority: Optional[float] output_cost_per_token_above_512k_tokens: Optional[ float ] # MiniMax-M3: prompts >512K priced at 2x output @@ -1572,7 +1578,7 @@ class Usage(SafeAttributeModel, CompletionUsage): prompt_tokens_details: Optional[PromptTokensDetailsWrapper] = None """Breakdown of tokens used in the prompt.""" - def __init__( # noqa: PLR0915 + def __init__( self, prompt_tokens: Optional[int] = None, completion_tokens: Optional[int] = None, @@ -1908,7 +1914,7 @@ class ModelResponse(ModelResponseBase): choices: List[Choices] """The list of completion choices the model generated for the input prompt.""" - def __init__( # noqa: PLR0915 + def __init__( self, id=None, choices=None, @@ -3093,6 +3099,8 @@ class CustomPricingLiteLLMParams(BaseModel): cache_read_input_token_cost_flex: Optional[float] = None cache_read_input_token_cost_priority: Optional[float] = None cache_read_input_token_cost_above_200k_tokens: Optional[float] = None + cache_read_input_token_cost_above_200k_tokens_priority: Optional[float] = None + cache_read_input_token_cost_above_272k_tokens_priority: Optional[float] = None cache_read_input_audio_token_cost: Optional[float] = None input_cost_per_character: Optional[float] = None input_cost_per_character_above_128k_tokens: Optional[float] = None @@ -3100,6 +3108,8 @@ class CustomPricingLiteLLMParams(BaseModel): input_cost_per_token_cache_hit: Optional[float] = None input_cost_per_token_above_128k_tokens: Optional[float] = None input_cost_per_token_above_200k_tokens: Optional[float] = None + input_cost_per_token_above_200k_tokens_priority: Optional[float] = None + input_cost_per_token_above_272k_tokens_priority: Optional[float] = None input_cost_per_query: Optional[float] = None input_cost_per_image: Optional[float] = None input_cost_per_image_above_128k_tokens: Optional[float] = None @@ -3117,6 +3127,8 @@ class CustomPricingLiteLLMParams(BaseModel): output_cost_per_audio_token: Optional[float] = None output_cost_per_token_above_128k_tokens: Optional[float] = None output_cost_per_token_above_200k_tokens: Optional[float] = None + output_cost_per_token_above_200k_tokens_priority: Optional[float] = None + output_cost_per_token_above_272k_tokens_priority: Optional[float] = None output_cost_per_character_above_128k_tokens: Optional[float] = None output_cost_per_image: Optional[float] = None output_cost_per_image_token: Optional[float] = None @@ -3234,6 +3246,11 @@ all_litellm_params = ( "order", "enable_json_schema_validation", "use_xai_oauth", + "_litellm_rate_limit_descriptors", + "_litellm_tpm_reserved_tokens", + "_litellm_tpm_reserved_model", + "_litellm_tpm_reserved_scopes", + "_litellm_tpm_reservation_released", ] + list(StandardCallbackDynamicParams.__annotations__.keys()) + list(CustomPricingLiteLLMParams.model_fields.keys()) @@ -3657,6 +3674,7 @@ class SpecialEnums(Enum): class ServiceTier(Enum): """Enum for service tier types used in cost calculations.""" + AUTO = "auto" FLEX = "flex" PRIORITY = "priority" diff --git a/litellm/utils.py b/litellm/utils.py index dd25de40835..d9f4e99dc9b 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -760,7 +760,7 @@ def _remove_thought_signatures_from_messages( return processed_messages -def function_setup( # noqa: PLR0915 +def function_setup( original_function: str, rules_obj, start_time, *args, **kwargs ): # just run once to check if user wants to send their data anywhere - PostHog/Sentry/Slack/etc. ### NOTICES ### @@ -1422,12 +1422,12 @@ def post_call_processing( raise e -def client(original_function): # noqa: PLR0915 +def client(original_function): Rules = getattr(sys.modules[__name__], "Rules") rules_obj = Rules() @wraps(original_function) - def wrapper(*args, **kwargs): # noqa: PLR0915 + def wrapper(*args, **kwargs): # DO NOT MOVE THIS. It always needs to run first # Check if this is an async function. If so only execute the async function call_type = original_function.__name__ @@ -1775,7 +1775,7 @@ def client(original_function): # noqa: PLR0915 raise e @wraps(original_function) - async def wrapper_async(*args, **kwargs): # noqa: PLR0915 + async def wrapper_async(*args, **kwargs): print_args_passed_to_litellm(original_function, args, kwargs) start_time = datetime.datetime.now() result = None @@ -2942,7 +2942,7 @@ def _resolve_builtin_model_cost_entry( return None -def register_model(model_cost: Union[str, dict]): # noqa: PLR0915 +def register_model(model_cost: Union[str, dict]): """ Register new / Override existing models (and their pricing) to specific providers. Provide EITHER a model cost dictionary or a url to a hosted json blob @@ -3365,7 +3365,7 @@ def get_optional_params_image_gen( return optional_params -def get_optional_params_embeddings( # noqa: PLR0915 +def get_optional_params_embeddings( # 2 optional params model: str, user: Optional[str] = None, @@ -4112,7 +4112,7 @@ def pre_process_optional_params( return optional_params -def get_optional_params( # noqa: PLR0915 +def get_optional_params( # use the openai defaults # https://platform.openai.com/docs/api-reference/chat/create model: str, @@ -5842,7 +5842,7 @@ def _is_potential_model_name_in_model_cost( ) -def _get_model_info_helper( # noqa: PLR0915 +def _get_model_info_helper( model: str, custom_llm_provider: Optional[str] = None, api_base: Optional[str] = None, @@ -6043,9 +6043,15 @@ def _get_model_info_helper( # noqa: PLR0915 cache_read_input_token_cost_above_200k_tokens=_model_info.get( "cache_read_input_token_cost_above_200k_tokens", None ), + cache_read_input_token_cost_above_200k_tokens_priority=_model_info.get( + "cache_read_input_token_cost_above_200k_tokens_priority", None + ), cache_read_input_token_cost_above_272k_tokens=_model_info.get( "cache_read_input_token_cost_above_272k_tokens", None ), + cache_read_input_token_cost_above_272k_tokens_priority=_model_info.get( + "cache_read_input_token_cost_above_272k_tokens_priority", None + ), cache_read_input_token_cost_above_512k_tokens=_model_info.get( "cache_read_input_token_cost_above_512k_tokens", None ), @@ -6067,9 +6073,15 @@ def _get_model_info_helper( # noqa: PLR0915 input_cost_per_token_above_200k_tokens=_model_info.get( "input_cost_per_token_above_200k_tokens", None ), + input_cost_per_token_above_200k_tokens_priority=_model_info.get( + "input_cost_per_token_above_200k_tokens_priority", None + ), input_cost_per_token_above_272k_tokens=_model_info.get( "input_cost_per_token_above_272k_tokens", None ), + input_cost_per_token_above_272k_tokens_priority=_model_info.get( + "input_cost_per_token_above_272k_tokens_priority", None + ), input_cost_per_token_above_512k_tokens=_model_info.get( "input_cost_per_token_above_512k_tokens", None ), @@ -6125,9 +6137,15 @@ def _get_model_info_helper( # noqa: PLR0915 output_cost_per_token_above_200k_tokens=_model_info.get( "output_cost_per_token_above_200k_tokens", None ), + output_cost_per_token_above_200k_tokens_priority=_model_info.get( + "output_cost_per_token_above_200k_tokens_priority", None + ), output_cost_per_token_above_272k_tokens=_model_info.get( "output_cost_per_token_above_272k_tokens", None ), + output_cost_per_token_above_272k_tokens_priority=_model_info.get( + "output_cost_per_token_above_272k_tokens_priority", None + ), output_cost_per_token_above_512k_tokens=_model_info.get( "output_cost_per_token_above_512k_tokens", None ), @@ -6566,7 +6584,7 @@ def create_proxy_transport_and_mounts(): return sync_proxy_mounts, async_proxy_mounts -def validate_environment( # noqa: PLR0915 +def validate_environment( model: Optional[str] = None, api_key: Optional[str] = None, api_base: Optional[str] = None, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 4bc74141c5d..0d0d879e98c 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -2528,6 +2528,100 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "azure_ai/gpt-5.5": { + "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_above_272k_tokens": 1e-06, + "cache_read_input_token_cost_priority": 1e-06, + "cache_read_input_token_cost_above_272k_tokens_priority": 2e-06, + "input_cost_per_token": 5e-06, + "input_cost_per_token_above_272k_tokens": 1e-05, + "input_cost_per_token_priority": 1e-05, + "input_cost_per_token_above_272k_tokens_priority": 2e-05, + "litellm_provider": "azure_ai", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 3e-05, + "output_cost_per_token_above_272k_tokens": 4.5e-05, + "output_cost_per_token_priority": 6e-05, + "output_cost_per_token_above_272k_tokens_priority": 9e-05, + "source": "https://ai.azure.com/catalog/models/gpt-5.5", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_service_tier": true, + "supports_vision": true, + "supports_web_search": true, + "supports_none_reasoning_effort": true, + "supports_xhigh_reasoning_effort": true, + "supports_minimal_reasoning_effort": false + }, + "azure_ai/gpt-5.5-2026-04-23": { + "cache_read_input_token_cost": 5e-07, + "cache_read_input_token_cost_above_272k_tokens": 1e-06, + "cache_read_input_token_cost_priority": 1e-06, + "cache_read_input_token_cost_above_272k_tokens_priority": 2e-06, + "input_cost_per_token": 5e-06, + "input_cost_per_token_above_272k_tokens": 1e-05, + "input_cost_per_token_priority": 1e-05, + "input_cost_per_token_above_272k_tokens_priority": 2e-05, + "litellm_provider": "azure_ai", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 3e-05, + "output_cost_per_token_above_272k_tokens": 4.5e-05, + "output_cost_per_token_priority": 6e-05, + "output_cost_per_token_above_272k_tokens_priority": 9e-05, + "source": "https://ai.azure.com/catalog/models/gpt-5.5", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/batch", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_service_tier": true, + "supports_vision": true, + "supports_web_search": true, + "supports_none_reasoning_effort": true, + "supports_xhigh_reasoning_effort": true, + "supports_minimal_reasoning_effort": false + }, "azure_ai/gpt-5.4": { "cache_read_input_token_cost": 2.5e-07, "cache_read_input_token_cost_above_272k_tokens": 5e-07, @@ -10068,6 +10162,8 @@ }, "claude-sonnet-4-5": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, + "cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, "input_cost_per_token_above_200k_tokens": 6e-06, @@ -10097,6 +10193,8 @@ }, "claude-sonnet-4-5-20250929": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, + "cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, "input_cost_per_token_above_200k_tokens": 6e-06, @@ -10127,6 +10225,7 @@ }, "claude-sonnet-4-6": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, "litellm_provider": "anthropic", @@ -10155,6 +10254,8 @@ }, "claude-sonnet-4-5-20250929-v1:0": { "cache_creation_input_token_cost": 3.75e-06, + "cache_creation_input_token_cost_above_1hr": 6e-06, + "cache_creation_input_token_cost_above_1hr_above_200k_tokens": 1.2e-05, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, "input_cost_per_token_above_200k_tokens": 6e-06, @@ -10811,13 +10912,13 @@ "supports_tool_choice": true }, "command-r7b-12-2024": { - "input_cost_per_token": 1.5e-07, + "input_cost_per_token": 3.75e-08, "litellm_provider": "cohere_chat", "max_input_tokens": 128000, "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_token": 3.75e-08, + "output_cost_per_token": 1.5e-07, "source": "https://docs.cohere.com/v2/docs/command-r7b", "supports_function_calling": true, "supports_tool_choice": true @@ -14511,6 +14612,38 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "fireworks_ai/accounts/fireworks/models/deepseek-v4-flash": { + "cache_read_input_token_cost": 2.8e-08, + "input_cost_per_token": 1.4e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 384000, + "max_tokens": 384000, + "mode": "chat", + "output_cost_per_token": 2.8e-07, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/accounts/fireworks/models/deepseek-v4-pro": { + "cache_read_input_token_cost": 1.45e-07, + "input_cost_per_token": 1.74e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 384000, + "max_tokens": 384000, + "mode": "chat", + "output_cost_per_token": 3.48e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, "fireworks_ai/accounts/fireworks/models/firefunction-v2": { "input_cost_per_token": 9e-07, "litellm_provider": "fireworks_ai", @@ -14586,43 +14719,64 @@ "input_cost_per_token": 1.4e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 202800, - "max_output_tokens": 202800, - "max_tokens": 202800, + "max_output_tokens": 131072, + "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 4.4e-06, - "source": "https://fireworks.ai/models/fireworks/glm-5p1", + "source": "https://docs.fireworks.ai/serverless/pricing", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/accounts/fireworks/models/glm-5p2": { + "cache_read_input_token_cost": 2.6e-07, + "input_cost_per_token": 1.4e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 4.4e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false }, "fireworks_ai/accounts/fireworks/models/gpt-oss-120b": { + "cache_read_input_token_cost": 1.5e-08, "input_cost_per_token": 1.5e-07, "litellm_provider": "fireworks_ai", "max_input_tokens": 131072, - "max_output_tokens": 131072, - "max_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, "mode": "chat", "output_cost_per_token": 6e-07, - "source": "https://fireworks.ai/pricing", + "source": "https://docs.fireworks.ai/serverless/pricing", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "supports_vision": false }, "fireworks_ai/accounts/fireworks/models/gpt-oss-20b": { - "input_cost_per_token": 5e-08, + "cache_read_input_token_cost": 3.5e-08, + "input_cost_per_token": 7e-08, "litellm_provider": "fireworks_ai", "max_input_tokens": 131072, - "max_output_tokens": 131072, - "max_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, "mode": "chat", - "output_cost_per_token": 2e-07, - "source": "https://fireworks.ai/pricing", + "output_cost_per_token": 3e-07, + "source": "https://docs.fireworks.ai/serverless/pricing", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "supports_vision": false }, "fireworks_ai/accounts/fireworks/models/kimi-k2-instruct": { "input_cost_per_token": 6e-07, @@ -14678,6 +14832,38 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "fireworks_ai/accounts/fireworks/models/kimi-k2p6": { + "cache_read_input_token_cost": 1.6e-07, + "input_cost_per_token": 9.5e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 4e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "fireworks_ai/accounts/fireworks/models/kimi-k2p7-code": { + "cache_read_input_token_cost": 1.9e-07, + "input_cost_per_token": 9.5e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 4e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "fireworks_ai/accounts/fireworks/models/llama-v3p1-405b-instruct": { "input_cost_per_token": 3e-06, "litellm_provider": "fireworks_ai", @@ -14795,6 +14981,38 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "fireworks_ai/accounts/fireworks/models/minimax-m2p7": { + "cache_read_input_token_cost": 6e-08, + "input_cost_per_token": 3e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 196608, + "max_output_tokens": 196608, + "max_tokens": 196608, + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/accounts/fireworks/models/minimax-m3": { + "cache_read_input_token_cost": 6e-08, + "input_cost_per_token": 3e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 512000, + "max_output_tokens": 512000, + "max_tokens": 512000, + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, "fireworks_ai/accounts/fireworks/models/mixtral-8x22b-instruct-hf": { "input_cost_per_token": 1.2e-06, "litellm_provider": "fireworks_ai", @@ -14847,6 +15065,38 @@ "supports_response_schema": true, "supports_tool_choice": false }, + "fireworks_ai/deepseek-v4-flash": { + "cache_read_input_token_cost": 2.8e-08, + "input_cost_per_token": 1.4e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 384000, + "max_tokens": 384000, + "mode": "chat", + "output_cost_per_token": 2.8e-07, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/deepseek-v4-pro": { + "cache_read_input_token_cost": 1.45e-07, + "input_cost_per_token": 1.74e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 384000, + "max_tokens": 384000, + "mode": "chat", + "output_cost_per_token": 3.48e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, "fireworks_ai/glm-4p7": { "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 6e-07, @@ -14867,15 +15117,80 @@ "input_cost_per_token": 1.4e-06, "litellm_provider": "fireworks_ai", "max_input_tokens": 202800, - "max_output_tokens": 202800, - "max_tokens": 202800, + "max_output_tokens": 131072, + "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 4.4e-06, - "source": "https://fireworks.ai/models/fireworks/glm-5p1", + "source": "https://docs.fireworks.ai/serverless/pricing", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/glm-5p1-fast": { + "cache_read_input_token_cost": 5.2e-07, + "input_cost_per_token": 2.8e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 202800, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 8.8e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/glm-5p2": { + "cache_read_input_token_cost": 2.6e-07, + "input_cost_per_token": 1.4e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 4.4e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/gpt-oss-120b": { + "cache_read_input_token_cost": 1.5e-08, + "input_cost_per_token": 1.5e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 6e-07, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/gpt-oss-20b": { + "cache_read_input_token_cost": 3.5e-08, + "input_cost_per_token": 7e-08, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 131072, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 3e-07, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false }, "fireworks_ai/kimi-k2p5": { "cache_read_input_token_cost": 1e-07, @@ -14891,6 +15206,70 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "fireworks_ai/kimi-k2p6": { + "cache_read_input_token_cost": 1.6e-07, + "input_cost_per_token": 9.5e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 4e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "fireworks_ai/kimi-k2p6-fast": { + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 2e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 8e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "fireworks_ai/kimi-k2p7-code": { + "cache_read_input_token_cost": 1.9e-07, + "input_cost_per_token": 9.5e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 4e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "fireworks_ai/kimi-k2p7-code-fast": { + "cache_read_input_token_cost": 3.8e-07, + "input_cost_per_token": 1.9e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 8e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "fireworks_ai/minimax-m2p1": { "cache_read_input_token_cost": 3e-08, "input_cost_per_token": 3e-07, @@ -14905,6 +15284,54 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "fireworks_ai/minimax-m2p7": { + "cache_read_input_token_cost": 6e-08, + "input_cost_per_token": 3e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 196608, + "max_output_tokens": 196608, + "max_tokens": 196608, + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/minimax-m3": { + "cache_read_input_token_cost": 6e-08, + "input_cost_per_token": 3e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 512000, + "max_output_tokens": 512000, + "max_tokens": 512000, + "mode": "chat", + "output_cost_per_token": 1.2e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/qwen3p7-plus": { + "cache_read_input_token_cost": 8e-08, + "input_cost_per_token": 4e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 262144, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_token": 1.6e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "fireworks_ai/nomic-ai/nomic-embed-text-v1": { "input_cost_per_token": 8e-09, "litellm_provider": "fireworks_ai-embedding-models", @@ -25112,6 +25539,21 @@ "supports_tool_choice": true, "supports_vision": true }, + "mistral/mistral-medium-3-5": { + "input_cost_per_token": 1.5e-06, + "litellm_provider": "mistral", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 7.5e-06, + "source": "https://docs.mistral.ai/models/model-cards/mistral-medium-3-5-26-04", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "mistral/mistral-small": { "input_cost_per_token": 1e-07, "litellm_provider": "mistral", @@ -39390,6 +39832,22 @@ "litellm_provider": "fireworks_ai", "mode": "chat" }, + "fireworks_ai/accounts/fireworks/models/qwen3p7-plus": { + "cache_read_input_token_cost": 8e-08, + "input_cost_per_token": 4e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 262144, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_token": 1.6e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "fireworks_ai/accounts/fireworks/models/qwq-32b": { "max_tokens": 131072, "max_input_tokens": 131072, @@ -39552,6 +40010,54 @@ "litellm_provider": "fireworks_ai", "mode": "chat" }, + "fireworks_ai/accounts/fireworks/routers/glm-5p1-fast": { + "cache_read_input_token_cost": 5.2e-07, + "input_cost_per_token": 2.8e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 202800, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 8.8e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "fireworks_ai/accounts/fireworks/routers/kimi-k2p6-fast": { + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 2e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 8e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "fireworks_ai/accounts/fireworks/routers/kimi-k2p7-code-fast": { + "cache_read_input_token_cost": 3.8e-07, + "input_cost_per_token": 1.9e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 8e-06, + "source": "https://docs.fireworks.ai/serverless/pricing", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "scaleway/qwen/qwen3.5-397b-a17b": { "input_cost_per_token": 6e-07, "litellm_provider": "scaleway", @@ -42839,4 +43345,105 @@ "supports_reasoning": true, "source": "https://serverless.tensormesh.ai/v1/models/openrouter" } - } + , + "deepseek-v4-flash": { + "cache_creation_input_token_cost": 0.0, + "cache_read_input_token_cost": 2.8e-09, + "input_cost_per_token": 1.4e-07, + "input_cost_per_token_cache_hit": 2.8e-09, + "litellm_provider": "deepseek", + "max_input_tokens": 1000000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 2.8e-07, + "source": "https://api-docs.deepseek.com/quick_start/pricing", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "deepseek-v4-pro": { + "cache_creation_input_token_cost": 0.0, + "cache_read_input_token_cost": 3.625e-09, + "input_cost_per_token": 4.35e-07, + "input_cost_per_token_cache_hit": 3.625e-09, + "litellm_provider": "deepseek", + "max_input_tokens": 1000000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 8.7e-07, + "source": "https://api-docs.deepseek.com/quick_start/pricing", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "deepseek/deepseek-v4-flash": { + "cache_creation_input_token_cost": 0.0, + "cache_read_input_token_cost": 2.8e-09, + "input_cost_per_token": 1.4e-07, + "input_cost_per_token_cache_hit": 2.8e-09, + "litellm_provider": "deepseek", + "max_input_tokens": 1000000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 2.8e-07, + "source": "https://api-docs.deepseek.com/quick_start/pricing", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": false + }, + "deepseek/deepseek-v4-pro": { + "cache_creation_input_token_cost": 0.0, + "cache_read_input_token_cost": 3.625e-09, + "input_cost_per_token": 4.35e-07, + "input_cost_per_token_cache_hit": 3.625e-09, + "litellm_provider": "deepseek", + "max_input_tokens": 1000000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 8.7e-07, + "source": "https://api-docs.deepseek.com/quick_start/pricing", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": false + } +} diff --git a/mypy-code-budget.json b/mypy-code-budget.json deleted file mode 100644 index 2cae0d661e9..00000000000 --- a/mypy-code-budget.json +++ /dev/null @@ -1,18 +0,0 @@ -{ - "import-not-found": { - "baseline": 8, - "slack": 3 - }, - "no-any-return": { - "baseline": 902, - "slack": 10 - }, - "no-untyped-def": { - "baseline": 4888, - "slack": 10 - }, - "valid-type": { - "baseline": 1, - "slack": 3 - } -} diff --git a/proxy_server_config.yaml b/proxy_server_config.yaml index f5f4e1956d4..6feffe036bd 100644 --- a/proxy_server_config.yaml +++ b/proxy_server_config.yaml @@ -64,39 +64,39 @@ model_list: litellm_params: model: openai/gpt-5-mini api_key: fake-key - api_base: https://exampleopenaiendpoint-production.up.railway.app/ + api_base: os.environ/FAKE_OPENAI_API_BASE - model_name: fake-openai-endpoint-2 litellm_params: model: openai/my-fake-model api_key: my-fake-key - api_base: https://exampleopenaiendpoint-production.up.railway.app/ + api_base: os.environ/FAKE_OPENAI_API_BASE stream_timeout: 0.001 rpm: 1 - model_name: fake-openai-endpoint-3 litellm_params: model: openai/my-fake-model api_key: my-fake-key - api_base: https://exampleopenaiendpoint-production.up.railway.app/ + api_base: os.environ/FAKE_OPENAI_API_BASE stream_timeout: 0.001 rpm: 1000 - model_name: fake-openai-endpoint-4 litellm_params: model: openai/my-fake-model api_key: my-fake-key - api_base: https://exampleopenaiendpoint-production.up.railway.app/ + api_base: os.environ/FAKE_OPENAI_API_BASE num_retries: 50 - model_name: fake-openai-endpoint-3 litellm_params: model: openai/my-fake-model-2 api_key: my-fake-key - api_base: https://exampleopenaiendpoint-production.up.railway.app/ + api_base: os.environ/FAKE_OPENAI_API_BASE stream_timeout: 0.001 rpm: 1000 - model_name: bad-model litellm_params: model: openai/bad-model api_key: os.environ/OPENAI_API_KEY - api_base: https://exampleopenaiendpoint-production.up.railway.app/ + api_base: os.environ/FAKE_OPENAI_API_BASE mock_timeout: True timeout: 60 rpm: 1000 @@ -106,7 +106,7 @@ model_list: litellm_params: model: openai/bad-model api_key: os.environ/OPENAI_API_KEY - api_base: https://exampleopenaiendpoint-production.up.railway.app/ + api_base: os.environ/FAKE_OPENAI_API_BASE rpm: 1000 model_info: health_check_timeout: 1 @@ -148,7 +148,7 @@ model_list: litellm_params: model: openai/my-fake-model api_key: my-fake-key - api_base: https://exampleopenaiendpoint-production.up.railway.app/ + api_base: os.environ/FAKE_OPENAI_API_BASE timeout: 1 - model_name: badly-configured-openai-endpoint litellm_params: diff --git a/pyproject.toml b/pyproject.toml index 8b1386aaf87..8ee2840b573 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -148,7 +148,6 @@ dev = [ "diff-cover==9.7.2", "flake8==7.3.0", "black==26.3.1", - "mypy==1.19.0", "basedpyright==1.39.7", "pytest==9.0.3", "pytest-mock==3.15.1", @@ -261,8 +260,6 @@ source-exclude = [ "litellm/proxy/enterprise", "**/__pycache__", "**/__pycache__/**", - "**/.mypy_cache", - "**/.mypy_cache/**", "**/.pytest_cache", "**/.pytest_cache/**", "**/.ruff_cache", @@ -278,9 +275,6 @@ version_files = [ "pyproject.toml:^version", ] -[tool.mypy] -plugins = "pydantic.mypy" - [tool.pytest.ini_options] asyncio_mode = "auto" asyncio_default_fixture_loop_scope = "session" diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index bb02ec01569..62ebdb559fc 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -1,27 +1,27 @@ { "ANN001": { "baseline": 2865, - "slack": 10 + "slack": 50 }, "ANN002": { "baseline": 64, - "slack": 3 + "slack": 5 }, "ANN003": { "baseline": 759, - "slack": 10 + "slack": 30 }, "ANN201": { "baseline": 1944, - "slack": 10 + "slack": 50 }, "ANN202": { "baseline": 858, - "slack": 10 + "slack": 30 }, "ANN204": { "baseline": 658, - "slack": 10 + "slack": 20 }, "ANN205": { "baseline": 117, @@ -33,7 +33,7 @@ }, "ANN401": { "baseline": 1886, - "slack": 10 + "slack": 50 }, "ASYNC230": { "baseline": 11, @@ -45,15 +45,15 @@ }, "B006": { "baseline": 180, - "slack": 3 + "slack": 10 }, "B008": { "baseline": 490, - "slack": 10 + "slack": 15 }, "B009": { "baseline": 79, - "slack": 10 + "slack": 5 }, "B010": { "baseline": 187, @@ -81,7 +81,7 @@ }, "BLE001": { "baseline": 2854, - "slack": 10 + "slack": 50 }, "C401": { "baseline": 8, @@ -109,7 +109,7 @@ }, "C901": { "baseline": 301, - "slack": 3 + "slack": 15 }, "D419": { "baseline": 6, @@ -125,7 +125,7 @@ }, "DTZ005": { "baseline": 229, - "slack": 10 + "slack": 15 }, "DTZ006": { "baseline": 10, @@ -165,7 +165,7 @@ }, "I001": { "baseline": 258, - "slack": 10 + "slack": 15 }, "LOG015": { "baseline": 5, @@ -189,11 +189,11 @@ }, "PERF403": { "baseline": 69, - "slack": 10 + "slack": 5 }, "PIE790": { "baseline": 263, - "slack": 10 + "slack": 15 }, "PIE800": { "baseline": 1, @@ -233,7 +233,7 @@ }, "PLR0913": { "baseline": 1813, - "slack": 3 + "slack": 50 }, "PLR1704": { "baseline": 3, @@ -245,7 +245,7 @@ }, "PLR1714": { "baseline": 252, - "slack": 10 + "slack": 15 }, "PLR1730": { "baseline": 7, @@ -265,11 +265,11 @@ }, "PLW0602": { "baseline": 215, - "slack": 10 + "slack": 15 }, "PLW0603": { "baseline": 183, - "slack": 3 + "slack": 10 }, "PLW1508": { "baseline": 188, @@ -301,15 +301,15 @@ }, "RET504": { "baseline": 709, - "slack": 10 + "slack": 20 }, "RUF010": { "baseline": 844, - "slack": 10 + "slack": 30 }, "RUF012": { "baseline": 158, - "slack": 3 + "slack": 10 }, "RUF015": { "baseline": 8, @@ -321,7 +321,7 @@ }, "RUF022": { "baseline": 80, - "slack": 10 + "slack": 5 }, "RUF023": { "baseline": 2, @@ -337,15 +337,15 @@ }, "RUF059": { "baseline": 69, - "slack": 10 + "slack": 5 }, "RUF100": { "baseline": 465, - "slack": 10 + "slack": 15 }, "S110": { "baseline": 222, - "slack": 10 + "slack": 15 }, "S112": { "baseline": 21, @@ -353,11 +353,11 @@ }, "SIM101": { "baseline": 58, - "slack": 10 + "slack": 5 }, "SIM102": { "baseline": 311, - "slack": 10 + "slack": 15 }, "SIM103": { "baseline": 119, @@ -412,20 +412,20 @@ "slack": 3 }, "TID251": { - "baseline": 2405, - "slack": 10 + "baseline": 2664, + "slack": 50 }, "TRY002": { "baseline": 528, - "slack": 10 + "slack": 20 }, "TRY004": { "baseline": 93, - "slack": 10 + "slack": 5 }, "TRY201": { "baseline": 409, - "slack": 10 + "slack": 15 }, "TRY203": { "baseline": 113, @@ -433,15 +433,15 @@ }, "TRY300": { "baseline": 853, - "slack": 10 + "slack": 30 }, "UP006": { "baseline": 12941, - "slack": 10 + "slack": 100 }, "UP007": { "baseline": 2520, - "slack": 10 + "slack": 50 }, "UP008": { "baseline": 2, @@ -469,7 +469,7 @@ }, "UP032": { "baseline": 609, - "slack": 10 + "slack": 20 }, "UP034": { "baseline": 1, @@ -477,7 +477,7 @@ }, "UP035": { "baseline": 2250, - "slack": 10 + "slack": 50 }, "UP036": { "baseline": 1, @@ -485,10 +485,10 @@ }, "UP037": { "baseline": 100, - "slack": 10 + "slack": 5 }, "UP045": { "baseline": 18417, - "slack": 10 + "slack": 100 } } diff --git a/ruff-strict.toml b/ruff-strict.toml index 8d517615244..1caa3567872 100644 --- a/ruff-strict.toml +++ b/ruff-strict.toml @@ -18,4 +18,16 @@ max-args = 5 "typing.Dict".msg = "Frozen dataclass / NamedTuple / ReadOnly TypedDict; create a Mapping alias with concrete value types if truly dynamic." "typing.Set".msg = "frozenset[X] or AbstractSet[X]." "typing.MutableSequence".msg = "Sequence[X]." -"typing.MutableMapping".msg = "See typing.Dict." \ No newline at end of file +"typing.MutableMapping".msg = "See typing.Dict." +# Unchecked casts: cast() lies to the type checker with no runtime guarantee. +# Validate into a concrete frozen type at the boundary (pydantic) instead. +# Per-call-site coverage lives in check_type_discipline.py (LIT006); this freezes +# new cast imports. Suppress (with a reason) via `# noqa: TID251 # `. +"typing.cast".msg = "No unchecked casts: validate into a frozen dataclass/NamedTuple/ReadOnly TypedDict at the boundary (pydantic)." +"typing_extensions.cast".msg = "Same as typing.cast." +# Unverified narrowing predicates: the checker never validates the guard body, so a +# wrong guard silently corrupts types. Banned outright (there are none today). +"typing.TypeGuard".msg = "Unverified narrowing. Parse into a concrete type, or use isinstance for a runtime-checked narrowing." +"typing_extensions.TypeGuard".msg = "Same as typing.TypeGuard." +"typing.TypeIs".msg = "Unverified narrowing (the body is trusted). Parse into a concrete type instead." +"typing_extensions.TypeIs".msg = "Same as typing.TypeIs." diff --git a/ruff.toml b/ruff.toml index 7baa1c5f92d..2db4122a30e 100644 --- a/ruff.toml +++ b/ruff.toml @@ -1,5 +1,5 @@ lint.ignore = ["F405", "E402", "E501", "F403"] -lint.extend-select = ["E501", "PLR0915", "T20", "PGH004", "RUF008", "RUF009", "RUF100"] +lint.extend-select = ["E501", "T20", "PGH004", "RUF008", "RUF009", "RUF100"] # RUF100 (unused-noqa) only knows the rules enabled in THIS config, so it would strip # `# noqa` directives that protect rules enforced elsewhere. List those codes as external # so RUF100 leaves their directives alone: the strict gate (ruff-strict.toml) and upstream @@ -23,9 +23,4 @@ exclude = ["litellm/types/*", "litellm/__init__.py", "litellm/proxy/example_conf "litellm/llms/azure_ai/embed/__init__.py" = ["F401"] "litellm/llms/azure_ai/rerank/__init__.py" = ["F401"] "litellm/llms/bedrock/chat/__init__.py" = ["F401"] -"litellm/proxy/utils.py" = ["F401", "PLR0915"] -"litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py" = ["PLR0915"] -"litellm/proxy/guardrails/guardrail_hooks/guardrail_benchmarks/test_eval.py" = ["PLR0915"] -"litellm/responses/streaming_iterator.py" = ["PLR0915"] -"litellm/files/main.py" = ["PLR0915"] -"litellm/llms/litellm_proxy/skills/sandbox_executor.py" = ["PLR0915"] +"litellm/proxy/utils.py" = ["F401"] diff --git a/scripts/budget_ratchet_check.py b/scripts/budget_ratchet_check.py new file mode 100644 index 00000000000..861d65489e8 --- /dev/null +++ b/scripts/budget_ratchet_check.py @@ -0,0 +1,160 @@ +#!/usr/bin/env python3 +"""Non-gating ratchet guard: budget ceilings may only fall, never rise. + +Every `*-budget.json` file (ruff-strict, type-discipline, basedpyright-code) is a +one-way ratchet: each rule's ceiling is `baseline + slack`, and the whole point is +to drive that number DOWN over time. This check compares every budget file against +its own content at the merge-base with the target branch and fails (exits 1, red) if: + + * a rule's ceiling went up, + * a rule was dropped from a budget (its ceiling effectively became infinite), or + * an entire budget file was deleted. + +New rules and lowered/equal ceilings are fine. + +This is deliberately NOT a gating check. It should turn the run red so that a +loosening is impossible to miss in review, but it must stay OUT of the +branch-protection required-checks list: a justified bump (e.g. banning a new API, +which mechanically raises a baseline) can then still be merged by a human who has +seen the red and accepted it. + +Usage: + python scripts/budget_ratchet_check.py [--base REF] [budget.json ...] + +Stdlib only. +""" + +from __future__ import annotations + +import argparse +import json +import subprocess +import sys +from pathlib import Path +from typing import NamedTuple + +REPO_ROOT = Path(__file__).resolve().parent.parent +DEFAULT_BASE = "origin/litellm_internal_staging" +DEFAULT_BUDGETS: tuple[str, ...] = ( + "ruff-strict-budget.json", + "type-discipline-budget.json", + "basedpyright-code-budget.json", +) + + +class Regression(NamedTuple): + budget: str + rule: str + detail: str + + +def _run(cmd: list[str]) -> subprocess.CompletedProcess[str]: + return subprocess.run(cmd, cwd=REPO_ROOT, capture_output=True, text=True) + + +def _merge_base(base: str) -> str: + """The common ancestor of `base` and HEAD, so unrelated base drift is ignored.""" + proc = _run(["git", "merge-base", base, "HEAD"]) + return proc.stdout.strip() or base + + +def _load_head(rel: str) -> dict | None: + path = REPO_ROOT / rel + if not path.exists(): + return None + return json.loads(path.read_text()) + + +def _ref_is_commit(ref: str) -> bool: + return _run(["git", "rev-parse", "--verify", "--quiet", f"{ref}^{{commit}}"]).returncode == 0 + + +def _load_base(rel: str, ref: str) -> dict | None: + """Budget content at `ref`, or None when the file did not exist there. + + `ref` is verified as a real commit by the caller, so a non-zero `git show` here means + the path was absent at that commit, not that the ref itself is unresolvable. + """ + proc = _run(["git", "show", f"{ref}:{rel}"]) + if proc.returncode != 0: + return None + return json.loads(proc.stdout) + + +def _caps(budget: dict) -> dict[str, int]: + """Map each rule to its ceiling (baseline + slack); skip malformed specs.""" + caps: dict[str, int] = {} + for rule, spec in budget.items(): + if isinstance(spec, dict): + caps[rule] = int(spec.get("baseline", 0)) + int(spec.get("slack", 0)) + return caps + + +def regressions_for(rel: str, base: dict | None, head: dict | None) -> list[Regression]: + if base is None: + return [] # new budget file: nothing to ratchet against yet + if head is None: + return [Regression(rel, "*", "budget file was deleted (every ceiling removed)")] + + base_caps = _caps(base) + head_caps = _caps(head) + return [ + Regression( + rel, + rule, + f"rule dropped (ceiling {base_cap} -> removed)" + if rule not in head_caps + else f"ceiling raised {base_cap} -> {head_caps[rule]}", + ) + for rule, base_cap in sorted(base_caps.items()) + if rule not in head_caps or head_caps[rule] > base_cap + ] + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--base", default=DEFAULT_BASE) + parser.add_argument("budgets", nargs="*", help="budget files to check") + args = parser.parse_args() + budgets = args.budgets or list(DEFAULT_BUDGETS) + + ref = _merge_base(args.base) + if not _ref_is_commit(ref): + print( + f"FAIL: base ref {ref!r} does not resolve to a commit, so the ratchet has nothing " + f"to compare against; refusing to pass vacuously (check the --base / BASE_SHA value)", + file=sys.stderr, + ) + return 1 + + regressions: list[Regression] = [] + checked: list[str] = [] + for rel in budgets: + base = _load_base(rel, ref) + head = _load_head(rel) + if base is None and head is None: + continue + if base is None: + print(f"skip {rel}: new file (no base at {args.base} to ratchet against)") + continue + checked.append(rel) + regressions.extend(regressions_for(rel, base, head)) + + if regressions: + print(f"FAIL: budget ceiling(s) loosened vs base {args.base} (merge-base {ref[:12]}):") + for reg in regressions: + print(f" {reg.budget} {reg.rule}: {reg.detail}") + print( + "Budgets are one-way ratchets and may only go down or stay flat. This " + "check is non-gating: if the increase is justified (e.g. a newly banned " + "API), a human can merge over the red after acknowledging it." + ) + return 1 + + suffix = f" ({', '.join(checked)})" if checked else "" + print(f"OK: no budget ceiling increased vs base {args.base}{suffix}") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/scripts/check_any_discipline.py b/scripts/check_any_discipline.py deleted file mode 100644 index 3185953d473..00000000000 --- a/scripts/check_any_discipline.py +++ /dev/null @@ -1,556 +0,0 @@ -#!/usr/bin/env python3 -"""Any-discipline gate: fail when a *changed* file holds a value typed `Any`. - -Where ruff, `mypy --strict`, and even basedpyright's `reportAny` stop short, this -catches the case that actually bites: a *union* hiding an `Any`. For example -`re.Match.group()` -> `str | Any`, `json.loads()` -> `Any`, and bare `list`/`dict` --> `list[Any]`/`dict[..., Any]`. Any value whose inferred type *contains* `Any` -(recursively, through unions / generics / tuples) is reported. - -Scope: changed-only, changed-lines ----------------------------------- -litellm already contains a large amount of pre-existing `Any` (a single legacy -file can have >100 findings), and a whole-tree scan would have to re-export types -for litellm's entire import closure on every run (~2 min, ~3 GB). So this gate is -*changed-only* and reports a finding only on a line that the diff against -`--base` actually adds or edits (untracked files count as wholly new). A brand -new file is therefore checked in full, while editing a legacy file only requires -*your* lines to be clean -- you can't introduce an `X | Any`, but you aren't -forced to clean the file's existing debt. This mirrors how `ruff_strict_gate.py` -blames a change only for the violations it introduces; cold legacy code is left -to the ratchet gates (mypy/basedpyright/ruff budgets). - -How it works ------------- -It loads `litellm/mypy.ini` (the same config `make lint-mypy` uses, so findings -match what developers already see), builds the changed files with mypy asking for -its exported expression->type map, and walks each file's AST applying a recursive -"contains Any" predicate -- the test `mypy --disallow-any-expr` uses internally -but applies inconsistently (python/mypy#12856). - -mypy only re-exports types for modules it re-type-checks, so for each target we -invalidate just its cached hash (deps stay warm) to force a fast re-check against -a persisted incremental cache (.mypy_cache_any). - -Rules ------ -Codes share the `LIT***` namespace with `scripts/check_type_discipline.py` (PR -#30500), which owns LIT001/002/003/004/006/007/008. This gate claims the rest: -LIT009 A value expression's inferred type is, or contains, `Any`. - Suppress with `# any-ok: ` on the offending line. -LIT005 An `# any-ok` suppression without a reason (the shared - suppression-needs-a-reason code, same as `# cast-ok` / `# guard-ok`). -LIT000 Setup failure: mypy could not build, or a target file could not be read. - -`Any`s produced purely by an already-reported error, and the special-form / -implementation-artifact internal `Any`s, are ignored. A bound method *reference* -whose signature mentions `Any` is not flagged -- only the value its call produces. - -Usage ------ - # gate mode (CI / pre-push): check changed lines under litellm/ - uv run --no-sync python scripts/check_any_discipline.py --changed --base origin/litellm_internal_staging - - # whole-file spot-check (no line filter), paths relative to repo root - uv run --no-sync python scripts/check_any_discipline.py litellm/budget_manager.py - -Exit code 1 if any Any-tainted value is found, 2 on a setup/usage error. -""" - -from __future__ import annotations - -import argparse -import json -import os -import re -import subprocess -import sys -import tokenize -from collections.abc import Iterable, Sequence -from pathlib import Path -from typing import NamedTuple - -try: - from mypy import build - from mypy.config_parser import parse_config_file - from mypy.find_sources import create_source_list - from mypy.fscache import FileSystemCache - from mypy.modulefinder import BuildSource - from mypy.nodes import AssignmentStmt, Expression, NameExpr, Node - from mypy.options import Options - from mypy.types import ( - AnyType, - CallableType, - Instance, - Overloaded, - TupleType, - Type, - TypeOfAny, - UnionType, - get_proper_type, - ) -except ImportError: # pragma: no cover - environment guard - sys.stderr.write( - "check_any_discipline: mypy is not importable in this interpreter.\n" - "Run it through the project environment, e.g.\n" - " uv run --no-sync python scripts/check_any_discipline.py --changed\n" - ) - raise SystemExit(2) - - -REPO_ROOT = Path(__file__).resolve().parent.parent -LITELLM_DIR = REPO_ROOT / "litellm" -MYPY_INI = LITELLM_DIR / "mypy.ini" -CACHE_DIR = REPO_ROOT / ".mypy_cache_any" -PY_TAG = f"{sys.version_info.major}.{sys.version_info.minor}" -DEFAULT_BASE = "origin/litellm_internal_staging" - -MIN_REASON_LEN = 3 -ANY_OK_RE = re.compile(r"#\s*any-ok(?::\s*(?P.*))?") -_HUNK_RE = re.compile(r"^@@ -\d+(?:,\d+)? \+(\d+)(?:,(\d+))? @@") - -# Files allowed to surface `Any` (the typed/untyped boundary). A finding is -# skipped if any fragment below is a substring of the file's posix path. Keep -# this tight -- prefer a line-level `# any-ok: ` over a blanket exemption. -BOUNDARY_PATHS: frozenset[str] = frozenset() - -# `Any` kinds that are not actionable: produced by an already-reported error, or -# an internal placeholder that never corresponds to a concrete runtime value. -# NOTE: `special_form` is deliberately NOT here. In mypy 1.19 the `Any` in -# typeshed unions like `re.Match.group() -> str | Any` is tagged `special_form`, -# and that union is the headline case this gate exists to catch. -_HARMLESS_ANY = frozenset( - kind - for kind in ( - TypeOfAny.from_error, - getattr(TypeOfAny, "implementation_artifact", None), - ) - if kind is not None -) - -# AST attributes that point OUTSIDE the syntactic subtree (a RefExpr's resolved -# definition, a node's TypeInfo). Skipping exactly these two makes a generic -# child-walk equivalent to mypy's TraverserVisitor -- validated to the node -# against ExtendedTraverserVisitor across the full grammar (see commit notes). -_NON_SYNTACTIC_ATTRS = frozenset({"node", "info"}) - - -class Violation(NamedTuple): - path: Path - line: int - col: int - code: str - message: str - - def render(self) -> str: - return f"{self.path}:{self.line}:{self.col}: {self.code} {self.message}" - - -# --------------------------------------------------------------------------- # -# The "contains Any" predicate -# --------------------------------------------------------------------------- # - - -def contains_any(t: Type, _seen: set[int] | None = None) -> bool: - """True if a *value* of type ``t`` carries `Any` anywhere meaningful.""" - seen = _seen if _seen is not None else set() - p = get_proper_type(t) - if id(p) in seen: - return False - seen.add(id(p)) - - # A function/method *reference* whose signature mentions Any is not itself an - # unsafe value -- only its eventual call result is. Don't recurse into it. - if isinstance(p, (CallableType, Overloaded)): - return False - if isinstance(p, AnyType): - return p.type_of_any not in _HARMLESS_ANY - if isinstance(p, UnionType): - return any(contains_any(item, seen) for item in p.items) - if isinstance(p, Instance): - return any(contains_any(arg, seen) for arg in p.args) - if isinstance(p, TupleType): - return any(contains_any(item, seen) for item in p.items) - return False - - -# --------------------------------------------------------------------------- # -# Generic, leak-free AST walk (works under a mypyc-compiled mypy, which forbids -# subclassing TraverserVisitor) -# --------------------------------------------------------------------------- # - - -def _walk_file(tree: Node) -> tuple[list[Expression], set[int]]: - """Return (every Expression in `tree`, ids of simple assignment-target names). - - The walk follows only syntactic children (every attribute except the two - non-syntactic back-references), so it never escapes the module. Simple - ``x = `` name targets are collected separately so we don't double-report - the assigned name as an echo of an Any rvalue. - """ - exprs: list[Expression] = [] - skip_lvalues: set[int] = set() - stack: list[object] = [tree] - seen: set[int] = set() - while stack: - n = stack.pop() - if isinstance(n, Node): - if id(n) in seen: - continue - seen.add(id(n)) - if isinstance(n, Expression): - exprs.append(n) - if isinstance(n, AssignmentStmt): - for lvalue in n.lvalues: - if isinstance(lvalue, NameExpr): - skip_lvalues.add(id(lvalue)) - for name in dir(n): - if name.startswith("__") or name in _NON_SYNTACTIC_ATTRS: - continue - try: - val = getattr(n, name) - except Exception: - continue - if callable(val): - continue - if isinstance(val, (Node, list, tuple)): - stack.append(val) - elif isinstance(n, (list, tuple)): - stack.extend(n) - return exprs, skip_lvalues - - -def find_any_in_tree(tree: Node, idmap: dict[int, Type]) -> list[tuple[int, int, str]]: - exprs, skip_lvalues = _walk_file(tree) - findings: list[tuple[int, int, str]] = [] - for expr in exprs: - if id(expr) in skip_lvalues: - continue - t = idmap.get(id(expr)) - if t is not None and contains_any(t): - findings.append((expr.line, expr.column, str(get_proper_type(t)))) - - out: list[tuple[int, int, str]] = [] - seen_pos: set[tuple[int, int]] = set() - for line, col, typ in sorted(findings): - if line < 1 or (line, col) in seen_pos: - continue - seen_pos.add((line, col)) - out.append((line, col, typ)) - return out - - -# --------------------------------------------------------------------------- # -# Comment scanning (LIT005 + any-ok suppression) -# --------------------------------------------------------------------------- # - - -def _reason_ok(reason: str | None) -> bool: - return reason is not None and len(reason.strip()) >= MIN_REASON_LEN - - -def scan_any_ok( - path: Path, source: str -) -> tuple[frozenset[int], tuple[Violation, ...]]: - """Return (lines with a valid any-ok suppression, LIT005 violations).""" - try: - tokens = tokenize.generate_tokens( - iter(source.splitlines(keepends=True)).__next__ - ) - comments = tuple( - (t.start[0], t.string) for t in tokens if t.type == tokenize.COMMENT - ) - except tokenize.TokenError: - return frozenset(), () - - ok_lines: set[int] = set() - violations: list[Violation] = [] - for line, text in comments: - m = ANY_OK_RE.search(text) - if m is None: - continue - if _reason_ok(m.group("reason")): - ok_lines.add(line) - else: - violations.append( - Violation( - path, - line, - 0, - "LIT005", - "any-ok requires a reason: `# any-ok: `", - ) - ) - return frozenset(ok_lines), tuple(violations) - - -# --------------------------------------------------------------------------- # -# mypy build (parity with `make lint-mypy`) + forced target re-check -# --------------------------------------------------------------------------- # - - -def _build_options() -> Options: - opts = Options() - if MYPY_INI.exists(): - parse_config_file(opts, lambda: None, str(MYPY_INI), sys.stdout, sys.stderr) - opts.export_types = True - opts.preserve_asts = True - opts.incremental = True - opts.cache_dir = str(CACHE_DIR) - opts.show_traceback = False - return opts - - -def _meta_path(module: str) -> Path: - return CACHE_DIR / PY_TAG / (module.replace(".", os.sep) + ".meta.json") - - -def _force_recheck(sources: Sequence[BuildSource]) -> None: - """Invalidate each target's cached entry so mypy re-type-checks (and thus - re-exports types + preserves the AST for) exactly these modules, while their - dependencies stay warm. A missing entry is a cold build for that module. - - mypy trusts a cache entry whenever the source mtime matches the cached one - (it never re-hashes on that fast path), so we must break BOTH: zero the - cached mtime to force a re-hash, and corrupt the cached hash so the re-hash - mismatches and the module is treated as changed.""" - for src in sources: - if not src.module: - continue - meta = _meta_path(src.module) - if not meta.exists(): - continue - try: - data = json.loads(meta.read_text()) - data["hash"] = "0" * 40 - data["mtime"] = 0 - meta.write_text(json.dumps(data)) - except (OSError, ValueError): - continue - - -def check_files(rel_paths: Sequence[str]) -> tuple[Violation, ...]: - """`rel_paths` are relative to the litellm package dir (the build cwd).""" - prev_cwd = Path.cwd() - os.chdir(LITELLM_DIR) - try: - opts = _build_options() - fscache = FileSystemCache() - sources = create_source_list(list(rel_paths), opts, fscache) - _force_recheck(sources) - try: - res = build.build(sources, options=opts, fscache=fscache) - except build.CompileError as exc: - joined = "; ".join(exc.messages[:3]) or "blocking error" - return ( - Violation( - Path(rel_paths[0]), - 0, - 0, - "LIT000", - f"mypy could not build: {joined}", - ), - ) - idmap = {id(expr): t for expr, t in res.types.items()} - # Resolve trees to absolute source paths while cwd is the build dir, since - # mypy stores the paths it was given (relative to this cwd). - trees: dict[str, Node] = {} - for state in res.graph.values(): - if state.path and state.tree is not None: - trees[os.path.realpath(state.path)] = state.tree - finally: - os.chdir(prev_cwd) - - out: list[Violation] = [] - for rel in rel_paths: - abs_path = (LITELLM_DIR / rel).resolve() - report_path = abs_path.relative_to(REPO_ROOT) - if _is_boundary(report_path): - continue - try: - source = abs_path.read_text(encoding="utf-8") - except (OSError, UnicodeDecodeError) as exc: - out.append( - Violation(report_path, 0, 0, "LIT000", f"could not read file: {exc}") - ) - continue - - ok_lines, ok_violations = scan_any_ok(report_path, source) - out.extend(ok_violations) - tree = trees.get(os.path.realpath(abs_path)) - if tree is None: - continue - for line, col, typ in find_any_in_tree(tree, idmap): - if line in ok_lines: - continue - out.append( - Violation( - report_path, - line, - col, - "LIT009", - f"value type contains Any -> {typ}", - ) - ) - return tuple(out) - - -# --------------------------------------------------------------------------- # -# File selection (changed-only, changed-lines) + driver -# --------------------------------------------------------------------------- # - - -class _AllLines: - """Sentinel: a wholly new / untracked file -- every line is in scope. - - A distinct object, not None, so that `line_map.get(path)` returning None for - a path absent from the map is never mistaken for "whole file in scope".""" - - -# A changed file's in-scope lines: a specific set, or every line. -LineScope = set[int] | _AllLines -ALL_LINES = _AllLines() - - -def _is_boundary(path: Path) -> bool: - posix = path.as_posix() - return any(frag in posix for frag in BOUNDARY_PATHS) - - -def _git(*args: str) -> list[str]: - result = subprocess.run( - ["git", "-C", str(REPO_ROOT), *args], - capture_output=True, - text=True, - check=True, - ) - return result.stdout.splitlines() - - -def _parse_added_lines(diff_text: str) -> dict[str, set[int]]: - """Map repo-relative path -> set of new-file line numbers the diff adds/edits.""" - changed: dict[str, set[int]] = {} - path: str | None = None - for line in diff_text.splitlines(): - if line.startswith("+++ b/"): - path = line[6:] - elif path and (m := _HUNK_RE.match(line)): - start = int(m.group(1)) - count = int(m.group(2)) if m.group(2) is not None else 1 - if count: - changed.setdefault(path, set()).update(range(start, start + count)) - return changed - - -def changed_line_map(base: str) -> dict[str, LineScope] | None: - """Repo-relative `.py` path under litellm/ -> changed line numbers (or - ALL_LINES for untracked files). Compares the working tree to the merge-base - with `base`, so it covers committed-on-branch + unstaged edits. None if git - is unavailable / not a repo.""" - try: - merge_base = _git("merge-base", base, "HEAD") - point = merge_base[0].strip() if merge_base else base - diff = "\n".join( - _git( - "diff", - "--unified=0", - "--no-color", - "--diff-filter=d", - point, - "--", - "litellm", - ) - ) - untracked = _git("ls-files", "--others", "--exclude-standard", "--", "litellm") - except (subprocess.CalledProcessError, FileNotFoundError): - return None - - out: dict[str, LineScope] = {} - for name, lines in _parse_added_lines(diff).items(): - if name.endswith(".py") and (REPO_ROOT / name).exists(): - out[name] = lines - for name in untracked: - if name.endswith(".py") and (REPO_ROOT / name).exists(): - out[name] = ALL_LINES - return out - - -def _to_litellm_relative(paths: Iterable[Path]) -> list[str]: - rels: list[str] = [] - for p in sorted(paths): - try: - rels.append(p.resolve().relative_to(LITELLM_DIR).as_posix()) - except ValueError: - continue - return rels - - -def _in_scope(v: Violation, line_map: dict[str, LineScope] | None) -> bool: - """A finding survives if line filtering is off (explicit paths), it's a build - error, or its line is one the diff added/edited.""" - if line_map is None or v.code == "LIT000": - return True - lines = line_map.get(v.path.as_posix()) - return lines is ALL_LINES or (lines is not None and v.line in lines) - - -def main(argv: Sequence[str]) -> int: - parser = argparse.ArgumentParser( - description="Any-discipline gate (changed-only, changed-lines)." - ) - parser.add_argument( - "paths", - nargs="*", - help="explicit files (repo-root relative); whole-file, no line filter", - ) - parser.add_argument( - "--changed", - action="store_true", - help="check changed lines under litellm/ vs --base", - ) - parser.add_argument("--base", default=os.environ.get("ANY_GATE_BASE", DEFAULT_BASE)) - args = parser.parse_args(list(argv)) - - line_map: dict[str, LineScope] | None = None - if args.changed: - line_map = changed_line_map(args.base) - if line_map is None: - print( - "check_any_discipline: not a git repository; nothing to check", - file=sys.stderr, - ) - return 0 - rel_paths = _to_litellm_relative( - (REPO_ROOT / name).resolve() for name in line_map - ) - elif args.paths: - rel_paths = _to_litellm_relative((REPO_ROOT / p).resolve() for p in args.paths) - else: - parser.error("pass --changed or explicit file paths") - return 2 - - if not rel_paths: - print("OK: no changed Python lines under litellm/ to check") - return 0 - - violations = tuple(v for v in check_files(rel_paths) if _in_scope(v, line_map)) - - for v in sorted(violations): - print(v.render()) - - if violations: - n = len(violations) - print( - f"\nFAIL: {n} Any-discipline violation(s) on changed lines.\n" - "Give the value a concrete type, or annotate the line `# any-ok: `.", - file=sys.stderr, - ) - return 1 - print( - f"OK: {len(rel_paths)} changed file(s) under litellm/ have no Any-typed values on changed lines" - ) - return 0 - - -if __name__ == "__main__": - raise SystemExit(main(sys.argv[1:])) diff --git a/scripts/check_type_discipline.py b/scripts/check_type_discipline.py new file mode 100644 index 00000000000..6e152541863 --- /dev/null +++ b/scripts/check_type_discipline.py @@ -0,0 +1,469 @@ +#!/usr/bin/env python3 +"""Type-discipline checker: the rules ruff can't enforce. + +Rules +----- +LIT001 Mutable collection in a type annotation, anywhere it appears: function + parameters, return types, class attributes, locals, and module globals. + Covers the builtins (dict/list/set, bare or parameterized), their typing + aliases (Dict/List/...), the collections concretes (deque/defaultdict/...), + and the mutable ABCs (MutableMapping/MutableSequence/MutableSet). A mutable + collection lets whoever holds it grow or rewrite it after the fact; annotate + a read-only view instead (Mapping/Sequence/AbstractSet/tuple[X, ...]/ + frozenset[X], or a frozen dataclass / NamedTuple / ReadOnly TypedDict) and + build it functionally (comprehension / map, not append-in-a-loop). + Suppress with `# mutable-ok: ` on the offending line. +LIT002 Mutable-collection *construction*: a list/dict/set literal or comprehension, or + a call to a mutable constructor (list/dict/set/deque/defaultdict/Counter/...). + Catches the unannotated seed-then-mutate pattern LIT001 cannot see (`acc = []`). + Build the value in one shot and freeze it: a `tuple`/`frozenset` wrapping a + generator (`tuple(f(x) for x in xs)`), a tuple literal, or a frozen dataclass / + NamedTuple / ReadOnly TypedDict. Generator expressions and `tuple`/`frozenset` + calls are not construction and pass. Annotation-internal lists (`Callable[[int], + str]`) are exempt. Suppress with `# mutable-ok: `. +LIT003 noqa suppression without rule codes or without a reason. + Required shape: `# noqa: TID251 # ` +LIT004 type/pyright/mypy ignore without bracketed codes or without a reason. + Required shape: `# pyright: ignore[reportArgumentType] # ` +LIT005 A `# mutable-ok` / `# cast-ok` / `# guard-ok` / `# kwargs-ok` + suppression without a reason. +LIT006 `cast(...)` call. typing.cast is an unchecked assertion (the moral equivalent + of TypeScript's `as`); it lies to the type checker with zero runtime guarantee. + Validate into a concrete frozen type at the boundary instead. + Suppress with `# cast-ok: ` on the call's first line. +LIT007 `TypeGuard[...]` / `TypeIs[...]` annotation. The narrowing predicate's body is + never verified by the checker, so a wrong guard silently corrupts types. + Prefer parsing into a concrete type. Suppress with `# guard-ok: `. +LIT008 `**kwargs` parameter. The keyword contract is erased and everything it carries + is effectively Any. ruff can force it to be typed (ANN003) but can't ban the + syntax. Declare explicit keyword params, or accept one frozen payload. `*args`, + by contrast, is fine when typed (it's just a tuple). Suppress: `# kwargs-ok: `. + +LIT000 Setup failure: a target file could not be read, or contains a syntax error. + Reported as a violation rather than crashing the run. + +Usage +----- + python check_type_discipline.py litellm/ tests/ + +Exit code 1 if any violation is found. Stdlib only. +""" + +from __future__ import annotations + +import ast +import io +import re +import sys +import tokenize +from dataclasses import dataclass +from pathlib import Path +from collections.abc import Iterable, Iterator, Sequence +from typing import NamedTuple + +# Mutable collection types, banned in *every* annotation. Name-based, so `dict`, +# `typing.Dict`, `collections.deque`, and `collections.abc.MutableMapping` all match +# however they were imported. The read-only interfaces (Mapping, Sequence, the +# immutable AbstractSet / `abc.Set`, Collection) and the immutable concretes (tuple, +# frozenset) are the escape hatch and are deliberately absent -- as is the bare name +# `Set`, which collides with the read-only `collections.abc.Set`. +MUTABLE_COLLECTIONS = frozenset(( + "dict", "list", "set", + "Dict", "List", "DefaultDict", "OrderedDict", "Counter", "Deque", "ChainMap", + "deque", "defaultdict", + "MutableMapping", "MutableSequence", "MutableSet", +)) + +# Callables whose result is a fresh *mutable* collection (LIT002). `tuple` and +# `frozenset` are deliberately absent -- they are the wrappers you reach for, and +# a generator expression fed to them is the blessed one-shot build. +MUTABLE_CONSTRUCTORS = frozenset(( + "dict", "list", "set", + "deque", "defaultdict", "OrderedDict", "Counter", "ChainMap", +)) +# A *qualified* call (`x.deque()`) counts as construction only for names that are rarely +# method names; `dict`/`list`/`set` are dropped here because `.dict()` / `.set()` / `.list()` +# are common methods (e.g. pydantic's `model.dict()`), not collection construction. A +# qualified `collections.deque(...)` still counts. +QUALIFIED_CONSTRUCTORS = MUTABLE_CONSTRUCTORS - frozenset(("dict", "list", "set")) +UNSAFE_GUARDS = frozenset(("TypeGuard", "TypeIs")) +MIN_REASON_LEN = 3 + +NOQA_RE = re.compile( + r"#\s*noqa" + r"(?P:\s*(?P[A-Z]+[0-9]+(?:\s*,\s*[A-Z]+[0-9]+)*))?" + r"(?P.*)", + re.IGNORECASE, +) +IGNORE_RE = re.compile( + r"#\s*(?:type|pyright|mypy):\s*ignore(?P\[[^\]]*\])?(?P.*)" +) +MUTABLE_OK_RE = re.compile(r"#\s*mutable-ok(?::\s*(?P.*))?") +CAST_OK_RE = re.compile(r"#\s*cast-ok(?::\s*(?P.*))?") +GUARD_OK_RE = re.compile(r"#\s*guard-ok(?::\s*(?P.*))?") +KWARGS_OK_RE = re.compile(r"#\s*kwargs-ok(?::\s*(?P.*))?") + +# Suppression tokens that must each carry a reason (LIT005). +OK_SUPPRESSIONS: tuple[tuple[str, re.Pattern[str]], ...] = ( + ("mutable-ok", MUTABLE_OK_RE), + ("cast-ok", CAST_OK_RE), + ("guard-ok", GUARD_OK_RE), + ("kwargs-ok", KWARGS_OK_RE), +) + + +class Violation(NamedTuple): + path: Path + line: int + code: str + message: str + + def render(self) -> str: + return f"{self.path}:{self.line}: {self.code} {self.message}" + + +@dataclass(frozen=True, slots=True) +class Comments: + """The lines carrying each valid `*-ok` suppression.""" + + mutable_ok_lines: frozenset[int] + cast_ok_lines: frozenset[int] + guard_ok_lines: frozenset[int] + kwargs_ok_lines: frozenset[int] + + +# --------------------------------------------------------------------------- # +# Comment scanning (LIT003 / LIT004 / LIT005) +# --------------------------------------------------------------------------- # + + +def _reason_of(rest: str) -> str: + return rest.strip().lstrip("#-").strip() + + +def _valid_ok(regex: re.Pattern[str], text: str) -> bool: + """True iff `text` carries this suppression with a reason of usable length.""" + m = regex.search(text) + return bool(m) and len((m.group("reason") or "").strip()) >= MIN_REASON_LEN + + +def _comment_violations(path: Path, line_no: int, text: str) -> Iterator[Violation]: + """Pure: all LIT003/004/005 findings for one comment.""" + for token, regex in OK_SUPPRESSIONS: + m = regex.search(text) + if m and len((m.group("reason") or "").strip()) < MIN_REASON_LEN: + yield Violation(path, line_no, "LIT005", f"{token} requires a reason: `# {token}: `") + + m = NOQA_RE.search(text) + if m: + if not m.group("codes"): + yield Violation(path, line_no, "LIT003", "noqa requires rule codes: `# noqa: XXX123 # `") + elif len(_reason_of(m.group("rest"))) < MIN_REASON_LEN: + yield Violation(path, line_no, "LIT003", "noqa requires a reason: `# noqa: XXX123 # `") + + m = IGNORE_RE.search(text) + if m: + codes = m.group("codes") + if not codes or codes == "[]": + yield Violation(path, line_no, "LIT004", + "ignore requires codes: `# pyright: ignore[ruleName] # `") + elif len(_reason_of(m.group("rest"))) < MIN_REASON_LEN: + yield Violation(path, line_no, "LIT004", + "ignore requires a reason: `# pyright: ignore[ruleName] # `") + + +def scan_comments(path: Path, source: str) -> tuple[Comments, tuple[Violation, ...]]: + try: + tokens = tokenize.generate_tokens(io.StringIO(source).readline) + comment_toks = tuple((t.start[0], t.string) for t in tokens if t.type == tokenize.COMMENT) + except (tokenize.TokenError, SyntaxError): + # tokenize raises TokenError (EOF mid-construct) or a SyntaxError subclass + # (IndentationError / TabError) on malformed source; defer to ast.parse below, + # which re-raises and is reported as LIT000 rather than crashing the run. + return Comments(frozenset(), frozenset(), frozenset(), frozenset()), () + + def _lines_with(regex: re.Pattern[str]) -> frozenset[int]: + return frozenset(line for line, text in comment_toks if _valid_ok(regex, text)) + + return ( + Comments( + mutable_ok_lines=_lines_with(MUTABLE_OK_RE), + cast_ok_lines=_lines_with(CAST_OK_RE), + guard_ok_lines=_lines_with(GUARD_OK_RE), + kwargs_ok_lines=_lines_with(KWARGS_OK_RE), + ), + tuple(v for line, text in comment_toks for v in _comment_violations(path, line, text)), + ) + + +# --------------------------------------------------------------------------- # + + +def mutable_names_in(annotation: ast.expr) -> Iterator[str]: + """Yield mutable-collection names anywhere inside an annotation expression. + + Matches bare names (`dict`, `MutableMapping`) and dotted access (`typing.Dict`, + `collections.deque`, `collections.abc.MutableMapping`), descends through nesting + (`Mapping[str, list[int]]`, `tuple[set[int], ...]`) and string forward references. + """ + for node in ast.walk(annotation): + if isinstance(node, ast.Name) and node.id in MUTABLE_COLLECTIONS: + yield node.id + elif isinstance(node, ast.Attribute) and node.attr in MUTABLE_COLLECTIONS: + yield node.attr + elif isinstance(node, ast.Constant): + value: object = node.value # forward references arrive as string constants + if isinstance(value, str): + try: + inner = ast.parse(value, mode="eval").body + except SyntaxError: + continue + yield from mutable_names_in(inner) + + +def _mutable_ann(path: Path, line: int, name: str, where: str) -> Violation: + return Violation( + path, line, "LIT001", + f"mutable `{name}` in {where}: a mutable collection can be grown or rewritten " + f"by whoever holds it. Annotate a read-only view -- Mapping[...], Sequence[...], " + f"AbstractSet[...], tuple[X, ...], frozenset[X], or a frozen dataclass / " + f"NamedTuple / ReadOnly TypedDict -- and build it functionally, not by " + f"append-in-a-loop (suppress: `# mutable-ok: `)", + ) + + +def _annotation_violations( + path: Path, annotation: ast.expr | None, line: int, where: str, ok_lines: frozenset[int] +) -> Iterator[Violation]: + if annotation is None or line in ok_lines: + return + yield from (_mutable_ann(path, line, name, where) for name in mutable_names_in(annotation)) + + +def _function_violations( + path: Path, node: ast.FunctionDef | ast.AsyncFunctionDef, comments: Comments +) -> Iterator[Violation]: + mutable_ok = comments.mutable_ok_lines + args = node.args + for arg in (*args.posonlyargs, *args.args, *args.kwonlyargs): + yield from _annotation_violations( + path, arg.annotation, arg.lineno, f"parameter `{arg.arg}` of `{node.name}`", mutable_ok + ) + + # *args is allowed when typed (it's just a tuple); ruff ANN002 forces the + # annotation, so here we only add the LIT001 mutable-collection check on the element type. + if args.vararg is not None: + yield from _annotation_violations( + path, args.vararg.annotation, args.vararg.lineno, f"`*args` of `{node.name}`", mutable_ok + ) + + # **kwargs is banned outright (LIT008): it erases the keyword contract and forces + # Any-typing on everything it carries. ruff can require it be typed (ANN003) but + # cannot ban the syntax, so this rule does. + if args.kwarg is not None and args.kwarg.lineno not in comments.kwargs_ok_lines: + yield Violation( + path, args.kwarg.lineno, "LIT008", + f"`**{args.kwarg.arg}` is banned: it erases the keyword contract and forces " + f"Any-typing; declare explicit keyword parameters, or accept one frozen payload " + f"(frozen dataclass / NamedTuple / ReadOnly TypedDict) " + f"(suppress: `# kwargs-ok: `)", + ) + + if node.returns is not None: + yield from _annotation_violations( + path, node.returns, node.returns.lineno, f"return type of `{node.name}`", mutable_ok + ) + + +def iter_annotation_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]: + # Every annotation is in scope: signatures (params / *args / return) plus every + # `x: T` -- class attribute, local, or module global. The latter three are all + # ast.AnnAssign, so one walk covers them; only the signature annotations (which + # are not AnnAssign) need the dedicated helper. + for node in ast.walk(tree): + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): + yield from _function_violations(path, node, comments) + elif isinstance(node, ast.AnnAssign): + target = node.target.id if isinstance(node.target, ast.Name) else "" + yield from _annotation_violations( + path, node.annotation, node.lineno, + f"the type of `{target}`", comments.mutable_ok_lines, + ) + + +# --------------------------------------------------------------------------- # +# Unchecked casts (LIT006) and unverified narrowing predicates (LIT007) +# --------------------------------------------------------------------------- # + + +def _is_cast_call(node: ast.Call) -> bool: + """`cast(...)` or `typing.cast(...)`, however the name was imported/aliased. + + Name-based like MUTABLE_COLLECTIONS: a stray method called `.cast()` is a rare + false positive, suppressible with `# cast-ok: `. + """ + func = node.func + return (isinstance(func, ast.Name) and func.id == "cast") or ( + isinstance(func, ast.Attribute) and func.attr == "cast" + ) + + +def iter_cast_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]: + for node in ast.walk(tree): + if isinstance(node, ast.Call) and _is_cast_call(node) and node.lineno not in comments.cast_ok_lines: + yield Violation( + path, node.lineno, "LIT006", + "cast() is an unchecked assertion (the type checker takes it on faith); " + "validate into a frozen dataclass/NamedTuple/ReadOnly TypedDict at the " + "boundary instead (suppress: `# cast-ok: `)", + ) + + +def iter_guard_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]: + # TypeGuard/TypeIs are legal only as a function's return annotation (`-> TypeGuard[int]`), + # so the walk is confined to `node.returns`; a runtime name that merely happens to read + # `TypeGuard` is not a narrowing predicate. ruff bans the import; this flags the use. + for node in ast.walk(tree): + if not isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)) or node.returns is None: + continue + for sub in ast.walk(node.returns): + name = ( + sub.id if isinstance(sub, ast.Name) + else sub.attr if isinstance(sub, ast.Attribute) + else None + ) + if name in UNSAFE_GUARDS and sub.lineno not in comments.guard_ok_lines: + yield Violation( + path, sub.lineno, "LIT007", + f"`{name}` narrowing predicate: the checker never verifies the body, so a " + f"wrong guard silently corrupts types; parse into a concrete type instead " + f"(suppress: `# guard-ok: `)", + ) + + +# --------------------------------------------------------------------------- # +# Mutable-collection construction (LIT002) +# --------------------------------------------------------------------------- # + + +def _annotations_of(node: ast.AST) -> tuple[ast.expr | None, ...]: + """The annotation expressions a node carries (signatures and `x: T`).""" + if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)): + a = node.args + params = (*a.posonlyargs, *a.args, *a.kwonlyargs, a.vararg, a.kwarg) + return (*(p.annotation for p in params if p is not None), node.returns) + if isinstance(node, ast.AnnAssign): + return (node.annotation,) + return () + + +def _annotation_node_ids(tree: ast.AST) -> frozenset[int]: + """ids() of every node living inside an annotation. + + A list display inside an annotation (`Callable[[int], str]`) is type syntax, + not construction, so the LIT002 walk must skip those subtrees. + """ + return frozenset( + id(sub) + for node in ast.walk(tree) + for ann in _annotations_of(node) + if ann is not None + for sub in ast.walk(ann) + ) + + +def _construction_kind(node: ast.expr) -> str | None: + """Human label if `node` builds a mutable collection, else None.""" + if isinstance(node, ast.List): + return "list literal" + if isinstance(node, ast.ListComp): + return "list comprehension" + if isinstance(node, ast.Set): + return "set literal" + if isinstance(node, ast.SetComp): + return "set comprehension" + if isinstance(node, ast.Dict): + return "dict literal" + if isinstance(node, ast.DictComp): + return "dict comprehension" + if isinstance(node, ast.Call): + func = node.func + if isinstance(func, ast.Name) and func.id in MUTABLE_CONSTRUCTORS: + return f"`{func.id}()` constructor" + if isinstance(func, ast.Attribute) and func.attr in QUALIFIED_CONSTRUCTORS: + return f"`{func.attr}()` constructor" + return None + + +def iter_construction_violations(path: Path, tree: ast.AST, comments: Comments) -> Iterator[Violation]: + in_annotation = _annotation_node_ids(tree) + for node in ast.walk(tree): + if not isinstance(node, ast.expr) or id(node) in in_annotation: + continue + kind = _construction_kind(node) + if kind is None or node.lineno in comments.mutable_ok_lines: + continue + yield Violation( + path, node.lineno, "LIT002", + f"mutable {kind}: this builds a collection that can be grown or rewritten. " + f"Build it in one shot and freeze it -- a tuple/frozenset wrapping a generator " + f"(`tuple(f(x) for x in xs)`), a tuple literal, or a frozen dataclass / NamedTuple " + f"/ ReadOnly TypedDict (suppress: `# mutable-ok: `)", + ) + + +# --------------------------------------------------------------------------- # +# Driver +# --------------------------------------------------------------------------- # + + +def check_file(path: Path) -> tuple[Violation, ...]: + try: + source = path.read_text(encoding="utf-8") + except (OSError, UnicodeDecodeError) as exc: + return (Violation(path, 0, "LIT000", f"could not read file: {exc}"),) + + comments, violations = scan_comments(path, source) + + try: + tree = ast.parse(source, filename=str(path)) + except SyntaxError as exc: + return (*violations, Violation(path, exc.lineno or 0, "LIT000", f"syntax error: {exc.msg}")) + + return ( + *violations, + *iter_annotation_violations(path, tree, comments), + *iter_cast_violations(path, tree, comments), + *iter_guard_violations(path, tree, comments), + *iter_construction_violations(path, tree, comments), + ) + + +def collect_paths(raw: Iterable[str]) -> Iterator[Path]: + for item in raw: + p = Path(item) + if p.is_dir(): + yield from sorted(p.rglob("*.py")) + elif p.suffix == ".py": + yield p + + +def main(argv: Sequence[str]) -> int: + paths = tuple(a for a in argv if not a.startswith("-")) + if not paths: + print("usage: check_type_discipline.py ...", file=sys.stderr) + return 2 + + violations = sorted(v for path in collect_paths(paths) for v in check_file(path)) + for v in violations: + print(v.render()) + + if violations: + print(f"\n{len(violations)} violation(s).", file=sys.stderr) + return 1 + return 0 + + +if __name__ == "__main__": + raise SystemExit(main(sys.argv[1:])) + \ No newline at end of file diff --git a/scripts/type_check_gate.py b/scripts/type_check_gate.py index 5ff485f0b0f..0f9a44703f9 100644 --- a/scripts/type_check_gate.py +++ b/scripts/type_check_gate.py @@ -1,47 +1,38 @@ #!/usr/bin/env python3 -"""Per-rule count gate for mypy and basedpyright. +"""Per-rule count gate for basedpyright. -Each tool's output is reduced to a count of errors per *rule* (mypy error codes -like ``arg-type``, basedpyright rules like ``reportAny``) and checked against a -committed budget of the form ``{rule: {baseline, slack}}``, the same shape as +basedpyright's ``--outputjson`` is reduced to a count of errors per *rule* +(``reportAny``, ``reportArgumentType``, ...) and checked against a committed +budget of the form ``{rule: {baseline, slack}}``, the same shape as ``ruff-strict-budget.json``. A rule fails when its codebase-wide total exceeds ``baseline + slack``. Counts ignore file, line, and column, so a violation moving anywhere in the tree is invisible; only the per-rule total moves the needle. Unlike ``ruff_strict_gate.py`` this does *not* re-run the tool on the merge base -to compute a delta: a second mypy/basedpyright pass is minutes and gigabytes, -whereas ruff is milliseconds. The committed budget is the baseline instead -- -exactly how the previous per-file gate worked -- so keep it fresh with -``--update`` (ratchet), which re-captures every rule's count from the current -tree while preserving each rule's slack. Tool output is read from stdin, so the -caller decides how to invoke the tool (and from which cwd). +to compute a delta: a second basedpyright pass is minutes and gigabytes, whereas +ruff is milliseconds. The committed budget is the baseline instead -- exactly +how the previous per-file gate worked -- so keep it fresh with ``--update`` +(ratchet), which re-captures every rule's count from the current tree while +preserving each rule's slack. Tool output is read from stdin, so the caller +decides how to invoke basedpyright (and from which cwd). -mypy is parsed from its text output (one error per line, the rule code in a -trailing ``[bracket]``). basedpyright is parsed from ``--outputjson``: its text -diagnostics routinely wrap across lines, leaving the ``(reportRule)`` on a -continuation line away from the ``- error:`` marker, so line parsing -mis-attributes ~60% of errors -- the JSON carries an unambiguous ``rule`` field. +``--outputjson`` is used rather than text diagnostics because the latter wrap +across lines, leaving the ``(reportRule)`` on a continuation line away from the +``- error:`` marker, so line parsing mis-attributes ~60% of errors -- the JSON +carries an unambiguous ``rule`` field. """ import argparse import json -import re import sys from collections import Counter from pathlib import Path -from typing import Iterable, Mapping, NamedTuple +from typing import Mapping, NamedTuple REPO_ROOT = Path(__file__).resolve().parent.parent -# mypy: one error per line, e.g. `path:12: error: msg [arg-type]`. ERROR_LINE -# recognizes the line; MYPY_CODE pulls the trailing [code]. Kept separate so an -# error emitted without a code is still counted (under UNCODED), never dropped. -MYPY_ERROR = re.compile(r"^(?P.+?):\d+: error:") -MYPY_CODE = re.compile(r"\[(?P[a-z][a-z0-9-]*)\]\s*$") - -# Bucket for an error whose rule code we couldn't read (a mypy error with no -# code, or a basedpyright diagnostic with no `rule`). Counted so it's gated. +# Bucket for a basedpyright diagnostic with no `rule`. Counted so it's gated. UNCODED = "" # Ceiling for a rule that shows up at HEAD but isn't in the budget at all -- a @@ -72,20 +63,6 @@ def _to_repo_relative(raw: str) -> str | None: return None -def count_mypy(lines: Iterable[str]) -> dict[str, int]: - """Count in-repo mypy errors per rule code from text output. Errors for - files outside the repo (third-party stubs) are ignored, as before.""" - counts: Counter[str] = Counter() - for raw in lines: - line = raw.rstrip("\n") - match = MYPY_ERROR.match(line) - if match is None or _to_repo_relative(match.group("file")) is None: - continue - code = MYPY_CODE.search(line) - counts[code.group("code") if code else UNCODED] += 1 - return dict(counts) - - def count_basedpyright(payload: str) -> dict[str, int]: """Count in-repo basedpyright errors per rule from `--outputjson`. Warnings and information are ignored; only `severity == "error"` is gated.""" @@ -108,12 +85,6 @@ def count_basedpyright(payload: str) -> dict[str, int]: return dict(counts) -def count_errors(stdin_text: str, tool: str) -> dict[str, int]: - if tool == "basedpyright": - return count_basedpyright(stdin_text) - return count_mypy(stdin_text.splitlines()) - - def evaluate( counts: Mapping[str, int], budget: Mapping[str, Mapping[str, int]] ) -> list[Breach]: @@ -136,13 +107,11 @@ def is_vacuous_run( return not counts and any(spec["baseline"] for spec in budget.values()) -def budget_path(tool: str) -> Path: - return REPO_ROOT / f"{tool}-code-budget.json" +BUDGET_PATH = REPO_ROOT / "basedpyright-code-budget.json" -def cmd_update(tool: str, counts: Mapping[str, int]) -> None: - path = budget_path(tool) - existing = json.loads(path.read_text()) if path.exists() else {} +def cmd_update(counts: Mapping[str, int]) -> None: + existing = json.loads(BUDGET_PATH.read_text()) if BUDGET_PATH.exists() else {} budget = { code: { "baseline": count, @@ -152,18 +121,18 @@ def cmd_update(tool: str, counts: Mapping[str, int]) -> None: } for code, count in sorted(counts.items()) } - path.write_text(json.dumps(budget, indent=2, sort_keys=True) + "\n") + BUDGET_PATH.write_text(json.dumps(budget, indent=2, sort_keys=True) + "\n") print( - f"Re-captured {tool} per-rule budget: {len(budget)} rules, {sum(counts.values())} errors total" + f"Re-captured basedpyright per-rule budget: {len(budget)} rules, {sum(counts.values())} errors total" ) -def cmd_check(tool: str, counts: Mapping[str, int]) -> None: - budget = json.loads(budget_path(tool).read_text()) +def cmd_check(counts: Mapping[str, int]) -> None: + budget = json.loads(BUDGET_PATH.read_text()) if is_vacuous_run(counts, budget): expected = sum(spec["baseline"] for spec in budget.values()) print( - f"FAIL: {tool} produced no errors, but {budget_path(tool).name} expects " + f"FAIL: basedpyright produced no errors, but {BUDGET_PATH.name} expects " f"~{expected}. The type checker almost certainly crashed or emitted " f"nothing; refusing to certify a vacuous run." ) @@ -171,25 +140,24 @@ def cmd_check(tool: str, counts: Mapping[str, int]) -> None: breaches = evaluate(counts, budget) if not breaches: print( - f"OK: every rule is within its {tool} ceiling ({sum(counts.values())} errors total)" + f"OK: every rule is within its basedpyright ceiling ({sum(counts.values())} errors total)" ) return - print(f"FAIL: {tool} errors exceed the per-rule ceiling:") + print("FAIL: basedpyright errors exceed the per-rule ceiling:") for breach in breaches: print(f" {breach.code}: {breach.total} errors over cap {breach.cap}") print( - f"Resolve the new errors, or run 'make lint-{tool}-budget-update' if the ceiling should move." + "Resolve the new errors, or run 'make lint-basedpyright-budget-update' if the ceiling should move." ) raise SystemExit(1) def main() -> None: parser = argparse.ArgumentParser(description=__doc__) - parser.add_argument("--tool", choices=("mypy", "basedpyright"), required=True) parser.add_argument("--update", action="store_true") args = parser.parse_args() - counts = count_errors(sys.stdin.read(), args.tool) - cmd_update(args.tool, counts) if args.update else cmd_check(args.tool, counts) + counts = count_basedpyright(sys.stdin.read()) + cmd_update(counts) if args.update else cmd_check(counts) if __name__ == "__main__": diff --git a/scripts/type_discipline_gate.py b/scripts/type_discipline_gate.py new file mode 100644 index 00000000000..c111486e56a --- /dev/null +++ b/scripts/type_discipline_gate.py @@ -0,0 +1,198 @@ +#!/usr/bin/env python3 +"""Total-count gate for the LIT* rules in scripts/check_type_discipline.py. + +Sibling of scripts/ruff_strict_gate.py. Each rule listed in +type-discipline-budget.json has a hard ceiling (baseline + slack). The gate counts +each rule across the whole `litellm` tree and fails when a rule is both over its +ceiling and higher than the base it merges into, so a change is blamed for the +violations it adds, never for drift that already exists in the base. + +Rules not present in the budget are ignored, but today every rule the checker +emits is gated: LIT001 (mutable collection in any annotation), LIT002 +(mutable-collection construction), LIT003/LIT004 (noqa / ignore without codes or +reason), LIT006 (cast), and LIT008 (`**kwargs`) carry slack-buffered ceilings to +ratchet down; LIT005 (`*-ok` suppression without a reason) is frozen at slack 0 +so any net-new reasonless suppression trips the gate; and LIT007 (TypeGuard/TypeIs) +is a hard zero. Re-baseline with `--update` to ratchet a ceiling down. +""" + +import argparse +import json +import re +import shutil +import subprocess +import sys +import tempfile +from collections import Counter +from pathlib import Path +from typing import NamedTuple + +REPO_ROOT = Path(__file__).resolve().parent.parent +CHECKER = REPO_ROOT / "scripts" / "check_type_discipline.py" +BUDGET_PATH = REPO_ROOT / "type-discipline-budget.json" +TARGET = "litellm" +DEFAULT_BASE = "origin/litellm_internal_staging" + +_HUNK = re.compile(r"^@@ -\d+(?:,\d+)? \+(\d+)(?:,(\d+))? @@") +_LINE = re.compile(r"^(?P.+?):(?P\d+): (?PLIT\d+) ") + + +class Violation(NamedTuple): + file: str + line: int + code: str + + +class Breach(NamedTuple): + rule: str + total: int + cap: int + added: int + + +def _run(cmd: list, cwd: Path = REPO_ROOT) -> str: + proc = subprocess.run(cmd, cwd=cwd, capture_output=True, text=True) + if proc.returncode not in (0, 1): + sys.stderr.write(proc.stderr) + raise SystemExit(f"{cmd[0]} exited {proc.returncode}") + return proc.stdout + + +def _check(root: Path, checker: Path) -> list: + # Resolve root first: on macOS tempfile dirs (/var/...) resolve to /private/var/..., + # and the checker prints already-resolved absolute paths, so relative_to would fail. + root = root.resolve() + out = _run([sys.executable, str(checker), str(root / TARGET)], cwd=root) + found = [] + for line in out.splitlines(): + m = _LINE.match(line) + if m is None: + continue + name = Path(m.group("file")) + full = name if name.is_absolute() else root / name + rel = full.resolve().relative_to(root).as_posix() + found.append(Violation(rel, int(m.group("line")), m.group("code"))) + return found + + +def head_violations() -> list: + return _check(REPO_ROOT, CHECKER) + + +def count_by_rule(violations: list) -> dict: + return dict(Counter(v.code for v in violations)) + + +def base_counts(ref: str) -> dict: + parent = Path(tempfile.mkdtemp(prefix="lit_base_")) + worktree = parent / "wt" + try: + _run(["git", "worktree", "add", "--detach", str(worktree), ref]) + # Measure the base with the *current* rule logic, not whatever shipped at base. + (worktree / "scripts").mkdir(parents=True, exist_ok=True) + checker = worktree / "scripts" / "check_type_discipline.py" + shutil.copy(CHECKER, checker) + return count_by_rule(_check(worktree, checker)) + finally: + # Best-effort teardown: cleanup must never raise, or it masks the real error when + # the body (or the `worktree add` itself) failed. rmtree is already best-effort. + subprocess.run( + ["git", "worktree", "remove", "--force", str(worktree)], + cwd=REPO_ROOT, capture_output=True, text=True, + ) + shutil.rmtree(parent, ignore_errors=True) + + +def over_ceiling(head: dict, budget: dict) -> frozenset: + """Rules whose head count already exceeds baseline + slack. + + A rule can only breach when it is over its ceiling, so when none are the base + comparison cannot change the verdict and the base worktree scan can be skipped. + """ + return frozenset( + rule for rule, spec in budget.items() + if head.get(rule, 0) > spec["baseline"] + spec["slack"] + ) + + +def evaluate(head: dict, base: dict, budget: dict) -> list: + breaches = [] + for rule, spec in budget.items(): + cap = spec["baseline"] + spec["slack"] + total = head.get(rule, 0) + if total > cap and total > base.get(rule, 0): + breaches.append(Breach(rule, total, cap, total - base.get(rule, 0))) + return sorted(breaches) + + +def parse_changed_lines(diff_text: str) -> dict: + changed: dict = {} + path = None + for line in diff_text.splitlines(): + if line.startswith("+++ b/"): + path = line[6:] + elif path and (match := _HUNK.match(line)): + start = int(match.group(1)) + count = int(match.group(2)) if match.group(2) is not None else 1 + changed.setdefault(path, set()).update(range(start, start + count)) + return changed + + +def introduced(violations: list, changed: dict) -> list: + return [v for v in violations if v.line in changed.get(v.file, set())] + + +def cmd_check(base: str) -> None: + budget = json.loads(BUDGET_PATH.read_text()) + head = head_violations() + head_counts = count_by_rule(head) + if not over_ceiling(head_counts, budget): + print(f"OK: every LIT rule is within its codebase ceiling (base {base})") + return + base_point = _run(["git", "merge-base", base, "HEAD"]).strip() or base + breaches = evaluate(head_counts, base_counts(base_point), budget) + if not breaches: + print(f"OK: every LIT rule is within its codebase ceiling (base {base})") + return + new = introduced( + head, + parse_changed_lines( + _run(["git", "diff", base_point, "--unified=0", "--no-color", "--", TARGET]) + ), + ) + print(f"FAIL: LIT-rule totals exceed their ceiling (base {base}):") + for breach in breaches: + print( + f" {breach.rule}: total {breach.total} over cap {breach.cap} (this change added {breach.added})" + ) + for violation in sorted(v for v in new if v.code == breach.rule): + print(f" {violation.file}:{violation.line}") + print( + "Remove the new violations, give each a reason (`# noqa: XXX # `, " + "`# pyright: ignore[rule] # `, `# mutable-ok: `, " + "`# cast-ok: `, `# guard-ok: `, `# kwargs-ok: `), or " + "remove an equal number elsewhere; the ceiling is baseline + slack in " + "type-discipline-budget.json." + ) + raise SystemExit(1) + + +def cmd_update() -> None: + budget = json.loads(BUDGET_PATH.read_text()) + head = count_by_rule(head_violations()) + for rule in budget: + budget[rule]["baseline"] = head.get(rule, 0) + BUDGET_PATH.write_text(json.dumps(budget, indent=2, sort_keys=True) + "\n") + print("Re-captured per-rule baselines from the current tree") + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--base", default=DEFAULT_BASE) + parser.add_argument("--update", action="store_true") + args = parser.parse_args() + cmd_update() if args.update else cmd_check(args.base) + + +if __name__ == "__main__": + main() diff --git a/tests/_fake_openai_endpoint_server.py b/tests/_fake_openai_endpoint_server.py new file mode 100644 index 00000000000..409f569070b --- /dev/null +++ b/tests/_fake_openai_endpoint_server.py @@ -0,0 +1,239 @@ +"""Canned OpenAI-shaped mock server for the CI proxy E2Es. + +Several CI jobs run the litellm proxy (often in its own Docker container) against +a model whose ``api_base`` is a fake OpenAI endpoint that returns canned +responses, so the run costs nothing and does not depend on a real provider. That +endpoint used to be a single shared deployment; when it went down every one of +those jobs failed with ``404 Application not found`` even though nothing in the +PR was broken. + +This process is the local stand-in. A model points its ``api_base`` here and +gets back a well-formed chat/text/embedding response with realistic ``usage`` so +cost tracking and spend accounting still exercise their real code paths. The one +behavioral special case mirrors the old hosted mock: a request whose ``model`` +is ``429`` returns HTTP 429 so rate-limit and cooldown tests still have +something to trip on. +""" + +from __future__ import annotations + +import json +import time +import uuid +from typing import AsyncIterator, Final + +import uvicorn +from starlette.applications import Starlette +from starlette.requests import Request +from starlette.responses import JSONResponse, PlainTextResponse, Response, StreamingResponse +from starlette.routing import Route + +_CANNED_CONTENT: Final = "Hello! This is a mock response from the fake OpenAI endpoint." +_RATE_LIMIT_MODEL: Final = "429" +_PROMPT_TOKENS: Final = 20 +_COMPLETION_TOKENS: Final = 20 + + +def _usage() -> dict[str, int]: + return { + "prompt_tokens": _PROMPT_TOKENS, + "completion_tokens": _COMPLETION_TOKENS, + "total_tokens": _PROMPT_TOKENS + _COMPLETION_TOKENS, + } + + +def _requested_model(body: dict[str, object]) -> str: + model = body.get("model") + return model if isinstance(model, str) else "mock-model" + + +def _wants_stream(body: dict[str, object]) -> bool: + return body.get("stream") is True + + +def _wants_stream_usage(body: dict[str, object]) -> bool: + options = body.get("stream_options") + return isinstance(options, dict) and options.get("include_usage") is True + + +async def _parse_body(request: Request) -> dict[str, object]: + raw = await request.body() + if not raw: + return {} + try: + parsed = json.loads(raw) + except ValueError: + return {} + return parsed if isinstance(parsed, dict) else {} + + +def _rate_limit_response(model: str) -> JSONResponse: + return JSONResponse( + status_code=429, + content={ + "error": { + "message": f"Rate limit reached for model `{model}` (mock).", + "type": "rate_limit_error", + "code": "429", + } + }, + ) + + +def _chat_completion_body(model: str) -> dict[str, object]: + return { + "id": f"chatcmpl-{uuid.uuid4().hex[:24]}", + "object": "chat.completion", + "created": int(time.time()), + "model": model, + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": _CANNED_CONTENT}, + "finish_reason": "stop", + } + ], + "usage": _usage(), + } + + +async def _chat_completion_stream(model: str, with_usage: bool) -> AsyncIterator[str]: + response_id = f"chatcmpl-{uuid.uuid4().hex[:24]}" + created = int(time.time()) + + def chunk(delta: dict[str, object], finish_reason: str | None) -> dict[str, object]: + return { + "id": response_id, + "object": "chat.completion.chunk", + "created": created, + "model": model, + "choices": [{"index": 0, "delta": delta, "finish_reason": finish_reason}], + } + + yield f"data: {json.dumps(chunk({'role': 'assistant', 'content': _CANNED_CONTENT}, None))}\n\n" + yield f"data: {json.dumps(chunk({}, 'stop'))}\n\n" + if with_usage: + final = chunk({}, None) | {"choices": [], "usage": _usage()} + yield f"data: {json.dumps(final)}\n\n" + yield "data: [DONE]\n\n" + + +async def chat_completions(request: Request) -> Response: + body = await _parse_body(request) + model = _requested_model(body) + if model == _RATE_LIMIT_MODEL: + return _rate_limit_response(model) + if _wants_stream(body): + return StreamingResponse( + _chat_completion_stream(model, _wants_stream_usage(body)), + media_type="text/event-stream", + ) + return JSONResponse(_chat_completion_body(model)) + + +def _text_completion_body(model: str) -> dict[str, object]: + return { + "id": f"cmpl-{uuid.uuid4().hex[:24]}", + "object": "text_completion", + "created": int(time.time()), + "model": model, + "choices": [ + { + "text": _CANNED_CONTENT, + "index": 0, + "logprobs": None, + "finish_reason": "stop", + } + ], + "usage": _usage(), + } + + +async def _text_completion_stream(model: str, with_usage: bool) -> AsyncIterator[str]: + response_id = f"cmpl-{uuid.uuid4().hex[:24]}" + created = int(time.time()) + + def chunk(text: str, finish_reason: str | None) -> dict[str, object]: + return { + "id": response_id, + "object": "text_completion", + "created": created, + "model": model, + "choices": [{"text": text, "index": 0, "logprobs": None, "finish_reason": finish_reason}], + } + + yield f"data: {json.dumps(chunk(_CANNED_CONTENT, None))}\n\n" + yield f"data: {json.dumps(chunk('', 'stop'))}\n\n" + if with_usage: + final = chunk("", None) | {"choices": [], "usage": _usage()} + yield f"data: {json.dumps(final)}\n\n" + yield "data: [DONE]\n\n" + + +async def completions(request: Request) -> Response: + body = await _parse_body(request) + model = _requested_model(body) + if model == _RATE_LIMIT_MODEL: + return _rate_limit_response(model) + if _wants_stream(body): + return StreamingResponse( + _text_completion_stream(model, _wants_stream_usage(body)), + media_type="text/event-stream", + ) + return JSONResponse(_text_completion_body(model)) + + +async def embeddings(request: Request) -> Response: + body = await _parse_body(request) + raw_input = body.get("input", "") + count = len(raw_input) if isinstance(raw_input, list) else 1 + return JSONResponse( + { + "object": "list", + "data": [{"object": "embedding", "index": i, "embedding": [0.0] * 1536} for i in range(max(count, 1))], + "model": _requested_model(body), + "usage": {"prompt_tokens": 5, "total_tokens": 5}, + } + ) + + +async def list_models(_request: Request) -> Response: + return JSONResponse( + { + "object": "list", + "data": [ + {"id": "fake", "object": "model", "owned_by": "mock"}, + {"id": "my-fake-model", "object": "model", "owned_by": "mock"}, + ], + } + ) + + +async def health(_request: Request) -> Response: + return PlainTextResponse("ok") + + +app = Starlette( + routes=[ + Route("/health", health, methods=["GET"]), + Route("/", health, methods=["GET"]), + Route("/chat/completions", chat_completions, methods=["POST"]), + Route("/v1/chat/completions", chat_completions, methods=["POST"]), + Route("/completions", completions, methods=["POST"]), + Route("/v1/completions", completions, methods=["POST"]), + Route("/embeddings", embeddings, methods=["POST"]), + Route("/v1/embeddings", embeddings, methods=["POST"]), + Route("/models", list_models, methods=["GET"]), + Route("/v1/models", list_models, methods=["GET"]), + ] +) + + +if __name__ == "__main__": + import argparse + + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--host", default="0.0.0.0") + parser.add_argument("--port", type=int, default=8190) + args = parser.parse_args() + uvicorn.run(app, host=args.host, port=args.port) diff --git a/tests/llm_translation/base_llm_unit_tests.py b/tests/llm_translation/base_llm_unit_tests.py index fef1d23d867..a184798b503 100644 --- a/tests/llm_translation/base_llm_unit_tests.py +++ b/tests/llm_translation/base_llm_unit_tests.py @@ -906,7 +906,10 @@ class BaseLLMChatTest(ABC): { "type": "image_url", "image_url": { - "url": "https://www.gstatic.com/webp/gallery/1.webp", + # sha-pinned in-repo logo via jsdelivr; gstatic's + # robots.txt blocks server-side fetchers (e.g. + # Anthropic), which 400s the request. + "url": "https://cdn.jsdelivr.net/gh/BerriAI/litellm@d769e81c90d453240c61fc572cdb27fae06a89d0/ui/litellm-dashboard/public/assets/logos/litellm_logo.jpg", "detail": detail, }, }, diff --git a/tests/local_testing/test_aim_guardrails.py b/tests/local_testing/test_aim_guardrails.py index 31416c565c1..2cb7f9cd357 100644 --- a/tests/local_testing/test_aim_guardrails.py +++ b/tests/local_testing/test_aim_guardrails.py @@ -6,10 +6,10 @@ import sys from unittest.mock import AsyncMock, patch, call import pytest -from fastapi.exceptions import HTTPException from httpx import Request, Response from litellm import DualCache +from litellm.proxy._types import ProxyException from litellm.proxy.guardrails.guardrail_hooks.aim.aim import ( AimGuardrail, AimGuardrailMissingSecrets, @@ -101,7 +101,7 @@ async def test_block_callback(mode: str): ], } - with pytest.raises(HTTPException, match="Jailbreak detected"): + with pytest.raises(ProxyException, match="Jailbreak detected") as exc_info: with patch( "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", return_value=Response( @@ -135,6 +135,137 @@ async def test_block_callback(mode: str): call_type="completion", ) + exc = exc_info.value + assert exc.code == "400" + assert exc.type == "invalid_request_error" + assert exc.param is None + assert exc.openai_code == "content_policy_violation" + + +@pytest.mark.asyncio +async def test_output_block_raises_proxy_exception(): + """An output-side block is a content-policy violation, like the input block: + it must surface a conformant ProxyException, not a bare HTTPException whose + type/param serialize as the literal string "None". Regression for LIT-3751.""" + init_guardrails_v2( + all_guardrails=[ + { + "guardrail_name": "gibberish-guard", + "litellm_params": { + "guardrail": "aim", + "mode": "post_call", + "api_key": "hs-aim-key", + }, + }, + ], + config_file_path="", + ) + aim_guardrails = [ + callback for callback in litellm.callbacks if isinstance(callback, AimGuardrail) + ] + assert len(aim_guardrails) == 1 + aim_guardrail = aim_guardrails[0] + + block_on_output = Response( + json={ + "analysis_result": {"policy_drill_down": {"PII": {}}}, + "required_action": { + "action_type": "block_action", + "detection_message": "Output blocked: leaked secret", + "policy_name": "blocking policy", + }, + }, + status_code=200, + request=Request(method="POST", url="http://aim"), + ) + response = ModelResponse( + choices=[ + { + "finish_reason": "stop", + "index": 0, + "message": {"content": "here is the secret", "role": "assistant"}, + } + ] + ) + + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + return_value=block_on_output, + ): + with pytest.raises(ProxyException, match="Output blocked") as exc_info: + await aim_guardrail.async_post_call_success_hook( + data={"messages": [{"role": "user", "content": "tell me a secret"}]}, + response=response, + user_api_key_dict=UserAPIKeyAuth(), + ) + + exc = exc_info.value + assert exc.code == "400" + assert exc.type == "invalid_request_error" + assert exc.param is None + assert exc.openai_code == "content_policy_violation" + + +@pytest.mark.asyncio +async def test_anonymize_multimodal_rejection_raises_proxy_exception(): + """Anonymize on multimodal input degrades to a 400 because mask-in-place would + drop non-text parts. That is a usage error, not a content-policy violation, so + it must raise a conformant ProxyException WITHOUT the content_policy_violation + code. Regression for LIT-3751.""" + init_guardrails_v2( + all_guardrails=[ + { + "guardrail_name": "gibberish-guard", + "litellm_params": { + "guardrail": "aim", + "mode": "pre_call", + "api_key": "hs-aim-key", + }, + }, + ], + config_file_path="", + ) + aim_guardrails = [ + callback for callback in litellm.callbacks if isinstance(callback, AimGuardrail) + ] + assert len(aim_guardrails) == 1 + aim_guardrail = aim_guardrails[0] + + data = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Hi my name is Brian"}, + { + "type": "image_url", + "image_url": {"url": "data:image/png;base64,iVBORw0KGgo="}, + }, + ], + }, + ], + } + + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + return_value=response_with_detections, + ): + with pytest.raises( + ProxyException, match="anonymize action requested for multimodal" + ) as exc_info: + await aim_guardrail.async_pre_call_hook( + data=data, + cache=DualCache(), + user_api_key_dict=UserAPIKeyAuth(), + call_type="completion", + ) + + exc = exc_info.value + assert exc.code == "400" + assert exc.type == "invalid_request_error" + assert exc.param is None + assert exc.openai_code != "content_policy_violation" + @pytest.mark.asyncio @pytest.mark.parametrize("mode", ["pre_call", "during_call"]) diff --git a/tests/local_testing/test_router.py b/tests/local_testing/test_router.py index 6d04e6ecaa5..6f1e367760f 100644 --- a/tests/local_testing/test_router.py +++ b/tests/local_testing/test_router.py @@ -1639,7 +1639,10 @@ async def test_router_text_completion_client(): "litellm_params": { "model": "text-completion-openai/gpt-3.5-turbo-instruct", "api_key": os.getenv("OPENAI_API_KEY", None), - "api_base": "https://exampleopenaiendpoint-production.up.railway.app/", + "api_base": os.getenv( + "FAKE_OPENAI_API_BASE", + "https://exampleopenaiendpoint-production.up.railway.app/", + ), }, } ] diff --git a/tests/pass_through_tests/ruby_passthrough_tests/spec/openai_assistants_passthrough_spec.rb b/tests/pass_through_tests/ruby_passthrough_tests/spec/openai_assistants_passthrough_spec.rb index 1cfaeb5e209..5a4dc0395f8 100644 --- a/tests/pass_through_tests/ruby_passthrough_tests/spec/openai_assistants_passthrough_spec.rb +++ b/tests/pass_through_tests/ruby_passthrough_tests/spec/openai_assistants_passthrough_spec.rb @@ -5,7 +5,8 @@ RSpec.describe 'OpenAI Assistants Passthrough' do let(:client) do OpenAI::Client.new( access_token: "sk-1234", - uri_base: "http://0.0.0.0:4000/openai" + uri_base: "http://0.0.0.0:4000/openai", + request_timeout: 600 ) end diff --git a/tests/pass_through_tests/test_vertex_ai.py b/tests/pass_through_tests/test_vertex_ai.py index 0ac66b470c6..e8223f2219c 100644 --- a/tests/pass_through_tests/test_vertex_ai.py +++ b/tests/pass_through_tests/test_vertex_ai.py @@ -126,15 +126,21 @@ async def test_basic_vertex_ai_pass_through_with_spendlog(): print("response", response) - # Poll for spend update instead of fixed sleep - spend logging is async/batched - max_wait = 120 # total seconds to wait + # Spend logging is async/batched and can lag under CI load, so poll instead of + # sleeping a fixed amount. A transient empty read is skipped, not counted as 0.0 + # spend, which would spuriously fail the assertion on an otherwise-billed call. + max_wait = 240 # total seconds to wait poll_interval = 10 # seconds between checks elapsed = 0 spend_after = spend_before while elapsed < max_wait: await asyncio.sleep(poll_interval) elapsed += poll_interval - spend_after = await call_spend_logs_endpoint() or 0.0 + latest_spend = await call_spend_logs_endpoint() + if latest_spend is None: + print(f"spend logs unavailable (elapsed={elapsed}s), retrying") + continue + spend_after = latest_spend print(f"spend_after (elapsed={elapsed}s)", spend_after) if spend_after > spend_before: break diff --git a/tests/test_keys.py b/tests/test_keys.py index 89977d43676..003e2711055 100644 --- a/tests/test_keys.py +++ b/tests/test_keys.py @@ -621,17 +621,23 @@ async def test_key_info_spend_values_image_generation(): assert spend > 0 # The record/replay proxy serves this identical second call from its - # cassette (free), but the proxy must still bill it. If the proxy's own - # response cache were on, the repeat would be a $0 cache hit and spend - # would not move, silently zeroing recorded-call spend; assert it grows. + # cassette (free), but the proxy must still bill it. Spend logging is + # async/batched, so poll for the increase rather than reading once after a + # fixed sleep; a spend that never grows means the repeat was not billed + # (e.g. the proxy response cache is on), which this still catches. await image_generation(session=session, key=key) - await asyncio.sleep(5) - key_info = await retry_request( - get_key_info, session=session, get_key=key, call_key=key - ) - assert key_info["info"]["spend"] > spend, ( - "spend did not increase on an identical repeat image call; the proxy " - "response cache appears to be ON, which would zero recorded-call spend" + spend_after = spend + for _ in range(12): + await asyncio.sleep(5) + key_info = await retry_request( + get_key_info, session=session, get_key=key, call_key=key + ) + spend_after = key_info["info"]["spend"] + if spend_after > spend: + break + assert spend_after > spend, ( + "spend did not increase on an identical repeat image call; the repeat " + "was not billed (the proxy response cache may be on)" ) diff --git a/tests/test_litellm/caching/test_check_and_fix_namespace_none_guard.py b/tests/test_litellm/caching/test_check_and_fix_namespace_none_guard.py new file mode 100644 index 00000000000..c049c3157f4 --- /dev/null +++ b/tests/test_litellm/caching/test_check_and_fix_namespace_none_guard.py @@ -0,0 +1,49 @@ +""" +Test that check_and_fix_namespace handles None key gracefully. + +Regression test for https://github.com/BerriAI/litellm/issues/30424 +""" +from unittest.mock import MagicMock + +from litellm.caching.redis_cache import RedisCache + + +def test_check_and_fix_namespace_with_none_key(): + """When key is None, check_and_fix_namespace should return None without raising.""" + cache = MagicMock(spec=RedisCache) + cache.namespace = "litellm" + # Call the real method + result = RedisCache.check_and_fix_namespace(cache, key=None) + assert result is None + + +def test_check_and_fix_namespace_with_none_key_no_namespace(): + """When key is None and namespace is None, should return None without raising.""" + cache = MagicMock(spec=RedisCache) + cache.namespace = None + result = RedisCache.check_and_fix_namespace(cache, key=None) + assert result is None + + +def test_check_and_fix_namespace_with_valid_key(): + """Normal behavior: prefix key with namespace if not already prefixed.""" + cache = MagicMock(spec=RedisCache) + cache.namespace = "litellm" + result = RedisCache.check_and_fix_namespace(cache, key="my_key") + assert result == "litellm:my_key" + + +def test_check_and_fix_namespace_with_already_prefixed_key(): + """If key already starts with namespace, don't double-prefix.""" + cache = MagicMock(spec=RedisCache) + cache.namespace = "litellm" + result = RedisCache.check_and_fix_namespace(cache, key="litellm:my_key") + assert result == "litellm:my_key" + + +def test_check_and_fix_namespace_no_namespace(): + """When namespace is None, return key as-is.""" + cache = MagicMock(spec=RedisCache) + cache.namespace = None + result = RedisCache.check_and_fix_namespace(cache, key="my_key") + assert result == "my_key" diff --git a/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py b/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py index 20824ca09e6..3447f5bdb7e 100644 --- a/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py +++ b/tests/test_litellm/integrations/otel/test_otel_v2_sources_of_truth.py @@ -528,6 +528,35 @@ def test_capture_span_content_resolves_modes(): ).capture_span_content is False ) + # V1 accepted UPPER_SNAKE_CASE; the env value is case-insensitive so an + # operator carrying ``SPAN_AND_EVENT`` forward still enables capture. + assert ( + OpenTelemetryV2Config( + capture_message_content="SPAN_AND_EVENT" + ).capture_span_content + is True + ) + assert ( + OpenTelemetryV2Config(capture_message_content="SPAN_ONLY").capture_span_content + is True + ) + assert ( + OpenTelemetryV2Config(capture_message_content="NO_CONTENT").capture_span_content + is False + ) + + +def test_capture_message_content_normalizer_only_touches_strings(): + """The casing normalizer lower-cases strings and leaves anything else + untouched, so a non-string value still fails the field's ``str`` validation + instead of being silently coerced into a bogus capture mode.""" + import pytest + from pydantic import ValidationError + + from litellm.integrations.otel.model.config import OpenTelemetryV2Config + + with pytest.raises(ValidationError): + OpenTelemetryV2Config(capture_message_content=123) def test_v2_flag_is_off_by_default(monkeypatch): diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index fe49b930c10..9b3152fae07 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -35,6 +35,7 @@ sys.path.insert( from litellm.litellm_core_utils.llm_cost_calc.utils import ( PromptTokensDetailsResult, _calculate_input_cost, + _get_token_base_cost, calculate_cache_writing_cost, generic_cost_per_token, ) @@ -298,6 +299,26 @@ def test_generic_cost_per_token_above_200k_tokens(): ) +def test_get_token_base_cost_picks_highest_crossed_tier(): + """Regression test for #30345. + + With graduated tiers at 90k and 128k whose keys have different digit lengths, a request + crossing both must be billed at the highest tier it crosses (128k), not the lower one that + happens to sort first lexicographically. + """ + model_info = { + "input_cost_per_token": 1e-6, + "output_cost_per_token": 2e-6, + "input_cost_per_token_above_90k_tokens": 5e-6, + "input_cost_per_token_above_128k_tokens": 9e-6, + } + usage = Usage(prompt_tokens=150_000, completion_tokens=10, total_tokens=150_010) + + prompt_base_cost = _get_token_base_cost(model_info, usage)[0] + + assert prompt_base_cost == 9e-6 + + def test_generic_cost_per_token_gpt54_above_272k_tokens(): """GPT-5.4/5.4-pro: prompts >272K input tokens priced at 2x input, 1.5x output.""" model = "gpt-5.4" @@ -1573,3 +1594,72 @@ def test_data_residency_composes_with_service_tier(_local_model_cost_map): assert priority_base_total > 0 assert priority_eu_total == pytest.approx(priority_base_total * 1.10, rel=1e-9) + + +def test_priority_service_tier_above_threshold_uses_priority_tier_rates_for_cached_tokens( + _local_model_cost_map, +): + """Regression: for a model that publishes both service_tier and above_threshold rate + variants, a priority request over the threshold must bill cached tokens at + cache_read_input_token_cost_above_200k_tokens_priority (and analogously for + input/output above-threshold), not the standard above-threshold rate.""" + usage = Usage( + prompt_tokens=250_000, + completion_tokens=1_000, + total_tokens=251_000, + prompt_tokens_details=PromptTokensDetailsWrapper( + cached_tokens=200_000, text_tokens=50_000 + ), + completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=1_000), + ) + + prompt_cost, completion_cost = generic_cost_per_token( + model="gemini-3-pro-preview", + usage=usage, + custom_llm_provider="gemini", + service_tier="priority", + ) + + # gemini-3-pro-preview priority + above_200k rates from the pricing JSON: + # input 7.2e-6, output 3.24e-5, cache_read 7.2e-7 + expected_prompt = 50_000 * 7.2e-6 + 200_000 * 7.2e-7 + expected_completion = 1_000 * 3.24e-5 + assert prompt_cost == pytest.approx(expected_prompt, rel=1e-9) + assert completion_cost == pytest.approx(expected_completion, rel=1e-9) + + +def test_priority_service_tier_above_threshold_falls_back_to_standard_for_cache_creation( + _local_model_cost_map, +): + """Regression: priority requests against models that publish standard above-threshold + cache_creation rates but no priority variant must fall back to the standard + above-threshold rate, not the priority-base rate. vertex_ai/claude-sonnet-4-5 + has cache_creation_input_token_cost_above_200k_tokens but no _priority sibling.""" + usage = Usage( + prompt_tokens=350_000, + completion_tokens=1_000, + total_tokens=351_000, + prompt_tokens_details=PromptTokensDetailsWrapper( + cached_tokens=200_000, + cache_creation_tokens=100_000, + text_tokens=50_000, + ), + completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=1_000), + ) + + prompt_cost, completion_cost = generic_cost_per_token( + model="vertex_ai/claude-sonnet-4-5", + usage=usage, + custom_llm_provider="vertex_ai", + service_tier="priority", + ) + + # vertex_ai/claude-sonnet-4-5 above_200k (no _priority variants): + # input 6e-6, output 2.25e-5, cache_read 6e-7, cache_creation 7.5e-6 + # text 50_000 * 6e-6 = 0.30 + # cache_read 200_000 * 6e-7 = 0.12 + # cache_creation 100_000 * 7.5e-6 = 0.75 + expected_prompt = 50_000 * 6e-6 + 200_000 * 6e-7 + 100_000 * 7.5e-6 + expected_completion = 1_000 * 2.25e-5 + assert prompt_cost == pytest.approx(expected_prompt, rel=1e-9) + assert completion_cost == pytest.approx(expected_completion, rel=1e-9) diff --git a/tests/test_litellm/litellm_core_utils/test_exception_mapping_utils.py b/tests/test_litellm/litellm_core_utils/test_exception_mapping_utils.py index 7dab0e02623..35c02184a51 100644 --- a/tests/test_litellm/litellm_core_utils/test_exception_mapping_utils.py +++ b/tests/test_litellm/litellm_core_utils/test_exception_mapping_utils.py @@ -11,6 +11,7 @@ sys.path.insert( from litellm.litellm_core_utils.exception_mapping_utils import ( ExceptionCheckers, + _get_body_error_code, exception_type, extract_and_raise_litellm_exception, ) @@ -359,6 +360,122 @@ def test_vertex_ai_rate_limit_error_mapping(error_message, should_raise_rate_lim ) +class TestGetBodyErrorCode: + """Unit tests for _get_body_error_code helper.""" + + def test_parses_int_code(self): + body = ( + '{"error":{"message":"high demand","type":"upstream_error",' + '"param":"","code":429}}' + ) + assert _get_body_error_code(body) == 429 + + def test_parses_string_code(self): + # some gateways serialize code as a string + body = '{"error":{"message":"x","code":"503"}}' + assert _get_body_error_code(body) == 503 + + def test_returns_none_on_non_json(self): + assert _get_body_error_code("not json") is None + + def test_returns_none_when_no_error_key(self): + assert _get_body_error_code('{"ok":true}') is None + + def test_returns_none_when_no_code_key(self): + assert _get_body_error_code('{"error":{"message":"x"}}') is None + + +# Test cases for Gemini upstream-error body-code mapping. +# +# Body code 429 wrapped in a 5xx HTTP envelope (e.g. new-api gateways) +# must map to RateLimitError so Router retries kick in. A 4xx HTTP +# envelope with body code:429 must NOT — it falls through to whatever +# the HTTP status code maps to (BadRequestError, AuthenticationError, +# etc.), matching upstream's existing semantics. +gemini_body_code_429_test_cases = [ + # (status_code, error_body, expected_exception_type, description) + ( + 500, + '{"error":{"message":" This model is currently experiencing high demand.' + " Spikes in demand are usually temporary. Please try again later." + ' (request id: x)","type":"upstream_error","param":"","code":429}}', + litellm.RateLimitError, + "HTTP 500 envelope with body code:429 -> RateLimitError", + ), + ( + 503, + '{"error":{"message":"upstream unavailable","type":"upstream_error",' + '"param":"","code":429}}', + litellm.RateLimitError, + "HTTP 503 envelope with body code:429 -> RateLimitError", + ), + ( + 502, + '{"error":{"message":"bad gateway","code":429}}', + litellm.RateLimitError, + "HTTP 502 envelope with body code:429 -> RateLimitError", + ), + ( + 500, + '{"error":{"message":"server boom","code":500}}', + litellm.InternalServerError, + "HTTP 500 with body code:500 stays InternalServerError", + ), + ( + 500, + "plain text 500 error", + litellm.InternalServerError, + "HTTP 500 with non-JSON body falls through to status_code mapping", + ), + ( + 400, + '{"error":{"message":"malformed","code":429}}', + litellm.BadRequestError, + "HTTP 400 with body code:429 must NOT be promoted to RateLimitError", + ), + ( + 401, + '{"error":{"message":"bad key","code":429}}', + litellm.AuthenticationError, + "HTTP 401 with body code:429 must NOT be promoted to RateLimitError", + ), +] + + +@pytest.mark.parametrize( + "status_code, error_body, expected_exception, description", + gemini_body_code_429_test_cases, +) +def test_gemini_upstream_error_body_code_429_maps_to_rate_limit( + status_code, error_body, expected_exception, description +): + """ + Body code 429 inside a 5xx envelope -> RateLimitError so Router + retries kick in. Body code 429 inside a 4xx envelope must fall + through to the HTTP-status-code branch (P1 from greptile review). + """ + model = "gemini/gemini-2.5-flash" + custom_llm_provider = "gemini" + + # Build an exception that looks like what _handle_error produces: + # a BaseLLMException-style object with .status_code and .message + class _FakeGeminiError(Exception): + def __init__(self, status_code, message): + self.status_code = status_code + self.message = message + super().__init__(message) + + original_exception = _FakeGeminiError(status_code=status_code, message=error_body) + + with pytest.raises(expected_exception) as excinfo: + exception_type( + model=model, + original_exception=original_exception, + custom_llm_provider=custom_llm_provider, + ) + assert isinstance(excinfo.value, expected_exception), description + + class TestExtractAndRaiseLitellmException: """Tests for extract_and_raise_litellm_exception function""" diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 228fb2dd984..e0d7f22f817 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -2116,6 +2116,88 @@ def test_get_error_information_error_code_priority(): assert result["error_class"] == "NoCodeException" +def test_get_error_information_prefers_message_attribute_over_str(): + """ + Regression for empty-error_message-in-spend-logs. + + ProxyException sets `self.message` but does NOT call + `super().__init__(message)` nor define `__str__`, so `str(exc)` + returns the empty string. Before the fix, get_error_information + used `str(original_exception)` and silently stripped the + human-readable message from spend_logs.metadata.error_information, + making dashboard "LLM Failure" rows un-triagable. + + Asserts the `.message` attribute is consulted first. + """ + from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup + + # Simulate a ProxyException-shaped exception: .message set, but + # super().__init__() NOT called and no __str__ override. + class ProxyExceptionLike(Exception): + def __init__(self, message, code): + self.message = str(message) + self.code = str(code) + # NOTE: deliberately NOT calling super().__init__(message) + + msg = "Authentication Error, Invalid proxy server token passed. key=..." + exc = ProxyExceptionLike(message=msg, code=401) + + # Sanity check: this exception type's str() really is empty + assert str(exc) == "", ( + "Test premise broken — bare-base Exception now returns message; " + "review whether ProxyException fix landed at the class level instead" + ) + + result = StandardLoggingPayloadSetup.get_error_information(exc) + assert ( + result["error_message"] == msg + ), f"expected message from .message attribute, got {result['error_message']!r}" + assert result["error_code"] == "401" + assert result["error_class"] == "ProxyExceptionLike" + + +def test_get_error_information_preserves_explicit_empty_message(): + """ + An exception that deliberately sets `.message = ""` must surface + the empty string verbatim, not fall through to `str(exc)`. + + Regression for greptile P2 finding on PR #30381: a truthiness + check (`if message_attr:`) would silently mask an explicit empty + message and substitute `str(original_exception)` — which for + ProxyException-shaped objects is also empty, but for plain + `Exception("boom")` would inject the wrong string and corrupt + the error_information signal. + """ + from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup + + class ProxyExceptionLike(Exception): + def __init__(self, message, code): + self.message = message + self.code = str(code) + super().__init__("unrelated-args-summary") + + exc = ProxyExceptionLike(message="", code=500) + result = StandardLoggingPayloadSetup.get_error_information(exc) + assert result["error_message"] == "", ( + "explicit empty .message must survive verbatim; got " + f"{result['error_message']!r}" + ) + + +def test_get_error_information_falls_back_to_str_when_no_message_attr(): + """ + Plain Exception (no `.message` attr) must still produce a useful + error_message via str(exc), preserving prior behavior for + non-litellm exception types. + """ + from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup + + exc = ValueError("boom") + result = StandardLoggingPayloadSetup.get_error_information(exc) + assert result["error_message"] == "boom" + assert result["error_class"] == "ValueError" + + # ────────────────────────────────────────────────────────────────────── # Tests for _get_assembled_streaming_response non-streaming early return # ────────────────────────────────────────────────────────────────────── @@ -3115,6 +3197,71 @@ class TestFirstApiCallStartTimeSetOnce: assert user_meta == {} +def test_get_error_information_for_logging_payload_ignores_spoofed_disconnect_without_flag(): + from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup + + baseline = StandardLoggingPayloadSetup.get_error_information( + original_exception=ValueError("provider failure"), + ) + error_information, error_str = ( + StandardLoggingPayloadSetup.get_error_information_for_logging_payload( + metadata={ + "error_information": { + "error_code": "499", + "error_message": "Client disconnected the request", + "error_class": "ClientDisconnected", + } + }, + original_exception=ValueError("provider failure"), + error_str="provider failure", + ) + ) + assert error_information == baseline + assert error_str == "provider failure" + + +def test_get_error_information_for_logging_payload_client_disconnect(): + from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup + + custom_error = { + "error_code": "499", + "error_message": "Client disconnected the request", + "error_class": "ClientDisconnected", + } + error_information, error_str = ( + StandardLoggingPayloadSetup.get_error_information_for_logging_payload( + metadata={"client_disconnected": True, "error_information": custom_error}, + original_exception=None, + error_str=None, + ) + ) + assert error_information == custom_error + assert error_str == "Client disconnected the request" + + error_information, error_str = ( + StandardLoggingPayloadSetup.get_error_information_for_logging_payload( + metadata={"client_disconnected": True}, + original_exception=None, + error_str="existing error", + ) + ) + assert error_information["error_code"] == "499" + assert error_str == "existing error" + + baseline = StandardLoggingPayloadSetup.get_error_information( + original_exception=None, + ) + error_information, error_str = ( + StandardLoggingPayloadSetup.get_error_information_for_logging_payload( + metadata={}, + original_exception=None, + error_str=None, + ) + ) + assert error_information == baseline + assert error_str is None + + def test_get_error_information_proxy_exception_preserves_message(): """ProxyException keeps its text in ``.message`` (str() was empty pre-fix), so error_information must still surface the message and code.""" diff --git a/tests/test_litellm/litellm_core_utils/test_token_counter.py b/tests/test_litellm/litellm_core_utils/test_token_counter.py index 92c070501b4..60e5a797627 100644 --- a/tests/test_litellm/litellm_core_utils/test_token_counter.py +++ b/tests/test_litellm/litellm_core_utils/test_token_counter.py @@ -523,7 +523,6 @@ from unittest.mock import MagicMock, patch from litellm.utils import _select_tokenizer_helper, claude_json_str, encoding - # Clear the cache at module load to ensure clean state _select_tokenizer_helper.cache_clear() @@ -1010,3 +1009,64 @@ def test_token_counter_with_thinking_content(): assert ( tokens_no_thinking < 15 ), f"Expected minimal token count for empty thinking block, got {tokens_no_thinking}" + + +def test_token_counter_with_tool_reference_block(): + """ + Regression test: a message containing an Anthropic tool-search + `tool_reference` content block must NOT raise. + + Before the fix, token_counter raised + `Invalid content item type: tool_reference`. On the streaming + anthropic_messages proxy path this nulled response_cost and caused the + SpendLogs row to be dropped, silently undercounting cost. token_counter + must instead count the referenced tool name and return a positive count. + """ + messages = [ + { + "role": "assistant", + "content": [ + {"type": "text", "text": "Let me look up the right tool."}, + {"type": "tool_reference", "tool_name": "search_knowledge_base"}, + ], + } + ] + + # Must not raise, and must produce a positive token count. + tokens = token_counter_new( + model="anthropic/claude-sonnet-4-5-20250929", messages=messages + ) + assert tokens > 0, f"Expected positive token count, got {tokens}" + + # A tool_reference with no/empty tool_name must also be handled gracefully. + messages_empty = [ + { + "role": "assistant", + "content": [{"type": "tool_reference", "tool_name": ""}], + } + ] + tokens_empty = token_counter_new( + model="anthropic/claude-sonnet-4-5-20250929", messages=messages_empty + ) + assert tokens_empty >= 0 + + +def test_count_content_list_rejects_unknown_type(): + """ + An unrecognized content block type must raise, and the error message must + enumerate the supported types (including `tool_reference`). This pins the + catch-all contract so a future block type isn't silently dropped. + """ + from litellm.litellm_core_utils.token_counter import _count_content_list + + with pytest.raises(ValueError) as exc_info: + _count_content_list( + count_function=len, + content_list=[{"type": "totally_unknown_block"}], + use_default_image_token_count=False, + default_token_count=None, + ) + + message = str(exc_info.value) + assert "Invalid content item type: totally_unknown_block" in message + assert "tool_reference" in message diff --git a/tests/test_litellm/litellm_core_utils/test_tool_search_spend_logging.py b/tests/test_litellm/litellm_core_utils/test_tool_search_spend_logging.py new file mode 100644 index 00000000000..813b4a5701f --- /dev/null +++ b/tests/test_litellm/litellm_core_utils/test_tool_search_spend_logging.py @@ -0,0 +1,131 @@ +""" +Integration / regression tests for Anthropic tool-search (`tool_reference`) +content blocks on the cost-calculation and streaming-assembly paths used by +Claude Code. + +Claude Code's tool-search feature emits assistant content blocks of the form +``{"type": "tool_reference", "tool_name": ...}`` -- a lightweight pointer to a +deferred tool. Before the fix, `token_counter` did not recognise this block +type and raised ``Invalid content item type: tool_reference``. + +Why this matters (the bug these tests guard against): + + * On the cost path, that exception propagates out of ``completion_cost`` -> + ``response_cost_calculator``. The proxy logging layer catches it and nulls + ``response_cost``; the spend-tracking callback then skips the request, so + the entire SpendLogs row is dropped. The request succeeds for the caller + but the spend is silently never recorded -- a cost undercount on ALL + tool-search traffic. + + * On the streaming-assembly path, ``stream_chunk_builder`` recomputes the + prompt tokens from the request messages when the provider stream does not + carry usage. The same exception there was swallowed and prompt tokens + silently collapsed to 0 -- a quieter undercount of the same traffic. + +These tests exercise the real public entry points (not the private +``_count_content_list`` helper) so the whole chain is covered end to end. +""" + +import os +import sys + +sys.path.insert(0, os.path.abspath("../../..")) + +import litellm +from litellm import stream_chunk_builder +from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + +ANTHROPIC_MODEL = "anthropic/claude-sonnet-4-5-20250929" + +# Mirrors a Claude Code tool-search turn: a normal text block followed by a +# `tool_reference` pointer to a deferred tool. +TOOL_SEARCH_MESSAGES = [ + { + "role": "assistant", + "content": [ + {"type": "text", "text": "Let me look up the right tool."}, + {"type": "tool_reference", "tool_name": "search_knowledge_base"}, + ], + } +] + + +def test_completion_cost_with_tool_reference_records_spend(): + """ + ``completion_cost`` must return a real, positive cost for messages that + contain a tool-search ``tool_reference`` block. + + This is the exact chain that fails on the streaming anthropic_messages + proxy path: before the fix ``completion_cost`` raised, the logging layer + caught the exception and set ``response_cost = None``, and the spend + callback then dropped the SpendLogs row. A positive cost here means the + row is recorded instead of silently dropped. + """ + cost = litellm.completion_cost(model=ANTHROPIC_MODEL, messages=TOOL_SEARCH_MESSAGES) + + assert cost is not None, "response_cost is None -> SpendLogs row would be dropped" + assert cost > 0, f"Expected a positive cost for tool-search traffic, got {cost}" + + +def test_completion_cost_with_empty_tool_name_records_spend(): + """A ``tool_reference`` with an empty/missing ``tool_name`` must also cost + out cleanly rather than raising and nulling the spend.""" + messages = [ + { + "role": "assistant", + "content": [{"type": "tool_reference", "tool_name": ""}], + } + ] + + cost = litellm.completion_cost(model=ANTHROPIC_MODEL, messages=messages) + + assert cost is not None + assert cost >= 0 + + +def test_stream_chunk_builder_counts_prompt_tokens_for_tool_reference(): + """ + On the streaming-assembly path used by Claude Code, when the provider + stream carries no prompt-token usage, ``stream_chunk_builder`` recomputes + prompt tokens from the request messages via ``token_counter``. + + With a ``tool_reference`` block in those messages the count must be + positive. Before the fix the underlying ``token_counter`` call raised and + the assembler swallowed it, collapsing ``prompt_tokens`` to 0 -- a silent + undercount of every tool-search request. + """ + model = "claude-sonnet-4-5-20250929" + # Chunks deliberately carry no usage, forcing the prompt-token fallback. + chunks = [ + ModelResponseStream( + id="chatcmpl-tool-search", + created=1700000000, + model=model, + object="chat.completion.chunk", + choices=[ + StreamingChoices( + finish_reason=None, + index=0, + delta=Delta(content="Searching...", role="assistant"), + ) + ], + ), + ModelResponseStream( + id="chatcmpl-tool-search", + created=1700000000, + model=model, + object="chat.completion.chunk", + choices=[ + StreamingChoices( + finish_reason="stop", index=0, delta=Delta(content="") + ), + ], + ), + ] + + response = stream_chunk_builder(chunks, messages=TOOL_SEARCH_MESSAGES) + + assert response is not None + assert ( + response.usage.prompt_tokens > 0 + ), "prompt_tokens collapsed to 0 -> tool-search traffic silently undercounted" diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py index abb162e9ddb..2876b56f516 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_transformation.py @@ -3702,6 +3702,39 @@ def test_fast_mode_with_inference_geo(): assert abs(completion_cost - base_completion * expected_multiplier) < 1e-10 +def test_calculate_usage_captures_service_tier(): + """ + Anthropic returns the assigned service tier on the response usage object + (e.g. ``"priority"``). It must be surfaced on the Usage object so it is + visible in logs and used to select tier-specific pricing. + """ + config = AnthropicConfig() + + usage_object = { + "input_tokens": 410, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "output_tokens": 585, + "service_tier": "priority", + } + + usage = config.calculate_usage(usage_object=usage_object, reasoning_content=None) + + assert usage.service_tier == "priority" + + +def test_calculate_usage_service_tier_defaults_to_none(): + """A response without a service tier must not invent one.""" + config = AnthropicConfig() + + usage = config.calculate_usage( + usage_object={"input_tokens": 10, "output_tokens": 5}, + reasoning_content=None, + ) + + assert usage.service_tier is None + + def test_fast_mode_parameter_in_supported_params(): """ Test that 'speed' is in the list of supported OpenAI params. diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py index a81261d5ffd..76aa3a9c6aa 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py @@ -1295,6 +1295,12 @@ CACHE_CONTROL_BEDROCK_CONVERSE_MODEL = ( "bedrock/converse/global.anthropic.claude-opus-4-5-20251101-v1:0" ) CACHE_CONTROL_NON_ANTHROPIC_MODEL = "gpt-4" +# Bedrock Application Inference Profile ARN: the string contains neither +# "anthropic" nor "claude", so the model can only be recognized via its ARN shape +CACHE_CONTROL_BEDROCK_ARN_MODEL = ( + "bedrock/converse/arn:aws:bedrock:us-east-1:123456789012:" + "application-inference-profile/abcdef123456" +) def test_should_add_cache_control_for_anthropic_model(): @@ -1411,6 +1417,68 @@ def test_cache_control_not_preserved_for_non_claude_model(): assert "cache_control" not in result[0]["content"][0] +@pytest.mark.parametrize( + "model, expected", + [ + (CACHE_CONTROL_BEDROCK_ARN_MODEL, True), + ( + "arn:aws-us-gov:bedrock:us-gov-west-1:123:application-inference-profile/x", + True, + ), + ("bedrock/amazon.titan-text-express-v1", False), + ("arn:aws:sagemaker:us-east-1:123:endpoint/my-endpoint", False), + ("arn:aws:sagemaker:us-east-1:123:endpoint/my-bedrock-transcriber", False), + (CACHE_CONTROL_NON_ANTHROPIC_MODEL, False), + ], +) +def test_is_bedrock_arn_model(model, expected): + """is_bedrock_arn_model requires an ARN with bedrock in the service field, not just anywhere.""" + assert LiteLLMAnthropicMessagesAdapter.is_bedrock_arn_model(model) is expected + + +def test_cache_control_preserved_for_bedrock_arn_inference_profile(): + """ + Regression for https://github.com/BerriAI/litellm/issues/26625 + + Bedrock Application Inference Profile ARNs hide the underlying Claude model + name, so cache_control must still be preserved through the /v1/messages adapter. + """ + anthropic_messages = [ + AnthropicMessagesUserMessageParam( + role="user", + content=[ + { + "type": "text", + "text": "This is cached content", + "cache_control": {"type": "ephemeral"}, + } + ], + ) + ] + + adapter = LiteLLMAnthropicMessagesAdapter() + result = adapter.translate_anthropic_messages_to_openai( + messages=anthropic_messages, model=CACHE_CONTROL_BEDROCK_ARN_MODEL + ) + + assert len(result) == 1 + assert result[0]["content"][0]["cache_control"] == {"type": "ephemeral"} + + +def test_cache_control_fix_does_not_broaden_claude_detection(): + """ + The cache_control fix is scoped to _add_cache_control_if_applicable; it must not + make is_anthropic_claude_model treat ARN profiles as Claude, which would route + thinking params through unmodified and break non-Claude Bedrock profiles. + """ + assert ( + LiteLLMAnthropicMessagesAdapter.is_anthropic_claude_model( + CACHE_CONTROL_BEDROCK_ARN_MODEL + ) + is False + ) + + def test_cache_control_preserved_in_image_content_for_claude(): """Cache control should be preserved in image content for Claude models.""" anthropic_messages = [ diff --git a/tests/test_litellm/llms/bedrock/chat/agentcore/test_agentcore_transformation.py b/tests/test_litellm/llms/bedrock/chat/agentcore/test_agentcore_transformation.py index 64b43b15dcd..448afd5f3a5 100644 --- a/tests/test_litellm/llms/bedrock/chat/agentcore/test_agentcore_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/agentcore/test_agentcore_transformation.py @@ -5,6 +5,7 @@ Tests: - Accept header fix (sign_request sets Accept: application/json, text/event-stream) - JSON response parsing fallback chain (_parse_json_response supports multiple schemas) - Streaming Content-Type fallback (JSON responses converted to single-chunk streams) +- Multimodal content preservation (transform_request forwards OpenAI content blocks) """ import json @@ -389,3 +390,249 @@ class TestAgentCoreStreamingJsonFallback: client=client, api_key="test-jwt-token", ) + + +class TestAgentCoreMultimodalContent: + """Tests for transform_request forwarding OpenAI multimodal content blocks. + + AgentCore Runtime is schemaless on the agent side — the agent author's + @app.entrypoint handler parses whatever JSON arrives. transform_request + only emits {"prompt": ""} by default and drops image_url, file, and + other non-text blocks. + + When the ``forward_multimodal_content`` litellm param is set, the OpenAI + content list is forwarded verbatim under a "content" field whenever the last + message contains a non-text block. This is opt-in: an agent must be written + to read payload["content"]. Without the flag, the payload is byte-identical + to the legacy {"prompt": "..."} shape. + """ + + @pytest.fixture + def config(self): + return AmazonAgentCoreConfig() + + @pytest.fixture + def transform_kwargs(self): + """Default kwargs — forwarding is OFF (no opt-in flag).""" + return { + "model": "bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:111111111111:runtime/test_agent", + "optional_params": {}, + "litellm_params": {}, + "headers": {}, + } + + @pytest.fixture + def opted_in_kwargs(self, transform_kwargs): + """Kwargs with the opt-in flag set in optional_params.""" + return { + **transform_kwargs, + "optional_params": {"forward_multimodal_content": True}, + } + + def test_string_content_payload_byte_identical_to_legacy( + self, config, transform_kwargs + ): + """String content → exactly {"prompt": ""}, no extra fields.""" + messages = [{"role": "user", "content": "hello agent"}] + payload = config.transform_request(messages=messages, **transform_kwargs) + assert payload == {"prompt": "hello agent"} + + def test_file_block_not_forwarded_by_default(self, config, transform_kwargs): + """Default (no opt-in flag): file blocks are NOT forwarded — backward compat.""" + content = [ + {"type": "text", "text": "summarize this report"}, + { + "type": "file", + "file": { + "filename": "report.pdf", + "file_data": "data:application/pdf;base64,JVBERi0xLjQK", + }, + }, + ] + messages = [{"role": "user", "content": content}] + payload = config.transform_request(messages=messages, **transform_kwargs) + assert payload == {"prompt": "summarize this report"} + assert "content" not in payload + + def test_text_only_list_content_no_content_field(self, config, opted_in_kwargs): + """All-text content list → no "content" field even when opted in.""" + messages = [ + { + "role": "user", + "content": [{"type": "text", "text": "hello agent"}], + } + ] + payload = config.transform_request(messages=messages, **opted_in_kwargs) + assert payload == {"prompt": "hello agent"} + assert "content" not in payload + + def test_file_data_block_passthrough(self, config, opted_in_kwargs): + """Opted in: a file block → "content" carries the original list verbatim.""" + content = [ + {"type": "text", "text": "summarize this report"}, + { + "type": "file", + "file": { + "filename": "report.pdf", + "file_data": "data:application/pdf;base64,JVBERi0xLjQK", + }, + }, + ] + messages = [{"role": "user", "content": content}] + payload = config.transform_request(messages=messages, **opted_in_kwargs) + assert payload["prompt"] == "summarize this report" + # Contents forwarded verbatim, but as a distinct list (no aliasing). + assert payload["content"] == content + assert payload["content"] is not content + + def test_image_url_block_passthrough(self, config, opted_in_kwargs): + """Opted in: an image_url block → "content" carries it verbatim.""" + content = [ + {"type": "text", "text": "what is in this image?"}, + { + "type": "image_url", + "image_url": {"url": "data:image/png;base64,iVBORw0KGgo="}, + }, + ] + messages = [{"role": "user", "content": content}] + payload = config.transform_request(messages=messages, **opted_in_kwargs) + assert payload["prompt"] == "what is in this image?" + assert payload["content"] == content + assert payload["content"] is not content + + def test_mixed_text_and_files_payload_shape(self, config, opted_in_kwargs): + """Opted in: text + file + image → both "prompt" (text-only) and "content".""" + content = [ + {"type": "text", "text": "first sentence."}, + { + "type": "file", + "file": { + "filename": "report.pdf", + "file_data": "data:application/pdf;base64,JVBERi0xLjQK", + }, + }, + {"type": "text", "text": "second sentence."}, + { + "type": "image_url", + "image_url": {"url": "data:image/png;base64,iVBORw0KGgo="}, + }, + ] + messages = [{"role": "user", "content": content}] + payload = config.transform_request(messages=messages, **opted_in_kwargs) + # prompt is the text-only flatten produced by convert_content_list_to_str. + assert "first sentence." in payload["prompt"] + assert "second sentence." in payload["prompt"] + assert "JVBERi0xLjQK" not in payload["prompt"] + assert "iVBORw0KGgo=" not in payload["prompt"] + # content carries every block in original order. + assert payload["content"] == content + + def test_forwarded_content_does_not_alias_message(self, config, opted_in_kwargs): + """Regression: the forwarded list is a shallow copy, so mutating the + returned payload before serialization must not leak back into the caller's + messages[-1]["content"].""" + content = [ + {"type": "text", "text": "describe this"}, + { + "type": "image_url", + "image_url": {"url": "data:image/png;base64,iVBORw0KGgo="}, + }, + ] + messages = [{"role": "user", "content": content}] + payload = config.transform_request(messages=messages, **opted_in_kwargs) + + payload["content"].append({"type": "text", "text": "injected"}) + + assert len(messages[-1]["content"]) == 2 + assert {"type": "text", "text": "injected"} not in messages[-1]["content"] + + def test_only_last_message_content_preserved(self, config, opted_in_kwargs): + """Opted in: file blocks in earlier messages don't trigger "content" — last only.""" + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "context"}, + { + "type": "file", + "file": { + "filename": "old.pdf", + "file_data": "data:application/pdf;base64,Zm9v", + }, + }, + ], + }, + {"role": "assistant", "content": "ok"}, + {"role": "user", "content": "follow-up question with no files"}, + ] + payload = config.transform_request(messages=messages, **opted_in_kwargs) + assert payload == {"prompt": "follow-up question with no files"} + assert "content" not in payload + + def test_unknown_non_text_block_type_passthrough(self, config, opted_in_kwargs): + """Opted in: unknown block types (e.g. input_audio) flow through.""" + content = [ + {"type": "text", "text": "transcribe this"}, + { + "type": "input_audio", + "input_audio": {"data": "U29tZUF1ZGlvQnl0ZXM=", "format": "wav"}, + }, + ] + messages = [{"role": "user", "content": content}] + payload = config.transform_request(messages=messages, **opted_in_kwargs) + assert payload["prompt"] == "transcribe this" + assert payload["content"] == content + assert payload["content"] is not content + + def test_forward_flag_as_string_true(self, config, transform_kwargs): + """The opt-in flag accepts config/env string values like "true".""" + content = [ + {"type": "text", "text": "hi"}, + { + "type": "image_url", + "image_url": {"url": "data:image/png;base64,iVBORw0KGgo="}, + }, + ] + messages = [{"role": "user", "content": content}] + kwargs = { + **transform_kwargs, + "optional_params": {"forward_multimodal_content": "true"}, + } + payload = config.transform_request(messages=messages, **kwargs) + assert payload["content"] == content + assert payload["content"] is not content + + def test_forward_flag_false_explicit(self, config, transform_kwargs): + """Explicit falsy flag → no content field.""" + content = [ + {"type": "text", "text": "hi"}, + { + "type": "image_url", + "image_url": {"url": "data:image/png;base64,iVBORw0KGgo="}, + }, + ] + messages = [{"role": "user", "content": content}] + kwargs = { + **transform_kwargs, + "optional_params": {"forward_multimodal_content": False}, + } + payload = config.transform_request(messages=messages, **kwargs) + assert "content" not in payload + + def test_forward_flag_via_litellm_params(self, config, transform_kwargs): + """The opt-in flag is also honored when set in litellm_params.""" + content = [ + {"type": "text", "text": "hi"}, + { + "type": "image_url", + "image_url": {"url": "data:image/png;base64,iVBORw0KGgo="}, + }, + ] + messages = [{"role": "user", "content": content}] + kwargs = { + **transform_kwargs, + "litellm_params": {"forward_multimodal_content": True}, + } + payload = config.transform_request(messages=messages, **kwargs) + assert payload["content"] == content + assert payload["content"] is not content diff --git a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_transformation.py b/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_transformation.py index 6fb02113a45..09437102d30 100644 --- a/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_transformation.py +++ b/tests/test_litellm/llms/bedrock_mantle/test_bedrock_mantle_transformation.py @@ -171,6 +171,13 @@ class TestBedrockMantleConfig: _, api_key = cfg._get_openai_compatible_provider_info(None, "explicit-key") assert api_key == "explicit-key" + def test_api_key_from_aws_bearer_token_bedrock_env(self, monkeypatch): + monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) + monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "standard-bearer") + cfg = BedrockMantleChatConfig() + _, api_key = cfg._get_openai_compatible_provider_info(None, None) + assert api_key == "standard-bearer" + def test_get_supported_openai_params(self): cfg = BedrockMantleChatConfig() params = cfg.get_supported_openai_params("openai.gpt-oss-120b") @@ -181,6 +188,260 @@ class TestBedrockMantleConfig: assert "max_tokens" in params +class TestBedrockMantleChatAuth: + """Chat Completions must use the same Bearer-or-SigV4 auth as the Responses + backend. These fail on a config that inherits the no-op default sign_request. + """ + + def _signer_that_forbids_credentials(self): + from unittest.mock import MagicMock + + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + signer = BaseAWSLLM() + signer.get_credentials = MagicMock( + side_effect=AssertionError("SigV4 must not run when a Bearer token exists") + ) + return signer + + def test_bearer_token_skips_sigv4(self, monkeypatch): + monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + signer = self._signer_that_forbids_credentials() + cfg = BedrockMantleChatConfig(aws_signer=signer) + + headers, signed_body = cfg.sign_request( + headers={}, + optional_params={}, + request_data={"model": "openai.gpt-oss-120b", "messages": []}, + api_base="https://bedrock-mantle.us-east-1.api.aws/v1/chat/completions", + api_key="bearer-from-arg", + ) + + assert headers["Authorization"] == "Bearer bearer-from-arg" + assert json.loads(signed_body) == { + "model": "openai.gpt-oss-120b", + "messages": [], + } + signer.get_credentials.assert_not_called() + + def test_mantle_env_key_used_as_bearer(self, monkeypatch): + monkeypatch.setenv("BEDROCK_MANTLE_API_KEY", "env-mantle-key") + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + signer = self._signer_that_forbids_credentials() + cfg = BedrockMantleChatConfig(aws_signer=signer) + + headers, _ = cfg.sign_request( + headers={}, + optional_params={}, + request_data={"input": "hi"}, + api_base="https://bedrock-mantle.us-east-1.api.aws/v1/chat/completions", + api_key=None, + ) + + assert headers["Authorization"] == "Bearer env-mantle-key" + signer.get_credentials.assert_not_called() + + def test_aws_bearer_token_bedrock_used_as_bearer(self, monkeypatch): + monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) + monkeypatch.setenv("AWS_BEARER_TOKEN_BEDROCK", "standard-bearer") + signer = self._signer_that_forbids_credentials() + cfg = BedrockMantleChatConfig(aws_signer=signer) + + headers, _ = cfg.sign_request( + headers={}, + optional_params={}, + request_data={"input": "hi"}, + api_base="https://bedrock-mantle.us-east-1.api.aws/v1/chat/completions", + api_key=None, + ) + + assert headers["Authorization"] == "Bearer standard-bearer" + signer.get_credentials.assert_not_called() + + def test_no_bearer_signs_with_sigv4(self, monkeypatch): + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + + cfg = BedrockMantleChatConfig(aws_signer=BaseAWSLLM()) + headers, signed_body = cfg.sign_request( + headers={}, + optional_params={ + "aws_access_key_id": "AKIAEXAMPLE", + "aws_secret_access_key": "c2VjcmV0LXRlc3Qtc2VjcmV0LXRlc3Qtc2VjcmV0", + "aws_session_token": "session-token-test", + "aws_region_name": "us-east-2", + }, + request_data={"model": "openai.gpt-oss-120b", "messages": []}, + api_base="https://bedrock-mantle.us-east-2.api.aws/v1/chat/completions", + api_key=None, + ) + + assert headers["Authorization"].startswith("AWS4-HMAC-SHA256") + assert "Credential=AKIAEXAMPLE/" in headers["Authorization"] + assert "/us-east-2/bedrock/aws4_request" in headers["Authorization"] + assert headers["X-Amz-Security-Token"] == "session-token-test" + assert json.loads(signed_body) == { + "model": "openai.gpt-oss-120b", + "messages": [], + } + + def test_sigv4_region_resolved_from_api_base_host(self, monkeypatch): + # Chat passes the OpenAI-mapped optional_params (no aws_region_name) to + # sign_request, so the SigV4 credential scope has to come from the already + # region-resolved api_base host or it would disagree with the URL -> 401. + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + for var in ( + "BEDROCK_MANTLE_API_KEY", + "AWS_BEARER_TOKEN_BEDROCK", + "BEDROCK_MANTLE_REGION", + "BEDROCK_MANTLE_API_BASE", + "AWS_REGION", + "AWS_REGION_NAME", + ): + monkeypatch.delenv(var, raising=False) + + cfg = BedrockMantleChatConfig(aws_signer=BaseAWSLLM()) + headers, _ = cfg.sign_request( + headers={}, + optional_params={ + "aws_access_key_id": "AKIAEXAMPLE", + "aws_secret_access_key": "c2VjcmV0LXRlc3Qtc2VjcmV0LXRlc3Qtc2VjcmV0", + }, + request_data={"input": "hi"}, + api_base="https://bedrock-mantle.eu-west-1.api.aws/v1/chat/completions", + api_key=None, + ) + + assert "/eu-west-1/bedrock/aws4_request" in headers["Authorization"] + + def test_sigv4_scope_matches_api_base_when_aws_region_name_disagrees( + self, monkeypatch + ): + # If a caller (e.g. proxy) passes a stale api_base in one region and an + # aws_region_name in a different region, the SigV4 credential scope must + # match the URL host or Bedrock rejects the request with 401. Without the + # fix, sign_request would prefer aws_region_name and sign for us-west-2 + # while POSTing to eu-west-1. + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + for var in ( + "BEDROCK_MANTLE_API_KEY", + "AWS_BEARER_TOKEN_BEDROCK", + "BEDROCK_MANTLE_REGION", + "BEDROCK_MANTLE_API_BASE", + "AWS_REGION", + "AWS_REGION_NAME", + ): + monkeypatch.delenv(var, raising=False) + + cfg = BedrockMantleChatConfig(aws_signer=BaseAWSLLM()) + headers, _ = cfg.sign_request( + headers={}, + optional_params={ + "aws_access_key_id": "AKIAEXAMPLE", + "aws_secret_access_key": "c2VjcmV0LXRlc3Qtc2VjcmV0LXRlc3Qtc2VjcmV0", + "aws_region_name": "us-west-2", + }, + request_data={"input": "hi"}, + api_base="https://bedrock-mantle.eu-west-1.api.aws/v1/chat/completions", + api_key=None, + ) + + assert "/eu-west-1/bedrock/aws4_request" in headers["Authorization"] + assert "/us-west-2/bedrock/aws4_request" not in headers["Authorization"] + + def test_no_bearer_and_no_credentials_raises_value_error(self, monkeypatch): + from unittest.mock import MagicMock + + from botocore.exceptions import NoCredentialsError + + from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM + + monkeypatch.delenv("BEDROCK_MANTLE_API_KEY", raising=False) + monkeypatch.delenv("AWS_BEARER_TOKEN_BEDROCK", raising=False) + + signer = BaseAWSLLM() + signer.get_credentials = MagicMock(side_effect=NoCredentialsError()) + cfg = BedrockMantleChatConfig(aws_signer=signer) + + with pytest.raises(ValueError) as exc: + cfg.sign_request( + headers={}, + optional_params={"aws_region_name": "us-east-2"}, + request_data={"input": "hi"}, + api_base="https://bedrock-mantle.us-east-2.api.aws/v1/chat/completions", + api_key=None, + ) + + msg = str(exc.value) + assert "Bearer" in msg + assert "SigV4" in msg or "IAM" in msg + + def test_completion_no_bearer_signs_with_sigv4_end_to_end(self, monkeypatch): + # The full completion chain (not just sign_request in isolation) must reach + # the SigV4 path when no Bearer token exists: with api_key=None the parent + # validate_environment must not short-circuit before sign_request runs. + for var in ( + "BEDROCK_MANTLE_API_KEY", + "AWS_BEARER_TOKEN_BEDROCK", + "BEDROCK_MANTLE_API_BASE", + "BEDROCK_MANTLE_REGION", + "AWS_REGION_NAME", + ): + monkeypatch.delenv(var, raising=False) + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "AKIAEXAMPLE") + monkeypatch.setenv( + "AWS_SECRET_ACCESS_KEY", "c2VjcmV0LXRlc3Qtc2VjcmV0LXRlc3Qtc2VjcmV0" + ) + monkeypatch.setenv("AWS_REGION", "us-east-2") + + requests = [] + + def mock_post(self, url, data=None, headers=None, **kwargs): + requests.append({"url": url, "headers": headers or {}}) + return httpx.Response( + status_code=200, + json={ + "id": "chatcmpl-test", + "object": "chat.completion", + "created": 1733529600, + "model": "openai.gpt-oss-120b", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "ok"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 1, + "completion_tokens": 1, + "total_tokens": 2, + }, + }, + request=httpx.Request("POST", url), + ) + + with patch( + "litellm.llms.custom_httpx.http_handler.HTTPHandler.post", mock_post + ): + response = litellm.completion( + model="bedrock_mantle/openai.gpt-oss-120b", + messages=[{"role": "user", "content": "hello"}], + ) + + assert response.choices[0].message.content == "ok" + assert len(requests) == 1 + authorization = requests[0]["headers"]["Authorization"] + assert authorization.startswith("AWS4-HMAC-SHA256") + assert "/us-east-2/bedrock/aws4_request" in authorization + assert requests[0]["url"].startswith("https://bedrock-mantle.us-east-2.api.aws") + + class TestBedrockMantleProjectHeader: def test_validate_environment_sets_openai_project_header(self): cfg = BedrockMantleChatConfig() diff --git a/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py b/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py index 474ffee3304..0c5e386c438 100644 --- a/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py +++ b/tests/test_litellm/llms/custom_httpx/test_aiohttp_transport.py @@ -63,10 +63,14 @@ class MockAiohttpResponse: ): self.status = status self.headers = headers or {} + self.closed = False self.content = MockContent( content_chunks, exception_to_raise, exception_at_chunk ) + def close(self): + self.closed = True + async def __aexit__(self, exc_type, exc_val, exc_tb): pass @@ -613,3 +617,64 @@ async def test_handle_session_closed_during_request(): assert counts["requests"] == 2 # First request failed, second succeeded assert counts["sessions"] == 2 # Created 2 sessions for retry assert response.status_code == 200 + + +@pytest.mark.asyncio +async def test_response_stream_closes_response_on_error(): + """ + Regression test for #30192: when body iteration ends with an error, the + underlying aiohttp response must be closed so its connector slot is + released. Leaked slots exhaust the pool and every later request times + out (408) until the proxy restarts, even after the backend recovers. + """ + mock_response = MockAiohttpResponse( + content_chunks=[b"chunk1", b"chunk2"], + exception_to_raise=aiohttp.ServerTimeoutError("read timeout"), + exception_at_chunk=1, + ) + + stream = AiohttpResponseStream(mock_response) # type: ignore + with pytest.raises(httpx.TimeoutException): + async for _ in stream: + pass + + assert mock_response.closed is True + + +@pytest.mark.asyncio +async def test_response_stream_closes_response_on_cancellation(): + """ + Regression test for #30192: a task cancelled mid-stream (e.g. the caller + disconnects during a traffic spike) must not leak its aiohttp connection. + """ + mock_response = MockAiohttpResponse( + content_chunks=[b"chunk1", b"chunk2", b"chunk3"], + exception_to_raise=asyncio.CancelledError(), + exception_at_chunk=1, + ) + + stream = AiohttpResponseStream(mock_response) # type: ignore + with pytest.raises(asyncio.CancelledError): + async for _ in stream: + pass + + assert mock_response.closed is True + + +@pytest.mark.asyncio +async def test_response_stream_closes_response_on_generator_exit(): + """ + Regression test for #30192: when the consumer stops iterating early and the + stream generator is closed (GeneratorExit), the underlying aiohttp response + must still be closed so its connector slot is released. + """ + mock_response = MockAiohttpResponse( + content_chunks=[b"chunk1", b"chunk2", b"chunk3"], + ) + + stream = AiohttpResponseStream(mock_response) # type: ignore + iterator = stream.__aiter__() + assert await iterator.__anext__() == b"chunk1" + await iterator.aclose() + + assert mock_response.closed is True diff --git a/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py b/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py index 0221db1b23d..683ec158f44 100644 --- a/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py +++ b/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py @@ -554,3 +554,31 @@ def test_map_response_format_json_object_unchanged(): drop_params=False, ) assert result == {"response_format": {"type": "json_object"}} + + +def test_transform_request_routes_short_form_router_to_routers_path(): + """A bare router model name ending in -fast must be rewritten to the + ``accounts/fireworks/routers/`` path, not the default ``models/`` path.""" + config = FireworksAIConfig() + result = config.transform_request( + model="glm-5p1-fast", + messages=[{"role": "user", "content": "Hi"}], + optional_params={}, + litellm_params={}, + headers={}, + ) + assert result["model"] == "accounts/fireworks/routers/glm-5p1-fast" + + +def test_transform_request_routes_short_form_model_to_models_path(): + """A bare direct-model name must still be rewritten to the + ``accounts/fireworks/models/`` path.""" + config = FireworksAIConfig() + result = config.transform_request( + model="glm-5p2", + messages=[{"role": "user", "content": "Hi"}], + optional_params={}, + litellm_params={}, + headers={}, + ) + assert result["model"] == "accounts/fireworks/models/glm-5p2" diff --git a/tests/test_litellm/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py b/tests/test_litellm/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py index 29ab3790609..5c1f1dcb63d 100644 --- a/tests/test_litellm/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py +++ b/tests/test_litellm/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py @@ -169,8 +169,8 @@ def test_hosted_vllm_supports_thinking(): def test_hosted_vllm_thinking_blocks_prepended_to_assistant_content(): """ - Test that thinking_blocks on assistant messages are converted to content - blocks prepended before the existing content. + Test that thinking_blocks on assistant messages are removed and content + stays a string for vLLM compatibility. """ config = HostedVLLMChatConfig() messages = [ @@ -203,21 +203,15 @@ def test_hosted_vllm_thinking_blocks_prepended_to_assistant_content(): ) assistant_msg = transformed["messages"][1] assert assistant_msg["role"] == "assistant" - assert isinstance(assistant_msg["content"], list) - assert assistant_msg["content"][0] == { - "type": "thinking", - "thinking": "Let me reason about this...", - } - assert assistant_msg["content"][1] == { - "type": "text", - "text": "Here is my answer.", - } + assert isinstance(assistant_msg["content"], str) + assert assistant_msg["content"] == "Here is my answer." assert "thinking_blocks" not in assistant_msg def test_hosted_vllm_thinking_blocks_with_list_content(): """ - Test thinking_blocks prepended when assistant content is already a list. + Test thinking_blocks are removed and assistant content list is converted + to a string. """ config = HostedVLLMChatConfig() messages = [ @@ -246,19 +240,125 @@ def test_hosted_vllm_thinking_blocks_with_list_content(): headers={}, ) assistant_msg = transformed["messages"][0] - assert len(assistant_msg["content"]) == 3 - assert assistant_msg["content"][0] == { - "type": "thinking", - "thinking": "Step 1 reasoning", - } - assert assistant_msg["content"][1] == { - "type": "thinking", - "thinking": "Step 2 reasoning", - } - assert assistant_msg["content"][2] == {"type": "text", "text": "Response text"} + assert isinstance(assistant_msg["content"], str) + assert assistant_msg["content"] == "Response text" assert "thinking_blocks" not in assistant_msg +def test_hosted_vllm_assistant_structured_content_is_preserved(): + config = HostedVLLMChatConfig() + image_block = { + "type": "image_url", + "image_url": {"url": "https://example.com/image.png"}, + } + messages = [ + { + "role": "assistant", + "content": [{"type": "text", "text": "Here is the image"}, image_block], + }, + ] + + transformed = config.transform_request( + model="hosted_vllm/llama-3.1-70b-instruct", + messages=messages, + optional_params={}, + litellm_params={}, + headers={}, + ) + + assistant_msg = transformed["messages"][0] + assert assistant_msg["content"] == [ + {"type": "text", "text": "Here is the image"}, + image_block, + ] + + +def test_hosted_vllm_assistant_tool_use_content_becomes_tool_calls(): + config = HostedVLLMChatConfig() + messages = [ + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "toolu_1", + "name": "get_weather", + "input": {"city": "Boston"}, + } + ], + }, + ] + + transformed = config.transform_request( + model="hosted_vllm/llama-3.1-70b-instruct", + messages=messages, + optional_params={}, + litellm_params={}, + headers={}, + ) + + assistant_msg = transformed["messages"][0] + assert assistant_msg["content"] == "" + assert assistant_msg["tool_calls"] == [ + { + "id": "toolu_1", + "type": "function", + "function": { + "name": "get_weather", + "arguments": json.dumps({"city": "Boston"}), + }, + } + ] + + +def test_hosted_vllm_assistant_tool_use_does_not_duplicate_existing_tool_calls(): + config = HostedVLLMChatConfig() + messages = [ + { + "role": "assistant", + "content": [ + { + "type": "tool_use", + "id": "toolu_1", + "name": "get_weather", + "input": {"city": "Boston"}, + } + ], + "tool_calls": [ + { + "id": "toolu_1", + "type": "function", + "function": { + "name": "get_weather", + "arguments": json.dumps({"city": "Boston"}), + }, + } + ], + }, + ] + + transformed = config.transform_request( + model="hosted_vllm/llama-3.1-70b-instruct", + messages=messages, + optional_params={}, + litellm_params={}, + headers={}, + ) + + assistant_msg = transformed["messages"][0] + assert assistant_msg["content"] == "" + assert assistant_msg["tool_calls"] == [ + { + "id": "toolu_1", + "type": "function", + "function": { + "name": "get_weather", + "arguments": json.dumps({"city": "Boston"}), + }, + } + ] + + def test_hosted_vllm_custom_tools_are_converted_to_function_tools(): config = HostedVLLMChatConfig() optional_params = config.map_openai_params( diff --git a/tests/test_litellm/llms/openai/chat/test_openai_gpt_transformation.py b/tests/test_litellm/llms/openai/chat/test_openai_gpt_transformation.py index bb9cda2584c..20a1bf85751 100644 --- a/tests/test_litellm/llms/openai/chat/test_openai_gpt_transformation.py +++ b/tests/test_litellm/llms/openai/chat/test_openai_gpt_transformation.py @@ -9,6 +9,7 @@ import pytest sys.path.insert(0, os.path.abspath("../../../../..")) +import litellm from litellm.llms.openai.chat.gpt_5_transformation import OpenAIGPT5Config from litellm.llms.openai.chat.gpt_transformation import ( OpenAIChatCompletionStreamingHandler, @@ -571,3 +572,162 @@ class TestGPT5ReasoningEffortPreservation: assert optional_params.get("temperature") == 0.5 assert non_default_params.get("reasoning_effort") == "none" + + +class TestCacheControlPreservationForCustomEndpoint: + """ + Regression tests for https://github.com/BerriAI/litellm/issues/30319 + + The AnthropicCacheControlHook injects cache_control when a user passes + cache_control_injection_points, but the base OpenAIGPTConfig used to strip + it unconditionally, making the feature a guaranteed no-op for the generic + openai provider pointed at a cache_control-aware endpoint (a LiteLLM proxy, + vLLM, an Anthropic-compatible gateway). cache_control must survive there + while still being stripped for real api.openai.com. + """ + + def setup_method(self): + self.config = OpenAIGPTConfig() + + @pytest.fixture(autouse=True) + def _clean_openai_base_env(self, monkeypatch): + monkeypatch.delenv("OPENAI_BASE_URL", raising=False) + monkeypatch.delenv("OPENAI_API_BASE", raising=False) + monkeypatch.setattr(litellm, "api_base", None, raising=False) + + @staticmethod + def _cache_controlled_messages(): + return [ + { + "role": "system", + "content": "You are helpful.", + "cache_control": {"type": "ephemeral"}, + }, + { + "role": "user", + "content": "Hello", + "cache_control": {"type": "ephemeral"}, + }, + ] + + def _transform(self, custom_llm_provider, api_base, optional_params=None): + return self.config.transform_request( + model="claude-sonnet-4", + messages=self._cache_controlled_messages(), + optional_params=optional_params or {}, + litellm_params={ + "custom_llm_provider": custom_llm_provider, + "api_base": api_base, + }, + headers={}, + ) + + def test_predicate_openai_provider_custom_api_base_preserves(self): + assert ( + self.config._should_preserve_cache_control_for_endpoint( + "openai", "http://localhost:4000/v1" + ) + is True + ) + + def test_predicate_real_openai_no_api_base_strips(self): + assert ( + self.config._should_preserve_cache_control_for_endpoint("openai", None) + is False + ) + + def test_predicate_explicit_openai_host_strips(self): + assert ( + self.config._should_preserve_cache_control_for_endpoint( + "openai", "https://api.openai.com/v1" + ) + is False + ) + + def test_predicate_non_openai_provider_strips(self): + assert ( + self.config._should_preserve_cache_control_for_endpoint( + "deepseek", "https://api.deepseek.com" + ) + is False + ) + + def test_predicate_resolves_openai_base_url_env(self, monkeypatch): + monkeypatch.setenv("OPENAI_BASE_URL", "http://localhost:4000/v1") + assert ( + self.config._should_preserve_cache_control_for_endpoint("openai", None) + is True + ) + + def test_predicate_resolves_openai_api_base_env(self, monkeypatch): + monkeypatch.setenv("OPENAI_API_BASE", "http://localhost:4000/v1") + assert ( + self.config._should_preserve_cache_control_for_endpoint("openai", None) + is True + ) + + def test_predicate_lookalike_host_is_not_treated_as_openai(self): + assert ( + self.config._should_preserve_cache_control_for_endpoint( + "openai", "https://api.openai.com.evil.example/v1" + ) + is True + ) + + def test_predicate_openai_subdomain_strips(self): + assert ( + self.config._should_preserve_cache_control_for_endpoint( + "openai", "https://eu.api.openai.com/v1" + ) + is False + ) + + def test_transform_request_preserves_for_custom_api_base(self): + body = self._transform("openai", "http://localhost:4000/v1") + assert all("cache_control" in m for m in body["messages"]) + + def test_transform_request_strips_for_real_openai(self): + body = self._transform("openai", None) + assert all("cache_control" not in m for m in body["messages"]) + + def test_transform_request_strips_for_non_openai_provider(self): + body = self._transform("fireworks_ai", "https://api.fireworks.ai/inference/v1") + assert all("cache_control" not in m for m in body["messages"]) + + def test_transform_request_preserves_tool_cache_control(self): + tools = [ + { + "type": "function", + "function": {"name": "f", "parameters": {}}, + "cache_control": {"type": "ephemeral"}, + } + ] + body = self._transform( + "openai", "http://localhost:4000/v1", optional_params={"tools": tools} + ) + assert "cache_control" in body["tools"][0] + + @pytest.mark.asyncio + async def test_async_transform_request_preserves_for_custom_api_base(self): + body = await self.config.async_transform_request( + model="claude-sonnet-4", + messages=self._cache_controlled_messages(), + optional_params={}, + litellm_params={ + "custom_llm_provider": "openai", + "api_base": "http://localhost:4000/v1", + }, + headers={}, + ) + assert all("cache_control" in m for m in body["messages"]) + + @pytest.mark.asyncio + async def test_async_transform_request_strips_for_real_openai(self): + body = await self.config.async_transform_request( + model="gpt-4o", + messages=self._cache_controlled_messages(), + optional_params={}, + litellm_params={"custom_llm_provider": "openai", "api_base": None}, + headers={}, + ) + assert all("cache_control" not in m for m in body["messages"]) diff --git a/tests/test_litellm/llms/openai/transcriptions/test_whisper_transformation.py b/tests/test_litellm/llms/openai/transcriptions/test_whisper_transformation.py new file mode 100644 index 00000000000..2dc24b1b313 --- /dev/null +++ b/tests/test_litellm/llms/openai/transcriptions/test_whisper_transformation.py @@ -0,0 +1,106 @@ +""" +Tests for OpenAIWhisperAudioTranscriptionConfig.transform_audio_transcription_request +and transform_audio_transcription_response. +""" + +import io +import json +from unittest.mock import MagicMock + +import pytest + +from litellm.llms.openai.transcriptions.whisper_transformation import ( + OpenAIWhisperAudioTranscriptionConfig, +) + + +class TestWhisperTransformRequestResponseFormat: + def _transform(self, optional_params: dict) -> dict: + config = OpenAIWhisperAudioTranscriptionConfig() + audio_file = io.BytesIO(b"fake audio") + audio_file.name = "test.wav" + result = config.transform_audio_transcription_request( + model="whisper-1", + audio_file=audio_file, + optional_params=optional_params, + litellm_params={}, + ) + return result.data + + def test_defaults_to_verbose_json_when_unset(self): + """When response_format is not specified, default to verbose_json for cost calculation.""" + data = self._transform({}) + assert data["response_format"] == "verbose_json" + + def test_respects_explicit_json(self): + """When response_format='json' is set, do not override to verbose_json.""" + data = self._transform({"response_format": "json"}) + assert data["response_format"] == "json" + + def test_respects_explicit_text(self): + """When response_format='text' is set, do not override to verbose_json.""" + data = self._transform({"response_format": "text"}) + assert data["response_format"] == "text" + + def test_preserves_verbose_json_when_set(self): + """verbose_json explicitly set by the caller stays as-is.""" + data = self._transform({"response_format": "verbose_json"}) + assert data["response_format"] == "verbose_json" + + +class TestWhisperTransformResponse: + def _make_response(self, *, text: str, content_type: str, is_json: bool): + mock = MagicMock() + mock.headers = {"content-type": content_type} + if is_json: + mock.json.return_value = {"text": text} + else: + mock.json.side_effect = json.JSONDecodeError("", "", 0) + mock.text = text + return mock + + def test_parses_json_response(self): + """JSON body (verbose_json or json format) is parsed into TranscriptionResponse.""" + config = OpenAIWhisperAudioTranscriptionConfig() + result = config.transform_audio_transcription_response( + self._make_response( + text="Hello world", content_type="application/json", is_json=True + ) + ) + assert result.text == "Hello world" + + def test_parses_plain_text_response(self): + """Plain-text body (response_format=text) is returned as TranscriptionResponse without error.""" + config = OpenAIWhisperAudioTranscriptionConfig() + result = config.transform_audio_transcription_response( + self._make_response( + text="Four score and seven years ago", + content_type="text/plain", + is_json=False, + ) + ) + assert result.text == "Four score and seven years ago" + + def test_malformed_json_body_with_json_content_type_raises(self): + """A non-JSON body labelled application/json is a genuine upstream error, not a transcription.""" + config = OpenAIWhisperAudioTranscriptionConfig() + with pytest.raises(json.JSONDecodeError): + config.transform_audio_transcription_response( + self._make_response( + text="502 Bad Gateway", + content_type="application/json", + is_json=False, + ) + ) + + def test_json_content_type_match_is_case_insensitive(self): + """Media types are case-insensitive (RFC 7231), so a mixed-case application/json still re-raises.""" + config = OpenAIWhisperAudioTranscriptionConfig() + with pytest.raises(json.JSONDecodeError): + config.transform_audio_transcription_response( + self._make_response( + text="502 Bad Gateway", + content_type="Application/JSON; charset=utf-8", + is_json=False, + ) + ) diff --git a/tests/test_litellm/llms/openrouter/chat/test_openrouter_chat_transformation.py b/tests/test_litellm/llms/openrouter/chat/test_openrouter_chat_transformation.py index f102319d6bf..8d1129cc5da 100644 --- a/tests/test_litellm/llms/openrouter/chat/test_openrouter_chat_transformation.py +++ b/tests/test_litellm/llms/openrouter/chat/test_openrouter_chat_transformation.py @@ -553,3 +553,68 @@ def test_openrouter_non_reasoning_models_do_not_add_reasoning_effort(): ) assert "reasoning_effort" not in supported_params + + +def test_openrouter_reasoning_effort_max_maps_to_xhigh(): + """ + OpenRouter expects 'xhigh' instead of 'max' for reasoning_effort. + """ + config = OpenrouterConfig() + + result = config.map_openai_params( + non_default_params={"reasoning_effort": "max"}, + optional_params={}, + model="openrouter/deepseek/deepseek-r1", + drop_params=False, + ) + + assert result["reasoning_effort"] == "xhigh" + + +def test_openrouter_reasoning_effort_max_does_not_mutate_caller_dict(): + """ + map_openai_params must not mutate the caller-supplied non_default_params dict. + """ + config = OpenrouterConfig() + original_params = {"reasoning_effort": "max"} + + config.map_openai_params( + non_default_params=original_params, + optional_params={}, + model="openrouter/deepseek/deepseek-r1", + drop_params=False, + ) + + assert original_params["reasoning_effort"] == "max" + + +def test_openrouter_reasoning_effort_xhigh_passes_through(): + """ + reasoning_effort='xhigh' should be forwarded unchanged. + """ + config = OpenrouterConfig() + + result = config.map_openai_params( + non_default_params={"reasoning_effort": "xhigh"}, + optional_params={}, + model="openrouter/deepseek/deepseek-r1", + drop_params=False, + ) + + assert result["reasoning_effort"] == "xhigh" + + +def test_openrouter_reasoning_effort_high_passes_through(): + """ + Non-max reasoning_effort values should be forwarded unchanged. + """ + config = OpenrouterConfig() + + result = config.map_openai_params( + non_default_params={"reasoning_effort": "high"}, + optional_params={}, + model="openrouter/deepseek/deepseek-r1", + drop_params=False, + ) + + assert result["reasoning_effort"] == "high" diff --git a/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py b/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py index d408f55c004..e2d1ab72c5e 100644 --- a/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py +++ b/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py @@ -1,7 +1,7 @@ """ Test file for Perplexity cost calculator functionality. -Tests the cost calculation for Perplexity models including citation tokens, +Tests the cost calculation for Perplexity models including citation tokens, search queries, and reasoning tokens. """ @@ -21,7 +21,11 @@ from litellm.cost_calculator import completion_cost, cost_per_token from litellm.llms.perplexity.cost_calculator import ( cost_per_token as perplexity_cost_per_token, ) -from litellm.types.utils import Usage, PromptTokensDetailsWrapper +from litellm.types.utils import ( + CompletionTokensDetailsWrapper, + Usage, + PromptTokensDetailsWrapper, +) from litellm.utils import get_model_info @@ -135,13 +139,14 @@ class TestPerplexityCostCalculator: model="sonar-deep-research", usage=usage ) - # Expected costs: - # Input: 100 tokens * $2e-6 = $0.0002 - # Output: 50 tokens * $8e-6 = $0.0004 - # Reasoning: 20 tokens * $3e-6 = $0.00006 - # Total completion cost: $0.00046 + # `completion_tokens` includes `reasoning_tokens` per the OpenAI/Perplexity + # convention codified in PR #18607. Non-reasoning portion = 50 - 20 = 30. + # Input: 100 tokens * $2e-6 = $0.0002 + # Output (text): 30 tokens * $8e-6 = $0.00024 + # Reasoning: 20 tokens * $3e-6 = $0.00006 + # Total completion cost = $0.0003 expected_prompt_cost = 100 * 2e-6 - expected_completion_cost = (50 * 8e-6) + (20 * 3e-6) + expected_completion_cost = ((50 - 20) * 8e-6) + (20 * 3e-6) assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6) assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-6) @@ -159,13 +164,10 @@ class TestPerplexityCostCalculator: model="sonar-deep-research", usage=usage ) - # Expected costs: - # Input: 100 tokens * $2e-6 = $0.0002 - # Output: 50 tokens * $8e-6 = $0.0004 - # Reasoning: 20 tokens * $3e-6 = $0.00006 - # Total completion cost: $0.00046 + # Same convention as the direct-attribute case above; reasoning is a subset of + # completion_tokens, so non-reasoning portion = 50 - 20 = 30. expected_prompt_cost = 100 * 2e-6 - expected_completion_cost = (50 * 8e-6) + (20 * 3e-6) + expected_completion_cost = ((50 - 20) * 8e-6) + (20 * 3e-6) assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6) assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-6) @@ -187,16 +189,16 @@ class TestPerplexityCostCalculator: model="sonar-deep-research", usage=usage ) - # Expected costs: - # Input: 100 tokens * $2e-6 = $0.0002 - # Citation: 30 tokens * $2e-6 = $0.00006 - # Total prompt cost: $0.00026 - # Output: 50 tokens * $8e-6 = $0.0004 - # Reasoning: 15 tokens * $3e-6 = $0.000045 - # Search: 2 queries * ($0.005 / 1000) = $0.00001 - # Total completion cost: $0.000455 + # Expected costs (reasoning is a subset of completion_tokens): + # Input: 100 tokens * $2e-6 = $0.0002 + # Citation: 30 tokens * $2e-6 = $0.00006 + # Total prompt cost = $0.00026 + # Output (text): (50 - 15) tokens * $8e-6 = $0.00028 + # Reasoning: 15 tokens * $3e-6 = $0.000045 + # Search: 2 queries * ($0.005 / 1000) = $0.00001 + # Total completion cost = $0.000335 expected_prompt_cost = (100 * 2e-6) + (30 * 2e-6) - expected_completion_cost = (50 * 8e-6) + (15 * 3e-6) + (2 / 1000 * 0.005) + expected_completion_cost = ((50 - 15) * 8e-6) + (15 * 3e-6) + (2 / 1000 * 0.005) assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6) assert math.isclose(completion_cost, expected_completion_cost, rel_tol=1e-6) @@ -306,11 +308,11 @@ class TestPerplexityCostCalculator: completion_response=response, custom_llm_provider="perplexity" ) - # Calculate expected total cost + # Calculate expected total cost (reasoning is a subset of completion_tokens) expected_prompt_cost = (100 * 2e-6) + (15 * 2e-6) # Input + citation expected_completion_cost = ( - (50 * 8e-6) + (10 * 3e-6) + (1 / 1000 * 0.005) - ) # Output + reasoning + search + ((50 - 10) * 8e-6) + (10 * 3e-6) + (1 / 1000 * 0.005) + ) # Output (text) + reasoning + search expected_total = expected_prompt_cost + expected_completion_cost assert math.isclose(total_cost, expected_total, rel_tol=1e-6) @@ -353,10 +355,13 @@ class TestPerplexityCostCalculator: model="sonar-deep-research", usage=usage ) - # Calculate expected costs + # Calculate expected costs. `completion_tokens` includes `reasoning_tokens`, + # so non-reasoning portion = 50 - reasoning_tokens. expected_prompt_cost = (100 * 2e-6) + (citation_tokens * 2e-6) expected_completion_cost = ( - (50 * 8e-6) + (reasoning_tokens * 3e-6) + (search_queries / 1000 * 0.005) + ((50 - reasoning_tokens) * 8e-6) + + (reasoning_tokens * 3e-6) + + (search_queries / 1000 * 0.005) ) assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6) @@ -413,3 +418,36 @@ class TestPerplexityCostCalculator: assert math.isclose(prompt_cost, expected_prompt, rel_tol=1e-6) assert math.isclose(completion_cost, expected_completion, rel_tol=1e-6) + + def test_reasoning_tokens_not_double_billed(self): + """ + Regression: `completion_tokens` includes `reasoning_tokens` per the + OpenAI/Perplexity usage convention (codified for the central path in PR #18607). + When `output_cost_per_reasoning_token` is configured the manual fallback must + subtract reasoning from completion before applying the output rate so the + reasoning tokens are not billed at BOTH the output rate and the reasoning rate. + + Uses the exact usage shape produced by the live response fixture in + `tests/llm_translation/test_perplexity_reasoning.py`. + """ + usage = Usage( + prompt_tokens=9, + completion_tokens=20, + total_tokens=29, + completion_tokens_details=CompletionTokensDetailsWrapper( + reasoning_tokens=15 + ), + ) + + prompt_cost, completion_cost = perplexity_cost_per_token( + model="sonar-deep-research", usage=usage + ) + + # sonar-deep-research rates: input 2e-6, output 8e-6, reasoning 3e-6. + # Non-reasoning portion of the 20 completion tokens = 20 - 15 = 5. + # Pre-fix this asserted 20 * 8e-6 + 15 * 3e-6 = 2.05e-4 (a 2.16x overcharge). + expected_prompt = 9 * 2e-6 + expected_completion = (20 - 15) * 8e-6 + 15 * 3e-6 + + assert math.isclose(prompt_cost, expected_prompt, rel_tol=1e-9) + assert math.isclose(completion_cost, expected_completion, rel_tol=1e-9) diff --git a/tests/test_litellm/llms/perplexity/test_perplexity_integration.py b/tests/test_litellm/llms/perplexity/test_perplexity_integration.py index 1b03fd7df88..e59fbc9f272 100644 --- a/tests/test_litellm/llms/perplexity/test_perplexity_integration.py +++ b/tests/test_litellm/llms/perplexity/test_perplexity_integration.py @@ -104,12 +104,10 @@ class TestPerplexityIntegration: ) citation_tokens = citation_chars // 4 - expected_prompt_cost = (100 * 2e-6) + ( - citation_tokens * 2e-6 - ) # Input + citation + expected_prompt_cost = (100 * 2e-6) + (citation_tokens * 2e-6) expected_completion_cost = ( - (50 * 8e-6) + (10 * 3e-6) + (2 / 1000 * 0.005) - ) # Output + reasoning + search + ((50 - 10) * 8e-6) + (10 * 3e-6) + (2 / 1000 * 0.005) + ) expected_total = expected_prompt_cost + expected_completion_cost assert math.isclose(total_cost, expected_total, rel_tol=1e-6) @@ -152,11 +150,10 @@ class TestPerplexityIntegration: usage_object=usage, ) - # Calculate expected costs - expected_prompt_cost = (200 * 2e-6) + (40 * 2e-6) # Input + citation + expected_prompt_cost = (200 * 2e-6) + (40 * 2e-6) expected_completion_cost = ( - (100 * 8e-6) + (25 * 3e-6) + (3 / 1000 * 0.005) - ) # Output + reasoning + search + ((100 - 25) * 8e-6) + (25 * 3e-6) + (3 / 1000 * 0.005) + ) assert math.isclose(prompt_cost, expected_prompt_cost, rel_tol=1e-6) assert math.isclose(completion_cost_val, expected_completion_cost, rel_tol=1e-6) @@ -263,15 +260,14 @@ class TestPerplexityIntegration: custom_llm_provider="perplexity", ) - # Calculate expected cost - expected_prompt_cost = (50000 * 2e-6) + (5000 * 2e-6) # $0.11 + expected_prompt_cost = (50000 * 2e-6) + (5000 * 2e-6) expected_completion_cost = ( - (25000 * 8e-6) + (10000 * 3e-6) + (100 / 1000 * 0.005) - ) # $0.23 - expected_total = expected_prompt_cost + expected_completion_cost # $0.34 + ((25000 - 10000) * 8e-6) + (10000 * 3e-6) + (100 / 1000 * 0.005) + ) + expected_total = expected_prompt_cost + expected_completion_cost assert math.isclose(total_cost, expected_total, rel_tol=1e-6) - assert total_cost > 0.3 # Sanity check for high-volume scenario + assert total_cost > 0.25 def test_transformation_preserves_existing_usage_fields(self): """Test that transformation doesn't overwrite existing standard usage fields.""" diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index 671d7355e8f..1a2d0d86810 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -4996,3 +4996,146 @@ def test_mid_stream_429_error_raises_during_iteration(): # Verify: 429 error is properly raised assert exc_info.value.status_code == 429 assert "RESOURCE_EXHAUSTED" in str(exc_info.value.message) + + +class TestModelResponseIteratorCleanup: + def _make_logging_obj(self): + from unittest.mock import Mock + + obj = Mock() + obj.optional_params = {} + return obj + + def test_aclose_closes_iterator_and_response(self): + import asyncio + from unittest.mock import AsyncMock, MagicMock + + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + ModelResponseIterator, + ) + + mock_response = MagicMock() + mock_response.aclose = AsyncMock() + + mock_iterator = MagicMock() + mock_iterator.aclose = AsyncMock() + + iterator = ModelResponseIterator( + streaming_response=MagicMock(), + sync_stream=False, + logging_obj=self._make_logging_obj(), + response=mock_response, + ) + iterator.async_response_iterator = mock_iterator + + asyncio.run(iterator.aclose()) + + mock_iterator.aclose.assert_awaited_once() + mock_response.aclose.assert_awaited_once() + + def test_close_closes_iterator_and_response(self): + from unittest.mock import MagicMock + + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + ModelResponseIterator, + ) + + mock_response = MagicMock() + mock_iterator = MagicMock() + + iterator = ModelResponseIterator( + streaming_response=MagicMock(), + sync_stream=True, + logging_obj=self._make_logging_obj(), + response=mock_response, + ) + iterator.response_iterator = mock_iterator + + iterator.close() + + mock_iterator.close.assert_called_once() + mock_response.close.assert_called_once() + + def test_aclose_without_response_does_not_raise(self): + import asyncio + from unittest.mock import AsyncMock, MagicMock + + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + ModelResponseIterator, + ) + + mock_iterator = MagicMock() + mock_iterator.aclose = AsyncMock() + + iterator = ModelResponseIterator( + streaming_response=MagicMock(), + sync_stream=False, + logging_obj=self._make_logging_obj(), + ) + iterator.async_response_iterator = mock_iterator + + asyncio.run(iterator.aclose()) + + mock_iterator.aclose.assert_awaited_once() + + def test_aclose_tolerates_iterator_error(self): + import asyncio + from unittest.mock import AsyncMock, MagicMock + + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + ModelResponseIterator, + ) + + mock_response = MagicMock() + mock_response.aclose = AsyncMock() + + mock_iterator = MagicMock() + mock_iterator.aclose = AsyncMock(side_effect=RuntimeError("transport error")) + + iterator = ModelResponseIterator( + streaming_response=MagicMock(), + sync_stream=False, + logging_obj=self._make_logging_obj(), + response=mock_response, + ) + iterator.async_response_iterator = mock_iterator + + asyncio.run(iterator.aclose()) + + mock_response.aclose.assert_awaited_once() + + def test_custom_stream_wrapper_aclose_triggers_model_response_iterator_aclose(self): + """CustomStreamWrapper.aclose() must propagate to ModelResponseIterator.aclose().""" + import asyncio + from unittest.mock import AsyncMock, MagicMock + + from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + ModelResponseIterator, + ) + + mock_response = MagicMock() + mock_response.aclose = AsyncMock() + + mock_iterator = MagicMock() + mock_iterator.aclose = AsyncMock() + + model_response_iter = ModelResponseIterator( + streaming_response=MagicMock(), + sync_stream=False, + logging_obj=self._make_logging_obj(), + response=mock_response, + ) + model_response_iter.async_response_iterator = mock_iterator + + wrapper = CustomStreamWrapper( + completion_stream=model_response_iter, + model="gemini-2.0-flash", + custom_llm_provider="vertex_ai", + logging_obj=MagicMock(), + ) + + asyncio.run(wrapper.aclose()) + + mock_iterator.aclose.assert_awaited_once() + mock_response.aclose.assert_awaited_once() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py index 1c31f437363..c86ae966f21 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server.py @@ -5179,6 +5179,12 @@ async def test_list_tools_with_legacy_db_m2m_server_resolves_oauth2_flow(): side_effect=lambda update: MCPServer( server_id=legacy_server.server_id, name=legacy_server.name, + # Carry alias/server_name forward so get_server_prefix resolves to + # "legacy_m2m" (not the server_id) when the request scope filter + # matches by alias. Without these, the filter relied on the now- + # removed silent fail-open fallback. + alias=legacy_server.alias, + server_name=legacy_server.server_name, transport=MCPTransport.http, auth_type=legacy_server.auth_type, oauth2_flow=update.get("oauth2_flow", legacy_server.oauth2_flow), @@ -6083,3 +6089,207 @@ async def test_execute_mcp_tool_rest_prefix_retry_resolution_still_enforces_serv assert exc_info.value.status_code == 403 assert exc_info.value.detail["error"] == "tool_server_mismatch" + + +# --------------------------------------------------------------------------- +# Regression tests for _get_allowed_mcp_servers_from_mcp_server_names +# +# Prior to the fail-closed fix, an unresolved scope filter (path- or +# header-derived) silently returned the caller's full allowed-server set, +# which made URL/header namespacing appear to work when it did not. +# --------------------------------------------------------------------------- + + +def _make_mcp_server_for_scope_filter(server_id: str, alias: str) -> MCPServer: + return MCPServer( + server_id=server_id, + name=alias, + alias=alias, + server_name=alias, + url=f"https://{alias}.test/mcp", + transport=MCPTransport.http, + mcp_info={"server_name": alias}, + ) + + +@pytest.mark.asyncio +async def test_get_allowed_mcp_servers_from_mcp_server_names_unknown_name_fails_closed(): + """ + Bug fix: requesting an unknown server name (e.g. ``/mcp//``) must + NOT silently fall back to the caller's full allowed-server set. + """ + try: + from litellm.proxy._experimental.mcp_server.server import ( + _get_allowed_mcp_servers_from_mcp_server_names, + ) + except ImportError: + pytest.skip("MCP server not available") + + allowed = [ + _make_mcp_server_for_scope_filter("id-a", "alpha"), + _make_mcp_server_for_scope_filter("id-b", "beta"), + ] + + with patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." + "MCPRequestHandler._get_mcp_servers_from_access_groups", + new_callable=AsyncMock, + return_value=[], + ): + result = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=["does-not-exist"], + allowed_mcp_servers=allowed, + ) + + assert result == [] + + +@pytest.mark.asyncio +async def test_get_allowed_mcp_servers_from_mcp_server_names_none_returns_all(): + """ + Regression: ``mcp_servers=None`` (no scope filter requested) must still + return the full allowed-server set. This is the legitimate "no scoping" + path that the fail-closed fix must not break. + """ + try: + from litellm.proxy._experimental.mcp_server.server import ( + _get_allowed_mcp_servers_from_mcp_server_names, + ) + except ImportError: + pytest.skip("MCP server not available") + + allowed = [ + _make_mcp_server_for_scope_filter("id-a", "alpha"), + _make_mcp_server_for_scope_filter("id-b", "beta"), + ] + + result = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=None, + allowed_mcp_servers=allowed, + ) + + assert {s.server_id for s in result} == {"id-a", "id-b"} + + +@pytest.mark.asyncio +async def test_get_allowed_mcp_servers_from_mcp_server_names_known_alias_returns_match(): + """ + Regression: a known server alias must still resolve to exactly that + server. Guards against the fix accidentally narrowing the happy path. + """ + try: + from litellm.proxy._experimental.mcp_server.server import ( + _get_allowed_mcp_servers_from_mcp_server_names, + ) + except ImportError: + pytest.skip("MCP server not available") + + allowed = [ + _make_mcp_server_for_scope_filter("id-a", "alpha"), + _make_mcp_server_for_scope_filter("id-b", "beta"), + ] + + with patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." + "MCPRequestHandler._get_mcp_servers_from_access_groups", + new_callable=AsyncMock, + return_value=[], + ): + result = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=["alpha"], + allowed_mcp_servers=allowed, + ) + + assert [s.server_id for s in result] == ["id-a"] + + +@pytest.mark.asyncio +async def test_get_allowed_mcp_servers_from_mcp_server_names_mixed_known_and_unknown(): + """ + Mixed scope (one valid + one unknown) returns only the resolved server, + not the full allowed set. Confirms the fail-closed branch only fires + when NOTHING resolves. + """ + try: + from litellm.proxy._experimental.mcp_server.server import ( + _get_allowed_mcp_servers_from_mcp_server_names, + ) + except ImportError: + pytest.skip("MCP server not available") + + allowed = [ + _make_mcp_server_for_scope_filter("id-a", "alpha"), + _make_mcp_server_for_scope_filter("id-b", "beta"), + ] + + with patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." + "MCPRequestHandler._get_mcp_servers_from_access_groups", + new_callable=AsyncMock, + return_value=[], + ): + result = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=["alpha", "does-not-exist"], + allowed_mcp_servers=allowed, + ) + + assert [s.server_id for s in result] == ["id-a"] + + +@pytest.mark.asyncio +async def test_get_allowed_mcp_servers_from_mcp_server_names_access_group_resolves(): + """ + Regression: when a requested name is not a server alias but IS an access + group, it must still resolve to the underlying servers (not be treated + as unresolved). + """ + try: + from litellm.proxy._experimental.mcp_server.server import ( + _get_allowed_mcp_servers_from_mcp_server_names, + ) + except ImportError: + pytest.skip("MCP server not available") + + allowed = [ + _make_mcp_server_for_scope_filter("id-a", "alpha"), + _make_mcp_server_for_scope_filter("id-b", "beta"), + ] + + with patch( + "litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp." + "MCPRequestHandler._get_mcp_servers_from_access_groups", + new_callable=AsyncMock, + return_value=["id-b"], + ): + result = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=["group-name"], + allowed_mcp_servers=allowed, + ) + + assert [s.server_id for s in result] == ["id-b"] + + +@pytest.mark.asyncio +async def test_get_allowed_mcp_servers_from_mcp_server_names_empty_list_fails_closed(): + """ + Edge case: ``mcp_servers=[]`` (explicit empty scope) is still an + explicit filter request. Fail closed rather than returning everything. + """ + try: + from litellm.proxy._experimental.mcp_server.server import ( + _get_allowed_mcp_servers_from_mcp_server_names, + ) + except ImportError: + pytest.skip("MCP server not available") + + allowed = [ + _make_mcp_server_for_scope_filter("id-a", "alpha"), + _make_mcp_server_for_scope_filter("id-b", "beta"), + ] + + result = await _get_allowed_mcp_servers_from_mcp_server_names( + mcp_servers=[], + allowed_mcp_servers=allowed, + ) + + assert result == [] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_identity_env.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_identity_env.py new file mode 100644 index 00000000000..ac7082c2668 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_identity_env.py @@ -0,0 +1,73 @@ +"""Regression tests for the configurable MCP gateway identity. + +``LITELLM_MCP_SERVER_NAME`` and ``LITELLM_MCP_SERVER_DESCRIPTION`` are read from +the environment at import time in +``litellm.proxy._experimental.mcp_server.utils`` and must flow through to every +consumer, including the well-known registry entry built in +``mcp_management_endpoints``. The env values are reloaded into the modules and +restored afterwards so the override does not leak into other tests. +""" + +import contextlib +import importlib +import os + +import pytest + +pytest.importorskip("mcp") + +UTILS_MODULE = "litellm.proxy._experimental.mcp_server.utils" +MGMT_MODULE = "litellm.proxy.management_endpoints.mcp_management_endpoints" + + +@contextlib.contextmanager +def _env_and_reload(**env): + saved = {key: os.environ.get(key) for key in env} + + def _apply_env(values): + for key, value in values.items(): + if value is None: + os.environ.pop(key, None) + else: + os.environ[key] = value + + def _reload(): + utils = importlib.reload(importlib.import_module(UTILS_MODULE)) + mgmt = importlib.reload(importlib.import_module(MGMT_MODULE)) + return utils, mgmt + + try: + _apply_env(env) + yield _reload() + finally: + _apply_env(saved) + _reload() + + +def test_defaults_used_when_env_unset(): + with _env_and_reload( + LITELLM_MCP_SERVER_NAME=None, LITELLM_MCP_SERVER_DESCRIPTION=None + ) as (utils, _mgmt): + assert utils.LITELLM_MCP_SERVER_NAME == "litellm-mcp-server" + assert utils.LITELLM_MCP_SERVER_DESCRIPTION == "MCP Server for LiteLLM" + + +def test_env_overrides_server_identity(): + with _env_and_reload( + LITELLM_MCP_SERVER_NAME="acme-gateway", + LITELLM_MCP_SERVER_DESCRIPTION="Acme internal MCP gateway", + ) as (utils, _mgmt): + assert utils.LITELLM_MCP_SERVER_NAME == "acme-gateway" + assert utils.LITELLM_MCP_SERVER_DESCRIPTION == "Acme internal MCP gateway" + + +def test_env_override_propagates_to_registry_entry(): + with _env_and_reload( + LITELLM_MCP_SERVER_NAME="acme-gateway", + LITELLM_MCP_SERVER_DESCRIPTION="Acme internal MCP gateway", + ) as (_utils, mgmt): + entry = mgmt._build_builtin_registry_entry("http://localhost:4000") + + assert entry["name"] == "acme-gateway" + assert entry["title"] == "acme-gateway" + assert entry["description"] == "Acme internal MCP gateway" diff --git a/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py b/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py index f6189382d74..a4da4587b7f 100644 --- a/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/anthropic_endpoints/test_endpoints.py @@ -86,3 +86,96 @@ class TestEventLoggingBatchEndpoint: assert response.status_code == 200 assert response.json() == {"status": "ok"} + + +class TestStripTotalTokens(unittest.TestCase): + """Cover ``_strip_total_tokens_from_anthropic_response``. + + The Anthropic /v1/messages spec does not define ``usage.total_tokens``. + LiteLLM injects it internally; the helper must remove it from the wire + response so the non-streaming path matches the streaming SSE shape and + direct Anthropic API responses. + """ + + def test_strips_total_tokens_when_present(self): + from litellm.proxy.anthropic_endpoints.endpoints import ( + _strip_total_tokens_from_anthropic_response, + ) + + response = { + "id": "msg_123", + "usage": { + "input_tokens": 100, + "output_tokens": 50, + "total_tokens": 150, + "cache_read_input_tokens": 0, + "cache_creation_input_tokens": 0, + }, + } + _strip_total_tokens_from_anthropic_response(response) + assert "total_tokens" not in response["usage"] + assert response["usage"]["input_tokens"] == 100 + assert response["usage"]["output_tokens"] == 50 + assert response["usage"]["cache_read_input_tokens"] == 0 + + def test_no_op_when_total_tokens_absent(self): + from litellm.proxy.anthropic_endpoints.endpoints import ( + _strip_total_tokens_from_anthropic_response, + ) + + response = {"usage": {"input_tokens": 100, "output_tokens": 50}} + _strip_total_tokens_from_anthropic_response(response) + assert response["usage"] == {"input_tokens": 100, "output_tokens": 50} + + def test_no_op_when_usage_missing(self): + from litellm.proxy.anthropic_endpoints.endpoints import ( + _strip_total_tokens_from_anthropic_response, + ) + + response = {"id": "msg_123"} + _strip_total_tokens_from_anthropic_response(response) + assert response == {"id": "msg_123"} + + def test_no_op_on_non_dict_response(self): + from litellm.proxy.anthropic_endpoints.endpoints import ( + _strip_total_tokens_from_anthropic_response, + ) + + # Streaming responses (StreamingResponse, async iterators) are not dicts. + # The helper must not raise or attempt to mutate them. + for value in (None, "stream", 42, [{"usage": {"total_tokens": 1}}]): + _strip_total_tokens_from_anthropic_response(value) # no raise + + def test_strips_total_tokens_on_pydantic_model_with_dict_usage(self): + """Greptile P1 on #30382: helper must not silently no-op when the + response is a Pydantic-shaped object whose `usage` attribute is a + plain dict (the common case for objects wrapping raw upstream JSON). + """ + from types import SimpleNamespace + + from litellm.proxy.anthropic_endpoints.endpoints import ( + _strip_total_tokens_from_anthropic_response, + ) + + # SimpleNamespace mimics the .usage attribute access pattern; the + # helper's contract: if .usage is dict-shaped, strip total_tokens. + response = SimpleNamespace( + usage={"input_tokens": 100, "output_tokens": 50, "total_tokens": 150} + ) + _strip_total_tokens_from_anthropic_response(response) + assert "total_tokens" not in response.usage + assert response.usage == {"input_tokens": 100, "output_tokens": 50} + + +class TestStripTotalTokensFeatureFlag(unittest.TestCase): + """The strip is gated behind `litellm.strip_anthropic_total_tokens`. + + Default off (backward compat). Greptile P1 on #30382 required a + user-controlled flag so existing clients reading the LiteLLM-shaped + `usage.total_tokens` continue to work after this PR lands. + """ + + def test_flag_defaults_off(self): + import litellm + + assert litellm.strip_anthropic_total_tokens is False diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 32b597376b4..e652c109987 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -13,6 +13,8 @@ from litellm.proxy.auth.auth_utils import ( _get_customer_id_from_standard_headers, abbreviate_api_key, check_complete_credentials, + custom_auth_common_checks_warning, + warn_once_if_custom_auth_skips_common_checks, get_end_user_id_from_request_body, get_key_mcp_rpm_limit, get_key_model_rpm_limit, @@ -25,6 +27,78 @@ from litellm.proxy.auth.auth_utils import ( ) +class TestCustomAuthCommonChecksWarning: + """custom_auth_common_checks_warning only warns when custom auth is configured + and the common-checks opt-in is off, since that is the only state where + project/team enforcement silently does nothing.""" + + def test_warns_when_custom_auth_configured_and_checks_off(self): + warning = custom_auth_common_checks_warning( + custom_auth_configured=True, + run_common_checks=False, + ) + assert warning is not None + assert "custom_auth_run_common_checks: true" in warning + assert "https://docs.litellm.ai/docs/proxy/custom_auth" in warning + + def test_no_warning_when_common_checks_enabled(self): + assert ( + custom_auth_common_checks_warning( + custom_auth_configured=True, + run_common_checks=True, + ) + is None + ) + + def test_no_warning_when_custom_auth_not_configured(self): + assert ( + custom_auth_common_checks_warning( + custom_auth_configured=False, + run_common_checks=False, + ) + is None + ) + assert ( + custom_auth_common_checks_warning( + custom_auth_configured=False, + run_common_checks=True, + ) + is None + ) + + +class TestWarnOnceIfCustomAuthSkipsCommonChecks: + """The startup warning must fire at most once per process, since load_config + re-runs on hot-reload / config refresh and would otherwise spam the log.""" + + @pytest.fixture(autouse=True) + def _reset_sentinel(self, monkeypatch): + monkeypatch.setattr( + "litellm.proxy.auth.auth_utils._custom_auth_common_checks_warning_emitted", + False, + ) + + def test_warns_only_once_across_repeated_calls(self): + logger = MagicMock() + for _ in range(3): + warn_once_if_custom_auth_skips_common_checks( + custom_auth_configured=True, + run_common_checks=False, + logger=logger, + ) + assert logger.warning.call_count == 1 + assert "custom_auth_run_common_checks" in logger.warning.call_args[0][0] + + def test_does_not_warn_when_common_checks_enabled(self): + logger = MagicMock() + warn_once_if_custom_auth_skips_common_checks( + custom_auth_configured=True, + run_common_checks=True, + logger=logger, + ) + assert logger.warning.call_count == 0 + + class TestGetKeyModelRpmLimit: """Tests for get_key_model_rpm_limit function.""" @@ -530,6 +604,59 @@ def test_get_model_from_request_handles_managed_id_decoder_failures(): ) +@pytest.mark.parametrize( + "route", + [ + "/realtime/client_secrets", + "/v1/realtime/client_secrets", + "/openai/v1/realtime/client_secrets", + "/realtime/calls", + "/v1/realtime/calls", + "/openai/v1/realtime/calls", + ], +) +def test_get_model_from_request_extracts_realtime_session_model(route): + """The effective realtime model lives in ``session.model`` (not the + top-level ``model``). It must be surfaced so can_key_call_model() can + validate the model a restricted key is actually requesting. + + Regression test for the model-access bypass on the GA Realtime WebRTC + HTTP routes (https://github.com/BerriAI/litellm/issues/29923). + """ + assert ( + get_model_from_request( + request_data={"session": {"type": "realtime", "model": "gpt-realtime"}}, + route=route, + ) + == "gpt-realtime" + ) + + +def test_get_model_from_request_realtime_includes_top_level_and_session_model(): + """When both top-level and session model are present, both are returned so + neither path can smuggle a disallowed model past the model-access check.""" + models = get_model_from_request( + request_data={ + "model": "gpt-4o-realtime-preview", + "session": {"type": "realtime", "model": "gpt-realtime"}, + }, + route="/v1/realtime/client_secrets", + ) + assert models == ["gpt-4o-realtime-preview", "gpt-realtime"] + + +def test_get_model_from_request_ignores_session_model_on_non_realtime_routes(): + """A nested ``session.model`` must not leak into model resolution for + unrelated routes.""" + assert ( + get_model_from_request( + request_data={"session": {"type": "realtime", "model": "gpt-realtime"}}, + route="/v1/chat/completions", + ) + is None + ) + + def test_abbreviate_api_key(): assert abbreviate_api_key("sk-test-1234") == "sk-...1234" diff --git a/tests/test_litellm/proxy/auth/test_model_checks.py b/tests/test_litellm/proxy/auth/test_model_checks.py index f38ac5c2000..8d686900ea6 100644 --- a/tests/test_litellm/proxy/auth/test_model_checks.py +++ b/tests/test_litellm/proxy/auth/test_model_checks.py @@ -388,6 +388,60 @@ def test_wildcard_credential_hydration_preserves_deployment_params( } +def test_wildcard_custom_prefix_does_not_stack_provider_prefix(monkeypatch): + """Regression test for #30358. + + A wildcard with a custom prefix (e.g. ``ollama_server1/*`` to distinguish multiple Ollama + instances) must not stack the provider's own prefix onto the expanded model ids. The expanded + ids should be ``ollama_server1/gemma3:1b`` rather than ``ollama_server1/ollama/gemma3:1b``. + """ + from litellm.proxy.auth import model_checks + from litellm.proxy.auth.model_checks import get_known_models_from_wildcard + from litellm.types.router import LiteLLM_Params + + monkeypatch.setattr( + model_checks, + "get_provider_models", + lambda provider, litellm_params=None: ["ollama/gemma3:1b", "ollama/llama3:8b"], + ) + + result = get_known_models_from_wildcard( + wildcard_model="ollama_server1/*", + litellm_params=LiteLLM_Params( + model="ollama_chat/*", custom_llm_provider="ollama_chat" + ), + ) + + assert result == ["ollama_server1/gemma3:1b", "ollama_server1/llama3:8b"] + + +def test_wildcard_custom_prefix_keeps_org_segment_for_non_provider_first_segment( + monkeypatch, +): + """Only a known provider prefix should be stripped before re-prefixing. + + If ``get_provider_models`` returns ids whose first segment is an org rather than a litellm + provider (e.g. ``meta-llama/Llama-3-8B``), stripping the first slash segment would drop the + org and produce an uncallable id. The org segment must be preserved. + """ + from litellm.proxy.auth import model_checks + from litellm.proxy.auth.model_checks import get_known_models_from_wildcard + from litellm.types.router import LiteLLM_Params + + monkeypatch.setattr( + model_checks, + "get_provider_models", + lambda provider, litellm_params=None: ["meta-llama/Llama-3-8B"], + ) + + result = get_known_models_from_wildcard( + wildcard_model="my_hf/*", + litellm_params=LiteLLM_Params(model="huggingface/*", custom_llm_provider="huggingface"), + ) + + assert result == ["my_hf/meta-llama/Llama-3-8B"] + + def test_wildcard_credential_hydration_preserves_missing_credential_name( monkeypatch, ): diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index 07b04961205..7a4597c4e02 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -460,6 +460,32 @@ def test_mcp_inference_routes_classified_as_llm_api(route): assert RouteChecks.is_management_route(route=route) is False +@pytest.mark.parametrize( + "route", + [ + "/realtime/client_secrets", + "/v1/realtime/client_secrets", + "/openai/v1/realtime/client_secrets", + "/realtime/calls", + "/v1/realtime/calls", + "/openai/v1/realtime/calls", + "/realtime/transcription_sessions", + "/v1/realtime/transcription_sessions", + "/openai/v1/realtime/transcription_sessions", + ], +) +def test_realtime_webrtc_http_routes_classified_as_llm_api(route): + """GA Realtime WebRTC HTTP routes must be classified as LLM API routes so + non-admin virtual keys can call them instead of hitting the admin-only + 401 branch in non_proxy_admin_allowed_routes_check. + + Regression test for https://github.com/BerriAI/litellm/issues/29923 + """ + + assert RouteChecks.is_llm_api_route(route=route) is True + assert RouteChecks.is_management_route(route=route) is False + + def test_virtual_key_allowed_routes_with_litellm_routes_member_name_denied(): """Test that virtual key is denied when route is not in the allowed LiteLLMRoutes group""" diff --git a/tests/test_litellm/proxy/db/db_transaction_queue/test_pod_lock_manager.py b/tests/test_litellm/proxy/db/db_transaction_queue/test_pod_lock_manager.py index 27fe9202276..f2745052faa 100644 --- a/tests/test_litellm/proxy/db/db_transaction_queue/test_pod_lock_manager.py +++ b/tests/test_litellm/proxy/db/db_transaction_queue/test_pod_lock_manager.py @@ -327,7 +327,7 @@ async def test_release_lock_uses_atomic_compare_delete_script_when_available( PodLockManager._COMPARE_AND_DELETE_LOCK_SCRIPT ) script_callable.assert_called_once_with( - keys=[lock_key], args=[pod_lock_manager.pod_id] + keys=[lock_key], args=[json.dumps(pod_lock_manager.pod_id)] ) mock_redis.async_get_cache.assert_not_called() mock_redis.async_delete_cache.assert_not_called() @@ -364,6 +364,78 @@ async def test_release_lock_lua_path_emits_released_event(pod_lock_manager, mock ) +class FakeRedisLockStore: + """ + Minimal stand-in that mirrors how RedisCache actually stores values: + async_set_cache JSON-encodes the value, and the compare-and-delete Lua + script compares against the raw stored bytes. This is what exposes the + quoted-vs-raw mismatch that a value-agnostic mock cannot catch. + """ + + def __init__(self): + self.store: dict = {} + + async def async_set_cache(self, key, value, nx=False, ttl=None, **kwargs): + if nx and key in self.store: + return None + self.store[key] = json.dumps(value) + return True + + async def async_get_cache(self, key, **kwargs): + raw = self.store.get(key) + return json.loads(raw) if raw is not None else None + + async def async_delete_cache(self, key, **kwargs): + return 1 if self.store.pop(key, None) is not None else 0 + + def async_register_script(self, script): + async def _run(keys, args): + key = keys[0] + if self.store.get(key) == args[0]: + del self.store[key] + return 1 + return 0 + + return _run + + +@pytest.mark.asyncio +async def test_release_lock_deletes_lock_held_by_same_pod(): + """ + Regression: acquire_lock stores the pod_id JSON-encoded, so release_lock's + Lua compare-and-delete must use the same encoding or the comparison never + matches and the lock leaks until its TTL expires (stalling the spend-update + drain and growing the Redis transaction buffers). + """ + redis = FakeRedisLockStore() + pod = PodLockManager(redis_cache=redis) + lock_key = PodLockManager.get_redis_lock_key("db_spend_update_job") + + acquired = await pod.acquire_lock(cronjob_id="db_spend_update_job") + assert acquired is True + assert lock_key in redis.store + + await pod.release_lock(cronjob_id="db_spend_update_job") + assert lock_key not in redis.store + + +@pytest.mark.asyncio +async def test_release_lock_preserves_lock_held_by_other_pod(): + """ + A pod must not release a lock currently held by a different pod, even with + the encoding fix in place. + """ + redis = FakeRedisLockStore() + holder = PodLockManager(redis_cache=redis) + other = PodLockManager(redis_cache=redis) + lock_key = PodLockManager.get_redis_lock_key("db_spend_update_job") + + assert await holder.acquire_lock(cronjob_id="db_spend_update_job") is True + + await other.release_lock(cronjob_id="db_spend_update_job") + assert redis.store.get(lock_key) == json.dumps(holder.pod_id) + + @pytest.mark.asyncio async def test_release_lock_falls_back_to_get_del_when_lua_execution_fails( pod_lock_manager, mock_redis diff --git a/tests/test_litellm/proxy/db/db_transaction_queue/test_redis_update_buffer.py b/tests/test_litellm/proxy/db/db_transaction_queue/test_redis_update_buffer.py index 0587e3bce1e..33372e7794a 100644 --- a/tests/test_litellm/proxy/db/db_transaction_queue/test_redis_update_buffer.py +++ b/tests/test_litellm/proxy/db/db_transaction_queue/test_redis_update_buffer.py @@ -1,7 +1,7 @@ import json import os import sys -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -11,7 +11,6 @@ sys.path.insert( from litellm.proxy.db.db_transaction_queue.redis_update_buffer import RedisUpdateBuffer from litellm.proxy.proxy_server import ProxyStartupEvent -from litellm.types.caching import RedisPipelineRpushOperation @pytest.fixture @@ -305,3 +304,73 @@ def test_validate_redis_transaction_buffer_passes_when_disabled(): general_settings={}, redis_usage_cache=None, ) + + +def test_get_transaction_buffer_redis_cache_builds_from_env(monkeypatch): + """ + When use_redis_transaction_buffer=true, a standalone RedisCache is built from + REDIS_* environment variables so the buffer works without a Redis cache backend. + """ + monkeypatch.setenv("REDIS_HOST", "localhost") + monkeypatch.setenv("REDIS_PORT", "6379") + + with patch("litellm.proxy.proxy_server.RedisCache") as mock_redis_cache: + result = ProxyStartupEvent._get_transaction_buffer_redis_cache( + general_settings={"use_redis_transaction_buffer": True}, + ) + + mock_redis_cache.assert_called_once() + assert mock_redis_cache.call_args.kwargs["host"] == "localhost" + assert result is mock_redis_cache.return_value + + +def test_get_transaction_buffer_redis_cache_none_when_disabled(): + """When use_redis_transaction_buffer is not enabled, no standalone cache is built.""" + result = ProxyStartupEvent._get_transaction_buffer_redis_cache( + general_settings={}, + ) + assert result is None + + +def test_get_transaction_buffer_redis_cache_none_without_redis_env(): + """ + When use_redis_transaction_buffer=true but no REDIS_* env vars are set, + no standalone cache is built (startup validation then raises the config error). + """ + with patch("litellm._redis._redis_kwargs_from_environment", return_value={}): + result = ProxyStartupEvent._get_transaction_buffer_redis_cache( + general_settings={"use_redis_transaction_buffer": True}, + ) + assert result is None + + +def test_get_transaction_buffer_redis_cache_none_without_host_or_url(): + """ + A REDIS_* var that is not a connection target (e.g. REDIS_SOCKET_TIMEOUT) must not + trigger a build. Without a host or url, get_redis_client raises, so return None and + let startup validation surface the config error instead of crashing. + """ + with patch( + "litellm._redis._redis_kwargs_from_environment", + return_value={"socket_timeout": 5.0}, + ): + result = ProxyStartupEvent._get_transaction_buffer_redis_cache( + general_settings={"use_redis_transaction_buffer": True}, + ) + assert result is None + + +def test_get_transaction_buffer_redis_cache_parses_string_flag(monkeypatch): + """ + use_redis_transaction_buffer accepts a string value (e.g. from env/YAML); "true" + is parsed to a bool before the standalone cache is built. + """ + monkeypatch.setenv("REDIS_HOST", "localhost") + + with patch("litellm.proxy.proxy_server.RedisCache") as mock_redis_cache: + result = ProxyStartupEvent._get_transaction_buffer_redis_cache( + general_settings={"use_redis_transaction_buffer": "true"}, + ) + + mock_redis_cache.assert_called_once() + assert result is mock_redis_cache.return_value diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 79e6494eab0..04c93f48ca9 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -239,6 +239,54 @@ async def test_update_daily_spend_sorting(): mock_table.upsert.assert_has_calls(upsert_calls) +@pytest.mark.asyncio +async def test_update_daily_spend_drains_all_batches_over_batch_size(): + """ + Regression for #30281: >BATCH_SIZE (100) unique entities in one flush must all + be written and the in-memory dict fully drained within a single call. Pre-fix, + only the first 100 sorted items were upserted then the method returned, silently + dropping the remaining entities. + """ + mock_prisma_client = MagicMock() + mock_batcher = MagicMock() + mock_table = MagicMock() + mock_prisma_client.db.batch_.return_value.__aenter__.return_value = mock_batcher + mock_batcher.litellm_dailyuserspend = mock_table + + num_entities = 250 + daily_spend_transactions = { + f"test_key_{i}": { + "user_id": f"user{i:04d}", + "date": "2024-01-01", + "api_key": "test-api-key", + "model": "gpt-4", + "custom_llm_provider": "openai", + "prompt_tokens": 10, + "completion_tokens": 20, + "spend": 0.1, + "api_requests": 1, + "successful_requests": 1, + "failed_requests": 0, + } + for i in range(num_entities) + } + + await DBSpendUpdateWriter._update_daily_spend( + n_retry_times=1, + prisma_client=mock_prisma_client, + proxy_logging_obj=MagicMock(), + daily_spend_transactions=daily_spend_transactions, + entity_type="user", + entity_id_field="user_id", + table_name="litellm_dailyuserspend", + unique_constraint_name="user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint", + ) + + assert mock_table.upsert.call_count == num_entities + assert mock_prisma_client.db.batch_.call_count == 3 + assert daily_spend_transactions == {} + + @pytest.mark.asyncio async def test_update_daily_spend_tag_with_request_id(): """ diff --git a/tests/test_litellm/proxy/discovery_endpoints/test_ui_discovery_endpoints.py b/tests/test_litellm/proxy/discovery_endpoints/test_ui_discovery_endpoints.py index 9199286e6fc..b3c3957548b 100644 --- a/tests/test_litellm/proxy/discovery_endpoints/test_ui_discovery_endpoints.py +++ b/tests/test_litellm/proxy/discovery_endpoints/test_ui_discovery_endpoints.py @@ -352,6 +352,79 @@ def test_ui_discovery_endpoints_is_control_plane_true_when_workers_configured(): assert data["workers"][0]["url"] == "https://worker-1:4001" +def test_ui_discovery_endpoints_hide_default_credentials_hint_default_false(): + """Default credentials hint is shown by default (flag false).""" + app = FastAPI() + app.include_router(router) + client = TestClient(app) + + with ( + patch("litellm.proxy.utils.get_server_root_path", return_value="/"), + patch("litellm.proxy.utils.get_proxy_base_url", return_value=None), + patch("litellm.proxy.auth.auth_utils._has_user_setup_sso", return_value=False), + patch.dict(os.environ, {"DISABLE_ADMIN_UI": "false"}, clear=False), + ): + os.environ.pop("LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT", None) + + response = client.get("/.well-known/litellm-ui-config") + + assert response.status_code == 200 + data = response.json() + assert data["hide_default_credentials_hint"] is False + + +def test_ui_discovery_endpoints_hide_default_credentials_hint_via_env_var(): + """LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT=true hides the login-page credentials card.""" + app = FastAPI() + app.include_router(router) + client = TestClient(app) + + with ( + patch("litellm.proxy.utils.get_server_root_path", return_value="/"), + patch("litellm.proxy.utils.get_proxy_base_url", return_value=None), + patch("litellm.proxy.auth.auth_utils._has_user_setup_sso", return_value=False), + patch.dict( + os.environ, + { + "LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT": "true", + "DISABLE_ADMIN_UI": "false", + }, + clear=False, + ), + ): + + response = client.get("/.well-known/litellm-ui-config") + + assert response.status_code == 200 + data = response.json() + assert data["hide_default_credentials_hint"] is True + + +def test_ui_discovery_endpoints_hide_default_credentials_hint_via_general_settings(): + """general_settings.hide_default_credentials_hint=true also hides the card.""" + app = FastAPI() + app.include_router(router) + client = TestClient(app) + + with ( + patch("litellm.proxy.utils.get_server_root_path", return_value="/"), + patch("litellm.proxy.utils.get_proxy_base_url", return_value=None), + patch("litellm.proxy.auth.auth_utils._has_user_setup_sso", return_value=False), + patch( + "litellm.proxy.proxy_server.general_settings", + {"hide_default_credentials_hint": True}, + ), + patch.dict(os.environ, {"DISABLE_ADMIN_UI": "false"}, clear=False), + ): + os.environ.pop("LITELLM_HIDE_DEFAULT_CREDENTIALS_HINT", None) + + response = client.get("/.well-known/litellm-ui-config") + + assert response.status_code == 200 + data = response.json() + assert data["hide_default_credentials_hint"] is True + + def test_ui_discovery_endpoints_is_control_plane_false_when_no_workers(): app = FastAPI() app.include_router(router) diff --git a/tests/test_litellm/proxy/google_endpoints/test_google_api_endpoints.py b/tests/test_litellm/proxy/google_endpoints/test_google_api_endpoints.py index 434f7953c21..99f587e87a3 100644 --- a/tests/test_litellm/proxy/google_endpoints/test_google_api_endpoints.py +++ b/tests/test_litellm/proxy/google_endpoints/test_google_api_endpoints.py @@ -2,6 +2,7 @@ """ Test to verify the Google GenAI proxy API endpoints """ + import os import sys from unittest.mock import AsyncMock, MagicMock, patch @@ -88,6 +89,8 @@ def test_google_stream_generate_content_endpoint(): # stream=True must be forced into the data the processor receives. init_kwargs = mock_init.call_args.kwargs assert init_kwargs["data"]["stream"] is True + assert init_kwargs["data"]["_litellm_raw_sse_stream"] is True + assert init_kwargs["data"]["_litellm_skip_openai_stream_done"] is True assert init_kwargs["data"]["model"] == "test-model" assert init_kwargs["data"]["contents"] == [ {"role": "user", "parts": [{"text": "Hello"}]} diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py index 16b5cbe8589..fe6cb98d1f5 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/openai/test_moderations.py @@ -2,6 +2,7 @@ """ Test OpenAI Moderation Guardrail """ + import os import sys @@ -822,6 +823,108 @@ def test_openai_moderation_process_error_metadata_none_edge_case(): assert "_openai_moderation_response" not in request_data["metadata"] +@pytest.mark.asyncio +async def test_openai_moderation_logs_violation_categories_harmful_content(): + """Flagged content surfaces only the violated category names in + StandardLoggingGuardrailInformation.violation_categories, so OTEL can index + a short ``guardrail_violation_categories`` attribute instead of the full + response blob (LIT-3801).""" + from fastapi import HTTPException + + from litellm.types.utils import GenericGuardrailAPIInputs + + with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}): + guardrail = OpenAIModerationGuardrail(guardrail_name="test-openai-moderation") + + mock_response = OpenAIModerationResponse( + id="modr-violations", + model="omni-moderation-latest", + results=[ + OpenAIModerationResult( + flagged=True, + categories={ + "sexual": False, + "hate": False, + "self-harm": True, + "self-harm/intent": True, + "violence": True, + }, + category_scores={ + "sexual": 0.0001, + "hate": 0.0001, + "self-harm": 0.97, + "self-harm/intent": 0.98, + "violence": 0.35, + }, + category_applied_input_types={}, + ) + ], + ) + + with patch.object(guardrail, "async_make_request", return_value=mock_response): + request_data = {"metadata": {}} + with pytest.raises(HTTPException): + await guardrail.apply_guardrail( + inputs=GenericGuardrailAPIInputs( + structured_messages=[{"role": "user", "content": "harmful"}] + ), + request_data=request_data, + input_type="request", + ) + + info = request_data["metadata"]["standard_logging_guardrail_information"][0] + + # Only the flagged categories, never the unflagged ones or the scores + assert info["violation_categories"] == [ + "self-harm", + "self-harm/intent", + "violence", + ] + + +@pytest.mark.asyncio +async def test_openai_moderation_no_violation_categories_safe_content(): + """Safe content carries no violation_categories key, so the short attribute + is absent rather than empty on allowed requests (LIT-3801).""" + from litellm.types.utils import GenericGuardrailAPIInputs + + with patch.dict(os.environ, {"OPENAI_API_KEY": "test-key"}): + guardrail = OpenAIModerationGuardrail(guardrail_name="test-openai-moderation") + + mock_response = OpenAIModerationResponse( + id="modr-safe", + model="omni-moderation-latest", + results=[ + OpenAIModerationResult( + flagged=False, + categories={"hate": False, "violence": False}, + category_scores={"hate": 0.001, "violence": 0.002}, + category_applied_input_types={}, + ) + ], + ) + + with patch.object(guardrail, "async_make_request", return_value=mock_response): + request_data = {"metadata": {}} + await guardrail.apply_guardrail( + inputs=GenericGuardrailAPIInputs( + structured_messages=[{"role": "user", "content": "hi"}] + ), + request_data=request_data, + input_type="request", + ) + + info = request_data["metadata"]["standard_logging_guardrail_information"][0] + assert "violation_categories" not in info + + +def test_openai_moderation_build_tracing_detail_non_dict_responses(): + """Non-dict guardrail responses (the "allow" sentinel, a raw Exception) yield + no tracing detail so logging never crashes when no moderation call ran.""" + assert OpenAIModerationGuardrail._build_tracing_detail("allow") is None + assert OpenAIModerationGuardrail._build_tracing_detail(ValueError("boom")) is None + + @pytest.mark.asyncio async def test_openai_moderation_guardrail_streaming_defaults(): """Defaults match the unified dispatcher: sampled in-stream, every 5th chunk.""" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py index 3efc42523f1..253d989f203 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_presidio.py @@ -584,6 +584,50 @@ async def test_logging_hook_multiple_content_items(presidio_guardrail): print("✓ Logging hook multiple content items test passed") +@pytest.mark.asyncio +async def test_logging_only_does_not_mask_pre_call_request( + mock_user_api_key, mock_cache +): + """ + A guardrail configured with `logging_only` must only mask PII for logs/traces, + never for the request sent to the model. `async_pre_call_hook` should leave the + request untouched so the model receives (and replies based on) the real input. + + Regression test for the case where the pre-call hook masked the live request, + causing the model's response to contain anonymization tokens (e.g. ) + instead of the real output. + """ + presidio_guardrail = _OPTIONAL_PresidioPIIMasking( + mock_testing=True, + logging_only=True, + pii_entities_config={PiiEntityType.PHONE_NUMBER: PiiAction.MASK}, + ) + + async def mock_check_pii(text, output_parse_pii, presidio_config, request_data): + return text.replace("555-123-4567", "[PHONE]") + + presidio_guardrail.check_pii = mock_check_pii + + original_text = "My phone is 555-123-4567" + test_data = { + "messages": [{"role": "user", "content": original_text}], + "model": "gpt-4", + } + + result = await presidio_guardrail.async_pre_call_hook( + user_api_key_dict=mock_user_api_key, + cache=mock_cache, + data=test_data, + call_type="completion", + ) + + # The live request must be unchanged: PII reaches the model intact. + assert result["messages"][0]["content"] == original_text + assert "[PHONE]" not in result["messages"][0]["content"] + + print("✓ logging_only leaves the pre-call request unmasked") + + @pytest.mark.asyncio async def test_presidio_sets_guardrail_information_in_request_data(): """ diff --git a/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py index d924d5ecdfe..3bdf9bafdc7 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_budget_endpoints.py @@ -3,6 +3,7 @@ import os import sys import types +from datetime import datetime, timedelta, timezone import pytest from unittest.mock import AsyncMock, MagicMock from fastapi.testclient import TestClient @@ -11,7 +12,6 @@ import litellm.proxy.proxy_server as ps from litellm.proxy.proxy_server import app from litellm.proxy._types import UserAPIKeyAuth, LitellmUserRoles, CommonProxyErrors - sys.path.insert( 0, os.path.abspath("../../../") ) # Adds the parent directory to the system path @@ -265,3 +265,98 @@ async def test_new_budget_invalid_model_max_budget(client_and_mocks, monkeypatch assert resp.status_code in (400, 422), resp.text detail = resp.json()["detail"] assert "model_max_budget" in str(detail) or "dictionary" in str(detail).lower() + + +def _capture_update_data(mock_table): + captured = {} + + async def capture(*, where, data): + captured.update(data) + return {**where, **data} + + mock_table.update = AsyncMock(side_effect=capture) + return captured + + +@pytest.mark.asyncio +async def test_update_budget_recomputes_reset_at_when_duration_changes( + client_and_mocks, +): + """ + Regression for LIT-3362: shortening budget_duration without an explicit + budget_reset_at must bring the reset forward instead of leaving it pinned + to the previous (longer) schedule. + """ + client, _, mock_table = client_and_mocks + captured = _capture_update_data(mock_table) + + before = datetime.now(timezone.utc) + resp = client.post( + "/budget/update", + json={"budget_id": "budget_reset_recompute", "budget_duration": "1d"}, + ) + assert resp.status_code == 200, resp.text + + assert ( + "budget_reset_at" in captured + ), "duration change must recompute budget_reset_at" + reset_at = captured["budget_reset_at"] + assert isinstance(reset_at, datetime) + assert reset_at > before, "recomputed reset must be in the future" + # "1d" resets at the next standardized day boundary, always within ~24h + assert reset_at <= before + timedelta(days=1, hours=1), reset_at + # and it must be far closer than a stale 30d schedule would have left it + assert reset_at < before + timedelta(days=29) + + +@pytest.mark.asyncio +async def test_update_budget_preserves_explicit_reset_at(client_and_mocks): + """An explicit budget_reset_at from the caller always wins over recompute.""" + client, _, mock_table = client_and_mocks + captured = _capture_update_data(mock_table) + + explicit = datetime(2027, 1, 1, tzinfo=timezone.utc) + resp = client.post( + "/budget/update", + json={ + "budget_id": "budget_explicit_reset", + "budget_duration": "1d", + "budget_reset_at": explicit.isoformat(), + }, + ) + assert resp.status_code == 200, resp.text + + assert captured["budget_reset_at"] == explicit + + +@pytest.mark.asyncio +async def test_update_budget_without_duration_leaves_reset_at_untouched( + client_and_mocks, +): + """Updates that do not touch budget_duration must not introduce budget_reset_at.""" + client, _, mock_table = client_and_mocks + captured = _capture_update_data(mock_table) + + resp = client.post( + "/budget/update", + json={"budget_id": "budget_other_field", "max_budget": 200.0}, + ) + assert resp.status_code == 200, resp.text + + assert "budget_reset_at" not in captured + + +@pytest.mark.asyncio +async def test_update_budget_duration_none_does_not_recompute(client_and_mocks): + """Clearing budget_duration (explicit null) must not recompute against a None duration.""" + client, _, mock_table = client_and_mocks + captured = _capture_update_data(mock_table) + + resp = client.post( + "/budget/update", + json={"budget_id": "budget_clear_duration", "budget_duration": None}, + ) + assert resp.status_code == 200, resp.text + + assert "budget_duration" in captured and captured["budget_duration"] is None + assert "budget_reset_at" not in captured diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index 8c26e9e4e1e..2d2c18bb46e 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -1,5 +1,6 @@ import os import sys +from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock import pytest @@ -9,11 +10,16 @@ sys.path.insert( ) # Adds the parent directory to the system path from litellm.proxy.management_endpoints.common_daily_activity import ( + _adjust_dates_for_timezone, + _build_aggregated_sql_query, _is_user_agent_tag, + _record_to_spend_metrics, get_api_key_metadata, get_daily_activity, get_daily_activity_aggregated, + update_metrics, ) +from litellm.types.proxy.management_endpoints.common_daily_activity import SpendMetrics @pytest.mark.asyncio @@ -632,6 +638,126 @@ async def test_aggregated_activity_preserves_metadata_for_deleted_keys(): assert key_data.metrics.spend == 10.0 +class TestAdjustDatesForTimezone: + """ + Regression tests for the timezone double-counting bug. + + Background: the previous implementation expanded the SQL date range by a full + UTC day on whichever side a non-UTC timezone offset pointed. Because spend is + bucketed in whole UTC days in the aggregation table, that expansion caused + single-day queries from non-UTC timezones to include a second full UTC day's + worth of data, producing approximately 2x over-counting. The sum of single-day + spends across a window then exceeded the equivalent multi-day aggregate, which + is mathematically impossible. + + These tests pin the function to a pass-through and assert the additivity + invariant that any future implementation must preserve. + """ + + @pytest.mark.parametrize( + "offset_minutes", + [ + None, + 0, + -330, # IST UTC+5:30 + -540, # JST UTC+9 + -60, # CET UTC+1 + 240, # AST UTC-4 + 300, # EST UTC-5 + 480, # PST UTC-8 + ], + ) + def test_returns_input_dates_unchanged_for_any_offset(self, offset_minutes): + start, end = _adjust_dates_for_timezone( + "2026-05-29", "2026-05-29", offset_minutes + ) + assert start == "2026-05-29" + assert end == "2026-05-29" + + def test_single_day_query_does_not_widen_to_two_utc_days(self): + """ + Pins the boundary that caused the original 2x bug: a single IST day must + not be translated into a SQL filter covering two UTC days. + """ + start, end = _adjust_dates_for_timezone("2026-05-29", "2026-05-29", -330) + assert start == end == "2026-05-29", ( + "Single-day IST query expanded to a multi-day UTC range; this is " + "the regression that produced approximately 2x over-counting." + ) + + def test_multi_day_range_endpoints_are_preserved(self): + start, end = _adjust_dates_for_timezone("2026-05-29", "2026-06-02", -330) + assert (start, end) == ("2026-05-29", "2026-06-02") + + @pytest.mark.parametrize("offset_minutes", [-330, 480]) + def test_single_day_sums_match_multi_day_window(self, offset_minutes): + """ + Additivity invariant: querying each day in a window separately and summing + the resulting SQL ranges must cover exactly the same range as querying the + whole window at once. The bug broke this; without it, single-day sums + exceeded the multi-day total by ~50% over a 5-day IST window. + """ + days = ["2026-05-29", "2026-05-30", "2026-05-31", "2026-06-01", "2026-06-02"] + single_day_ranges = [ + _adjust_dates_for_timezone(d, d, offset_minutes) for d in days + ] + multi_day_range = _adjust_dates_for_timezone(days[0], days[-1], offset_minutes) + + per_day_starts = [r[0] for r in single_day_ranges] + per_day_ends = [r[1] for r in single_day_ranges] + assert min(per_day_starts) == multi_day_range[0] + assert max(per_day_ends) == multi_day_range[1] + assert per_day_starts == days + assert per_day_ends == days + + +class TestBuildAggregatedSqlQuery: + """ + Asserts the SQL emitted by the aggregated query path stays anchored to the + user-supplied date range. The original bug shipped a function that returned + expanded dates from _adjust_dates_for_timezone, so the regression surface is + not just the helper but the SQL it feeds into. + """ + + @pytest.mark.parametrize("offset_minutes", [None, 0, -330, 480]) + def test_sql_date_bounds_are_user_supplied_dates(self, offset_minutes): + sql, params = _build_aggregated_sql_query( + table_name="litellm_dailyuserspend", + entity_id_field="user_id", + entity_id="user-1", + start_date="2026-05-29", + end_date="2026-05-29", + model=None, + api_key=None, + timezone_offset_minutes=offset_minutes, + ) + + assert params[0] == "2026-05-29" + assert params[1] == "2026-05-29" + assert "date >= $1" in sql + assert "date <= $2" in sql + + def test_optional_filters_appear_in_params_in_order(self): + sql, params = _build_aggregated_sql_query( + table_name="litellm_dailyuserspend", + entity_id_field="user_id", + entity_id="user-1", + start_date="2026-05-29", + end_date="2026-06-02", + model="bedrock/global.anthropic.claude-opus-4-8", + api_key="sk-test", + timezone_offset_minutes=-330, + ) + + assert params == [ + "2026-05-29", + "2026-06-02", + "user-1", + "bedrock/global.anthropic.claude-opus-4-8", + "sk-test", + ] + assert "model = $4" in sql + assert "api_key = $5" in sql @pytest.mark.asyncio async def test_get_daily_activity_aggregated_empty_result_set(): """Regression test for the empty-range 500. @@ -688,3 +814,45 @@ async def test_get_daily_activity_aggregated_empty_result_set(): assert result.metadata.total_failed_requests == 0 assert result.metadata.total_cache_read_input_tokens == 0 assert result.metadata.total_cache_creation_input_tokens == 0 + + +def _no_spend_record(): + """A rollup row for a key with no spend, where SUM() returns NULL (None).""" + return SimpleNamespace( + spend=None, + prompt_tokens=None, + completion_tokens=None, + cache_read_input_tokens=None, + cache_creation_input_tokens=None, + api_requests=None, + successful_requests=None, + failed_requests=None, + ) + + +def test_record_to_spend_metrics_handles_none_values(): + """Keys with no spend produce NULL aggregates; treat them as zero, not a crash.""" + metrics = _record_to_spend_metrics(_no_spend_record()) + assert metrics.spend == 0 + assert metrics.prompt_tokens == 0 + assert metrics.completion_tokens == 0 + assert metrics.total_tokens == 0 + assert metrics.api_requests == 0 + assert metrics.successful_requests == 0 + assert metrics.failed_requests == 0 + assert metrics.cache_read_input_tokens == 0 + assert metrics.cache_creation_input_tokens == 0 + + +def test_update_metrics_handles_none_values(): + """update_metrics should coalesce NULL aggregates instead of raising TypeError.""" + metrics = update_metrics(SpendMetrics(), _no_spend_record()) + assert metrics.spend == 0 + assert metrics.prompt_tokens == 0 + assert metrics.completion_tokens == 0 + assert metrics.total_tokens == 0 + assert metrics.api_requests == 0 + assert metrics.successful_requests == 0 + assert metrics.failed_requests == 0 + assert metrics.cache_read_input_tokens == 0 + assert metrics.cache_creation_input_tokens == 0 diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_utils.py b/tests/test_litellm/proxy/management_endpoints/test_common_utils.py index d53ea6fa34d..7a8d04507dc 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_utils.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_utils.py @@ -157,6 +157,32 @@ class TestUpdateMetadataFieldsEmptyCollections: assert "guardrails" not in updated_kv assert updated_kv["metadata"]["guardrails"] == ["my-guardrail"] + @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + def test_false_boolean_does_not_trigger_premium_check(self, mock_premium_check): + """ + Regression #30285: /team/update sends disable_global_guardrails=False + (the UI's unchanged default). A falsy boolean must not trigger the + premium check, so non-premium users are not wrongly 403'd. + """ + updated_kv = {"team_id": "test-team", "disable_global_guardrails": False} + _update_metadata_fields(updated_kv=updated_kv) + mock_premium_check.assert_not_called() + + @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + def test_false_boolean_still_updates_metadata(self, mock_premium_check): + """A falsy boolean must still be moved into metadata so it persists.""" + updated_kv = {"team_id": "test-team", "disable_global_guardrails": False} + _update_metadata_fields(updated_kv=updated_kv) + assert "disable_global_guardrails" not in updated_kv + assert updated_kv["metadata"]["disable_global_guardrails"] is False + + @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") + def test_true_boolean_triggers_premium_check(self, mock_premium_check): + """Control: enabling the premium feature (True) still requires a license.""" + updated_kv = {"team_id": "test-team", "disable_global_guardrails": True} + _update_metadata_fields(updated_kv=updated_kv) + mock_premium_check.assert_called() + @patch("litellm.proxy.management_endpoints.common_utils._premium_user_check") def test_ui_typical_payload_does_not_trigger_premium_check( self, mock_premium_check diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index ed04b9e30dd..cc8b4c7f5cc 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -1547,6 +1547,51 @@ async def test_prepare_key_update_data_budget_limits_serializes_windows(): assert windows[0]["reset_at"] is not None +@pytest.mark.asyncio +async def test_prepare_key_update_data_disable_global_guardrails_false_no_premium( + monkeypatch, +): + """ + Regression #30285: editing a key via the UI sends disable_global_guardrails=False + (unchanged default). A non-premium user must NOT get a 403, and False must persist. + """ + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False) + data = UpdateKeyRequest(key="sk-1", disable_global_guardrails=False) + existing_key = LiteLLM_VerificationToken(token="hashed") + + result = await prepare_key_update_data(data=data, existing_key_row=existing_key) + + assert result["metadata"]["disable_global_guardrails"] is False + + +@pytest.mark.asyncio +async def test_prepare_key_update_data_disable_global_guardrails_true_requires_premium( + monkeypatch, +): + """Control: enabling the premium feature (True) without a license still 403s.""" + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False) + data = UpdateKeyRequest(key="sk-1", disable_global_guardrails=True) + existing_key = LiteLLM_VerificationToken(token="hashed") + + with pytest.raises(HTTPException) as exc_info: + await prepare_key_update_data(data=data, existing_key_row=existing_key) + assert exc_info.value.status_code == 403 + + +@pytest.mark.asyncio +async def test_prepare_key_update_data_disable_global_guardrails_true_premium_persists( + monkeypatch, +): + """A premium user enabling the feature (True) succeeds and the value persists.""" + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True) + data = UpdateKeyRequest(key="sk-1", disable_global_guardrails=True) + existing_key = LiteLLM_VerificationToken(token="hashed") + + result = await prepare_key_update_data(data=data, existing_key_row=existing_key) + + assert result["metadata"]["disable_global_guardrails"] is True + + @pytest.mark.asyncio async def test_validate_team_id_used_in_service_account_request_requires_team_id(): """ @@ -11862,7 +11907,6 @@ async def test_ghsa_q775_default_team_id_does_not_grant_session_token_exemption( assert "cannot exceed" in msg.lower() - @pytest.mark.asyncio async def test_prepare_key_update_data_budget_duration_null_clears_fields(): """ @@ -11941,3 +11985,511 @@ async def test_prepare_key_update_data_budget_duration_valid_sets_reset(): assert result["budget_reset_at"] is not None +@pytest.mark.asyncio +async def test_info_key_fn_includes_model_max_budget_usage(monkeypatch): + """ + /key/info should include model_max_budget_usage showing current-period spend + for each model that has a per-model budget configured. + """ + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._types import LiteLLM_VerificationToken + from litellm.proxy.management_endpoints.key_management_endpoints import info_key_fn + + test_key_token = "hashed_token_budget_test" + model_max_budget = { + "gpt-4o": {"budget_limit": 0.50, "time_period": "1d"}, + } + + mock_prisma_client = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mock_user_api_key_cache = AsyncMock() + mock_user_api_key_cache.async_get_cache = AsyncMock(return_value=0.23) + monkeypatch.setattr( + "litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache + ) + + mock_key_info = MagicMock(spec=LiteLLM_VerificationToken) + mock_key_info.token = test_key_token + mock_key_info.object_permission_id = None + mock_key_info.user_id = "user-x" + mock_key_info.team_id = None + mock_key_info.litellm_budget_table = None + mock_key_info.model_dump.return_value = { + "token": test_key_token, + "model_max_budget": model_max_budget, + "user_id": "user-x", + "team_id": None, + "object_permission_id": None, + "litellm_budget_table": None, + } + mock_key_info.dict.return_value = mock_key_info.model_dump.return_value + + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=mock_key_info + ) + mock_prisma_client.db.query_raw = AsyncMock() + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-test-budget-key", + ) + + result = await info_key_fn( + key="sk-test-budget-key", + user_api_key_dict=user_api_key_dict, + ) + + assert "model_max_budget_usage" in result["info"] + usage = result["info"]["model_max_budget_usage"] + assert usage["gpt-4o"]["current_spend"] == 0.23 + assert usage["gpt-4o"]["budget_limit"] == 0.50 + assert usage["gpt-4o"]["time_period"] == "1d" + mock_prisma_client.db.query_raw.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_info_key_fn_no_model_max_budget_skips_usage(monkeypatch): + """Keys with no model_max_budget should not include model_max_budget_usage.""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._types import LiteLLM_VerificationToken + from litellm.proxy.management_endpoints.key_management_endpoints import info_key_fn + + test_key_token = "hashed_token_no_budget" + + mock_prisma_client = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mock_user_api_key_cache = AsyncMock() + mock_user_api_key_cache.async_get_cache = AsyncMock() + monkeypatch.setattr( + "litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache + ) + + mock_key_info = MagicMock(spec=LiteLLM_VerificationToken) + mock_key_info.token = test_key_token + mock_key_info.object_permission_id = None + mock_key_info.user_id = "user-y" + mock_key_info.team_id = None + mock_key_info.litellm_budget_table = None + mock_key_info.model_dump.return_value = { + "token": test_key_token, + "model_max_budget": {}, + "user_id": "user-y", + "team_id": None, + "object_permission_id": None, + "litellm_budget_table": None, + } + mock_key_info.dict.return_value = mock_key_info.model_dump.return_value + + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=mock_key_info + ) + mock_prisma_client.db.query_raw = AsyncMock() + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-test-no-budget", + ) + + result = await info_key_fn( + key="sk-test-no-budget", + user_api_key_dict=user_api_key_dict, + ) + + assert "model_max_budget_usage" not in result["info"] + mock_prisma_client.db.query_raw.assert_not_awaited() + mock_user_api_key_cache.async_get_cache.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_info_key_fn_v2_includes_model_max_budget_usage(monkeypatch): + """/v2/key/info should include model_max_budget_usage for keys with per-model budgets.""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._types import KeyRequest, LiteLLM_VerificationToken + from litellm.proxy.management_endpoints.key_management_endpoints import ( + info_key_fn_v2, + ) + + test_key_token = "hashed_token_v2_test" + model_max_budget = {"gpt-4o": {"budget_limit": 1.00, "time_period": "7d"}} + + mock_prisma_client = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mock_user_api_key_cache = AsyncMock() + mock_user_api_key_cache.async_get_cache = AsyncMock(return_value=0.55) + monkeypatch.setattr( + "litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache + ) + + mock_key = MagicMock(spec=LiteLLM_VerificationToken) + mock_key.token = test_key_token + mock_key.user_id = "user-v2" + mock_key.team_id = None + mock_key.model_dump.return_value = { + "token": test_key_token, + "model_max_budget": model_max_budget, + "user_id": "user-v2", + "team_id": None, + "litellm_budget_table": None, + } + mock_key.dict.return_value = mock_key.model_dump.return_value + + mock_prisma_client.get_data = AsyncMock(return_value=[mock_key]) + mock_prisma_client.db.query_raw = AsyncMock() + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin", + ) + + result = await info_key_fn_v2( + data=KeyRequest(keys=[test_key_token]), + user_api_key_dict=user_api_key_dict, + ) + + assert len(result["info"]) == 1 + key_info = result["info"][0] + assert "model_max_budget_usage" in key_info + usage = key_info["model_max_budget_usage"] + assert usage["gpt-4o"]["current_spend"] == 0.55 + assert usage["gpt-4o"]["budget_limit"] == 1.00 + assert usage["gpt-4o"]["time_period"] == "7d" + mock_prisma_client.db.query_raw.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_info_key_fn_budget_table_fallback(monkeypatch): + """When model_max_budget is empty on the key but set in litellm_budget_table, + /key/info should still populate model_max_budget_usage. + """ + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._types import LiteLLM_VerificationToken + from litellm.proxy.management_endpoints.key_management_endpoints import info_key_fn + + test_key_token = "hashed_token_budget_table_test" + budget_table_model_max_budget = { + "bedrock/anthropic.claude-opus-4": {"max_budget": 5, "budget_duration": "30d"}, + } + + mock_prisma_client = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mock_user_api_key_cache = AsyncMock() + mock_user_api_key_cache.async_get_cache = AsyncMock(return_value=1.20) + monkeypatch.setattr( + "litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache + ) + + mock_key_info = MagicMock(spec=LiteLLM_VerificationToken) + mock_key_info.token = test_key_token + mock_key_info.object_permission_id = None + mock_key_info.user_id = "user-bt" + mock_key_info.team_id = None + mock_key_info.litellm_budget_table = None + mock_key_info.model_dump.return_value = { + "token": test_key_token, + "model_max_budget": {}, + "user_id": "user-bt", + "team_id": None, + "object_permission_id": None, + "litellm_budget_table": { + "budget_id": "bt-123", + "budget_duration": "30d", + "budget_reset_at": "2026-07-01T00:00:00+00:00", + "model_max_budget": budget_table_model_max_budget, + }, + } + mock_key_info.dict.return_value = mock_key_info.model_dump.return_value + + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=mock_key_info + ) + mock_prisma_client.db.query_raw = AsyncMock() + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-test-bt-key", + ) + + result = await info_key_fn( + key="sk-test-bt-key", + user_api_key_dict=user_api_key_dict, + ) + + assert "model_max_budget_usage" in result["info"] + usage = result["info"]["model_max_budget_usage"] + assert usage["bedrock/anthropic.claude-opus-4"]["current_spend"] == 1.20 + assert usage["bedrock/anthropic.claude-opus-4"]["budget_limit"] == 5 + assert usage["bedrock/anthropic.claude-opus-4"]["time_period"] == "30d" + mock_prisma_client.db.query_raw.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_info_key_fn_v2_budget_table_fallback(monkeypatch): + """When model_max_budget is empty on the key but set in litellm_budget_table, + /v2/key/info should still populate model_max_budget_usage.""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._types import KeyRequest, LiteLLM_VerificationToken + from litellm.proxy.management_endpoints.key_management_endpoints import ( + info_key_fn_v2, + ) + + test_key_token = "hashed_token_v2_bt_test" + budget_table_model_max_budget = { + "bedrock/anthropic.claude-opus-4": {"max_budget": 5, "budget_duration": "30d"}, + } + + mock_prisma_client = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mock_user_api_key_cache = AsyncMock() + mock_user_api_key_cache.async_get_cache = AsyncMock(return_value=2.50) + monkeypatch.setattr( + "litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache + ) + + mock_key = MagicMock(spec=LiteLLM_VerificationToken) + mock_key.token = test_key_token + mock_key.user_id = "user-v2-bt" + mock_key.team_id = None + mock_key.model_dump.return_value = { + "token": test_key_token, + "model_max_budget": {}, + "user_id": "user-v2-bt", + "team_id": None, + "litellm_budget_table": { + "budget_id": "bt-456", + "budget_duration": "30d", + "budget_reset_at": "2026-07-01T00:00:00+00:00", + "model_max_budget": budget_table_model_max_budget, + }, + } + mock_key.dict.return_value = mock_key.model_dump.return_value + + mock_prisma_client.get_data = AsyncMock(return_value=[mock_key]) + mock_prisma_client.db.query_raw = AsyncMock() + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-admin-v2-bt", + ) + + result = await info_key_fn_v2( + data=KeyRequest(keys=[test_key_token]), + user_api_key_dict=user_api_key_dict, + ) + + assert len(result["info"]) == 1 + key_info = result["info"][0] + assert "model_max_budget_usage" in key_info + usage = key_info["model_max_budget_usage"] + assert usage["bedrock/anthropic.claude-opus-4"]["current_spend"] == 2.50 + mock_prisma_client.db.query_raw.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_info_key_fn_provider_prefix_spend_fallback(monkeypatch): + """Cached spend for 'gpt-4o' matches budget key 'openai/gpt-4o' via suffix match.""" + from unittest.mock import AsyncMock, MagicMock + + from litellm.proxy._types import LiteLLM_VerificationToken + from litellm.proxy.management_endpoints.key_management_endpoints import info_key_fn + + test_key_token = "hashed_token_prefix_test" + model_max_budget = { + "openai/gpt-4o": {"budget_limit": 2.00, "time_period": "7d"}, + } + + mock_prisma_client = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mock_user_api_key_cache = AsyncMock() + mock_user_api_key_cache.async_get_cache = AsyncMock(side_effect=[None, 0.75]) + monkeypatch.setattr( + "litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache + ) + + mock_key_info = MagicMock(spec=LiteLLM_VerificationToken) + mock_key_info.token = test_key_token + mock_key_info.object_permission_id = None + mock_key_info.user_id = "user-prefix" + mock_key_info.team_id = None + mock_key_info.litellm_budget_table = None + mock_key_info.model_dump.return_value = { + "token": test_key_token, + "model_max_budget": model_max_budget, + "user_id": "user-prefix", + "team_id": None, + "object_permission_id": None, + "litellm_budget_table": None, + } + mock_key_info.dict.return_value = mock_key_info.model_dump.return_value + + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=mock_key_info + ) + mock_prisma_client.db.query_raw = AsyncMock() + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-prefix-test", + ) + + result = await info_key_fn( + key="sk-prefix-test", + user_api_key_dict=user_api_key_dict, + ) + + assert "model_max_budget_usage" in result["info"] + usage = result["info"]["model_max_budget_usage"] + assert usage["openai/gpt-4o"]["current_spend"] == 0.75 + assert mock_user_api_key_cache.async_get_cache.await_count == 2 + + +@pytest.mark.asyncio +async def test_build_model_max_budget_usage_no_cache_returns_empty(): + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _build_model_max_budget_usage, + ) + + result = await _build_model_max_budget_usage( + api_key_hash="some-hash", + model_max_budget={"gpt-4o": {"budget_limit": 1.0, "time_period": "1d"}}, + user_api_key_cache=None, + ) + assert result == {} + + +@pytest.mark.asyncio +async def test_build_model_max_budget_usage_reads_current_cache_window(): + from unittest.mock import AsyncMock + + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _build_model_max_budget_usage, + ) + + mock_user_api_key_cache = AsyncMock() + mock_user_api_key_cache.async_get_cache = AsyncMock(return_value=0.30) + + result = await _build_model_max_budget_usage( + api_key_hash="some-hash", + model_max_budget={"gpt-4o": {"budget_limit": 1.0, "time_period": "30d"}}, + user_api_key_cache=mock_user_api_key_cache, + ) + + assert result["gpt-4o"]["current_spend"] == 0.30 + mock_user_api_key_cache.async_get_cache.assert_awaited_once_with( + key="virtual_key_spend:some-hash:gpt-4o:30d" + ) + + +@pytest.mark.asyncio +async def test_build_model_max_budget_usage_no_duration_in_budget_returns_empty(): + from unittest.mock import AsyncMock + + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _build_model_max_budget_usage, + ) + + mock_user_api_key_cache = AsyncMock() + mock_user_api_key_cache.async_get_cache = AsyncMock() + + result = await _build_model_max_budget_usage( + api_key_hash="some-hash", + model_max_budget={"gpt-4o": {"budget_limit": 1.0}}, + user_api_key_cache=mock_user_api_key_cache, + ) + assert result == {} + mock_user_api_key_cache.async_get_cache.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_build_model_max_budget_usage_skips_model_without_duration(): + from unittest.mock import AsyncMock + + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _build_model_max_budget_usage, + ) + + mock_user_api_key_cache = AsyncMock() + mock_user_api_key_cache.async_get_cache = AsyncMock(return_value=0.10) + + result = await _build_model_max_budget_usage( + api_key_hash="some-hash", + model_max_budget={ + "gpt-4o": {"budget_limit": 1.0, "time_period": "1d"}, + "gpt-3.5-turbo": {"budget_limit": 0.5}, + }, + user_api_key_cache=mock_user_api_key_cache, + ) + assert "gpt-4o" in result + assert "gpt-3.5-turbo" not in result + assert mock_user_api_key_cache.async_get_cache.await_count == 1 + + +@pytest.mark.asyncio +async def test_build_model_max_budget_usage_unparseable_duration_skipped(): + from unittest.mock import AsyncMock + + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _build_model_max_budget_usage, + ) + + mock_user_api_key_cache = AsyncMock() + mock_user_api_key_cache.async_get_cache = AsyncMock() + + result = await _build_model_max_budget_usage( + api_key_hash="some-hash", + model_max_budget={ + "gpt-4o": {"budget_limit": 1.0, "budget_duration": "not-valid"} + }, + user_api_key_cache=mock_user_api_key_cache, + ) + assert result == {} + mock_user_api_key_cache.async_get_cache.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_build_model_max_budget_usage_invalid_budget_config_skipped(): + from unittest.mock import AsyncMock + + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _build_model_max_budget_usage, + ) + + mock_user_api_key_cache = AsyncMock() + mock_user_api_key_cache.async_get_cache = AsyncMock(return_value=0.20) + + result = await _build_model_max_budget_usage( + api_key_hash="some-hash", + model_max_budget={ + "gpt-4o": {"max_budget": "not-a-number", "budget_duration": "1d"}, + "gpt-3.5-turbo": {"budget_limit": 0.5, "time_period": "7d"}, + }, + user_api_key_cache=mock_user_api_key_cache, + ) + assert "gpt-4o" not in result + assert "gpt-3.5-turbo" in result + assert mock_user_api_key_cache.async_get_cache.await_count == 1 + + +@pytest.mark.asyncio +async def test_build_model_max_budget_usage_provider_prefix_cache_fallback(): + from unittest.mock import AsyncMock + + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _build_model_max_budget_usage, + ) + + mock_user_api_key_cache = AsyncMock() + mock_user_api_key_cache.async_get_cache = AsyncMock(side_effect=[None, 0.55]) + + result = await _build_model_max_budget_usage( + api_key_hash="test-hash", + model_max_budget={"openai/gpt-4o": {"budget_limit": 2.0, "time_period": "7d"}}, + user_api_key_cache=mock_user_api_key_cache, + ) + + assert result["openai/gpt-4o"]["current_spend"] == 0.55 + assert mock_user_api_key_cache.async_get_cache.await_count == 2 diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index b81807ee19e..a649bc7225e 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -9313,3 +9313,34 @@ async def test_clear_team_member_budget_fields_no_budget_row_skips_update(): mock_update_budget.assert_not_awaited() assert "team_member_budget" not in result assert "team_member_rpm_limit" not in result + + +@pytest.mark.asyncio +async def test_team_info_forwards_key_limit_to_get_data(): + """/team/info must thread its ``key_limit`` query param into the key + lookup so the database caps how many keys are returned for the team. + """ + from fastapi import Request + + from litellm.proxy.management_endpoints import team_endpoints + + mock_prisma = MagicMock() + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock( + return_value=LiteLLM_TeamTable(team_id="team-1") + ) + mock_prisma.get_data = AsyncMock(return_value=[]) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), + patch.object( + team_endpoints, "get_all_team_memberships", AsyncMock(return_value=[]) + ), + ): + await team_endpoints.team_info( + http_request=MagicMock(spec=Request), + team_id="team-1", + key_limit=7, + user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), + ) + + assert mock_prisma.get_data.await_args.kwargs["limit"] == 7 diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 2efec3e0b34..acca357e641 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -2120,6 +2120,292 @@ class TestCLIKeyRegenerationFlow: assert exc_info.value.status_code == 429 mock_cache.set_cache.assert_not_called() + @pytest.mark.asyncio + async def test_cli_sso_start_returns_verification_uri_complete_when_enabled(self): + """Test CLI SSO start returns a verification_uri_complete that round-trips the user_code only when the operator opts in""" + from urllib.parse import parse_qs, urlparse + + from litellm.constants import LITELLM_CLI_SOURCE_IDENTIFIER + from litellm.proxy.management_endpoints.ui_sso import cli_sso_start + + mock_request = MagicMock(spec=Request) + mock_request.client = SimpleNamespace(host="127.0.0.1") + mock_request.headers = {} + mock_request.base_url = "https://proxy.example.com/" + mock_cache = MagicMock() + mock_cache.increment_cache.return_value = 1 + + with ( + patch.dict( + os.environ, + {"PROXY_BASE_URL": "https://proxy.example.com", "SERVER_ROOT_PATH": ""}, + ), + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch( + "litellm.proxy.proxy_server.general_settings", + {"allow_cli_sso_verification_uri_complete": True}, + ), + ): + result = await cli_sso_start(request=mock_request) + + verification_uri_complete = result["verification_uri_complete"] + parsed = urlparse(verification_uri_complete) + query = parse_qs(parsed.query) + + assert parsed.path.endswith("/sso/key/generate") + assert query["source"] == [LITELLM_CLI_SOURCE_IDENTIFIER] + assert query["key"] == [result["login_id"]] + assert query["user_code"] == [result["user_code"]] + + @pytest.mark.asyncio + async def test_cli_sso_start_omits_verification_uri_complete_by_default(self): + """Test CLI SSO start does NOT advertise verification_uri_complete unless the operator enables it (default off)""" + from litellm.proxy.management_endpoints.ui_sso import cli_sso_start + + mock_request = MagicMock(spec=Request) + mock_request.client = SimpleNamespace(host="127.0.0.1") + mock_request.headers = {} + mock_request.base_url = "https://proxy.example.com/" + mock_cache = MagicMock() + mock_cache.increment_cache.return_value = 1 + + with ( + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch("litellm.proxy.proxy_server.general_settings", {}), + ): + result = await cli_sso_start(request=mock_request) + + assert "verification_uri_complete" not in result + assert result["user_code"] + assert result["login_id"].startswith("cli-") + + def test_cli_sso_verification_uri_complete_enabled_reads_general_settings(self): + """Test the operator opt-in flag is read from general_settings and defaults off""" + from litellm.proxy.management_endpoints.ui_sso import ( + _cli_sso_verification_uri_complete_enabled, + ) + + with patch("litellm.proxy.proxy_server.general_settings", {}): + assert _cli_sso_verification_uri_complete_enabled() is False + with patch( + "litellm.proxy.proxy_server.general_settings", + {"allow_cli_sso_verification_uri_complete": True}, + ): + assert _cli_sso_verification_uri_complete_enabled() is True + + @pytest.mark.asyncio + async def test_google_login_only_threads_user_code_when_enabled(self): + """Test google_login forwards user_code into the OAuth state only when the operator opt-in is on, dropping it otherwise""" + from litellm.proxy.management_endpoints.ui_sso import google_login + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://proxy.example.com/" + mock_cache = MagicMock() + mock_cache.get_cache.return_value = {"poll_secret_hash": "h"} + + async def drive(enabled: bool): + with ( + patch.dict(os.environ, {}, clear=True), + patch("litellm.proxy.proxy_server.premium_user", True), + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.user_api_key_cache", mock_cache), + patch( + "litellm.proxy.proxy_server.user_custom_ui_sso_sign_in_handler", + None, + ), + patch( + "litellm.proxy.proxy_server.general_settings", + {"allow_cli_sso_verification_uri_complete": enabled}, + ), + patch( + "litellm.proxy.management_endpoints.ui_sso.show_missing_vars_in_env", + return_value=None, + ), + patch( + "litellm.proxy.management_endpoints.ui_sso.SSOAuthenticationHandler.get_redirect_url_for_sso", + return_value="https://proxy.example.com/sso/callback", + ), + patch( + "litellm.proxy.management_endpoints.ui_sso.SSOAuthenticationHandler._get_cli_state", + return_value=None, + ) as mock_get_cli_state, + ): + try: + await google_login( + request=mock_request, + source="litellm-cli", + key="cli-validsessionkey123456", + user_code="WXYZ-2345", + ) + except Exception: + pass + return mock_get_cli_state.call_args.kwargs["user_code"] + + assert await drive(enabled=True) == "WXYZ-2345" + assert await drive(enabled=False) is None + + def test_get_cli_state_appends_user_code_for_prefill(self): + """Test the OAuth state carries the user_code only for the opt-in prefill flow""" + from litellm.constants import ( + LITELLM_CLI_SESSION_TOKEN_PREFIX, + LITELLM_CLI_SOURCE_IDENTIFIER, + ) + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + manual_state = SSOAuthenticationHandler._get_cli_state( + source=LITELLM_CLI_SOURCE_IDENTIFIER, key="cli-abc123" + ) + prefill_state = SSOAuthenticationHandler._get_cli_state( + source=LITELLM_CLI_SOURCE_IDENTIFIER, + key="cli-abc123", + user_code="WXYZ-2345", + ) + + assert manual_state == f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-abc123" + assert ( + prefill_state == f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-abc123:WXYZ-2345" + ) + assert ( + SSOAuthenticationHandler._get_cli_state( + source="not-cli", key="cli-abc123", user_code="WXYZ-2345" + ) + is None + ) + + def test_get_cli_state_drops_malformed_user_code(self): + """Test a user_code that is not a server-issued code is dropped before reaching the size-limited OAuth state""" + from litellm.constants import ( + LITELLM_CLI_SESSION_TOKEN_PREFIX, + LITELLM_CLI_SOURCE_IDENTIFIER, + ) + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + manual_only = f"{LITELLM_CLI_SESSION_TOKEN_PREFIX}:cli-abc123" + for bad_user_code in ("A" * 4096, "not-a-code", "WXYZ2345", "WXYZ-234", ""): + assert ( + SSOAuthenticationHandler._get_cli_state( + source=LITELLM_CLI_SOURCE_IDENTIFIER, + key="cli-abc123", + user_code=bad_user_code, + ) + == manual_only + ) + + def test_is_valid_cli_sso_user_code_matches_generated_format(self): + """Test the user_code validator accepts a freshly generated code and rejects malformed input""" + from litellm.proxy.management_endpoints.ui_sso import ( + _generate_cli_sso_user_code, + _is_valid_cli_sso_user_code, + ) + + assert _is_valid_cli_sso_user_code(_generate_cli_sso_user_code()) + assert _is_valid_cli_sso_user_code("WXYZ-2345") + assert not _is_valid_cli_sso_user_code("WXYZ-2340") # 0 is not in the alphabet + assert not _is_valid_cli_sso_user_code("wxyz-2345") + assert not _is_valid_cli_sso_user_code("WXYZ2345") + assert not _is_valid_cli_sso_user_code("A" * 64) + assert not _is_valid_cli_sso_user_code(None) + + def test_cli_state_round_trips_user_code_to_callback_parser(self): + """Test the callback's state parser recovers login_id and user_code from the state _get_cli_state builds""" + from litellm.constants import LITELLM_CLI_SOURCE_IDENTIFIER + from litellm.proxy.management_endpoints.ui_sso import SSOAuthenticationHandler + + state = SSOAuthenticationHandler._get_cli_state( + source=LITELLM_CLI_SOURCE_IDENTIFIER, + key="cli-abc123", + user_code="WXYZ-2345", + ) + + state_parts = state.split(":", 2) + key_id = state_parts[1] if len(state_parts) > 1 else None + prefill_user_code = state_parts[2] if len(state_parts) > 2 else None + + assert key_id == "cli-abc123" + assert prefill_user_code == "WXYZ-2345" + + def test_render_cli_sso_verification_page_prefills_user_code(self): + """Test the verify page pre-fills the user_code input (HTML-escaped) when provided""" + from litellm.proxy.management_endpoints.ui_sso import ( + _render_cli_sso_verification_page, + ) + + html = _render_cli_sso_verification_page( + verify_url="https://proxy.example.com/sso/cli/complete/cli-abc123", + browser_complete_token="browser-token", + prefill_user_code='WXYZ-2345">