mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge remote-tracking branch 'origin/main' into litellm_role_permissions_normalization
This commit is contained in:
commit
800b09ba41
55 changed files with 3865 additions and 422 deletions
130
AGENTS.md
130
AGENTS.md
|
|
@ -1,3 +1,131 @@
|
|||
Read @CLAUDE.md for coding guidelines
|
||||
Do not write comments unless they are any of:
|
||||
- absolutely necessary to explain some very complex business logic (in which case, keep it concise and clear)
|
||||
- used as an input for tools to read and act on. For example:
|
||||
- entries in `.git-blame-ignore-revs` saying which commit is excluded from git blame
|
||||
- a lint or type checker suppression like `# mutable-ok` or `# pyright: ignore[reportArgumentType] # <reason>` when introducing a truly unavoidable violation
|
||||
- a TODO or FIXME
|
||||
- Not great to have those, but if it's unavoidable, make sure to include a strong, concise reason for why it's there or, better yet, link to a GitHub issue for the follow-up work
|
||||
|
||||
Explanation: The point of this rule is to keep out AI slop comments. AI writes way too many and way too verbose comments. Code comments are, in a way, a violation of DRY code. You must update logic in two locations to change the code, and "hard to change" is literally the definition of tech debt. We should instead aim to write code that is intuitive and clear, even at a glance, to the reader, being both easy to maintain and high performance
|
||||
|
||||
Don't assume that the existing code is correct or the right way of doing things / good coding patterns. In fact, there are a lot of bad coding practices, overly complex code, code smells, etc. If something doesn't look right, speak up. Feel free to break existing patterns or question weird existing code to make new code high quality, as in:
|
||||
|
||||
- correct
|
||||
- secure
|
||||
- performant
|
||||
- readable
|
||||
- easy to maintain/change
|
||||
- modern
|
||||
|
||||
In descending order of importance
|
||||
|
||||
When adding new features, add meaningful tests. Don't add tests that don't check anything substantial and is there just to make the code coverage pass. Yes, code coverage is important, but I'd rather have no signal whether the code is working than tests that don't fail when code is broken. The goal is to have tests that would fail before the feature was added/if the code was mutated in a way that breaks the feature and succeed only when the feature is fully working. I should run mutation testing and see > 90% kill rate
|
||||
|
||||
Same thing for bug fixes. The tests should make it so that this specific bug can never happen again without failing tests (i.e., regression)
|
||||
|
||||
Never test structure of code only function of it
|
||||
|
||||
A test must only fail when litellm code changes. Never pin facts we don't own (a vendor's price, a third party's field, an upstream default, today's date) as literals or as "X must be absent"; assert the invariant our code guarantees instead, e.g. two rows agree, a value is within range, a field is derived from another. If an outside fact is truly load-bearing, cite its source and date next to the assertion so a reader can tell stale from broken
|
||||
|
||||
`tests/test_litellm/` mirrors `litellm/` in a parallel path (see `tests/test_litellm/readme.md`). Name tests `test_<filename>.py`, but always match the existing test file in the directory you touch — many provider dirs use longer descriptive names (e.g. `test_anthropic_chat_transformation.py`) to avoid ambiguity across sibling folders. For bug fixes, extend the existing mapped test file rather than creating a new one. Only create a new test file for a new feature (provider, endpoint, or transformation module) that has no mapped test yet, following that directory's naming convention (or `test_<filename>.py` if you're the first test there). One focused regression test beats many shallow ones
|
||||
|
||||
End-to-end tests belong in `tests/e2e/` and must follow the harness conventions documented in that directory's `AGENTS.md`
|
||||
|
||||
When creating PRs, target the repository's current default branch for both internal and external / OSS contributions. Check it with `python3 scripts/default_branch.py --branch` instead of assuming a branch name or relying on cached `origin/HEAD`
|
||||
|
||||
When writing a PR body, treat the comments and imperative instructions inside .github/pull_request_template.md as rules to follow, not just layout. Agent harnesses may strip HTML comments from copies of that file injected into context, so read .github/pull_request_template.md from disk before writing a PR body to make sure you see every comment rule
|
||||
|
||||
Same applies for filing bug reports and feature requests, with .github/ISSUE_TEMPLATE/bug_report.yml and .github/ISSUE_TEMPLATE/feature_request.yml, respectively
|
||||
|
||||
If you're resolving a linear ticket, in the "## Linear ticket" section of the PR, say "Resolves LIT-1234", replacing "LIT-1234" with the actual ticket id that you're resolving. If you don't have the ticket id, don't make one up or search for it. Just leave the section blank
|
||||
|
||||
Never use `pytest` commands or the like as "Screenshots / Proof of Fix". We prefer curl'ing a live proxy instance running on localhost:4000 (I like to run it with `python litellm/proxy/proxy_cli.py --config litellm/proxy/dev_config.yaml --detailed_debug --reload --use_v2_migration_resolver 2>&1 | tee litellm.log`; the Admin UI dev server is `npm run dev` in `ui/litellm-dashboard`, served on port 3000) and showing both the command run and the output. Also, it should hit real LLM provider APIs, not mocks, and cost real $$$ because that is the most realistic test. The proof of fix should be exactly what the end user / customer would see / do. The run logs in PR #27703 is a prime example of how to do it (not a huge fan of using a python test script that future me and the team will have no visibility into; I prefer just curl commands or a short list of bash commands (e.g., using `for`)). If it's a UI thing, or the main use case runs through a headful agentic coding tool like Claude Code or Codex, drive that surface yourself and embed your own before and after screenshots of it in the PR (the Admin UI page, or what the coding tool shows), next to an ordered list of the URLs to go to (e.g., http://localhost:4000/ui/?page=logs), where to click, and what fields to fill out so a reviewer can reproduce it
|
||||
|
||||
If you ever write any human-facing text (pull requests, issues, commit messages, discussion posts, github comments, release notes, docs, etc.), always follow these guidelines to sound less AI-y:
|
||||
- don't use emojis
|
||||
- don't use "—". Instead, reach for ",", ".", conjunction words, ":", ";", etc. in descending order of preference: vary among them, weighted toward the front of the list, and skip "," where it would cause a comma splice or the sentence is getting long. Overusing any one of them, ";" especially, also feels AI-y. A word cap does not penalize you for adding more sentences: when writing under tight word budgets, prefer a period split or a conjunction over ";", and keep to at most one ";" per message
|
||||
- don't use the pattern "It's not X, it's Y", "You're not X, you're Y", etc.
|
||||
- unless explicitly asked, don't use bulleted or numbered lists unless it would be nonsensical not to. Instead, prefer prose
|
||||
- don't add a trailing "." at the end of paragraphs (just like this file). That means every paragraph, not just the last one (of the markdown file, PR description, GitHub comment, etc.). Rule of thumb: if you're adding new line(s) before the next sentence, don't add a "."
|
||||
- don't use →. Instead, prefer not to use arrows, and if need be, use -> instead
|
||||
- use plain, simple, everyday engineering language: the common phrase engineers actually say over rare compact phrasing, in grammatically complete sentences. When explicitly asked to use bullets or ordered lists and structure legitimately helps the reader, prefer nested bullets (any depth is fine) over dense lines in a flat structure
|
||||
|
||||
Don't hesitate to use values in .env to get needed API keys and other secrets, as long as you never add them to conversation history, commit them, or include them in GitHub issues / PRs
|
||||
|
||||
Python max line length is 120, not 88
|
||||
|
||||
Never edit or commit `ruff-strict-budget.json`, `type-discipline-budget.json`, `basedpyright-code-budget.json`, or `test-quality-budget.json` on a PR branch, and don't run `make lint-budget-update` there. A scheduled Devin automation lowers the limits on the default branch in its own PR by exactly what landed since the last ratchet, so concurrent PRs don't fight over the same `"limit"` lines. Keep the hosted automation's target in sync when the repository default changes. If your branch already carries a budget edit, drop it before opening the PR
|
||||
|
||||
`make check` (f.k.a. `make pre-commit`, which still works identically as an alias) saves its complete output to a log file in .git (overwriting previous logs) and prints that path as its first and last output lines. To inspect a run, read or grep that log instead of re-running the multi-minute checks just to see a different slice
|
||||
|
||||
`make check`, `make lint`, `scripts/pre_commit_lint.sh`, and the standalone budget gates (`scripts/ruff_strict_gate.py`, `scripts/type_discipline_gate.py`, `scripts/type_check_gate.py`) each hold one of 2 machine-wide slots, so when other sessions or worktrees on the same box are already running heavy work, yours prints "all N machine-wide slots are busy; queueing" and then stays quiet until a slot frees. Give the command a long timeout and let it wait rather than killing it, retrying it, or assuming it hung. Don't change the # of machine-wide slots or make it unlimited by setting `LITELLM_GATE_SLOTS=0`
|
||||
|
||||
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
|
||||
|
||||
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()` / `MappingProxyType()` / `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
|
||||
|
||||
Every lint or type suppression must name the exact rule inside brackets and carry a reason comment, e.g. `# pyright: ignore[reportArgumentType] # stubs lack async overload` or `# noqa: TID251 # <reason>`. `# type: ignore` is banned (LIT009): pyrightconfig.json sets `enableTypeIgnoreComments` to false, so it silently does nothing
|
||||
|
||||
Commit and push your work when you're done without asking
|
||||
|
||||
When referencing or running models (coding, QA'ing, writing docs, writing tests, etc.), use the latest model in that model family unless otherwise specified; treat your training knowledge, memories, configs, and tests as stale, and determine the family's latest with model_prices_and_context_window.json or the web
|
||||
|
||||
Always pull before starting any work. The checkout or worktree may be sitting on a stale branch
|
||||
|
||||
If you're an internal contributor, when creating a new PR, the typical flow is to branch off the repository's current default branch and create a branch prefixed with litellm_. Do not create a branch prefixed with claude/ and generally do not have / in your branch names
|
||||
|
||||
Do not add `Co-Authored-By: Claude` or any Claude attribution to commit messages. Never use a `claude/` prefix or put a `/` in a branch name. Do not add "Generated with Claude Code" (or any similar attribution) to PR descriptions or comments. Do not create a new PR/branch off the existing PR to fix/add something that is related and could've just been committed directly to the existing PR's branch
|
||||
|
||||
When working on a PR, keep the PR description in sync with new commits being made
|
||||
|
||||
All GitHub comments must be human-readable and 15-25 words max
|
||||
|
||||
Monkeypatching attributes of a class to do testing is an anti-pattern. Prefer dependency-injecting things into classes. That way, at unit test time, you can pass a mocked dependency in
|
||||
|
||||
Do not put names of customers or customer company names in code, PR descriptions, issue bodies, etc. This means never mention literally any company name. Especially if you're about to say a sentence mentioning that the reason the PR exists was a feature/model/bug fix/etc. requested by a company. That's the indication that you should replace that company name with "the customer". e.g. not "Model request from Acme (Pylon #1234)" but "Model request from a customer (Pylon #1234)". This is because the codebase is public. The only exception is for publicly known providers or vendors such as OpenAI, Anthropic, AWS Bedrock, etc. only IF we're adding support for that provider/vendor in general and NOT if that PR or whatnot was a request by one of them, and they're actually one of our customers.
|
||||
|
||||
CI supply-chain safety: Never pipe a remote script into a shell (`curl ... | bash`, `wget ... | sh`); download the artifact to a file, verify its SHA-256 checksum, then install. Pin every external tool to a specific version with a full URL (not `latest` or `stable`). Verify checksums for all downloaded binaries, using the provider's official `.sha256` / `.sha256sum` sidecar when available. These rules apply to every download in CI
|
||||
|
||||
Prisma migrations apply synchronously at proxy boot, before it serves traffic, so a migration must only change schema, never rewrite rows. No `UPDATE`, `DELETE` or `MERGE`, and no `INSERT ... SELECT`: on a spend-log-sized table any of those is minutes of downtime plus a doubled heap that plain autovacuum won't give back. `tests/code_coverage_tests/check_migrations_no_data_rewrites.py` enforces this. When a rewrite is genuinely bounded and has to ship inside the migration, mark the statement `-- data-migration-ok: <what bounds it>`
|
||||
|
||||
Follow these coding conventions for new/updated code (a three-line fix in a legacy file shouldn't trigger huge drive-by refactors):
|
||||
|
||||
- Composition over inheritance
|
||||
- Never-nester: early returns over deep nesting
|
||||
- Don't throw; model failures as values (One function (e.g., raise_public) maps error union to existing public exception contracts via exhaustive match + assert_never)
|
||||
- No mutation; don't reassign variables, global or local. Instead of mutable lists and dicts, prefer tuples, frozen dataclasses (with slots=True), `MappingProxyType`, etc.
|
||||
- Annotate every variable with `: Final` (LIT010). Unpacking and walrus targets cannot carry the annotation, so they are implicitly final. Don't rebind them. Never rebind or mutate function parameters (LIT011); `self`/`cls` attribute stores are the exception. If rebinding or in-place mutation is truly unavoidable, suppress with `# rebind-ok: <reason>`
|
||||
- Qualify every TypedDict field with `ReadOnly[...]` (LIT012), which nests freely with `Required` / `NotRequired` / `Annotated` in any order. If making the key writable is truly unavoidable, suppress with `# writable-ok: <reason>`
|
||||
- Use dependency injection
|
||||
- Fully typed; no `Any` or coarse types like `dict[str, Any]` or just `dict`. Every function parameter must be strongly typed
|
||||
- Use tagged unions + match
|
||||
- No monster files or god objects
|
||||
- No file sprawl: deliberate file and folder structure
|
||||
- Standard over hand-rolled: use the official SDK or a library where one exists; where none does, follow industry standards instead of inventing local conventions
|
||||
- API-fragmentation-aware: when logic must branch on which API surface produced or consumes data (e.g. chat completions vs Anthropic Messages vs Responses API shapes), proactively look for an existing shared helper (e.g. `litellm_core_utils/prompt_templates/factory.py`) before writing per-surface parsing in the new module; if none exists, add one there instead of duplicating the same format-detection logic in every new guardrail/integration
|
||||
|
||||
Follow conventional commits for commit names and PR titles
|
||||
|
||||
## Think Before Coding
|
||||
|
||||
**Don't assume. Don't hide confusion. Surface tradeoffs**
|
||||
|
||||
Before implementing:
|
||||
- State your assumptions explicitly. If uncertain, ask
|
||||
- If multiple interpretations exist, present them. Don't pick silently
|
||||
- If a simpler approach exists, say so. Push back when warranted
|
||||
- If something is unclear, stop. Name what's confusing. Ask
|
||||
|
||||
## Simplicity First
|
||||
|
||||
**Minimum code that solves the problem. Nothing speculative**
|
||||
|
||||
- No features beyond what was asked
|
||||
- No abstractions for single-use code
|
||||
- No "flexibility" or "configurability" that wasn't requested
|
||||
- No error handling for impossible scenarios
|
||||
- If you write 200 lines and it could be 50, rewrite it
|
||||
|
||||
Ask yourself: "Would a senior engineer say this is overcomplicated?" If yes, simplify
|
||||
|
||||
Before requesting maintainer review, verify the current PR tip passes required CI and code coverage, meets Greptile confidence of at least 4/5, and has acceptable Veria and Bugbot reviews. Inspect warnings and findings, fix actionable issues, and rerun the affected checks and reviewers after changes. Record evidence for any false positive or unavailable review; never treat a pending or missing bot result as a pass. Do not lower coverage thresholds or lint budgets to satisfy a check
|
||||
|
|
|
|||
129
CLAUDE.md
129
CLAUDE.md
|
|
@ -1,129 +0,0 @@
|
|||
Do not write comments unless they are any of:
|
||||
- absolutely necessary to explain some very complex business logic (in which case, keep it concise and clear)
|
||||
- used as an input for tools to read and act on. For example:
|
||||
- entries in `.git-blame-ignore-revs` saying which commit is excluded from git blame
|
||||
- a lint or type checker suppression like `# mutable-ok` or `# pyright: ignore[reportArgumentType] # <reason>` when introducing a truly unavoidable violation
|
||||
- a TODO or FIXME
|
||||
- Not great to have those, but if it's unavoidable, make sure to include a strong, concise reason for why it's there or, better yet, link to a GitHub issue for the follow-up work
|
||||
|
||||
Explanation: The point of this rule is to keep out AI slop comments. AI writes way too many and way too verbose comments. Code comments are, in a way, a violation of DRY code. You must update logic in two locations to change the code, and "hard to change" is literally the definition of tech debt. We should instead aim to write code that is intuitive and clear, even at a glance, to the reader, being both easy to maintain and high performance
|
||||
|
||||
Don't assume that the existing code is correct or the right way of doing things / good coding patterns. In fact, there are a lot of bad coding practices, overly complex code, code smells, etc. If something doesn't look right, speak up. Feel free to break existing patterns or question weird existing code to make new code high quality, as in:
|
||||
|
||||
- correct
|
||||
- secure
|
||||
- performant
|
||||
- readable
|
||||
- easy to maintain/change
|
||||
- modern
|
||||
|
||||
In descending order of importance
|
||||
|
||||
When adding new features, add meaningful tests. Don't add tests that don't check anything substantial and is there just to make the code coverage pass. Yes, code coverage is important, but I'd rather have no signal whether the code is working than tests that don't fail when code is broken. The goal is to have tests that would fail before the feature was added/if the code was mutated in a way that breaks the feature and succeed only when the feature is fully working. I should run mutation testing and see > 90% kill rate
|
||||
|
||||
Same thing for bug fixes. The tests should make it so that this specific bug can never happen again without failing tests (i.e., regression)
|
||||
|
||||
Never test structure of code only function of it
|
||||
|
||||
A test must only fail when litellm code changes. Never pin facts we don't own (a vendor's price, a third party's field, an upstream default, today's date) as literals or as "X must be absent"; assert the invariant our code guarantees instead, e.g. two rows agree, a value is within range, a field is derived from another. If an outside fact is truly load-bearing, cite its source and date next to the assertion so a reader can tell stale from broken
|
||||
|
||||
`tests/test_litellm/` mirrors `litellm/` in a parallel path (see `tests/test_litellm/readme.md`). Name tests `test_<filename>.py`, but always match the existing test file in the directory you touch — many provider dirs use longer descriptive names (e.g. `test_anthropic_chat_transformation.py`) to avoid ambiguity across sibling folders. For bug fixes, extend the existing mapped test file rather than creating a new one. Only create a new test file for a new feature (provider, endpoint, or transformation module) that has no mapped test yet, following that directory's naming convention (or `test_<filename>.py` if you're the first test there). One focused regression test beats many shallow ones
|
||||
|
||||
End-to-end tests belong in `tests/e2e/` and must follow the harness conventions documented in that directory's `CLAUDE.md`
|
||||
|
||||
When creating PRs, target the repository's current default branch for both internal and external / OSS contributions. Check it with `python3 scripts/default_branch.py --branch` instead of assuming a branch name or relying on cached `origin/HEAD`
|
||||
|
||||
When writing a PR body, treat the comments and imperative instructions inside .github/pull_request_template.md as rules to follow, not just layout. Agent harnesses may strip HTML comments from copies of that file injected into context, so read .github/pull_request_template.md from disk before writing a PR body to make sure you see every comment rule
|
||||
|
||||
Same applies for filing bug reports and feature requests, with .github/ISSUE_TEMPLATE/bug_report.yml and .github/ISSUE_TEMPLATE/feature_request.yml, respectively
|
||||
|
||||
If you're resolving a linear ticket, in the "## Linear ticket" section of the PR, say "Resolves LIT-1234", replacing "LIT-1234" with the actual ticket id that you're resolving. If you don't have the ticket id, don't make one up or search for it. Just leave the section blank
|
||||
|
||||
Never use `pytest` commands or the like as "Screenshots / Proof of Fix". We prefer curl'ing a live proxy instance running on localhost:4000 (I like to run it with `python litellm/proxy/proxy_cli.py --config litellm/proxy/dev_config.yaml --detailed_debug --reload --use_v2_migration_resolver 2>&1 | tee litellm.log`; the Admin UI dev server is `npm run dev` in `ui/litellm-dashboard`, served on port 3000) and showing both the command run and the output. Also, it should hit real LLM provider APIs, not mocks, and cost real $$$ because that is the most realistic test. The proof of fix should be exactly what the end user / customer would see / do. The run logs in PR #27703 is a prime example of how to do it (not a huge fan of using a python test script that future me and the team will have no visibility into; I prefer just curl commands or a short list of bash commands (e.g., using `for`)). If it's a UI thing, or the main use case runs through a headful agentic coding tool like Claude Code or Codex, drive that surface yourself and embed your own before and after screenshots of it in the PR (the Admin UI page, or what the coding tool shows), next to an ordered list of the URLs to go to (e.g., http://localhost:4000/ui/?page=logs), where to click, and what fields to fill out so a reviewer can reproduce it
|
||||
|
||||
If you ever write any human-facing text (pull requests, issues, commit messages, discussion posts, github comments, release notes, docs, etc.), always follow these guidelines to sound less AI-y:
|
||||
- don't use emojis
|
||||
- don't use "—". Instead, reach for ",", ".", conjunction words, ":", ";", etc. in descending order of preference: vary among them, weighted toward the front of the list, and skip "," where it would cause a comma splice or the sentence is getting long. Overusing any one of them, ";" especially, also feels AI-y. A word cap does not penalize you for adding more sentences: when writing under tight word budgets, prefer a period split or a conjunction over ";", and keep to at most one ";" per message
|
||||
- don't use the pattern "It's not X, it's Y", "You're not X, you're Y", etc.
|
||||
- unless explicitly asked, don't use bulleted or numbered lists unless it would be nonsensical not to. Instead, prefer prose
|
||||
- don't add a trailing "." at the end of paragraphs (just like this file). That means every paragraph, not just the last one (of the markdown file, PR description, GitHub comment, etc.). Rule of thumb: if you're adding new line(s) before the next sentence, don't add a "."
|
||||
- don't use →. Instead, prefer not to use arrows, and if need be, use -> instead
|
||||
- use plain, simple, everyday engineering language: the common phrase engineers actually say over rare compact phrasing, in grammatically complete sentences. When explicitly asked to use bullets or ordered lists and structure legitimately helps the reader, prefer nested bullets (any depth is fine) over dense lines in a flat structure
|
||||
|
||||
Don't hesitate to use values in .env to get needed API keys and other secrets, as long as you never add them to conversation history, commit them, or include them in GitHub issues / PRs
|
||||
|
||||
Python max line length is 120, not 88
|
||||
|
||||
Never edit or commit `ruff-strict-budget.json`, `type-discipline-budget.json`, `basedpyright-code-budget.json`, or `test-quality-budget.json` on a PR branch, and don't run `make lint-budget-update` there. A scheduled Devin automation lowers the limits on the default branch in its own PR by exactly what landed since the last ratchet, so concurrent PRs don't fight over the same `"limit"` lines. Keep the hosted automation's target in sync when the repository default changes. If your branch already carries a budget edit, drop it before opening the PR
|
||||
|
||||
`make check` (f.k.a. `make pre-commit`, which still works identically as an alias) saves its complete output to a log file in .git (overwriting previous logs) and prints that path as its first and last output lines. To inspect a run, read or grep that log instead of re-running the multi-minute checks just to see a different slice
|
||||
|
||||
`make check`, `make lint`, `scripts/pre_commit_lint.sh`, and the standalone budget gates (`scripts/ruff_strict_gate.py`, `scripts/type_discipline_gate.py`, `scripts/type_check_gate.py`) each hold one of 2 machine-wide slots, so when other sessions or worktrees on the same box are already running heavy work, yours prints "all N machine-wide slots are busy; queueing" and then stays quiet until a slot frees. Give the command a long timeout and let it wait rather than killing it, retrying it, or assuming it hung. Don't change the # of machine-wide slots or make it unlimited by setting `LITELLM_GATE_SLOTS=0`
|
||||
|
||||
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
|
||||
|
||||
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()` / `MappingProxyType()` / `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
|
||||
|
||||
Every lint or type suppression must name the exact rule inside brackets and carry a reason comment, e.g. `# pyright: ignore[reportArgumentType] # stubs lack async overload` or `# noqa: TID251 # <reason>`. `# type: ignore` is banned (LIT009): pyrightconfig.json sets `enableTypeIgnoreComments` to false, so it silently does nothing
|
||||
|
||||
Commit and push your work when you're done without asking
|
||||
|
||||
When referencing or running models (coding, QA'ing, writing docs, writing tests, etc.), use the latest model in that model family unless otherwise specified; treat your training knowledge, memories, configs, and tests as stale, and determine the family's latest with model_prices_and_context_window.json or the web
|
||||
|
||||
Always pull before starting any work. The checkout or worktree may be sitting on a stale branch
|
||||
|
||||
If you're an internal contributor, when creating a new PR, the typical flow is to branch off the repository's current default branch and create a branch prefixed with litellm_. Do not create a branch prefixed with claude/ and generally do not have / in your branch names
|
||||
|
||||
Do not add `Co-Authored-By: Claude` or any Claude attribution to commit messages. Never use a `claude/` prefix or put a `/` in a branch name. Do not add "Generated with Claude Code" (or any similar attribution) to PR descriptions or comments. Do not create a new PR/branch off the existing PR to fix/add something that is related and could've just been committed directly to the existing PR's branch
|
||||
|
||||
When working on a PR, keep the PR description in sync with new commits being made
|
||||
|
||||
All GitHub comments must be human-readable and 15-25 words max
|
||||
|
||||
Monkeypatching attributes of a class to do testing is an anti-pattern. Prefer dependency-injecting things into classes. That way, at unit test time, you can pass a mocked dependency in
|
||||
|
||||
Do not put names of customers or customer company names in code, PR descriptions, issue bodies, etc. This means never mention literally any company name. Especially if you're about to say a sentence mentioning that the reason the PR exists was a feature/model/bug fix/etc. requested by a company. That's the indication that you should replace that company name with "the customer". e.g. not "Model request from Acme (Pylon #1234)" but "Model request from a customer (Pylon #1234)". This is because the codebase is public. The only exception is for publicly known providers or vendors such as OpenAI, Anthropic, AWS Bedrock, etc. only IF we're adding support for that provider/vendor in general and NOT if that PR or whatnot was a request by one of them, and they're actually one of our customers.
|
||||
|
||||
CI supply-chain safety: Never pipe a remote script into a shell (`curl ... | bash`, `wget ... | sh`); download the artifact to a file, verify its SHA-256 checksum, then install. Pin every external tool to a specific version with a full URL (not `latest` or `stable`). Verify checksums for all downloaded binaries, using the provider's official `.sha256` / `.sha256sum` sidecar when available. These rules apply to every download in CI
|
||||
|
||||
Prisma migrations apply synchronously at proxy boot, before it serves traffic, so a migration must only change schema, never rewrite rows. No `UPDATE`, `DELETE` or `MERGE`, and no `INSERT ... SELECT`: on a spend-log-sized table any of those is minutes of downtime plus a doubled heap that plain autovacuum won't give back. `tests/code_coverage_tests/check_migrations_no_data_rewrites.py` enforces this. When a rewrite is genuinely bounded and has to ship inside the migration, mark the statement `-- data-migration-ok: <what bounds it>`
|
||||
|
||||
Follow these coding conventions for new/updated code (a three-line fix in a legacy file shouldn't trigger huge drive-by refactors):
|
||||
|
||||
- Composition over inheritance
|
||||
- Never-nester: early returns over deep nesting
|
||||
- Don't throw; model failures as values (One function (e.g., raise_public) maps error union to existing public exception contracts via exhaustive match + assert_never)
|
||||
- No mutation; don't reassign variables, global or local. Instead of mutable lists and dicts, prefer tuples, frozen dataclasses (with slots=True), `MappingProxyType`, etc.
|
||||
- Annotate every variable with `: Final` (LIT010). Unpacking and walrus targets cannot carry the annotation, so they are implicitly final. Don't rebind them. Never rebind or mutate function parameters (LIT011); `self`/`cls` attribute stores are the exception. If rebinding or in-place mutation is truly unavoidable, suppress with `# rebind-ok: <reason>`
|
||||
- Qualify every TypedDict field with `ReadOnly[...]` (LIT012), which nests freely with `Required` / `NotRequired` / `Annotated` in any order. If making the key writable is truly unavoidable, suppress with `# writable-ok: <reason>`
|
||||
- Use dependency injection
|
||||
- Fully typed; no `Any` or coarse types like `dict[str, Any]` or just `dict`. Every function parameter must be strongly typed
|
||||
- Use tagged unions + match
|
||||
- No monster files or god objects
|
||||
- No file sprawl: deliberate file and folder structure
|
||||
- Standard over hand-rolled: use the official SDK or a library where one exists; where none does, follow industry standards instead of inventing local conventions
|
||||
- API-fragmentation-aware: when logic must branch on which API surface produced or consumes data (e.g. chat completions vs Anthropic Messages vs Responses API shapes), proactively look for an existing shared helper (e.g. `litellm_core_utils/prompt_templates/factory.py`) before writing per-surface parsing in the new module; if none exists, add one there instead of duplicating the same format-detection logic in every new guardrail/integration
|
||||
|
||||
Follow conventional commits for commit names and PR titles
|
||||
|
||||
## Think Before Coding
|
||||
|
||||
**Don't assume. Don't hide confusion. Surface tradeoffs**
|
||||
|
||||
Before implementing:
|
||||
- State your assumptions explicitly. If uncertain, ask
|
||||
- If multiple interpretations exist, present them. Don't pick silently
|
||||
- If a simpler approach exists, say so. Push back when warranted
|
||||
- If something is unclear, stop. Name what's confusing. Ask
|
||||
|
||||
## Simplicity First
|
||||
|
||||
**Minimum code that solves the problem. Nothing speculative**
|
||||
|
||||
- No features beyond what was asked
|
||||
- No abstractions for single-use code
|
||||
- No "flexibility" or "configurability" that wasn't requested
|
||||
- No error handling for impossible scenarios
|
||||
- If you write 200 lines and it could be 50, rewrite it
|
||||
|
||||
Ask yourself: "Would a senior engineer say this is overcomplicated?" If yes, simplify
|
||||
|
|
@ -148,7 +148,7 @@ make lint
|
|||
|
||||
Individual linting commands:
|
||||
```bash
|
||||
make format-check # Check Black formatting
|
||||
make format-check # Check ruff format formatting
|
||||
make lint-ruff # Run Ruff linting
|
||||
make lint-basedpyright # Run basedpyright type checking
|
||||
make check-circular-imports # Check for circular imports
|
||||
|
|
@ -160,14 +160,14 @@ Apply formatting (auto-fixes issues):
|
|||
make format
|
||||
```
|
||||
|
||||
> **Black formatting is enforced in CI.** All PRs must pass the Black formatting check.
|
||||
> **Formatting is enforced in CI.** All PRs must pass the `ruff format --check` step.
|
||||
>
|
||||
> - **AI coding agents** (Claude Code, Copilot, Cursor, etc.): `AGENTS.md` and `CLAUDE.md` instruct agents to run `poetry run black .` before committing.
|
||||
> - **VS Code users**: Install the [Black Formatter extension](https://marketplace.visualstudio.com/items?itemName=ms-python.black-formatter) and enable format-on-save:
|
||||
> - **AI coding agents** (Claude Code, Copilot, Cursor, etc.): follow `AGENTS.md` and run `make format` before committing.
|
||||
> - **VS Code users**: Install the [Ruff extension](https://marketplace.visualstudio.com/items?itemName=charliermarsh.ruff) and enable format-on-save:
|
||||
> ```json
|
||||
> {
|
||||
> "[python]": {
|
||||
> "editor.defaultFormatter": "ms-python.black-formatter",
|
||||
> "editor.defaultFormatter": "charliermarsh.ruff",
|
||||
> "editor.formatOnSave": true
|
||||
> }
|
||||
> }
|
||||
|
|
@ -197,8 +197,8 @@ make help # Show all available commands
|
|||
make install-dev # Install development dependencies
|
||||
make install-proxy-dev # Install proxy development dependencies
|
||||
make install-test-deps # Install the full local test environment
|
||||
make format # Apply Black code formatting
|
||||
make format-check # Check Black formatting (matches CI)
|
||||
make format # Apply ruff format code formatting
|
||||
make format-check # Check ruff format formatting (matches CI)
|
||||
make lint # Run all linting checks
|
||||
make test-unit # Run unit tests
|
||||
make test-integration # Run integration tests
|
||||
|
|
@ -210,8 +210,7 @@ make test-unit-helm # Run Helm unit tests
|
|||
LiteLLM follows the [Google Python Style Guide](https://google.github.io/styleguide/pyguide.html).
|
||||
|
||||
Our automated quality checks include:
|
||||
- **Black** for consistent code formatting
|
||||
- **Ruff** for linting and code quality
|
||||
- **Ruff** for formatting, linting, and code quality
|
||||
- **basedpyright** for static type checking
|
||||
- **Circular import detection**
|
||||
- **Import safety validation**
|
||||
|
|
|
|||
|
|
@ -1 +1 @@
|
|||
Read @CLAUDE.md for coding guidelines
|
||||
Read @AGENTS.md for coding guidelines
|
||||
|
|
|
|||
|
|
@ -633,9 +633,8 @@ For detailed contributing guidelines, see [CONTRIBUTING.md](CONTRIBUTING.md).
|
|||
LiteLLM follows the [Google Python Style Guide](https://google.github.io/styleguide/pyguide.html).
|
||||
|
||||
Our automated checks include:
|
||||
- **Black** for code formatting
|
||||
- **Ruff** for linting and code quality
|
||||
- **MyPy** for type checking
|
||||
- **Ruff** for formatting, linting, and code quality
|
||||
- **basedpyright** for type checking
|
||||
- **Circular import detection**
|
||||
- **Import safety checks**
|
||||
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
- Target invariants, not completion claims; these supersede older conflicting bridge guidance
|
||||
- Target invariants, not completion claims; these supersede the crate guidance below where they conflict
|
||||
- Keep this crate the product-specific PyO3 consumer of `litellm-host-python`
|
||||
- Own registration, input projection, the route host and the caller callables it answers operations with (file readers, token providers), public response/error construction and the per-call composition of machine, route host and callback contract
|
||||
- Legacy callback sharing (the caller's args, kwargs and request object, body/header roots, re-aliasing unchanged body keys) lives in `litellm-callbacks-legacy` behind `PublicCall` and `run_legacy_call`; the bridge hands the public call over and keeps no copy
|
||||
|
|
@ -34,3 +34,45 @@
|
|||
- References: [ownership](https://pyo3.rs/v0.29.2/types.html), [GC](https://pyo3.rs/v0.29.2/class/protocols.html#garbage-collector-integration), [exception transfer](https://docs.rs/pyo3/0.29.2/pyo3/struct.PyErr.html#method.into_value), [re-entry](https://pyo3.rs/v0.29.2/class/call.html)
|
||||
- [GIL policy](https://pyo3.rs/v0.29.2/free-threading.html), [experimental async limits](https://pyo3.rs/v0.29.2/async-await.html), [task conversion](https://docs.rs/pyo3-async-runtimes/0.29.0/pyo3_async_runtimes/fn.into_future_with_locals.html), [native cancellation/delivery](https://docs.rs/pyo3-async-runtimes/0.29.0/pyo3_async_runtimes/tokio/fn.future_into_py.html)
|
||||
- [performance](https://pyo3.rs/v0.29.2/performance.html), [PyBackedBytes](https://docs.rs/pyo3/0.29.2/pyo3/pybacked/struct.PyBackedBytes.html), [typing](https://pyo3.rs/v0.29.2/python-typing-hints.html)
|
||||
|
||||
Rules for `litellm-rust/crates/python-bridge`.
|
||||
|
||||
## Responsibility
|
||||
|
||||
`python-bridge` is the PyO3 boundary between Python LiteLLM and Rust transforms.
|
||||
Keep this crate thin. It exposes LiteLLM Rust APIs, assembles domain requests,
|
||||
maps domain errors to Python exceptions, and delegates generic conversion and
|
||||
GIL handling to `litellm-host-python`.
|
||||
|
||||
## Bridge Shape
|
||||
|
||||
- Prefer one stable method per top-level LiteLLM route, for example
|
||||
`messages(...)`, calling the matching `litellm-core` entrypoint.
|
||||
- Do not add one exported PyO3 function per provider helper unless there is a
|
||||
measured reason.
|
||||
- Provider dispatch belongs in the `litellm-core` route module (e.g.
|
||||
`litellm_core::messages`), not in this PyO3 crate.
|
||||
- Python owns rollout state and fallback. Rust should return errors; Python
|
||||
decides whether to raise or fall back. For a rust-only provider/route (no
|
||||
Python reference), the Python side is a thin dispatch that calls Rust and
|
||||
raises when the bridge is unavailable, with no fallback.
|
||||
- Keep the Python interface minimal (well under 100 lines per route): it only
|
||||
marshals inputs and calls Rust. Do not add per-route feature flags, and do
|
||||
not put provider dispatch in `litellm/main.py`; it lives in a thin dispatch
|
||||
class under `litellm/llms/<provider>/<route>/`.
|
||||
|
||||
## Data Handling
|
||||
|
||||
- OCR payloads can contain personal data and large base64 images. Do not log
|
||||
payloads or provider responses.
|
||||
- Avoid copying large payloads more than needed. The current JSON round-trip is
|
||||
acceptable for the first scaffold, but future performance work should evaluate
|
||||
direct PyO3 conversion before expanding Rust coverage to image-heavy paths.
|
||||
- Do not expose raw Rust errors that include document contents or upstream
|
||||
bodies.
|
||||
|
||||
## Tests
|
||||
|
||||
- `cargo test --workspace` must compile this crate.
|
||||
- Python tests must cover bridge disabled, bridge enabled, and module-missing
|
||||
fallback behavior for every exposed route.
|
||||
|
|
|
|||
|
|
@ -1,43 +0,0 @@
|
|||
# CLAUDE.md
|
||||
|
||||
Rules for `litellm-rust/crates/python-bridge`.
|
||||
|
||||
## Responsibility
|
||||
|
||||
`python-bridge` is the PyO3 boundary between Python LiteLLM and Rust transforms.
|
||||
Keep this crate thin. It exposes LiteLLM Rust APIs, assembles domain requests,
|
||||
maps domain errors to Python exceptions, and delegates generic conversion and
|
||||
GIL handling to `litellm-host-python`.
|
||||
|
||||
## Bridge Shape
|
||||
|
||||
- Prefer one stable method per top-level LiteLLM route, for example
|
||||
`messages(...)`, calling the matching `litellm-core` entrypoint.
|
||||
- Do not add one exported PyO3 function per provider helper unless there is a
|
||||
measured reason.
|
||||
- Provider dispatch belongs in the `litellm-core` route module (e.g.
|
||||
`litellm_core::messages`), not in this PyO3 crate.
|
||||
- Python owns rollout state and fallback. Rust should return errors; Python
|
||||
decides whether to raise or fall back. For a rust-only provider/route (no
|
||||
Python reference), the Python side is a thin dispatch that calls Rust and
|
||||
raises when the bridge is unavailable, with no fallback.
|
||||
- Keep the Python interface minimal (well under 100 lines per route): it only
|
||||
marshals inputs and calls Rust. Do not add per-route feature flags, and do
|
||||
not put provider dispatch in `litellm/main.py`; it lives in a thin dispatch
|
||||
class under `litellm/llms/<provider>/<route>/`.
|
||||
|
||||
## Data Handling
|
||||
|
||||
- OCR payloads can contain personal data and large base64 images. Do not log
|
||||
payloads or provider responses.
|
||||
- Avoid copying large payloads more than needed. The current JSON round-trip is
|
||||
acceptable for the first scaffold, but future performance work should evaluate
|
||||
direct PyO3 conversion before expanding Rust coverage to image-heavy paths.
|
||||
- Do not expose raw Rust errors that include document contents or upstream
|
||||
bodies.
|
||||
|
||||
## Tests
|
||||
|
||||
- `cargo test --workspace` must compile this crate.
|
||||
- Python tests must cover bridge disabled, bridge enabled, and module-missing
|
||||
fallback behavior for every exposed route.
|
||||
|
|
@ -12,7 +12,7 @@ import litellm
|
|||
from litellm._logging import redact_internal_details_from_client_message, verbose_logger
|
||||
from litellm.constants import REALTIME_SESSION_FAILURE_LOGGED_KEY, REALTIME_SESSION_SUCCESS_LOGGED_KEY
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
|
||||
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig, RealtimeBackend
|
||||
from litellm.types.llms.openai import (
|
||||
OpenAIRealtimeEvents,
|
||||
OpenAIRealtimeOutputItemDone,
|
||||
|
|
@ -127,7 +127,7 @@ class RealTimeStreaming:
|
|||
def __init__(
|
||||
self,
|
||||
websocket: Any,
|
||||
backend_ws: CLIENT_CONNECTION_CLASS,
|
||||
backend_ws: CLIENT_CONNECTION_CLASS | RealtimeBackend,
|
||||
logging_obj: LiteLLMLogging,
|
||||
provider_config: BaseRealtimeConfig | None = None,
|
||||
model: str = "",
|
||||
|
|
|
|||
275
litellm/llms/base_llm/realtime/transcription_protocol.py
Normal file
275
litellm/llms/base_llm/realtime/transcription_protocol.py
Normal file
|
|
@ -0,0 +1,275 @@
|
|||
import base64
|
||||
import binascii
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal
|
||||
|
||||
from pydantic import JsonValue, TypeAdapter, ValidationError
|
||||
|
||||
from litellm._uuid import uuid
|
||||
from litellm.types.llms.openai import (
|
||||
OpenAIRealtimeErrorEvent,
|
||||
OpenAIRealtimeInputAudioBufferSpeechEvent,
|
||||
OpenAIRealtimeInputAudioTranscriptionCompleted,
|
||||
OpenAIRealtimeInputAudioTranscriptionDelta,
|
||||
OpenAIRealtimeServerVadTurnDetection,
|
||||
OpenAIRealtimeTranscriptionSession,
|
||||
OpenAIRealtimeTranscriptionSessionCreated,
|
||||
OpenAIRealtimeTranscriptionSettings,
|
||||
)
|
||||
from litellm.types.realtime import RealtimeInputAudioTranscriptionDurationUsage, RealtimeInputAudioTranscriptionUsage
|
||||
|
||||
SESSION_UPDATE_EVENT_TYPES: Final = frozenset(("session.update", "transcription_session.update"))
|
||||
PCM16_ENCODINGS: Final = frozenset(("pcm16", "audio/pcm"))
|
||||
SERVER_VAD_TURN_DETECTION: Final[OpenAIRealtimeServerVadTurnDetection] = {"type": "server_vad"}
|
||||
EMPTY_JSON_OBJECT: Final[Mapping[str, JsonValue]] = MappingProxyType({})
|
||||
_SUPPORTED_TRANSCRIPTION_KEYS: Final = frozenset(("model", "language"))
|
||||
_JSON_ADAPTER: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
|
||||
|
||||
|
||||
class RealtimeTranscriptionProtocolError(ValueError):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TranscriptionAudioFormat:
|
||||
layout: Literal["beta", "ga"]
|
||||
encoding: str | None
|
||||
rate: int | None
|
||||
channels: int | None
|
||||
|
||||
@property
|
||||
def is_pcm16(self) -> bool:
|
||||
return self.encoding in PCM16_ENCODINGS
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class TranscriptionSessionUpdate:
|
||||
session_type: str | None
|
||||
audio_format: TranscriptionAudioFormat | None
|
||||
model: str | None
|
||||
language: str | None
|
||||
unsupported_transcription_keys: tuple[str, ...]
|
||||
turn_detection: Mapping[str, JsonValue] | None
|
||||
turn_detection_disabled: bool
|
||||
|
||||
@property
|
||||
def turn_detection_type(self) -> JsonValue | None:
|
||||
return None if self.turn_detection is None else self.turn_detection.get("type")
|
||||
|
||||
|
||||
ProtocolErrorType = type[RealtimeTranscriptionProtocolError]
|
||||
|
||||
|
||||
def json_object(payload: str, error: ProtocolErrorType = RealtimeTranscriptionProtocolError) -> Mapping[str, JsonValue]:
|
||||
try:
|
||||
value: Final = _JSON_ADAPTER.validate_json(payload)
|
||||
except ValidationError:
|
||||
raise error("invalid JSON object") from None
|
||||
if not isinstance(value, dict):
|
||||
raise error("message must be a JSON object")
|
||||
return value
|
||||
|
||||
|
||||
def json_mapping(
|
||||
value: JsonValue | None, name: str, error: ProtocolErrorType = RealtimeTranscriptionProtocolError
|
||||
) -> Mapping[str, JsonValue]:
|
||||
if value is None:
|
||||
return EMPTY_JSON_OBJECT
|
||||
if not isinstance(value, dict):
|
||||
raise error(f"{name} must be an object")
|
||||
return value
|
||||
|
||||
|
||||
def json_string(
|
||||
value: JsonValue | None, name: str, error: ProtocolErrorType = RealtimeTranscriptionProtocolError
|
||||
) -> str | None:
|
||||
if value is None:
|
||||
return None
|
||||
if not isinstance(value, str):
|
||||
raise error(f"{name} must be a string")
|
||||
return value
|
||||
|
||||
|
||||
def json_integer(
|
||||
value: JsonValue | None, name: str, error: ProtocolErrorType = RealtimeTranscriptionProtocolError
|
||||
) -> int | None:
|
||||
if value is None:
|
||||
return None
|
||||
if isinstance(value, bool) or not isinstance(value, int):
|
||||
raise error(f"{name} must be an integer")
|
||||
return value
|
||||
|
||||
|
||||
def new_event_id() -> str:
|
||||
return f"event_{uuid.uuid4().hex}"
|
||||
|
||||
|
||||
def parse_transcription_session_update(
|
||||
payload: str, error: ProtocolErrorType = RealtimeTranscriptionProtocolError
|
||||
) -> TranscriptionSessionUpdate:
|
||||
message: Final = json_object(payload, error)
|
||||
if message.get("type") not in SESSION_UPDATE_EVENT_TYPES:
|
||||
raise error("expected session.update")
|
||||
session: Final = json_mapping(message.get("session"), "session", error)
|
||||
if not session:
|
||||
raise error("session.update requires a session object")
|
||||
audio: Final = json_mapping(session.get("audio"), "session.audio", error)
|
||||
audio_input: Final = json_mapping(audio.get("input"), "session.audio.input", error)
|
||||
beta_transcription: Final = session.get("input_audio_transcription")
|
||||
ga_transcription: Final = audio_input.get("transcription")
|
||||
if beta_transcription is not None and ga_transcription is not None:
|
||||
raise error("input transcription must use either beta or GA layout")
|
||||
transcription: Final = json_mapping(
|
||||
beta_transcription if beta_transcription is not None else ga_transcription,
|
||||
"input audio transcription",
|
||||
error,
|
||||
)
|
||||
turn_detection_present: Final = "turn_detection" in session or "turn_detection" in audio_input
|
||||
turn_detection: Final = session.get("turn_detection", audio_input.get("turn_detection"))
|
||||
return TranscriptionSessionUpdate(
|
||||
session_type=json_string(session.get("type"), "session.type", error),
|
||||
audio_format=_parse_audio_format(session, audio_input, error),
|
||||
model=json_string(transcription.get("model"), "transcription model", error),
|
||||
language=json_string(transcription.get("language"), "language", error),
|
||||
unsupported_transcription_keys=tuple(
|
||||
sorted(key for key in transcription if key not in _SUPPORTED_TRANSCRIPTION_KEYS)
|
||||
),
|
||||
turn_detection=None if turn_detection is None else json_mapping(turn_detection, "turn_detection", error),
|
||||
turn_detection_disabled=turn_detection_present and turn_detection is None,
|
||||
)
|
||||
|
||||
|
||||
def _parse_audio_format(
|
||||
session: Mapping[str, JsonValue], audio_input: Mapping[str, JsonValue], error: ProtocolErrorType
|
||||
) -> TranscriptionAudioFormat | None:
|
||||
beta_format: Final = session.get("input_audio_format")
|
||||
ga_format: Final = audio_input.get("format")
|
||||
if beta_format is not None and ga_format is not None:
|
||||
raise error("input audio format must use either beta or GA layout")
|
||||
if beta_format is not None:
|
||||
return TranscriptionAudioFormat(
|
||||
layout="beta",
|
||||
encoding=json_string(beta_format, "session.input_audio_format", error),
|
||||
rate=None,
|
||||
channels=None,
|
||||
)
|
||||
if ga_format is None:
|
||||
return None
|
||||
if isinstance(ga_format, str):
|
||||
return TranscriptionAudioFormat(layout="ga", encoding=ga_format, rate=None, channels=None)
|
||||
format_mapping: Final = json_mapping(ga_format, "session.audio.input.format", error)
|
||||
return TranscriptionAudioFormat(
|
||||
layout="ga",
|
||||
encoding=json_string(format_mapping.get("type"), "session.audio.input.format.type", error),
|
||||
rate=json_integer(format_mapping.get("rate"), "session.audio.input.format.rate", error),
|
||||
channels=json_integer(format_mapping.get("channels"), "session.audio.input.format.channels", error),
|
||||
)
|
||||
|
||||
|
||||
def decode_pcm16_append(
|
||||
audio: JsonValue | None,
|
||||
max_encoded_bytes: int | None = None,
|
||||
error: ProtocolErrorType = RealtimeTranscriptionProtocolError,
|
||||
) -> bytes:
|
||||
if not isinstance(audio, str):
|
||||
raise error("Audio must be a base64 string")
|
||||
if max_encoded_bytes is not None and len(audio) > max_encoded_bytes:
|
||||
raise error("Audio append exceeds the four-second backlog limit")
|
||||
try:
|
||||
decoded: Final = base64.b64decode(audio, validate=True)
|
||||
except (binascii.Error, ValueError):
|
||||
raise error("Audio must be valid base64") from None
|
||||
if len(decoded) % 2:
|
||||
raise error("PCM16 audio must contain complete samples")
|
||||
return decoded
|
||||
|
||||
|
||||
def _transcription_settings(model: str, language: str | None) -> OpenAIRealtimeTranscriptionSettings:
|
||||
if language is None:
|
||||
model_only: Final[OpenAIRealtimeTranscriptionSettings] = {"model": model}
|
||||
return model_only
|
||||
with_language: Final[OpenAIRealtimeTranscriptionSettings] = {"model": model, "language": language}
|
||||
return with_language
|
||||
|
||||
|
||||
def transcription_session(
|
||||
*, session_id: str, model: str, sample_rate: int, language: str | None, server_vad: bool
|
||||
) -> OpenAIRealtimeTranscriptionSession:
|
||||
settings: Final = _transcription_settings(model, language)
|
||||
session: Final[OpenAIRealtimeTranscriptionSession] = {
|
||||
"id": session_id,
|
||||
"object": "realtime.transcription_session",
|
||||
"type": "transcription",
|
||||
"audio": {
|
||||
"input": {
|
||||
"format": {"type": "audio/pcm", "rate": sample_rate},
|
||||
"transcription": settings,
|
||||
"turn_detection": SERVER_VAD_TURN_DETECTION if server_vad else None,
|
||||
}
|
||||
},
|
||||
}
|
||||
return session
|
||||
|
||||
|
||||
def transcription_session_created_event(
|
||||
session: OpenAIRealtimeTranscriptionSession,
|
||||
) -> OpenAIRealtimeTranscriptionSessionCreated:
|
||||
event: Final[OpenAIRealtimeTranscriptionSessionCreated] = {
|
||||
"type": "session.created",
|
||||
"event_id": new_event_id(),
|
||||
"session": session,
|
||||
}
|
||||
return event
|
||||
|
||||
|
||||
def error_event(message: str) -> OpenAIRealtimeErrorEvent:
|
||||
event: Final[OpenAIRealtimeErrorEvent] = {
|
||||
"type": "error",
|
||||
"error": {"type": "server_error", "message": message},
|
||||
}
|
||||
return event
|
||||
|
||||
|
||||
def speech_event(
|
||||
event_type: Literal["input_audio_buffer.speech_started", "input_audio_buffer.speech_stopped"], item_id: str
|
||||
) -> OpenAIRealtimeInputAudioBufferSpeechEvent:
|
||||
event: Final[OpenAIRealtimeInputAudioBufferSpeechEvent] = {
|
||||
"type": event_type,
|
||||
"event_id": new_event_id(),
|
||||
"item_id": item_id,
|
||||
}
|
||||
return event
|
||||
|
||||
|
||||
def delta_event(item_id: str, delta: str) -> OpenAIRealtimeInputAudioTranscriptionDelta:
|
||||
event: Final[OpenAIRealtimeInputAudioTranscriptionDelta] = {
|
||||
"type": "conversation.item.input_audio_transcription.delta",
|
||||
"event_id": new_event_id(),
|
||||
"item_id": item_id,
|
||||
"content_index": 0,
|
||||
"delta": delta,
|
||||
}
|
||||
return event
|
||||
|
||||
|
||||
def completed_event(
|
||||
item_id: str, transcript: str, usage: RealtimeInputAudioTranscriptionUsage | None
|
||||
) -> OpenAIRealtimeInputAudioTranscriptionCompleted:
|
||||
event: Final[OpenAIRealtimeInputAudioTranscriptionCompleted] = {
|
||||
"type": "conversation.item.input_audio_transcription.completed",
|
||||
"event_id": new_event_id(),
|
||||
"item_id": item_id,
|
||||
"content_index": 0,
|
||||
"transcript": transcript,
|
||||
}
|
||||
if usage is None:
|
||||
return event
|
||||
billed: Final[OpenAIRealtimeInputAudioTranscriptionCompleted] = {**event, "usage": usage}
|
||||
return billed
|
||||
|
||||
|
||||
def duration_usage(seconds: float) -> RealtimeInputAudioTranscriptionUsage:
|
||||
usage: Final[RealtimeInputAudioTranscriptionDurationUsage] = {"type": "duration", "seconds": seconds}
|
||||
return usage
|
||||
|
|
@ -1,8 +1,10 @@
|
|||
from abc import ABC, abstractmethod
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from types import TracebackType
|
||||
from typing import TYPE_CHECKING, Any, Protocol
|
||||
|
||||
import httpx
|
||||
from typing_extensions import Self
|
||||
|
||||
from litellm.types.llms.openai import OpenAIRealtimeStreamSessionEvents
|
||||
from litellm.types.realtime import (
|
||||
|
|
@ -21,6 +23,23 @@ else:
|
|||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
class RealtimeBackend(Protocol):
|
||||
async def __aenter__(self) -> Self: ...
|
||||
|
||||
async def __aexit__(
|
||||
self,
|
||||
exc_type: type[BaseException] | None,
|
||||
exc_value: BaseException | None,
|
||||
traceback: TracebackType | None,
|
||||
) -> None: ...
|
||||
|
||||
async def send(self, message: str | bytes) -> None: ...
|
||||
|
||||
async def recv(self, decode: bool | None = None) -> str | bytes: ...
|
||||
|
||||
async def close(self) -> None: ...
|
||||
|
||||
|
||||
class BaseRealtimeConfig(ABC):
|
||||
@abstractmethod
|
||||
def validate_environment(
|
||||
|
|
@ -78,6 +97,9 @@ class BaseRealtimeConfig(ABC):
|
|||
def unbilled_usage_on_session_close(self, model: str) -> RealtimeInputAudioTranscriptionUsage | None:
|
||||
return None
|
||||
|
||||
async def open_backend(self, url: str, headers: Mapping[str, str]) -> RealtimeBackend | None:
|
||||
return None
|
||||
|
||||
def transform_session_created_event(
|
||||
self,
|
||||
model: str,
|
||||
|
|
|
|||
|
|
@ -6287,7 +6287,12 @@ class BaseLLMHTTPHandler:
|
|||
ssl_context = ssl.SSLContext(ssl.PROTOCOL_TLS_CLIENT)
|
||||
ssl_context.check_hostname = False
|
||||
ssl_context.verify_mode = ssl.CERT_NONE
|
||||
backend_ws: Final = await self._open_realtime_backend_ws(websockets, url, headers, ssl_context)
|
||||
provider_backend: Final = await provider_config.open_backend(url, headers)
|
||||
backend_ws: Final = (
|
||||
provider_backend
|
||||
if provider_backend is not None
|
||||
else await self._open_realtime_backend_ws(websockets, url, headers, ssl_context)
|
||||
)
|
||||
async with backend_ws:
|
||||
_request_data: Final[dict[str, object]] = {}
|
||||
if litellm_metadata:
|
||||
|
|
|
|||
|
|
@ -1,36 +1,41 @@
|
|||
import asyncio
|
||||
import base64
|
||||
import binascii
|
||||
import json
|
||||
import math
|
||||
import time
|
||||
from collections.abc import Awaitable, Callable, Iterator, Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final, Literal
|
||||
from typing import Final
|
||||
from urllib.parse import urlparse, urlunparse
|
||||
|
||||
from pydantic import JsonValue, TypeAdapter, ValidationError
|
||||
from pydantic import JsonValue
|
||||
|
||||
from litellm import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.realtime.transcription_protocol import (
|
||||
RealtimeTranscriptionProtocolError,
|
||||
TranscriptionSessionUpdate,
|
||||
completed_event,
|
||||
decode_pcm16_append,
|
||||
delta_event,
|
||||
duration_usage,
|
||||
error_event,
|
||||
json_object,
|
||||
parse_transcription_session_update,
|
||||
speech_event,
|
||||
transcription_session,
|
||||
transcription_session_created_event,
|
||||
)
|
||||
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.meta import MuseAudioEncoding, MuseHandshake, MuseMode, MuseSampleRate
|
||||
from litellm.types.llms.openai import (
|
||||
OpenAIRealtimeErrorEvent,
|
||||
OpenAIRealtimeEvents,
|
||||
OpenAIRealtimeInputAudioBufferSpeechEvent,
|
||||
OpenAIRealtimeInputAudioTranscriptionCompleted,
|
||||
OpenAIRealtimeInputAudioTranscriptionDelta,
|
||||
OpenAIRealtimeServerVadTurnDetection,
|
||||
OpenAIRealtimeTranscriptionSession,
|
||||
OpenAIRealtimeTranscriptionSessionCreated,
|
||||
OpenAIRealtimeTranscriptionSettings,
|
||||
)
|
||||
from litellm.types.realtime import (
|
||||
RealtimeInputAudioTranscriptionDurationUsage,
|
||||
RealtimeInputAudioTranscriptionUsage,
|
||||
RealtimeResponseTransformInput,
|
||||
RealtimeResponseTypedDict,
|
||||
|
|
@ -98,17 +103,13 @@ _LANGUAGE_CODES: Final = MappingProxyType(
|
|||
"zh": "Mandarin Chinese",
|
||||
}
|
||||
)
|
||||
_SUPPORTED_TRANSCRIPTION_KEYS: Final = frozenset(("model", "language"))
|
||||
_MAX_AUDIO_BACKLOG_SECONDS: Final = 4
|
||||
_PACKET_MS: Final = 80
|
||||
_END_STREAM: Final = '{"type":"endStream"}'
|
||||
_PROVIDER_ERROR_MESSAGE: Final = "Meta Muse realtime transcription failed"
|
||||
_JSON_ADAPTER: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue)
|
||||
_EMPTY_OBJECT: Final[Mapping[str, JsonValue]] = MappingProxyType({})
|
||||
_SERVER_VAD: Final[OpenAIRealtimeServerVadTurnDetection] = {"type": "server_vad"}
|
||||
|
||||
|
||||
class MuseProtocolError(ValueError):
|
||||
class MuseProtocolError(RealtimeTranscriptionProtocolError):
|
||||
pass
|
||||
|
||||
|
||||
|
|
@ -150,26 +151,13 @@ class MuseSessionConfig:
|
|||
return biased
|
||||
|
||||
def openai_session(self, session_id: str) -> OpenAIRealtimeTranscriptionSession:
|
||||
session: Final[OpenAIRealtimeTranscriptionSession] = {
|
||||
"id": session_id,
|
||||
"object": "realtime.transcription_session",
|
||||
"type": "transcription",
|
||||
"audio": {
|
||||
"input": {
|
||||
"format": {"type": "audio/pcm", "rate": self.sample_rate},
|
||||
"transcription": self._transcription_settings(),
|
||||
"turn_detection": None if self.mode == "PUSH_TO_TALK" else _SERVER_VAD,
|
||||
}
|
||||
},
|
||||
}
|
||||
return session
|
||||
|
||||
def _transcription_settings(self) -> OpenAIRealtimeTranscriptionSettings:
|
||||
base: Final[OpenAIRealtimeTranscriptionSettings] = {"model": self.model}
|
||||
if not self.language_bias:
|
||||
return base
|
||||
localized: Final[OpenAIRealtimeTranscriptionSettings] = {**base, "language": self.language_bias[0]}
|
||||
return localized
|
||||
return transcription_session(
|
||||
session_id=session_id,
|
||||
model=self.model,
|
||||
sample_rate=self.sample_rate,
|
||||
language=self.language_bias[0] if self.language_bias else None,
|
||||
server_vad=self.mode != "PUSH_TO_TALK",
|
||||
)
|
||||
|
||||
|
||||
_DEFAULT_SESSION_CONFIG: Final = MuseSessionConfig(
|
||||
|
|
@ -177,40 +165,10 @@ _DEFAULT_SESSION_CONFIG: Final = MuseSessionConfig(
|
|||
)
|
||||
|
||||
|
||||
def _json_object(payload: str) -> Mapping[str, JsonValue]:
|
||||
try:
|
||||
value: Final = _JSON_ADAPTER.validate_json(payload)
|
||||
except ValidationError:
|
||||
raise MuseProtocolError("invalid JSON object") from None
|
||||
if not isinstance(value, dict):
|
||||
raise MuseProtocolError("message must be a JSON object")
|
||||
return value
|
||||
|
||||
|
||||
def _mapping(value: JsonValue | None, name: str) -> Mapping[str, JsonValue]:
|
||||
if value is None:
|
||||
return _EMPTY_OBJECT
|
||||
if not isinstance(value, dict):
|
||||
raise MuseProtocolError(f"{name} must be an object")
|
||||
return value
|
||||
|
||||
|
||||
def _string(value: JsonValue | None, name: str) -> str | None:
|
||||
if value is None:
|
||||
return None
|
||||
if not isinstance(value, str):
|
||||
raise MuseProtocolError(f"{name} must be a string")
|
||||
return value
|
||||
|
||||
|
||||
def _normalize_model(model: str) -> str:
|
||||
return model.removeprefix("meta/").strip()
|
||||
|
||||
|
||||
def _event_id() -> str:
|
||||
return f"event_{uuid.uuid4().hex}"
|
||||
|
||||
|
||||
def normalize_language(language: str) -> str:
|
||||
value: Final = language.strip()
|
||||
if not value:
|
||||
|
|
@ -254,138 +212,55 @@ def build_muse_realtime_url(api_base: str | None) -> str:
|
|||
return urlunparse((scheme, netloc, "/v1/asr/realtime", "", "", ""))
|
||||
|
||||
|
||||
def _parse_sample_rate(session: Mapping[str, JsonValue]) -> MuseSampleRate:
|
||||
beta_format: Final = session.get("input_audio_format")
|
||||
audio: Final = _mapping(session.get("audio"), "session.audio")
|
||||
audio_input: Final = _mapping(audio.get("input"), "session.audio.input")
|
||||
ga_format: Final = audio_input.get("format")
|
||||
if beta_format is not None and ga_format is not None:
|
||||
raise MuseProtocolError("input audio format must use either beta or GA layout")
|
||||
if beta_format is not None:
|
||||
if beta_format != "pcm16":
|
||||
def _parse_sample_rate(update: TranscriptionSessionUpdate) -> MuseSampleRate:
|
||||
audio_format: Final = update.audio_format
|
||||
if audio_format is None:
|
||||
return 24_000
|
||||
if audio_format.layout == "beta":
|
||||
if audio_format.encoding != "pcm16":
|
||||
raise MuseProtocolError("Muse Voice requires pcm16 input audio")
|
||||
return 24_000
|
||||
if ga_format is None:
|
||||
return 24_000
|
||||
if isinstance(ga_format, str):
|
||||
if ga_format != "pcm16":
|
||||
raise MuseProtocolError("Muse Voice requires audio/pcm input audio")
|
||||
return 24_000
|
||||
format_mapping: Final = _mapping(ga_format, "session.audio.input.format")
|
||||
if format_mapping.get("type") != "audio/pcm":
|
||||
if not audio_format.is_pcm16:
|
||||
raise MuseProtocolError("Muse Voice requires audio/pcm input audio")
|
||||
channels: Final = format_mapping.get("channels", 1)
|
||||
if isinstance(channels, bool) or channels != 1:
|
||||
if audio_format.channels not in (None, 1):
|
||||
raise MuseProtocolError("Muse Voice requires mono input audio")
|
||||
rate: Final = format_mapping.get("rate", 24_000)
|
||||
if isinstance(rate, bool) or not isinstance(rate, int) or rate not in SUPPORTED_SAMPLE_RATES:
|
||||
rate: Final = 24_000 if audio_format.rate is None else audio_format.rate
|
||||
if rate not in SUPPORTED_SAMPLE_RATES:
|
||||
raise MuseProtocolError("Muse Voice supports PCM16 at 16000 Hz or 24000 Hz")
|
||||
return 16_000 if rate == 16_000 else 24_000
|
||||
|
||||
|
||||
def _parse_mode(session: Mapping[str, JsonValue], audio_input: Mapping[str, JsonValue]) -> MuseMode:
|
||||
turn_detection_present: Final = "turn_detection" in session or "turn_detection" in audio_input
|
||||
turn_detection: Final = session.get("turn_detection", audio_input.get("turn_detection"))
|
||||
if turn_detection_present and turn_detection is None:
|
||||
def _parse_mode(update: TranscriptionSessionUpdate) -> MuseMode:
|
||||
if update.turn_detection_disabled:
|
||||
return "PUSH_TO_TALK"
|
||||
if turn_detection is None:
|
||||
return "ENDPOINTING"
|
||||
turn_detection_mapping: Final = _mapping(turn_detection, "turn_detection")
|
||||
if turn_detection_mapping.get("type") not in (None, "server_vad"):
|
||||
if update.turn_detection_type not in (None, "server_vad"):
|
||||
raise MuseProtocolError("Muse Voice supports server_vad turn detection or null")
|
||||
return "ENDPOINTING"
|
||||
|
||||
|
||||
def parse_session_update(payload: str, expected_model: str) -> MuseSessionConfig:
|
||||
message: Final = _json_object(payload)
|
||||
if message.get("type") not in ("session.update", "transcription_session.update"):
|
||||
raise MuseProtocolError("expected session.update")
|
||||
session: Final = _mapping(message.get("session"), "session")
|
||||
if not session:
|
||||
raise MuseProtocolError("session.update requires a session object")
|
||||
if session.get("type") not in (None, "transcription", "realtime"):
|
||||
update: Final = parse_transcription_session_update(payload, MuseProtocolError)
|
||||
if update.session_type not in (None, "transcription", "realtime"):
|
||||
raise MuseProtocolError("Muse Voice supports transcription sessions only")
|
||||
audio: Final = _mapping(session.get("audio"), "session.audio")
|
||||
audio_input: Final = _mapping(audio.get("input"), "session.audio.input")
|
||||
beta_transcription: Final = session.get("input_audio_transcription")
|
||||
ga_transcription: Final = audio_input.get("transcription")
|
||||
if beta_transcription is not None and ga_transcription is not None:
|
||||
raise MuseProtocolError("input transcription must use either beta or GA layout")
|
||||
transcription: Final = _mapping(
|
||||
beta_transcription if beta_transcription is not None else ga_transcription,
|
||||
"input audio transcription",
|
||||
)
|
||||
unsupported: Final = tuple(sorted(key for key in transcription if key not in _SUPPORTED_TRANSCRIPTION_KEYS))
|
||||
if unsupported:
|
||||
verbose_logger.warning("Meta realtime: dropping unsupported transcription settings %s", unsupported)
|
||||
requested_model: Final = _string(transcription.get("model"), "transcription model")
|
||||
if update.unsupported_transcription_keys:
|
||||
verbose_logger.warning(
|
||||
"Meta realtime: dropping unsupported transcription settings %s", update.unsupported_transcription_keys
|
||||
)
|
||||
normalized_model: Final = _normalize_model(expected_model)
|
||||
if normalized_model != MUSE_MODEL:
|
||||
raise MuseProtocolError("unsupported Meta realtime model")
|
||||
if requested_model is not None and _normalize_model(requested_model) != normalized_model:
|
||||
if update.model is not None and _normalize_model(update.model) != normalized_model:
|
||||
raise MuseProtocolError("realtime session model cannot be changed")
|
||||
language: Final = _string(transcription.get("language"), "language")
|
||||
return MuseSessionConfig(
|
||||
model=normalized_model,
|
||||
mode=_parse_mode(session, audio_input),
|
||||
sample_rate=_parse_sample_rate(session),
|
||||
language_bias=() if language is None else (normalize_language(language),),
|
||||
mode=_parse_mode(update),
|
||||
sample_rate=_parse_sample_rate(update),
|
||||
language_bias=() if update.language is None else (normalize_language(update.language),),
|
||||
)
|
||||
|
||||
|
||||
def session_created_event(config: MuseSessionConfig, session_id: str) -> OpenAIRealtimeTranscriptionSessionCreated:
|
||||
event: Final[OpenAIRealtimeTranscriptionSessionCreated] = {
|
||||
"type": "session.created",
|
||||
"event_id": _event_id(),
|
||||
"session": config.openai_session(session_id),
|
||||
}
|
||||
return event
|
||||
|
||||
|
||||
def error_event(message: str) -> OpenAIRealtimeErrorEvent:
|
||||
event: Final[OpenAIRealtimeErrorEvent] = {
|
||||
"type": "error",
|
||||
"error": {"type": "server_error", "message": message},
|
||||
}
|
||||
return event
|
||||
|
||||
|
||||
def _speech_event(
|
||||
event_type: Literal["input_audio_buffer.speech_started", "input_audio_buffer.speech_stopped"], item_id: str
|
||||
) -> OpenAIRealtimeInputAudioBufferSpeechEvent:
|
||||
event: Final[OpenAIRealtimeInputAudioBufferSpeechEvent] = {
|
||||
"type": event_type,
|
||||
"event_id": _event_id(),
|
||||
"item_id": item_id,
|
||||
}
|
||||
return event
|
||||
|
||||
|
||||
def _delta_event(item_id: str, delta: str) -> OpenAIRealtimeInputAudioTranscriptionDelta:
|
||||
event: Final[OpenAIRealtimeInputAudioTranscriptionDelta] = {
|
||||
"type": "conversation.item.input_audio_transcription.delta",
|
||||
"event_id": _event_id(),
|
||||
"item_id": item_id,
|
||||
"content_index": 0,
|
||||
"delta": delta,
|
||||
}
|
||||
return event
|
||||
|
||||
|
||||
def _completed_event(
|
||||
item_id: str, transcript: str, usage: RealtimeInputAudioTranscriptionUsage | None
|
||||
) -> OpenAIRealtimeInputAudioTranscriptionCompleted:
|
||||
event: Final[OpenAIRealtimeInputAudioTranscriptionCompleted] = {
|
||||
"type": "conversation.item.input_audio_transcription.completed",
|
||||
"event_id": _event_id(),
|
||||
"item_id": item_id,
|
||||
"content_index": 0,
|
||||
"transcript": transcript,
|
||||
}
|
||||
if usage is None:
|
||||
return event
|
||||
billed: Final[OpenAIRealtimeInputAudioTranscriptionCompleted] = {**event, "usage": usage}
|
||||
return billed
|
||||
return transcription_session_created_event(config.openai_session(session_id))
|
||||
|
||||
|
||||
def _required_turn_id(message: Mapping[str, JsonValue], event: str) -> str:
|
||||
|
|
@ -424,18 +299,18 @@ class _TurnState:
|
|||
has_content: Final = self.latest_partial is not None or self.final_text is not None
|
||||
if (self.started or has_content) and not self.start_emitted:
|
||||
self.start_emitted = True
|
||||
yield _speech_event("input_audio_buffer.speech_started", self.item_id)
|
||||
yield speech_event("input_audio_buffer.speech_started", self.item_id)
|
||||
if self.latest_partial is not None and self.final_text is None:
|
||||
delta: Final = _new_suffix(self.emitted_partial, self.latest_partial)
|
||||
if delta:
|
||||
self.emitted_partial = self.latest_partial
|
||||
yield _delta_event(self.item_id, delta)
|
||||
yield delta_event(self.item_id, delta)
|
||||
if self.stopped and not self.stopped_emitted:
|
||||
self.stopped_emitted = True
|
||||
yield _speech_event("input_audio_buffer.speech_stopped", self.item_id)
|
||||
yield speech_event("input_audio_buffer.speech_stopped", self.item_id)
|
||||
if self.final_text is not None and self.stopped_emitted and not self.completed_emitted:
|
||||
self.completed_emitted = True
|
||||
yield _completed_event(self.item_id, self.final_text, take_usage())
|
||||
yield completed_event(self.item_id, self.final_text, take_usage())
|
||||
|
||||
|
||||
class MuseEventTransformer:
|
||||
|
|
@ -467,8 +342,7 @@ class MuseEventTransformer:
|
|||
if seconds <= 0:
|
||||
return None
|
||||
self._unbilled_seconds = 0.0
|
||||
usage: Final[RealtimeInputAudioTranscriptionDurationUsage] = {"type": "duration", "seconds": seconds}
|
||||
return usage
|
||||
return duration_usage(seconds)
|
||||
|
||||
def _apply_turn_event(self, event_type: JsonValue | None, message: Mapping[str, JsonValue]) -> _TurnState | None:
|
||||
match event_type:
|
||||
|
|
@ -612,7 +486,7 @@ class MetaRealtimeConfig(BaseRealtimeConfig):
|
|||
model: str,
|
||||
session_configuration_request: str | None = None,
|
||||
) -> tuple[str | bytes, ...]:
|
||||
request: Final = _json_object(message)
|
||||
request: Final = json_object(message, MuseProtocolError)
|
||||
event_type: Final = request.get("type")
|
||||
if event_type in ("session.update", "transcription_session.update"):
|
||||
return self._configure(message, model)
|
||||
|
|
@ -664,7 +538,7 @@ class MetaRealtimeConfig(BaseRealtimeConfig):
|
|||
return result
|
||||
|
||||
def _backend_events(self, payload: str) -> tuple[OpenAIRealtimeEvents, ...]:
|
||||
frame: Final = _json_object(payload)
|
||||
frame: Final = json_object(payload, MuseProtocolError)
|
||||
session_id: Final = frame.get("sessionId")
|
||||
if session_id is None:
|
||||
return self._transformer.transform(frame)
|
||||
|
|
@ -686,17 +560,7 @@ class MetaRealtimeConfig(BaseRealtimeConfig):
|
|||
|
||||
def _append_audio(self, request: Mapping[str, JsonValue]) -> tuple[bytes, ...]:
|
||||
config: Final = self._require_config()
|
||||
encoded: Final = request.get("audio")
|
||||
if not isinstance(encoded, str):
|
||||
raise MuseProtocolError("Audio must be a base64 string")
|
||||
if len(encoded) > config.max_encoded_append_bytes:
|
||||
raise MuseProtocolError("Audio append exceeds the four-second backlog limit")
|
||||
try:
|
||||
audio: Final = base64.b64decode(encoded, validate=True)
|
||||
except (binascii.Error, ValueError):
|
||||
raise MuseProtocolError("Audio must be valid base64") from None
|
||||
if len(audio) % 2:
|
||||
raise MuseProtocolError("PCM16 audio must contain complete samples")
|
||||
audio: Final = decode_pcm16_append(request.get("audio"), config.max_encoded_append_bytes, MuseProtocolError)
|
||||
buffered: Final = self._pending_audio + audio
|
||||
packet_end: Final = len(buffered) - len(buffered) % config.packet_bytes
|
||||
self._pending_audio = buffered[packet_end:]
|
||||
|
|
|
|||
418
litellm/llms/vertex_ai/audio_transcription/realtime_backend.py
Normal file
418
litellm/llms/vertex_ai/audio_transcription/realtime_backend.py
Normal file
|
|
@ -0,0 +1,418 @@
|
|||
import asyncio
|
||||
import time
|
||||
from collections.abc import AsyncIterable, AsyncIterator, Awaitable, Callable
|
||||
from dataclasses import dataclass
|
||||
from datetime import timedelta
|
||||
from types import MappingProxyType, TracebackType
|
||||
from typing import TYPE_CHECKING, Final, Literal, Protocol
|
||||
|
||||
from pydantic import TypeAdapter
|
||||
from typing_extensions import Self, assert_never
|
||||
from websockets.exceptions import ConnectionClosedError, ConnectionClosedOK
|
||||
from websockets.frames import Close
|
||||
|
||||
from litellm import verbose_logger
|
||||
from litellm.llms.vertex_ai.audio_transcription.realtime_transformation import SpeechStreamingTarget
|
||||
from litellm.types.llms.vertex_ai_speech_to_text import (
|
||||
VertexSpeechStreamingCommand,
|
||||
VertexSpeechStreamingCommandUnion,
|
||||
VertexSpeechStreamingConfigure,
|
||||
VertexSpeechStreamingConfigured,
|
||||
VertexSpeechStreamingDiscardTurn,
|
||||
VertexSpeechStreamingFinishTurn,
|
||||
VertexSpeechStreamingResponse,
|
||||
VertexSpeechStreamingResult,
|
||||
VertexSpeechStreamingTurnDiscarded,
|
||||
VertexSpeechStreamingTurnFinished,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from google.cloud.speech_v2.types import (
|
||||
StreamingRecognitionConfig,
|
||||
StreamingRecognizeRequest,
|
||||
StreamingRecognizeResponse,
|
||||
)
|
||||
|
||||
SPEECH_SDK_INSTALL_HINT: Final = (
|
||||
"google-cloud-speech is not installed. Install with `pip install 'litellm[stt-vertex-chirp]'`."
|
||||
)
|
||||
STREAM_FAILURE_CLOSE_CODE: Final = 1011
|
||||
STREAM_ROTATION_SECONDS: Final = 240.0
|
||||
STREAM_ROTATION_DEADLINE_SECONDS: Final = 280.0
|
||||
REQUEST_QUEUE_SIZE: Final = 64
|
||||
OUTBOX_SIZE: Final = 256
|
||||
_LINK_QUEUE_SIZE: Final = 64
|
||||
_CLOSE_REASON_MAX_CHARS: Final = 120
|
||||
_CONFIGURED_EVENT: Final = VertexSpeechStreamingConfigured().model_dump_json()
|
||||
_TURN_FINISHED_EVENT: Final = VertexSpeechStreamingTurnFinished().model_dump_json()
|
||||
_COMMAND_ADAPTER: Final = TypeAdapter[VertexSpeechStreamingCommandUnion](VertexSpeechStreamingCommand)
|
||||
_TIMEDELTA_ADAPTER: Final = TypeAdapter(timedelta)
|
||||
_SPEECH_EVENTS: Final[MappingProxyType[str, Literal["begin", "end"]]] = MappingProxyType(
|
||||
{
|
||||
"SPEECH_ACTIVITY_BEGIN": "begin",
|
||||
"SPEECH_ACTIVITY_END": "end",
|
||||
"END_OF_SINGLE_UTTERANCE": "end",
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class ClosableTransport(Protocol):
|
||||
def close(self) -> Awaitable[None]: ...
|
||||
|
||||
|
||||
class SpeechStreamingClient(Protocol):
|
||||
def streaming_recognize(
|
||||
self, requests: "AsyncIterator[StreamingRecognizeRequest] | None" = None
|
||||
) -> "Awaitable[AsyncIterable[StreamingRecognizeResponse]]": ...
|
||||
|
||||
@property
|
||||
def transport(self) -> ClosableTransport: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _StreamFailure:
|
||||
reason: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Closed:
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _TurnResult:
|
||||
turn: int
|
||||
event: str
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _TurnDiscarded:
|
||||
turn: int
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _TurnDiscardedEvent:
|
||||
turn: int
|
||||
event: str
|
||||
|
||||
|
||||
_OutboxItem = str | _TurnResult | _TurnDiscardedEvent | _StreamFailure | _Closed
|
||||
|
||||
|
||||
def open_speech_client(target: SpeechStreamingTarget, access_token: str) -> SpeechStreamingClient:
|
||||
try:
|
||||
from google.api_core.client_options import ClientOptions
|
||||
from google.cloud.speech_v2 import SpeechAsyncClient
|
||||
from google.oauth2.credentials import Credentials
|
||||
except ImportError as e:
|
||||
raise ImportError(SPEECH_SDK_INSTALL_HINT) from e
|
||||
return SpeechAsyncClient(
|
||||
credentials=Credentials(token=access_token),
|
||||
transport="grpc_asyncio",
|
||||
client_options=ClientOptions(api_endpoint=target.api_endpoint),
|
||||
)
|
||||
|
||||
|
||||
def _streaming_config(command: VertexSpeechStreamingConfigure) -> "StreamingRecognitionConfig":
|
||||
from google.cloud.speech_v2.types import (
|
||||
ExplicitDecodingConfig,
|
||||
RecognitionConfig,
|
||||
StreamingRecognitionConfig,
|
||||
StreamingRecognitionFeatures,
|
||||
)
|
||||
|
||||
return StreamingRecognitionConfig(
|
||||
config=RecognitionConfig(
|
||||
explicit_decoding_config=ExplicitDecodingConfig(
|
||||
encoding=ExplicitDecodingConfig.AudioEncoding.LINEAR16,
|
||||
sample_rate_hertz=command.sample_rate_hertz,
|
||||
audio_channel_count=1,
|
||||
),
|
||||
model=command.model,
|
||||
language_codes=command.language_codes,
|
||||
),
|
||||
streaming_features=StreamingRecognitionFeatures(interim_results=True, enable_voice_activity_events=True),
|
||||
)
|
||||
|
||||
|
||||
def _response_event(response: "StreamingRecognizeResponse", billed_seconds: float) -> str:
|
||||
return VertexSpeechStreamingResponse(
|
||||
speech_event=_SPEECH_EVENTS.get(response.speech_event_type.name, "none"),
|
||||
results=tuple(
|
||||
VertexSpeechStreamingResult(
|
||||
transcript=result.alternatives[0].transcript if result.alternatives else "",
|
||||
is_final=result.is_final,
|
||||
)
|
||||
for result in response.results
|
||||
),
|
||||
billed_seconds=billed_seconds,
|
||||
).model_dump_json()
|
||||
|
||||
|
||||
def _billed_seconds(response: "StreamingRecognizeResponse") -> float:
|
||||
return _TIMEDELTA_ADAPTER.validate_python(response.metadata.total_billed_duration).total_seconds()
|
||||
|
||||
|
||||
def _normal_closure() -> ConnectionClosedOK:
|
||||
return ConnectionClosedOK(rcvd=Close(1000, ""), sent=None)
|
||||
|
||||
|
||||
class _RecognizeStream:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
client: SpeechStreamingClient,
|
||||
request_type: "type[StreamingRecognizeRequest]",
|
||||
first_request: "StreamingRecognizeRequest",
|
||||
opened_at: float,
|
||||
turn: int,
|
||||
) -> None:
|
||||
self._client: Final = client
|
||||
self._request_type: Final = request_type
|
||||
self.opened_at: Final = opened_at
|
||||
self.turn: Final = turn
|
||||
self._requests: Final[asyncio.Queue[StreamingRecognizeRequest | None]] = asyncio.Queue(
|
||||
maxsize=REQUEST_QUEUE_SIZE
|
||||
)
|
||||
self._requests.put_nowait(first_request)
|
||||
self.speech_active: bool = False
|
||||
self.billed_seconds: float = 0.0
|
||||
self._cancelled: bool = False
|
||||
self._closed: bool = False
|
||||
self._task: asyncio.Task[None] | None = None
|
||||
|
||||
async def send_audio(self, audio: bytes) -> None:
|
||||
await self._requests.put(self._request_type(audio=audio))
|
||||
|
||||
async def half_close(self) -> None:
|
||||
await self._requests.put(None)
|
||||
|
||||
def cancel(self) -> None:
|
||||
self._cancelled = True
|
||||
if self._task is not None:
|
||||
self._task.cancel()
|
||||
|
||||
async def close(self) -> None:
|
||||
if self._closed:
|
||||
return
|
||||
self._closed = True
|
||||
await self._client.transport.close()
|
||||
|
||||
async def relay(self, outbox: asyncio.Queue[_OutboxItem], billed_before: float) -> float:
|
||||
if self._cancelled:
|
||||
await self.close()
|
||||
return 0.0
|
||||
task: Final = asyncio.create_task(self._forward(outbox, billed_before))
|
||||
self._task = task
|
||||
try:
|
||||
await asyncio.wait((task,))
|
||||
except asyncio.CancelledError:
|
||||
task.cancel()
|
||||
await asyncio.wait((task,))
|
||||
raise
|
||||
finally:
|
||||
await self.close()
|
||||
return self.billed_seconds
|
||||
|
||||
async def _forward(self, outbox: asyncio.Queue[_OutboxItem], billed_before: float) -> None:
|
||||
try:
|
||||
responses: Final = await self._client.streaming_recognize(self._drain())
|
||||
async for response in responses:
|
||||
self._note(response)
|
||||
await outbox.put(
|
||||
_TurnResult(turn=self.turn, event=_response_event(response, billed_before + self.billed_seconds))
|
||||
)
|
||||
except Exception as e: # noqa: BLE001 # task boundary: a swallowed failure would hang the client session
|
||||
verbose_logger.warning("Google Speech-to-Text streaming failed: %s", e)
|
||||
await outbox.put(_StreamFailure(reason=f"Google Speech-to-Text streaming failed: {e}"))
|
||||
|
||||
def _note(self, response: "StreamingRecognizeResponse") -> None:
|
||||
activity: Final = _SPEECH_EVENTS.get(response.speech_event_type.name)
|
||||
if activity is not None:
|
||||
self.speech_active = activity == "begin"
|
||||
self.billed_seconds = max(self.billed_seconds, _billed_seconds(response))
|
||||
|
||||
async def _drain(self) -> "AsyncIterator[StreamingRecognizeRequest]":
|
||||
while (request := await self._requests.get()) is not None:
|
||||
yield request
|
||||
|
||||
|
||||
_Link = _RecognizeStream | str | _TurnDiscarded
|
||||
|
||||
|
||||
class SpeechStreamingBackend:
|
||||
def __init__(
|
||||
self,
|
||||
target: SpeechStreamingTarget,
|
||||
*,
|
||||
client_factory: Callable[[SpeechStreamingTarget, str], SpeechStreamingClient] = open_speech_client,
|
||||
clock: Callable[[], float] = time.monotonic,
|
||||
rotation_seconds: float = STREAM_ROTATION_SECONDS,
|
||||
rotation_deadline_seconds: float = STREAM_ROTATION_DEADLINE_SECONDS,
|
||||
) -> None:
|
||||
self._target: Final = target
|
||||
self._client_factory: Final = client_factory
|
||||
self._clock: Final = clock
|
||||
self._rotation_seconds: Final = rotation_seconds
|
||||
self._rotation_deadline_seconds: Final = rotation_deadline_seconds
|
||||
self._outbox: Final[asyncio.Queue[_OutboxItem]] = asyncio.Queue(maxsize=OUTBOX_SIZE)
|
||||
self._links: Final[asyncio.Queue[_Link]] = asyncio.Queue(maxsize=_LINK_QUEUE_SIZE)
|
||||
self._pump: asyncio.Task[None] | None = None
|
||||
self._config: StreamingRecognitionConfig | None = None
|
||||
self._turn: tuple[_RecognizeStream, ...] = ()
|
||||
self._turn_index: int = 0
|
||||
self._discarded_turns: frozenset[int] = frozenset()
|
||||
self._billed_before: float = 0.0
|
||||
self._closed: bool = False
|
||||
|
||||
async def __aenter__(self) -> Self:
|
||||
return self
|
||||
|
||||
async def __aexit__(
|
||||
self,
|
||||
exc_type: type[BaseException] | None,
|
||||
exc_value: BaseException | None,
|
||||
traceback: TracebackType | None,
|
||||
) -> None:
|
||||
await self.close()
|
||||
|
||||
async def send(self, message: str | bytes) -> None:
|
||||
if self._closed:
|
||||
raise _normal_closure()
|
||||
if isinstance(message, bytes):
|
||||
await self._send_audio(message)
|
||||
return
|
||||
command: Final = _COMMAND_ADAPTER.validate_json(message)
|
||||
match command:
|
||||
case VertexSpeechStreamingConfigure():
|
||||
self._config = _streaming_config(command)
|
||||
await self._link(_CONFIGURED_EVENT)
|
||||
case VertexSpeechStreamingFinishTurn():
|
||||
await self._finish_turn()
|
||||
case VertexSpeechStreamingDiscardTurn():
|
||||
await self._discard_turn()
|
||||
case _:
|
||||
assert_never(command)
|
||||
|
||||
async def recv(self, decode: bool | None = None) -> str | bytes:
|
||||
while not (self._closed and self._outbox.empty()):
|
||||
if (event := self._deliverable(await self._outbox.get())) is not None:
|
||||
return event
|
||||
raise _normal_closure()
|
||||
|
||||
def _deliverable(self, item: _OutboxItem) -> str | None:
|
||||
match item:
|
||||
case _StreamFailure():
|
||||
raise ConnectionClosedError(
|
||||
rcvd=Close(STREAM_FAILURE_CLOSE_CODE, item.reason[:_CLOSE_REASON_MAX_CHARS]), sent=None
|
||||
)
|
||||
case _Closed():
|
||||
raise _normal_closure()
|
||||
case _TurnResult():
|
||||
return None if item.turn in self._discarded_turns else item.event
|
||||
case _TurnDiscardedEvent():
|
||||
self._discarded_turns -= {item.turn}
|
||||
return item.event
|
||||
case str():
|
||||
return item
|
||||
case _:
|
||||
assert_never(item)
|
||||
|
||||
async def close(self) -> None:
|
||||
if self._closed:
|
||||
return
|
||||
self._closed = True
|
||||
self._turn = ()
|
||||
pump: Final = self._pump
|
||||
if pump is not None:
|
||||
pump.cancel()
|
||||
await asyncio.wait((pump,))
|
||||
await self._close_unrelayed_streams()
|
||||
if not self._outbox.full():
|
||||
self._outbox.put_nowait(_Closed())
|
||||
|
||||
async def _close_unrelayed_streams(self) -> None:
|
||||
unrelayed: Final = tuple(self._links.get_nowait() for _ in range(self._links.qsize()))
|
||||
for link in unrelayed:
|
||||
if isinstance(link, _RecognizeStream):
|
||||
await link.close()
|
||||
|
||||
async def _link(self, item: _Link) -> None:
|
||||
if self._pump is None:
|
||||
self._pump = asyncio.create_task(self._pump_links())
|
||||
await self._links.put(item)
|
||||
|
||||
async def _pump_links(self) -> None:
|
||||
while True:
|
||||
await self._relay(await self._links.get())
|
||||
|
||||
async def _relay(self, link: _Link) -> None:
|
||||
match link:
|
||||
case str():
|
||||
await self._outbox.put(link)
|
||||
case _RecognizeStream():
|
||||
self._billed_before += await link.relay(self._outbox, self._billed_before)
|
||||
case _TurnDiscarded():
|
||||
await self._outbox.put(
|
||||
_TurnDiscardedEvent(
|
||||
turn=link.turn,
|
||||
event=VertexSpeechStreamingTurnDiscarded(billed_seconds=self._billed_before).model_dump_json(),
|
||||
)
|
||||
)
|
||||
case _:
|
||||
assert_never(link)
|
||||
|
||||
async def _send_audio(self, audio: bytes) -> None:
|
||||
stream: Final = await self._turn_stream()
|
||||
await stream.send_audio(audio)
|
||||
|
||||
async def _turn_stream(self) -> _RecognizeStream:
|
||||
current: Final = self._turn[-1] if self._turn else None
|
||||
if current is not None and not self._expired(current):
|
||||
return current
|
||||
if current is not None:
|
||||
await current.half_close()
|
||||
stream: Final = await self._open_stream()
|
||||
self._turn = (*self._turn, stream)
|
||||
return stream
|
||||
|
||||
def _expired(self, stream: _RecognizeStream) -> bool:
|
||||
elapsed: Final = self._clock() - stream.opened_at
|
||||
if elapsed >= self._rotation_deadline_seconds:
|
||||
return True
|
||||
return elapsed >= self._rotation_seconds and not stream.speech_active
|
||||
|
||||
async def _open_stream(self) -> _RecognizeStream:
|
||||
from google.cloud.speech_v2.types import StreamingRecognizeRequest
|
||||
|
||||
config: Final = self._config
|
||||
if config is None:
|
||||
raise RuntimeError("audio was sent before the Speech-to-Text stream was configured")
|
||||
access_token: Final = await self._target.resolve_access_token()
|
||||
stream: Final = _RecognizeStream(
|
||||
client=self._client_factory(self._target, access_token),
|
||||
request_type=StreamingRecognizeRequest,
|
||||
first_request=StreamingRecognizeRequest(recognizer=self._target.recognizer, streaming_config=config),
|
||||
opened_at=self._clock(),
|
||||
turn=self._turn_index,
|
||||
)
|
||||
await self._link(stream)
|
||||
return stream
|
||||
|
||||
async def _finish_turn(self) -> None:
|
||||
turn: Final = self._turn
|
||||
self._turn = ()
|
||||
self._turn_index += 1
|
||||
if turn:
|
||||
await turn[-1].half_close()
|
||||
await self._link(_TURN_FINISHED_EVENT)
|
||||
|
||||
async def _discard_turn(self) -> None:
|
||||
streams: Final = self._turn
|
||||
turn: Final = self._turn_index
|
||||
self._turn = ()
|
||||
self._discarded_turns |= {turn}
|
||||
self._turn_index += 1
|
||||
for stream in streams:
|
||||
stream.cancel()
|
||||
await self._link(_TurnDiscarded(turn=turn))
|
||||
|
|
@ -0,0 +1,446 @@
|
|||
from collections.abc import Awaitable, Callable, Mapping
|
||||
from dataclasses import dataclass, replace
|
||||
from typing import Final
|
||||
|
||||
from pydantic import JsonValue, TypeAdapter
|
||||
from typing_extensions import assert_never
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.litellm_core_utils.audio_utils.utils import normalize_transcription_language_to_bcp47
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.base_llm.realtime.transcription_protocol import (
|
||||
RealtimeTranscriptionProtocolError,
|
||||
TranscriptionAudioFormat,
|
||||
TranscriptionSessionUpdate,
|
||||
completed_event,
|
||||
decode_pcm16_append,
|
||||
delta_event,
|
||||
duration_usage,
|
||||
json_object,
|
||||
parse_transcription_session_update,
|
||||
speech_event,
|
||||
transcription_session,
|
||||
transcription_session_created_event,
|
||||
)
|
||||
from litellm.llms.base_llm.realtime.transformation import BaseRealtimeConfig, RealtimeBackend
|
||||
from litellm.llms.vertex_ai.audio_transcription.transformation import (
|
||||
AUTO_LANGUAGE_CODE,
|
||||
DEFAULT_SPEECH_TO_TEXT_LOCATION,
|
||||
speech_to_text_host,
|
||||
validate_vertex_transcription_location,
|
||||
validate_vertex_transcription_project_id,
|
||||
)
|
||||
from litellm.types.llms.openai import (
|
||||
OpenAIRealtimeEvents,
|
||||
OpenAIRealtimeTranscriptionSession,
|
||||
OpenAIRealtimeTranscriptionSessionCreated,
|
||||
)
|
||||
from litellm.types.llms.vertex_ai_speech_to_text import (
|
||||
VertexSpeechStreamingConfigure,
|
||||
VertexSpeechStreamingConfigured,
|
||||
VertexSpeechStreamingDiscardTurn,
|
||||
VertexSpeechStreamingEvent,
|
||||
VertexSpeechStreamingEventUnion,
|
||||
VertexSpeechStreamingFinishTurn,
|
||||
VertexSpeechStreamingResponse,
|
||||
VertexSpeechStreamingTurnDiscarded,
|
||||
VertexSpeechStreamingTurnFinished,
|
||||
)
|
||||
from litellm.types.realtime import (
|
||||
RealtimeInputAudioTranscriptionUsage,
|
||||
RealtimeResponseTransformInput,
|
||||
RealtimeResponseTypedDict,
|
||||
)
|
||||
|
||||
DEFAULT_SAMPLE_RATE_HERTZ: Final = 24_000
|
||||
MIN_SAMPLE_RATE_HERTZ: Final = 8_000
|
||||
MAX_SAMPLE_RATE_HERTZ: Final = 48_000
|
||||
MAX_AUDIO_MESSAGE_BYTES: Final = 25_000
|
||||
_SPEECH_TO_TEXT_ENDPOINTS: Final = frozenset({"/v1/audio/transcriptions", "/v1/realtime"})
|
||||
_VERTEX_MODEL_PREFIX: Final = "vertex_ai/"
|
||||
_STREAMING_EVENT_ADAPTER: Final = TypeAdapter[VertexSpeechStreamingEventUnion](VertexSpeechStreamingEvent)
|
||||
_FINISH_TURN_COMMAND: Final = VertexSpeechStreamingFinishTurn().model_dump_json()
|
||||
_DISCARD_TURN_COMMAND: Final = VertexSpeechStreamingDiscardTurn().model_dump_json()
|
||||
|
||||
|
||||
class ChirpProtocolError(RealtimeTranscriptionProtocolError):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SpeechStreamingTarget:
|
||||
api_endpoint: str
|
||||
recognizer: str
|
||||
resolve_access_token: Callable[[], Awaitable[str]]
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ChirpSessionConfig:
|
||||
model: str
|
||||
language: str | None
|
||||
sample_rate: int
|
||||
server_vad: bool
|
||||
|
||||
def openai_session(self, session_id: str) -> OpenAIRealtimeTranscriptionSession:
|
||||
return transcription_session(
|
||||
session_id=session_id,
|
||||
model=self.model,
|
||||
sample_rate=self.sample_rate,
|
||||
language=self.language,
|
||||
server_vad=self.server_vad,
|
||||
)
|
||||
|
||||
def configure_command(self) -> str:
|
||||
return VertexSpeechStreamingConfigure(
|
||||
model=self.model,
|
||||
language_codes=(AUTO_LANGUAGE_CODE,) if self.language is None else (self.language,),
|
||||
sample_rate_hertz=self.sample_rate,
|
||||
).model_dump_json()
|
||||
|
||||
|
||||
def is_vertex_speech_to_text_model(model: str) -> bool:
|
||||
try:
|
||||
info: Final = litellm.get_model_info(
|
||||
model=normalize_speech_to_text_model(model), custom_llm_provider="vertex_ai"
|
||||
)
|
||||
except Exception: # noqa: BLE001 # get_model_info raises for unmapped models, which are not Speech-to-Text models
|
||||
return False
|
||||
if info.get("mode") != "audio_transcription":
|
||||
return False
|
||||
return _SPEECH_TO_TEXT_ENDPOINTS <= frozenset(info.get("supported_endpoints") or ())
|
||||
|
||||
|
||||
def normalize_speech_to_text_model(model: str) -> str:
|
||||
return model.removeprefix(_VERTEX_MODEL_PREFIX)
|
||||
|
||||
|
||||
def default_session_config(model: str) -> ChirpSessionConfig:
|
||||
return ChirpSessionConfig(
|
||||
model=normalize_speech_to_text_model(model),
|
||||
language=None,
|
||||
sample_rate=DEFAULT_SAMPLE_RATE_HERTZ,
|
||||
server_vad=True,
|
||||
)
|
||||
|
||||
|
||||
def parse_chirp_session_update(payload: str, expected_model: str) -> ChirpSessionConfig:
|
||||
update: Final = parse_transcription_session_update(payload, ChirpProtocolError)
|
||||
if update.session_type not in (None, "transcription", "realtime"):
|
||||
raise ChirpProtocolError("Speech-to-Text streaming supports transcription sessions only")
|
||||
if update.unsupported_transcription_keys:
|
||||
verbose_logger.debug(
|
||||
"Speech-to-Text streaming: ignoring unsupported transcription settings %s",
|
||||
update.unsupported_transcription_keys,
|
||||
)
|
||||
model: Final = normalize_speech_to_text_model(expected_model)
|
||||
if update.model is not None and normalize_speech_to_text_model(update.model) != model:
|
||||
raise ChirpProtocolError("realtime session model cannot be changed")
|
||||
return ChirpSessionConfig(
|
||||
model=model,
|
||||
language=None if update.language is None else normalize_transcription_language_to_bcp47(update.language),
|
||||
sample_rate=_parse_sample_rate(update.audio_format),
|
||||
server_vad=_parse_server_vad(update),
|
||||
)
|
||||
|
||||
|
||||
def _parse_sample_rate(audio_format: TranscriptionAudioFormat | None) -> int:
|
||||
if audio_format is None:
|
||||
return DEFAULT_SAMPLE_RATE_HERTZ
|
||||
if not audio_format.is_pcm16:
|
||||
raise ChirpProtocolError("Speech-to-Text streaming requires pcm16 input audio")
|
||||
if audio_format.channels not in (None, 1):
|
||||
raise ChirpProtocolError("Speech-to-Text streaming requires mono input audio")
|
||||
rate: Final = DEFAULT_SAMPLE_RATE_HERTZ if audio_format.rate is None else audio_format.rate
|
||||
if not MIN_SAMPLE_RATE_HERTZ <= rate <= MAX_SAMPLE_RATE_HERTZ:
|
||||
raise ChirpProtocolError(
|
||||
f"Speech-to-Text streaming supports sample rates from {MIN_SAMPLE_RATE_HERTZ} Hz"
|
||||
f" to {MAX_SAMPLE_RATE_HERTZ} Hz"
|
||||
)
|
||||
return rate
|
||||
|
||||
|
||||
def _parse_server_vad(update: TranscriptionSessionUpdate) -> bool:
|
||||
if update.turn_detection_disabled:
|
||||
return False
|
||||
if update.turn_detection_type not in (None, "server_vad"):
|
||||
raise ChirpProtocolError("Speech-to-Text streaming supports server_vad turn detection or null")
|
||||
return True
|
||||
|
||||
|
||||
def session_created_event(config: ChirpSessionConfig, session_id: str) -> OpenAIRealtimeTranscriptionSessionCreated:
|
||||
return transcription_session_created_event(config.openai_session(session_id))
|
||||
|
||||
|
||||
def _normalize_word(word: str) -> str:
|
||||
return "".join(char for char in word if char.isalnum()).casefold()
|
||||
|
||||
|
||||
def new_words(previous: str, current: str) -> str:
|
||||
previous_words: Final = previous.split()
|
||||
current_words: Final = current.split()
|
||||
common: Final = next(
|
||||
(
|
||||
index
|
||||
for index, (old, new) in enumerate(zip(previous_words, current_words, strict=False))
|
||||
if _normalize_word(old) != _normalize_word(new)
|
||||
),
|
||||
min(len(previous_words), len(current_words)),
|
||||
)
|
||||
appended: Final = " ".join(current_words[common:])
|
||||
if not appended:
|
||||
return ""
|
||||
return f" {appended}" if common else appended
|
||||
|
||||
|
||||
def _join_transcript(committed: str, tail: str) -> str:
|
||||
return " ".join(part for part in (committed, tail) if part)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Turn:
|
||||
item_id: str
|
||||
committed: str = ""
|
||||
preview: str = ""
|
||||
started_emitted: bool = False
|
||||
stopped_emitted: bool = False
|
||||
|
||||
|
||||
class ChirpEventTransformer:
|
||||
def __init__(self, *, new_item_id: Callable[[], str] = lambda: f"item_{uuid.uuid4().hex}") -> None:
|
||||
self._new_item_id: Final = new_item_id
|
||||
self._config: ChirpSessionConfig | None = None
|
||||
self._session_id: str | None = None
|
||||
self._turn: _Turn | None = None
|
||||
self._billed_seconds: float = 0.0
|
||||
self._reported_seconds: float = 0.0
|
||||
|
||||
def configure(self, config: ChirpSessionConfig, session_id: str) -> None:
|
||||
self._config = config
|
||||
self._session_id = session_id
|
||||
|
||||
def take_unbilled_usage(self) -> RealtimeInputAudioTranscriptionUsage | None:
|
||||
unreported: Final = self._billed_seconds - self._reported_seconds
|
||||
if unreported <= 0:
|
||||
return None
|
||||
self._reported_seconds = self._billed_seconds
|
||||
return duration_usage(unreported)
|
||||
|
||||
def transform(self, frame: VertexSpeechStreamingEventUnion) -> tuple[OpenAIRealtimeEvents, ...]:
|
||||
match frame:
|
||||
case VertexSpeechStreamingConfigured():
|
||||
return (session_created_event(self._require_config(), self._require_session_id()),)
|
||||
case VertexSpeechStreamingResponse():
|
||||
return self._response(frame)
|
||||
case VertexSpeechStreamingTurnFinished():
|
||||
return self._finish_turn()
|
||||
case VertexSpeechStreamingTurnDiscarded():
|
||||
self._billed_seconds = max(self._billed_seconds, frame.billed_seconds)
|
||||
self._turn = None
|
||||
return ()
|
||||
case _:
|
||||
assert_never(frame)
|
||||
|
||||
def _response(self, frame: VertexSpeechStreamingResponse) -> tuple[OpenAIRealtimeEvents, ...]:
|
||||
self._billed_seconds = max(self._billed_seconds, frame.billed_seconds)
|
||||
interim: Final = " ".join(
|
||||
result.transcript.strip() for result in frame.results if not result.is_final and result.transcript.strip()
|
||||
)
|
||||
finals: Final = tuple(
|
||||
result.transcript.strip() for result in frame.results if result.is_final and result.transcript.strip()
|
||||
)
|
||||
begin_events: Final = self._begin() if frame.speech_event == "begin" else ()
|
||||
final_events: Final = tuple(event for final in finals for event in self._final(final))
|
||||
interim_events: Final = self._hypothesis(interim) if interim else ()
|
||||
end_events: Final = self._stop() if frame.speech_event == "end" else ()
|
||||
return (*begin_events, *final_events, *interim_events, *end_events)
|
||||
|
||||
def _begin(self) -> tuple[OpenAIRealtimeEvents, ...]:
|
||||
turn: Final = self._require_turn()
|
||||
if turn.started_emitted or not self._require_config().server_vad:
|
||||
return ()
|
||||
self._turn = replace(turn, started_emitted=True)
|
||||
return (speech_event("input_audio_buffer.speech_started", turn.item_id),)
|
||||
|
||||
def _stop(self) -> tuple[OpenAIRealtimeEvents, ...]:
|
||||
turn: Final = self._turn
|
||||
if turn is None or turn.stopped_emitted or not self._require_config().server_vad:
|
||||
return ()
|
||||
self._turn = replace(turn, stopped_emitted=True)
|
||||
return (speech_event("input_audio_buffer.speech_stopped", turn.item_id),)
|
||||
|
||||
def _hypothesis(self, text: str) -> tuple[OpenAIRealtimeEvents, ...]:
|
||||
begin_events: Final = self._begin()
|
||||
turn: Final = self._require_turn()
|
||||
hypothesis: Final = _join_transcript(turn.committed, text)
|
||||
delta: Final = new_words(turn.preview, hypothesis)
|
||||
self._turn = replace(turn, preview=hypothesis)
|
||||
return (*begin_events, delta_event(turn.item_id, delta)) if delta else begin_events
|
||||
|
||||
def _final(self, text: str) -> tuple[OpenAIRealtimeEvents, ...]:
|
||||
begin_events: Final = self._begin()
|
||||
turn: Final = self._require_turn()
|
||||
committed: Final = _join_transcript(turn.committed, text)
|
||||
delta: Final = new_words(turn.preview, committed)
|
||||
self._turn = replace(turn, committed=committed, preview=committed)
|
||||
delta_events: Final[tuple[OpenAIRealtimeEvents, ...]] = (delta_event(turn.item_id, delta),) if delta else ()
|
||||
if not self._require_config().server_vad:
|
||||
return (*begin_events, *delta_events)
|
||||
return (*begin_events, *delta_events, *self._complete())
|
||||
|
||||
def _finish_turn(self) -> tuple[OpenAIRealtimeEvents, ...]:
|
||||
if self._turn is None:
|
||||
return ()
|
||||
return self._complete()
|
||||
|
||||
def _complete(self) -> tuple[OpenAIRealtimeEvents, ...]:
|
||||
turn: Final = self._require_turn()
|
||||
stop_events: Final = self._stop()
|
||||
transcript: Final = turn.committed or turn.preview
|
||||
self._turn = None
|
||||
return (*stop_events, completed_event(turn.item_id, transcript, self.take_unbilled_usage()))
|
||||
|
||||
def _require_turn(self) -> _Turn:
|
||||
if self._turn is None:
|
||||
self._turn = _Turn(item_id=self._new_item_id())
|
||||
return self._turn
|
||||
|
||||
def _require_config(self) -> ChirpSessionConfig:
|
||||
if self._config is None:
|
||||
raise ChirpProtocolError("session.update must configure the session before the backend responds")
|
||||
return self._config
|
||||
|
||||
def _require_session_id(self) -> str:
|
||||
if self._session_id is None:
|
||||
raise ChirpProtocolError("session.update must configure the session before the backend responds")
|
||||
return self._session_id
|
||||
|
||||
|
||||
def _default_backend_factory(target: SpeechStreamingTarget) -> RealtimeBackend:
|
||||
from litellm.llms.vertex_ai.audio_transcription.realtime_backend import SpeechStreamingBackend
|
||||
|
||||
return SpeechStreamingBackend(target)
|
||||
|
||||
|
||||
class VertexChirpRealtimeConfig(BaseRealtimeConfig):
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
resolve_access_token: Callable[[], Awaitable[str]],
|
||||
project: str,
|
||||
location: str | None,
|
||||
backend_factory: Callable[[SpeechStreamingTarget], RealtimeBackend] = _default_backend_factory,
|
||||
) -> None:
|
||||
self._resolve_access_token: Final = resolve_access_token
|
||||
self._project: Final = validate_vertex_transcription_project_id(project)
|
||||
self._location: Final = validate_vertex_transcription_location(location, DEFAULT_SPEECH_TO_TEXT_LOCATION)
|
||||
self._backend_factory: Final = backend_factory
|
||||
self._transformer: Final = ChirpEventTransformer()
|
||||
self._config: ChirpSessionConfig | None = None
|
||||
self._session_id: str | None = None
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict[str, str], # mutable-ok: BaseRealtimeConfig contract
|
||||
model: str,
|
||||
api_key: str | None = None,
|
||||
) -> dict[str, str]: # mutable-ok: BaseRealtimeConfig contract
|
||||
return headers
|
||||
|
||||
def get_complete_url(self, api_base: str | None, model: str, api_key: str | None = None) -> str:
|
||||
if not is_vertex_speech_to_text_model(model):
|
||||
raise ValueError(f"Unsupported Speech-to-Text streaming model: {model}")
|
||||
return _api_endpoint(api_base) if api_base else speech_to_text_host(self._location)
|
||||
|
||||
async def open_backend(self, url: str, headers: Mapping[str, str]) -> RealtimeBackend | None:
|
||||
return self._backend_factory(
|
||||
SpeechStreamingTarget(
|
||||
api_endpoint=url,
|
||||
recognizer=f"projects/{self._project}/locations/{self._location}/recognizers/_",
|
||||
resolve_access_token=self._resolve_access_token,
|
||||
)
|
||||
)
|
||||
|
||||
def is_setup_message(self, msg_obj: Mapping[str, object]) -> bool:
|
||||
return msg_obj.get("kind") == "configure"
|
||||
|
||||
def transform_session_created_event(
|
||||
self,
|
||||
model: str,
|
||||
logging_session_id: str,
|
||||
session_configuration_request: str | None = None,
|
||||
) -> OpenAIRealtimeTranscriptionSessionCreated:
|
||||
self._session_id = logging_session_id
|
||||
return session_created_event(default_session_config(model), logging_session_id)
|
||||
|
||||
def transform_realtime_request(
|
||||
self,
|
||||
message: str,
|
||||
model: str,
|
||||
session_configuration_request: str | None = None,
|
||||
) -> tuple[str | bytes, ...]:
|
||||
request: Final = json_object(message, ChirpProtocolError)
|
||||
event_type: Final = request.get("type")
|
||||
if event_type in ("session.update", "transcription_session.update"):
|
||||
return self._configure(message, model)
|
||||
if event_type == "input_audio_buffer.append":
|
||||
return self._append_audio(request)
|
||||
if event_type in ("input_audio_buffer.commit", "input_audio_buffer.end"):
|
||||
self._require_config()
|
||||
return (_FINISH_TURN_COMMAND,)
|
||||
if event_type == "input_audio_buffer.clear":
|
||||
self._require_config()
|
||||
return (_DISCARD_TURN_COMMAND,)
|
||||
verbose_logger.debug("Speech-to-Text streaming: dropping unsupported client event %s", event_type)
|
||||
return ()
|
||||
|
||||
def unbilled_usage_on_session_close(self, model: str) -> RealtimeInputAudioTranscriptionUsage | None:
|
||||
return self._transformer.take_unbilled_usage()
|
||||
|
||||
def transform_realtime_response(
|
||||
self,
|
||||
message: str | bytes,
|
||||
model: str,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
realtime_response_transform_input: RealtimeResponseTransformInput,
|
||||
) -> RealtimeResponseTypedDict:
|
||||
frame: Final = _STREAMING_EVENT_ADAPTER.validate_json(message)
|
||||
events: Final = list(self._transformer.transform(frame)) # mutable-ok: response field is a list
|
||||
result: Final[RealtimeResponseTypedDict] = {
|
||||
"response": events,
|
||||
"current_output_item_id": realtime_response_transform_input.get("current_output_item_id"),
|
||||
"current_response_id": realtime_response_transform_input.get("current_response_id"),
|
||||
"current_delta_chunks": realtime_response_transform_input.get("current_delta_chunks"),
|
||||
"current_conversation_id": realtime_response_transform_input.get("current_conversation_id"),
|
||||
"current_item_chunks": realtime_response_transform_input.get("current_item_chunks"),
|
||||
"current_delta_type": realtime_response_transform_input.get("current_delta_type"),
|
||||
"session_configuration_request": realtime_response_transform_input.get("session_configuration_request"),
|
||||
}
|
||||
return result
|
||||
|
||||
def _configure(self, message: str, model: str) -> tuple[str, ...]:
|
||||
if self._config is not None:
|
||||
verbose_logger.debug("Speech-to-Text streaming: ignoring session.update after the stream was configured")
|
||||
return ()
|
||||
config: Final = parse_chirp_session_update(message, model)
|
||||
self._config = config
|
||||
self._transformer.configure(config, self._session_id or f"sess_{uuid.uuid4().hex}")
|
||||
return (config.configure_command(),)
|
||||
|
||||
def _append_audio(self, request: Mapping[str, JsonValue]) -> tuple[bytes, ...]:
|
||||
self._require_config()
|
||||
audio: Final = decode_pcm16_append(request.get("audio"), error=ChirpProtocolError)
|
||||
return tuple(
|
||||
audio[start : start + MAX_AUDIO_MESSAGE_BYTES] for start in range(0, len(audio), MAX_AUDIO_MESSAGE_BYTES)
|
||||
)
|
||||
|
||||
def _require_config(self) -> ChirpSessionConfig:
|
||||
if self._config is None:
|
||||
raise ChirpProtocolError("session.update must configure the session before audio is sent")
|
||||
return self._config
|
||||
|
||||
|
||||
def _api_endpoint(api_base: str) -> str:
|
||||
without_scheme: Final = api_base.split("://", 1)[-1]
|
||||
return without_scheme.split("/", 1)[0]
|
||||
|
|
@ -42,6 +42,10 @@ def validate_vertex_transcription_location(location: str | None, default_locatio
|
|||
raise VertexAIError(status_code=400, message=str(e)) from e
|
||||
|
||||
|
||||
def speech_to_text_host(location: str) -> str:
|
||||
return "speech.googleapis.com" if location == "global" else f"{location}-speech.googleapis.com"
|
||||
|
||||
|
||||
def validate_vertex_transcription_project_id(project_id: str) -> str:
|
||||
if not project_id or ".." in project_id or any(c in project_id for c in _URL_UNSAFE_PROJECT_CHARS):
|
||||
raise VertexAIError(status_code=400, message=f"Invalid vertex_project format: {project_id!r}")
|
||||
|
|
@ -122,8 +126,7 @@ class VertexAIAudioTranscriptionConfig(BaseAudioTranscriptionConfig, VertexBase)
|
|||
project_id: Final = validate_vertex_transcription_project_id(
|
||||
self.safe_get_vertex_ai_project(litellm_params) or self._resolve_project_id_from_credentials(litellm_params)
|
||||
)
|
||||
host: Final = "speech.googleapis.com" if location == "global" else f"{location}-speech.googleapis.com"
|
||||
base_url: Final = (api_base or f"https://{host}").rstrip("/")
|
||||
base_url: Final = (api_base or f"https://{speech_to_text_host(location)}").rstrip("/")
|
||||
return f"{base_url}/v2/projects/{project_id}/locations/{location}/recognizers/_:recognize"
|
||||
|
||||
def _resolve_project_id_from_credentials(self, litellm_params: dict) -> str:
|
||||
|
|
|
|||
|
|
@ -12,10 +12,16 @@ Auth: OAuth2 Bearer token (not an API key).
|
|||
"""
|
||||
|
||||
import json
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Final
|
||||
|
||||
from litellm import verbose_logger
|
||||
from litellm.llms.gemini.realtime.transformation import GeminiRealtimeConfig
|
||||
from litellm.llms.vertex_ai.audio_transcription.realtime_transformation import (
|
||||
VertexChirpRealtimeConfig,
|
||||
is_vertex_speech_to_text_model,
|
||||
)
|
||||
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
|
||||
|
||||
|
||||
class VertexAIRealtimeConfig(GeminiRealtimeConfig):
|
||||
|
|
@ -232,3 +238,20 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig):
|
|||
return []
|
||||
|
||||
return super().transform_realtime_request(message, model, session_configuration_request)
|
||||
|
||||
|
||||
def vertex_realtime_config(
|
||||
model: str,
|
||||
*,
|
||||
access_token: str,
|
||||
resolve_access_token: Callable[[], Awaitable[str]],
|
||||
project: str,
|
||||
location: str | None,
|
||||
) -> VertexAIRealtimeConfig | VertexChirpRealtimeConfig:
|
||||
if is_vertex_speech_to_text_model(model):
|
||||
return VertexChirpRealtimeConfig(resolve_access_token=resolve_access_token, project=project, location=location)
|
||||
return VertexAIRealtimeConfig(
|
||||
access_token=access_token,
|
||||
project=project,
|
||||
location=VertexBase.get_vertex_region(vertex_region=location, model=model),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -48905,7 +48905,8 @@
|
|||
"mode": "audio_transcription",
|
||||
"source": "https://cloud.google.com/speech-to-text/pricing",
|
||||
"supported_endpoints": [
|
||||
"/v1/audio/transcriptions"
|
||||
"/v1/audio/transcriptions",
|
||||
"/v1/realtime"
|
||||
]
|
||||
},
|
||||
"vertex_ai/claude-3-5-haiku": {
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
# Experimental MCP Server Change Guidelines
|
||||
|
||||
Read @../../../../CLAUDE.md and @CLAUDE.md before changing this package.
|
||||
Read @../../../../AGENTS.md before changing this package.
|
||||
|
||||
This directory owns the proxy-hosted MCP server implementation. Keep changes
|
||||
inside the module that owns the behavior, and only reach outside this package
|
||||
|
|
@ -14,7 +14,6 @@ Respect the current package boundaries:
|
|||
```text
|
||||
litellm/proxy/_experimental/mcp_server/
|
||||
AGENTS.md
|
||||
CLAUDE.md
|
||||
server.py # ASGI/MCP route handling, sessions, tool calls [PR7: 7-arm only — move BYOK/OAuth pre-fetch into resolver]
|
||||
mcp_server_manager.py # upstream server registry, clients, tool routing [PR7: _create_mcp_client swaps resolve_mcp_auth -> resolve_credentials]
|
||||
auth/
|
||||
|
|
@ -68,8 +67,10 @@ module materially harder to understand.
|
|||
auth, SSE, streamable HTTP, and stdio as separate flows. Do not collapse them
|
||||
behind a single generic branch unless tests prove every mode still behaves
|
||||
correctly.
|
||||
- Be especially careful with legacy `delegate_auth_to_upstream: true`. The local
|
||||
`CLAUDE.md` explains its admitted replacement and public discovery contract.
|
||||
- Be especially careful with legacy `delegate_auth_to_upstream: true`. `auth_type: oauth2`
|
||||
with `delegate_auth_to_upstream: true` is deprecated: LiteLLM admission is required
|
||||
for matching MCP routes. Use `auth_type: oauth_delegate` for client-forwarded OAuth.
|
||||
OAuth discovery endpoints stay public so clients can start the RFC 9728 flow.
|
||||
- Keep database-backed fields in sync across migrations, typed models under
|
||||
`litellm/types/mcp.py` or `litellm/types/mcp_server/`, config loading, this
|
||||
package, and dashboard state when the field is user-visible.
|
||||
|
|
|
|||
|
|
@ -1 +0,0 @@
|
|||
MCP note: **`auth_type: oauth2` with `delegate_auth_to_upstream: true` is deprecated** - LiteLLM admission is required for matching MCP routes. Use `auth_type: oauth_delegate` for client-forwarded OAuth. OAuth discovery endpoints stay public so clients can start the RFC 9728 flow
|
||||
|
|
@ -233,6 +233,11 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = (
|
|||
module_path="litellm.proxy.anthropic_endpoints.skills_endpoints",
|
||||
path_prefixes=("/v1/skills", "/skills"),
|
||||
),
|
||||
LazyFeature(
|
||||
name="claude_code_gateway",
|
||||
module_path="litellm.proxy.anthropic_endpoints.gateway_endpoints",
|
||||
path_prefixes=("/claude_code_gateway",),
|
||||
),
|
||||
LazyFeature(
|
||||
name="langfuse_passthrough",
|
||||
module_path="litellm.proxy.vertex_ai_endpoints.langfuse_endpoints",
|
||||
|
|
|
|||
|
|
@ -5235,6 +5235,12 @@
|
|||
}
|
||||
}
|
||||
},
|
||||
"claude_code_gateway": {
|
||||
"components": {
|
||||
"schemas": {}
|
||||
},
|
||||
"paths": {}
|
||||
},
|
||||
"claude_code_marketplace": {
|
||||
"components": {
|
||||
"schemas": {
|
||||
|
|
|
|||
|
|
@ -512,6 +512,8 @@ class LiteLLMRoutes(enum.Enum):
|
|||
anthropic_routes = [
|
||||
"/v1/messages",
|
||||
"/v1/messages/count_tokens",
|
||||
"/claude_code_gateway/v1/messages",
|
||||
"/claude_code_gateway/v1/messages/count_tokens",
|
||||
"/v1/skills",
|
||||
"/v1/skills/{skill_id}",
|
||||
"/claude-code/marketplace.json",
|
||||
|
|
@ -889,6 +891,11 @@ class LiteLLMRoutes(enum.Enum):
|
|||
# of; a caller who administers none gets an empty result set.
|
||||
"/organization/daily/activity",
|
||||
"/user/available_roles", # read-only role metadata; any authenticated user may read
|
||||
# Claude Code gateway: the signed-in CLI fetches its managed settings and posts its own telemetry
|
||||
"/claude_code_gateway/managed/settings",
|
||||
"/claude_code_gateway/v1/metrics",
|
||||
"/claude_code_gateway/v1/logs",
|
||||
"/claude_code_gateway/v1/traces",
|
||||
"/user/list", # org admins checked in endpoint; non-admins get 403
|
||||
"/management/v1/users/bulk_delete", # proxy admins delete anyone, org admins only their orgs' users; others 403
|
||||
"/model/{model_id}/update",
|
||||
|
|
@ -2604,6 +2611,14 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase):
|
|||
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",
|
||||
)
|
||||
enable_claude_code_gateway: bool | None = Field(
|
||||
None,
|
||||
description="serve the Claude Code gateway protocol (https://code.claude.com/docs/en/claude-apps-gateway) under /claude_code_gateway: OAuth device-flow sign-in reusing proxy SSO, plus managed settings and OTLP telemetry ingestion. Off by default",
|
||||
)
|
||||
claude_code_gateway_managed_settings: dict[str, Any] | None = Field(
|
||||
None,
|
||||
description="Claude Code managed-settings.json served verbatim at the gateway's /claude_code_gateway/managed/settings endpoint. When unset the endpoint returns 404 (no managed policy)",
|
||||
)
|
||||
database_url: str | None = Field(
|
||||
None,
|
||||
description="connect to a postgres db - needed for generating temporary keys + tracking spend / key",
|
||||
|
|
|
|||
397
litellm/proxy/anthropic_endpoints/gateway_endpoints.py
Normal file
397
litellm/proxy/anthropic_endpoints/gateway_endpoints.py
Normal file
|
|
@ -0,0 +1,397 @@
|
|||
"""
|
||||
Claude Code gateway protocol.
|
||||
|
||||
Implements the wire contract the Claude Code CLI uses to talk to a gateway:
|
||||
OAuth 2.0 device-authorization sign-in (RFC 8414 / RFC 8628), inference via the
|
||||
Anthropic Messages API, managed settings, and OTLP telemetry ingestion. See
|
||||
https://code.claude.com/docs/en/claude-apps-gateway.
|
||||
|
||||
Everything lives under the ``/claude_code_gateway`` base so operators point
|
||||
Claude Code at ``https://<proxy-host>/claude_code_gateway`` via ``/login``. The
|
||||
device flow reuses the proxy's existing SSO login machinery: the browser leg is
|
||||
served by ``/sso/key/generate`` and the shared ``cli_sso_session_cache`` flow,
|
||||
so the bearer token minted here is the same session JWT the LiteLLM CLI uses and
|
||||
is accepted by every bearer-authenticated proxy route.
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import secrets
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from fastapi import APIRouter, Depends, Request, Response
|
||||
from fastapi.responses import JSONResponse
|
||||
from pydantic import BaseModel, Field, TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.constants import (
|
||||
CLI_JWT_EXPIRATION_HOURS,
|
||||
CLI_SSO_SESSION_TTL_SECONDS,
|
||||
LITELLM_CLI_SOURCE_IDENTIFIER,
|
||||
)
|
||||
from litellm.proxy._types import LiteLLM_UserTable, LitellmUserRoles
|
||||
from litellm.proxy.anthropic_endpoints.endpoints import anthropic_response, count_tokens
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.common_utils.http_parsing_utils import _safe_set_request_parsed_body
|
||||
from litellm.proxy.management_endpoints.ui_sso import CliSsoTeamDetail
|
||||
|
||||
GATEWAY_PREFIX: Final = "/claude_code_gateway"
|
||||
_DEVICE_CODE_GRANT: Final = "urn:ietf:params:oauth:grant-type:device_code"
|
||||
_REFRESH_TOKEN_GRANT: Final = "refresh_token"
|
||||
_DEVICE_CODE_SEPARATOR: Final = "."
|
||||
_DEVICE_POLL_INTERVAL_SECONDS: Final = 5
|
||||
_SECONDS_PER_HOUR: Final = 3600
|
||||
_MANAGED_SETTINGS_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
_NO_SETTINGS: Final = MappingProxyType({})
|
||||
_POST_ONLY: Final = ["POST"] # mutable-ok: FastAPI's add_api_route only accepts a list of methods
|
||||
|
||||
|
||||
class _GatewaySessionData(BaseModel):
|
||||
user_id: str
|
||||
user_role: LitellmUserRoles
|
||||
models: list[str] = Field(default_factory=list)
|
||||
teams: tuple[str, ...] = ()
|
||||
team_details: object | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _GatewayLogin:
|
||||
user_info: LiteLLM_UserTable
|
||||
team_id: str | None
|
||||
team: CliSsoTeamDetail
|
||||
|
||||
|
||||
class _OAuthErrorBody(BaseModel):
|
||||
error: str
|
||||
error_description: str | None = None
|
||||
|
||||
|
||||
class _AuthorizationServerMetadata(BaseModel):
|
||||
issuer: str
|
||||
device_authorization_endpoint: str
|
||||
token_endpoint: str
|
||||
grant_types_supported: tuple[str, ...]
|
||||
|
||||
|
||||
class _DeviceAuthorizationBody(BaseModel):
|
||||
device_code: str
|
||||
user_code: str
|
||||
verification_uri: str
|
||||
verification_uri_complete: str | None = None
|
||||
expires_in: int
|
||||
interval: int
|
||||
|
||||
|
||||
class _AccessTokenBody(BaseModel):
|
||||
access_token: str
|
||||
expires_in: int
|
||||
token_type: str = "Bearer"
|
||||
|
||||
|
||||
class _ManagedSettingsBody(BaseModel):
|
||||
uuid: str
|
||||
checksum: str
|
||||
settings: dict[str, object]
|
||||
|
||||
|
||||
def _general_settings() -> Mapping[str, object]:
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
return general_settings or _NO_SETTINGS
|
||||
|
||||
|
||||
def _is_gateway_enabled() -> bool:
|
||||
return bool(_general_settings().get("enable_claude_code_gateway", False))
|
||||
|
||||
|
||||
def ensure_gateway_enabled() -> None:
|
||||
from fastapi import HTTPException
|
||||
|
||||
if not _is_gateway_enabled():
|
||||
raise HTTPException(status_code=404, detail="Claude Code gateway is not enabled")
|
||||
|
||||
|
||||
def _managed_settings() -> dict[str, object] | None:
|
||||
settings: Final[object] = _general_settings().get("claude_code_gateway_managed_settings")
|
||||
if not isinstance(settings, dict):
|
||||
return None
|
||||
return _MANAGED_SETTINGS_ADAPTER.validate_python(settings)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _OAuthError:
|
||||
status_code: int
|
||||
error: str
|
||||
description: str | None = None
|
||||
|
||||
|
||||
def _oauth_error_response(err: _OAuthError) -> JSONResponse:
|
||||
body: Final = _OAuthErrorBody(error=err.error, error_description=err.description)
|
||||
return JSONResponse(status_code=err.status_code, content=body.model_dump(exclude_none=True))
|
||||
|
||||
|
||||
router: Final = APIRouter(
|
||||
prefix=GATEWAY_PREFIX,
|
||||
tags=["Claude Code gateway"], # mutable-ok: FastAPI's APIRouter only accepts a list of tags
|
||||
)
|
||||
_GATEWAY_ENABLED: Final = (Depends(ensure_gateway_enabled),)
|
||||
_AUTHENTICATED: Final = (Depends(user_api_key_auth),)
|
||||
|
||||
router.add_api_route(
|
||||
"/v1/messages",
|
||||
anthropic_response,
|
||||
methods=_POST_ONLY,
|
||||
dependencies=_GATEWAY_ENABLED,
|
||||
include_in_schema=False,
|
||||
)
|
||||
router.add_api_route(
|
||||
"/v1/messages/count_tokens",
|
||||
count_tokens,
|
||||
methods=_POST_ONLY,
|
||||
dependencies=_GATEWAY_ENABLED,
|
||||
include_in_schema=False,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/.well-known/oauth-authorization-server", include_in_schema=False)
|
||||
async def oauth_authorization_server(request: Request) -> JSONResponse:
|
||||
if not _is_gateway_enabled():
|
||||
return _oauth_error_response(_OAuthError(status_code=404, error="not_found"))
|
||||
|
||||
from litellm.proxy.utils import get_custom_url
|
||||
|
||||
request_base_url: Final = str(request.base_url)
|
||||
metadata: Final = _AuthorizationServerMetadata(
|
||||
issuer=get_custom_url(request_base_url=request_base_url, route="claude_code_gateway"),
|
||||
device_authorization_endpoint=get_custom_url(
|
||||
request_base_url=request_base_url, route="claude_code_gateway/oauth/device_authorization"
|
||||
),
|
||||
token_endpoint=get_custom_url(request_base_url=request_base_url, route="claude_code_gateway/oauth/token"),
|
||||
grant_types_supported=(_DEVICE_CODE_GRANT, _REFRESH_TOKEN_GRANT),
|
||||
)
|
||||
return JSONResponse(content=metadata.model_dump())
|
||||
|
||||
|
||||
@router.post("/oauth/device_authorization", include_in_schema=False)
|
||||
async def device_authorization(request: Request) -> JSONResponse:
|
||||
from urllib.parse import urlencode
|
||||
|
||||
from litellm.proxy.management_endpoints.ui_sso import (
|
||||
_check_cli_sso_start_rate_limit, # pyright: ignore[reportPrivateUsage] # shared device-flow helper
|
||||
_cli_sso_verification_uri_complete_enabled, # pyright: ignore[reportPrivateUsage] # shared device-flow helper
|
||||
_generate_cli_sso_user_code, # pyright: ignore[reportPrivateUsage] # shared device-flow helper
|
||||
_hash_cli_sso_secret, # pyright: ignore[reportPrivateUsage] # shared device-flow helper
|
||||
_normalize_cli_sso_user_code, # pyright: ignore[reportPrivateUsage] # shared device-flow helper
|
||||
_set_cli_sso_flow, # pyright: ignore[reportPrivateUsage] # shared device-flow helper
|
||||
)
|
||||
from litellm.proxy.proxy_server import cli_sso_session_cache
|
||||
from litellm.proxy.utils import get_custom_url
|
||||
|
||||
if not _is_gateway_enabled():
|
||||
return _oauth_error_response(_OAuthError(status_code=404, error="not_found"))
|
||||
|
||||
_check_cli_sso_start_rate_limit(
|
||||
request=request,
|
||||
cache=cli_sso_session_cache,
|
||||
use_x_forwarded_for=bool(_general_settings().get("use_x_forwarded_for", False)),
|
||||
)
|
||||
|
||||
login_id: Final = f"cli-{secrets.token_urlsafe(24)}"
|
||||
poll_secret: Final = secrets.token_urlsafe(32)
|
||||
user_code: Final = _generate_cli_sso_user_code()
|
||||
flow: Final = { # mutable-ok: the shared CLI SSO cache entry is a dict the browser leg mutates
|
||||
"poll_secret_hash": _hash_cli_sso_secret(poll_secret),
|
||||
"user_code_hash": _hash_cli_sso_secret(_normalize_cli_sso_user_code(user_code)),
|
||||
"sso_complete": False,
|
||||
"user_code_verified": False,
|
||||
"session_data": None,
|
||||
}
|
||||
_set_cli_sso_flow(login_id=login_id, cache=cli_sso_session_cache, flow=flow)
|
||||
|
||||
request_base_url: Final = str(request.base_url)
|
||||
verification_uri: Final = get_custom_url(request_base_url=request_base_url, route="sso/key/generate")
|
||||
query: Final = MappingProxyType({"source": LITELLM_CLI_SOURCE_IDENTIFIER, "key": login_id})
|
||||
body: Final = _DeviceAuthorizationBody(
|
||||
device_code=f"{login_id}{_DEVICE_CODE_SEPARATOR}{poll_secret}",
|
||||
user_code=user_code,
|
||||
verification_uri=f"{verification_uri}?{urlencode(query)}",
|
||||
verification_uri_complete=(
|
||||
f"{verification_uri}?{urlencode(MappingProxyType({**query, 'user_code': user_code}))}"
|
||||
if _cli_sso_verification_uri_complete_enabled()
|
||||
else None
|
||||
),
|
||||
expires_in=CLI_SSO_SESSION_TTL_SECONDS,
|
||||
interval=_DEVICE_POLL_INTERVAL_SECONDS,
|
||||
)
|
||||
return JSONResponse(content=body.model_dump(exclude_none=True))
|
||||
|
||||
|
||||
def _validate_login(flow: Mapping[str, object]) -> _GatewayLogin | _OAuthError:
|
||||
from litellm.proxy.management_endpoints.ui_sso import selected_cli_sso_team_detail
|
||||
|
||||
try:
|
||||
session_data: Final = _GatewaySessionData.model_validate(flow.get("session_data"))
|
||||
except ValidationError as err:
|
||||
verbose_proxy_logger.warning("Claude Code gateway login session is malformed: %s", err)
|
||||
return _OAuthError(
|
||||
status_code=400, error="invalid_grant", description="The login session is malformed; sign in again"
|
||||
)
|
||||
|
||||
team_id: Final = session_data.teams[0] if session_data.teams else None
|
||||
selected_team: Final = selected_cli_sso_team_detail(team_details=session_data.team_details, team_id=team_id)
|
||||
if selected_team is None:
|
||||
return _OAuthError(
|
||||
status_code=400,
|
||||
error="invalid_grant",
|
||||
description=f"Could not resolve the model grants for team {team_id}; sign in again",
|
||||
)
|
||||
|
||||
user_info: Final = LiteLLM_UserTable(
|
||||
user_id=session_data.user_id,
|
||||
user_role=session_data.user_role.value,
|
||||
models=session_data.models,
|
||||
)
|
||||
return _GatewayLogin(user_info=user_info, team_id=team_id, team=selected_team)
|
||||
|
||||
|
||||
def _mint_access_token(login: _GatewayLogin) -> str:
|
||||
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken
|
||||
|
||||
return ExperimentalUIJWTToken.get_cli_jwt_auth_token(
|
||||
user_info=login.user_info,
|
||||
team_id=login.team_id,
|
||||
team_alias=login.team.team_alias,
|
||||
team_models=login.team.team_models,
|
||||
team_model_aliases=login.team.team_model_aliases,
|
||||
max_budget=None,
|
||||
)
|
||||
|
||||
|
||||
async def _claim_device_code(login_id: str, cache: DualCache) -> bool:
|
||||
from litellm.proxy.management_endpoints.ui_sso import (
|
||||
_get_cli_sso_flow_cache_key, # pyright: ignore[reportPrivateUsage] # shared device-flow helper
|
||||
)
|
||||
|
||||
claims: Final = await cache.async_increment_cache(
|
||||
key=f"{_get_cli_sso_flow_cache_key(login_id)}:claimed",
|
||||
value=1,
|
||||
ttl=CLI_SSO_SESSION_TTL_SECONDS,
|
||||
)
|
||||
return claims == 1
|
||||
|
||||
|
||||
async def _handle_device_code_grant(device_code: str | None) -> JSONResponse:
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.proxy.management_endpoints.ui_sso import (
|
||||
_get_cli_sso_flow_cache_key, # pyright: ignore[reportPrivateUsage] # shared device-flow helper
|
||||
_get_cli_sso_flow_or_raise, # pyright: ignore[reportPrivateUsage] # shared device-flow helper
|
||||
_verify_cli_sso_poll_secret, # pyright: ignore[reportPrivateUsage] # shared device-flow helper
|
||||
)
|
||||
from litellm.proxy.proxy_server import cli_sso_session_cache
|
||||
|
||||
if not device_code:
|
||||
return _oauth_error_response(
|
||||
_OAuthError(status_code=400, error="invalid_request", description="device_code is required")
|
||||
)
|
||||
|
||||
login_id, _, poll_secret = device_code.partition(_DEVICE_CODE_SEPARATOR)
|
||||
try:
|
||||
flow: Final = _get_cli_sso_flow_or_raise(login_id=login_id, cache=cli_sso_session_cache)
|
||||
except HTTPException:
|
||||
return _oauth_error_response(_OAuthError(status_code=400, error="expired_token"))
|
||||
|
||||
if not _verify_cli_sso_poll_secret(flow, poll_secret):
|
||||
return _oauth_error_response(_OAuthError(status_code=400, error="expired_token"))
|
||||
|
||||
if not flow.get("sso_complete") or not flow.get("user_code_verified"):
|
||||
return _oauth_error_response(_OAuthError(status_code=400, error="authorization_pending"))
|
||||
|
||||
login: Final = _validate_login(flow)
|
||||
if isinstance(login, _OAuthError):
|
||||
return _oauth_error_response(login)
|
||||
|
||||
access_token: Final = _mint_access_token(login)
|
||||
if not await _claim_device_code(login_id, cli_sso_session_cache):
|
||||
return _oauth_error_response(_OAuthError(status_code=400, error="expired_token"))
|
||||
|
||||
await cli_sso_session_cache.async_delete_cache(key=_get_cli_sso_flow_cache_key(login_id))
|
||||
body: Final = _AccessTokenBody(access_token=access_token, expires_in=CLI_JWT_EXPIRATION_HOURS * _SECONDS_PER_HOUR)
|
||||
return JSONResponse(content=body.model_dump())
|
||||
|
||||
|
||||
@router.post("/oauth/token", include_in_schema=False)
|
||||
async def oauth_token(request: Request) -> JSONResponse:
|
||||
if not _is_gateway_enabled():
|
||||
return _oauth_error_response(_OAuthError(status_code=404, error="not_found"))
|
||||
|
||||
form: Final = await request.form()
|
||||
grant_type: Final = form.get("grant_type")
|
||||
|
||||
if grant_type == _DEVICE_CODE_GRANT:
|
||||
device_code: Final = form.get("device_code")
|
||||
return await _handle_device_code_grant(device_code if isinstance(device_code, str) else None)
|
||||
|
||||
if grant_type == _REFRESH_TOKEN_GRANT:
|
||||
return _oauth_error_response(
|
||||
_OAuthError(
|
||||
status_code=401,
|
||||
error="invalid_grant",
|
||||
description="This gateway does not issue refresh tokens; sign in again",
|
||||
)
|
||||
)
|
||||
|
||||
return _oauth_error_response(
|
||||
_OAuthError(
|
||||
status_code=400, error="unsupported_grant_type", description=f"Unsupported grant_type: {grant_type}"
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@router.get("/managed/settings", include_in_schema=False, dependencies=_AUTHENTICATED)
|
||||
async def managed_settings(request: Request) -> Response:
|
||||
ensure_gateway_enabled()
|
||||
|
||||
settings: Final = _managed_settings()
|
||||
if settings is None:
|
||||
return Response(status_code=404)
|
||||
|
||||
canonical: Final = json.dumps(settings, sort_keys=True, separators=(",", ":"))
|
||||
checksum: Final = "sha256:" + hashlib.sha256(canonical.encode("utf-8")).hexdigest()
|
||||
etag: Final = f'"{checksum}"'
|
||||
headers: Final = MappingProxyType({"ETag": etag})
|
||||
if request.headers.get("If-None-Match") == etag:
|
||||
return Response(status_code=304, headers=headers)
|
||||
body: Final = _ManagedSettingsBody(uuid=checksum, checksum=checksum, settings=settings)
|
||||
return Response(content=body.model_dump_json(), media_type="application/json", headers=headers)
|
||||
|
||||
|
||||
async def _skip_otlp_body_parsing(request: Request) -> None:
|
||||
_safe_set_request_parsed_body(request=request, parsed_body={})
|
||||
|
||||
|
||||
_OTLP_AUTHENTICATED: Final = (Depends(_skip_otlp_body_parsing), *_AUTHENTICATED)
|
||||
|
||||
|
||||
def _accept_otlp() -> Response:
|
||||
ensure_gateway_enabled()
|
||||
return Response(status_code=200)
|
||||
|
||||
|
||||
@router.post("/v1/metrics", include_in_schema=False, dependencies=_OTLP_AUTHENTICATED)
|
||||
async def otlp_metrics() -> Response:
|
||||
return _accept_otlp()
|
||||
|
||||
|
||||
@router.post("/v1/logs", include_in_schema=False, dependencies=_OTLP_AUTHENTICATED)
|
||||
async def otlp_logs() -> Response:
|
||||
return _accept_otlp()
|
||||
|
||||
|
||||
@router.post("/v1/traces", include_in_schema=False, dependencies=_OTLP_AUTHENTICATED)
|
||||
async def otlp_traces() -> Response:
|
||||
return _accept_otlp()
|
||||
|
|
@ -7,6 +7,7 @@ from collections.abc import MutableMapping
|
|||
from typing import Any, Final
|
||||
|
||||
from fastapi import Request
|
||||
from starlette.routing import get_route_path
|
||||
from starlette.types import ASGIApp, Receive, Scope, Send
|
||||
|
||||
import litellm
|
||||
|
|
@ -15,6 +16,12 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
|||
|
||||
# Cache the header name at module level to avoid repeated enum attribute access
|
||||
_AUTHORIZATION_HEADER: Final = SpecialHeaders.openai_authorization.value # "Authorization"
|
||||
_METRICS_MOUNT: Final = "/metrics"
|
||||
|
||||
|
||||
def _is_metrics_route(scope: Scope) -> bool:
|
||||
route_path: Final = get_route_path(scope)
|
||||
return route_path == _METRICS_MOUNT or route_path.startswith(_METRICS_MOUNT + "/")
|
||||
|
||||
|
||||
class PrometheusAuthMiddleware:
|
||||
|
|
@ -36,7 +43,7 @@ class PrometheusAuthMiddleware:
|
|||
|
||||
async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None:
|
||||
# Fast path: only inspect HTTP requests; pass through websocket/lifespan immediately
|
||||
if scope["type"] != "http" or "/metrics" not in scope.get("path", ""):
|
||||
if scope["type"] != "http" or not _is_metrics_route(scope):
|
||||
await self.app(scope, receive, send)
|
||||
return
|
||||
|
||||
|
|
|
|||
|
|
@ -38,7 +38,8 @@ from ..llms.azure.realtime.handler import AzureOpenAIRealtime, azure_realtime_pr
|
|||
from ..llms.bedrock.realtime.handler import BedrockRealtime
|
||||
from ..llms.custom_httpx.http_handler import get_shared_realtime_ssl_context
|
||||
from ..llms.openai.realtime.handler import OpenAIRealtime
|
||||
from ..llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig
|
||||
from ..llms.vertex_ai.audio_transcription.realtime_transformation import is_vertex_speech_to_text_model
|
||||
from ..llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig, vertex_realtime_config
|
||||
from ..llms.vertex_ai.vertex_llm_base import VertexBase
|
||||
from ..llms.xai.realtime.handler import XAIRealtime
|
||||
from ..utils import client as wrapper_client
|
||||
|
|
@ -541,8 +542,6 @@ async def _arealtime(
|
|||
or get_secret_str("VERTEXAI_LOCATION")
|
||||
)
|
||||
|
||||
resolved_location: Final = vertex_llm_base.get_vertex_region(vertex_region=vertex_location, model=model)
|
||||
|
||||
(
|
||||
access_token,
|
||||
resolved_project,
|
||||
|
|
@ -553,17 +552,28 @@ async def _arealtime(
|
|||
timeout_seconds=REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS,
|
||||
)
|
||||
|
||||
vertex_realtime_config: Final = VertexAIRealtimeConfig(
|
||||
async def resolve_vertex_access_token() -> str:
|
||||
refreshed_token, _ = await _resolve_vertex_access_token_bounded(
|
||||
credentials=vertex_credentials,
|
||||
project_id=resolved_project,
|
||||
resolver=vertex_access_token_resolver,
|
||||
timeout_seconds=REALTIME_CREDENTIAL_RESOLUTION_TIMEOUT_SECONDS,
|
||||
)
|
||||
return refreshed_token
|
||||
|
||||
vertex_provider_config: Final = vertex_realtime_config(
|
||||
model,
|
||||
access_token=access_token,
|
||||
resolve_access_token=resolve_vertex_access_token,
|
||||
project=resolved_project,
|
||||
location=resolved_location,
|
||||
location=vertex_location,
|
||||
)
|
||||
|
||||
await base_llm_http_handler.async_realtime(
|
||||
model=model,
|
||||
websocket=websocket,
|
||||
logging_obj=litellm_logging_obj,
|
||||
provider_config=vertex_realtime_config,
|
||||
provider_config=vertex_provider_config,
|
||||
api_base=dynamic_api_base or litellm_params.api_base,
|
||||
api_key=None,
|
||||
client=client,
|
||||
|
|
@ -684,6 +694,11 @@ async def _realtime_health_check(
|
|||
api_base=resolved_api_base or "https://api.x.ai/v1", query_params={"model": model}
|
||||
)
|
||||
elif custom_llm_provider == "vertex_ai":
|
||||
if is_vertex_speech_to_text_model(model):
|
||||
raise ValueError(
|
||||
f"Realtime health checks are not supported for Speech-to-Text streaming model {model};"
|
||||
" health check it with mode audio_transcription"
|
||||
)
|
||||
vertex_model_params: Final = dict(resolved_params)
|
||||
resolved_location: Final = vertex_llm_base.get_vertex_region(
|
||||
vertex_region=VertexBase.safe_get_vertex_ai_location(vertex_model_params),
|
||||
|
|
|
|||
|
|
@ -1,4 +1,6 @@
|
|||
from pydantic import BaseModel
|
||||
from typing import Annotated, Literal
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
|
||||
|
|
@ -38,3 +40,66 @@ class VertexSpeechToTextResponseMetadata(BaseModel):
|
|||
class VertexSpeechToTextRecognizeResponse(BaseModel):
|
||||
results: list[VertexSpeechToTextResult] = []
|
||||
metadata: VertexSpeechToTextResponseMetadata | None = None
|
||||
|
||||
|
||||
class VertexSpeechStreamingConfigure(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
kind: Literal["configure"] = "configure"
|
||||
model: str
|
||||
language_codes: tuple[str, ...]
|
||||
sample_rate_hertz: int
|
||||
|
||||
|
||||
class VertexSpeechStreamingFinishTurn(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
kind: Literal["finish_turn"] = "finish_turn"
|
||||
|
||||
|
||||
class VertexSpeechStreamingDiscardTurn(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
kind: Literal["discard_turn"] = "discard_turn"
|
||||
|
||||
|
||||
VertexSpeechStreamingCommandUnion = (
|
||||
VertexSpeechStreamingConfigure | VertexSpeechStreamingFinishTurn | VertexSpeechStreamingDiscardTurn
|
||||
)
|
||||
VertexSpeechStreamingCommand = Annotated[VertexSpeechStreamingCommandUnion, Field(discriminator="kind")]
|
||||
|
||||
|
||||
class VertexSpeechStreamingResult(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
transcript: str
|
||||
is_final: bool
|
||||
|
||||
|
||||
class VertexSpeechStreamingResponse(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
kind: Literal["response"] = "response"
|
||||
speech_event: Literal["none", "begin", "end"]
|
||||
results: tuple[VertexSpeechStreamingResult, ...]
|
||||
billed_seconds: float
|
||||
|
||||
|
||||
class VertexSpeechStreamingConfigured(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
kind: Literal["configured"] = "configured"
|
||||
|
||||
|
||||
class VertexSpeechStreamingTurnFinished(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
kind: Literal["turn_finished"] = "turn_finished"
|
||||
|
||||
|
||||
class VertexSpeechStreamingTurnDiscarded(BaseModel):
|
||||
model_config = ConfigDict(frozen=True)
|
||||
kind: Literal["turn_discarded"] = "turn_discarded"
|
||||
billed_seconds: float
|
||||
|
||||
|
||||
VertexSpeechStreamingEventUnion = (
|
||||
VertexSpeechStreamingResponse
|
||||
| VertexSpeechStreamingConfigured
|
||||
| VertexSpeechStreamingTurnFinished
|
||||
| VertexSpeechStreamingTurnDiscarded
|
||||
)
|
||||
VertexSpeechStreamingEvent = Annotated[VertexSpeechStreamingEventUnion, Field(discriminator="kind")]
|
||||
|
|
|
|||
|
|
@ -48905,7 +48905,8 @@
|
|||
"mode": "audio_transcription",
|
||||
"source": "https://cloud.google.com/speech-to-text/pricing",
|
||||
"supported_endpoints": [
|
||||
"/v1/audio/transcriptions"
|
||||
"/v1/audio/transcriptions",
|
||||
"/v1/realtime"
|
||||
]
|
||||
},
|
||||
"vertex_ai/claude-3-5-haiku": {
|
||||
|
|
|
|||
|
|
@ -129,6 +129,12 @@ grpc = [
|
|||
# Newest non-yanked release older than the 30-day cutoff.
|
||||
"grpcio==1.78.0",
|
||||
]
|
||||
stt-vertex-chirp = [
|
||||
# Google Cloud Speech-to-Text v2 streaming (gRPC) for Chirp models on
|
||||
# /v1/realtime. Imported lazily inside the backend so litellm core stays
|
||||
# usable without it.
|
||||
"google-cloud-speech>=2.40.0,<3.0",
|
||||
]
|
||||
stt-nvidia-riva = [
|
||||
# NVIDIA Riva STT provider (gRPC). These are imported lazily inside the
|
||||
# provider handler so litellm core remains usable without them.
|
||||
|
|
@ -152,6 +158,7 @@ proxy-runtime = [
|
|||
# Keep these in a dedicated extra so uv-based images preserve the same
|
||||
# feature surface without forcing the base SDK install to grow.
|
||||
"google-cloud-aiplatform>=1.133.0,<2.0",
|
||||
"google-cloud-speech>=2.40.0,<3.0",
|
||||
"google-genai>=1.37.0,<2.0",
|
||||
"anthropic[vertex]>=0.84.0,<1.0",
|
||||
"grpcio==1.78.0",
|
||||
|
|
@ -270,6 +277,7 @@ ci = [
|
|||
"langgraph>=1.2.4,<1.3.0",
|
||||
"langgraph-prebuilt>=1.1.0,<1.3.0",
|
||||
"claude-agent-sdk==0.1.44",
|
||||
"google-cloud-speech==2.40.0",
|
||||
]
|
||||
healthcheck = [
|
||||
"httpx==0.28.1",
|
||||
|
|
@ -329,9 +337,6 @@ litellm-enterprise = { workspace = true }
|
|||
[tool.uv.workspace]
|
||||
members = ["enterprise", "litellm-proxy-extras"]
|
||||
|
||||
[tool.isort]
|
||||
profile = "black"
|
||||
|
||||
[tool.commitizen]
|
||||
version = "1.103.0"
|
||||
version_files = [
|
||||
|
|
|
|||
|
|
@ -3,10 +3,10 @@
|
|||
Standalone script to test tool allowlist enforcement and tool name extraction.
|
||||
|
||||
Run from repo root:
|
||||
poetry run python scripts/test_tool_allowlist_script.py
|
||||
uv run python scripts/test_tool_allowlist_script.py
|
||||
|
||||
Or run the unit tests:
|
||||
poetry run pytest tests/test_litellm/proxy/test_tools_allowlist_enforcement.py -v
|
||||
uv run pytest tests/test_litellm/proxy/test_tools_allowlist_enforcement.py -v
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
|
|
@ -148,7 +148,7 @@ def main():
|
|||
asyncio.run(test_check_tools_allowlist())
|
||||
print("Done. For full unit tests run:")
|
||||
print(
|
||||
" poetry run pytest tests/test_litellm/proxy/test_tools_allowlist_enforcement.py -v"
|
||||
" uv run pytest tests/test_litellm/proxy/test_tools_allowlist_enforcement.py -v"
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ anywhere; a small allowlist grandfathers the files that legitimately make raw ca
|
|||
(the transport itself, the root conftest liveness probe, the claude_code version
|
||||
resolver's constant registry URL fetch, and the mcp OAuth client, whose httpx
|
||||
client is the object the official mcp SDK's streamable_http_client requires and so
|
||||
cannot go through the sync requests transport). Referenced by tests/e2e/CLAUDE.md."""
|
||||
cannot go through the sync requests transport). Referenced by tests/e2e/AGENTS.md."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
|
|
|
|||
|
|
@ -142,7 +142,7 @@ def select_tests(changed: tuple[str, ...]) -> tuple[str, ...]:
|
|||
("tests/e2e/guardrails/test_bedrock_guardrail_e2e.py",),
|
||||
("tests/e2e/guardrails/test_bedrock_guardrail_e2e.py",),
|
||||
),
|
||||
(("tests/e2e/logging/helpers.py", "docs/my-website/docs/index.md", "tests/e2e/CLAUDE.md"), ()),
|
||||
(("tests/e2e/logging/helpers.py", "docs/my-website/docs/index.md", "tests/e2e/AGENTS.md"), ()),
|
||||
(
|
||||
("tests/e2e/logging/test_datadog_e2e.py", "tests/e2e/logging/test_datadog_e2e.py"),
|
||||
("tests/e2e/logging/test_datadog_e2e.py",),
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
# e2e harness conventions
|
||||
|
||||
Code-style rules for writing tests under `tests/e2e/`. The harness already encodes the plumbing; your job is the feature-specific behavior, not reinventing it. For what a complete test must do (the lifecycle contract, asserting both recorded state and enforced behavior) and how to run a suite, see `CONTRIBUTING.md` in this directory. Repo-wide conventions live in the root `CLAUDE.md`
|
||||
Code-style rules for writing tests under `tests/e2e/`. The harness already encodes the plumbing; your job is the feature-specific behavior, not reinventing it. For what a complete test must do (the lifecycle contract, asserting both recorded state and enforced behavior) and how to run a suite, see `CONTRIBUTING.md` in this directory. Repo-wide conventions live in the root `AGENTS.md`
|
||||
|
||||
## Suite folders
|
||||
|
||||
|
|
@ -2,7 +2,7 @@
|
|||
|
||||
This directory holds the live end-to-end suites that prove product correctness against a real running proxy and real provider APIs. The goal of this guide is simple: when you ship a feature, you add e2e coverage that walks that feature the way production does, across every route and edge case it touches, so a later change that breaks it fails here first
|
||||
|
||||
Read this before adding a test and i recommend reading through CLAUDE.md
|
||||
Read this before adding a test and i recommend reading through AGENTS.md
|
||||
|
||||
When contributing to this directory, please first discuss the change you wish to make via issue or pull request. We require screenshots and proof of your tests working on a live proxy.
|
||||
|
||||
|
|
@ -134,7 +134,7 @@ One sharp edge: a replayed response reuses the recorded provider response id, an
|
|||
|
||||
Another sharp edge, same root: record and replay derive every per-test token deterministically (the model name included, so a replay regenerates the exact requests the record run sent), which means an edge-wired deployment left in the database by an interrupted earlier run carries the same model name as the fresh one the current run registers. The proxy then holds two deployments under one model group and load-balances across both, and because the leftover's `api_base` points at the earlier run's edge process, which is gone, the calls that land on it fail with a connection error that reads like a transport bug rather than the stale row it is. Give each record or replay run a fresh database, or let a run finish so its own teardown deletes what it registered, and never reuse one long-lived proxy across back-to-back record/replay sessions. CI hands every job its own empty database and its own proxy, so it never sees this
|
||||
|
||||
Replay answers any provider call that drifted from the recording with an HTTP 599 whose body names the computed and closest recorded keys, so the test fails loudly instead of silently going live, and a bundle older than seven days fails at collection time naming its age; either way the fix is to re-record. Only tests that register edge-wired deployments participate: everything else hits its provider live in every mode, so record exactly the suite you replay. If the proxy runs in a container, set `E2E_PROVIDER_EDGE_ADVERTISE_HOST` (e.g. `host.docker.internal`) so the api_base the proxy stores can reach the edge on the pytest host, and `E2E_PROVIDER_EDGE_BIND_HOST=0.0.0.0` so the edge accepts it. The suites wired to the edge today are `quota_management/spend_tracking/test_provider_edge_spend_e2e.py`, `llm_translation/test_chat_completions_contract_e2e.py`, the OpenAI registrations in `llm_translation/test_embeddings_endpoint_e2e.py`, the Anthropic tests in `llm_translation/test_messages_e2e.py`, streamed and not, and the OpenAI batch deployment behind `batches/`. A streamed response replays as the chunk sequence the provider sent rather than one buffered body. See `CLAUDE.md` in this directory for the bundle format, the edge design, and the current limits (Bedrock). The scheduled CI record/replay lane is described above
|
||||
Replay answers any provider call that drifted from the recording with an HTTP 599 whose body names the computed and closest recorded keys, so the test fails loudly instead of silently going live, and a bundle older than seven days fails at collection time naming its age; either way the fix is to re-record. Only tests that register edge-wired deployments participate: everything else hits its provider live in every mode, so record exactly the suite you replay. If the proxy runs in a container, set `E2E_PROVIDER_EDGE_ADVERTISE_HOST` (e.g. `host.docker.internal`) so the api_base the proxy stores can reach the edge on the pytest host, and `E2E_PROVIDER_EDGE_BIND_HOST=0.0.0.0` so the edge accepts it. The suites wired to the edge today are `quota_management/spend_tracking/test_provider_edge_spend_e2e.py`, `llm_translation/test_chat_completions_contract_e2e.py`, the OpenAI registrations in `llm_translation/test_embeddings_endpoint_e2e.py`, the Anthropic tests in `llm_translation/test_messages_e2e.py`, streamed and not, and the OpenAI batch deployment behind `batches/`. A streamed response replays as the chunk sequence the provider sent rather than one buffered body. See `AGENTS.md` in this directory for the bundle format, the edge design, and the current limits (Bedrock). The scheduled CI record/replay lane is described above
|
||||
|
||||
Tests marked `@pytest.mark.e2e` hard-fail when no proxy answers `/health/liveliness`, so a run that goes red with `No live proxy` at setup means the proxy isn't up; they never skip for a missing proxy, so an absent proxy can't be mistaken for a pass
|
||||
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ cost write-back via a cross-run marker baton (design below).
|
|||
Only supported cells are tested. The capability table in `capabilities.py` holds one
|
||||
row per supported (provider, scenario) pair, so there are no skipped cells in the
|
||||
parametrized run. The batches suite never skips: missing provider creds or upstream
|
||||
failures are hard test failures (see `tests/e2e/CLAUDE.md`).
|
||||
failures are hard test failures (see `tests/e2e/AGENTS.md`).
|
||||
|
||||
| Provider | create | retrieve | cancel | list | content download | file backing |
|
||||
|-----------|--------|----------|--------|------|------------------|--------------|
|
||||
|
|
|
|||
|
|
@ -288,7 +288,7 @@ else
|
|||
# Download the tarball and Astral's official .sha256 sidecar to disk
|
||||
# and verify the digest before extracting/executing anything. This
|
||||
# closes the supply-chain trust gap of piping a remote binary
|
||||
# straight into `tar -xzO ... > file ; chmod +x` (see CLAUDE.md
|
||||
# straight into `tar -xzO ... > file ; chmod +x` (see AGENTS.md
|
||||
# "CI Supply-Chain Safety").
|
||||
curl -fsSL --output "${UV_TMPDIR}/${UV_TARBALL_NAME}" "${UV_DOWNLOAD_URL}"
|
||||
curl -fsSL --output "${UV_TMPDIR}/${UV_TARBALL_NAME}.sha256" "${UV_DOWNLOAD_URL}.sha256"
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@
|
|||
This directory is the **denominator** for e2e test coverage: the set of behaviors we
|
||||
want covered, one row per behavior, checked into the repo so coverage is a number we
|
||||
can track instead of a guess. It implements the plan in the "E2E Coverage Tracking"
|
||||
note; the naming grammar lives in `tests/e2e/CLAUDE.md`.
|
||||
note; the naming grammar lives in `tests/e2e/AGENTS.md`.
|
||||
|
||||
## The model
|
||||
|
||||
|
|
|
|||
|
|
@ -3,6 +3,6 @@
|
|||
`schema.py` defines one validated row per customer-noticeable behavior (a "cell").
|
||||
The `*.yaml` files hold the rows, one file per id-prefix. `registry.py` loads and
|
||||
validates them; `collector.py` diffs the registry against the `@pytest.mark.covers`
|
||||
markers on the live tests and reports coverage per module. See tests/e2e/CLAUDE.md
|
||||
markers on the live tests and reports coverage per module. See tests/e2e/AGENTS.md
|
||||
for the naming grammar.
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
# MCP module. Grounded in litellm/proxy/_experimental/mcp_server/. See tests/e2e/CLAUDE.md for the grammar.
|
||||
# MCP module. Grounded in litellm/proxy/_experimental/mcp_server/. See tests/e2e/AGENTS.md for the grammar.
|
||||
- id: mcp.list_tools.api_key.succeeds
|
||||
module: mcp
|
||||
tier: P0
|
||||
|
|
|
|||
|
|
@ -49,7 +49,7 @@ kept commented out in `PROVIDERS` until they pass end-to-end here; re-enable the
|
|||
uncommenting their entry.
|
||||
|
||||
Every provider is provisioned and asserted; the suite never skips a provider. Per
|
||||
`tests/e2e/CLAUDE.md` there is no sanctioned skip: the whole-suite proxy-liveness
|
||||
`tests/e2e/AGENTS.md` there is no sanctioned skip: the whole-suite proxy-liveness
|
||||
probe hard-fails when no proxy answers, and a provider whose credentials or upstream
|
||||
realtime model are missing on the gateway is likewise a hard failure, not a skip.
|
||||
Give the gateway each provider's credentials to turn its tests green.
|
||||
|
|
|
|||
|
|
@ -29,7 +29,7 @@ def realtime_models(client: RealtimeClient) -> Iterator[dict[str, str]]:
|
|||
provider-id -> model-name map the tests connect with; delete them on teardown.
|
||||
Every provider is provisioned (never skipped): a provider whose credentials or
|
||||
upstream model are missing on the gateway hard-fails its test, per the suite's
|
||||
fail-on-behavior contract in tests/e2e/CLAUDE.md."""
|
||||
fail-on-behavior contract in tests/e2e/AGENTS.md."""
|
||||
records = tuple((provider.id, *client.provision(provider)) for provider in PROVIDERS)
|
||||
try:
|
||||
yield {provider_id: model_name for provider_id, model_name, _ in records}
|
||||
|
|
|
|||
|
|
@ -38,7 +38,7 @@ class RealtimeProvider:
|
|||
the suite registers through /model/new (the gateway resolves the os.environ/*
|
||||
credential refs), so the suite is self-contained and never depends on a static
|
||||
gateway model_list. Every provider here is provisioned and asserted: per
|
||||
tests/e2e/CLAUDE.md the suite never skips a provider, so a provider whose
|
||||
tests/e2e/AGENTS.md the suite never skips a provider, so a provider whose
|
||||
credentials or upstream realtime model are missing on the gateway is a hard
|
||||
failure, not a skip."""
|
||||
|
||||
|
|
@ -98,7 +98,7 @@ PROVIDERS = (
|
|||
def realtime_model(provider: RealtimeProvider, provisioned: Mapping[str, str]) -> str:
|
||||
"""Return the provisioned deployment name for this provider. Every provider in
|
||||
PROVIDERS is provisioned at session start, so a missing entry is a harness bug,
|
||||
never an environment skip - the suite hard-fails instead (see tests/e2e/CLAUDE.md)."""
|
||||
never an environment skip - the suite hard-fails instead (see tests/e2e/AGENTS.md)."""
|
||||
model = provisioned.get(provider.id)
|
||||
assert model is not None, (
|
||||
f"{provider.id} was not provisioned; the realtime_models fixture is broken"
|
||||
|
|
|
|||
0
tests/test_litellm/llms/base_llm/realtime/__init__.py
Normal file
0
tests/test_litellm/llms/base_llm/realtime/__init__.py
Normal file
|
|
@ -0,0 +1,128 @@
|
|||
import base64
|
||||
import json
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.base_llm.realtime.transcription_protocol import (
|
||||
RealtimeTranscriptionProtocolError,
|
||||
completed_event,
|
||||
decode_pcm16_append,
|
||||
parse_transcription_session_update,
|
||||
transcription_session,
|
||||
)
|
||||
|
||||
|
||||
def _session_update(session: dict[str, object]) -> str:
|
||||
return json.dumps({"type": "session.update", "session": session})
|
||||
|
||||
|
||||
def test_ga_layout_parses_format_language_and_turn_detection():
|
||||
update = parse_transcription_session_update(
|
||||
_session_update(
|
||||
{
|
||||
"type": "transcription",
|
||||
"audio": {
|
||||
"input": {
|
||||
"format": {"type": "audio/pcm", "rate": 16_000, "channels": 1},
|
||||
"transcription": {"model": "chirp_3", "language": "pt-BR", "prompt": "names"},
|
||||
"turn_detection": {"type": "server_vad", "threshold": 0.5},
|
||||
}
|
||||
},
|
||||
}
|
||||
)
|
||||
)
|
||||
assert update.session_type == "transcription"
|
||||
assert update.audio_format is not None
|
||||
assert (update.audio_format.layout, update.audio_format.rate, update.audio_format.channels) == ("ga", 16_000, 1)
|
||||
assert update.audio_format.is_pcm16
|
||||
assert (update.model, update.language) == ("chirp_3", "pt-BR")
|
||||
assert update.unsupported_transcription_keys == ("prompt",)
|
||||
assert update.turn_detection_type == "server_vad"
|
||||
assert not update.turn_detection_disabled
|
||||
|
||||
|
||||
def test_beta_layout_parses_flat_fields():
|
||||
update = parse_transcription_session_update(
|
||||
json.dumps(
|
||||
{
|
||||
"type": "transcription_session.update",
|
||||
"session": {
|
||||
"input_audio_format": "pcm16",
|
||||
"input_audio_transcription": {"model": "whisper-1"},
|
||||
"turn_detection": None,
|
||||
},
|
||||
}
|
||||
)
|
||||
)
|
||||
assert update.audio_format is not None
|
||||
assert (update.audio_format.layout, update.audio_format.encoding) == ("beta", "pcm16")
|
||||
assert update.audio_format.is_pcm16
|
||||
assert update.model == "whisper-1"
|
||||
assert update.turn_detection_disabled
|
||||
|
||||
|
||||
def test_absent_turn_detection_is_not_disabled():
|
||||
update = parse_transcription_session_update(_session_update({"audio": {"input": {"transcription": {}}}}))
|
||||
assert update.turn_detection is None
|
||||
assert not update.turn_detection_disabled
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("payload", "message"),
|
||||
[
|
||||
("not json", "invalid JSON object"),
|
||||
("[]", "must be a JSON object"),
|
||||
(json.dumps({"type": "response.create"}), "expected session.update"),
|
||||
(_session_update({}), "requires a session object"),
|
||||
(_session_update({"input_audio_format": "pcm16", "audio": {"input": {"format": "pcm16"}}}), "either beta or GA"),
|
||||
(_session_update({"input_audio_transcription": {}, "audio": {"input": {"transcription": {}}}}), "either beta or GA"),
|
||||
(_session_update({"audio": {"input": {"format": {"rate": "fast"}}}}), "must be an integer"),
|
||||
(_session_update({"audio": {"input": {"format": {"rate": True}}}}), "must be an integer"),
|
||||
(_session_update({"audio": {"input": {"transcription": {"language": 7}}}}), "must be a string"),
|
||||
(_session_update({"audio": {"input": {"transcription": []}}}), "must be an object"),
|
||||
],
|
||||
)
|
||||
def test_malformed_session_updates_are_rejected(payload: str, message: str):
|
||||
with pytest.raises(RealtimeTranscriptionProtocolError, match=message):
|
||||
parse_transcription_session_update(payload)
|
||||
|
||||
|
||||
def test_decode_pcm16_append_returns_the_raw_samples():
|
||||
assert decode_pcm16_append(base64.b64encode(b"\x01\x02\x03\x04").decode()) == b"\x01\x02\x03\x04"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("audio", "message"),
|
||||
[
|
||||
(None, "must be a base64 string"),
|
||||
("@@@", "must be valid base64"),
|
||||
(base64.b64encode(b"\x01\x02\x03").decode(), "complete samples"),
|
||||
],
|
||||
)
|
||||
def test_decode_pcm16_append_rejects_bad_audio(audio: object, message: str):
|
||||
with pytest.raises(RealtimeTranscriptionProtocolError, match=message):
|
||||
decode_pcm16_append(audio)
|
||||
|
||||
|
||||
def test_decode_pcm16_append_enforces_the_backlog_limit():
|
||||
with pytest.raises(RealtimeTranscriptionProtocolError, match="backlog limit"):
|
||||
decode_pcm16_append(base64.b64encode(b"\x00" * 8).decode(), max_encoded_bytes=4)
|
||||
|
||||
|
||||
def test_transcription_session_reflects_negotiated_settings():
|
||||
manual = transcription_session(session_id="sess_1", model="chirp_3", sample_rate=16_000, language=None, server_vad=False)
|
||||
assert manual["id"] == "sess_1"
|
||||
assert manual["audio"]["input"] == {
|
||||
"format": {"type": "audio/pcm", "rate": 16_000},
|
||||
"transcription": {"model": "chirp_3"},
|
||||
"turn_detection": None,
|
||||
}
|
||||
vad = transcription_session(session_id="sess_1", model="chirp_3", sample_rate=24_000, language="en-US", server_vad=True)
|
||||
assert vad["audio"]["input"]["transcription"] == {"model": "chirp_3", "language": "en-US"}
|
||||
assert vad["audio"]["input"]["turn_detection"] == {"type": "server_vad"}
|
||||
|
||||
|
||||
def test_completed_event_carries_usage_only_when_billed():
|
||||
assert "usage" not in completed_event("item_1", "hello", None)
|
||||
billed = completed_event("item_1", "hello", {"type": "duration", "seconds": 2.5})
|
||||
assert (billed["item_id"], billed["transcript"], billed["usage"]) == ("item_1", "hello", {"type": "duration", "seconds": 2.5})
|
||||
|
|
@ -1,4 +1,5 @@
|
|||
import asyncio
|
||||
import base64
|
||||
import json
|
||||
import logging
|
||||
import threading
|
||||
|
|
@ -2064,6 +2065,7 @@ async def _run_async_realtime_with_backend_failure(client_ws):
|
|||
provider_config = Mock()
|
||||
provider_config.get_complete_url.return_value = "wss://backend.example/live"
|
||||
provider_config.validate_environment.return_value = {}
|
||||
provider_config.open_backend = AsyncMock(return_value=None)
|
||||
|
||||
with patch.object(
|
||||
handler,
|
||||
|
|
@ -3707,3 +3709,149 @@ def test_image_edit_handler_keeps_the_sync_transform():
|
|||
assert config.transform_calls == ["sync"]
|
||||
assert captured["body"] == {"transformed_by": "sync"}
|
||||
assert response.data[0].b64_json == "sync"
|
||||
|
||||
|
||||
class _ScriptedClientWebSocket(_FakeClientWebSocket):
|
||||
def __init__(self, messages: list[str], last_event_type: str) -> None:
|
||||
super().__init__()
|
||||
self._messages: Final = list(messages)
|
||||
self._last_event_type: Final = last_event_type
|
||||
self._backend_done: Final = asyncio.Event()
|
||||
|
||||
async def receive_text(self) -> str:
|
||||
if self._messages:
|
||||
return self._messages.pop(0)
|
||||
await asyncio.wait_for(self._backend_done.wait(), timeout=5)
|
||||
raise RuntimeError("client went away")
|
||||
|
||||
async def send_text(self, payload: str) -> None:
|
||||
await super().send_text(payload)
|
||||
if json.loads(payload).get("type") == self._last_event_type:
|
||||
self._backend_done.set()
|
||||
|
||||
def sent_events(self) -> list[dict[str, object]]:
|
||||
return [json.loads(payload) for name, payload in self.events if name == "send_text"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_realtime_bridges_a_transcription_session_through_the_provider_backend():
|
||||
import websockets.exceptions # noqa: F401 # binds the submodule so async_realtime's except clause resolves, as in the proxy process
|
||||
|
||||
from datetime import timedelta
|
||||
|
||||
from google.cloud.speech_v2.types import (
|
||||
RecognitionResponseMetadata,
|
||||
SpeechRecognitionAlternative,
|
||||
StreamingRecognitionResult,
|
||||
StreamingRecognizeResponse,
|
||||
)
|
||||
|
||||
from litellm.llms.vertex_ai.audio_transcription.realtime_backend import SpeechStreamingBackend
|
||||
from litellm.llms.vertex_ai.audio_transcription.realtime_transformation import VertexChirpRealtimeConfig
|
||||
|
||||
def google_response(transcript: str, is_final: bool, billed: float) -> StreamingRecognizeResponse:
|
||||
return StreamingRecognizeResponse(
|
||||
results=[
|
||||
StreamingRecognitionResult(
|
||||
alternatives=[SpeechRecognitionAlternative(transcript=transcript)], is_final=is_final
|
||||
)
|
||||
],
|
||||
metadata=RecognitionResponseMetadata(total_billed_duration=timedelta(seconds=billed)),
|
||||
)
|
||||
|
||||
class FakeTransport:
|
||||
async def close(self) -> None:
|
||||
return None
|
||||
|
||||
class FakeSpeechClient:
|
||||
transport = FakeTransport()
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.requests: Final[list[object]] = []
|
||||
|
||||
async def streaming_recognize(self, requests=None):
|
||||
return self._respond(requests)
|
||||
|
||||
async def _respond(self, requests):
|
||||
script = [google_response("four score", False, 0.0), google_response("Four score and seven", True, 2.0)]
|
||||
async for request in requests:
|
||||
self.requests.append(request)
|
||||
if request.audio and script:
|
||||
yield script.pop(0)
|
||||
|
||||
speech_client = FakeSpeechClient()
|
||||
|
||||
async def resolve_access_token() -> str:
|
||||
return "token"
|
||||
|
||||
provider_config = VertexChirpRealtimeConfig(
|
||||
resolve_access_token=resolve_access_token,
|
||||
project="proj-1",
|
||||
location="us",
|
||||
backend_factory=lambda target: SpeechStreamingBackend(
|
||||
target, client_factory=lambda target, access_token: speech_client
|
||||
),
|
||||
)
|
||||
audio = base64.b64encode(b"\x00\x01" * 800).decode()
|
||||
client_ws = _ScriptedClientWebSocket(
|
||||
[
|
||||
json.dumps(
|
||||
{
|
||||
"type": "session.update",
|
||||
"session": {
|
||||
"type": "transcription",
|
||||
"audio": {
|
||||
"input": {
|
||||
"format": {"type": "audio/pcm", "rate": 16000},
|
||||
"transcription": {"model": "chirp_3", "language": "en"},
|
||||
"turn_detection": {"type": "server_vad"},
|
||||
}
|
||||
},
|
||||
},
|
||||
}
|
||||
),
|
||||
json.dumps({"type": "input_audio_buffer.append", "audio": audio}),
|
||||
json.dumps({"type": "input_audio_buffer.append", "audio": audio}),
|
||||
json.dumps({"type": "input_audio_buffer.commit"}),
|
||||
],
|
||||
last_event_type="conversation.item.input_audio_transcription.completed",
|
||||
)
|
||||
logging_obj = Mock()
|
||||
logging_obj.litellm_trace_id = "trace_1"
|
||||
logging_obj.model_call_details = {}
|
||||
logging_obj.dispatch_success_handlers = AsyncMock()
|
||||
logging_obj.dispatch_failure_handlers = AsyncMock()
|
||||
handler = BaseLLMHTTPHandler()
|
||||
|
||||
with patch.object(handler, "_open_realtime_backend_ws", AsyncMock(side_effect=AssertionError("dialed a websocket"))) as dial:
|
||||
await handler.async_realtime(
|
||||
model="chirp_3",
|
||||
websocket=client_ws,
|
||||
logging_obj=logging_obj,
|
||||
provider_config=provider_config,
|
||||
headers={},
|
||||
query_params={"model": "chirp_3", "intent": "transcription"},
|
||||
)
|
||||
|
||||
dial.assert_not_awaited()
|
||||
events = client_ws.sent_events()
|
||||
assert [event["type"] for event in events] == [
|
||||
"session.created",
|
||||
"session.updated",
|
||||
"input_audio_buffer.speech_started",
|
||||
"conversation.item.input_audio_transcription.delta",
|
||||
"conversation.item.input_audio_transcription.delta",
|
||||
"input_audio_buffer.speech_stopped",
|
||||
"conversation.item.input_audio_transcription.completed",
|
||||
]
|
||||
assert events[0]["session"]["audio"]["input"]["transcription"] == {"model": "chirp_3"}
|
||||
assert events[1]["session"]["audio"]["input"] == {
|
||||
"format": {"type": "audio/pcm", "rate": 16000},
|
||||
"transcription": {"model": "chirp_3", "language": "en-US"},
|
||||
"turn_detection": {"type": "server_vad"},
|
||||
}
|
||||
assert [event["delta"] for event in events[3:5]] == ["four score", " and seven"]
|
||||
assert events[6]["transcript"] == "Four score and seven"
|
||||
assert events[6]["usage"] == {"type": "duration", "seconds": 2.0}
|
||||
assert speech_client.requests[0].streaming_config.config.model == "chirp_3"
|
||||
assert [bytes(request.audio) for request in speech_client.requests[1:]] == [b"\x00\x01" * 800, b"\x00\x01" * 800]
|
||||
|
|
|
|||
|
|
@ -0,0 +1,491 @@
|
|||
import asyncio
|
||||
import json
|
||||
from collections.abc import AsyncIterator, Callable, Sequence
|
||||
from dataclasses import replace
|
||||
from datetime import timedelta
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
from google.cloud.speech_v2.types import (
|
||||
RecognitionResponseMetadata,
|
||||
SpeechRecognitionAlternative,
|
||||
StreamingRecognitionResult,
|
||||
StreamingRecognizeRequest,
|
||||
StreamingRecognizeResponse,
|
||||
)
|
||||
from websockets.exceptions import ConnectionClosedError, ConnectionClosedOK
|
||||
|
||||
from litellm.llms.vertex_ai.audio_transcription.realtime_backend import REQUEST_QUEUE_SIZE, SpeechStreamingBackend
|
||||
from litellm.llms.vertex_ai.audio_transcription.realtime_transformation import SpeechStreamingTarget
|
||||
|
||||
|
||||
async def _static_token() -> str:
|
||||
return "token"
|
||||
|
||||
|
||||
TARGET: Final = SpeechStreamingTarget(
|
||||
api_endpoint="us-speech.googleapis.com",
|
||||
recognizer="projects/proj-1/locations/us/recognizers/_",
|
||||
resolve_access_token=_static_token,
|
||||
)
|
||||
CONFIGURE: Final = json.dumps(
|
||||
{"kind": "configure", "model": "chirp_3", "language_codes": ["en-US"], "sample_rate_hertz": 16_000}
|
||||
)
|
||||
FINISH_TURN: Final = json.dumps({"kind": "finish_turn"})
|
||||
DISCARD_TURN: Final = json.dumps({"kind": "discard_turn"})
|
||||
ScriptItem = StreamingRecognizeResponse | Exception | asyncio.Event
|
||||
|
||||
|
||||
def _response(
|
||||
transcript: str | None,
|
||||
*,
|
||||
is_final: bool = False,
|
||||
billed: float = 0.0,
|
||||
event: str = "SPEECH_EVENT_TYPE_UNSPECIFIED",
|
||||
) -> StreamingRecognizeResponse:
|
||||
results = (
|
||||
[]
|
||||
if transcript is None
|
||||
else [
|
||||
StreamingRecognitionResult(
|
||||
alternatives=[SpeechRecognitionAlternative(transcript=transcript)], is_final=is_final
|
||||
)
|
||||
]
|
||||
)
|
||||
return StreamingRecognizeResponse(
|
||||
results=results,
|
||||
speech_event_type=event,
|
||||
metadata=RecognitionResponseMetadata(total_billed_duration=timedelta(seconds=billed)),
|
||||
)
|
||||
|
||||
|
||||
class _FakeTransport:
|
||||
def __init__(self) -> None:
|
||||
self.closed = False
|
||||
|
||||
async def close(self) -> None:
|
||||
self.closed = True
|
||||
|
||||
|
||||
class _FakeSpeechClient:
|
||||
def __init__(self, *scripts: Sequence[ScriptItem]) -> None:
|
||||
self.transport: Final = _FakeTransport()
|
||||
self.streams: Final[list[list[StreamingRecognizeRequest]]] = []
|
||||
self._scripts: Final = [list(script) for script in scripts]
|
||||
|
||||
async def streaming_recognize(
|
||||
self, requests: AsyncIterator[StreamingRecognizeRequest] | None = None
|
||||
) -> AsyncIterator[StreamingRecognizeResponse]:
|
||||
assert requests is not None
|
||||
script: Final = self._scripts.pop(0) if self._scripts else []
|
||||
received: Final[list[StreamingRecognizeRequest]] = []
|
||||
self.streams.append(received)
|
||||
return self._respond(requests, script, received)
|
||||
|
||||
async def _respond(
|
||||
self,
|
||||
requests: AsyncIterator[StreamingRecognizeRequest],
|
||||
script: list[ScriptItem],
|
||||
received: list[StreamingRecognizeRequest],
|
||||
) -> AsyncIterator[StreamingRecognizeResponse]:
|
||||
async for request in requests:
|
||||
received.append(request)
|
||||
if request.audio and script:
|
||||
yield await self._next(script)
|
||||
while script:
|
||||
yield await self._next(script)
|
||||
|
||||
@staticmethod
|
||||
async def _next(script: list[ScriptItem]) -> StreamingRecognizeResponse:
|
||||
item: Final = script.pop(0)
|
||||
if isinstance(item, asyncio.Event):
|
||||
await item.wait()
|
||||
return await _FakeSpeechClient._next(script)
|
||||
if isinstance(item, Exception):
|
||||
raise item
|
||||
return item
|
||||
|
||||
|
||||
def _backend(client: _FakeSpeechClient, **kwargs: object) -> SpeechStreamingBackend:
|
||||
return SpeechStreamingBackend(TARGET, client_factory=lambda target, access_token: client, **kwargs)
|
||||
|
||||
|
||||
async def _recv(backend: SpeechStreamingBackend) -> dict[str, object]:
|
||||
message: Final = await asyncio.wait_for(backend.recv(), timeout=2)
|
||||
assert isinstance(message, str)
|
||||
return json.loads(message)
|
||||
|
||||
|
||||
async def _transcript(backend: SpeechStreamingBackend) -> str:
|
||||
event: Final = await _recv(backend)
|
||||
assert event["kind"] == "response", event
|
||||
(result,) = event["results"]
|
||||
return result["transcript"]
|
||||
|
||||
|
||||
async def _configure(backend: SpeechStreamingBackend) -> None:
|
||||
await backend.send(CONFIGURE)
|
||||
assert await _recv(backend) == {"kind": "configured"}
|
||||
|
||||
|
||||
async def _until(condition: Callable[[], bool]) -> None:
|
||||
async def poll() -> None:
|
||||
while not condition():
|
||||
await asyncio.sleep(0)
|
||||
|
||||
await asyncio.wait_for(poll(), timeout=2)
|
||||
|
||||
|
||||
def _audio(stream: list[StreamingRecognizeRequest]) -> list[bytes]:
|
||||
return [bytes(request.audio) for request in stream[1:]]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_audio_streams_through_one_recognize_call_with_the_config_first():
|
||||
client = _FakeSpeechClient([_response("hello"), _response("hello world", is_final=True, billed=2.0)])
|
||||
async with _backend(client) as backend:
|
||||
await _configure(backend)
|
||||
await backend.send(b"\x01\x02")
|
||||
await backend.send(b"\x03\x04")
|
||||
await backend.send(FINISH_TURN)
|
||||
first, second, finished = [await _recv(backend) for _ in range(3)]
|
||||
assert first == {
|
||||
"kind": "response",
|
||||
"speech_event": "none",
|
||||
"results": [{"transcript": "hello", "is_final": False}],
|
||||
"billed_seconds": 0.0,
|
||||
}
|
||||
assert second["results"] == [{"transcript": "hello world", "is_final": True}]
|
||||
assert second["billed_seconds"] == 2.0
|
||||
assert finished == {"kind": "turn_finished"}
|
||||
(requests,) = client.streams
|
||||
assert requests[0].recognizer == TARGET.recognizer
|
||||
config = requests[0].streaming_config
|
||||
assert config.config.model == "chirp_3"
|
||||
assert list(config.config.language_codes) == ["en-US"]
|
||||
assert config.config.explicit_decoding_config.sample_rate_hertz == 16_000
|
||||
assert config.config.explicit_decoding_config.audio_channel_count == 1
|
||||
assert config.config.explicit_decoding_config.encoding.name == "LINEAR16"
|
||||
assert config.streaming_features.interim_results
|
||||
assert config.streaming_features.enable_voice_activity_events
|
||||
assert _audio(requests) == [b"\x01\x02", b"\x03\x04"]
|
||||
assert client.transport.closed
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_voice_activity_events_are_relayed():
|
||||
client = _FakeSpeechClient(
|
||||
[_response(None, event="SPEECH_ACTIVITY_BEGIN"), _response(None, event="SPEECH_ACTIVITY_END")]
|
||||
)
|
||||
async with _backend(client) as backend:
|
||||
await _configure(backend)
|
||||
await backend.send(b"\x00\x00")
|
||||
await backend.send(b"\x00\x00")
|
||||
begin, end = [await _recv(backend) for _ in range(2)]
|
||||
assert (begin["speech_event"], begin["results"]) == ("begin", [])
|
||||
assert end["speech_event"] == "end"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_audio_before_configure_is_rejected():
|
||||
backend = _backend(_FakeSpeechClient())
|
||||
with pytest.raises(RuntimeError, match="before the Speech-to-Text stream was configured"):
|
||||
await backend.send(b"\x00\x00")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stream_failure_closes_the_session_with_1011_and_the_reason():
|
||||
client = _FakeSpeechClient([PermissionError("IAM_PERMISSION_DENIED: speech.recognizers.recognize")])
|
||||
async with _backend(client) as backend:
|
||||
await _configure(backend)
|
||||
await backend.send(b"\x00\x00")
|
||||
with pytest.raises(ConnectionClosedError) as excinfo:
|
||||
await backend.recv()
|
||||
assert excinfo.value.rcvd is not None
|
||||
assert excinfo.value.rcvd.code == 1011
|
||||
assert "IAM_PERMISSION_DENIED" in excinfo.value.rcvd.reason
|
||||
assert client.transport.closed
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_close_reports_a_normal_closure_to_both_directions():
|
||||
client = _FakeSpeechClient([_response("hi")])
|
||||
backend = _backend(client)
|
||||
await _configure(backend)
|
||||
await backend.send(b"\x00\x00")
|
||||
assert await _transcript(backend) == "hi"
|
||||
await backend.close()
|
||||
with pytest.raises(ConnectionClosedOK):
|
||||
await backend.recv()
|
||||
with pytest.raises(ConnectionClosedOK):
|
||||
await backend.send(b"\x00\x00")
|
||||
assert client.transport.closed
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_turn_commands_without_audio_answer_immediately():
|
||||
backend = _backend(_FakeSpeechClient())
|
||||
await _configure(backend)
|
||||
await backend.send(FINISH_TURN)
|
||||
assert await _recv(backend) == {"kind": "turn_finished"}
|
||||
await backend.send(DISCARD_TURN)
|
||||
assert await _recv(backend) == {"kind": "turn_discarded", "billed_seconds": 0.0}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discard_turn_cancels_the_open_stream_and_the_next_turn_starts_fresh():
|
||||
client = _FakeSpeechClient([_response("draft")], [_response("again", is_final=True)])
|
||||
async with _backend(client) as backend:
|
||||
await _configure(backend)
|
||||
await backend.send(b"\x01\x01")
|
||||
assert await _transcript(backend) == "draft"
|
||||
await backend.send(DISCARD_TURN)
|
||||
assert await _recv(backend) == {"kind": "turn_discarded", "billed_seconds": 0.0}
|
||||
await backend.send(b"\x02\x02")
|
||||
assert await _transcript(backend) == "again"
|
||||
assert [_audio(stream) for stream in client.streams] == [[b"\x01\x01"], [b"\x02\x02"]]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discard_turn_drops_its_queued_results_and_keeps_google_billed_seconds():
|
||||
client = _FakeSpeechClient(
|
||||
[_response("draft"), _response("leftover", is_final=True, billed=2.0)],
|
||||
[_response("fresh", is_final=True, billed=1.0)],
|
||||
)
|
||||
async with _backend(client) as backend:
|
||||
await _configure(backend)
|
||||
await backend.send(b"\x01\x01")
|
||||
assert await _transcript(backend) == "draft"
|
||||
await backend.send(b"\x02\x02")
|
||||
await _until(lambda: len(client.streams[0]) == 3)
|
||||
await backend.send(DISCARD_TURN)
|
||||
assert await _recv(backend) == {"kind": "turn_discarded", "billed_seconds": 2.0}
|
||||
assert backend._discarded_turns == frozenset()
|
||||
await backend.send(b"\x03\x03")
|
||||
fresh = await _recv(backend)
|
||||
assert fresh["results"] == [{"transcript": "fresh", "is_final": True}]
|
||||
assert fresh["billed_seconds"] == 3.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discard_turn_keeps_the_queued_results_of_the_turn_finished_before_it():
|
||||
client = _FakeSpeechClient([_response("one", is_final=True, billed=2.0)], [_response("two")])
|
||||
async with _backend(client) as backend:
|
||||
await _configure(backend)
|
||||
await backend.send(b"\x01\x01")
|
||||
await backend.send(FINISH_TURN)
|
||||
await backend.send(b"\x02\x02")
|
||||
await _until(lambda: len(client.streams) == 2 and len(client.streams[1]) == 2)
|
||||
await backend.send(DISCARD_TURN)
|
||||
assert await _transcript(backend) == "one"
|
||||
assert await _recv(backend) == {"kind": "turn_finished"}
|
||||
assert await _recv(backend) == {"kind": "turn_discarded", "billed_seconds": 2.0}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_billed_seconds_accumulate_across_turns():
|
||||
client = _FakeSpeechClient(
|
||||
[_response("one", is_final=True, billed=2.0)], [_response("two", is_final=True, billed=3.0)]
|
||||
)
|
||||
async with _backend(client) as backend:
|
||||
await _configure(backend)
|
||||
await backend.send(b"\x00\x00")
|
||||
await backend.send(FINISH_TURN)
|
||||
first = await _recv(backend)
|
||||
assert await _recv(backend) == {"kind": "turn_finished"}
|
||||
await backend.send(b"\x00\x00")
|
||||
await backend.send(FINISH_TURN)
|
||||
second = await _recv(backend)
|
||||
assert await _recv(backend) == {"kind": "turn_finished"}
|
||||
assert (first["billed_seconds"], second["billed_seconds"]) == (2.0, 5.0)
|
||||
assert len(client.streams) == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_streams_rotate_before_the_five_minute_limit_without_ending_the_turn():
|
||||
now = [0.0]
|
||||
client = _FakeSpeechClient(
|
||||
[_response("first"), _response("first half", is_final=True, billed=239.0)],
|
||||
[_response("second", billed=1.0)],
|
||||
)
|
||||
async with _backend(client, clock=lambda: now[0], rotation_seconds=240.0) as backend:
|
||||
await _configure(backend)
|
||||
await backend.send(b"\x01\x01")
|
||||
assert await _transcript(backend) == "first"
|
||||
now[0] = 239.0
|
||||
await backend.send(b"\x02\x02")
|
||||
assert await _transcript(backend) == "first half"
|
||||
now[0] = 240.0
|
||||
await backend.send(b"\x03\x03")
|
||||
second = await _recv(backend)
|
||||
assert second["results"] == [{"transcript": "second", "is_final": False}]
|
||||
assert second["billed_seconds"] == 240.0
|
||||
await backend.send(FINISH_TURN)
|
||||
assert await _recv(backend) == {"kind": "turn_finished"}
|
||||
assert [_audio(stream) for stream in client.streams] == [[b"\x01\x01", b"\x02\x02"], [b"\x03\x03"]]
|
||||
assert client.streams[1][0].streaming_config.config.model == "chirp_3"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_turn_finished_follows_results_that_arrive_after_a_rotation():
|
||||
now = [0.0]
|
||||
client = _FakeSpeechClient(
|
||||
[_response("one"), _response("one two", is_final=True)],
|
||||
[_response("three")],
|
||||
)
|
||||
async with _backend(client, clock=lambda: now[0], rotation_seconds=240.0) as backend:
|
||||
await _configure(backend)
|
||||
await backend.send(b"\x01\x01")
|
||||
assert await _transcript(backend) == "one"
|
||||
now[0] = 240.0
|
||||
await backend.send(b"\x02\x02")
|
||||
await backend.send(FINISH_TURN)
|
||||
assert await _transcript(backend) == "one two"
|
||||
assert await _transcript(backend) == "three"
|
||||
assert await _recv(backend) == {"kind": "turn_finished"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rotation_waits_for_a_pause_in_speech():
|
||||
now = [0.0]
|
||||
client = _FakeSpeechClient(
|
||||
[
|
||||
_response(None, event="SPEECH_ACTIVITY_BEGIN"),
|
||||
_response("still talking"),
|
||||
_response("still talking", is_final=True, event="SPEECH_ACTIVITY_END"),
|
||||
],
|
||||
[_response("next")],
|
||||
)
|
||||
async with _backend(
|
||||
client, clock=lambda: now[0], rotation_seconds=240.0, rotation_deadline_seconds=280.0
|
||||
) as backend:
|
||||
await _configure(backend)
|
||||
await backend.send(b"\x01\x01")
|
||||
assert (await _recv(backend))["speech_event"] == "begin"
|
||||
now[0] = 250.0
|
||||
await backend.send(b"\x02\x02")
|
||||
assert await _transcript(backend) == "still talking"
|
||||
now[0] = 260.0
|
||||
await backend.send(b"\x03\x03")
|
||||
assert (await _recv(backend))["speech_event"] == "end"
|
||||
now[0] = 261.0
|
||||
await backend.send(b"\x04\x04")
|
||||
assert await _transcript(backend) == "next"
|
||||
assert [_audio(stream) for stream in client.streams] == [[b"\x01\x01", b"\x02\x02", b"\x03\x03"], [b"\x04\x04"]]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rotation_is_forced_at_the_deadline_during_continuous_speech():
|
||||
now = [0.0]
|
||||
client = _FakeSpeechClient(
|
||||
[_response(None, event="SPEECH_ACTIVITY_BEGIN"), _response("still talking")],
|
||||
[_response("cut off")],
|
||||
)
|
||||
async with _backend(
|
||||
client, clock=lambda: now[0], rotation_seconds=240.0, rotation_deadline_seconds=280.0
|
||||
) as backend:
|
||||
await _configure(backend)
|
||||
await backend.send(b"\x01\x01")
|
||||
assert (await _recv(backend))["speech_event"] == "begin"
|
||||
now[0] = 279.0
|
||||
await backend.send(b"\x02\x02")
|
||||
assert await _transcript(backend) == "still talking"
|
||||
now[0] = 280.0
|
||||
await backend.send(b"\x03\x03")
|
||||
assert await _transcript(backend) == "cut off"
|
||||
assert [_audio(stream) for stream in client.streams] == [[b"\x01\x01", b"\x02\x02"], [b"\x03\x03"]]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_every_stream_opens_its_own_client_with_a_freshly_resolved_token():
|
||||
now = [0.0]
|
||||
tokens = iter(("token-1", "token-2"))
|
||||
seen_tokens: list[str] = []
|
||||
clients = [_FakeSpeechClient([_response("first")]), _FakeSpeechClient([_response("second")])]
|
||||
unopened = iter(clients)
|
||||
|
||||
async def resolve_access_token() -> str:
|
||||
return next(tokens)
|
||||
|
||||
def open_client(target: SpeechStreamingTarget, access_token: str) -> _FakeSpeechClient:
|
||||
seen_tokens.append(access_token)
|
||||
return next(unopened)
|
||||
|
||||
backend = SpeechStreamingBackend(
|
||||
replace(TARGET, resolve_access_token=resolve_access_token),
|
||||
client_factory=open_client,
|
||||
clock=lambda: now[0],
|
||||
rotation_seconds=240.0,
|
||||
)
|
||||
async with backend:
|
||||
await _configure(backend)
|
||||
await backend.send(b"\x01\x01")
|
||||
assert await _transcript(backend) == "first"
|
||||
now[0] = 240.0
|
||||
await backend.send(b"\x02\x02")
|
||||
assert await _transcript(backend) == "second"
|
||||
assert clients[0].transport.closed
|
||||
assert not clients[1].transport.closed
|
||||
assert seen_tokens == ["token-1", "token-2"]
|
||||
assert [len(client.streams) for client in clients] == [1, 1]
|
||||
assert clients[1].transport.closed
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_close_releases_a_rotated_stream_that_never_started_relaying():
|
||||
now = [0.0]
|
||||
hold = asyncio.Event()
|
||||
clients = [_FakeSpeechClient([_response("first"), hold]), _FakeSpeechClient([_response("never")])]
|
||||
unopened = iter(clients)
|
||||
backend = SpeechStreamingBackend(
|
||||
TARGET,
|
||||
client_factory=lambda target, access_token: next(unopened),
|
||||
clock=lambda: now[0],
|
||||
rotation_seconds=240.0,
|
||||
)
|
||||
await _configure(backend)
|
||||
await backend.send(b"\x01\x01")
|
||||
assert await _transcript(backend) == "first"
|
||||
now[0] = 240.0
|
||||
await backend.send(b"\x02\x02")
|
||||
await asyncio.sleep(0)
|
||||
assert clients[1].streams == []
|
||||
await backend.close()
|
||||
assert [client.transport.closed for client in clients] == [True, True]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_discard_turn_cancels_every_stream_of_the_turn():
|
||||
now = [0.0]
|
||||
hold = asyncio.Event()
|
||||
client = _FakeSpeechClient(
|
||||
[_response("draft"), hold, _response("never delivered")],
|
||||
[_response("fresh", is_final=True)],
|
||||
)
|
||||
async with _backend(client, clock=lambda: now[0], rotation_seconds=240.0) as backend:
|
||||
await _configure(backend)
|
||||
await backend.send(b"\x01\x01")
|
||||
assert await _transcript(backend) == "draft"
|
||||
now[0] = 240.0
|
||||
await backend.send(b"\x02\x02")
|
||||
await backend.send(DISCARD_TURN)
|
||||
assert await _recv(backend) == {"kind": "turn_discarded", "billed_seconds": 0.0}
|
||||
await backend.send(b"\x03\x03")
|
||||
assert await _transcript(backend) == "fresh"
|
||||
assert [_audio(stream) for stream in client.streams] == [[b"\x01\x01"], [b"\x03\x03"]]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_audio_sends_block_once_the_request_queue_is_full():
|
||||
hold = asyncio.Event()
|
||||
client = _FakeSpeechClient([hold, _response("late", is_final=True)])
|
||||
async with _backend(client) as backend:
|
||||
await _configure(backend)
|
||||
for _ in range(REQUEST_QUEUE_SIZE + 1):
|
||||
await backend.send(b"\x00\x00")
|
||||
blocked = asyncio.create_task(backend.send(b"\x00\x00"))
|
||||
await asyncio.sleep(0)
|
||||
assert not blocked.done()
|
||||
hold.set()
|
||||
await asyncio.wait_for(blocked, timeout=2)
|
||||
assert await _transcript(backend) == "late"
|
||||
|
|
@ -0,0 +1,424 @@
|
|||
import base64
|
||||
import json
|
||||
from typing import Final
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.llms.base_llm.realtime.transformation import RealtimeBackend
|
||||
from litellm.llms.vertex_ai.audio_transcription.realtime_transformation import (
|
||||
MAX_AUDIO_MESSAGE_BYTES,
|
||||
ChirpProtocolError,
|
||||
ChirpSessionConfig,
|
||||
SpeechStreamingTarget,
|
||||
VertexChirpRealtimeConfig,
|
||||
is_vertex_speech_to_text_model,
|
||||
new_words,
|
||||
parse_chirp_session_update,
|
||||
)
|
||||
from litellm.llms.vertex_ai.common_utils import VertexAIError
|
||||
from litellm.types.llms.vertex_ai_speech_to_text import (
|
||||
VertexSpeechStreamingConfigured,
|
||||
VertexSpeechStreamingResponse,
|
||||
VertexSpeechStreamingResult,
|
||||
VertexSpeechStreamingTurnDiscarded,
|
||||
VertexSpeechStreamingTurnFinished,
|
||||
)
|
||||
from litellm.types.realtime import RealtimeResponseTransformInput
|
||||
|
||||
MODEL: Final = "chirp_3"
|
||||
EMPTY_TRANSFORM_INPUT: Final[RealtimeResponseTransformInput] = {
|
||||
"session_configuration_request": None,
|
||||
"current_output_item_id": None,
|
||||
"current_response_id": None,
|
||||
"current_delta_chunks": None,
|
||||
"current_item_chunks": None,
|
||||
"current_conversation_id": None,
|
||||
"current_delta_type": None,
|
||||
}
|
||||
DELTA: Final = "conversation.item.input_audio_transcription.delta"
|
||||
COMPLETED: Final = "conversation.item.input_audio_transcription.completed"
|
||||
|
||||
|
||||
def _event(event_type: str, **fields: object) -> str:
|
||||
return json.dumps({"type": event_type, **fields})
|
||||
|
||||
|
||||
def _ga_session_update(
|
||||
rate: int = 24_000, turn_detection: str | None = "server_vad", language: str | None = "en", model: str = MODEL
|
||||
) -> str:
|
||||
transcription = {"model": model} if language is None else {"model": model, "language": language}
|
||||
return _event(
|
||||
"session.update",
|
||||
session={
|
||||
"type": "transcription",
|
||||
"audio": {
|
||||
"input": {
|
||||
"format": {"type": "audio/pcm", "rate": rate},
|
||||
"turn_detection": None if turn_detection is None else {"type": turn_detection},
|
||||
"transcription": transcription,
|
||||
}
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
async def _token() -> str:
|
||||
return "token"
|
||||
|
||||
|
||||
def _config(location: str | None = "us") -> VertexChirpRealtimeConfig:
|
||||
return VertexChirpRealtimeConfig(resolve_access_token=_token, project="proj-1", location=location)
|
||||
|
||||
|
||||
def _configured(
|
||||
rate: int = 24_000, turn_detection: str | None = "server_vad", language: str | None = "en"
|
||||
) -> VertexChirpRealtimeConfig:
|
||||
config = _config()
|
||||
config.transform_session_created_event(MODEL, "sess_1")
|
||||
config.transform_realtime_request(_ga_session_update(rate, turn_detection, language), MODEL)
|
||||
return config
|
||||
|
||||
|
||||
def _backend_events(config: VertexChirpRealtimeConfig, frame: object) -> list[dict[str, object]]:
|
||||
assert hasattr(frame, "model_dump_json")
|
||||
response = config.transform_realtime_response(frame.model_dump_json(), MODEL, MagicMock(), EMPTY_TRANSFORM_INPUT)[
|
||||
"response"
|
||||
]
|
||||
assert isinstance(response, list)
|
||||
return response
|
||||
|
||||
|
||||
def _response(
|
||||
*results: tuple[str, bool], speech_event: str = "none", billed_seconds: float = 0.0
|
||||
) -> VertexSpeechStreamingResponse:
|
||||
return VertexSpeechStreamingResponse(
|
||||
speech_event=speech_event,
|
||||
results=tuple(VertexSpeechStreamingResult(transcript=text, is_final=final) for text, final in results),
|
||||
billed_seconds=billed_seconds,
|
||||
)
|
||||
|
||||
|
||||
def _types(events: list[dict[str, object]]) -> list[object]:
|
||||
return [event["type"] for event in events]
|
||||
|
||||
|
||||
def _commands(config: VertexChirpRealtimeConfig, payload: str) -> list[object]:
|
||||
return [
|
||||
json.loads(command) if isinstance(command, str) else command
|
||||
for command in config.transform_realtime_request(payload, MODEL)
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("model", "expected"),
|
||||
[
|
||||
("vertex_ai/chirp_3", True),
|
||||
("chirp_3", True),
|
||||
("chirp_2", False),
|
||||
("gemini-live-2.5-flash", False),
|
||||
("vertex_ai/gemini-2.0-flash-live-preview-04-09", False),
|
||||
("vertex_ai/gemini-3.5-transcribe-live-preview", False),
|
||||
("gemini-3.5-transcribe-preview", False),
|
||||
],
|
||||
)
|
||||
def test_is_vertex_speech_to_text_model(model: str, expected: bool):
|
||||
assert is_vertex_speech_to_text_model(model) is expected
|
||||
|
||||
|
||||
def test_ga_session_update_maps_to_a_speech_config():
|
||||
config = parse_chirp_session_update(_ga_session_update(16_000, "server_vad", "pt"), "vertex_ai/chirp_3")
|
||||
assert config == ChirpSessionConfig(model=MODEL, language="pt-BR", sample_rate=16_000, server_vad=True)
|
||||
assert json.loads(config.configure_command()) == {
|
||||
"kind": "configure",
|
||||
"model": MODEL,
|
||||
"language_codes": ["pt-BR"],
|
||||
"sample_rate_hertz": 16_000,
|
||||
}
|
||||
|
||||
|
||||
def test_beta_session_update_defaults_the_rate_and_auto_detects_the_language():
|
||||
config = parse_chirp_session_update(
|
||||
_event(
|
||||
"transcription_session.update",
|
||||
session={
|
||||
"input_audio_format": "pcm16",
|
||||
"input_audio_transcription": {"model": MODEL},
|
||||
"turn_detection": None,
|
||||
},
|
||||
),
|
||||
MODEL,
|
||||
)
|
||||
assert config == ChirpSessionConfig(model=MODEL, language=None, sample_rate=24_000, server_vad=False)
|
||||
assert json.loads(config.configure_command())["language_codes"] == ["auto"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("payload", "message"),
|
||||
[
|
||||
(_event("session.update", session={"type": "realtime_voice"}), "transcription sessions only"),
|
||||
(_ga_session_update(model="gemini-live-2.5-flash"), "cannot be changed"),
|
||||
(_event("session.update", session={"audio": {"input": {"format": {"type": "audio/pcmu"}}}}), "pcm16"),
|
||||
(
|
||||
_event("session.update", session={"audio": {"input": {"format": {"type": "audio/pcm", "channels": 2}}}}),
|
||||
"mono",
|
||||
),
|
||||
(_ga_session_update(rate=4_000), "sample rates"),
|
||||
(_ga_session_update(rate=96_000), "sample rates"),
|
||||
(_ga_session_update(turn_detection="semantic_vad"), "server_vad"),
|
||||
],
|
||||
)
|
||||
def test_unsupported_session_settings_are_rejected(payload: str, message: str):
|
||||
with pytest.raises(ChirpProtocolError, match=message):
|
||||
parse_chirp_session_update(payload, MODEL)
|
||||
|
||||
|
||||
def test_session_update_configures_once_and_later_updates_are_ignored():
|
||||
config = _config()
|
||||
config.transform_session_created_event(MODEL, "sess_1")
|
||||
first = _commands(config, _ga_session_update(16_000))
|
||||
assert first == [{"kind": "configure", "model": MODEL, "language_codes": ["en-US"], "sample_rate_hertz": 16_000}]
|
||||
assert config.is_setup_message(first[0])
|
||||
assert _commands(config, _ga_session_update(8_000)) == []
|
||||
|
||||
|
||||
def test_audio_and_commits_before_session_update_are_rejected():
|
||||
config = _config()
|
||||
with pytest.raises(ChirpProtocolError, match=r"session\.update must configure"):
|
||||
config.transform_realtime_request(
|
||||
_event("input_audio_buffer.append", audio=base64.b64encode(b"\x00\x00").decode()), MODEL
|
||||
)
|
||||
with pytest.raises(ChirpProtocolError, match=r"session\.update must configure"):
|
||||
config.transform_realtime_request(_event("input_audio_buffer.commit"), MODEL)
|
||||
|
||||
|
||||
def test_append_is_split_into_google_sized_chunks():
|
||||
config = _configured()
|
||||
audio = bytes(range(256)) * 250
|
||||
chunks = config.transform_realtime_request(
|
||||
_event("input_audio_buffer.append", audio=base64.b64encode(audio).decode()), MODEL
|
||||
)
|
||||
assert [len(chunk) for chunk in chunks] == [
|
||||
MAX_AUDIO_MESSAGE_BYTES,
|
||||
MAX_AUDIO_MESSAGE_BYTES,
|
||||
64_000 - 2 * MAX_AUDIO_MESSAGE_BYTES,
|
||||
]
|
||||
assert b"".join(chunk for chunk in chunks if isinstance(chunk, bytes)) == audio
|
||||
|
||||
|
||||
def test_commit_end_and_clear_map_to_turn_commands():
|
||||
config = _configured()
|
||||
assert _commands(config, _event("input_audio_buffer.commit")) == [{"kind": "finish_turn"}]
|
||||
assert _commands(config, _event("input_audio_buffer.end")) == [{"kind": "finish_turn"}]
|
||||
assert _commands(config, _event("input_audio_buffer.clear")) == [{"kind": "discard_turn"}]
|
||||
|
||||
|
||||
def test_unsupported_client_events_are_dropped():
|
||||
assert _commands(_configured(), _event("response.create")) == []
|
||||
|
||||
|
||||
def test_connect_announces_a_session_with_chirp_defaults():
|
||||
event = _config().transform_session_created_event(MODEL, "sess_1")
|
||||
assert event["type"] == "session.created"
|
||||
assert event["session"]["id"] == "sess_1"
|
||||
assert event["session"]["audio"]["input"] == {
|
||||
"format": {"type": "audio/pcm", "rate": 24_000},
|
||||
"transcription": {"model": MODEL},
|
||||
"turn_detection": {"type": "server_vad"},
|
||||
}
|
||||
|
||||
|
||||
def test_configured_backend_reports_the_negotiated_session():
|
||||
config = _configured(rate=16_000, turn_detection=None, language="pt-BR")
|
||||
events = _backend_events(config, VertexSpeechStreamingConfigured())
|
||||
assert _types(events) == ["session.created"]
|
||||
session = events[0]["session"]
|
||||
assert isinstance(session, dict)
|
||||
assert session["id"] == "sess_1"
|
||||
assert session["audio"]["input"] == {
|
||||
"format": {"type": "audio/pcm", "rate": 16_000},
|
||||
"transcription": {"model": MODEL, "language": "pt-BR"},
|
||||
"turn_detection": None,
|
||||
}
|
||||
|
||||
|
||||
def test_backend_frames_before_session_update_are_an_error():
|
||||
config = _config()
|
||||
config.transform_session_created_event(MODEL, "sess_1")
|
||||
with pytest.raises(ChirpProtocolError, match=r"session\.update must configure"):
|
||||
_backend_events(config, VertexSpeechStreamingConfigured())
|
||||
|
||||
|
||||
def test_server_vad_turn_streams_new_words_then_completes_with_usage():
|
||||
config = _configured()
|
||||
assert _types(_backend_events(config, _response(speech_event="begin"))) == ["input_audio_buffer.speech_started"]
|
||||
first = _backend_events(config, _response(("four score", False)))
|
||||
assert [(event["type"], event["delta"]) for event in first] == [(DELTA, "four score")]
|
||||
second = _backend_events(config, _response(("four score and seven", False)))
|
||||
assert [event["delta"] for event in second] == [" and seven"]
|
||||
final = _backend_events(config, _response(("Four score and seven years ago.", True), billed_seconds=3.5))
|
||||
assert _types(final) == [DELTA, "input_audio_buffer.speech_stopped", COMPLETED]
|
||||
assert final[0]["delta"] == " years ago."
|
||||
assert final[2]["transcript"] == "Four score and seven years ago."
|
||||
assert final[2]["usage"] == {"type": "duration", "seconds": 3.5}
|
||||
assert {event["item_id"] for event in (*first, *second, *final)} == {first[0]["item_id"]}
|
||||
assert _backend_events(config, _response(speech_event="end")) == []
|
||||
|
||||
|
||||
def test_server_vad_final_result_completes_before_the_interim_that_follows_it():
|
||||
config = _configured()
|
||||
_backend_events(config, _response(speech_event="begin"))
|
||||
events = _backend_events(config, _response(("four score", True), ("and seven", False)))
|
||||
assert _types(events) == [
|
||||
DELTA,
|
||||
"input_audio_buffer.speech_stopped",
|
||||
COMPLETED,
|
||||
"input_audio_buffer.speech_started",
|
||||
DELTA,
|
||||
]
|
||||
assert events[2]["transcript"] == "four score"
|
||||
assert events[4]["delta"] == "and seven"
|
||||
assert events[4]["item_id"] != events[2]["item_id"]
|
||||
assert events[4]["item_id"] == events[3]["item_id"]
|
||||
finished = _backend_events(config, _response(("and seven years", True)))
|
||||
assert [(event["type"], event.get("delta", event.get("transcript"))) for event in finished] == [
|
||||
(DELTA, " years"),
|
||||
("input_audio_buffer.speech_stopped", None),
|
||||
(COMPLETED, "and seven years"),
|
||||
]
|
||||
assert {event["item_id"] for event in finished} == {events[4]["item_id"]}
|
||||
|
||||
|
||||
def test_manual_turn_keeps_the_interim_that_follows_a_final_in_the_same_frame():
|
||||
config = _configured(turn_detection=None)
|
||||
first = _backend_events(config, _response(("four score", True), ("and seven", False)))
|
||||
assert [(event["type"], event["delta"]) for event in first] == [(DELTA, "four score"), (DELTA, " and seven")]
|
||||
second = _backend_events(config, _response(("and seven years", True)))
|
||||
assert [event["delta"] for event in second] == [" years"]
|
||||
completed = _backend_events(config, VertexSpeechStreamingTurnFinished())
|
||||
assert [(event["type"], event["transcript"]) for event in completed] == [(COMPLETED, "four score and seven years")]
|
||||
assert {event["item_id"] for event in (*first, *second, *completed)} == {first[0]["item_id"]}
|
||||
|
||||
|
||||
def test_manual_turns_complete_on_commit_without_speech_events():
|
||||
config = _configured(turn_detection=None)
|
||||
assert _backend_events(config, _response(speech_event="begin")) == []
|
||||
first = _backend_events(config, _response(("hello there", True), billed_seconds=1.25))
|
||||
assert [(event["type"], event["delta"]) for event in first] == [(DELTA, "hello there")]
|
||||
second = _backend_events(config, _response(("world", True)))
|
||||
assert [event["delta"] for event in second] == [" world"]
|
||||
completed = _backend_events(config, VertexSpeechStreamingTurnFinished())
|
||||
assert _types(completed) == [COMPLETED]
|
||||
assert completed[0]["transcript"] == "hello there world"
|
||||
assert completed[0]["usage"] == {"type": "duration", "seconds": 1.25}
|
||||
assert _backend_events(config, VertexSpeechStreamingTurnFinished()) == []
|
||||
|
||||
|
||||
def test_clear_discards_the_open_turn():
|
||||
config = _configured(turn_detection=None)
|
||||
draft = _backend_events(config, _response(("draft", False)))
|
||||
assert _backend_events(config, VertexSpeechStreamingTurnDiscarded(billed_seconds=0.0)) == []
|
||||
assert _backend_events(config, VertexSpeechStreamingTurnFinished()) == []
|
||||
fresh = _backend_events(config, _response(("again", False)))
|
||||
assert fresh[0]["delta"] == "again"
|
||||
assert fresh[0]["item_id"] != draft[0]["item_id"]
|
||||
|
||||
|
||||
def test_cleared_audio_keeps_google_billed_seconds_for_the_close_flush():
|
||||
config = _configured(turn_detection=None)
|
||||
assert _backend_events(config, _response(("draft", False), billed_seconds=1.0)) != []
|
||||
assert _backend_events(config, VertexSpeechStreamingTurnDiscarded(billed_seconds=2.5)) == []
|
||||
assert config.unbilled_usage_on_session_close(MODEL) == {"type": "duration", "seconds": 2.5}
|
||||
|
||||
|
||||
def test_usage_is_billed_once_across_turns_and_flushed_on_close():
|
||||
config = _configured()
|
||||
first = _backend_events(config, _response(("one", True), billed_seconds=2.0))
|
||||
second = _backend_events(config, _response(("two", True), billed_seconds=5.0))
|
||||
assert first[-1]["usage"] == {"type": "duration", "seconds": 2.0}
|
||||
assert second[-1]["usage"] == {"type": "duration", "seconds": 3.0}
|
||||
assert config.unbilled_usage_on_session_close(MODEL) is None
|
||||
assert _backend_events(config, _response(billed_seconds=6.5)) == []
|
||||
assert config.unbilled_usage_on_session_close(MODEL) == {"type": "duration", "seconds": 1.5}
|
||||
assert config.unbilled_usage_on_session_close(MODEL) is None
|
||||
|
||||
|
||||
class _NullBackend:
|
||||
async def __aenter__(self) -> "_NullBackend":
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type: object, exc_value: object, traceback: object) -> None:
|
||||
return None
|
||||
|
||||
async def send(self, message: str | bytes) -> None:
|
||||
return None
|
||||
|
||||
async def recv(self, decode: bool | None = None) -> str | bytes:
|
||||
return ""
|
||||
|
||||
async def close(self) -> None:
|
||||
return None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_open_backend_targets_the_regional_speech_endpoint():
|
||||
targets: list[SpeechStreamingTarget] = []
|
||||
|
||||
def factory(target: SpeechStreamingTarget) -> RealtimeBackend:
|
||||
targets.append(target)
|
||||
return _NullBackend()
|
||||
|
||||
config = VertexChirpRealtimeConfig(
|
||||
resolve_access_token=_token, project="proj-1", location=None, backend_factory=factory
|
||||
)
|
||||
url = config.get_complete_url(None, "vertex_ai/chirp_3")
|
||||
assert url == "us-speech.googleapis.com"
|
||||
assert config.validate_environment({}, MODEL, "https://" + url) == {}
|
||||
backend = await config.open_backend(url, {})
|
||||
assert isinstance(backend, _NullBackend)
|
||||
assert targets == [
|
||||
SpeechStreamingTarget(
|
||||
api_endpoint="us-speech.googleapis.com",
|
||||
recognizer="projects/proj-1/locations/us/recognizers/_",
|
||||
resolve_access_token=_token,
|
||||
)
|
||||
]
|
||||
assert await targets[0].resolve_access_token() == "token"
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("location", "api_base", "endpoint"),
|
||||
[
|
||||
("global", None, "speech.googleapis.com"),
|
||||
("europe-west4", None, "europe-west4-speech.googleapis.com"),
|
||||
("us", "https://speech-proxy.internal:8443/v2", "speech-proxy.internal:8443"),
|
||||
],
|
||||
)
|
||||
def test_get_complete_url_honors_location_and_api_base(location: str, api_base: str | None, endpoint: str):
|
||||
assert _config(location).get_complete_url(api_base, MODEL) == endpoint
|
||||
|
||||
|
||||
def test_get_complete_url_rejects_non_speech_models():
|
||||
with pytest.raises(ValueError, match="Unsupported Speech-to-Text streaming model"):
|
||||
_config().get_complete_url(None, "gemini-live-2.5-flash")
|
||||
|
||||
|
||||
@pytest.mark.parametrize("location", ["bad loc", "../us"])
|
||||
def test_invalid_locations_are_rejected_up_front(location: str):
|
||||
with pytest.raises(VertexAIError):
|
||||
_config(location)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("previous", "current", "delta"),
|
||||
[
|
||||
("", "hello", "hello"),
|
||||
("hello", "hello world", " world"),
|
||||
("hello", "Hello, world", " world"),
|
||||
("hello world", "hello world", ""),
|
||||
("hello there", "hello world", " world"),
|
||||
("hello world", "hello", ""),
|
||||
],
|
||||
)
|
||||
def test_new_words(previous: str, current: str, delta: str):
|
||||
assert new_words(previous, current) == delta
|
||||
|
|
@ -0,0 +1,475 @@
|
|||
"""
|
||||
Tests for the Claude Code gateway protocol (anthropic_endpoints/gateway_endpoints.py).
|
||||
|
||||
Covers the OAuth device-flow surface (RFC 8414 discovery, RFC 8628 device
|
||||
authorization + token), managed settings, OTLP ingestion, and the enable flag.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Iterator, Mapping
|
||||
from contextlib import ExitStack, contextmanager
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
from fastapi import FastAPI
|
||||
from fastapi.testclient import TestClient
|
||||
|
||||
from litellm.caching.dual_cache import DualCache
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.anthropic_endpoints import gateway_endpoints
|
||||
from litellm.proxy.management_endpoints.ui_sso import (
|
||||
_get_cli_sso_flow_cache_key,
|
||||
_hash_cli_sso_secret,
|
||||
_set_cli_sso_flow,
|
||||
)
|
||||
from litellm.proxy.middleware.prometheus_auth_middleware import PrometheusAuthMiddleware
|
||||
|
||||
_DEVICE_CODE_GRANT: Final = "urn:ietf:params:oauth:grant-type:device_code"
|
||||
_MASTER_KEY: Final = "sk-master-key"
|
||||
_SHARED_LOGIN_ID: Final = "cli-shared-login-code"
|
||||
_SHARED_POLL_SECRET: Final = "shared-poll-secret"
|
||||
_SHARED_DEVICE_CODE: Final = f"{_SHARED_LOGIN_ID}.{_SHARED_POLL_SECRET}"
|
||||
_MINT: Final = "litellm.proxy.auth.auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token"
|
||||
_PROTOBUF_BODY: Final = b"\x0a\x05hello\x12\x03{{{"
|
||||
_COMPLETED_SESSION: Final = MappingProxyType(
|
||||
{
|
||||
"user_id": "user-123",
|
||||
"user_role": "internal_user",
|
||||
"models": ["claude-sonnet-4-5"],
|
||||
"teams": ["team-a"],
|
||||
"team_details": [
|
||||
{
|
||||
"team_id": "team-a",
|
||||
"team_alias": "Team A",
|
||||
"team_models": ["claude-sonnet-4-5"],
|
||||
"team_model_aliases": None,
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
class _SharedRedisFake:
|
||||
def __init__(self) -> None:
|
||||
self.values: Mapping[str, object] = MappingProxyType({})
|
||||
self.counters: Mapping[str, float] = MappingProxyType({})
|
||||
|
||||
def set_cache(self, key: str, value: object, **kwargs: object) -> None:
|
||||
self.values = MappingProxyType({**self.values, key: value})
|
||||
|
||||
def get_cache(self, key: str, **kwargs: object) -> object:
|
||||
return self.values.get(key)
|
||||
|
||||
def delete_cache(self, key: str) -> None:
|
||||
self.values = MappingProxyType({name: value for name, value in self.values.items() if name != key})
|
||||
|
||||
async def async_delete_cache(self, key: str) -> None:
|
||||
self.delete_cache(key)
|
||||
|
||||
async def async_increment(self, key: str, value: float, **kwargs: object) -> float:
|
||||
incremented: Final = self.counters.get(key, 0) + value
|
||||
self.counters = MappingProxyType({**self.counters, key: incremented})
|
||||
return incremented
|
||||
|
||||
|
||||
def _replica(redis: _SharedRedisFake) -> DualCache:
|
||||
return DualCache(redis_cache=redis, default_in_memory_ttl=600) # pyright: ignore[reportArgumentType] # duck-typed Redis double
|
||||
|
||||
|
||||
def _real_auth_proxy_attrs() -> Mapping[str, object]:
|
||||
proxy_logging_obj: Final = MagicMock()
|
||||
proxy_logging_obj.internal_usage_cache.dual_cache = AsyncMock()
|
||||
proxy_logging_obj.post_call_failure_hook = AsyncMock(return_value=None)
|
||||
return MappingProxyType(
|
||||
{
|
||||
"master_key": _MASTER_KEY,
|
||||
"prisma_client": None,
|
||||
"user_api_key_cache": DualCache(),
|
||||
"proxy_logging_obj": proxy_logging_obj,
|
||||
"llm_router": None,
|
||||
"llm_model_list": [],
|
||||
"user_custom_auth": None,
|
||||
"litellm_proxy_admin_name": "admin",
|
||||
"jwt_handler": None,
|
||||
"open_telemetry_logger": None,
|
||||
"model_max_budget_limiter": MagicMock(),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _gateway_env(
|
||||
*,
|
||||
enabled: bool = True,
|
||||
managed_settings: Mapping[str, object] | None = None,
|
||||
cache: DualCache | None = None,
|
||||
real_auth: bool = False,
|
||||
extra_settings: Mapping[str, object] = MappingProxyType({}),
|
||||
) -> Iterator[tuple[TestClient, DualCache]]:
|
||||
general_settings: Final = {
|
||||
"enable_claude_code_gateway": enabled,
|
||||
**({} if managed_settings is None else {"claude_code_gateway_managed_settings": dict(managed_settings)}),
|
||||
**extra_settings,
|
||||
}
|
||||
session_cache: Final = cache or DualCache(default_in_memory_ttl=600)
|
||||
|
||||
app: Final = FastAPI()
|
||||
app.add_middleware(PrometheusAuthMiddleware)
|
||||
app.include_router(gateway_endpoints.router)
|
||||
|
||||
async def _fake_auth() -> object:
|
||||
return object()
|
||||
|
||||
with ExitStack() as stack:
|
||||
stack.enter_context(
|
||||
patch( # test-quality-ok: the gateway reads this proxy_server module global and has no injection seam
|
||||
"litellm.proxy.proxy_server.general_settings", general_settings
|
||||
)
|
||||
)
|
||||
stack.enter_context(
|
||||
patch( # test-quality-ok: the CLI SSO flow cache is this proxy_server module global shared with ui_sso
|
||||
"litellm.proxy.proxy_server.cli_sso_session_cache", session_cache
|
||||
)
|
||||
)
|
||||
if real_auth:
|
||||
for name, value in _real_auth_proxy_attrs().items():
|
||||
stack.enter_context(patch(f"litellm.proxy.proxy_server.{name}", value))
|
||||
else:
|
||||
app.dependency_overrides[gateway_endpoints.user_api_key_auth] = _fake_auth
|
||||
with TestClient(app) as client:
|
||||
yield client, session_cache
|
||||
|
||||
|
||||
def _start_device_flow(client: TestClient) -> str:
|
||||
return client.post("/claude_code_gateway/oauth/device_authorization").json()["device_code"]
|
||||
|
||||
|
||||
def _request_token(client: TestClient, device_code: str) -> httpx.Response:
|
||||
return client.post(
|
||||
"/claude_code_gateway/oauth/token",
|
||||
data={"grant_type": _DEVICE_CODE_GRANT, "device_code": device_code},
|
||||
)
|
||||
|
||||
|
||||
def _completed_flow(session_data: Mapping[str, object] = _COMPLETED_SESSION) -> dict[str, object]:
|
||||
return {
|
||||
"poll_secret_hash": _hash_cli_sso_secret(_SHARED_POLL_SECRET),
|
||||
"user_code_hash": "unused",
|
||||
"sso_complete": True,
|
||||
"user_code_verified": True,
|
||||
"session_data": dict(session_data),
|
||||
}
|
||||
|
||||
|
||||
def _login_id(device_code: str) -> str:
|
||||
return device_code.partition(".")[0]
|
||||
|
||||
|
||||
def _complete_flow(
|
||||
cache: DualCache, device_code: str, session_data: Mapping[str, object] = _COMPLETED_SESSION
|
||||
) -> None:
|
||||
key: Final = _get_cli_sso_flow_cache_key(_login_id(device_code))
|
||||
flow: Final = cache.get_cache(key=key)
|
||||
assert isinstance(flow, dict)
|
||||
completed: Final = {**flow, **_completed_flow(session_data), "poll_secret_hash": flow["poll_secret_hash"]}
|
||||
cache.set_cache(key=key, value=completed, ttl=600)
|
||||
|
||||
|
||||
def test_discovery_shape():
|
||||
with _gateway_env() as (client, _):
|
||||
resp = client.get("/claude_code_gateway/.well-known/oauth-authorization-server")
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert body["device_authorization_endpoint"].endswith("/claude_code_gateway/oauth/device_authorization")
|
||||
assert body["token_endpoint"].endswith("/claude_code_gateway/oauth/token")
|
||||
assert body["grant_types_supported"] == [
|
||||
"urn:ietf:params:oauth:grant-type:device_code",
|
||||
"refresh_token",
|
||||
]
|
||||
# authorization_endpoint is intentionally absent (device flow only).
|
||||
assert "authorization_endpoint" not in body
|
||||
# Both endpoints must be same-origin with the issuer.
|
||||
assert body["device_authorization_endpoint"].startswith(body["issuer"])
|
||||
assert body["token_endpoint"].startswith(body["issuer"])
|
||||
|
||||
|
||||
def test_discovery_404_when_disabled():
|
||||
with _gateway_env(enabled=False) as (client, _):
|
||||
resp = client.get("/claude_code_gateway/.well-known/oauth-authorization-server")
|
||||
assert resp.status_code == 404
|
||||
|
||||
|
||||
def test_device_authorization_returns_rfc8628_shape_and_persists_flow():
|
||||
with _gateway_env() as (client, cache):
|
||||
resp = client.post("/claude_code_gateway/oauth/device_authorization")
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
device_code = body["device_code"]
|
||||
login_id, separator, poll_secret = device_code.partition(".")
|
||||
assert login_id.startswith("cli-")
|
||||
assert separator == "."
|
||||
assert len(poll_secret) >= 32
|
||||
assert body["user_code"]
|
||||
assert body["expires_in"] == 600
|
||||
assert body["interval"] == 5
|
||||
assert "verification_uri_complete" not in body
|
||||
assert body["verification_uri"].endswith(f"/sso/key/generate?source=litellm-cli&key={login_id}")
|
||||
assert poll_secret not in body["verification_uri"]
|
||||
stored = cache.get_cache(key=_get_cli_sso_flow_cache_key(login_id))
|
||||
assert isinstance(stored, dict)
|
||||
assert stored["sso_complete"] is False
|
||||
assert stored["poll_secret_hash"] == _hash_cli_sso_secret(poll_secret)
|
||||
assert cache.get_cache(key=_get_cli_sso_flow_cache_key(device_code)) is None
|
||||
|
||||
|
||||
@pytest.mark.parametrize("opted_in", [True, False])
|
||||
def test_verification_uri_complete_carries_the_user_code_only_when_the_operator_opts_in(opted_in: bool):
|
||||
with _gateway_env(extra_settings={"allow_cli_sso_verification_uri_complete": opted_in}) as (client, _):
|
||||
body = client.post("/claude_code_gateway/oauth/device_authorization").json()
|
||||
login_id = _login_id(body["device_code"])
|
||||
if not opted_in:
|
||||
assert "verification_uri_complete" not in body
|
||||
return
|
||||
assert body["verification_uri_complete"].endswith(
|
||||
f"/sso/key/generate?source=litellm-cli&key={login_id}&user_code={body['user_code']}"
|
||||
)
|
||||
assert "user_code=" not in body["verification_uri"]
|
||||
|
||||
|
||||
def test_token_authorization_pending_before_browser_completes():
|
||||
with _gateway_env() as (client, _):
|
||||
resp = _request_token(client, _start_device_flow(client))
|
||||
assert resp.status_code == 400
|
||||
assert resp.json()["error"] == "authorization_pending"
|
||||
|
||||
|
||||
@pytest.mark.parametrize("tamper", ["login_id_only", "wrong_secret"])
|
||||
def test_token_refuses_the_browser_login_id_without_the_client_secret(tamper: str):
|
||||
with _gateway_env() as (client, cache):
|
||||
device_code = _start_device_flow(client)
|
||||
_complete_flow(cache, device_code)
|
||||
login_id = _login_id(device_code)
|
||||
presented = login_id if tamper == "login_id_only" else f"{login_id}.not-the-secret"
|
||||
with patch(_MINT, return_value="sk-session") as mint:
|
||||
resp = _request_token(client, presented)
|
||||
assert resp.status_code == 400
|
||||
assert resp.json()["error"] == "expired_token"
|
||||
mint.assert_not_called()
|
||||
with_secret = _request_token(client, device_code)
|
||||
assert with_secret.status_code == 200
|
||||
|
||||
|
||||
def test_token_success_mints_bearer_and_is_single_use():
|
||||
with _gateway_env() as (client, cache):
|
||||
device_code = _start_device_flow(client)
|
||||
_complete_flow(cache, device_code)
|
||||
|
||||
with patch(_MINT, return_value="sk-litellm-session-token") as mint:
|
||||
resp = _request_token(client, device_code)
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert body["access_token"] == "sk-litellm-session-token"
|
||||
assert body["token_type"] == "Bearer"
|
||||
assert body["expires_in"] > 0
|
||||
|
||||
called_user = mint.call_args.kwargs["user_info"]
|
||||
assert called_user.user_id == "user-123"
|
||||
assert mint.call_args.kwargs["team_id"] == "team-a"
|
||||
assert mint.call_args.kwargs["team_alias"] == "Team A"
|
||||
assert mint.call_args.kwargs["team_models"] == ("claude-sonnet-4-5",)
|
||||
|
||||
# Single-use: the flow is deleted, so a replay returns expired_token.
|
||||
replay = _request_token(client, device_code)
|
||||
assert replay.status_code == 400
|
||||
assert replay.json()["error"] == "expired_token"
|
||||
|
||||
|
||||
def test_token_teamless_user_mints_without_a_team():
|
||||
with _gateway_env() as (client, cache):
|
||||
device_code = _start_device_flow(client)
|
||||
_complete_flow(cache, device_code, session_data={**_COMPLETED_SESSION, "teams": [], "team_details": []})
|
||||
with patch(_MINT, return_value="sk-litellm-session-token") as mint:
|
||||
resp = _request_token(client, device_code)
|
||||
assert resp.status_code == 200
|
||||
assert mint.call_args.kwargs["team_id"] is None
|
||||
assert mint.call_args.kwargs["team_models"] == ()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"session_data",
|
||||
[
|
||||
{"user_role": "internal_user"},
|
||||
{**_COMPLETED_SESSION, "user_role": None},
|
||||
{**_COMPLETED_SESSION, "user_role": "not-a-role"},
|
||||
],
|
||||
ids=["missing_user_id", "no_role", "unknown_role"],
|
||||
)
|
||||
def test_token_malformed_session_is_invalid_grant_and_does_not_consume_the_login(session_data: Mapping[str, object]):
|
||||
with _gateway_env() as (client, cache):
|
||||
device_code = _start_device_flow(client)
|
||||
_complete_flow(cache, device_code, session_data=session_data)
|
||||
with patch(_MINT) as mint:
|
||||
resp = _request_token(client, device_code)
|
||||
again = _request_token(client, device_code)
|
||||
assert resp.status_code == 400
|
||||
assert resp.json()["error"] == "invalid_grant"
|
||||
assert again.json()["error"] == "invalid_grant"
|
||||
mint.assert_not_called()
|
||||
|
||||
|
||||
def test_token_mint_failure_leaves_the_login_unconsumed():
|
||||
with _gateway_env() as (client, cache):
|
||||
device_code = _start_device_flow(client)
|
||||
_complete_flow(cache, device_code)
|
||||
with patch(_MINT, side_effect=RuntimeError("signing key unavailable")), pytest.raises(RuntimeError):
|
||||
_request_token(client, device_code)
|
||||
with patch(_MINT, return_value="sk-session"):
|
||||
retry = _request_token(client, device_code)
|
||||
assert retry.status_code == 200
|
||||
assert retry.json()["access_token"] == "sk-session"
|
||||
|
||||
|
||||
def test_token_unknown_team_grants_is_invalid_grant():
|
||||
with _gateway_env() as (client, cache):
|
||||
device_code = _start_device_flow(client)
|
||||
_complete_flow(cache, device_code, session_data={**_COMPLETED_SESSION, "team_details": []})
|
||||
with patch(_MINT) as mint:
|
||||
resp = _request_token(client, device_code)
|
||||
assert resp.status_code == 400
|
||||
assert resp.json()["error"] == "invalid_grant"
|
||||
mint.assert_not_called()
|
||||
|
||||
|
||||
def test_token_mints_on_a_replica_that_did_not_start_the_login():
|
||||
redis: Final = _SharedRedisFake()
|
||||
_set_cli_sso_flow(login_id=_SHARED_LOGIN_ID, cache=_replica(redis), flow=_completed_flow())
|
||||
|
||||
with _gateway_env(cache=_replica(redis)) as (client, _), patch(_MINT, return_value="sk-session") as mint:
|
||||
resp = _request_token(client, _SHARED_DEVICE_CODE)
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["access_token"] == "sk-session"
|
||||
assert mint.call_args.kwargs["team_id"] == "team-a"
|
||||
assert mint.call_args.kwargs["user_info"].user_role == "internal_user"
|
||||
|
||||
|
||||
def test_token_refuses_a_device_code_another_replica_already_claimed():
|
||||
redis: Final = _SharedRedisFake()
|
||||
replica_a: Final = _replica(redis)
|
||||
_set_cli_sso_flow(login_id=_SHARED_LOGIN_ID, cache=replica_a, flow=_completed_flow())
|
||||
assert asyncio.run(gateway_endpoints._claim_device_code(_SHARED_LOGIN_ID, replica_a)) is True
|
||||
|
||||
with _gateway_env(cache=_replica(redis)) as (client, _), patch(_MINT, return_value="sk-session"):
|
||||
resp = _request_token(client, _SHARED_DEVICE_CODE)
|
||||
assert resp.status_code == 400
|
||||
assert resp.json() == {"error": "expired_token"}
|
||||
|
||||
|
||||
def test_token_unknown_device_code_is_expired_token():
|
||||
with _gateway_env() as (client, _):
|
||||
resp = _request_token(client, "cli-does-not-exist")
|
||||
assert resp.status_code == 400
|
||||
assert resp.json()["error"] == "expired_token"
|
||||
|
||||
|
||||
def test_refresh_grant_forces_relogin():
|
||||
with _gateway_env() as (client, _):
|
||||
resp = client.post(
|
||||
"/claude_code_gateway/oauth/token",
|
||||
data={"grant_type": "refresh_token", "refresh_token": "whatever"},
|
||||
)
|
||||
assert resp.status_code == 401
|
||||
assert resp.json()["error"] == "invalid_grant"
|
||||
|
||||
|
||||
def test_unsupported_grant_type():
|
||||
with _gateway_env() as (client, _):
|
||||
resp = client.post("/claude_code_gateway/oauth/token", data={"grant_type": "password"})
|
||||
assert resp.status_code == 400
|
||||
assert resp.json()["error"] == "unsupported_grant_type"
|
||||
|
||||
|
||||
def test_managed_settings_404_when_unset():
|
||||
with _gateway_env() as (client, _):
|
||||
resp = client.get("/claude_code_gateway/managed/settings")
|
||||
assert resp.status_code == 404
|
||||
|
||||
|
||||
def test_managed_settings_returns_client_envelope_and_304_on_cached_checksum():
|
||||
settings = {"permissions": {"defaultMode": "acceptEdits"}, "env": {"FOO": "bar"}}
|
||||
with _gateway_env(managed_settings=settings) as (client, _):
|
||||
resp = client.get("/claude_code_gateway/managed/settings")
|
||||
assert resp.status_code == 200
|
||||
body = resp.json()
|
||||
assert body["settings"] == settings
|
||||
checksum = body["checksum"]
|
||||
assert checksum.startswith("sha256:")
|
||||
assert body["uuid"] == checksum
|
||||
assert resp.headers["ETag"] == f'"{checksum}"'
|
||||
|
||||
not_modified = client.get(
|
||||
"/claude_code_gateway/managed/settings", headers={"If-None-Match": f'"{checksum}"'}
|
||||
)
|
||||
assert not_modified.status_code == 304
|
||||
assert not_modified.headers["ETag"] == f'"{checksum}"'
|
||||
|
||||
stale = client.get("/claude_code_gateway/managed/settings", headers={"If-None-Match": '"sha256:stale"'})
|
||||
assert stale.status_code == 200
|
||||
assert stale.json()["checksum"] == checksum
|
||||
|
||||
|
||||
def test_managed_settings_checksum_tracks_policy_content():
|
||||
with _gateway_env(managed_settings={"env": {"FOO": "bar"}}) as (client, _):
|
||||
first = client.get("/claude_code_gateway/managed/settings").json()["checksum"]
|
||||
with _gateway_env(managed_settings={"env": {"FOO": "baz"}}) as (client, _):
|
||||
second = client.get("/claude_code_gateway/managed/settings").json()["checksum"]
|
||||
assert first != second
|
||||
|
||||
|
||||
def test_managed_settings_404_when_gateway_disabled():
|
||||
with _gateway_env(enabled=False, managed_settings={"env": {}}) as (client, _):
|
||||
resp = client.get("/claude_code_gateway/managed/settings")
|
||||
assert resp.status_code == 404
|
||||
|
||||
|
||||
@pytest.mark.parametrize("signal", ["metrics", "logs", "traces"])
|
||||
def test_otlp_endpoints_accept_and_return_200(signal: str):
|
||||
with _gateway_env() as (client, _):
|
||||
resp = client.post(f"/claude_code_gateway/v1/{signal}", content=b"\x00\x01binary-otlp")
|
||||
assert resp.status_code == 200
|
||||
|
||||
|
||||
@pytest.mark.parametrize("signal", ["metrics", "logs", "traces"])
|
||||
def test_otlp_endpoints_404_when_disabled(signal: str):
|
||||
with _gateway_env(enabled=False) as (client, _):
|
||||
resp = client.post(f"/claude_code_gateway/v1/{signal}", content=b"payload")
|
||||
assert resp.status_code == 404
|
||||
|
||||
|
||||
@pytest.mark.parametrize("signal", ["metrics", "logs", "traces"])
|
||||
def test_otlp_protobuf_body_is_accepted_through_real_auth(signal: str):
|
||||
with _gateway_env(real_auth=True) as (client, _):
|
||||
resp = client.post(
|
||||
f"/claude_code_gateway/v1/{signal}",
|
||||
content=_PROTOBUF_BODY,
|
||||
headers={"Authorization": f"Bearer {_MASTER_KEY}", "Content-Type": "application/x-protobuf"},
|
||||
)
|
||||
assert resp.status_code == 200
|
||||
|
||||
|
||||
def test_otlp_without_a_bearer_is_rejected_by_real_auth():
|
||||
with _gateway_env(real_auth=True) as (client, _), pytest.raises(ProxyException) as exc_info:
|
||||
client.post(
|
||||
"/claude_code_gateway/v1/metrics",
|
||||
content=_PROTOBUF_BODY,
|
||||
headers={"Content-Type": "application/x-protobuf"},
|
||||
)
|
||||
assert exc_info.value.code == "401"
|
||||
|
||||
|
||||
def test_messages_gated_by_enable_flag():
|
||||
with _gateway_env(enabled=False) as (client, _):
|
||||
resp = client.post("/claude_code_gateway/v1/messages", json={"model": "claude-sonnet-4-5", "messages": []})
|
||||
assert resp.status_code == 404
|
||||
|
|
@ -910,6 +910,36 @@ def test_anthropic_count_tokens_route_accessible_to_internal_users():
|
|||
assert RouteChecks.is_llm_api_route("/v1/messages") is True
|
||||
|
||||
|
||||
_CLAUDE_CODE_GATEWAY_ROUTES: Final = (
|
||||
"/claude_code_gateway/v1/messages",
|
||||
"/claude_code_gateway/v1/messages/count_tokens",
|
||||
"/claude_code_gateway/managed/settings",
|
||||
"/claude_code_gateway/v1/metrics",
|
||||
"/claude_code_gateway/v1/logs",
|
||||
"/claude_code_gateway/v1/traces",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("route", _CLAUDE_CODE_GATEWAY_ROUTES)
|
||||
@pytest.mark.parametrize(
|
||||
"role", [LitellmUserRoles.INTERNAL_USER.value, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value]
|
||||
)
|
||||
def test_claude_code_gateway_routes_open_to_signed_in_cli_users(role: str, route: str):
|
||||
user_obj: Final = LiteLLM_UserTable(user_id="test_user", user_email="test@example.com", user_role=role)
|
||||
valid_token: Final = UserAPIKeyAuth(user_id="test_user", user_role=role)
|
||||
request: Final = MagicMock(spec=Request)
|
||||
request.query_params = {}
|
||||
|
||||
RouteChecks.non_proxy_admin_allowed_routes_check(
|
||||
user_obj=user_obj,
|
||||
_user_role=role,
|
||||
route=route,
|
||||
request=request,
|
||||
valid_token=valid_token,
|
||||
request_data={},
|
||||
)
|
||||
|
||||
|
||||
def test_virtual_key_llm_api_routes_allows_registered_pass_through_endpoints():
|
||||
"""
|
||||
Virtual keys with llm_api_routes can access auth=true pass-through endpoints only when
|
||||
|
|
|
|||
|
|
@ -573,6 +573,13 @@ async def test_lone_surrogate_escape_is_rejected_with_400(content: bytes):
|
|||
assert parsed["messages"][0]["content"] == "say ok \U0001F600"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("media_type", ["application/x-protobuf", "application/protobuf", "application/octet-stream"])
|
||||
async def test_json_body_under_a_binary_content_type_is_still_parsed(media_type: str):
|
||||
request = _starlette_request(b'{"model": "claude-sonnet-5"}', media_type)
|
||||
assert await _read_request_body(request) == {"model": "claude-sonnet-5"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_form_data():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -51,6 +51,14 @@ def app_with_middleware():
|
|||
async def embeddings():
|
||||
return {"msg": "embeddings OK"}
|
||||
|
||||
@app.post("/claude_code_gateway/v1/metrics")
|
||||
async def gateway_telemetry():
|
||||
return {"msg": "gateway telemetry OK"}
|
||||
|
||||
@app.get("/metrics/detail")
|
||||
async def metrics_detail():
|
||||
return {"msg": "metrics detail OK"}
|
||||
|
||||
return app
|
||||
|
||||
|
||||
|
|
@ -240,3 +248,63 @@ def test_non_metrics_requests_dont_trigger_auth(app_with_middleware, monkeypatch
|
|||
response = client.get("/embeddings")
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.json() == {"msg": "embeddings OK"}
|
||||
|
||||
|
||||
def test_gateway_telemetry_path_is_not_treated_as_the_metrics_endpoint(app_with_middleware, monkeypatch):
|
||||
monkeypatch.setattr(litellm, "require_auth_for_metrics_endpoint", True)
|
||||
|
||||
def should_not_be_called(*args, **kwargs):
|
||||
raise Exception("Auth should not be called for the gateway telemetry route")
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.middleware.prometheus_auth_middleware.user_api_key_auth",
|
||||
should_not_be_called,
|
||||
)
|
||||
|
||||
client = TestClient(app_with_middleware)
|
||||
|
||||
response = client.post("/claude_code_gateway/v1/metrics", content=b"\x0a\x05hello")
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.json() == {"msg": "gateway telemetry OK"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("path", ["/metrics", "/metrics/", "/metrics/detail"])
|
||||
def test_metrics_paths_still_require_auth(app_with_middleware, monkeypatch, path):
|
||||
monkeypatch.setattr(litellm, "require_auth_for_metrics_endpoint", True)
|
||||
|
||||
async def reject(*args, **kwargs):
|
||||
raise Exception("Invalid API key")
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.middleware.prometheus_auth_middleware.user_api_key_auth",
|
||||
reject,
|
||||
)
|
||||
|
||||
client = TestClient(app_with_middleware)
|
||||
|
||||
response = client.get(path)
|
||||
assert response.status_code == 401, response.text
|
||||
|
||||
|
||||
def test_metrics_under_a_root_path_still_requires_auth(monkeypatch):
|
||||
monkeypatch.setattr(litellm, "require_auth_for_metrics_endpoint", True)
|
||||
|
||||
async def reject(*args, **kwargs):
|
||||
raise Exception("Invalid API key")
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.middleware.prometheus_auth_middleware.user_api_key_auth",
|
||||
reject,
|
||||
)
|
||||
|
||||
app = FastAPI(root_path="/litellm")
|
||||
app.add_middleware(PrometheusAuthMiddleware)
|
||||
|
||||
@app.get("/metrics")
|
||||
async def metrics():
|
||||
return {"msg": "metrics OK"}
|
||||
|
||||
client = TestClient(app, root_path="/litellm")
|
||||
|
||||
response = client.get("/metrics")
|
||||
assert response.status_code == 401, response.text
|
||||
|
|
|
|||
|
|
@ -499,3 +499,69 @@ async def test_arealtime_azure_env_beta_protocol_wins_over_a_ga_client(monkeypat
|
|||
assert await _azure_backend_url_dialed_for(_GA_CLIENT) == (
|
||||
"wss://my-endpoint.openai.azure.com/openai/realtime?api-version=2024-10-01-preview&deployment=gpt-realtime"
|
||||
)
|
||||
|
||||
|
||||
async def _vertex_provider_config_for(monkeypatch, model: str, vertex_location: str | None):
|
||||
from litellm.llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig
|
||||
from litellm.llms.vertex_ai.audio_transcription.realtime_transformation import VertexChirpRealtimeConfig
|
||||
|
||||
captured: dict[str, object] = {}
|
||||
|
||||
def mock_get_llm_provider(model, api_base, api_key):
|
||||
return model.removeprefix("vertex_ai/"), "vertex_ai", None, api_base
|
||||
|
||||
async def mock_token_resolver(**kwargs):
|
||||
return "access-token", kwargs["project_id"]
|
||||
|
||||
async def mock_async_realtime(**kwargs):
|
||||
captured.update(kwargs)
|
||||
|
||||
monkeypatch.setattr(realtime_main, "get_llm_provider", mock_get_llm_provider)
|
||||
monkeypatch.setattr(realtime_main, "vertex_access_token_resolver", mock_token_resolver)
|
||||
monkeypatch.setattr(realtime_main.base_llm_http_handler, "async_realtime", mock_async_realtime)
|
||||
monkeypatch.setattr(litellm, "vertex_location", None)
|
||||
monkeypatch.delenv("VERTEXAI_LOCATION", raising=False)
|
||||
await realtime_main._arealtime.__wrapped__(
|
||||
model=model,
|
||||
websocket=MagicMock(),
|
||||
litellm_logging_obj=FakeLogging(),
|
||||
query_params={"model": model, "intent": "transcription"},
|
||||
vertex_credentials="fake-credentials",
|
||||
vertex_project="proj-1",
|
||||
vertex_location=vertex_location,
|
||||
)
|
||||
provider_config = captured["provider_config"]
|
||||
assert isinstance(provider_config, (VertexAIRealtimeConfig, VertexChirpRealtimeConfig))
|
||||
return provider_config, captured["model"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_arealtime_routes_chirp_models_to_the_speech_to_text_backend(monkeypatch):
|
||||
from litellm.llms.vertex_ai.audio_transcription.realtime_transformation import VertexChirpRealtimeConfig
|
||||
|
||||
provider_config, model = await _vertex_provider_config_for(monkeypatch, "vertex_ai/chirp_3", None)
|
||||
assert isinstance(provider_config, VertexChirpRealtimeConfig)
|
||||
assert model == "chirp_3"
|
||||
assert provider_config.get_complete_url(None, model) == "us-speech.googleapis.com"
|
||||
assert provider_config.validate_environment({}, model, "https://us-speech.googleapis.com") == {}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_arealtime_routes_chirp_models_to_the_configured_speech_region(monkeypatch):
|
||||
provider_config, model = await _vertex_provider_config_for(monkeypatch, "vertex_ai/chirp_3", "europe-west4")
|
||||
assert provider_config.get_complete_url(None, model) == "europe-west4-speech.googleapis.com"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_arealtime_keeps_gemini_live_on_the_vertex_realtime_websocket(monkeypatch):
|
||||
from litellm.llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig
|
||||
|
||||
provider_config, model = await _vertex_provider_config_for(monkeypatch, "vertex_ai/gemini-live-2.5-flash", None)
|
||||
assert isinstance(provider_config, VertexAIRealtimeConfig)
|
||||
assert provider_config.get_complete_url(None, model).startswith("wss://us-central1-aiplatform.googleapis.com/")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_realtime_health_check_names_the_batch_mode_for_chirp_models():
|
||||
with pytest.raises(ValueError, match="mode audio_transcription"):
|
||||
await realtime_main._realtime_health_check(model="chirp_3", custom_llm_provider="vertex_ai", api_key=None)
|
||||
|
|
|
|||
|
|
@ -1060,7 +1060,7 @@ def test_max_tokens_consistency():
|
|||
if len(inconsistencies) > 10:
|
||||
error_msg += f"\n ... and {len(inconsistencies) - 10} more\n"
|
||||
|
||||
error_msg += "\nTo fix these inconsistencies, run: poetry run python fix_max_tokens_inconsistencies.py"
|
||||
error_msg += "\nTo fix these inconsistencies, run: uv run python fix_max_tokens_inconsistencies.py"
|
||||
raise AssertionError(error_msg)
|
||||
|
||||
|
||||
|
|
|
|||
12
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
12
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -26686,6 +26686,13 @@ export interface components {
|
|||
* @description cancel the in-flight upstream LLM request (non-streaming) when the client disconnects, freeing backend capacity (e.g. a vLLM GPU slot); the request is logged as a 499 failure
|
||||
*/
|
||||
cancel_on_disconnect?: boolean | null;
|
||||
/**
|
||||
* Claude Code Gateway Managed Settings
|
||||
* @description Claude Code managed-settings.json served verbatim at the gateway's /claude_code_gateway/managed/settings endpoint. When unset the endpoint returns 404 (no managed policy)
|
||||
*/
|
||||
claude_code_gateway_managed_settings?: {
|
||||
[key: string]: unknown;
|
||||
} | null;
|
||||
/**
|
||||
* Completion Model
|
||||
* @description proxy level default model for all chat completion calls
|
||||
|
|
@ -26780,6 +26787,11 @@ export interface components {
|
|||
* @description If True, disables ownership enforcement on Responses API ids. Keys may then retrieve, cancel, delete, and chain from any response id, including ids belonging to another user or team and ids this proxy never issued. WARNING: this removes tenant isolation on /v1/responses
|
||||
*/
|
||||
disable_responses_id_security?: boolean | null;
|
||||
/**
|
||||
* Enable Claude Code Gateway
|
||||
* @description serve the Claude Code gateway protocol (https://code.claude.com/docs/en/claude-apps-gateway) under /claude_code_gateway: OAuth device-flow sign-in reusing proxy SSO, plus managed settings and OTLP telemetry ingestion. Off by default
|
||||
*/
|
||||
enable_claude_code_gateway?: boolean | null;
|
||||
/**
|
||||
* Enable Openai Websocket Passthrough
|
||||
* @description Serve the OpenAI pass-through WebSocket route, which relays frames to OpenAI under the proxy's own provider credential without reading them. Off by default.
|
||||
|
|
|
|||
27
uv.lock
generated
27
uv.lock
generated
|
|
@ -2615,6 +2615,23 @@ wheels = [
|
|||
{ url = "https://files.pythonhosted.org/packages/3d/f7/661d7a9023e877a226b5683429c3662f75a29ef45cb1464cf39adb689218/google_cloud_resource_manager-1.17.0-py3-none-any.whl", hash = "sha256:e479baf4b014a57f298e01b8279e3290b032e3476d69c8e5e1427af8f82739a5", size = 404403, upload-time = "2026-03-26T22:15:26.57Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "google-cloud-speech"
|
||||
version = "2.40.0"
|
||||
source = { registry = "https://pypi.org/simple" }
|
||||
dependencies = [
|
||||
{ name = "google-api-core", version = "2.25.2", source = { registry = "https://pypi.org/simple" }, extra = ["grpc"], marker = "python_full_version >= '3.14'" },
|
||||
{ name = "google-api-core", version = "2.30.3", source = { registry = "https://pypi.org/simple" }, extra = ["grpc"], marker = "python_full_version < '3.14'" },
|
||||
{ name = "google-auth" },
|
||||
{ name = "grpcio" },
|
||||
{ name = "proto-plus" },
|
||||
{ name = "protobuf" },
|
||||
]
|
||||
sdist = { url = "https://files.pythonhosted.org/packages/5a/c1/5dc9795314f4aefea0b01b02e9f5486a198341ecc15fe47f89a61c68df63/google_cloud_speech-2.40.0.tar.gz", hash = "sha256:e89e688e4ce0b926754038bf992d0d0f065c5f1c3503bb20e6c46d08b63658fc", size = 404366, upload-time = "2026-06-03T16:13:59.506Z" }
|
||||
wheels = [
|
||||
{ url = "https://files.pythonhosted.org/packages/cc/78/afeca8d597fab54bdd823f857aad15d6f9c4628ff3cb72aa237d01700721/google_cloud_speech-2.40.0-py3-none-any.whl", hash = "sha256:7cc0302b3b9ca33d2eae9669da94a44316601a240942895362ac70e765b9f39c", size = 345427, upload-time = "2026-06-03T16:12:40.909Z" },
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "google-cloud-storage"
|
||||
version = "3.4.1"
|
||||
|
|
@ -4566,6 +4583,7 @@ proxy-runtime = [
|
|||
{ name = "ddtrace" },
|
||||
{ name = "detect-secrets" },
|
||||
{ name = "google-cloud-aiplatform" },
|
||||
{ name = "google-cloud-speech" },
|
||||
{ name = "google-genai" },
|
||||
{ name = "grpcio" },
|
||||
{ name = "langfuse" },
|
||||
|
|
@ -4593,6 +4611,9 @@ stt-nvidia-riva = [
|
|||
{ name = "nvidia-riva-client" },
|
||||
{ name = "soundfile" },
|
||||
]
|
||||
stt-vertex-chirp = [
|
||||
{ name = "google-cloud-speech" },
|
||||
]
|
||||
utils = [
|
||||
{ name = "numpydoc" },
|
||||
]
|
||||
|
|
@ -4607,6 +4628,7 @@ ci = [
|
|||
{ name = "blockbuster" },
|
||||
{ name = "claude-agent-sdk" },
|
||||
{ name = "detect-secrets" },
|
||||
{ name = "google-cloud-speech" },
|
||||
{ name = "google-generativeai" },
|
||||
{ name = "jsonlines" },
|
||||
{ name = "langchain" },
|
||||
|
|
@ -4721,6 +4743,8 @@ requires-dist = [
|
|||
{ name = "google-cloud-aiplatform", marker = "extra == 'proxy-runtime'", specifier = ">=1.133.0,<2.0" },
|
||||
{ name = "google-cloud-iam", marker = "extra == 'extra-proxy'", specifier = ">=2.19.1,<3.0" },
|
||||
{ name = "google-cloud-kms", marker = "extra == 'extra-proxy'", specifier = ">=2.24.2,<3.0" },
|
||||
{ name = "google-cloud-speech", marker = "extra == 'proxy-runtime'", specifier = ">=2.40.0,<3.0" },
|
||||
{ name = "google-cloud-speech", marker = "extra == 'stt-vertex-chirp'", specifier = ">=2.40.0,<3.0" },
|
||||
{ name = "google-genai", marker = "extra == 'proxy-runtime'", specifier = ">=1.37.0,<2.0" },
|
||||
{ name = "granian", marker = "extra == 'proxy'", specifier = ">=2.7.4,<3.0" },
|
||||
{ name = "grpcio", marker = "extra == 'grpc'", specifier = "==1.78.0" },
|
||||
|
|
@ -4787,7 +4811,7 @@ requires-dist = [
|
|||
{ name = "uvloop", marker = "sys_platform != 'win32' and extra == 'proxy'", specifier = ">=0.22.1,<1.0" },
|
||||
{ name = "websockets", marker = "extra == 'proxy'", specifier = ">=15.0.1,<16.0" },
|
||||
]
|
||||
provides-extras = ["proxy", "cli", "extra-proxy", "utils", "caching", "mcp", "saml", "semantic-router", "mlflow", "grpc", "stt-nvidia-riva", "google", "bedrock-realtime", "proxy-runtime"]
|
||||
provides-extras = ["proxy", "cli", "extra-proxy", "utils", "caching", "mcp", "saml", "semantic-router", "mlflow", "grpc", "stt-vertex-chirp", "stt-nvidia-riva", "google", "bedrock-realtime", "proxy-runtime"]
|
||||
|
||||
[package.metadata.requires-dev]
|
||||
ci = [
|
||||
|
|
@ -4799,6 +4823,7 @@ ci = [
|
|||
{ name = "blockbuster", specifier = "==1.5.26" },
|
||||
{ name = "claude-agent-sdk", specifier = "==0.1.44" },
|
||||
{ name = "detect-secrets", specifier = "==1.5.0" },
|
||||
{ name = "google-cloud-speech", specifier = "==2.40.0" },
|
||||
{ name = "google-generativeai", specifier = "==0.8.6" },
|
||||
{ name = "jsonlines", specifier = "==4.0.0" },
|
||||
{ name = "langchain", specifier = "==1.3.9" },
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue