Merge branch 'BerriAI:litellm_internal_staging' into fix-bedrock-converse-tool-pairing-sanitize

This commit is contained in:
lance-cognichip 2026-05-29 14:00:07 -04:00 • committed by GitHub
commit cfadb72519
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
162 changed files with 10924 additions and 5081 deletions

View file

@ -10,9 +10,9 @@
**Please complete all items before asking a LiteLLM maintainer to review your PR**
- [ ] I have Added testing in the [`tests/test_litellm/`](https://github.com/BerriAI/litellm/tree/main/tests/test_litellm) directory, **Adding at least 1 test is a hard requirement** - [see details](https://docs.litellm.ai/docs/extras/contributing_code)
- [ ] I have added meaningful tests
- [ ] My PR passes all unit tests on [`make test-unit`](https://docs.litellm.ai/docs/extras/contributing_code)
- [ ] My PR's scope is as isolated as possible, it only solves 1 specific problem
- [ ] My PR's scope is as isolated as possible; it only solves 1 specific problem
- [ ] I have requested a Greptile review by commenting `@greptileai` and received a **Confidence Score of at least 4/5** before requesting a maintainer review
## Delays in PR merge?

View file

@ -101,6 +101,31 @@ jobs:
docker logs litellm-test
exit 1
- name: Setup Node for Playwright
uses: actions/setup-node@49933ea5288caeca8642d1e84afbd3f7d6820020 # v4.4.0
with:
node-version: "20"
- name: Install UI deps and Chromium
working-directory: ui/litellm-dashboard
run: |
npm ci
npx playwright install --with-deps chromium
- name: Run SERVER_ROOT_PATH redirect e2e
working-directory: ui/litellm-dashboard
env:
SERVER_ROOT_PATH: ${{ matrix.root_path }}
run: npx playwright test --config=e2e_tests/serverRootPath.config.ts
- name: Upload Playwright artifacts on failure
if: failure()
uses: actions/upload-artifact@ea165f8d65b6e75b540449e92b4886f43607fa02 # v4.6.2
with:
name: playwright-trace-${{ strategy.job-index }}
path: ui/litellm-dashboard/test-results/
retention-days: 7
- name: Cleanup
if: always()
run: |

294
AGENTS.md
View file

@ -1,293 +1 @@
# INSTRUCTIONS FOR LITELLM
This document provides comprehensive instructions for AI agents working in the LiteLLM repository.
## OVERVIEW
LiteLLM is a unified interface for 100+ LLMs that:
- Translates inputs to provider-specific completion, embedding, and image generation endpoints
- Provides consistent OpenAI-format output across all providers
- Includes retry/fallback logic across multiple deployments (Router)
- Offers a proxy server (LLM Gateway) with budgets, rate limits, and authentication
- Supports advanced features like function calling, streaming, caching, and observability
## REPOSITORY STRUCTURE
### Core Components
- `litellm/` - Main library code
- `llms/` - Provider-specific implementations (OpenAI, Anthropic, Azure, etc.)
- `proxy/` - Proxy server implementation (LLM Gateway)
- `router_utils/` - Load balancing and fallback logic
- `types/` - Type definitions and schemas
- `integrations/` - Third-party integrations (observability, caching, etc.)
### Key Directories
- `tests/` - Comprehensive test suites
- `ui/litellm-dashboard/` - Admin dashboard UI
- `enterprise/` - Enterprise-specific features
Documentation lives in the separate [BerriAI/litellm-docs](https://github.com/BerriAI/litellm-docs) repository and is served at [docs.litellm.ai](https://docs.litellm.ai).
## DEVELOPMENT GUIDELINES
### MAKING CODE CHANGES
1. **Provider Implementations**: When adding/modifying LLM providers:
- Follow existing patterns in `litellm/llms/{provider}/`
- Implement proper transformation classes that inherit from `BaseConfig`
- Support both sync and async operations
- Handle streaming responses appropriately
- Include proper error handling with provider-specific exceptions
2. **Type Safety**:
- Use proper type hints throughout
- Update type definitions in `litellm/types/`
- Ensure compatibility with both Pydantic v1 and v2
3. **Testing**:
- Add tests in appropriate `tests/` subdirectories
- Include both unit tests and integration tests
- Test provider-specific functionality thoroughly
- Consider adding load tests for performance-critical changes
### MAKING CODE CHANGES FOR THE UI (IGNORE FOR BACKEND)
1. **Always use `antd` for new UI components — Tremor is DEPRECATED**
- We are migrating off of `@tremor/react`. Do not introduce new `Badge`, `Text`, `Card`, `Grid`, `Title`, or other imports from `@tremor/react` in any new or modified file.
- Use `antd` equivalents: `Tag` for labels, plain `<span>`/`<div>` with Tailwind classes (or `Typography.Text`) for text, `Card` from `antd`, etc. Note that `antd` has no `"yellow"` Tag color — use `"gold"` for amber/yellow.
- The only exception is the Tremor Table component and its required Tremor Table sub components.
2. **Use Common Components as much as possible**:
- These are usually defined in the `common_components` directory
- Use these components as much as possible and avoid building new components unless needed
3. **Testing**:
- The codebase uses **Vitest** and **React Testing Library**
- **Query Priority Order**: Use query methods in this order: `getByRole`, `getByLabelText`, `getByPlaceholderText`, `getByText`, `getByTestId`
- **Always use `screen`** instead of destructuring from `render()` (e.g., use `screen.getByText()` not `getByText`)
- **Wrap user interactions in `act()`**: Always wrap `fireEvent` calls with `act()` to ensure React state updates are properly handled
- **Use `query` methods for absence checks**: Use `queryBy*` methods (not `getBy*`) when expecting an element to NOT be present
- **Test names must start with "should"**: All test names should follow the pattern `it("should ...")`
- **Mock external dependencies**: Check `setupTests.ts` for global mocks and mock child components/networking calls as needed
- **Structure tests properly**:
- First test should verify the component renders successfully
- Subsequent tests should focus on functionality and user interactions
- Use `waitFor` for async operations that aren't already awaited
- **Avoid using `querySelector`**: Prefer React Testing Library queries over direct DOM manipulation
### IMPORTANT PATTERNS
1. **Function/Tool Calling**:
- LiteLLM standardizes tool calling across providers
- OpenAI format is the standard, with transformations for other providers
- See `litellm/llms/anthropic/chat/transformation.py` for complex tool handling
2. **Streaming**:
- All providers should support streaming where possible
- Use consistent chunk formatting across providers
- Handle both sync and async streaming
3. **Error Handling**:
- Use provider-specific exception classes
- Maintain consistent error formats across providers
- Include proper retry logic and fallback mechanisms
4. **Configuration**:
- Support both environment variables and programmatic configuration
- Use `BaseConfig` classes for provider configurations
- Allow dynamic parameter passing
## PROXY SERVER (LLM GATEWAY)
The proxy server is a critical component that provides:
- Authentication and authorization
- Rate limiting and budget management
- Load balancing across multiple models/deployments
- Observability and logging
- Admin dashboard UI
- Enterprise features
Key files:
- `litellm/proxy/proxy_server.py` - Main server implementation
- `litellm/proxy/auth/` - Authentication logic
- `litellm/proxy/management_endpoints/` - Admin API endpoints
**Database (proxy)**: Use Prisma model methods (`prisma_client.db.<model>.upsert`, `.find_many`, `.find_unique`, etc.), not raw SQL (`execute_raw`/`query_raw`). See COMMON PITFALLS for details.
## MCP (MODEL CONTEXT PROTOCOL) SUPPORT
LiteLLM supports MCP for agent workflows:
- MCP server integration for tool calling
- Transformation between OpenAI and MCP tool formats
- Support for external MCP servers (Zapier, Jira, Linear, etc.)
- See `litellm/experimental_mcp_client/` and `litellm/proxy/_experimental/mcp_server/`
## RUNNING SCRIPTS
Use `uv run python script.py` to run Python scripts in the project environment (for non-test files).
## GITHUB TEMPLATES
When opening issues or pull requests, follow these templates:
### Bug Reports (`.github/ISSUE_TEMPLATE/bug_report.yml`)
- Describe what happened vs. expected behavior
- Include relevant log output
- Specify LiteLLM version
- Indicate if you're part of an ML Ops team (helps with prioritization)
### Feature Requests (`.github/ISSUE_TEMPLATE/feature_request.yml`)
- Clearly describe the feature
- Explain motivation and use case with concrete examples
### Pull Requests (`.github/pull_request_template.md`)
- Add at least 1 test in `tests/litellm/`
- Ensure `make test-unit` passes
## TESTING CONSIDERATIONS
1. **Provider Tests**: Test against real provider APIs when possible
2. **Proxy Tests**: Include authentication, rate limiting, and routing tests
3. **Performance Tests**: Load testing for high-throughput scenarios
4. **Integration Tests**: End-to-end workflows including tool calling
## DOCUMENTATION
- Keep documentation in sync with code changes
- Update provider documentation when adding new providers
- Include code examples for new features
- Update changelog and release notes
## SECURITY CONSIDERATIONS
- Handle API keys securely
- Validate all inputs, especially for proxy endpoints
- Consider rate limiting and abuse prevention
- Follow security best practices for authentication
## ENTERPRISE FEATURES
- Some features are enterprise-only
- Check `enterprise/` directory for enterprise-specific code
- Maintain compatibility between open-source and enterprise versions
## COMMON PITFALLS TO AVOID
1. **Breaking Changes**: LiteLLM has many users - avoid breaking existing APIs
2. **Provider Specifics**: Each provider has unique quirks - handle them properly
3. **Rate Limits**: Respect provider rate limits in tests
4. **Memory Usage**: Be mindful of memory usage in streaming scenarios
5. **Dependencies**: Keep dependencies minimal and well-justified
6. **UI/Backend Contract Mismatch**: When adding a new entity type to the UI, always check whether the backend endpoint accepts a single value or an array. Match the UI control accordingly (single-select vs. multi-select) to avoid silently dropping user selections
7. **Missing Tests for New Entity Types**: When adding a new entity type (e.g., in `EntityUsage`, `UsageViewSelect`), always add corresponding tests in the existing test files and update any icon/component mocks
8. **Raw SQL in proxy DB code**: Do not use `execute_raw` or `query_raw` for proxy database access. Use Prisma model methods (e.g. `prisma_client.db.litellm_tooltable.upsert()`, `.find_many()`, `.find_unique()`) so behavior stays consistent with the schema, the client stays mockable in tests, and you avoid the pitfalls of hand-written SQL (parameter ordering, type casting, schema drift)
8. **Do not hardcode model-specific flags**: Put model-specific capability flags in `model_prices_and_context_window.json` and read them via `get_model_info` (or existing helpers like `supports_reasoning`). This prevents users from needing to upgrade LiteLLM each time a new model supports a feature.
**Example of BAD** (hardcoded model checks):
```python
@staticmethod
def _is_effort_supported_model(model: str) -> bool:
"""Check if the model supports the output_config.effort parameter..."""
model_lower = model.lower()
if AnthropicConfig._is_claude_4_6_model(model):
return True
return any(
v in model_lower for v in ("opus-4-5", "opus_4_5", "opus-4.5", "opus_4.5")
)
```
**Example of GOOD** (config-driven or helper that reads from config):
```python
if (
"claude-3-7-sonnet" in model
or AnthropicConfig._is_claude_4_6_model(model)
or supports_reasoning(
model=model,
custom_llm_provider=self.custom_llm_provider,
)
):
...
```
Using helpers like `supports_reasoning` (which read from `model_prices_and_context_window.json` / `get_model_info`) allows future model updates to "just work" without code changes.
9. **Never close HTTP/SDK clients on cache eviction**: Do not add `close()`, `aclose()`, or `create_task(close_fn())` inside `LLMClientCache._remove_key()` or any cache eviction path. Evicted clients may still be held by in-flight requests; closing them causes `RuntimeError: Cannot send a request, as the client has been closed.` in production after the cache TTL (1 hour) expires. Connection cleanup is handled at shutdown by `close_litellm_async_clients()`. See PR #22247 for the full incident history.
## HELPFUL RESOURCES
- Main documentation: https://docs.litellm.ai/ (source: [BerriAI/litellm-docs](https://github.com/BerriAI/litellm-docs))
- Provider-specific docs: https://docs.litellm.ai/docs/providers/
- Admin UI for testing proxy features
## WHEN IN DOUBT
- Follow existing patterns in the codebase
- Check similar provider implementations
- Ensure comprehensive test coverage
- Update documentation appropriately
- Consider backward compatibility impact
## Cursor Cloud specific instructions
### Environment
- uv is installed in `~/.local/bin`; the update script ensures it is on `PATH`.
- Python 3.12, Node 22 are pre-installed.
- The project virtual environment lives under `.venv/`.
### Running the proxy server
Create a minimal config file and start the proxy:
```yaml
# config.yaml
model_list:
- model_name: fake-openai-endpoint
litellm_params:
model: openai/fake-model
api_key: fake-key
api_base: https://fake-api.example.com
general_settings:
master_key: sk-1234
litellm_settings:
drop_params: True
telemetry: False
```
```bash
uv run litellm --config config.yaml --port 4000
```
The proxy takes ~15-20 seconds to fully start (it runs Prisma migrations on boot). Wait for `/health` to return before sending requests. Without a PostgreSQL `DATABASE_URL`, the proxy connects to a default Neon dev database embedded in the `litellm-proxy-extras` package.
### Running tests
See `CLAUDE.md` and the `Makefile` for standard commands. Key notes:
- `uv sync --group proxy-dev --extra proxy` installs the Prisma and proxy-side test dependencies used by the standard local workflow.
- The `--timeout` pytest flag is NOT available; don't pass it.
- Unit tests: `uv run pytest tests/test_litellm/ -x -vv -n 4`
- **Before committing, always run `uv run black .` to format your code.** Black formatting is enforced in CI.
- If `uv sync` fails because the lockfile is outdated, run `uv lock` and retry.
### Lint
```bash
cd litellm && uv run ruff check .
```
Ruff is the primary fast linter. For the full lint suite (including mypy, black, circular imports), run `make lint` per `CLAUDE.md`.
### UI Dashboard development
- The UI is at `ui/litellm-dashboard/`. Run `npm run dev` from that directory for the Next.js dev server on port 3000.
- The proxy at port 4000 serves a **pre-built** static UI from `litellm/proxy/_experimental/out/`. After making UI code changes, you must run `npm run build` in the dashboard directory and copy the output: `cp -r ui/litellm-dashboard/out/* litellm/proxy/_experimental/out/` for the proxy to serve the updated UI.
- SVGs used as provider logos (loaded via `<img>` tags) must NOT use `fill="currentColor"` — replace with an explicit color like `#000000` or use the `-color` variant from lobehub icons, since CSS color inheritance does not work inside `<img>` elements.
- Provider logos live in `ui/litellm-dashboard/public/assets/logos/` (source) and `litellm/proxy/_experimental/out/assets/logos/` (pre-built). Both locations must have the file for it to work in dev and proxy-served modes.
- UI Vitest tests: `cd ui/litellm-dashboard && npx vitest run`
Read @CLAUDE.md for coding guidelines

201
CLAUDE.md
View file

@ -1,181 +1,70 @@
# CLAUDE.md
Do not write comments unless they are absolutely necessary to explain some very complex business logic. Please clean up if there are comments that are not absolutely necessary. Do not remove comments that are unrelated to the addition of the code of this PR
This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository.
Explanation: 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 to the reader, while being both easy to maintain and high performance
## Documentation
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 that order of importance
Documentation lives in a separate repository: [BerriAI/litellm-docs](https://github.com/BerriAI/litellm-docs). It is served at [docs.litellm.ai](https://docs.litellm.ai). Do not create or edit documentation files in this repository — open doc PRs against `BerriAI/litellm-docs` instead.
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
## Development Commands
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)
### Installation
- `make install-dev` - Install core development dependencies
- `make install-proxy-dev` - Install proxy development dependencies with full feature set
- `make install-test-deps` - Install the full local test environment and generate the Prisma client
When creating PRs, don't set base to `main`. `litellm_internal_staging` serves that purpose
### Testing
- `make test` - Run all tests
- `make test-unit` - Run unit tests (tests/test_litellm) with 4 parallel workers
- `make test-integration` - Run integration tests (excludes unit tests)
- `pytest tests/` - Direct pytest execution
Always use @.github/pull_request_template.md as a guide for your PR body
### Code Quality
- `make lint` - Run all linting (Ruff, MyPy, Black, circular imports, import safety)
- `make format` - Apply Black code formatting
- `make lint-ruff` - Run Ruff linting only
- `make lint-mypy` - Run MyPy type checking only
- **Before committing, always run `uv run black .` to format your code.** Black formatting is enforced in CI.
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`) 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, just tell me which URLs to go to (e.g., http://localhost:4000/ui/?page=logs), where to click, what fields to fill out, etc. along with the other commands to run in an ordered list, and I'll do it myself and post the screenshots after you make the PR
### Single Test Files
- `uv run pytest tests/path/to/test_file.py -v` - Run specific test file
- `uv run pytest tests/path/to/test_file.py::test_function -v` - Run specific test
If you ever make public-facing PR descriptions, comments, issues, commit messages, etc., always follow these guidelines to sound less AI-y:
- don't use emojis
- don't use "—". Instead, reach for ";", ".", etc.
- don't use the pattern "It's not X, it's Y", "You're not X, you're Y", etc.
- don't use bulleted or numbered lists unless it would be nonsensical not to. Instead, prefer prose
### Running Scripts
- `uv run python script.py` - Run Python scripts (use for non-test files)
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
### GitHub Issue & PR Templates
When contributing to the project, use the appropriate templates:
Run tests, format your code, and lint your code before each commit
**Bug Reports** (`.github/ISSUE_TEMPLATE/bug_report.yml`):
- Describe what happened vs. what you expected
- Include relevant log output
- Specify your LiteLLM version
Ask to commit and push your work when you're done (or if you're confident that your code is good and works, just do it)
**Feature Requests** (`.github/ISSUE_TEMPLATE/feature_request.yml`):
- Describe the feature clearly
- Explain the motivation and use case
When you must use real LLM models to, for example, write e2e tests, write a QA runbook, etc., make sure to use the latest models (doesn't have to be smartest, can also be a modern small, fast one. No strong preference for smart vs fast here, just use something modern) as of the year and month of the current date. Do a web search as necessary to figure that out
**Pull Requests** (`.github/pull_request_template.md`):
- Add at least 1 test in `tests/litellm/`
- Ensure `make test-unit` passes
If you're an internal contributor, when creating a new PR, the typical flow is to branch off litellm_internal_staging and create a branch prefixed with litellm_. Do not create a branch prefixed with claude/ and generally do not have / in your branch names
## Architecture Overview
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. 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
LiteLLM is a unified interface for 100+ LLM providers with two main components:
When working on a PR, keep the PR description in sync with new commits being made
### Core Library (`litellm/`)
- **Main entry point**: `litellm/main.py` - Contains core completion() function
- **Provider implementations**: `litellm/llms/` - Each provider has its own subdirectory
- **Router system**: `litellm/router.py` + `litellm/router_utils/` - Load balancing and fallback logic
- **Type definitions**: `litellm/types/` - Pydantic models and type hints
- **Integrations**: `litellm/integrations/` - Third-party observability, caching, logging
- **Caching**: `litellm/caching/` - Multiple cache backends (Redis, in-memory, S3, etc.)
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
### Proxy Server (`litellm/proxy/`)
- **Main server**: `proxy_server.py` - FastAPI application
- **Authentication**: `auth/` - API key management, JWT, OAuth2
- **Database**: `db/` - Prisma ORM with PostgreSQL/SQLite support
- **Management endpoints**: `management_endpoints/` - Admin APIs for keys, teams, models
- **Pass-through endpoints**: `pass_through_endpoints/` - Provider-specific API forwarding
- **Guardrails**: `guardrails/` - Safety and content filtering hooks
- **UI Dashboard**: Served from `_experimental/out/` (Next.js build)
Do not put names of customers or customer company names in code, PRs, and issues. The codebase is public
## Key Patterns
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
### Provider Implementation
- Providers inherit from base classes in `litellm/llms/base.py`
- Each provider has transformation functions for input/output formatting
- Support both sync and async operations
- Handle streaming responses and function calling
## Think Before Coding
### Error Handling
- Provider-specific exceptions mapped to OpenAI-compatible errors
- Fallback logic handled by Router system
- Comprehensive logging through `litellm/_logging.py`
**Don't assume. Don't hide confusion. Surface tradeoffs.**
### Configuration
- YAML config files for proxy server (see `proxy/example_config_yaml/`)
- Environment variables for API keys and settings
- Database schema managed via Prisma (`proxy/schema.prisma`)
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.
## Development Notes
## Simplicity First
### Code Style
- Uses Black formatter, Ruff linter, MyPy type checker
- Pydantic v2 for data validation
- Async/await patterns throughout
- Type hints required for all public APIs
- **Avoid imports within methods** — place all imports at the top of the file (module-level). Inline imports inside functions/methods make dependencies harder to trace and hurt readability. The only exception is avoiding circular imports where absolutely necessary.
- **Use dict spread for immutable copies** — prefer `{**original, "key": new_value}` over `dict(obj)` + mutation. The spread produces the final dict in one step and makes intent clear.
- **Guard at resolution time** — when resolving an optional value through a fallback chain (`a or b or ""`), raise immediately if the resolved result being empty is an error. Don't pass empty strings or sentinel values downstream for the callee to deal with.
- **Extract complex comprehensions to named helpers** — a set/dict comprehension that calls into the DB or manager (e.g. "which of these server IDs are OAuth2?") belongs in a named helper function, not inline in the caller.
- **FastAPI parameter declarations** — mark required query/form params with `= Query(...)` / `= Form(...)` explicitly when other params in the same handler are optional. Mixing `str` (required) with `Optional[str] = None` in the same signature causes silent 422s when the required param is missing.
**Minimum code that solves the problem. Nothing speculative.**
### Testing Strategy
- Unit tests in `tests/test_litellm/`
- Integration tests for each provider in `tests/llm_translation/`
- Proxy tests in `tests/proxy_unit_tests/`
- Load tests in `tests/load_tests/`
- **Always add tests when adding new entity types or features** — if the existing test file covers other entity types, add corresponding tests for the new one
- **Keep monkeypatch stubs in sync with real signatures** — when a function gains a new optional parameter, update every `fake_*` / `stub_*` in tests that patch it to also accept that kwarg (even as `**kwargs`). Stale stubs fail with `unexpected keyword argument` and mask real bugs.
- **Test all branches of name→ID resolution** — when adding server/resource lookup that resolves names to UUIDs, test: (1) name resolves and UUID is allowed, (2) name resolves but UUID is not allowed, (3) name does not resolve at all. The silent-fallback path is where access-control bugs hide.
- 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.
### UI / Backend Consistency
- When wiring a new UI entity type to an existing backend endpoint, verify the backend API contract (single value vs. array, required vs. optional params) and ensure the UI controls match — e.g., use a single-select dropdown when the backend accepts a single value, not a multi-select
### UI Component Library
- **Always use `antd` for new UI components** — we are migrating off of `@tremor/react`. Do not introduce new `Badge`, `Text`, `Card`, `Grid`, `Title`, or other imports from `@tremor/react` in any new or modified file. Use `antd` equivalents: `Tag` for labels, `Typography.Text` / `Typography.Title` / `Typography.Paragraph` for textual content (avoid plain text-only `<span>`, `<p>`, `<h*>` when Typography fits), and `Card` from `antd`. Note that `antd` has no `"yellow"` Tag color — use `"gold"` for amber/yellow.
### MCP OAuth / OpenAPI Transport Mapping
- **`available_on_public_internet: false` with `delegate_auth_to_upstream: true` (oauth2, interactive — not `client_credentials`)** — LiteLLM still allows the anonymous upstream PKCE path (no proxy API key for `/authorize` and matching MCP routes). The internal-only flag mainly affects other surfaces (e.g. IP-based discovery). Rely on the upstream IdP and network policy; the dashboard shows a warning when both are set, and the proxy logs a warning when the server is loaded from config or the database.
- `TRANSPORT.OPENAPI` is a UI-only concept. The backend only accepts `"http"`, `"sse"`, or `"stdio"`. Always map it to `"http"` before any API call (including pre-OAuth temp-session calls).
- FastAPI validation errors return `detail` as an array of `{loc, msg, type}` objects. Error extractors must handle: array (map `.msg`), string, nested `{error: string}`, and fallback.
- When an MCP server already has `authorization_url` stored, skip OAuth discovery (`_discovery_metadata`) — the server URL for OpenAPI MCPs is the spec file, not the API base, and fetching it causes timeouts.
- `client_id` should be optional in the `/authorize` endpoint — if the server has a stored `client_id` in credentials, use that. Never require callers to re-supply it.
### MCP Credential Storage
- OAuth credentials and BYOK credentials share the `litellm_mcpusercredentials` table, distinguished by a `"type"` field in the JSON payload (`"oauth2"` vs plain string).
- When deleting OAuth credentials, check type before deleting to avoid accidentally deleting a BYOK credential for the same `(user_id, server_id)` pair.
- Always pass the raw `expires_at` timestamp to the client — never set it to `None` for expired credentials. Let the frontend compute the "Expired" display state from the timestamp.
- Use `RecordNotFoundError` (not bare `except Exception`) when catching "already deleted" in credential delete endpoints.
### Browser Storage Safety (UI)
- Never write LiteLLM access tokens or API keys to `localStorage` — use `sessionStorage` only. `localStorage` survives browser close and is readable by any injected script (XSS).
- Shared utility functions (e.g. `extractErrorMessage`) belong in `src/utils/` — never define them inline in hooks or duplicate them across files.
### Database Migrations
- Prisma handles schema migrations
- Migration files auto-generated with `prisma migrate dev`
- Always test migrations against both PostgreSQL and SQLite
### Proxy database access
- **Do not write raw SQL** for proxy DB operations. Use Prisma model methods instead of `execute_raw` / `query_raw`.
- Use the generated client: `prisma_client.db.<model>` (e.g. `litellm_tooltable`, `litellm_usertable`) with `.upsert()`, `.find_many()`, `.find_unique()`, `.update()`, `.update_many()` as appropriate. This avoids schema/client drift, keeps code testable with simple mocks, and matches patterns used in spend logs and other proxy code.
- **No N+1 queries.** Never query the DB inside a loop. Batch-fetch with `{"in": ids}` and distribute in-memory.
- **Batch writes.** Use `create_many`/`update_many`/`delete_many` instead of individual calls (these return counts only; `update_many`/`delete_many` no-op silently on missing rows). When multiple separate writes target the same table (e.g. in `batch_()`), order by primary key to avoid deadlocks.
- **Push work to the DB.** Filter, sort, group, and aggregate in SQL, not Python. Verify Prisma generates the expected SQL — e.g. prefer `group_by` over `find_many(distinct=...)` which does client-side processing.
- **Bound large result sets.** Prisma materializes full results in memory. For results over ~10 MB, paginate with `take`/`skip` or `cursor`/`take`, always with an explicit `order`. Prefer cursor-based pagination (`skip` is O(n)). Don't paginate naturally small result sets.
- **Limit fetched columns on wide tables.** Use `select` to fetch only needed fields — returns a partial object, so downstream code must not access unselected fields.
- **Check index coverage.** For new or modified queries, check `schema.prisma` for a supporting index. Prefer extending an existing index (e.g. `@@index([a])` → `@@index([a, b])`) over adding a new one, unless it's a `@@unique`. Only add indexes for large/frequent queries.
- **Keep schema files in sync.** Apply schema changes to all `schema.prisma` copies (`schema.prisma`, `litellm/proxy/`, `litellm-proxy-extras/`) with a migration under `litellm-proxy-extras/litellm_proxy_extras/migrations/`.
### Setup Wizard (`litellm/setup_wizard.py`)
- The wizard is implemented as a single `SetupWizard` class with `@staticmethod` methods — keep it that way. No module-level functions except `run_setup_wizard()` (the public entrypoint) and pure helpers (color, ANSI).
- Use `litellm.utils.check_valid_key(model, api_key)` for credential validation — never roll a custom completion call.
- Do not hardcode provider env-key names or model lists that already exist in the codebase. Add a `test_model` field to each provider entry to drive `check_valid_key`; set it to `None` for providers that can't be validated with a single API key (Azure, Bedrock, Ollama).
### Enterprise Features
- Enterprise-specific code in `enterprise/` directory
- Optional features enabled via environment variables
- Separate licensing and authentication for enterprise features
### 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`). Unversioned downloads silently change under you.
- **Verify checksums for all downloaded binaries.** Use the provider's official `.sha256` / `.sha256sum` sidecar file when available; otherwise compute and hardcode the digest.
- **Prefer reusable CircleCI commands** (`commands:` section) so a tool is installed and verified in exactly one place, then referenced everywhere with `- install_<tool>` or `- wait_for_service`.
- **Don't add tools just because they were there before.** Audit whether an external dependency is still needed. If it can be replaced with a shell one-liner or a tool already in the image, remove it.
- These rules apply to every download in CI: binaries, install scripts, language version managers, package repos. No exceptions.
### HTTP Client Cache Safety
- **Never close HTTP/SDK clients on cache eviction.** `LLMClientCache._remove_key()` must not call `close()`/`aclose()` on evicted clients — they may still be used by in-flight requests. Doing so causes `RuntimeError: Cannot send a request, as the client has been closed.` after the 1-hour TTL expires. Cleanup happens at shutdown via `close_litellm_async_clients()`.
### Troubleshooting: DB schema out of sync after proxy restart
`litellm-proxy-extras` runs `prisma migrate deploy` on startup using **its own** bundled migration files, which may lag behind schema changes in the current worktree. Symptoms: `Unknown column`, `Invalid prisma invocation`, or missing data on new fields.
**Diagnose:** Run `\d "TableName"` in psql and compare against `schema.prisma` — missing columns confirm the issue.
**Fix options:**
1. **Create a Prisma migration** (permanent) — run `prisma migrate dev --name <description>` in the worktree. The generated file will be picked up by `prisma migrate deploy` on next startup.
2. **Apply manually for local dev** — `psql -d litellm -c "ALTER TABLE ... ADD COLUMN IF NOT EXISTS ..."` after each proxy start. Fine for dev, not for production.
3. **Update litellm-proxy-extras** — if the package is installed from PyPI, its migration directory must include the new file. Either update the package or run the migration manually until the next release ships it.
Ask yourself: "Would a senior engineer say this is overcomplicated?" If yes, simplify.

109
GEMINI.md
View file

@ -1,108 +1 @@
# GEMINI.md
This file provides guidance to Gemini when working with code in this repository.
## Development Commands
### Installation
- `make install-dev` - Install core development dependencies
- `make install-proxy-dev` - Install proxy development dependencies with full feature set
- `make install-test-deps` - Install all test dependencies
### Testing
- `make test` - Run all tests
- `make test-unit` - Run unit tests (tests/test_litellm) with 4 parallel workers
- `make test-integration` - Run integration tests (excludes unit tests)
- `pytest tests/` - Direct pytest execution
### Code Quality
- `make lint` - Run all linting (Ruff, MyPy, Black, circular imports, import safety)
- `make format` - Apply Black code formatting
- `make lint-ruff` - Run Ruff linting only
- `make lint-mypy` - Run MyPy type checking only
### Single Test Files
- `uv run pytest tests/path/to/test_file.py -v` - Run specific test file
- `uv run pytest tests/path/to/test_file.py::test_function -v` - Run specific test
### Running Scripts
- `uv run python script.py` - Run Python scripts (use for non-test files)
### GitHub Issue & PR Templates
When contributing to the project, use the appropriate templates:
**Bug Reports** (`.github/ISSUE_TEMPLATE/bug_report.yml`):
- Describe what happened vs. what you expected
- Include relevant log output
- Specify your LiteLLM version
**Feature Requests** (`.github/ISSUE_TEMPLATE/feature_request.yml`):
- Describe the feature clearly
- Explain the motivation and use case
**Pull Requests** (`.github/pull_request_template.md`):
- Add at least 1 test in `tests/litellm/`
- Ensure `make test-unit` passes
## Architecture Overview
LiteLLM is a unified interface for 100+ LLM providers with two main components:
### Core Library (`litellm/`)
- **Main entry point**: `litellm/main.py` - Contains core completion() function
- **Provider implementations**: `litellm/llms/` - Each provider has its own subdirectory
- **Router system**: `litellm/router.py` + `litellm/router_utils/` - Load balancing and fallback logic
- **Type definitions**: `litellm/types/` - Pydantic models and type hints
- **Integrations**: `litellm/integrations/` - Third-party observability, caching, logging
- **Caching**: `litellm/caching/` - Multiple cache backends (Redis, in-memory, S3, etc.)
### Proxy Server (`litellm/proxy/`)
- **Main server**: `proxy_server.py` - FastAPI application
- **Authentication**: `auth/` - API key management, JWT, OAuth2
- **Database**: `db/` - Prisma ORM with PostgreSQL/SQLite support
- **Management endpoints**: `management_endpoints/` - Admin APIs for keys, teams, models
- **Pass-through endpoints**: `pass_through_endpoints/` - Provider-specific API forwarding
- **Guardrails**: `guardrails/` - Safety and content filtering hooks
- **UI Dashboard**: Served from `_experimental/out/` (Next.js build)
## Key Patterns
### Provider Implementation
- Providers inherit from base classes in `litellm/llms/base.py`
- Each provider has transformation functions for input/output formatting
- Support both sync and async operations
- Handle streaming responses and function calling
### Error Handling
- Provider-specific exceptions mapped to OpenAI-compatible errors
- Fallback logic handled by Router system
- Comprehensive logging through `litellm/_logging.py`
### Configuration
- YAML config files for proxy server (see `proxy/example_config_yaml/`)
- Environment variables for API keys and settings
- Database schema managed via Prisma (`proxy/schema.prisma`)
## Development Notes
### Code Style
- Uses Black formatter, Ruff linter, MyPy type checker
- Pydantic v2 for data validation
- Async/await patterns throughout
- Type hints required for all public APIs
### Testing Strategy
- Unit tests in `tests/test_litellm/`
- Integration tests for each provider in `tests/llm_translation/`
- Proxy tests in `tests/proxy_unit_tests/`
- Load tests in `tests/load_tests/`
### Database Migrations
- Prisma handles schema migrations
- Migration files auto-generated with `prisma migrate dev`
- Always test migrations against both PostgreSQL and SQLite
### Enterprise Features
- Enterprise-specific code in `enterprise/` directory
- Optional features enabled via environment variables
- Separate licensing and authentication for enterprise features
Read @CLAUDE.md for coding guidelines

View file

@ -12,17 +12,27 @@ USER root
COPY --from=uvbin /uv /uvx /usr/local/bin/
RUN apk add --no-cache bash gcc python3 python3-dev openssl openssl-dev libsndfile
# nodejs/npm so `prisma generate` uses Wolfi's Node via PRISMA_USE_GLOBAL_NODE
# instead of nodeenv downloading one whose dynamic deps may not be in Wolfi
# (e.g. Node 26.2.0 needs libatomic). Retry for transient apk.cgr.dev flakes.
RUN for i in 1 2 3; do \
apk add --no-cache bash gcc python3 python3-dev openssl openssl-dev libsndfile nodejs npm && break; \
[ $i = 3 ] && { echo "apk add failed after 3 retries" >&2; exit 1; }; \
sleep 5; \
done
# UV_COMPILE_BYTECODE=1 precompiles .pyc at install time → faster cold start.
# UV_LINK_MODE=copy avoids hardlink warnings when uv installs from a
# BuildKit cache mount (different filesystem).
# UV_PYTHON_DOWNLOADS=0 force uv to use the apk-installed CPython instead of
# silently pulling a managed interpreter.
# PRISMA_USE_GLOBAL_NODE explicit (matches default) so an env override can't
# silently re-enable nodeenv's Node download.
ENV UV_PROJECT_ENVIRONMENT=/app/.venv \
UV_LINK_MODE=copy \
UV_COMPILE_BYTECODE=1 \
UV_PYTHON_DOWNLOADS=0 \
PRISMA_USE_GLOBAL_NODE=true \
PATH="/app/.venv/bin:${PATH}"
# Stage 1 — install dependencies only.
@ -58,7 +68,11 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime
USER root
RUN apk add --no-cache bash openssl tzdata python3 libsndfile libatomic
RUN for i in 1 2 3; do \
apk add --no-cache bash openssl tzdata python3 libsndfile libatomic && break; \
[ $i = 3 ] && { echo "apk add failed after 3 retries" >&2; exit 1; }; \
sleep 5; \
done
# wolfi-base ships an unprivileged `nonroot` account (UID/GID 65532) with
# /home/nonroot. We run the backend as that user

View file

@ -1,5 +1,5 @@
module github.com/BerriAI/litellm/cookbook/gollem_go_agent_framework
go 1.25.1
go 1.26.3
require github.com/fugue-labs/gollem v0.1.0

View file

@ -12,17 +12,27 @@ USER root
COPY --from=uvbin /uv /uvx /usr/local/bin/
RUN apk add --no-cache bash gcc python3 python3-dev openssl openssl-dev libsndfile
# nodejs/npm so `prisma generate` uses Wolfi's Node via PRISMA_USE_GLOBAL_NODE
# instead of nodeenv downloading one whose dynamic deps may not be in Wolfi
# (e.g. Node 26.2.0 needs libatomic). Retry for transient apk.cgr.dev flakes.
RUN for i in 1 2 3; do \
apk add --no-cache bash gcc python3 python3-dev openssl openssl-dev libsndfile nodejs npm && break; \
[ $i = 3 ] && { echo "apk add failed after 3 retries" >&2; exit 1; }; \
sleep 5; \
done
# UV_COMPILE_BYTECODE=1 precompiles .pyc at install time → faster cold start.
# UV_LINK_MODE=copy avoids hardlink warnings when uv installs from a
# BuildKit cache mount (different filesystem).
# UV_PYTHON_DOWNLOADS=0 force uv to use the apk-installed CPython instead of
# silently pulling a managed interpreter.
# PRISMA_USE_GLOBAL_NODE explicit (matches default) so an env override can't
# silently re-enable nodeenv's Node download.
ENV UV_PROJECT_ENVIRONMENT=/app/.venv \
UV_LINK_MODE=copy \
UV_COMPILE_BYTECODE=1 \
UV_PYTHON_DOWNLOADS=0 \
PRISMA_USE_GLOBAL_NODE=true \
PATH="/app/.venv/bin:${PATH}"
# Stage 1 — install dependencies only.
@ -58,7 +68,11 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime
USER root
RUN apk add --no-cache bash openssl tzdata python3 libsndfile libatomic
RUN for i in 1 2 3; do \
apk add --no-cache bash openssl tzdata python3 libsndfile libatomic && break; \
[ $i = 3 ] && { echo "apk add failed after 3 retries" >&2; exit 1; }; \
sleep 5; \
done
# wolfi-base ships an unprivileged `nonroot` account (UID/GID 65532) with
# /home/nonroot. We run the proxy as that user.

View file

@ -56,16 +56,34 @@ app.kubernetes.io/component: ui
{{- end -}}
{{/*
Shared ServiceAccount name used by all three component Deployments. When
`serviceAccount.create` is true and `serviceAccount.name` is empty, default
to the chart fullname. When `create` is false, fall back to the provided
name or the namespace's `default` SA.
Per-component ServiceAccount name helpers.
Each component (gateway, backend, ui) has its own SA config under
.Values.serviceAccounts.<component>. When `create` is true and `name` is
empty the chart defaults to "<release>-litellm-<component>". When `create`
is false the chart uses the provided name, or the namespace `default` SA.
*/}}
{{- define "litellm.serviceAccountName" -}}
{{- if .Values.serviceAccount.create -}}
{{ default (include "litellm.fullname" .) .Values.serviceAccount.name }}
{{- define "litellm.gateway.serviceAccountName" -}}
{{- if .Values.serviceAccounts.gateway.create -}}
{{ default (include "litellm.gateway.fullname" .) .Values.serviceAccounts.gateway.name }}
{{- else -}}
{{ default "default" .Values.serviceAccount.name }}
{{ default "default" .Values.serviceAccounts.gateway.name }}
{{- end -}}
{{- end -}}
{{- define "litellm.backend.serviceAccountName" -}}
{{- if .Values.serviceAccounts.backend.create -}}
{{ default (include "litellm.backend.fullname" .) .Values.serviceAccounts.backend.name }}
{{- else -}}
{{ default "default" .Values.serviceAccounts.backend.name }}
{{- end -}}
{{- end -}}
{{- define "litellm.ui.serviceAccountName" -}}
{{- if .Values.serviceAccounts.ui.create -}}
{{ default (include "litellm.ui.fullname" .) .Values.serviceAccounts.ui.name }}
{{- else -}}
{{ default "default" .Values.serviceAccounts.ui.name }}
{{- end -}}
{{- end -}}

View file

@ -19,7 +19,8 @@ spec:
labels:
{{- include "litellm.backend.selectorLabels" . | nindent 8 }}
spec:
serviceAccountName: {{ include "litellm.serviceAccountName" . }}
serviceAccountName: {{ include "litellm.backend.serviceAccountName" . }}
automountServiceAccountToken: {{ .Values.serviceAccounts.backend.automount }}
{{- with .Values.imagePullSecrets }}
imagePullSecrets:
{{- toYaml . | nindent 8 }}

View file

@ -22,7 +22,8 @@ spec:
labels:
{{- include "litellm.gateway.selectorLabels" . | nindent 8 }}
spec:
serviceAccountName: {{ include "litellm.serviceAccountName" . }}
serviceAccountName: {{ include "litellm.gateway.serviceAccountName" . }}
automountServiceAccountToken: {{ .Values.serviceAccounts.gateway.automount }}
{{- with .Values.imagePullSecrets }}
imagePullSecrets:
{{- toYaml . | nindent 8 }}

View file

@ -28,7 +28,7 @@ spec:
app.kubernetes.io/component: migrations
spec:
restartPolicy: Never
serviceAccountName: {{ include "litellm.serviceAccountName" . }}
serviceAccountName: {{ include "litellm.backend.serviceAccountName" . }}
{{- with .Values.imagePullSecrets }}
imagePullSecrets:
{{- toYaml . | nindent 8 }}

View file

@ -1,13 +1,51 @@
{{- if .Values.serviceAccount.create -}}
{{- $prev := false -}}
{{- if .Values.serviceAccounts.gateway.create -}}
{{- $prev = true }}
apiVersion: v1
kind: ServiceAccount
metadata:
name: {{ include "litellm.serviceAccountName" . }}
name: {{ include "litellm.gateway.serviceAccountName" . }}
labels:
{{- include "litellm.commonLabels" . | nindent 4 }}
{{- with .Values.serviceAccount.annotations }}
app.kubernetes.io/component: gateway
{{- with .Values.serviceAccounts.gateway.annotations }}
annotations:
{{- toYaml . | nindent 4 }}
{{- end }}
automountServiceAccountToken: {{ .Values.serviceAccount.automount }}
automountServiceAccountToken: {{ .Values.serviceAccounts.gateway.automount }}
{{- end }}
{{- if .Values.serviceAccounts.backend.create }}
{{- if $prev }}
---
{{- end }}
{{- $prev = true }}
apiVersion: v1
kind: ServiceAccount
metadata:
name: {{ include "litellm.backend.serviceAccountName" . }}
labels:
{{- include "litellm.commonLabels" . | nindent 4 }}
app.kubernetes.io/component: backend
{{- with .Values.serviceAccounts.backend.annotations }}
annotations:
{{- toYaml . | nindent 4 }}
{{- end }}
automountServiceAccountToken: {{ .Values.serviceAccounts.backend.automount }}
{{- end }}
{{- if .Values.serviceAccounts.ui.create }}
{{- if $prev }}
---
{{- end }}
apiVersion: v1
kind: ServiceAccount
metadata:
name: {{ include "litellm.ui.serviceAccountName" . }}
labels:
{{- include "litellm.commonLabels" . | nindent 4 }}
app.kubernetes.io/component: ui
{{- with .Values.serviceAccounts.ui.annotations }}
annotations:
{{- toYaml . | nindent 4 }}
{{- end }}
automountServiceAccountToken: {{ .Values.serviceAccounts.ui.automount }}
{{- end }}

View file

@ -19,7 +19,8 @@ spec:
labels:
{{- include "litellm.ui.selectorLabels" . | nindent 8 }}
spec:
serviceAccountName: {{ include "litellm.serviceAccountName" . }}
serviceAccountName: {{ include "litellm.ui.serviceAccountName" . }}
automountServiceAccountToken: {{ .Values.serviceAccounts.ui.automount }}
{{- with .Values.imagePullSecrets }}
imagePullSecrets:
{{- toYaml . | nindent 8 }}

View file

@ -14,16 +14,33 @@ ingress:
host: "" # optional; if set, becomes the rule's host
tls: []
# Shared ServiceAccount used by all three component Deployments. Set
# `create: true` to have the chart provision it (e.g. when wiring an EKS
# Pod Identity association by SA name). Set `name` to use an existing SA
# (chart-created or out-of-band). When both are empty / false, pods run
# with the namespace's `default` SA.
serviceAccount:
create: false
automount: true
annotations: {}
name: ""
# Per-component ServiceAccounts for gateway, backend, and ui.
#
# Each section mirrors the old shared serviceAccount shape. Set `create:
# true` to have the chart provision the SA (useful for EKS Pod Identity /
# GKE Workload Identity annotations). Set `name` to bind an existing SA.
# When both are unset the component pod runs with the namespace `default` SA.
#
# The UI SA deliberately defaults to `automount: false` — the static nginx
# container does not need the K8s API and should not carry a projected
# ServiceAccount token that a compromised container could use to call the
# cloud-provider metadata service or the K8s API.
serviceAccounts:
gateway:
create: false
automount: true
annotations: {}
name: ""
backend:
create: false
automount: true
annotations: {}
name: ""
ui:
create: false
automount: false
annotations: {}
name: ""
# Pre-install / pre-upgrade Helm hook that runs `prisma migrate deploy`
# against the writer database, creating the LiteLLM schema (tables that

View file

@ -225,6 +225,11 @@ use_chat_completions_url_for_anthropic_messages: bool = bool(
route_all_chat_openai_to_responses: bool = (
os.getenv("LITELLM_ROUTE_ALL_CHAT_OPENAI_TO_RESPONSES", "false").lower() == "true"
) # When True, routes all OpenAI /chat/completions requests through the Responses API bridge
# When True, Gemini/Vertex Live setup is deferred until client `session.update`.
# Default False preserves historical behavior (auto-send setup on connect).
gemini_live_defer_setup: bool = (
os.getenv("LITELLM_GEMINI_LIVE_DEFER_SETUP", "false").lower() == "true"
)
use_legacy_interactions_schema: bool = (
os.getenv("LITELLM_USE_LEGACY_INTERACTIONS_SCHEMA", "false").lower() == "true"
) # When True, sends Api-Revision: 2026-05-07 to Google so responses use the legacy `outputs`

View file

@ -1147,6 +1147,7 @@ BEDROCK_CONVERSE_MODELS = [
"openai.gpt-oss-120b-1:0",
"anthropic.claude-haiku-4-5-20251001-v1:0",
"anthropic.claude-sonnet-4-5-20250929-v1:0",
"anthropic.claude-opus-4-8",
"anthropic.claude-opus-4-7",
"anthropic.claude-opus-4-6-v1:0",
"anthropic.claude-opus-4-6-v1",

View file

@ -133,6 +133,8 @@ _VIDEO_CALL_TYPES = frozenset(
{
CallTypes.create_video.value,
CallTypes.acreate_video.value,
CallTypes.video_edit.value,
CallTypes.avideo_edit.value,
CallTypes.video_remix.value,
CallTypes.avideo_remix.value,
}

View file

@ -2,10 +2,17 @@ import asyncio
import os
import time
from datetime import datetime
from typing import Dict, List, Optional, Tuple
from typing import Any, Dict, List, Optional, Tuple, cast
from litellm._logging import verbose_logger
from litellm.integrations.custom_batch_logger import CustomBatchLogger
from litellm.integrations.datadog.datadog_handler import (
get_datadog_env,
get_datadog_hostname,
get_datadog_pod_name,
get_datadog_service,
)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
@ -15,9 +22,30 @@ from litellm.types.integrations.datadog_cost_management import (
)
from litellm.types.utils import StandardLoggingPayload
# Reserved tag keys whose values come from trusted sources (infra env, LiteLLM
# core payload fields, or proxy-controlled auth metadata). User-supplied
# request_tags / metadata cannot overwrite these, even when the key is
# allowlisted via cost_tag_keys, because that would let an authenticated caller
# spoof cost attribution (e.g. request_tags=["team:victim-team"]).
_RESERVED_TAG_KEYS: frozenset = frozenset(
{
"env",
"service",
"host",
"pod_name",
"provider",
"model",
"model_id",
"team",
"user",
"model_group",
}
)
class DatadogCostManagementLogger(CustomBatchLogger):
def __init__(self, **kwargs):
def __init__(self, cost_tag_keys: Optional[List[str]] = None, **kwargs):
self.cost_tag_keys: List[str] = list(cost_tag_keys) if cost_tag_keys else []
self.dd_api_key = os.getenv("DD_API_KEY")
self.dd_app_key = os.getenv("DD_APP_KEY")
self.dd_site = os.getenv("DD_SITE", "datadoghq.com")
@ -68,20 +96,21 @@ class DatadogCostManagementLogger(CustomBatchLogger):
if not self.log_queue:
return
batch_to_send = self.log_queue[:]
self.log_queue = []
try:
# Aggregate costs from the batch
aggregated_entries = self._aggregate_costs(self.log_queue)
aggregated_entries = self._aggregate_costs(batch_to_send)
if not aggregated_entries:
verbose_logger.debug(
"Datadog Cost Management: batch produced no aggregable entries; "
"dropping %d log(s) from queue.",
len(batch_to_send),
)
return
# Send to Datadog
await self._upload_to_datadog(aggregated_entries)
# Clear queue only on success (or if we decide to drop on failure)
# CustomBatchLogger clears queue in flush_queue, so we just process here
except Exception as e:
self.log_queue = batch_to_send + self.log_queue
verbose_logger.exception(
f"Datadog Cost Management: Error in async_send_batch: {str(e)}"
)
@ -151,45 +180,81 @@ class DatadogCostManagementLogger(CustomBatchLogger):
return list(aggregator.values())
def _extract_tags(self, log: StandardLoggingPayload) -> Dict[str, str]:
from litellm.integrations.datadog.datadog_handler import (
get_datadog_env,
get_datadog_hostname,
get_datadog_pod_name,
get_datadog_service,
)
tags = {
tags: Dict[str, str] = {
"env": get_datadog_env(),
"service": get_datadog_service(),
"host": get_datadog_hostname(),
"pod_name": get_datadog_pod_name(),
}
# Add metadata as tags
metadata = log.get("metadata", {})
if metadata:
# Add user info
# Add user info
if metadata.get("user_api_key_alias"):
tags["user"] = str(metadata["user_api_key_alias"])
# Always-on canonical FOCUS dimensions from top-level payload fields.
# Non-sensitive and required for Datadog Custom Costs per-model attribution.
self._add_tag(tags, "provider", log.get("custom_llm_provider"))
self._add_tag(tags, "model", log.get("model"))
self._add_tag(tags, "model_id", log.get("model_id"))
# Add Team Tag
team_tag = (
metadata.get("user_api_key_team_alias")
or metadata.get("team_alias") # type: ignore
or metadata.get("user_api_key_team_id")
or metadata.get("team_id") # type: ignore
)
# cast because StandardLoggingMetadata is a TypedDict; we iterate it
# as a generic mapping below.
metadata: Dict[str, Any] = cast(Dict[str, Any], log.get("metadata") or {})
if team_tag:
tags["team"] = str(team_tag)
# model_group is not in StandardLoggingMetadata TypedDict, so we need to access it via dict.get()
model_group = metadata.get("model_group") # type: ignore[misc]
if model_group:
tags["model_group"] = str(model_group)
# Backwards-compat: team/user/model_group preserved regardless of allowlist.
if metadata.get("user_api_key_alias"):
tags["user"] = str(metadata["user_api_key_alias"])
team_tag = (
metadata.get("user_api_key_team_alias")
or metadata.get("team_alias")
or metadata.get("user_api_key_team_id")
or metadata.get("team_id")
)
if team_tag:
tags["team"] = str(team_tag)
if metadata.get("model_group"):
tags["model_group"] = str(metadata["model_group"])
# Allowlist-gated: request_tags (split on `:`) and arbitrary metadata.*.
# Reserved keys are hard-blocked here regardless of allowlist membership —
# see _RESERVED_TAG_KEYS for the rationale.
if self.cost_tag_keys:
allow = set(self.cost_tag_keys)
for rt in log.get("request_tags") or []:
if not isinstance(rt, str) or ":" not in rt:
continue
k, _, v = rt.partition(":")
if k in allow and v:
self._set_custom_tag(tags, k, v)
for k, v in metadata.items():
if k in allow and v is not None and not isinstance(v, (dict, list)):
self._set_custom_tag(tags, k, str(v))
for nested_key in ("spend_logs_metadata", "requester_metadata"):
nested = metadata.get(nested_key)
if isinstance(nested, dict):
for k, v in nested.items():
if (
k in allow
and v is not None
and not isinstance(v, (dict, list))
):
self._set_custom_tag(tags, k, str(v))
return tags
@staticmethod
def _set_custom_tag(tags: Dict[str, str], key: str, value: str) -> None:
if key in _RESERVED_TAG_KEYS:
verbose_logger.debug(
"Datadog Cost Management: dropping user-supplied tag %r=%r — "
"key is reserved for trusted cost attribution.",
key,
value,
)
return
tags[key] = value
@staticmethod
def _add_tag(tags: Dict[str, str], key: str, value: Any) -> None:
if value:
tags[key] = str(value)
async def _upload_to_datadog(self, payload: List[Dict]):
if not self.dd_api_key or not self.dd_app_key:
return
@ -201,8 +266,6 @@ class DatadogCostManagementLogger(CustomBatchLogger):
}
# The API endpoint expects a list of objects directly in the body (file content behavior)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
data_json = safe_dumps(payload)
response = await self.async_client.put(

View file

@ -1,18 +1,29 @@
import json
import os
from typing import Any, Dict, List, Optional
import re
from typing import Any, Dict, List, Optional, Tuple, cast
from pydantic import BaseModel, Field
import litellm
from litellm._logging import verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.prompt_templates.common_utils import (
convert_content_list_to_str,
get_content_from_model_response,
)
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.types.llms.openai import AllMessageValues
GALILEO_CLOUD_API_BASE_URL = "https://api.galileo.ai"
# Cap the in-memory buffer so persistent flush failures (e.g. Galileo
# unavailable, invalid credentials) cannot leak memory unboundedly.
GALILEO_MAX_IN_MEMORY_RECORDS = 1000
# from here: https://docs.rungalileo.io/galileo/gen-ai-studio-products/galileo-observe/how-to/logging-data-via-restful-apis#structuring-your-records
class LLMResponse(BaseModel):
latency_ms: int
status_code: int
@ -37,65 +48,190 @@ class GalileoObserve(CustomLogger):
def __init__(self) -> None:
self.in_memory_records: List[dict] = []
self.batch_size = 1
self.base_url = os.getenv("GALILEO_BASE_URL", None)
self.project_id = os.getenv("GALILEO_PROJECT_ID", None)
self.api_key = os.getenv("GALILEO_API_KEY")
self.project_id = os.getenv("GALILEO_PROJECT_ID")
self.log_stream_id = os.getenv("GALILEO_LOG_STREAM_ID")
self.username = os.getenv("GALILEO_USERNAME")
self.password = os.getenv("GALILEO_PASSWORD")
self.base_url = self._normalize_base_url(os.getenv("GALILEO_BASE_URL"))
if self.api_key and not self.base_url:
self.base_url = GALILEO_CLOUD_API_BASE_URL
self.use_v2_api = bool(self.api_key)
self.headers: Optional[Dict[str, str]] = None
self.async_httpx_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.LoggingCallback
)
pass
def set_galileo_headers(self):
# following https://docs.rungalileo.io/galileo/gen-ai-studio-products/galileo-observe/how-to/logging-data-via-restful-apis#logging-your-records
@staticmethod
def _normalize_base_url(base_url: Optional[str]) -> Optional[str]:
if base_url:
return base_url.rstrip("/")
return None
headers = {
"accept": "application/json",
"Content-Type": "application/x-www-form-urlencoded",
}
galileo_login_response = litellm.module_level_client.post(
def _is_configured(self) -> bool:
if not self.project_id or not self.base_url:
return False
if self.use_v2_api:
return bool(self.api_key)
return bool(self.username and self.password)
async def async_set_galileo_headers(self) -> None:
galileo_login_response = await self.async_httpx_handler.post(
url=f"{self.base_url}/login",
headers=headers,
headers={
"accept": "application/json",
"Content-Type": "application/x-www-form-urlencoded",
},
data={
"username": os.getenv("GALILEO_USERNAME"),
"password": os.getenv("GALILEO_PASSWORD"),
"username": self.username,
"password": self.password,
},
)
galileo_login_response.raise_for_status()
access_token = galileo_login_response.json()["access_token"]
self.headers = {
"accept": "application/json",
"Content-Type": "application/json",
"Authorization": f"Bearer {access_token}",
}
def get_output_str_from_response(self, response_obj, kwargs):
output = None
if response_obj is not None and (
kwargs.get("call_type", None) == "embedding"
or isinstance(response_obj, litellm.EmbeddingResponse)
):
output = None
elif response_obj is not None and isinstance(
response_obj, litellm.ModelResponse
):
output = response_obj["choices"][0]["message"].json()
elif response_obj is not None and isinstance(
response_obj, litellm.TextCompletionResponse
):
output = response_obj.choices[0].text
elif response_obj is not None and isinstance(
response_obj, litellm.ImageResponse
):
output = response_obj["data"]
async def _ensure_headers(self) -> bool:
if self.headers is not None:
return True
return output
if self.use_v2_api:
if not self.api_key:
return False
self.headers = {
"accept": "application/json",
"Content-Type": "application/json",
"Galileo-API-Key": self.api_key,
}
return True
if not (self.username and self.password and self.base_url):
return False
try:
await self.async_set_galileo_headers()
return True
except Exception as e:
verbose_logger.debug("Galileo Logger: failed to authenticate: %s", e)
return False
@staticmethod
def _galileo_input_messages(
messages: Optional[List[Any]], input_text: str
) -> List[Dict[str, str]]:
if not messages:
return [{"role": "user", "content": input_text}]
galileo_messages: List[Dict[str, str]] = []
for message in messages:
if not isinstance(message, dict):
continue
role = message.get("role")
if not role:
continue
galileo_messages.append(
{
"role": str(role),
"content": convert_content_list_to_str(
message=cast(AllMessageValues, message)
),
}
)
if galileo_messages:
return galileo_messages
return [{"role": "user", "content": input_text}]
@staticmethod
def _record_to_v2_span(record: Dict[str, Any]) -> Dict[str, Any]:
created_at = record.get("created_at", "")
if created_at and not re.search(r"(Z|[+-]\d{2}:?\d{2})$", created_at):
created_at = f"{created_at}Z"
span: Dict[str, Any] = {
"type": "llm",
"name": record.get("node_type", "litellm"),
"created_at": created_at,
"input": GalileoObserve._galileo_input_messages(
record.get("messages"), record.get("input_text", "")
),
"output": {
"role": "assistant",
"content": record.get("output_text", ""),
},
"status_code": record.get("status_code", 200),
"model": record.get("model"),
"metrics": {
"duration_ns": int(record.get("latency_ms", 0)) * 1_000_000,
"num_input_tokens": record.get("num_input_tokens"),
"num_output_tokens": record.get("num_output_tokens"),
},
}
if record.get("tags"):
span["tags"] = record["tags"]
return span
def _get_ingest_request(self) -> Optional[Tuple[str, Dict[str, Any]]]:
if not self.base_url or not self.project_id:
return None
# Snapshot the records to be sent into a new list so concurrent appends
# during the network round-trip (across the await points in
# flush_in_memory_records) aren't silently dropped when we later clear
# the in-memory buffer.
records = list(self.in_memory_records)
if self.use_v2_api:
payload: Dict[str, Any] = {
"spans": [self._record_to_v2_span(record) for record in records],
"reliable": False,
}
if self.log_stream_id:
payload["log_stream_id"] = self.log_stream_id
return (
f"{self.base_url}/v2/projects/{self.project_id}/spans",
payload,
)
return (
f"{self.base_url}/projects/{self.project_id}/observe/ingest",
{"records": records},
)
def get_output_str_from_response(
self, response_obj: Any, kwargs: Dict[str, Any]
) -> Optional[str]:
if response_obj is None:
return None
if kwargs.get("call_type", None) == "embedding" or isinstance(
response_obj, litellm.EmbeddingResponse
):
return None
if isinstance(response_obj, litellm.TextCompletionResponse):
return response_obj.choices[0].text
if isinstance(response_obj, litellm.ImageResponse):
return json.dumps(response_obj["data"], default=str)
if isinstance(response_obj, (litellm.ModelResponse, dict)):
return get_content_from_model_response(response_obj)
return None
async def async_log_success_event(
self, kwargs: Any, response_obj: Any, start_time: Any, end_time: Any
):
verbose_logger.debug("On Async Success")
if not self._is_configured():
verbose_logger.debug(
"Galileo Logger: skipping flush — set GALILEO_PROJECT_ID and "
"either GALILEO_API_KEY (hosted) or GALILEO_USERNAME/GALILEO_PASSWORD "
"(enterprise Observe)."
)
return
_latency_ms = int((end_time - start_time).total_seconds() * 1000)
_call_type = kwargs.get("call_type", "litellm")
input_text = litellm.utils.get_formatted_prompt(
@ -125,26 +261,69 @@ class GalileoObserve(CustomLogger):
), # timestamp str constructed in "%Y-%m-%dT%H:%M:%S" format
)
# dump to dict
request_dict = request_record.model_dump()
messages = kwargs.get("messages")
if messages:
request_dict["messages"] = messages
self.in_memory_records.append(request_dict)
# Bound the buffer so persistent flush failures cannot grow it
# without limit. Drop the oldest records once we exceed the cap.
if len(self.in_memory_records) > GALILEO_MAX_IN_MEMORY_RECORDS:
dropped = len(self.in_memory_records) - GALILEO_MAX_IN_MEMORY_RECORDS
self.in_memory_records = self.in_memory_records[
-GALILEO_MAX_IN_MEMORY_RECORDS:
]
verbose_logger.warning(
"Galileo Logger: in-memory buffer exceeded %s records; "
"dropped %s oldest record(s). Check Galileo connectivity/credentials.",
GALILEO_MAX_IN_MEMORY_RECORDS,
dropped,
)
if len(self.in_memory_records) >= self.batch_size:
await self.flush_in_memory_records()
async def flush_in_memory_records(self):
verbose_logger.debug("flushing in memory records")
response = await self.async_httpx_handler.post(
url=f"{self.base_url}/projects/{self.project_id}/observe/ingest",
headers=self.headers,
json={"records": self.in_memory_records},
)
if not self.in_memory_records:
return
if response.status_code == 200:
# Capture the number of records that will be sent BEFORE any await so
# that concurrent appends made by other asyncio tasks during the
# network round-trip aren't silently dropped on the success-clear.
records_in_payload = len(self.in_memory_records)
ingest_request = self._get_ingest_request()
if ingest_request is None:
verbose_logger.debug(
"Galileo Logger:successfully flushed in memory records"
"Galileo Logger: missing GALILEO_BASE_URL or GALILEO_PROJECT_ID"
)
self.in_memory_records = []
return
if not await self._ensure_headers():
verbose_logger.debug("Galileo Logger: could not set request headers")
return
url, payload = ingest_request
verbose_logger.debug("flushing in memory records to %s", url)
try:
response = await self.async_httpx_handler.post(
url=url,
headers=self.headers,
json=payload,
)
except Exception as e:
verbose_logger.debug(
"Galileo Logger: failed to flush in memory records: %s", e
)
return
if response.is_success:
verbose_logger.debug(
"Galileo Logger: successfully flushed in memory records"
)
del self.in_memory_records[:records_in_payload]
else:
verbose_logger.debug("Galileo Logger: failed to flush in memory records")
verbose_logger.debug(
@ -152,6 +331,13 @@ class GalileoObserve(CustomLogger):
response.text,
response.status_code,
)
# Legacy enterprise auth caches a bearer token obtained from
# /login. If the request was rejected for auth reasons, drop the
# cached headers so the next flush re-authenticates instead of
# silently failing forever on a stale token. The v2 API key path
# uses a long-lived static key, so leave its headers in place.
if not self.use_v2_api and response.status_code in (401, 403):
self.headers = None
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
verbose_logger.debug("On Async Failure")

View file

@ -86,6 +86,12 @@ class RealTimeStreaming:
# When a text message is blocked, hold the guardrail reason so the next
# response.create can be rewritten to include the failure context.
self._pending_guardrail_message: Optional[str] = None
# Track whether session.created has already been sent to the client
# (e.g. synthetic event in deferred setup mode).
self._session_created_sent_to_client: bool = False
# Track whether we have already sent the guardrail turn-detection update
# that disables provider auto-response for transcription guardrails.
self._guardrail_turn_detection_update_sent: bool = False
_SESSION_EVENT_TYPES = frozenset(["session.created", "session.updated"])
_AUDIO_FORMAT_MAP: Dict[str, Dict[str, Any]] = {
@ -248,22 +254,52 @@ class RealTimeStreaming:
## SYNC LOGGING
executor.submit(self.logging_obj.success_handler(self.messages))
async def _send_to_backend(self, message: str) -> None:
async def _send_to_backend(self, message: str) -> bool:
"""Send a message to the backend WebSocket.
If a provider_config is set the message is first passed through
transform_realtime_request so that provider-specific translation
(e.g. dropping session.update for Vertex AI) is applied even for
guardrail-injected messages.
Returns True if at least one message was actually delivered to the
backend, False if the provider transformation produced no output and
the message was effectively dropped.
"""
if self.provider_config:
transformed = self.provider_config.transform_realtime_request(
message, self.model, self.session_configuration_request
)
sent = False
for msg in transformed:
# Send first; only cache the setup payload once the backend
# has actually accepted it. Caching before send would leave
# ``session_configuration_request`` populated after a failed
# send, causing subsequent client session.update messages to
# be treated as "subsequent" and dropped even though the
# backend never received the original setup.
await self.backend_ws.send(msg) # type: ignore[union-attr, attr-defined]
else:
await self.backend_ws.send(message) # type: ignore[union-attr, attr-defined]
self._cache_session_configuration_request(msg)
sent = True
return sent
await self.backend_ws.send(message) # type: ignore[union-attr, attr-defined]
return True
def _cache_session_configuration_request(self, transformed_message: str) -> None:
"""Store setup payload once sent to backend.
Updates the cached setup on every successful setup send so follow-up
``session.update`` messages (which produce a merged setup with new
``generationConfig`` / ``systemInstruction`` / etc.) are reflected in
the cache used by downstream readers (``transform_session_created_event``,
``return_new_content_delta_events`` modality lookup, ...).
"""
try:
message_obj = json.loads(transformed_message)
if "setup" in message_obj:
self.session_configuration_request = transformed_message
except (json.JSONDecodeError, TypeError):
return
def _make_disable_auto_response_message(self) -> str:
"""Return a session.update that disables VAD auto-response."""
@ -280,6 +316,20 @@ class RealTimeStreaming:
}
return json.dumps({"type": "session.update", "session": session})
async def _maybe_send_guardrail_turn_detection_update(self) -> None:
"""Disable provider auto-response once when transcription guardrails are enabled."""
if self._guardrail_turn_detection_update_sent:
return
if not self._has_audio_transcription_guardrails():
return
sent = await self._send_to_backend(self._make_disable_auto_response_message())
# Only mark as sent when the provider transformation actually delivered
# the update to the backend. Otherwise (e.g. Gemini drops session.update
# after the initial setup), leave the flag unset so future opportunities
# — such as a duplicate session.created — can retry.
if sent:
self._guardrail_turn_detection_update_sent = True
def _has_realtime_guardrails(self) -> bool:
"""Return True if any callback is registered for realtime guardrail event types."""
from litellm.integrations.custom_guardrail import CustomGuardrail
@ -318,12 +368,20 @@ class RealTimeStreaming:
self,
transcript: str,
item_id: Optional[str] = None,
pre_block_backend_message: Optional[str] = None,
) -> bool:
"""
Run registered guardrails on a completed speech transcription.
Returns True if blocked (synthetic warning already sent to client).
Returns False if clean (caller should send response.create to the backend).
``pre_block_backend_message`` (if provided) is sent to the backend
BEFORE any of the guardrail's own backend messages when a block is
triggered. This is needed for protocol contracts that require a
specific message to be sent first — e.g. Gemini Live requires a
matching ``toolResponse`` immediately after a ``toolCall`` before any
other client messages can be accepted.
"""
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.types.guardrails import GuardrailEventHooks
@ -383,6 +441,13 @@ class RealTimeStreaming:
getattr(callback, "realtime_violation_message", None) or safe_msg
)
# Deliver any caller-supplied backend message FIRST so that
# protocol contracts requiring a specific ordering (e.g.
# Gemini Live's mandatory ``toolResponse`` after a
# ``toolCall``) are honored before the guardrail's own
# clientContent / cancel messages are sent.
if pre_block_backend_message is not None:
await self._send_to_backend(pre_block_backend_message)
# Cancel any in-progress LLM response (e.g. VAD auto-response).
await self._send_to_backend(json.dumps({"type": "response.cancel"}))
# Send the policy violation hint (shows as small gray status text in UI).
@ -478,16 +543,34 @@ class RealTimeStreaming:
else [transformed_response]
)
for event in events:
is_session_created_event = (
isinstance(event, dict) and event.get("type") == "session.created"
)
if is_session_created_event:
if self._session_created_sent_to_client:
# A synthetic session.created (with placeholder defaults) was
# already forwarded to the client when we connected. The
# provider's real session.created (e.g. emitted from Gemini
# `setupComplete`) carries the authoritative modalities/model
# from the client's session.update. Re-emit it as
# `session.updated` so the client learns the corrected
# configuration without seeing two `session.created` events.
event = {**event, "type": "session.updated"}
else:
self._session_created_sent_to_client = True
event_str = json.dumps(event)
## For audio/VAD guardrail path: forward session.created first, then inject.
if (
isinstance(event, dict)
and event.get("type") == "session.created"
and self._has_audio_transcription_guardrails()
):
## For audio/VAD guardrail path: forward the (possibly retyped)
## session.created first, then invoke the one-time guardrail
## turn-detection update. ``_maybe_send_guardrail_turn_detection_update``
## is idempotent (gated by ``_guardrail_turn_detection_update_sent``),
## so duplicate session.created events — including those emitted
## after a synthetic session.created from ``llm_http_handler`` in
## deferred-setup mode — still get a single chance to inject the
## update if a prior attempt was dropped by the provider transform.
if is_session_created_event and self._has_audio_transcription_guardrails():
self.store_message(event_str)
await self.websocket.send_text(event_str)
await self._send_to_backend(self._make_disable_auto_response_message())
await self._maybe_send_guardrail_turn_detection_update()
continue
## GUARDRAIL: run on transcription events in provider_config path too
if (
@ -790,12 +873,13 @@ class RealTimeStreaming:
item["content"] = new_content
return item
async def client_ack_messages(self):
async def client_ack_messages(self): # noqa: PLR0915
try:
while True:
message = await self.websocket.receive_text()
## GUARDRAIL: intercept conversation.item.create for text-based injection.
guardrail_turn_detection_injected = False
try:
msg_obj = json.loads(message)
msg_type = msg_obj.get("type")
@ -803,7 +887,68 @@ class RealTimeStreaming:
if msg_type == "conversation.item.create":
# Check user text messages for prompt injection
item = msg_obj.get("item", {})
if item.get("role") == "user":
# Check function_call_output first so a client cannot
# bypass the tool-result guardrail by also setting
# role="user" on a function_call_output item.
if item.get("type") == "function_call_output":
# Tool results are client-controlled and fed to the
# model; check them with the same guardrail used for
# user text so an attacker cannot smuggle blocked
# content into a function_call_output.
output = item.get("output", "")
output_text = (
output
if isinstance(output, str)
else json.dumps(output)
)
if output_text:
# Build the sanitized function_call_output up
# front so we can hand it to the guardrail
# runner as the pre-block message. Providers
# that pair every toolCall with a toolResponse
# (e.g. Gemini/Vertex Live) require the
# toolResponse to arrive BEFORE any other
# client message — otherwise the guardrail's
# own clientContent would violate the
# pending-tool-call protocol contract and the
# backend could close the connection before
# the sanitized response ever lands. Dropping
# the blocked item outright would similarly
# leave such providers waiting indefinitely.
# The sanitized payload carries no blocked
# content — only a generic policy marker.
sanitized_msg = json.dumps(
{
**msg_obj,
"item": {
**item,
"output": json.dumps(
{
"error": "Tool output blocked by content policy",
}
),
},
}
)
blocked = await self.run_realtime_guardrails(
output_text,
pre_block_backend_message=sanitized_msg,
)
if blocked:
# ``_pending_guardrail_message`` is
# intentionally NOT set here. That flag
# exists to swallow the reflexive
# ``response.create`` an OpenAI client
# sends immediately after a user text
# message. In a tool-calling flow the
# client may not send a ``response.create``
# at all (e.g. Gemini SDKs auto-respond),
# so leaving the flag set would
# incorrectly drop an unrelated
# ``response.create`` from a later
# interaction turn.
continue
elif item.get("role") == "user":
content_list = item.get("content", [])
texts = [
c.get("text", "")
@ -831,6 +976,89 @@ class RealTimeStreaming:
self._pending_guardrail_message = None
continue
## GUARDRAIL: Inject turn_detection into first session.update
# if needed. Done BEFORE the GA remap so the injected
# ``create_response`` rides along with any client-provided
# turn_detection fields (e.g. silence_duration_ms) into the
# nested ``audio.input.turn_detection`` path produced by the
# remap. Doing this after the remap would create a separate
# minimal root-level ``turn_detection`` and silently drop
# the client's nested settings.
if (
msg_type == "session.update"
and self.session_configuration_request is None
and not self._guardrail_turn_detection_update_sent
and self._has_audio_transcription_guardrails()
):
session = msg_obj.setdefault("session", {})
if isinstance(session, dict):
existing_td = session.get("turn_detection")
if not isinstance(existing_td, dict):
existing_td = {}
existing_td["create_response"] = False
session["turn_detection"] = existing_td
message = json.dumps(msg_obj)
guardrail_turn_detection_injected = True
verbose_logger.debug(
"Injected turn_detection into first session.update for audio transcription guardrails"
)
## GUARDRAIL: Force ``create_response`` to False in any
# client-provided ``turn_detection`` so a later
# ``session.update`` cannot re-enable VAD auto-response
# and bypass the transcription guardrail after the
# initial disable. Covers both the flat beta key and the
# nested GA ``audio.input.turn_detection`` shape, since
# the GA remap below also accepts either form. Skipped
# when the injection block above already ran for this
# message, to avoid redundant double-serialization.
if (
msg_type == "session.update"
and not guardrail_turn_detection_injected
and self._has_audio_transcription_guardrails()
):
session = msg_obj.get("session")
if isinstance(session, dict):
td_overridden = False
flat_td = session.get("turn_detection")
flat_td_present = flat_td is not None
if flat_td_present:
if not isinstance(flat_td, dict):
flat_td = {}
if flat_td.get("create_response") is not False:
flat_td["create_response"] = False
session["turn_detection"] = flat_td
td_overridden = True
nested_td_present = False
audio = session.get("audio")
if isinstance(audio, dict):
audio_input = audio.get("input")
if isinstance(audio_input, dict):
nested_td = audio_input.get("turn_detection")
if nested_td is not None:
nested_td_present = True
if not isinstance(nested_td, dict):
nested_td = {}
if (
nested_td.get("create_response")
is not False
):
nested_td["create_response"] = False
audio_input["turn_detection"] = nested_td
td_overridden = True
# Symmetric with the first-update injection block:
# if the client omitted turn_detection entirely on
# a subsequent session.update, still inject the
# ``create_response: False`` override so the
# transcription guardrail cannot be re-enabled by
# any downstream merge that drops the original
# disable.
if not flat_td_present and not nested_td_present:
session["turn_detection"] = {"create_response": False}
td_overridden = True
if td_overridden:
message = json.dumps(msg_obj)
# GA compatibility: remap beta-style session fields only when
# the upstream is in GA mode. Beta upstreams expect the flat
# session shape unchanged.
@ -848,17 +1076,20 @@ class RealTimeStreaming:
pass
## LOGGING
# Log after any in-place modifications (GA remap, guardrail
# turn_detection injection) so audit logs reflect what we
# actually forward to the backend.
self.store_input(message=message)
## FORWARD TO BACKEND
if self.provider_config:
message = self.provider_config.transform_realtime_request(
message, self.model
)
for msg in message:
await self.backend_ws.send(msg) # type: ignore[union-attr]
else:
await self.backend_ws.send(message) # type: ignore[union-attr]
## FORWARD TO BACKEND
# Only mark the guardrail turn_detection update as sent after the
# backend actually accepted the message. Setting the flag earlier
# would permanently disable the injection if ``_send_to_backend``
# raised — neither this loop nor
# ``_maybe_send_guardrail_turn_detection_update`` would retry.
sent = await self._send_to_backend(message)
if guardrail_turn_detection_injected and sent:
self._guardrail_turn_detection_update_sent = True
except Exception as e:
verbose_logger.debug(f"Error in client ack messages: {e}")

View file

@ -59,6 +59,8 @@ FUNCTION_CALL_ATTRIBUTE = "function_call"
_SYNC_ITER_EXHAUSTED = object()
_GCHUNK_FIELDS: frozenset = frozenset(GChunk.__annotations__)
def _next_sync_or_exhausted(it: Any) -> Any:
"""
@ -181,6 +183,30 @@ class CustomStreamWrapper:
self.created: Optional[int] = None
self._last_returned_hidden_params: Optional[dict] = None
_cached_logging_provider = self.logging_obj.model_call_details.get(
"custom_llm_provider", None
)
self._cached_logging_llm_provider: Optional[str] = _cached_logging_provider
_effective_model = model or ""
if (
custom_llm_provider == "openai"
and custom_llm_provider != _cached_logging_provider
):
_effective_model = "{}/{}".format(
_cached_logging_provider, _effective_model
)
self._cached_model_name: str = _effective_model
# Snapshot assumes self._hidden_params is populated from litellm_params
# at init and never mutated during the stream. If that ever changes,
# this cache must be removed.
self._base_hidden_params: Dict[str, Any] = {
**self._hidden_params,
"response_cost": None,
}
self._post_streaming_hooks: Optional[List] = None
def _check_max_streaming_duration(self) -> None:
"""Raise litellm.Timeout if the stream has exceeded LITELLM_MAX_STREAMING_DURATION_SECONDS."""
from litellm.constants import LITELLM_MAX_STREAMING_DURATION_SECONDS
@ -681,29 +707,16 @@ class CustomStreamWrapper:
def model_response_creator(
self, chunk: Optional[dict] = None, hidden_params: Optional[dict] = None
):
_model = self.model
_received_llm_provider = self.custom_llm_provider
_logging_obj_llm_provider = self.logging_obj.model_call_details.get("custom_llm_provider", None) # type: ignore
if (
_received_llm_provider == "openai"
and _received_llm_provider != _logging_obj_llm_provider
):
_model = "{}/{}".format(_logging_obj_llm_provider, _model)
_model = self._cached_model_name
_logging_obj_llm_provider = self._cached_logging_llm_provider
if chunk is None:
chunk = {}
args: Dict[str, Any] = {"model": _model}
else:
# pop model keyword
chunk.pop("model", None)
chunk_dict = {}
for key, value in chunk.items():
if key != "stream":
chunk_dict[key] = value
args = {
"model": _model,
**chunk_dict,
}
args = {"model": _model}
if chunk:
args.update({k: v for k, v in chunk.items() if k != "stream"})
model_response = ModelResponseStream(**args)
if self.response_id is not None:
@ -717,15 +730,23 @@ class CustomStreamWrapper:
model_response.created = self.created
else:
self.created = model_response.created
# Spread order is load-bearing: _base_hidden_params (model_id, api_base, ...)
# must win over both caller-supplied hidden_params and the computed
# custom_llm_provider/created_at values, so it comes last.
if hidden_params is not None:
model_response._hidden_params = hidden_params
model_response._hidden_params["custom_llm_provider"] = _logging_obj_llm_provider
model_response._hidden_params["created_at"] = time.time()
model_response._hidden_params = {
**model_response._hidden_params,
**self._hidden_params,
"response_cost": None,
}
model_response._hidden_params = {
**hidden_params,
"custom_llm_provider": _logging_obj_llm_provider,
"created_at": time.time(),
**self._base_hidden_params,
}
else:
model_response._hidden_params = {
"custom_llm_provider": _logging_obj_llm_provider,
"created_at": time.time(),
**self._base_hidden_params,
}
if (
len(model_response.choices) > 0
@ -1627,7 +1648,17 @@ class CustomStreamWrapper:
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.utils import CallTypes
# Get request kwargs from logging object
if self._post_streaming_hooks is None:
self._post_streaming_hooks = [
cb
for cb in litellm.callbacks
if isinstance(cb, CustomLogger)
and hasattr(cb, "async_post_call_streaming_deployment_hook")
]
if not self._post_streaming_hooks:
return chunk
request_data = self.logging_obj.model_call_details
call_type_str = self.logging_obj.call_type
@ -1636,18 +1667,14 @@ class CustomStreamWrapper:
except ValueError:
typed_call_type = None
# Call hooks for all callbacks
for callback in litellm.callbacks:
if isinstance(callback, CustomLogger) and hasattr(
callback, "async_post_call_streaming_deployment_hook"
):
result = await callback.async_post_call_streaming_deployment_hook(
request_data=request_data,
response_chunk=chunk,
call_type=typed_call_type,
)
if result is not None:
chunk = result
for callback in self._post_streaming_hooks:
result = await callback.async_post_call_streaming_deployment_hook(
request_data=request_data,
response_chunk=chunk,
call_type=typed_call_type,
)
if result is not None:
chunk = result
return chunk
except Exception as e:
@ -1888,17 +1915,15 @@ class CustomStreamWrapper:
response = self._add_mcp_list_tools_to_first_chunk(response)
self.sent_first_chunk = True
if hasattr(
response, "usage"
): # remove usage from chunk, only send on final chunk
# Convert the object to a dictionary
# ModelResponseStream declares `usage` as a field, so
# hasattr(response, "usage") is always True — must check
# `is not None` to avoid running this path on every chunk.
if getattr(response, "usage", None) is not None:
obj_dict = response.model_dump()
# Remove an attribute (e.g., 'attr2')
if "usage" in obj_dict:
del obj_dict["usage"]
# Create a new object without the removed attribute
response = self.model_response_creator(
chunk=obj_dict, hidden_params=response._hidden_params
)
@ -2398,10 +2423,7 @@ def generic_chunk_has_all_required_fields(chunk: dict) -> bool:
:param chunk: The dictionary to check.
:return: True if all required fields are present, False otherwise.
"""
_all_fields = GChunk.__annotations__
decision = all(key in _all_fields for key in chunk)
return decision
return all(key in _GCHUNK_FIELDS for key in chunk)
def convert_generic_chunk_to_model_response_stream(

View file

@ -337,13 +337,12 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
)
@staticmethod
def _supports_effort_level(model: str, level: str) -> bool:
"""Check ``supports_{level}_reasoning_effort`` in the model map.
def _supports_model_capability(model: str, key: str) -> bool:
"""Check a boolean capability ``key`` in the model map.
Strips bedrock/vertex prefixes so a provider-routed Claude still
resolves to the Anthropic model-map entry.
"""
key = f"supports_{level}_reasoning_effort"
try:
if _supports_factory(
model=model,
@ -372,8 +371,6 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
except Exception:
pass
try:
import litellm
for cand in candidates:
if cand in litellm.model_cost and (
litellm.model_cost[cand].get(key) is True
@ -383,6 +380,13 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
pass
return False
@staticmethod
def _supports_effort_level(model: str, level: str) -> bool:
"""Check ``supports_{level}_reasoning_effort`` in the model map."""
return AnthropicConfig._supports_model_capability(
model, f"supports_{level}_reasoning_effort"
)
@staticmethod
def _validate_effort_for_model(model: str, effort: Optional[str]) -> Optional[str]:
"""Return ``None`` if ``effort`` is allowed on ``model``, else an error message."""
@ -400,7 +404,15 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
@staticmethod
def _model_supports_effort_param(model: str) -> bool:
"""Whether the model accepts ``output_config.effort`` at all."""
"""Whether the model accepts ``output_config.effort`` at all.
A model qualifies if its map entry advertises ``supports_output_config``
or any ``supports_*_reasoning_effort`` flag. The two are independent
signals: e.g. Claude Opus 4.5 supports ``output_config`` without
advertising a non-default (max/xhigh) effort level.
"""
if AnthropicConfig._supports_model_capability(model, "supports_output_config"):
return True
return any(
AnthropicConfig._supports_effort_level(model, level)
for level in ("low", "minimal", "medium", "high", "xhigh", "max")
@ -1793,7 +1805,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
self._ensure_context_management_beta_header(
headers, optional_params["context_management"]
)
if optional_params.get("output_format") is not None:
output_config = optional_params.get("output_config")
if optional_params.get("output_format") is not None or (
isinstance(output_config, dict) and output_config.get("format") is not None
):
self._ensure_beta_header(
headers, ANTHROPIC_BETA_HEADER_VALUES.STRUCTURED_OUTPUT_2025_09_25.value
)

View file

@ -427,8 +427,13 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
ANTHROPIC_BETA_HEADER_VALUES.CONTEXT_MANAGEMENT_2025_06_27.value
)
# Check for structured outputs
if optional_params.get("output_format") is not None:
# Check for structured outputs. Anthropic's newer request shape nests
# the schema under output_config.format; the older top-level
# output_format remains supported for backwards compatibility.
output_config = optional_params.get("output_config")
if optional_params.get("output_format") is not None or (
isinstance(output_config, dict) and output_config.get("format") is not None
):
beta_values.add(
ANTHROPIC_BETA_HEADER_VALUES.STRUCTURED_OUTPUT_2025_09_25.value
)

View file

@ -3,6 +3,7 @@ from typing import TYPE_CHECKING, Any, List, Optional, Union
import httpx
from litellm.types.llms.openai import OpenAIRealtimeStreamSessionEvents
from litellm.types.realtime import (
RealtimeResponseTransformInput,
RealtimeResponseTypedDict,
@ -69,6 +70,20 @@ class BaseRealtimeConfig(ABC):
) -> Optional[str]: # message sent to setup the realtime session
return None
def transform_session_created_event(
self,
model: str,
logging_session_id: str,
session_configuration_request: Optional[str] = None,
) -> Optional[Union[dict, OpenAIRealtimeStreamSessionEvents]]:
"""
Optional hook for providers that defer session setup until client `session.update`.
Return an OpenAI-compatible `session.created` payload when the proxy should
emit a synthetic event immediately after backend websocket connection.
"""
return None
@abstractmethod
def transform_realtime_response(
self,

View file

@ -321,6 +321,23 @@ class BaseVideoConfig(ABC):
"video get character is not supported for this provider"
)
def get_video_edit_prefetch_params(
self,
video_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> Optional[Tuple[str, Dict]]:
"""
Return (url, body) for a pre-fetch HTTP call that must be made before
transform_video_edit_request, or None if no pre-fetch is required.
Providers that need to retrieve the source video before constructing the
edit request (e.g. Vertex AI) should override this method. The handler
uses the existing shared httpx client so the call is properly async.
"""
return None
def transform_video_edit_request(
self,
prompt: str,
@ -329,6 +346,7 @@ class BaseVideoConfig(ABC):
litellm_params: GenericLiteLLMParams,
headers: dict,
extra_body: Optional[Dict[str, Any]] = None,
prefetched_source_data: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict]:
"""
Transform the video edit request into a URL and JSON data.
@ -343,6 +361,7 @@ class BaseVideoConfig(ABC):
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: Optional[str] = None,
request_data: Optional[Dict] = None,
) -> VideoObject:
raise NotImplementedError("video edit is not supported for this provider")

View file

@ -78,6 +78,7 @@ from ..common_utils import (
get_anthropic_beta_from_headers,
get_bedrock_tool_name,
is_claude_4_5_on_bedrock,
normalize_bedrock_opus_output_config_effort,
)
# Computer use tool prefixes supported by Bedrock
@ -448,10 +449,20 @@ class AmazonConverseConfig(BaseConfig):
value=reasoning_effort,
llm_provider="bedrock_converse",
)
existing_output_config = optional_params.get("output_config")
if not isinstance(existing_output_config, dict):
existing_output_config = {}
existing_output_config.setdefault("effort", mapped_effort)
normalize_bedrock_opus_output_config_effort(
model=model,
output_config=existing_output_config,
)
mapped_effort = existing_output_config["effort"]
self._validate_anthropic_adaptive_effort(
model=model, effort=mapped_effort
)
optional_params["output_config"] = {"effort": mapped_effort}
optional_params["output_config"] = existing_output_config
optional_params["_output_config_normalized"] = True
@staticmethod
def _validate_anthropic_adaptive_effort(model: str, effort: str) -> None:
@ -1201,6 +1212,12 @@ class AmazonConverseConfig(BaseConfig):
self, optional_params: dict, model: str
) -> Tuple[dict, dict, dict, Optional[OutputConfigBlock]]:
"""Prepare and separate request parameters."""
# Consume the internal ``_output_config_normalized`` marker set by
# ``_handle_reasoning_effort_parameter`` so it does not linger on the
# caller's ``optional_params`` after the transformation returns.
anthropic_output_config_already_normalized = bool(
optional_params.pop("_output_config_normalized", False)
)
# Filter out exception objects before deepcopy to prevent deepcopy failures
# Exceptions should not be stored in optional_params (this is a defensive fix)
cleaned_params = filter_exceptions_from_params(optional_params)
@ -1219,8 +1236,17 @@ class AmazonConverseConfig(BaseConfig):
# Anthropic-only ``output_config`` (snake_case) — re-attached to
# ``additionalModelRequestFields`` for Anthropic models below. The
# Bedrock-native ``outputConfig`` (camelCase) is handled separately.
# structured-output ``format`` subfield is consumed into Bedrock's
# native ``outputConfig`` (camelCase), which is handled separately.
anthropic_output_config = inference_params.pop("output_config", None)
output_config_format = None
if isinstance(anthropic_output_config, dict):
anthropic_output_config = dict(anthropic_output_config)
candidate_output_config_format = anthropic_output_config.pop("format", None)
if isinstance(candidate_output_config_format, dict):
output_config_format = candidate_output_config_format
if not anthropic_output_config:
anthropic_output_config = None
# Extract requestMetadata before processing other parameters
request_metadata = inference_params.pop("requestMetadata", None)
@ -1230,6 +1256,30 @@ class AmazonConverseConfig(BaseConfig):
output_config: Optional[OutputConfigBlock] = inference_params.pop(
"outputConfig", None
)
base_model = BedrockModelInfo.get_base_model(model)
if (
output_config is None
and output_config_format is not None
and output_config_format.get("type") == "json_schema"
and base_model.startswith("anthropic")
and self._supports_native_structured_outputs(
model, self.custom_llm_provider
)
):
output_config = self._create_output_config_for_response_format(
json_schema=output_config_format.get("schema"),
name=output_config_format.get("name"),
description=output_config_format.get("description"),
)
elif output_config is None and output_config_format is not None:
litellm.verbose_logger.warning(
"Bedrock Converse: dropping `output_config.format` for model=%s — "
"model does not advertise `supports_native_structured_output` in "
"model_prices_and_context_window.json. The schema will not be "
"enforced; pass `response_format` to use the synthetic tool-call "
"fallback.",
model,
)
# keep supported params in 'inference_params', and set all model-specific params in 'additional_request_params'
additional_request_params = {
@ -1275,7 +1325,6 @@ class AmazonConverseConfig(BaseConfig):
if anthropic_output_config is not None and isinstance(
anthropic_output_config, dict
):
base_model = BedrockModelInfo.get_base_model(model)
if base_model.startswith("anthropic"):
if (
litellm.drop_params is True
@ -1286,6 +1335,11 @@ class AmazonConverseConfig(BaseConfig):
model,
)
else:
if not anthropic_output_config_already_normalized:
normalize_bedrock_opus_output_config_effort(
model=model,
output_config=anthropic_output_config,
)
effort = anthropic_output_config.get("effort")
if effort is not None:
self._validate_anthropic_adaptive_effort(

View file

@ -16,8 +16,11 @@ from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation
AmazonInvokeConfig,
)
from litellm.llms.bedrock.common_utils import (
convert_bedrock_invoke_output_format_to_inline_schema,
get_anthropic_beta_from_headers,
normalize_bedrock_opus_output_config_effort,
normalize_tool_input_schema_types_for_bedrock_invoke,
pop_bedrock_invoke_output_config_format,
remove_custom_field_from_tools,
)
from litellm.types.llms.anthropic import ANTHROPIC_TOOL_SEARCH_BETA_HEADER
@ -75,6 +78,17 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
# Use a model name that forces tool-based approach
model = "claude-3-sonnet-20240229"
# Clamp ``reasoning_effort`` to the Bedrock effort ceiling before the
# parent mapping converts it to ``output_config.effort`` and the
# downstream effort gate runs. Mirrors the converse path's
# ``_handle_reasoning_effort_parameter`` and the messages path's
# ``_clamp_adaptive_reasoning_effort_for_bedrock`` so adaptive Claude
# requests degrade ``xhigh`` -> ``max`` rather than 400-ing on
# models like Opus 4.6 that don't natively advertise xhigh.
self._clamp_adaptive_reasoning_effort_for_bedrock(
model=original_model, params=non_default_params
)
optional_params = AnthropicConfig.map_openai_params(
self,
non_default_params,
@ -88,6 +102,27 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
return optional_params
@staticmethod
def _clamp_adaptive_reasoning_effort_for_bedrock(model: str, params: dict) -> None:
"""Lower ``reasoning_effort`` to the Bedrock effort ceiling before mapping.
Bedrock's adaptive Claude models accept the OpenAI-style
``reasoning_effort`` tier, but the request validator can reject tiers
the model does not natively advertise (e.g. ``xhigh`` on Opus 4.6).
Clamp the raw tier to the model's
``bedrock_output_config_effort_ceiling`` so Claude Code "goal mode"
keeps working. Non-adaptive models and models without a ceiling are
left untouched.
"""
if not AnthropicConfig._is_adaptive_thinking_model(model):
return
effort = params.get("reasoning_effort")
if not isinstance(effort, str):
return
clamped = {"effort": effort}
normalize_bedrock_opus_output_config_effort(model=model, output_config=clamped)
params["reasoning_effort"] = clamped["effort"]
def transform_request(
self,
model: str,
@ -157,6 +192,13 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
for k, v in optional_params.items()
if k not in self.aws_authentication_params
}
output_config = filtered_params.get("output_config")
if isinstance(output_config, dict):
filtered_params["output_config"] = dict(output_config)
normalize_bedrock_opus_output_config_effort(
model=model,
output_config=filtered_params["output_config"],
)
filtered_params = self._normalize_bedrock_tool_search_tools(filtered_params)
anthropic_request = AnthropicConfig.transform_request(
@ -170,7 +212,20 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
anthropic_request.pop("model", None)
anthropic_request.pop("stream", None)
anthropic_request.pop("output_format", None)
output_format = anthropic_request.pop("output_format", None)
output_config_format = pop_bedrock_invoke_output_config_format(
anthropic_request
)
if output_format:
convert_bedrock_invoke_output_format_to_inline_schema(
output_format=output_format,
request_body=anthropic_request,
)
elif output_config_format:
convert_bedrock_invoke_output_format_to_inline_schema(
output_format=output_config_format,
request_body=anthropic_request,
)
if not (
_supports_factory(
model=model,

View file

@ -34,6 +34,15 @@ class BedrockError(BaseLLMException):
# Lazy import cache to avoid circular imports and performance impact
_get_model_info = None
BedrockOutputConfigEffort = Literal["low", "medium", "high", "max", "xhigh"]
_BEDROCK_OUTPUT_CONFIG_EFFORT_ORDER: Dict[BedrockOutputConfigEffort, int] = {
"low": 0,
"medium": 1,
"high": 2,
"max": 3,
"xhigh": 4,
}
def get_cached_model_info():
"""
@ -51,6 +60,79 @@ def get_cached_model_info():
return _get_model_info
@functools.lru_cache(maxsize=1)
def _get_local_model_cost_map() -> Dict:
from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap
return GetModelCostMap.load_local_model_cost_map()
def pop_bedrock_invoke_output_config_format(request_body: Dict) -> Optional[Dict]:
"""
Remove and return Anthropic's nested ``output_config.format`` field.
Bedrock Invoke paths convert the schema to inline message text. Any remaining
``output_config`` keys, such as ``effort``, are left in place.
"""
output_config = request_body.get("output_config")
if not isinstance(output_config, dict):
return None
output_format = output_config.pop("format", None)
if not output_config:
request_body.pop("output_config", None)
if isinstance(output_format, dict):
return output_format
return None
def convert_bedrock_invoke_output_format_to_inline_schema(
output_format: Dict,
request_body: Dict,
) -> None:
"""
Embed an Anthropic structured-output schema into the last user message.
Bedrock Invoke does not support ``output_format`` directly, so the schema is
appended to the final user message for prompt-engineered structured output.
The caller's ``messages`` list, message dict, and content list are not
mutated; a fresh ``messages`` list with a copied final user message is
written back to ``request_body``.
"""
schema = output_format.get("schema")
if not schema:
return
messages = request_body.get("messages")
if not isinstance(messages, list) or not messages:
return
last_user_idx = None
for i in range(len(messages) - 1, -1, -1):
message = messages[i]
if isinstance(message, dict) and message.get("role") == "user":
last_user_idx = i
break
if last_user_idx is None:
return
original = messages[last_user_idx]
content = original.get("content", [])
schema_block = {"type": "text", "text": json.dumps(schema)}
if isinstance(content, str):
new_content = [{"type": "text", "text": content}, schema_block]
elif isinstance(content, list):
new_content = [*content, schema_block]
else:
return
new_messages = list(messages)
new_messages[last_user_idx] = {**original, "content": new_content}
request_body["messages"] = new_messages
def remove_custom_field_from_tools(request_body: dict) -> None:
"""
Remove ``custom`` field from each tool in the request body.
@ -603,6 +685,62 @@ def is_claude_4_5_on_bedrock(model: str) -> bool:
return any(pattern in model_lower for pattern in claude_4_5_patterns)
def normalize_bedrock_opus_output_config_effort(model: str, output_config: Any) -> None:
"""
Normalize Anthropic ``output_config.effort`` values for Bedrock Opus ids.
Bedrock's Claude Opus request validator can accept a narrower effort
vocabulary than Anthropic's compatibility surface. The Bedrock ceiling is
read from ``model_prices_and_context_window.json`` via
``bedrock_output_config_effort_ceiling``.
Mutates ``output_config`` in place so callers can accept Claude Code's
``xhigh`` input without forwarding a provider-invalid value.
"""
if not isinstance(output_config, dict):
return
effort = output_config.get("effort")
if effort not in _BEDROCK_OUTPUT_CONFIG_EFFORT_ORDER:
return
ceiling = _get_bedrock_output_config_effort_ceiling(model)
if ceiling is None:
return
if (
_BEDROCK_OUTPUT_CONFIG_EFFORT_ORDER[effort]
> _BEDROCK_OUTPUT_CONFIG_EFFORT_ORDER[ceiling]
):
output_config["effort"] = ceiling
def _get_bedrock_output_config_effort_ceiling(
model: str,
) -> Optional[BedrockOutputConfigEffort]:
try:
model_info = get_cached_model_info()(
model=model,
custom_llm_provider="bedrock",
)
except Exception:
return None
ceiling = model_info.get("bedrock_output_config_effort_ceiling")
if isinstance(ceiling, str) and ceiling in _BEDROCK_OUTPUT_CONFIG_EFFORT_ORDER:
return ceiling # type: ignore[return-value]
model_cost_key = model_info.get("key")
if not isinstance(model_cost_key, str):
return None
local_model_info = _get_local_model_cost_map().get(model_cost_key, {})
ceiling = local_model_info.get("bedrock_output_config_effort_ceiling")
if isinstance(ceiling, str) and ceiling in _BEDROCK_OUTPUT_CONFIG_EFFORT_ORDER:
return ceiling # type: ignore[return-value]
return None
# Import after standalone functions to avoid circular imports
from litellm.llms.bedrock.count_tokens.bedrock_token_counter import BedrockTokenCounter

View file

@ -32,10 +32,13 @@ from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation
AmazonInvokeConfig,
)
from litellm.llms.bedrock.common_utils import (
convert_bedrock_invoke_output_format_to_inline_schema,
ensure_bedrock_anthropic_messages_tool_names,
get_anthropic_beta_from_headers,
is_claude_4_5_on_bedrock,
normalize_bedrock_opus_output_config_effort,
normalize_tool_input_schema_types_for_bedrock_invoke,
pop_bedrock_invoke_output_config_format,
remove_custom_field_from_tools,
)
from litellm.types.llms.anthropic import ANTHROPIC_TOOL_SEARCH_BETA_HEADER
@ -450,145 +453,15 @@ class AmazonAnthropicClaudeMessagesConfig(
else:
anthropic_messages_request.pop("context_management", None)
def _convert_output_format_to_inline_schema(
self,
output_format: Dict,
anthropic_messages_request: Dict,
) -> None:
"""
Convert Anthropic output_format to inline schema in message content.
Bedrock Invoke doesn't support the output_format parameter, so we embed
the schema directly into the user message content as text instructions.
This approach adds the schema to the last user message, instructing the model
to respond in the specified JSON format.
Args:
output_format: The output_format dict with 'type' and 'schema'
anthropic_messages_request: The request dict to modify in-place
Ref: https://aws.amazon.com/blogs/machine-learning/structured-data-response-with-amazon-bedrock-prompt-engineering-and-tool-use/
"""
import json
# Extract schema from output_format
schema = output_format.get("schema")
if not schema:
return
# Get messages from the request
messages = anthropic_messages_request.get("messages", [])
if not messages:
return
# Find the last user message
last_user_message_idx = None
for idx in range(len(messages) - 1, -1, -1):
if messages[idx].get("role") == "user":
last_user_message_idx = idx
break
if last_user_message_idx is None:
return
last_user_message = messages[last_user_message_idx]
content = last_user_message.get("content", [])
# Ensure content is a list
if isinstance(content, str):
content = [{"type": "text", "text": content}]
last_user_message["content"] = content
# Add schema as text content to the message
schema_text = {"type": "text", "text": json.dumps(schema)}
content.append(schema_text)
def transform_anthropic_messages_request(
def _get_bedrock_invoke_anthropic_beta_headers(
self,
model: str,
messages: List[Dict],
anthropic_messages_optional_request_params: Dict,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> Dict:
anthropic_messages_request = AnthropicMessagesConfig.transform_anthropic_messages_request(
self=self,
model=model,
messages=messages,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
litellm_params=litellm_params,
headers=headers,
)
#########################################################
############## BEDROCK Invoke SPECIFIC TRANSFORMATION ###
#########################################################
# 1. anthropic_version is required for all claude models
if "anthropic_version" not in anthropic_messages_request:
anthropic_messages_request["anthropic_version"] = (
self.DEFAULT_BEDROCK_ANTHROPIC_API_VERSION
)
# 2. `stream` is not allowed in request body for bedrock invoke
if "stream" in anthropic_messages_request:
anthropic_messages_request.pop("stream", None)
# 3. `model` is not allowed in request body for bedrock invoke
if "model" in anthropic_messages_request:
anthropic_messages_request.pop("model", None)
injected_thinking_for_clear_thinking = (
self._ensure_thinking_for_clear_thinking_context_management(
anthropic_messages_request=anthropic_messages_request,
model=model,
)
)
# 4. Remove `ttl` field from cache_control in messages (Bedrock doesn't support it for older models)
self._remove_ttl_from_cache_control(
anthropic_messages_request=anthropic_messages_request, model=model
)
# 5. Convert `output_format` to inline schema (Bedrock invoke doesn't support output_format)
output_format = anthropic_messages_request.pop("output_format", None)
if output_format:
self._convert_output_format_to_inline_schema(
output_format=output_format,
anthropic_messages_request=anthropic_messages_request,
)
# 5a. Bedrock Invoke supports output_config (effort) for Claude 4.6+ models,
# but older models do not — strip it to avoid request rejection.
# Ref: https://github.com/BerriAI/litellm/issues/22797
if not (
_supports_factory(
model=model,
custom_llm_provider="bedrock",
key="supports_output_config",
)
or AnthropicConfig._model_supports_effort_param(model)
):
if anthropic_messages_request.pop("output_config", None) is not None:
verbose_logger.warning(
"Bedrock Invoke: stripping unsupported `output_config` for "
"model=%s — neither `supports_output_config` nor any "
"`supports_*_reasoning_effort` flag is set in "
"model_prices_and_context_window.json. Add the capability "
"flag to the model JSON entry if this model accepts "
"`output_config`.",
model,
)
# 5b. Remove `custom` field from tools (Bedrock doesn't support it)
# Claude Code sends `custom: {defer_loading: true}` on tool definitions,
# which causes Bedrock to reject the request with "Extra inputs are not permitted"
# Ref: https://github.com/BerriAI/litellm/issues/22847
remove_custom_field_from_tools(anthropic_messages_request)
normalize_tool_input_schema_types_for_bedrock_invoke(anthropic_messages_request)
ensure_bedrock_anthropic_messages_tool_names(anthropic_messages_request)
# 6. AUTO-INJECT beta headers based on features used
anthropic_messages_request: Dict,
injected_thinking_for_clear_thinking: bool,
) -> List[str]:
anthropic_model_info = AnthropicModelInfo()
tools = anthropic_messages_optional_request_params.get("tools")
messages_typed = cast(List[AllMessageValues], messages)
@ -651,6 +524,160 @@ class AmazonAnthropicClaudeMessagesConfig(
dropped_user_betas,
)
return filtered_betas
def _strip_unsupported_bedrock_invoke_fields(
self,
anthropic_messages_request: Dict,
) -> Dict:
allowed = self.BEDROCK_INVOKE_ALLOWED_TOP_LEVEL_FIELDS
stripped = sorted(k for k in anthropic_messages_request if k not in allowed)
if stripped:
verbose_logger.debug(
"Bedrock Invoke: stripping unsupported top-level request fields: %s",
stripped,
)
return {k: v for k, v in anthropic_messages_request.items() if k in allowed}
@staticmethod
def _clamp_adaptive_reasoning_effort_for_bedrock(
model: str, optional_params: Dict
) -> None:
"""Lower ``reasoning_effort`` to the Bedrock effort ceiling before validation.
The shared ``/v1/messages`` effort gate rejects tiers a model does not
natively support (e.g. ``xhigh`` on Opus 4.6). Bedrock's chat paths instead
clamp the tier to the model's ``bedrock_output_config_effort_ceiling`` so
Claude Code "goal mode" keeps working; mirror that here so the messages
path degrades ``xhigh`` -> ``max`` rather than 400-ing. Non-adaptive models
and models without a ceiling are left untouched.
"""
if not AnthropicModelInfo._is_adaptive_thinking_model(model):
return
effort = optional_params.get("reasoning_effort")
if not isinstance(effort, str):
return
clamped = {"effort": effort}
normalize_bedrock_opus_output_config_effort(model=model, output_config=clamped)
optional_params["reasoning_effort"] = clamped["effort"]
def transform_anthropic_messages_request(
self,
model: str,
messages: List[Dict],
anthropic_messages_optional_request_params: Dict,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> Dict:
self._clamp_adaptive_reasoning_effort_for_bedrock(
model=model,
optional_params=anthropic_messages_optional_request_params,
)
anthropic_messages_request = AnthropicMessagesConfig.transform_anthropic_messages_request(
self=self,
model=model,
messages=messages,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
litellm_params=litellm_params,
headers=headers,
)
#########################################################
############## BEDROCK Invoke SPECIFIC TRANSFORMATION ###
#########################################################
# 1. anthropic_version is required for all claude models
if "anthropic_version" not in anthropic_messages_request:
anthropic_messages_request["anthropic_version"] = (
self.DEFAULT_BEDROCK_ANTHROPIC_API_VERSION
)
# 2. `stream` is not allowed in request body for bedrock invoke
if "stream" in anthropic_messages_request:
anthropic_messages_request.pop("stream", None)
# 3. `model` is not allowed in request body for bedrock invoke
if "model" in anthropic_messages_request:
anthropic_messages_request.pop("model", None)
injected_thinking_for_clear_thinking = (
self._ensure_thinking_for_clear_thinking_context_management(
anthropic_messages_request=anthropic_messages_request,
model=model,
)
)
# 4. Remove `ttl` field from cache_control in messages (Bedrock doesn't support it for older models)
self._remove_ttl_from_cache_control(
anthropic_messages_request=anthropic_messages_request, model=model
)
# 5. Convert structured-output params to inline schema.
# Bedrock Invoke doesn't support top-level `output_format`; its
# accepted `output_config` subset is also narrower than Anthropic's, so
# consume the newer `output_config.format` shape here instead of
# forwarding it as an unknown nested key.
existing_output_config = anthropic_messages_request.get("output_config")
if isinstance(existing_output_config, dict):
anthropic_messages_request["output_config"] = dict(existing_output_config)
output_format = anthropic_messages_request.pop("output_format", None)
output_config_format = pop_bedrock_invoke_output_config_format(
anthropic_messages_request
)
if output_format:
convert_bedrock_invoke_output_format_to_inline_schema(
output_format=output_format,
request_body=anthropic_messages_request,
)
elif output_config_format:
convert_bedrock_invoke_output_format_to_inline_schema(
output_format=output_config_format,
request_body=anthropic_messages_request,
)
normalize_bedrock_opus_output_config_effort(
model=model,
output_config=anthropic_messages_request.get("output_config"),
)
# 5a. Bedrock Invoke supports output_config (effort) for Claude 4.6+ models,
# but older models do not — strip it to avoid request rejection.
# Ref: https://github.com/BerriAI/litellm/issues/22797
if not (
_supports_factory(
model=model,
custom_llm_provider="bedrock",
key="supports_output_config",
)
or AnthropicConfig._model_supports_effort_param(model)
):
if anthropic_messages_request.pop("output_config", None) is not None:
verbose_logger.warning(
"Bedrock Invoke: stripping unsupported `output_config` for "
"model=%s — neither `supports_output_config` nor any "
"`supports_*_reasoning_effort` flag is set in "
"model_prices_and_context_window.json. Add the capability "
"flag to the model JSON entry if this model accepts "
"`output_config`.",
model,
)
# 5b. Remove `custom` field from tools (Bedrock doesn't support it)
# Claude Code sends `custom: {defer_loading: true}` on tool definitions,
# which causes Bedrock to reject the request with "Extra inputs are not permitted"
# Ref: https://github.com/BerriAI/litellm/issues/22847
remove_custom_field_from_tools(anthropic_messages_request)
normalize_tool_input_schema_types_for_bedrock_invoke(anthropic_messages_request)
ensure_bedrock_anthropic_messages_tool_names(anthropic_messages_request)
# 6. AUTO-INJECT beta headers based on features used
filtered_betas = self._get_bedrock_invoke_anthropic_beta_headers(
model=model,
messages=messages,
anthropic_messages_optional_request_params=anthropic_messages_optional_request_params,
headers=headers,
anthropic_messages_request=anthropic_messages_request,
injected_thinking_for_clear_thinking=injected_thinking_for_clear_thinking,
)
if filtered_betas:
anthropic_messages_request["anthropic_beta"] = filtered_betas
@ -669,16 +696,9 @@ class AmazonAnthropicClaudeMessagesConfig(
# Catches Anthropic-only extensions (output_config, speed, mcp_servers, ...)
# and any future additions Claude Code may start sending. ``context_management``
# has already been pre-filtered to its Bedrock-supported subset above.
allowed = self.BEDROCK_INVOKE_ALLOWED_TOP_LEVEL_FIELDS
stripped = sorted(k for k in anthropic_messages_request if k not in allowed)
if stripped:
verbose_logger.debug(
"Bedrock Invoke: stripping unsupported top-level request fields: %s",
stripped,
)
anthropic_messages_request = {
k: v for k, v in anthropic_messages_request.items() if k in allowed
}
anthropic_messages_request = self._strip_unsupported_bedrock_invoke_fields(
anthropic_messages_request
)
return anthropic_messages_request

View file

@ -5316,6 +5316,28 @@ class BaseLLMHTTPHandler:
)
if _session_config:
realtime_streaming.session_configuration_request = _session_config
# For providers that defer setup until client session.update, optionally
# send synthetic session.created to unblock clients waiting on connect.
if not provider_config.requires_session_configuration():
synthetic_session = provider_config.transform_session_created_event(
model=model,
logging_session_id=logging_obj.litellm_trace_id,
session_configuration_request=None,
)
if synthetic_session is not None:
synthetic_session_str = json.dumps(synthetic_session)
# Record before sending so the synthetic session.created is
# captured in the session log alongside provider-driven
# events; without this it would be silently absent from
# success_handler / async_success_handler payloads.
realtime_streaming.store_message(synthetic_session_str)
await websocket.send_text(synthetic_session_str)
realtime_streaming._session_created_sent_to_client = True
verbose_logger.debug(
"Sent synthetic session.created to client to unblock connection"
)
await realtime_streaming.bidirectional_forward()
except websockets.exceptions.InvalidStatusCode as e: # type: ignore
@ -6538,6 +6560,7 @@ class BaseLLMHTTPHandler:
api_key=api_key or litellm_params.get("api_key", None),
headers=extra_headers or {},
model="",
litellm_params=litellm_params,
)
if extra_headers:
@ -6620,6 +6643,7 @@ class BaseLLMHTTPHandler:
api_key=api_key or litellm_params.get("api_key", None),
headers=extra_headers or {},
model="",
litellm_params=litellm_params,
)
if extra_headers:
@ -6712,6 +6736,7 @@ class BaseLLMHTTPHandler:
api_key=api_key or litellm_params.get("api_key", None),
headers=extra_headers or {},
model="",
litellm_params=litellm_params,
)
if extra_headers:
headers.update(extra_headers)
@ -6783,6 +6808,7 @@ class BaseLLMHTTPHandler:
api_key=api_key or litellm_params.get("api_key", None),
headers=extra_headers or {},
model="",
litellm_params=litellm_params,
)
if extra_headers:
headers.update(extra_headers)
@ -6866,6 +6892,7 @@ class BaseLLMHTTPHandler:
api_key=api_key or litellm_params.get("api_key", None),
headers=extra_headers or {},
model="",
litellm_params=litellm_params,
)
if extra_headers:
headers.update(extra_headers)
@ -6923,6 +6950,7 @@ class BaseLLMHTTPHandler:
api_key=api_key or litellm_params.get("api_key", None),
headers=extra_headers or {},
model="",
litellm_params=litellm_params,
)
if extra_headers:
headers.update(extra_headers)
@ -6999,6 +7027,7 @@ class BaseLLMHTTPHandler:
api_key=api_key or litellm_params.get("api_key", None),
headers=extra_headers or {},
model="",
litellm_params=litellm_params,
)
if extra_headers:
headers.update(extra_headers)
@ -7009,27 +7038,49 @@ class BaseLLMHTTPHandler:
litellm_params=dict(litellm_params),
)
url, data = video_provider_config.transform_video_edit_request(
prompt=prompt,
prefetched_source_data = None
prefetch_params = video_provider_config.get_video_edit_prefetch_params(
video_id=video_id,
api_base=api_base,
litellm_params=litellm_params,
headers=headers,
extra_body=extra_body,
)
logging_obj.pre_call(
input=prompt,
api_key="",
additional_args={
"complete_input_dict": data,
"api_base": url,
"headers": headers,
"video_id": video_id,
},
)
if prefetch_params is not None:
prefetch_url, prefetch_body = prefetch_params
try:
prefetch_resp = sync_httpx_client.post(
url=prefetch_url,
headers=headers,
json=prefetch_body,
timeout=timeout,
)
prefetch_resp.raise_for_status()
except Exception as e:
raise self._handle_error(e=e, provider_config=video_provider_config)
prefetched_source_data = prefetch_resp.json()
try:
url, data = video_provider_config.transform_video_edit_request(
prompt=prompt,
video_id=video_id,
api_base=api_base,
litellm_params=litellm_params,
headers=headers,
extra_body=extra_body,
prefetched_source_data=prefetched_source_data,
)
logging_obj.pre_call(
input=prompt,
api_key="",
additional_args={
"complete_input_dict": data,
"api_base": url,
"headers": headers,
"video_id": video_id,
},
)
response = sync_httpx_client.post(
url=url,
headers=headers,
@ -7041,6 +7092,7 @@ class BaseLLMHTTPHandler:
raw_response=response,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
request_data=data,
)
except Exception as e:
raise self._handle_error(e=e, provider_config=video_provider_config)
@ -7071,6 +7123,7 @@ class BaseLLMHTTPHandler:
api_key=api_key or litellm_params.get("api_key", None),
headers=extra_headers or {},
model="",
litellm_params=litellm_params,
)
if extra_headers:
headers.update(extra_headers)
@ -7081,27 +7134,49 @@ class BaseLLMHTTPHandler:
litellm_params=dict(litellm_params),
)
url, data = video_provider_config.transform_video_edit_request(
prompt=prompt,
prefetched_source_data = None
prefetch_params = video_provider_config.get_video_edit_prefetch_params(
video_id=video_id,
api_base=api_base,
litellm_params=litellm_params,
headers=headers,
extra_body=extra_body,
)
logging_obj.pre_call(
input=prompt,
api_key="",
additional_args={
"complete_input_dict": data,
"api_base": url,
"headers": headers,
"video_id": video_id,
},
)
if prefetch_params is not None:
prefetch_url, prefetch_body = prefetch_params
try:
prefetch_resp = await async_httpx_client.post(
url=prefetch_url,
headers=headers,
json=prefetch_body,
timeout=timeout,
)
prefetch_resp.raise_for_status()
except Exception as e:
raise self._handle_error(e=e, provider_config=video_provider_config)
prefetched_source_data = prefetch_resp.json()
try:
url, data = video_provider_config.transform_video_edit_request(
prompt=prompt,
video_id=video_id,
api_base=api_base,
litellm_params=litellm_params,
headers=headers,
extra_body=extra_body,
prefetched_source_data=prefetched_source_data,
)
logging_obj.pre_call(
input=prompt,
api_key="",
additional_args={
"complete_input_dict": data,
"api_base": url,
"headers": headers,
"video_id": video_id,
},
)
response = await async_httpx_client.post(
url=url,
headers=headers,
@ -7113,6 +7188,7 @@ class BaseLLMHTTPHandler:
raw_response=response,
logging_obj=logging_obj,
custom_llm_provider=custom_llm_provider,
request_data=data,
)
except Exception as e:
raise self._handle_error(e=e, provider_config=video_provider_config)
@ -7160,6 +7236,7 @@ class BaseLLMHTTPHandler:
api_key=api_key or litellm_params.get("api_key", None),
headers=extra_headers or {},
model="",
litellm_params=litellm_params,
)
if extra_headers:
headers.update(extra_headers)
@ -7234,6 +7311,7 @@ class BaseLLMHTTPHandler:
api_key=api_key or litellm_params.get("api_key", None),
headers=extra_headers or {},
model="",
litellm_params=litellm_params,
)
if extra_headers:
headers.update(extra_headers)
@ -7445,6 +7523,7 @@ class BaseLLMHTTPHandler:
api_key=api_key,
headers=extra_headers or {},
model="",
litellm_params=litellm_params,
)
if extra_headers:

File diff suppressed because it is too large Load diff

View file

@ -581,12 +581,23 @@ class GeminiVideoConfig(BaseVideoConfig):
raise NotImplementedError("video get character is not supported for Gemini")
def transform_video_edit_request(
self, prompt, video_id, api_base, litellm_params, headers, extra_body=None
self,
prompt,
video_id,
api_base,
litellm_params,
headers,
extra_body=None,
prefetched_source_data=None,
):
raise NotImplementedError("video edit is not supported for Gemini")
def transform_video_edit_response(
self, raw_response, logging_obj, custom_llm_provider=None
self,
raw_response,
logging_obj,
custom_llm_provider=None,
request_data=None,
):
raise NotImplementedError("video edit is not supported for Gemini")

View file

@ -534,6 +534,7 @@ class OpenAIVideoConfig(BaseVideoConfig):
litellm_params: GenericLiteLLMParams,
headers: dict,
extra_body: Optional[Dict[str, Any]] = None,
prefetched_source_data: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict]:
original_video_id = extract_original_video_id(video_id)
url = f"{api_base.rstrip('/')}/edits"
@ -547,6 +548,7 @@ class OpenAIVideoConfig(BaseVideoConfig):
raw_response: httpx.Response,
logging_obj: Any,
custom_llm_provider: Optional[str] = None,
request_data: Optional[Dict] = None,
) -> VideoObject:
video_obj = VideoObject(**raw_response.json())
if custom_llm_provider and video_obj.id:

View file

@ -623,12 +623,23 @@ class RunwayMLVideoConfig(BaseVideoConfig):
raise NotImplementedError("video get character is not supported for RunwayML")
def transform_video_edit_request(
self, prompt, video_id, api_base, litellm_params, headers, extra_body=None
self,
prompt,
video_id,
api_base,
litellm_params,
headers,
extra_body=None,
prefetched_source_data=None,
):
raise NotImplementedError("video edit is not supported for RunwayML")
def transform_video_edit_response(
self, raw_response, logging_obj, custom_llm_provider=None
self,
raw_response,
logging_obj,
custom_llm_provider=None,
request_data=None,
):
raise NotImplementedError("video edit is not supported for RunwayML")

View file

@ -14,6 +14,7 @@ Auth: OAuth2 Bearer token (not an API key).
import json
from typing import List, Optional
from litellm import verbose_logger
from litellm.llms.gemini.realtime.transformation import GeminiRealtimeConfig
@ -26,6 +27,7 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig):
"""
def __init__(self, access_token: str, project: str, location: str) -> None:
super().__init__()
self._access_token = access_token
self._project = project
self._location = location
@ -138,6 +140,62 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig):
# Request translation
# ------------------------------------------------------------------
def _vertex_model_path(self, model: str) -> str:
"""Return the fully-qualified Vertex AI model resource path."""
return (
f"projects/{self._project}"
f"/locations/{self._location}"
f"/publishers/google/models/{model}"
)
def _build_vertex_ai_setup_config(self, model: str, session_params: dict) -> dict:
"""Build Vertex AI setup configuration with proper model path and defaults."""
# Normalize GA-remapped fields (``output_modalities``, nested
# ``audio.input.transcription``, ``audio.input.turn_detection``) back to
# their flat beta keys so ``map_openai_params`` picks them up. Without
# this, GA clients' explicit modality / transcription / turn-detection
# settings would be silently dropped because ``map_openai_params`` only
# recognises the flat OpenAI-beta key names.
session_params = self._normalize_session_payload_for_mapping(session_params)
setup_config = self.map_openai_params(
optional_params={}, non_default_params=session_params
)
# Use full Vertex AI model path
setup_config["model"] = self._vertex_model_path(model)
# Add Vertex AI specific defaults if not provided
generation_config = setup_config.setdefault("generationConfig", {})
generation_config.setdefault("responseModalities", ["AUDIO"])
# Ensure Vertex defaults for realtimeInputConfig apply even when
# the client provided a partial ``turn_detection`` (e.g. only
# ``silence_duration_ms``). ``map_automatic_turn_detection`` sets
# ``disabled=True`` whenever ``create_response`` is absent or
# ``False``. Force ``disabled=False`` only when the client did
# not explicitly request ``create_response: False`` — that path
# is how transcription guardrails suppress automatic responses,
# and overriding it here would silently bypass the guardrail.
# Vertex Live has no "VAD on, no auto-response" mode, so callers
# that need that behaviour must accept that VAD is off.
client_turn_detection = session_params.get("turn_detection")
client_disabled_auto_response = (
isinstance(client_turn_detection, dict)
and client_turn_detection.get("create_response") is False
)
realtime_input_config = setup_config.setdefault("realtimeInputConfig", {})
automatic_detection = realtime_input_config.setdefault(
"automaticActivityDetection", {}
)
if not client_disabled_auto_response:
automatic_detection["disabled"] = False
automatic_detection.setdefault("silenceDurationMs", 800)
setup_config.setdefault("inputAudioTranscription", {})
setup_config.setdefault("outputAudioTranscription", {})
return setup_config
def transform_realtime_request(
self,
message: str,
@ -147,16 +205,50 @@ class VertexAIRealtimeConfig(GeminiRealtimeConfig):
"""
Translate OpenAI realtime client messages to Vertex AI format.
``session.update`` is intentionally ignored (returns []) because
Vertex AI only accepts a single ``setup`` message at the start of
the connection — sending a second one causes a 1007 close error.
The initial setup (sent automatically before bidirectional_forward)
already includes AUDIO modality and server VAD, so there is nothing
more to configure.
On the first ``session.update`` (when no setup has been sent yet) the
full ``BidiGenerateContentSetup`` is built with Vertex AI's model path
and forwarded. Any later ``session.update`` is dropped: Vertex AI
documents ``setup`` as the first-and-only client message, and a second
``setup`` closes the connection with a 1007 policy error.
"""
json_message = json.loads(message)
if json_message.get("type") == "session.update":
# Do not forward as a second setup — Vertex AI rejects it.
msg_type = json_message.get("type")
if msg_type == "session.update":
if session_configuration_request is None:
setup_config = self._build_vertex_ai_setup_config(
model, json_message.get("session") or {}
)
gemini_setup_msg = json.dumps({"setup": setup_config})
verbose_logger.debug(
"Vertex AI Realtime: Sending initial setup with tools to backend"
)
return [gemini_setup_msg]
# A follow-up session.update can't be forwarded as a second setup
# (Vertex Live closes the WebSocket with 1007). If this drop is
# silencing the audio-transcription guardrail's create_response
# disable, surface a warning so operators know the model will
# auto-respond before the guardrail can gate it on Vertex AI.
client_turn_detection = GeminiRealtimeConfig._extract_turn_detection(
json_message.get("session") or {}
)
if (
isinstance(client_turn_detection, dict)
and client_turn_detection.get("create_response") is False
):
verbose_logger.warning(
"Vertex AI Realtime: Dropping subsequent session.update "
"(turn_detection.create_response=False) — Vertex Live "
"rejects a second setup message. Audio-transcription "
"guardrails cannot suppress the model's auto-response on "
"Vertex AI in non-deferred mode."
)
else:
verbose_logger.debug(
"Vertex AI Realtime: Ignoring session.update (setup already sent)"
)
return []
return super().transform_realtime_request(

View file

@ -40,6 +40,29 @@ else:
BaseLLMException = Any
def _build_vertex_video_usage_from_request_data(
request_data: Optional[Dict[str, Any]],
) -> Dict[str, Any]:
"""Build usage metadata (duration, resolution) for video cost calculation."""
usage_data: Dict[str, Any] = {}
if not request_data:
return usage_data
parameters = request_data.get("parameters", {})
duration = (
parameters.get("durationSeconds") or DEFAULT_GOOGLE_VIDEO_DURATION_SECONDS
)
if duration is not None:
try:
usage_data["duration_seconds"] = float(duration)
except (ValueError, TypeError):
pass
res = parameters.get("resolution")
if res is not None and str(res).strip() != "":
usage_data["video_resolution"] = str(res).strip().lower()
return usage_data
def _convert_image_to_vertex_format(image_file) -> Dict[str, str]:
"""
Convert image file to Vertex AI format with base64 encoding and MIME type.
@ -363,23 +386,7 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
id=video_id, object="video", status="processing", model=model
)
usage_data: Dict[str, Any] = {}
if request_data:
parameters = request_data.get("parameters", {})
duration = (
parameters.get("durationSeconds")
or DEFAULT_GOOGLE_VIDEO_DURATION_SECONDS
)
if duration is not None:
try:
usage_data["duration_seconds"] = float(duration)
except (ValueError, TypeError):
pass
res = parameters.get("resolution")
if res is not None and str(res).strip() != "":
usage_data["video_resolution"] = str(res).strip().lower()
video_obj.usage = usage_data
video_obj.usage = _build_vertex_video_usage_from_request_data(request_data)
return video_obj
def transform_video_status_retrieve_request(
@ -647,15 +654,123 @@ class VertexAIVideoConfig(BaseVideoConfig, VertexBase):
def transform_video_get_character_response(self, raw_response, logging_obj):
raise NotImplementedError("video get character is not supported for Vertex AI")
def get_video_edit_prefetch_params(
self,
video_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
) -> Tuple[str, Dict]:
"""Return the fetchPredictOperation URL and body needed to retrieve the source video."""
return self.transform_video_status_retrieve_request(
video_id=video_id,
api_base=api_base,
litellm_params=litellm_params,
headers=headers,
)
def transform_video_edit_request(
self, prompt, video_id, api_base, litellm_params, headers, extra_body=None
):
raise NotImplementedError("video edit is not supported for Vertex AI")
self,
prompt: str,
video_id: str,
api_base: str,
litellm_params: GenericLiteLLMParams,
headers: dict,
extra_body: Optional[Dict[str, Any]] = None,
prefetched_source_data: Optional[Dict[str, Any]] = None,
) -> Tuple[str, Dict]:
"""
Build a predictLongRunning edit request from the pre-fetched source video.
The actual fetchPredictOperation HTTP call is hoisted into the handler so
it can use the shared async/sync httpx client instead of blocking the loop.
"""
if prefetched_source_data is None:
raise ValueError(
"prefetched_source_data is required for Vertex AI video edit. "
"Ensure get_video_edit_prefetch_params is called by the handler."
)
if not prefetched_source_data.get("done", False):
raise ValueError(
"Source video generation is not complete yet. "
"Check the video status before editing."
)
videos = prefetched_source_data.get("response", {}).get("videos", [])
if not videos:
raise ValueError("No videos found in the completed operation. Cannot edit.")
source_video = videos[0]
video_input: Dict[str, Any] = {}
if "gcsUri" in source_video:
video_input["gcsUri"] = source_video["gcsUri"]
elif "bytesBase64Encoded" in source_video:
video_input["bytesBase64Encoded"] = source_video["bytesBase64Encoded"]
video_input["mimeType"] = source_video.get("mimeType", "video/mp4")
else:
raise ValueError(
"Source video has neither gcsUri nor bytesBase64Encoded. Cannot edit."
)
operation_name = extract_original_video_id(video_id)
model = self.extract_model_from_operation_name(operation_name) or ""
instance_dict: Dict[str, Any] = {"prompt": prompt, "video": video_input}
request_data: Dict[str, Any] = {"instances": [instance_dict]}
if extra_body:
extra_body_copy = dict(extra_body)
nested_params = extra_body_copy.pop("parameters", None)
vertex_params: Dict[str, Any] = {}
if isinstance(nested_params, dict):
vertex_params.update(nested_params)
vertex_params.update(extra_body_copy)
if vertex_params:
request_data["parameters"] = vertex_params
edit_url = f"{api_base.rstrip('/')}/{model}:predictLongRunning"
return edit_url, request_data
def transform_video_edit_response(
self, raw_response, logging_obj, custom_llm_provider=None
):
raise NotImplementedError("video edit is not supported for Vertex AI")
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
custom_llm_provider: Optional[str] = None,
request_data: Optional[Dict] = None,
) -> VideoObject:
"""
Transform the Veo video edit response.
Veo returns the same operation response as video generation:
{"name": "projects/.../operations/OPERATION_ID"}
usage includes duration_seconds and optional video_resolution from the
edit request parameters for cost calculation.
"""
response_data = raw_response.json()
operation_name = response_data.get("name")
if not operation_name:
raise ValueError(f"No operation name in Veo edit response: {response_data}")
model = self.extract_model_from_operation_name(operation_name) or ""
if custom_llm_provider:
video_id = encode_video_id_with_provider(
operation_name, custom_llm_provider, model
)
else:
video_id = operation_name
video_obj = VideoObject(
id=video_id,
object="video",
status="processing",
model=model,
)
video_obj.usage = _build_vertex_video_usage_from_request_data(request_data)
return video_obj
def transform_video_extension_request(
self,

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1 @@
MCP note: **`available_on_public_internet: false` with `delegate_auth_to_upstream: true` (oauth2, interactive - not `client_credentials`)** - LiteLLM still allows the anonymous upstream PKCE path (no proxy API key for `/authorize` and matching MCP routes). The internal-only flag mainly affects other surfaces (e.g. IP-based discovery). Rely on the upstream IdP and network policy; the dashboard shows a warning when both are set, and the proxy logs a warning when the server is loaded from config or the database

View file

@ -992,42 +992,78 @@ class MCPRequestHandler:
"""
Get allowed MCP servers for a team.
Note: object_permission is automatically loaded by get_team_object() in main auth flow.
Unions two sources:
- Legacy team.object_permission (mcp_servers, mcp_access_groups,
mcp_tool_permissions).
- Unified team.access_group_ids → access_group.access_mcp_server_ids.
Mirrors the model-side pattern in can_team_access_model — the group
is already attached to the team, so the team relationship is itself
the gate (no assigned_team_ids check needed here).
"""
try:
# Get team object permission (already loaded in main auth flow)
object_permissions = await MCPRequestHandler._get_team_object_permission(
user_api_key_auth
)
if object_permissions is None:
return []
# Permission entries may be server_ids OR names/aliases — expand to ids.
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy.auth.auth_checks import (
_get_mcp_server_ids_from_access_groups,
get_team_object,
)
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if (
user_api_key_auth is None
or not user_api_key_auth.team_id
or prisma_client is None
):
return []
team_obj: Optional[LiteLLM_TeamTable] = await get_team_object(
team_id=user_api_key_auth.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
parent_otel_span=user_api_key_auth.parent_otel_span,
proxy_logging_obj=proxy_logging_obj,
)
if team_obj is None:
return []
team_access_group_servers = await _get_mcp_server_ids_from_access_groups(
access_group_ids=team_obj.access_group_ids or [],
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
object_permissions = team_obj.object_permission
if object_permissions is None:
return list(set(team_access_group_servers))
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(
object_permissions.mcp_servers or []
)
# Get MCP servers from access groups
access_group_servers = (
legacy_access_group_servers = (
await MCPRequestHandler._get_mcp_servers_from_access_groups(
object_permissions.mcp_access_groups or []
)
)
# servers referenced in tool permissions should also be accessible
tool_perm_servers = list(
global_mcp_server_manager.expand_tool_permissions(
object_permissions.mcp_tool_permissions
).keys()
)
# Combine all lists
all_servers = direct_mcp_servers + access_group_servers + tool_perm_servers
all_servers = (
direct_mcp_servers
+ legacy_access_group_servers
+ tool_perm_servers
+ team_access_group_servers
)
return list(set(all_servers))
except Exception as e:
verbose_logger.warning(

View file

@ -1067,6 +1067,7 @@ class KeyRequestBase(GenerateRequestBase):
key: Optional[str] = None
budget_id: Optional[str] = None
tags: Optional[List[str]] = None
disable_global_guardrails: Optional[bool] = None
enforced_params: Optional[List[str]] = None
allowed_routes: Optional[list] = []
allowed_passthrough_routes: Optional[list] = None
@ -1832,6 +1833,7 @@ class NewTeamRequest(TeamBase):
prompts: Optional[List[str]] = None
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None
allowed_passthrough_routes: Optional[list] = None
disable_global_guardrails: Optional[bool] = None
secret_manager_settings: Optional[dict] = None
model_rpm_limit: Optional[Dict[str, int]] = None
rpm_limit_type: Optional[
@ -1900,6 +1902,7 @@ class UpdateTeamRequest(LiteLLMPydanticObjectBase):
guardrails: Optional[List[str]] = None
policies: Optional[List[str]] = None
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None
disable_global_guardrails: Optional[bool] = None
team_member_budget: Optional[float] = None
team_member_budget_duration: Optional[str] = None
team_member_rpm_limit: Optional[int] = None
@ -4281,6 +4284,7 @@ LiteLLM_ManagementEndpoint_MetadataFields = [
]
LiteLLM_ManagementEndpoint_MetadataFields_Premium = [
"disable_global_guardrails",
"guardrails",
"policies",
"tags",

View file

@ -619,6 +619,9 @@ async def common_checks( # noqa: PLR0915
proxy_logging_obj=proxy_logging_obj,
)
# Run before apply_key_tags_pre_auth injects key metadata.tags into request_body.
_reject_clientside_metadata_tags_check(general_settings, request_body, route)
# If this is a free model, skip all budget checks
if not skip_budget_checks:
# 3. If team is in budget
@ -660,6 +663,14 @@ async def common_checks( # noqa: PLR0915
proxy_logging_obj=proxy_logging_obj,
)
if valid_token is not None:
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
request_data=request_body,
user_api_key_dict=valid_token,
)
with tracer.trace("litellm.proxy.auth.common_checks.tag_max_budget_check"):
await _tag_max_budget_check(
request_body=request_body,
@ -709,7 +720,6 @@ async def common_checks( # noqa: PLR0915
await _check_end_user_budget(end_user_obj=end_user_object, route=route)
_enforce_user_param_check(general_settings, request, request_body, route)
_reject_clientside_metadata_tags_check(general_settings, request_body, route)
_global_proxy_budget_check(global_proxy_spend, skip_budget_checks, route)
_guardrail_modification_check(request_body, team_object)

View file

@ -839,6 +839,7 @@ class ProxyBaseLLMRequestProcessing:
"aget_run",
"acancel_run",
"adelete_run",
"apply_guardrail",
],
version: Optional[str] = None,
user_model: Optional[str] = None,
@ -1368,6 +1369,21 @@ class ProxyBaseLLMRequestProcessing:
user_api_key_dict=user_api_key_dict,
request_data=self.data,
)
if route_type == "aresponses":
# Streaming /v1/responses returns here without
# reaching the non-streaming ownership tail below.
# Wrap the SSE generator so container ownership is
# written once the upstream iterator finishes
# assembling ``completed_response`` — otherwise
# code-interpreter containers created during the
# stream stay unregistered and follow-up file API
# calls 403. Covers the background-polling path
# too, which loops ``body_iterator`` end-to-end.
selected_data_generator = ProxyBaseLLMRequestProcessing._wrap_responses_stream_for_container_ownership(
original_stream_response=response,
wrapped_generator=selected_data_generator,
user_api_key_dict=user_api_key_dict,
)
return await create_response(
generator=selected_data_generator,
media_type="text/event-stream",
@ -1483,8 +1499,93 @@ class ProxyBaseLLMRequestProcessing:
await check_response_size_is_safe(response=response)
if route_type in {"aresponses", "aget_responses"}:
await ProxyBaseLLMRequestProcessing._record_container_owners_from_responses_if_needed(
response=response,
user_api_key_dict=user_api_key_dict,
)
return response
@staticmethod
async def _record_container_owners_from_responses_if_needed(
response: Any,
user_api_key_dict: UserAPIKeyAuth,
) -> None:
"""Register code-interpreter containers so follow-up file APIs pass ownership checks."""
from litellm.proxy.container_endpoints.ownership import (
record_container_owners_from_responses_response,
)
if response is None:
return
try:
await record_container_owners_from_responses_response(
response=response,
user_api_key_dict=user_api_key_dict,
)
except Exception as e:
verbose_proxy_logger.exception(
"Container ownership recording failed after responses call: %s",
e,
)
@staticmethod
def _extract_completed_responses_response(stream_response: Any) -> Any:
"""Pull the assembled ``ResponsesAPIResponse`` off a streaming iterator.
``ResponsesAPIStreamingIterator`` stores the terminal stream event
(``response.completed`` / ``response.incomplete`` / ``response.failed``)
in ``completed_response``; the actual response body hangs off
that event's ``.response`` attribute. Some iterators store the
``ResponsesAPIResponse`` directly. Handle both shapes so the
container-ownership recording path can walk ``.output`` either way.
"""
completed = getattr(stream_response, "completed_response", None)
if completed is None:
return None
response_obj = getattr(completed, "response", None)
if response_obj is not None:
return response_obj
return completed
@staticmethod
async def _wrap_responses_stream_for_container_ownership(
original_stream_response: Any,
wrapped_generator: Any,
user_api_key_dict: UserAPIKeyAuth,
):
"""Forward SSE chunks, then record container ownership at stream end.
Streaming ``/v1/responses`` short-circuits out of
``base_process_llm_request`` before the non-streaming ownership
tail runs, so without this wrap the
``LiteLLM_ManagedObjectTable`` row for any container created
during the stream is never written and follow-up file API calls
return 403.
"""
try:
async for chunk in wrapped_generator:
yield chunk
finally:
try:
completed_obj = (
ProxyBaseLLMRequestProcessing._extract_completed_responses_response(
original_stream_response
)
)
if completed_obj is not None:
await ProxyBaseLLMRequestProcessing._record_container_owners_from_responses_if_needed(
response=completed_obj,
user_api_key_dict=user_api_key_dict,
)
except Exception as e:
verbose_proxy_logger.exception(
"Container ownership recording failed after streaming responses call: %s",
e,
)
async def base_passthrough_process_llm_request(
self,
request: Request,

View file

@ -317,7 +317,15 @@ def initialize_callbacks_on_proxy( # noqa: PLR0915
DatadogCostManagementLogger,
)
datadog_cost_management_obj = DatadogCostManagementLogger()
init_params = {}
if (
"datadog_cost_management" in callback_specific_params
and isinstance(
callback_specific_params["datadog_cost_management"], dict
)
):
init_params = callback_specific_params["datadog_cost_management"]
datadog_cost_management_obj = DatadogCostManagementLogger(**init_params)
imported_list.append(datadog_cost_management_obj)
elif isinstance(callback, CustomLogger):
imported_list.append(callback)

View file

@ -117,6 +117,58 @@ async def _get_prisma_client():
return prisma_client
def _custom_llm_provider_from_responses_response(
response: Any,
default: str = "openai",
) -> str:
hidden_params: Dict[str, Any] = {}
if isinstance(response, dict):
hidden_params = response.get("_hidden_params") or {}
else:
hidden_params = getattr(response, "_hidden_params", None) or {}
provider = hidden_params.get("custom_llm_provider")
if isinstance(provider, str) and provider:
return provider
return default
async def record_container_owners_from_responses_response(
response: Any,
user_api_key_dict: UserAPIKeyAuth,
custom_llm_provider: Optional[str] = None,
) -> None:
"""Track containers created implicitly by code interpreter in /v1/responses."""
container_ids = (
ResponsesAPIRequestUtils.collect_container_ids_from_responses_response(response)
)
if not container_ids:
return
resolved_provider = (
custom_llm_provider or _custom_llm_provider_from_responses_response(response)
)
for container_id in container_ids:
try:
await record_container_owner(
response={"id": container_id, "object": "container"},
user_api_key_dict=user_api_key_dict,
custom_llm_provider=resolved_provider,
)
except Exception as e:
# Per-container errors (including ``HTTPException`` from
# conflicting/forbidden ownership rows) must not abort the
# batch — other containers in the same response should still
# get recorded so their follow-up file API calls don't 403.
verbose_proxy_logger.exception(
"Failed to record container ownership from responses output "
"for container_id=%s: %s",
container_id,
e,
)
async def record_container_owner(
response: Any,
user_api_key_dict: UserAPIKeyAuth,
@ -151,6 +203,8 @@ async def record_container_owner(
file_object = _dump_response(response)
file_object["custom_llm_provider"] = resolved_provider
file_object["provider_container_id"] = original_container_id
# Prisma Python requires Json fields to be serialized as a JSON string.
file_object_json: str = json.dumps(file_object)
prisma_client = await _get_prisma_client()
if prisma_client is None:
@ -172,7 +226,7 @@ async def record_container_owner(
where={"model_object_id": model_object_id},
data={
"unified_object_id": container_id,
"file_object": file_object,
"file_object": file_object_json,
"updated_by": owner,
},
)
@ -181,7 +235,7 @@ async def record_container_owner(
data={
"unified_object_id": container_id,
"model_object_id": model_object_id,
"file_object": file_object,
"file_object": file_object_json,
"file_purpose": CONTAINER_OBJECT_PURPOSE,
"created_by": owner,
"updated_by": owner,

View file

@ -23,11 +23,11 @@ model_list:
model: bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0
#########################################################
########## batch specific params ########################
s3_bucket_name: litellm-proxy-123456789012
s3_bucket_name: litellm-proxy-941277531214
s3_region_name: us-west-2
s3_access_key_id: os.environ/AWS_ACCESS_KEY_ID
s3_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
aws_batch_role_arn: arn:aws:iam::123456789012:role/service-role/AmazonBedrockExecutionRoleForAgents_EXAMPLE
aws_batch_role_arn: arn:aws:iam::941277531214:role/service-role/AmazonBedrockExecutionRoleForAgents_BB9HNW6V4CV
model_info:
mode: batch

View file

@ -10,7 +10,7 @@ from datetime import datetime, timezone
from typing import Any, Dict, List, Literal, Optional, Type, TypeVar, Union, cast
from urllib.parse import urlparse
from fastapi import APIRouter, Depends, HTTPException
from fastapi import APIRouter, Depends, HTTPException, Request
from pydantic import BaseModel
from litellm.proxy.common_utils.path_utils import safe_join
@ -2187,9 +2187,97 @@ async def test_custom_code_guardrail(
)
def _resolve_guardrail_input_type(
active_guardrail: CustomGuardrail, input_type: str
) -> Literal["request", "response"]:
"""Return the effective input_type, auto-upgrading to 'response' for post_call guardrails."""
if input_type == "request":
hook = getattr(active_guardrail, "event_hook", None)
if hook == GuardrailEventHooks.post_call or hook == "post_call":
return "response"
return "response" if input_type == "response" else "request"
def _patch_logging_obj_for_guardrail(
litellm_logging_obj: Any, request: ApplyGuardrailRequest
) -> None:
"""Configure the logging object so Langfuse/OTEL extract input and output correctly."""
litellm_logging_obj.call_type = "pass_through_endpoint"
litellm_logging_obj.model_call_details["call_type"] = "pass_through_endpoint"
litellm_logging_obj.update_messages(
request.messages
if request.messages
else [{"role": "user", "content": request.text}]
)
async def _emit_guardrail_success_logs(
proxy_logging_obj: Any,
litellm_logging_obj: Any,
data: dict,
user_api_key_dict: UserAPIKeyAuth,
response: ApplyGuardrailResponse,
start_time: datetime,
) -> ApplyGuardrailResponse:
"""Fire proxy and LiteLLM success hooks after a successful guardrail run.
Each hook is wrapped defensively so a callback failure never prevents the
caller from receiving the guardrail response. Returns the (possibly
hook-modified) response.
"""
from litellm.litellm_core_utils.thread_pool_executor import (
executor as thread_pool_executor,
)
try:
modified = await proxy_logging_obj.post_call_success_hook(
data=data,
user_api_key_dict=user_api_key_dict,
response=response,
)
if isinstance(modified, ApplyGuardrailResponse):
response = modified
except Exception:
verbose_proxy_logger.exception("apply_guardrail: post_call_success_hook failed")
# Build the logging payload after post_call_success_hook so that logged
# data matches what the caller actually receives if the hook modified
# the response.
response_for_logging = {"response": response.model_dump(exclude_none=True)}
if litellm_logging_obj is not None:
end_time = datetime.now(timezone.utc)
try:
await litellm_logging_obj.async_success_handler(
result=response_for_logging,
start_time=start_time,
end_time=end_time,
cache_hit=False,
)
except Exception:
verbose_proxy_logger.exception(
"apply_guardrail: async_success_handler failed"
)
try:
thread_pool_executor.submit(
litellm_logging_obj.success_handler,
response_for_logging,
start_time,
end_time,
False,
)
except Exception:
verbose_proxy_logger.exception(
"apply_guardrail: success_handler submit failed"
)
return response
@router.post("/guardrails/apply_guardrail", response_model=ApplyGuardrailResponse)
@router.post("/apply_guardrail", response_model=ApplyGuardrailResponse)
async def apply_guardrail(
fastapi_request: Request,
request: ApplyGuardrailRequest,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
):
@ -2198,8 +2286,29 @@ async def apply_guardrail(
This endpoint allows testing guardrails by applying them to custom text inputs.
"""
import traceback
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
from litellm.litellm_core_utils.thread_pool_executor import (
executor as thread_pool_executor,
)
from litellm.proxy.proxy_server import (
general_settings,
proxy_config,
proxy_logging_obj,
version,
)
from litellm.proxy.utils import handle_exception_on_proxy
data: dict = {
"guardrail_name": request.guardrail_name,
"input": [request.text],
"messages": request.messages or [],
"metadata": {"route": "/apply_guardrail"},
}
litellm_logging_obj = None
start_time = datetime.now(timezone.utc)
try:
active_guardrail: Optional[CustomGuardrail] = (
GUARDRAIL_REGISTRY.get_initialized_guardrail_callback(
@ -2212,23 +2321,25 @@ async def apply_guardrail(
detail=f"Guardrail '{request.guardrail_name}' not found. Please ensure the guardrail is configured in your LiteLLM proxy.",
)
request_data: dict = {}
if request.messages:
request_data["messages"] = request.messages
request_processor = ProxyBaseLLMRequestProcessing(data=data)
data, litellm_logging_obj = (
await request_processor.common_processing_pre_call_logic(
request=fastapi_request,
general_settings=general_settings,
user_api_key_dict=user_api_key_dict,
version=version,
proxy_logging_obj=proxy_logging_obj,
proxy_config=proxy_config,
route_type="apply_guardrail",
)
)
# Auto-detect input_type: if the caller didn't specify "response" but the
# guardrail only runs post_call (e.g. LLM-as-a-judge), use "response" so
# the test actually exercises the guardrail logic.
from litellm.types.guardrails import GuardrailEventHooks
if litellm_logging_obj is not None:
_patch_logging_obj_for_guardrail(litellm_logging_obj, request)
resolved_input_type = request.input_type
if resolved_input_type == "request":
hook = getattr(active_guardrail, "event_hook", None)
if hook == GuardrailEventHooks.post_call or hook == "post_call":
resolved_input_type = "response"
_input_type: Literal["request", "response"] = (
"response" if resolved_input_type == "response" else "request"
request_data: dict = {"messages": request.messages} if request.messages else {}
_input_type = _resolve_guardrail_input_type(
active_guardrail, request.input_type
)
guardrailed_inputs = await active_guardrail.apply_guardrail(
inputs={"texts": [request.text]},
@ -2236,13 +2347,55 @@ async def apply_guardrail(
input_type=_input_type,
)
response_text = guardrailed_inputs.get("texts", [])
return ApplyGuardrailResponse(
response = ApplyGuardrailResponse(
response_text=response_text[0] if response_text else request.text
)
except Exception as e:
if litellm_logging_obj is not None and not isinstance(e, HTTPException):
try:
await litellm_logging_obj.async_failure_handler(
exception=e,
traceback_exception=traceback.format_exc(),
)
except Exception:
verbose_proxy_logger.exception(
"apply_guardrail: async_failure_handler failed"
)
try:
thread_pool_executor.submit(
litellm_logging_obj.failure_handler,
e,
traceback.format_exc(),
)
except Exception:
verbose_proxy_logger.exception(
"apply_guardrail: failure_handler submit failed"
)
try:
transformed_exception = await proxy_logging_obj.post_call_failure_hook(
user_api_key_dict=user_api_key_dict,
original_exception=e,
request_data=data,
)
if isinstance(transformed_exception, Exception):
e = transformed_exception
except Exception:
verbose_proxy_logger.exception(
"apply_guardrail: post_call_failure_hook failed"
)
raise handle_exception_on_proxy(e)
# Success logging outside except so a hook error never triggers failure handlers.
response = await _emit_guardrail_success_logs(
proxy_logging_obj=proxy_logging_obj,
litellm_logging_obj=litellm_logging_obj,
data=data,
user_api_key_dict=user_api_key_dict,
response=response,
start_time=start_time,
)
return response
# Usage (dashboard) endpoints: overview, detail, logs
router.include_router(guardrails_usage_router)

View file

@ -1193,6 +1193,36 @@ class LiteLLMProxyRequestSetup:
return tags
@staticmethod
def apply_key_tags_pre_auth(
request_data: dict,
user_api_key_dict: UserAPIKeyAuth,
) -> None:
"""Merge key metadata tags into request_data before _tag_max_budget_check."""
key_metadata = user_api_key_dict.metadata
if not key_metadata:
return
key_tags = key_metadata.get("tags")
if not key_tags or not isinstance(key_tags, list):
return
_metadata_variable_name = get_metadata_variable_name_from_kwargs(request_data)
metadata = request_data.get(_metadata_variable_name)
if isinstance(metadata, str):
parsed = safe_json_loads(metadata)
metadata = parsed if isinstance(parsed, dict) else {}
request_data[_metadata_variable_name] = metadata
elif not isinstance(metadata, dict):
metadata = {}
request_data[_metadata_variable_name] = metadata
existing_tags = metadata.get("tags")
metadata["tags"] = LiteLLMProxyRequestSetup._merge_tags(
request_tags=existing_tags if isinstance(existing_tags, list) else None,
tags_to_add=key_tags,
)
@staticmethod
def apply_client_tag_policy_pre_auth(
request: Request,
@ -1513,10 +1543,16 @@ async def add_litellm_data_to_request( # noqa: PLR0915
# spend_tracking_utils, streaming_iterator) read `body` to audit the
# request; taking the snapshot here ensures they see cleaned metadata.
#
# Exclude secret_fields (which contains raw_headers with Authorization
# tokens) from the snapshot — they must never be persisted in spend logs
# or any other audit trail.
_body_snapshot = {k: v for k, v in data.items() if k != "secret_fields"}
# Exclude:
# - secret_fields: contains raw_headers with Authorization tokens; must
# never be persisted in spend logs or any other audit trail.
# - proxy_server_request: already a key on `data` at this point (set
# earlier in this function); including it would make the snapshot
# self-reference — body.proxy_server_request.body would be the same
# dict as body, producing an infinite traversal loop for any consumer
# that walks the structure.
_body_snapshot_exclude = {"secret_fields", "proxy_server_request"}
_body_snapshot = {k: v for k, v in data.items() if k not in _body_snapshot_exclude}
data["proxy_server_request"]["body"] = _body_snapshot
# Snapshot the requester-supplied metadata for downstream consumers.

View file

@ -3978,11 +3978,13 @@ async def _batch_resolve_access_group_resources(
def _convert_teams_to_response_models(
teams: list,
use_deleted_table: bool,
keys_count_by_team: Optional[Dict[str, int]] = None,
) -> List[Union[TeamListItem, LiteLLM_TeamTable, LiteLLM_DeletedTeamTable]]:
"""Convert raw Prisma team rows to response models."""
team_list: List[
Union[TeamListItem, LiteLLM_TeamTable, LiteLLM_DeletedTeamTable]
] = []
counts = keys_count_by_team or {}
for team in teams:
try:
team_dict = team.model_dump()
@ -3997,10 +3999,45 @@ def _convert_teams_to_response_models(
members_with_roles = []
team_dict["members_with_roles"] = members_with_roles
members_count = len(members_with_roles)
team_list.append(TeamListItem(**team_dict, members_count=members_count))
keys_count = counts.get(team_dict.get("team_id") or "", 0)
team_list.append(
TeamListItem(
**team_dict,
members_count=members_count,
keys_count=keys_count,
)
)
return team_list
async def _get_keys_count_by_team(
prisma_client: Any,
teams: list,
) -> Dict[str, int]:
"""Aggregate virtual-key counts per team for the given page of teams.
Runs a single GROUP BY against LiteLLM_VerificationToken. The IN clause is
bounded by page_size and uses the existing @@index([team_id]), so this is
one DB round-trip per page. Returns an empty map when the page has no teams.
"""
page_team_ids = [
getattr(t, "team_id", None) for t in teams if getattr(t, "team_id", None)
]
if not page_team_ids:
return {}
grouped = await prisma_client.db.litellm_verificationtoken.group_by(
by=["team_id"],
where={"team_id": {"in": page_team_ids}},
count={"team_id": True},
)
return {
row["team_id"]: row.get("_count", {}).get("team_id", 0)
for row in grouped
if row.get("team_id")
}
async def _enforce_list_team_v2_access(
user_api_key_dict: UserAPIKeyAuth,
user_id: Optional[str],
@ -4228,8 +4265,16 @@ async def list_team_v2(
# Calculate total pages
total_pages = -(-total_count // page_size) # Ceiling division
# Convert Prisma models to response models with members_count
team_list = _convert_teams_to_response_models(teams, use_deleted_table)
# Aggregate virtual-key counts per team for the current page. The deleted
# table does not carry keys_count, so it is skipped.
keys_count_by_team: Dict[str, int] = {}
if not use_deleted_table:
keys_count_by_team = await _get_keys_count_by_team(prisma_client, teams)
# Convert Prisma models to response models with members_count and keys_count
team_list = _convert_teams_to_response_models(
teams, use_deleted_table, keys_count_by_team=keys_count_by_team
)
# Resolve resources inherited from access groups (single batch query)
if not use_deleted_table:

View file

@ -568,14 +568,22 @@ def management_endpoint_wrapper(func):
)
parent_otel_span = getattr(user_api_key_dict, "parent_otel_span", None)
if parent_otel_span is not None:
await _emit_management_endpoint_otel_span(
func=func,
kwargs=kwargs,
parent_otel_span=parent_otel_span,
start_time=start_time,
end_time=end_time,
exception=e,
)
try:
await _emit_management_endpoint_otel_span(
func=func,
kwargs=kwargs,
parent_otel_span=parent_otel_span,
start_time=start_time,
end_time=end_time,
exception=e,
)
except Exception as otel_exc:
# Non-Blocking Exception - never let OTEL failures swallow
# the original management-endpoint exception.
verbose_logger.debug(
"Error emitting OTEL span in management endpoint wrapper failure path: %s",
str(otel_exc),
)
raise e

View file

@ -1071,6 +1071,7 @@ app = FastAPI(
root_path=server_root_path,
lifespan=proxy_startup_event, # type: ignore[reportGeneralTypeIssues]
generate_unique_id_function=_generate_stable_operation_id,
strict_content_type=False,
)
vertex_live_passthrough_vertex_base = VertexBase()

View file

@ -738,6 +738,98 @@ class ResponsesAPIRequestUtils:
model_id,
)
@staticmethod
def _collect_container_ids_from_annotations(
annotations: Any,
collected: set[str],
) -> None:
if not annotations or not isinstance(annotations, list):
return
for ann in annotations:
ResponsesAPIRequestUtils._collect_container_ids_from_output_item(
ann, collected
)
@staticmethod
def _collect_container_ids_from_message_content(
content: Any,
collected: set[str],
) -> None:
if not content:
return
if isinstance(content, list):
for part in content:
if isinstance(part, dict):
ResponsesAPIRequestUtils._collect_container_ids_from_annotations(
part.get("annotations"),
collected,
)
else:
ResponsesAPIRequestUtils._collect_container_ids_from_annotations(
getattr(part, "annotations", None),
collected,
)
@staticmethod
def _collect_container_ids_from_output_item(
item: Any,
collected: set[str],
) -> None:
"""Collect managed or raw ``container_id`` values from one output item."""
if item is None:
return
if isinstance(item, dict):
cid = item.get("container_id")
if isinstance(cid, str) and cid:
collected.add(cid)
nested = item.get("code_interpreter_call")
if isinstance(nested, dict):
nc = nested.get("container_id")
if isinstance(nc, str) and nc:
collected.add(nc)
if item.get("type") == "message":
ResponsesAPIRequestUtils._collect_container_ids_from_message_content(
item.get("content"),
collected,
)
return
cid_attr = getattr(item, "container_id", None)
if isinstance(cid_attr, str) and cid_attr:
collected.add(cid_attr)
nested_obj = getattr(item, "code_interpreter_call", None)
if nested_obj is not None:
ResponsesAPIRequestUtils._collect_container_ids_from_output_item(
nested_obj, collected
)
if getattr(item, "type", None) == "message":
ResponsesAPIRequestUtils._collect_container_ids_from_message_content(
getattr(item, "content", None),
collected,
)
@staticmethod
def collect_container_ids_from_responses_response(response: Any) -> list[str]:
"""Return unique container IDs referenced in a Responses API payload."""
if response is None:
return []
if isinstance(response, dict):
output = response.get("output", [])
else:
output = getattr(response, "output", []) or []
collected: set[str] = set()
if output:
for item in output:
ResponsesAPIRequestUtils._collect_container_ids_from_output_item(
item, collected
)
return list(collected)
@staticmethod
def _update_container_ids_in_response(
responses_api_response: Union[ResponsesAPIResponse, Dict[str, Any]],

View file

@ -52,11 +52,12 @@ PROVIDERS: List[Dict] = [
{
"id": "anthropic",
"name": "Anthropic",
"description": "Claude Opus 4.7, Opus 4.6, Sonnet 4.6, Haiku 4.5",
"description": "Claude Opus 4.8, Opus 4.7, Opus 4.6, Sonnet 4.6, Haiku 4.5",
"env_key": "ANTHROPIC_API_KEY",
"key_hint": "sk-ant-...",
"test_model": "claude-haiku-4-5-20251001",
"models": [
"claude-opus-4-8",
"claude-opus-4-7",
"claude-opus-4-6",
"claude-sonnet-4-6",

View file

@ -1,4 +1,4 @@
from typing import Dict, Optional, TypedDict
from typing import Dict, List, Optional, TypedDict
from litellm.types.integrations.custom_logger import StandardCustomLoggerInitParams
@ -9,7 +9,7 @@ class DatadogCostManagementInitParams(StandardCustomLoggerInitParams):
Init params for Datadog Cost Management
"""
datadog_cost_management_params: Optional[Dict] = None
cost_tag_keys: Optional[List[str]] = None
class DatadogFOCUSCostEntry(TypedDict):

View file

@ -39,7 +39,8 @@ class AnthropicOutputSchema(TypedDict, total=False):
class AnthropicOutputConfig(TypedDict, total=False):
"""Configuration for controlling Claude's output behavior."""
effort: Literal["high", "medium", "low"]
effort: Literal["high", "medium", "low", "xhigh", "max"]
format: AnthropicOutputSchema
class AnthropicMessagesTool(TypedDict, total=False):

View file

@ -133,7 +133,7 @@ class BidiGenerateContentSetup(TypedDict, total=False):
tools: List[Tools]
"""The tools to be used for the realtime session."""
realtimeInputConfig: dict
realtimeInputConfig: BidiGenerateContentRealtimeInputConfig
"""The realtime config to be used for the realtime session."""
sessionResumption: dict

View file

@ -79,7 +79,14 @@ from pydantic import (
field_serializer,
field_validator,
)
from typing_extensions import Annotated, Dict, Required, TypedDict, override
from typing_extensions import (
Annotated,
Dict,
NotRequired,
Required,
TypedDict,
override,
)
from litellm.types.llms.base import BaseLiteLLMOpenAIResponseObject
from litellm.types.responses.main import (
@ -1935,6 +1942,7 @@ class OpenAIRealtimeStreamResponseOutputItemAdded(TypedDict):
response_id: str
output_index: int
item: OpenAIRealtimeStreamResponseOutputItem
event_id: NotRequired[str]
class OpenAIRealtimeStreamResponseBaseObject(TypedDict):
@ -2061,6 +2069,17 @@ class OpenAIRealtimeContentPartDone(TypedDict):
type: Literal["response.content_part.done"]
class OpenAIRealtimeFunctionCallArgumentsDone(TypedDict):
type: Literal["response.function_call_arguments.done"]
event_id: str
response_id: str
item_id: str
output_index: int
call_id: str
name: str
arguments: str
class OpenAIRealtimeOutputItemDone(TypedDict):
event_id: str
item: OpenAIRealtimeStreamResponseOutputItem
@ -2126,6 +2145,7 @@ OpenAIRealtimeEvents = Union[
OpenAIRealtimeResponseAudioDone,
OpenAIRealtimeContentPartDone,
OpenAIRealtimeOutputItemDone,
OpenAIRealtimeFunctionCallArgumentsDone,
OpenAIRealtimeDoneEvent,
]

View file

@ -69,6 +69,7 @@ class TeamListItem(LiteLLM_TeamTable):
"""A team item in the paginated list response, enriched with computed fields."""
members_count: int = 0
keys_count: int = 0
# Resources inherited from access groups (separate from direct assignments)
access_group_models: Optional[List[str]] = None
access_group_mcp_server_ids: Optional[List[str]] = None

View file

@ -148,6 +148,9 @@ class ProviderSpecificModelInfo(TypedDict, total=False):
supports_xhigh_reasoning_effort: Optional[bool]
supports_max_reasoning_effort: Optional[bool]
supports_output_config: Optional[bool]
bedrock_output_config_effort_ceiling: Optional[
Literal["low", "medium", "high", "max", "xhigh"]
]
class SearchContextCostPerQuery(TypedDict, total=False):

View file

@ -6036,6 +6036,9 @@ def _get_model_info_helper( # noqa: PLR0915
supports_max_reasoning_effort=_model_info.get(
"supports_max_reasoning_effort", None
),
bedrock_output_config_effort_ceiling=_model_info.get(
"bedrock_output_config_effort_ceiling", None
),
supports_computer_use=_model_info.get("supports_computer_use", None),
search_context_cost_per_query=_model_info.get(
"search_context_cost_per_query", None

View file

@ -31,12 +31,20 @@ USER root
COPY --from=uvbin /uv /uvx /usr/local/bin/
RUN apk add --no-cache bash gcc python3 python3-dev openssl openssl-dev libsndfile
# nodejs/npm so `prisma generate` uses Wolfi's Node via PRISMA_USE_GLOBAL_NODE
# instead of nodeenv downloading one whose dynamic deps may not be in Wolfi
# (e.g. Node 26.2.0 needs libatomic). Retry for transient apk.cgr.dev flakes.
RUN for i in 1 2 3; do \
apk add --no-cache bash gcc python3 python3-dev openssl openssl-dev libsndfile nodejs npm && break; \
[ $i = 3 ] && { echo "apk add failed after 3 retries" >&2; exit 1; }; \
sleep 5; \
done
ENV UV_PROJECT_ENVIRONMENT=/app/.venv \
UV_LINK_MODE=copy \
UV_COMPILE_BYTECODE=1 \
UV_PYTHON_DOWNLOADS=0 \
PRISMA_USE_GLOBAL_NODE=true \
PATH="/app/.venv/bin:${PATH}"
# Stage 1 — install third-party deps only (cached by pyproject.toml/uv.lock).
@ -78,7 +86,11 @@ FROM $LITELLM_RUNTIME_IMAGE AS runtime
USER root
RUN apk add --no-cache bash openssl tzdata python3 libsndfile libatomic
RUN for i in 1 2 3; do \
apk add --no-cache bash openssl tzdata python3 libsndfile libatomic && break; \
[ $i = 3 ] && { echo "apk add failed after 3 retries" >&2; exit 1; }; \
sleep 5; \
done
# wolfi-base ships an unprivileged `nonroot` account (UID/GID 65532). The
# Prisma engine binaries are dynamically linked against libssl/libcrypto, so

File diff suppressed because it is too large Load diff

View file

@ -1,6 +1,6 @@
[project]
name = "litellm"
version = "1.87.0"
version = "1.88.0"
description = "Library to easily interface with LLM API providers"
readme = "README.md"
requires-python = ">=3.10, <3.14"
@ -33,62 +33,66 @@ Homepage = "https://litellm.ai"
Repository = "https://github.com/BerriAI/litellm"
Documentation = "https://docs.litellm.ai"
# Optional extras retain exact pins because they are consumed by Docker images
# where exact reproducibility matters. The core SDK uses ranges so downstream
# consumers can coexist with other packages without forced downgrades.
# Optional extras use compatible ranges (like the core SDK above) so downstream
# consumers can coexist with other packages and pick up security patches without
# forking. Reproducibility for our Docker/CI comes from `uv.lock` (images install
# via `uv sync --frozen`). A few deps stay exact-pinned: litellm's own
# sub-packages and the opentelemetry trio move in lockstep, and grpcio is
# supply-chain-pinned to a vetted, aged release.
[project.optional-dependencies]
proxy = [
"gunicorn==23.0.0",
"uvicorn==0.33.0",
"granian==2.5.7",
"uvloop==0.21.0; sys_platform != 'win32'",
"fastapi==0.124.4",
"backoff==2.2.1",
"pyyaml==6.0.3",
"rq==2.7.0",
"orjson==3.11.6",
"apscheduler==3.11.2",
"fastapi-sso==0.19.0",
"PyJWT==2.12.0",
"python-multipart==0.0.27",
"cryptography==46.0.7",
"pynacl==1.6.2",
"websockets==15.0.1",
"boto3==1.43.1",
"azure-identity==1.25.2",
"azure-storage-blob==12.28.0",
"mcp==1.26.0",
"gunicorn>=23.0.0,<24.0",
"uvicorn>=0.33.0,<1.0",
"granian>=2.7.4,<3.0",
"uvloop>=0.21.0,<1.0; sys_platform != 'win32'",
"fastapi>=0.136.3,<1.0",
"starlette>=1.0.1,<2.0",
"backoff>=2.2.1,<3.0",
"pyyaml>=6.0.3,<7.0",
"rq>=2.7.0,<3.0",
"orjson>=3.11.6,<4.0",
"apscheduler>=3.11.2,<4.0",
"fastapi-sso>=0.19.0,<1.0",
"PyJWT>=2.12.0,<3.0",
"python-multipart>=0.0.27,<1.0",
"cryptography>=46.0.7,<47.0",
"pynacl>=1.6.2,<2.0",
"websockets>=15.0.1,<16.0",
"boto3>=1.43.1,<2.0",
"azure-identity>=1.25.2,<2.0",
"azure-storage-blob>=12.28.0,<13.0",
"mcp>=1.26.0,<2.0",
"litellm-proxy-extras==0.4.73",
"litellm-enterprise==0.1.41",
"RestrictedPython==8.1",
"rich==13.9.4",
"polars==1.38.1",
"soundfile==0.12.1",
"pyroscope-io==0.8.16; sys_platform != 'win32'",
"pydantic-settings>=2.14.1",
"RestrictedPython>=8.1,<9.0",
"rich>=13.9.4,<14.0",
"polars>=1.38.1,<2.0",
"soundfile>=0.12.1,<1.0",
"pyroscope-io>=0.8.16,<1.0; sys_platform != 'win32'",
"pydantic-settings>=2.14.1,<3.0",
]
extra_proxy = [
"prisma==0.11.0",
"azure-identity==1.25.2",
"azure-keyvault-secrets==4.10.0",
"prisma>=0.11.0,<1.0",
"azure-identity>=1.25.2,<2.0",
"azure-keyvault-secrets>=4.10.0,<5.0",
# Not in PyPI proxy extra.
"google-cloud-kms==2.24.2",
"google-cloud-iam==2.19.1",
"google-cloud-kms>=2.24.2,<3.0",
"google-cloud-iam>=2.19.1,<3.0",
# Not in PyPI proxy extra.
"resend==2.23.0",
"redisvl==0.4.1; python_version < '3.14'",
"a2a-sdk==0.3.24",
"resend>=2.23.0,<3.0",
"redisvl>=0.4.1,<1.0; python_version < '3.14'",
"a2a-sdk>=0.3.24,<1.0",
]
utils = [
# Not in Docker or PyPI proxy extra.
"numpydoc==1.8.0",
"numpydoc>=1.8.0,<2.0",
]
caching = ["diskcache==5.6.3"]
caching = ["diskcache>=5.6.3,<6.0"]
semantic-router = [
"semantic-router==0.1.12; python_version < '3.14'",
"aurelio-sdk==0.0.19; python_version < '3.14'",
"semantic-router>=0.1.15,<1.0; python_version < '3.14'",
"aurelio-sdk>=0.0.19,<1.0; python_version < '3.14'",
]
mlflow = ["mlflow==3.11.1"]
mlflow = ["mlflow>=3.11.1,<4.0"]
grpc = [
# Newest non-yanked release older than the 30-day cutoff.
"grpcio==1.78.0",
@ -101,28 +105,28 @@ stt-nvidia-riva = [
"audioread>=3.0.1",
"numpy>=1.26.0",
]
google = ["google-cloud-aiplatform==1.133.0"]
google = ["google-cloud-aiplatform>=1.133.0,<2.0"]
proxy-runtime = [
# Historically bundled in the proxy Docker images via requirements.txt.
# 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",
"google-genai==1.37.0",
"anthropic[vertex]==0.84.0",
"google-cloud-aiplatform>=1.133.0,<2.0",
"google-genai>=1.37.0,<2.0",
"anthropic[vertex]>=0.84.0,<1.0",
"grpcio==1.78.0",
"prometheus-client==0.20.0",
"langfuse==2.59.7",
"prometheus-client>=0.20.0,<1.0",
"langfuse>=2.59.7,<3.0",
"opentelemetry-api==1.28.0",
"opentelemetry-sdk==1.28.0",
"opentelemetry-exporter-otlp==1.28.0",
"ddtrace==2.19.0",
"sentry-sdk==2.21.0",
"mangum==0.17.0",
"azure-ai-contentsafety==1.0.0",
"azure-storage-file-datalake==12.20.0",
"pypdf==6.10.2; python_version < '3.14'",
"llm-sandbox==0.3.39",
"detect-secrets==1.5.0",
"ddtrace>=2.19.0,<3.0",
"sentry-sdk>=2.21.0,<3.0",
"mangum>=0.17.0,<1.0",
"azure-ai-contentsafety>=1.0.0,<2.0",
"azure-storage-file-datalake>=12.20.0,<13.0",
"pypdf>=6.10.2,<7.0; python_version < '3.14'",
"llm-sandbox>=0.3.39,<1.0",
"detect-secrets>=1.5.0,<2.0",
]
[project.scripts]
@ -188,7 +192,7 @@ ci = [
"psycopg2-binary==2.9.11",
"pytest-codspeed==4.3.0",
"pytest-retry==1.7.0",
"pyarrow==22.0.0",
"pyarrow==23.0.1",
"langchain==1.2.10",
"lunary==1.4.36; python_version == '3.10'",
"lunary==1.4.37; python_version >= '3.11'",
@ -253,7 +257,7 @@ source-exclude = [
profile = "black"
[tool.commitizen]
version = "1.87.0"
version = "1.88.0"
version_files = [
"pyproject.toml:^version",
]

View file

@ -0,0 +1,191 @@
#!/usr/bin/env python3
"""Tight microbenchmark for CustomStreamWrapper.model_response_creator.
Calls model_response_creator() in a tight loop on a pre-built wrapper to
isolate per-call cost. Driving the full wrapper adds threadpool logging,
gc, and other noise that swamps microsecond-scale changes here.
Example:
uv run python scripts/benchmark_model_response_creator.py --label baseline
uv run python scripts/benchmark_model_response_creator.py --label optimized
"""
from __future__ import annotations
import argparse
import gc
import json
import logging
import os
import statistics
import time
from dataclasses import asdict, dataclass
from typing import List
from unittest.mock import MagicMock
os.environ.setdefault("LITELLM_LOG", "ERROR")
logging.getLogger("LiteLLM").setLevel(logging.ERROR)
import litellm # noqa: E402
litellm.suppress_debug_info = True
from litellm.litellm_core_utils.streaming_handler import (
CustomStreamWrapper,
) # noqa: E402
def _make_logging_obj(provider: str) -> MagicMock:
logging_obj = MagicMock()
logging_obj.model_call_details = {
"custom_llm_provider": provider,
"litellm_params": {},
}
logging_obj.call_type = "completion"
logging_obj.stream_options = None
logging_obj.messages = [{"role": "user", "content": "hi"}]
logging_obj.completion_start_time = None
logging_obj._llm_caching_handler = None
return logging_obj
def _make_wrapper(provider: str, model: str) -> CustomStreamWrapper:
return CustomStreamWrapper(
completion_stream=iter([]),
model=model,
logging_obj=_make_logging_obj(provider),
custom_llm_provider=provider,
)
@dataclass
class Result:
label: str
scenario: str
iterations: int
elapsed_min_s: float
elapsed_median_s: float
per_call_us: float
calls_per_sec: float
SCENARIOS = {
"no_chunk": {
"description": "model_response_creator() — no chunk arg (most common path)",
"chunk_factory": lambda i: None,
},
"text_chunk": {
"description": "model_response_creator(chunk={'text': '...'}) — text delta path",
"chunk_factory": lambda i: {"text": f"token{i}"},
},
"rich_chunk": {
"description": "model_response_creator(chunk={...}) — full chunk dict path",
"chunk_factory": lambda i: {
"id": f"id-{i}",
"object": "chat.completion.chunk",
"created": 1234567890,
},
},
}
def bench_no_chunk(wrapper: CustomStreamWrapper, iterations: int) -> float:
gc.collect()
gc.disable()
try:
start = time.perf_counter()
for _ in range(iterations):
wrapper.model_response_creator()
elapsed = time.perf_counter() - start
finally:
gc.enable()
return elapsed
def bench_with_chunk(wrapper: CustomStreamWrapper, factory, iterations: int) -> float:
# Pre-build chunks so we don't measure their construction cost.
chunks = [factory(i) for i in range(iterations)]
gc.collect()
gc.disable()
try:
start = time.perf_counter()
for chunk in chunks:
wrapper.model_response_creator(chunk=dict(chunk)) # copy because mutated
elapsed = time.perf_counter() - start
finally:
gc.enable()
return elapsed
def run_scenario(
label: str,
scenario_key: str,
iterations: int,
repeats: int,
warmup: int,
) -> Result:
spec = SCENARIOS[scenario_key]
wrapper = _make_wrapper(provider="anthropic", model="claude-3-5-sonnet")
if scenario_key == "no_chunk":
runner = lambda: bench_no_chunk(wrapper, iterations) # noqa: E731
else:
runner = lambda: bench_with_chunk(
wrapper, spec["chunk_factory"], iterations
) # noqa: E731
for _ in range(warmup):
runner()
samples = [runner() for _ in range(repeats)]
elapsed_min = min(samples)
elapsed_median = statistics.median(samples)
per_call_us = (elapsed_min * 1_000_000) / iterations
calls_per_sec = iterations / elapsed_min if elapsed_min > 0 else 0.0
return Result(
label=label,
scenario=scenario_key,
iterations=iterations,
elapsed_min_s=elapsed_min,
elapsed_median_s=elapsed_median,
per_call_us=per_call_us,
calls_per_sec=calls_per_sec,
)
def main() -> None:
ap = argparse.ArgumentParser(description=__doc__.splitlines()[0])
ap.add_argument("--label", required=True)
ap.add_argument("--iterations", type=int, default=200_000)
ap.add_argument("--warmup", type=int, default=2)
ap.add_argument("--repeats", type=int, default=8)
ap.add_argument("--json", dest="json_out")
args = ap.parse_args()
print(
f"\n=== label={args.label} iterations={args.iterations:,} "
f"warmup={args.warmup} repeats={args.repeats} (min reported) ==="
)
results: List[Result] = []
for scenario in SCENARIOS:
r = run_scenario(
args.label, scenario, args.iterations, args.repeats, args.warmup
)
results.append(r)
print(
f" {r.scenario:12s}: "
f"min={r.elapsed_min_s*1000:8.2f} ms "
f"median={r.elapsed_median_s*1000:8.2f} ms "
f"per-call={r.per_call_us:7.3f} μs "
f"calls/s={r.calls_per_sec:>12,.0f}"
)
if args.json_out:
with open(args.json_out, "w", encoding="utf-8") as f:
json.dump([asdict(r) for r in results], f, indent=2)
print(f"\nWrote {len(results)} results to {args.json_out}")
if __name__ == "__main__":
main()

View file

@ -0,0 +1,369 @@
#!/usr/bin/env python3
"""Benchmark CustomStreamWrapper per-chunk overhead.
Drives CustomStreamWrapper directly with synthetic in-memory chunks for
Anthropic (GenericStreamingChunk), Bedrock Invoke (GenericStreamingChunk),
and Bedrock Converse (ModelResponseStream). A full proxy benchmark adds
FastAPI, HTTP, and TCP latency, which dilutes the per-chunk CPU signal.
Example:
uv run python scripts/benchmark_streaming_chunk_overhead.py \\
--streams 500 --chunks 200 --warmup 50 --repeats 5
"""
from __future__ import annotations
import argparse
import asyncio
import gc
import json
import logging
import os
import statistics
import time
from dataclasses import asdict, dataclass
from typing import Callable, List, Optional
from unittest.mock import MagicMock
# Silence litellm's "Provider List" warnings emitted by get_llm_provider
# when it sees synthetic model names — we're not exercising provider
# routing, only the per-chunk wrapper hot path.
os.environ.setdefault("LITELLM_LOG", "ERROR")
logging.getLogger("LiteLLM").setLevel(logging.ERROR)
import litellm # noqa: E402
litellm.suppress_debug_info = True
from litellm.litellm_core_utils.streaming_handler import (
CustomStreamWrapper,
) # noqa: E402
from litellm.types.utils import ( # noqa: E402
Delta,
GenericStreamingChunk as GChunk,
ModelResponseStream,
StreamingChoices,
Usage,
)
# ---------------------------------------------------------------------------
# Synthetic chunk fixtures
# ---------------------------------------------------------------------------
def _make_logging_obj(provider: str) -> MagicMock:
logging_obj = MagicMock()
logging_obj.model_call_details = {
"custom_llm_provider": provider,
"litellm_params": {},
}
logging_obj.call_type = "completion"
logging_obj.stream_options = None
logging_obj.messages = [{"role": "user", "content": "hi"}]
logging_obj.completion_start_time = None
logging_obj._llm_caching_handler = None
return logging_obj
def _make_generic_chunk(
text: str,
is_finished: bool = False,
finish_reason: str = "",
usage: Optional[dict] = None,
) -> GChunk:
return GChunk(
text=text,
is_finished=is_finished,
finish_reason=finish_reason,
usage=usage,
index=0,
tool_use=None,
)
def _make_converse_chunk(
text: str = "",
finish_reason: str = "",
usage: Optional[Usage] = None,
) -> ModelResponseStream:
return ModelResponseStream(
choices=[
StreamingChoices(
finish_reason=finish_reason or None,
index=0,
delta=Delta(content=text, role="assistant"),
)
],
id="msg-bench",
model="anthropic.claude-3-5-sonnet",
usage=usage,
)
# ---------------------------------------------------------------------------
# Provider stream factories
# ---------------------------------------------------------------------------
def anthropic_chunks(n: int) -> List[GChunk]:
out: List[GChunk] = [_make_generic_chunk(f"tok{i} ") for i in range(n)]
out.append(
_make_generic_chunk(
"",
is_finished=True,
finish_reason="stop",
usage={"prompt_tokens": 10, "completion_tokens": n, "total_tokens": 10 + n},
)
)
return out
def bedrock_invoke_chunks(n: int) -> List[GChunk]:
# Bedrock Invoke surfaces GChunk-shaped dicts, same shape as Anthropic.
return anthropic_chunks(n)
def bedrock_converse_chunks(n: int) -> List[ModelResponseStream]:
out: List[ModelResponseStream] = [
_make_converse_chunk(f"tok{i} ") for i in range(n)
]
out.append(
_make_converse_chunk(
text="",
finish_reason="stop",
usage=Usage(prompt_tokens=10, completion_tokens=n, total_tokens=10 + n),
)
)
return out
PROVIDERS: dict[str, tuple[str, Callable[[int], list]]] = {
"anthropic": ("anthropic", anthropic_chunks),
"bedrock_invoke": ("bedrock", bedrock_invoke_chunks),
"bedrock_converse": ("bedrock", bedrock_converse_chunks),
}
# ---------------------------------------------------------------------------
# Drive a single stream end-to-end
# ---------------------------------------------------------------------------
def _make_wrapper(
chunks: list, provider: str, async_stream: bool
) -> CustomStreamWrapper:
logging_obj = _make_logging_obj(provider)
if async_stream:
async def _agen():
for c in chunks:
yield c
stream = _agen()
else:
stream = iter(chunks)
return CustomStreamWrapper(
completion_stream=stream,
model="claude-3-5-sonnet",
logging_obj=logging_obj,
custom_llm_provider=provider,
)
def drive_sync(provider_key: str, chunks_per_stream: int, n_streams: int) -> float:
provider, factory = PROVIDERS[provider_key]
# Pre-build the chunk lists; we only measure wrapper iteration cost.
chunk_lists = [factory(chunks_per_stream) for _ in range(n_streams)]
gc.collect()
gc.disable()
try:
start = time.perf_counter()
for chunks in chunk_lists:
wrapper = _make_wrapper(chunks, provider, async_stream=False)
for _ in wrapper:
pass
elapsed = time.perf_counter() - start
finally:
gc.enable()
return elapsed
async def drive_async(
provider_key: str, chunks_per_stream: int, n_streams: int
) -> float:
provider, factory = PROVIDERS[provider_key]
chunk_lists = [factory(chunks_per_stream) for _ in range(n_streams)]
gc.collect()
gc.disable()
try:
start = time.perf_counter()
for chunks in chunk_lists:
wrapper = _make_wrapper(chunks, provider, async_stream=True)
async for _ in wrapper:
pass
elapsed = time.perf_counter() - start
finally:
gc.enable()
return elapsed
# ---------------------------------------------------------------------------
# Repeat × take-min runner
# ---------------------------------------------------------------------------
@dataclass
class Result:
label: str
provider: str
mode: str
streams: int
chunks_per_stream: int
total_chunks: int
elapsed_min_s: float
elapsed_median_s: float
per_chunk_us: float
chunks_per_sec: float
streams_per_sec: float
def run_case(
label: str,
provider_key: str,
mode: str,
chunks_per_stream: int,
n_streams: int,
repeats: int,
warmup: int,
) -> Result:
if mode == "sync":
# Warmup runs amortize import-time and JIT-y caches.
for _ in range(warmup):
drive_sync(provider_key, chunks_per_stream, max(1, n_streams // 10))
samples = [
drive_sync(provider_key, chunks_per_stream, n_streams)
for _ in range(repeats)
]
elif mode == "async":
async def _warm():
for _ in range(warmup):
await drive_async(
provider_key, chunks_per_stream, max(1, n_streams // 10)
)
asyncio.run(_warm())
samples = [
asyncio.run(drive_async(provider_key, chunks_per_stream, n_streams))
for _ in range(repeats)
]
else:
raise ValueError(f"unknown mode {mode!r}")
elapsed_min = min(samples)
elapsed_median = statistics.median(samples)
# Each stream emits chunks_per_stream text chunks + 1 finish/usage chunk.
total_chunks = n_streams * (chunks_per_stream + 1)
per_chunk_us = (elapsed_min * 1_000_000) / total_chunks
chunks_per_sec = total_chunks / elapsed_min if elapsed_min > 0 else 0.0
streams_per_sec = n_streams / elapsed_min if elapsed_min > 0 else 0.0
return Result(
label=label,
provider=provider_key,
mode=mode,
streams=n_streams,
chunks_per_stream=chunks_per_stream,
total_chunks=total_chunks,
elapsed_min_s=elapsed_min,
elapsed_median_s=elapsed_median,
per_chunk_us=per_chunk_us,
chunks_per_sec=chunks_per_sec,
streams_per_sec=streams_per_sec,
)
def format_result(r: Result) -> str:
return (
f" {r.provider:18s} {r.mode:5s}: "
f"min={r.elapsed_min_s*1000:8.2f} ms "
f"median={r.elapsed_median_s*1000:8.2f} ms "
f"per-chunk={r.per_chunk_us:7.2f} μs "
f"chunks/s={r.chunks_per_sec:>10,.0f} "
f"streams/s={r.streams_per_sec:>8,.1f}"
)
# ---------------------------------------------------------------------------
# CLI
# ---------------------------------------------------------------------------
def main() -> None:
ap = argparse.ArgumentParser(description=__doc__.splitlines()[0])
ap.add_argument(
"--label", required=True, help="Run label (e.g. baseline / optimized)"
)
ap.add_argument("--streams", type=int, default=500, help="Streams per run")
ap.add_argument(
"--chunks",
type=int,
default=200,
help="Text chunks per stream (excl. finish chunk)",
)
ap.add_argument("--warmup", type=int, default=2, help="Warmup runs")
ap.add_argument(
"--repeats", type=int, default=5, help="Measured runs (we report min)"
)
ap.add_argument(
"--providers",
default="anthropic,bedrock_invoke,bedrock_converse",
help="Comma-separated provider list",
)
ap.add_argument(
"--modes",
default="sync,async",
help="Comma-separated iteration modes (sync/async)",
)
ap.add_argument(
"--json", dest="json_out", help="Write results as JSON to this path"
)
args = ap.parse_args()
providers = [p.strip() for p in args.providers.split(",") if p.strip()]
modes = [m.strip() for m in args.modes.split(",") if m.strip()]
for p in providers:
if p not in PROVIDERS:
raise SystemExit(f"unknown provider {p!r}; choose from {list(PROVIDERS)}")
for m in modes:
if m not in {"sync", "async"}:
raise SystemExit(f"unknown mode {m!r}; choose from sync/async")
print(
f"\n=== label={args.label} streams={args.streams} chunks/stream={args.chunks} "
f"warmup={args.warmup} repeats={args.repeats} (min reported) ==="
)
results: List[Result] = []
for provider_key in providers:
for mode in modes:
r = run_case(
label=args.label,
provider_key=provider_key,
mode=mode,
chunks_per_stream=args.chunks,
n_streams=args.streams,
repeats=args.repeats,
warmup=args.warmup,
)
results.append(r)
print(format_result(r))
if args.json_out:
with open(args.json_out, "w", encoding="utf-8") as f:
json.dump([asdict(r) for r in results], f, indent=2)
print(f"\nWrote {len(results)} results to {args.json_out}")
if __name__ == "__main__":
main()

View file

@ -644,6 +644,23 @@ def _should_drop_telemetry_record(request) -> bool:
return not _current_test_records_telemetry()
def _should_passthrough_credential_exchange(request) -> bool:
"""Force the Google OAuth2/STS token mint to run live, never from cassette.
The mint returns a short-lived ``ya29.*`` access token. Recording it lets a
*stale* token replay on a later run; litellm caches it (the recorded
``expires_in`` keeps ``credentials.expired`` False, so it is never
refreshed) and sends it to a live Vertex/Gemini endpoint, which rejects it
with ``ACCESS_TOKEN_EXPIRED``. The token body carries nothing a test asserts
on, so always mint it live: returning ``None`` from ``before_record_request``
makes vcrpy neither store nor replay the call. Inert during
``Cassette._load`` for the same reason as ``_should_drop_telemetry_record``.
"""
if _vcr_load_in_progress():
return False
return _is_credential_exchange_request(request)
# Google APIs (Vertex AI, Gemini, OAuth2/STS). Auth is a ``ya29.*`` OAuth2
# access token minted fresh on every run, so the per-request key fingerprint
# rotates and never matches a recording. The logical credential — the GCP
@ -931,6 +948,8 @@ def _before_record_request(request):
# store the interaction; the request passes through live (fire-and-forget).
if _should_drop_telemetry_record(request):
return None
if _should_passthrough_credential_exchange(request):
return None
headers = getattr(request, "headers", None)
if headers is None:
return request

View file

@ -376,8 +376,23 @@ class LicenseChecker:
all_compliant = True
for req in requirements:
# Prefer a lower-bound/exact version (a real released version) for the
# PyPI license lookup. ``next(iter(req.specifier))`` returns an
# arbitrary clause; for a range like ``>=1.0,<2.0`` that can be the
# upper bound (``2.0``) — a version that may not exist on PyPI and
# would 404 to an "unknown" license.
try:
version = next(iter(req.specifier)).version if req.specifier else None
floor_versions = [
spec.version
for spec in req.specifier
if spec.operator in (">=", "==", "===", "~=", ">")
]
if floor_versions:
version = floor_versions[0]
else:
version = (
next(iter(req.specifier)).version if req.specifier else None
)
except StopIteration:
version = None

View file

@ -0,0 +1,42 @@
"""Shared fixtures for guardrail apply_guardrail tests."""
from contextlib import contextmanager
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@contextmanager
def _mock_proxy_logging():
"""Patch the proxy-server globals that apply_guardrail imports at call time."""
mock_proxy_logging = MagicMock()
mock_proxy_logging.post_call_success_hook = AsyncMock(return_value=None)
mock_proxy_logging.post_call_failure_hook = AsyncMock(return_value=None)
mock_logging_obj = MagicMock()
mock_logging_obj.async_success_handler = AsyncMock(return_value=None)
mock_logging_obj.async_failure_handler = AsyncMock(return_value=None)
mock_logging_obj.success_handler = MagicMock(return_value=None)
mock_logging_obj.failure_handler = MagicMock(return_value=None)
mock_logging_obj.model_call_details = {}
with (
patch(
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing"
) as mock_proc_cls,
patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging),
patch("litellm.proxy.proxy_server.general_settings", {}),
patch("litellm.proxy.proxy_server.proxy_config", MagicMock()),
patch("litellm.proxy.proxy_server.version", "0.0.0"),
):
mock_proc = MagicMock()
mock_proc.common_processing_pre_call_logic = AsyncMock(
return_value=({}, mock_logging_obj)
)
mock_proc_cls.return_value = mock_proc
yield mock_proxy_logging
@pytest.fixture
def mock_proxy_logging_ctx():
"""Return the proxy-logging context manager factory for use as `with ctx():`."""
return _mock_proxy_logging

View file

@ -18,14 +18,19 @@ from litellm.types.guardrails import ApplyGuardrailRequest, ApplyGuardrailRespon
@pytest.mark.asyncio
async def test_apply_guardrail_endpoint_returns_correct_response():
async def test_apply_guardrail_endpoint_returns_correct_response(
mock_proxy_logging_ctx,
):
"""Test that apply_guardrail endpoint returns ApplyGuardrailResponse object"""
from litellm.proxy.guardrails.guardrail_endpoints import apply_guardrail
# Mock the guardrail registry
with patch(
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY"
) as mock_registry:
with (
patch(
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY"
) as mock_registry,
mock_proxy_logging_ctx(),
):
# Create a mock guardrail
mock_guardrail = Mock(spec=CustomGuardrail)
# Apply guardrail returns GenericGuardrailAPIInputs (dict with texts key)
@ -49,7 +54,9 @@ async def test_apply_guardrail_endpoint_returns_correct_response():
# Call the endpoint
response = await apply_guardrail(
request=request, user_api_key_dict=user_api_key_dict
fastapi_request=Mock(),
request=request,
user_api_key_dict=user_api_key_dict,
)
# Verify the response is of the correct type
@ -65,15 +72,18 @@ async def test_apply_guardrail_endpoint_returns_correct_response():
@pytest.mark.asyncio
async def test_apply_guardrail_endpoint_guardrail_not_found():
async def test_apply_guardrail_endpoint_guardrail_not_found(mock_proxy_logging_ctx):
"""Test that apply_guardrail endpoint raises exception when guardrail not found"""
from litellm.proxy._types import ProxyException
from litellm.proxy.guardrails.guardrail_endpoints import apply_guardrail
# Mock the guardrail registry to return None
with patch(
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY"
) as mock_registry:
with (
patch(
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY"
) as mock_registry,
mock_proxy_logging_ctx(),
):
mock_registry.get_initialized_guardrail_callback.return_value = None
# Create the request
@ -86,26 +96,35 @@ async def test_apply_guardrail_endpoint_guardrail_not_found():
# Verify exception is raised
with pytest.raises(ProxyException) as exc_info:
await apply_guardrail(request=request, user_api_key_dict=user_api_key_dict)
await apply_guardrail(
fastapi_request=Mock(),
request=request,
user_api_key_dict=user_api_key_dict,
)
assert "non-existent-guardrail" in exc_info.value.message
assert "not found" in exc_info.value.message
@pytest.mark.asyncio
async def test_apply_guardrail_endpoint_with_presidio_guardrail():
async def test_apply_guardrail_endpoint_with_presidio_guardrail(mock_proxy_logging_ctx):
"""Test apply_guardrail endpoint with a Presidio-like guardrail"""
from litellm.proxy.guardrails.guardrail_endpoints import apply_guardrail
# Mock the guardrail registry
with patch(
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY"
) as mock_registry:
with (
patch(
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY"
) as mock_registry,
mock_proxy_logging_ctx(),
):
# Create a mock guardrail that simulates Presidio behavior
mock_guardrail = Mock(spec=CustomGuardrail)
# Simulate masking PII entities - returns GenericGuardrailAPIInputs (dict with texts key)
mock_guardrail.apply_guardrail = AsyncMock(
return_value={"texts": ["My name is [PERSON] and my email is [EMAIL_ADDRESS]"]}
return_value={
"texts": ["My name is [PERSON] and my email is [EMAIL_ADDRESS]"]
}
)
# Configure the registry to return our mock guardrail
@ -124,7 +143,9 @@ async def test_apply_guardrail_endpoint_with_presidio_guardrail():
# Call the endpoint
response = await apply_guardrail(
request=request, user_api_key_dict=user_api_key_dict
fastapi_request=Mock(),
request=request,
user_api_key_dict=user_api_key_dict,
)
# Verify the response is of the correct type
@ -138,14 +159,17 @@ async def test_apply_guardrail_endpoint_with_presidio_guardrail():
@pytest.mark.asyncio
async def test_apply_guardrail_endpoint_without_optional_params():
async def test_apply_guardrail_endpoint_without_optional_params(mock_proxy_logging_ctx):
"""Test apply_guardrail endpoint without optional language and entities parameters"""
from litellm.proxy.guardrails.guardrail_endpoints import apply_guardrail
# Mock the guardrail registry
with patch(
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY"
) as mock_registry:
with (
patch(
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY"
) as mock_registry,
mock_proxy_logging_ctx(),
):
# Create a mock guardrail
mock_guardrail = Mock(spec=CustomGuardrail)
# Returns GenericGuardrailAPIInputs (dict with texts key)
@ -166,7 +190,9 @@ async def test_apply_guardrail_endpoint_without_optional_params():
# Call the endpoint
response = await apply_guardrail(
request=request, user_api_key_dict=user_api_key_dict
fastapi_request=Mock(),
request=request,
user_api_key_dict=user_api_key_dict,
)
# Verify the response is of the correct type

View file

@ -4,7 +4,7 @@ Test the Bedrock guardrail apply_guardrail functionality
import os
import sys
from unittest.mock import AsyncMock, patch
from unittest.mock import AsyncMock, Mock, patch
import pytest
@ -153,7 +153,7 @@ async def test_bedrock_apply_guardrail_api_failure():
@pytest.mark.asyncio
async def test_bedrock_apply_guardrail_endpoint_integration():
async def test_bedrock_apply_guardrail_endpoint_integration(mock_proxy_logging_ctx):
"""Test the full endpoint integration with Bedrock guardrail"""
from litellm.proxy.guardrails.guardrail_endpoints import apply_guardrail
@ -165,9 +165,12 @@ async def test_bedrock_apply_guardrail_endpoint_integration():
)
# Mock the guardrail registry
with patch(
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY"
) as mock_registry:
with (
patch(
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY"
) as mock_registry,
mock_proxy_logging_ctx(),
):
# Mock the make_bedrock_api_request method
with patch.object(
guardrail, "make_bedrock_api_request", new_callable=AsyncMock
@ -194,7 +197,9 @@ async def test_bedrock_apply_guardrail_endpoint_integration():
# Call the endpoint
response = await apply_guardrail(
request=request, user_api_key_dict=user_api_key_dict
fastapi_request=Mock(),
request=request,
user_api_key_dict=user_api_key_dict,
)
# Verify the response

View file

@ -22,6 +22,7 @@ class ModelEntry:
required_env: FrozenSet[str] = field(default_factory=frozenset)
caps: FrozenSet[str] = field(default_factory=frozenset)
fail_reason: Optional[str] = None
bedrock_effort_ceiling: Optional[str] = None
def params(self) -> Dict[str, str]:
return dict(self.extra_params)
@ -59,9 +60,31 @@ _ADAPTIVE_EFFORT_LABEL: Dict[str, str] = {
"max": "max",
}
_EFFORT_RANK: Dict[str, int] = {
"low": 0,
"medium": 1,
"high": 2,
"max": 3,
"xhigh": 4,
}
_BAD_REQUEST_EFFORTS: FrozenSet[str] = frozenset({"disabled", "invalid", ""})
def _bedrock_clamps_effort(model: "ModelEntry", effort: str) -> bool:
"""Whether Bedrock will clamp ``effort`` down to ``bedrock_effort_ceiling``.
Bedrock chat/messages paths clamp unsupported high tiers (e.g. ``xhigh``
on Opus 4.6) to the model's ceiling rather than rejecting them, so the
missing native capability is OK — the wire effort just degrades.
"""
if model.bedrock_effort_ceiling is None:
return False
if effort not in _EFFORT_RANK or model.bedrock_effort_ceiling not in _EFFORT_RANK:
return False
return _EFFORT_RANK[effort] > _EFFORT_RANK[model.bedrock_effort_ceiling]
def expected(model: ModelEntry, effort: str) -> CellExpectation:
if effort in ("__omit__", "none"):
if model.mode == "budget":
@ -73,14 +96,20 @@ def expected(model: ModelEntry, effort: str) -> CellExpectation:
if effort in ("xhigh", "max"):
cap = f"supports_{effort}_reasoning_effort"
if cap not in model.caps:
if cap not in model.caps and not _bedrock_clamps_effort(model, effort):
return CellExpectation(status=400, thinking_type=OMIT)
if model.mode == "adaptive":
wire_effort = _ADAPTIVE_EFFORT_LABEL[effort]
if model.bedrock_effort_ceiling is not None:
wire_rank = _EFFORT_RANK[wire_effort]
ceiling_rank = _EFFORT_RANK[model.bedrock_effort_ceiling]
if wire_rank > ceiling_rank:
wire_effort = model.bedrock_effort_ceiling
return CellExpectation(
status=200,
thinking_type="adaptive",
output_config_effort=_ADAPTIVE_EFFORT_LABEL[effort],
output_config_effort=wire_effort,
)
return CellExpectation(
@ -219,6 +248,7 @@ BEDROCK_CONVERSE_MODELS: Tuple[ModelEntry, ...] = (
extra_params=(("aws_region_name", "us-east-1"),),
required_env=_BEDROCK_REQ,
caps=_CAPS_4_6,
bedrock_effort_ceiling="max",
),
ModelEntry(
alias="bedrock-claude-sonnet-4-6",
@ -247,6 +277,7 @@ BEDROCK_INVOKE_CHAT_MODELS: Tuple[ModelEntry, ...] = (
extra_params=(("aws_region_name", "us-east-1"),),
required_env=_BEDROCK_REQ,
caps=_CAPS_4_6,
bedrock_effort_ceiling="max",
),
ModelEntry(
alias="bedrock-invoke-claude-sonnet-4-6",

View file

@ -21,11 +21,13 @@ sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "..",
from tests._vcr_conftest_common import ( # noqa: E402
VCR_FIXED_MULTIPART_BOUNDARY,
VCR_IMAGE_B64_PLACEHOLDER,
_before_record_request,
_normalize_multipart_boundary,
_should_passthrough_credential_exchange,
_strip_image_b64_payloads,
_vcr_load_guard,
)
# ---------------------------------------------------------------------------
# Image b64 stripper
# ---------------------------------------------------------------------------
@ -218,3 +220,55 @@ def test_normalize_multipart_handles_quoted_boundary():
_normalize_multipart_boundary(req)
assert b"quoted-boundary" not in req.body
assert VCR_FIXED_MULTIPART_BOUNDARY.encode("utf-8") in req.body
# ---------------------------------------------------------------------------
# Credential-exchange passthrough (Google OAuth2/STS token mint must run live)
# ---------------------------------------------------------------------------
def _oauth_token_request() -> Request:
return Request(
method="POST",
uri="https://oauth2.googleapis.com/token",
body=b"assertion=eyJhbGciOiJSUzI1NiJ9.signed-jwt&grant_type=urn",
headers={"content-type": "application/x-www-form-urlencoded"},
)
def test_before_record_request_drops_oauth_token_mint():
# The token mint must never be stored or replayed, else a stale ya29.* token
# gets sent to a live Vertex/Gemini endpoint -> ACCESS_TOKEN_EXPIRED.
assert _before_record_request(_oauth_token_request()) is None
def test_before_record_request_keeps_normal_request():
req = Request(
method="POST",
uri="https://api.openai.com/v1/chat/completions",
body=b'{"model":"gpt-4o"}',
headers={"content-type": "application/json"},
)
assert _before_record_request(req) is req
def test_credential_exchange_passthrough_inert_during_cassette_load():
# During Cassette._load stored episodes are replayed through this hook;
# dropping there would mutate the cassette on read. The guard makes it inert.
_vcr_load_guard.active = True
try:
assert _should_passthrough_credential_exchange(_oauth_token_request()) is False
assert _before_record_request(_oauth_token_request()) is not None
finally:
_vcr_load_guard.active = False
def test_credential_exchange_passthrough_covers_sts_and_metadata_hosts():
for host in ("sts.googleapis.com", "metadata.google.internal", "169.254.169.254"):
req = Request(
method="POST",
uri=f"https://{host}/token",
body=b"grant_type=urn",
headers={},
)
assert _should_passthrough_credential_exchange(req) is True

View file

@ -853,6 +853,18 @@ def test_personal_key_generation_check():
{"tags": ["old_tag"]},
{"metadata": {"tags": ["old_tag"], "enforced_params": ["metadata.tags"]}},
),
(
{"disable_global_guardrails": True},
{},
{},
{"metadata": {"disable_global_guardrails": True}},
),
(
{"disable_global_guardrails": False},
{},
{"disable_global_guardrails": True},
{"metadata": {"disable_global_guardrails": False}},
),
],
)
def test_prepare_metadata_fields(

View file

@ -1,3 +1,4 @@
import json
import sys
from types import SimpleNamespace
from unittest.mock import AsyncMock
@ -91,8 +92,9 @@ async def test_should_not_mutate_dict_container_response_when_recording_owner(
assert returned == {"id": "cntr_provider", "object": "container"}
data = table.create.await_args.kwargs["data"]
assert data["file_object"]["custom_llm_provider"] == "openai"
assert data["file_object"]["provider_container_id"] == "cntr_provider"
file_obj = json.loads(data["file_object"])
assert file_obj["custom_llm_provider"] == "openai"
assert file_obj["provider_container_id"] == "cntr_provider"
@pytest.mark.asyncio
@ -913,3 +915,195 @@ async def test_admin_with_identity_records_container_ownership(monkeypatch):
table.create.assert_awaited_once()
created_data = table.create.await_args.kwargs["data"]
assert created_data["created_by"] == "proxy-admin"
@pytest.mark.asyncio
async def test_should_record_containers_from_responses_output_for_service_account(
monkeypatch,
):
table = AsyncMock()
table.find_unique.return_value = None
prisma_client = SimpleNamespace(
db=SimpleNamespace(litellm_managedobjecttable=table)
)
monkeypatch.setattr(
ownership,
"_get_prisma_client",
AsyncMock(return_value=prisma_client),
)
auth = UserAPIKeyAuth(team_id="team-1")
encoded_container_id = (
"cntr_bGl0ZWxsbTpjdXN0b21fbGxtX3Byb3ZpZGVyOmF6dXJlO21vZGVsX2lkOmR"
"lZi0xMjM7Y29udGFpbmVyX2lkOmNudHJfbmF0aXZl"
)
responses_payload = {
"output": [
{
"type": "message",
"content": [
{
"type": "output_text",
"annotations": [
{
"type": "container_file_citation",
"container_id": encoded_container_id,
"file_id": "cfile_abc",
}
],
}
],
}
],
"_hidden_params": {"custom_llm_provider": "azure"},
}
await ownership.record_container_owners_from_responses_response(
response=responses_payload,
user_api_key_dict=auth,
)
table.create.assert_awaited_once()
created_data = table.create.await_args.kwargs["data"]
assert created_data["created_by"] == "team:team-1"
assert created_data["unified_object_id"] == encoded_container_id
@pytest.mark.asyncio
async def test_service_account_can_access_container_after_responses_tracking(
monkeypatch,
):
encoded_container_id = (
"cntr_bGl0ZWxsbTpjdXN0b21fbGxtX3Byb3ZpZGVyOmF6dXJlO21vZGVsX2lkOmR"
"lZi0xMjM7Y29udGFpbmVyX2lkOmNudHJfbmF0aXZl"
)
table = AsyncMock()
table.find_unique.return_value = None
prisma_client = SimpleNamespace(
db=SimpleNamespace(litellm_managedobjecttable=table)
)
monkeypatch.setattr(
ownership,
"_get_prisma_client",
AsyncMock(return_value=prisma_client),
)
auth = UserAPIKeyAuth(team_id="team-1")
await ownership.record_container_owners_from_responses_response(
response={
"output": [
{
"type": "code_interpreter_call",
"container_id": encoded_container_id,
}
],
"_hidden_params": {"custom_llm_provider": "azure"},
},
user_api_key_dict=auth,
)
original_id, provider = await ownership.assert_user_can_access_container(
container_id=encoded_container_id,
user_api_key_dict=auth,
custom_llm_provider="azure",
)
assert original_id == "cntr_native"
assert provider == "azure"
@pytest.mark.asyncio
async def test_should_record_container_ownership_after_streaming_responses_finish(
monkeypatch,
):
"""Streaming /v1/responses calls return through the
``select_data_generator`` branch and never reach the non-streaming
container-ownership tail. The wrapper must read
``completed_response`` off the upstream iterator once iteration
finishes and write the row, otherwise code-interpreter containers
created during the stream stay unregistered and follow-up file API
calls 403.
"""
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
encoded_container_id = (
"cntr_bGl0ZWxsbTpjdXN0b21fbGxtX3Byb3ZpZGVyOmF6dXJlO21vZGVsX2lkOmR"
"lZi0xMjM7Y29udGFpbmVyX2lkOmNudHJfbmF0aXZl"
)
response_body = SimpleNamespace(
output=[
SimpleNamespace(
type="code_interpreter_call",
container_id=encoded_container_id,
code_interpreter_call=None,
)
]
)
stream_response = SimpleNamespace(
completed_response=SimpleNamespace(response=response_body),
_hidden_params={"custom_llm_provider": "azure"},
)
async def fake_sse_generator():
yield "data: chunk-1\n\n"
yield "data: chunk-2\n\n"
table = AsyncMock()
table.find_unique.return_value = None
prisma_client = SimpleNamespace(
db=SimpleNamespace(litellm_managedobjecttable=table)
)
monkeypatch.setattr(
ownership,
"_get_prisma_client",
AsyncMock(return_value=prisma_client),
)
auth = UserAPIKeyAuth(team_id="team-1")
wrapped = (
ProxyBaseLLMRequestProcessing._wrap_responses_stream_for_container_ownership(
original_stream_response=stream_response,
wrapped_generator=fake_sse_generator(),
user_api_key_dict=auth,
)
)
chunks = [chunk async for chunk in wrapped]
assert chunks == ["data: chunk-1\n\n", "data: chunk-2\n\n"]
table.create.assert_awaited_once()
created_data = table.create.await_args.kwargs["data"]
assert created_data["created_by"] == "team:team-1"
assert created_data["unified_object_id"] == encoded_container_id
@pytest.mark.asyncio
async def test_streaming_ownership_wrap_no_op_when_stream_did_not_complete(
monkeypatch,
):
"""If the stream errored before ``response.completed``,
``completed_response`` is ``None`` — we must skip the ownership
write rather than crash the response generator."""
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
stream_response = SimpleNamespace(completed_response=None)
async def fake_sse_generator():
yield "data: chunk-1\n\n"
record = AsyncMock()
monkeypatch.setattr(
ownership,
"record_container_owners_from_responses_response",
record,
)
wrapped = (
ProxyBaseLLMRequestProcessing._wrap_responses_stream_for_container_ownership(
original_stream_response=stream_response,
wrapped_generator=fake_sse_generator(),
user_api_key_dict=UserAPIKeyAuth(user_id="user-1"),
)
)
chunks = [chunk async for chunk in wrapped]
assert chunks == ["data: chunk-1\n\n"]
record.assert_not_awaited()

View file

@ -3,7 +3,7 @@ import time
from unittest.mock import AsyncMock
import pytest
from httpx import Response
from httpx import Request, Response
from litellm.integrations.datadog.datadog_cost_management import (
DatadogCostManagementLogger,
@ -167,3 +167,230 @@ async def test_async_send_batch(clean_env):
content = json.loads(call_args[1]["content"])
assert content[0]["ProviderName"] == "openai"
assert content[0]["BilledCost"] == 0.01
_PUT_REQUEST = Request("PUT", "https://api.test.datadoghq.com/api/v2/cost/custom_costs")
@pytest.mark.asyncio
async def test_async_send_batch_clears_queue_on_success(clean_env):
"""Bug 1 regression: log_queue must be empty after a successful upload."""
logger = DatadogCostManagementLogger()
logger.async_client = AsyncMock()
logger.async_client.put.return_value = Response(
202, json={"status": "ok"}, request=_PUT_REQUEST
)
logger.log_queue = [
StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4",
response_cost=0.01,
startTime=time.time(),
)
]
await logger.async_send_batch()
assert logger.log_queue == []
@pytest.mark.asyncio
async def test_async_send_batch_preserves_events_added_during_upload(clean_env):
"""Events appended while the upload is in flight survive (land on the cleared queue)."""
logger = DatadogCostManagementLogger()
later_event = StandardLoggingPayload(
custom_llm_provider="anthropic",
model="claude-3",
response_cost=0.02,
startTime=time.time(),
)
async def slow_put(*args, **kwargs):
logger.log_queue.append(later_event)
return Response(202, json={"status": "ok"}, request=_PUT_REQUEST)
logger.async_client = AsyncMock()
logger.async_client.put.side_effect = slow_put
logger.log_queue = [
StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4",
response_cost=0.01,
startTime=time.time(),
)
]
await logger.async_send_batch()
assert logger.log_queue == [later_event]
@pytest.mark.asyncio
async def test_async_send_batch_requeues_on_upload_failure(clean_env):
"""Failed upload requeues the original batch (no data loss)."""
logger = DatadogCostManagementLogger()
logger.async_client = AsyncMock()
logger.async_client.put.side_effect = Exception("boom")
original = StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4",
response_cost=0.01,
startTime=time.time(),
)
logger.log_queue = [original]
await logger.async_send_batch()
assert logger.log_queue == [original]
@pytest.mark.asyncio
async def test_extract_tags_emits_canonical_focus_dimensions(clean_env):
"""provider, model, model_id always emitted regardless of cost_tag_keys."""
logger = DatadogCostManagementLogger()
log = StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4o",
model_id="router-id-123",
response_cost=0.01,
startTime=time.time(),
)
tags = logger._extract_tags(log)
assert tags["provider"] == "openai"
assert tags["model"] == "gpt-4o"
assert tags["model_id"] == "router-id-123"
@pytest.mark.asyncio
async def test_extract_tags_allowlist_filters_request_tags(clean_env):
"""Only request_tags whose key is in cost_tag_keys reach the Tags dict."""
logger = DatadogCostManagementLogger(cost_tag_keys=["capability", "tier"])
log = StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4",
response_cost=0.01,
startTime=time.time(),
request_tags=["capability:chat", "tier:gold", "secret:disallowed"],
)
tags = logger._extract_tags(log)
assert tags["capability"] == "chat"
assert tags["tier"] == "gold"
assert "secret" not in tags
@pytest.mark.asyncio
async def test_extract_tags_allowlist_filters_metadata(clean_env):
"""Only metadata keys in cost_tag_keys flow through; others (and dict/list values) are dropped."""
logger = DatadogCostManagementLogger(cost_tag_keys=["capability", "owner"])
log = StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4",
response_cost=0.01,
startTime=time.time(),
metadata={
"capability": "chat",
"owner": "team-x",
"secret_field": "sensitive",
"nested_obj": {"a": 1},
},
)
tags = logger._extract_tags(log)
assert tags["capability"] == "chat"
assert tags["owner"] == "team-x"
assert "secret_field" not in tags
assert "nested_obj" not in tags
@pytest.mark.asyncio
async def test_extract_tags_empty_allowlist_default(clean_env):
"""With no cost_tag_keys, request_tags and arbitrary metadata.* do NOT leak into Tags."""
logger = DatadogCostManagementLogger()
log = StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4",
response_cost=0.01,
startTime=time.time(),
request_tags=["capability:chat"],
metadata={"capability": "chat", "user_api_key_alias": "alice"},
)
tags = logger._extract_tags(log)
assert "capability" not in tags
# Backwards-compat keys still flow:
assert tags["user"] == "alice"
@pytest.mark.asyncio
async def test_extract_tags_nested_metadata_allowlisted(clean_env):
"""spend_logs_metadata and requester_metadata get spread one level under the allowlist."""
logger = DatadogCostManagementLogger(cost_tag_keys=["env", "platform"])
log = StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4",
response_cost=0.01,
startTime=time.time(),
metadata={
"spend_logs_metadata": {"platform": "web", "ignored": "x"},
"requester_metadata": {"env": "prod"},
},
)
tags = logger._extract_tags(log)
assert tags["platform"] == "web"
# "env" is a reserved trusted dimension — requester_metadata.env must NOT
# overwrite the value sourced from get_datadog_env().
assert tags["env"] != "prod"
assert "ignored" not in tags
@pytest.mark.asyncio
async def test_extract_tags_allowlist_cannot_override_reserved_dimensions(clean_env):
"""
Reserved tag keys (env, service, host, pod_name, provider, model, model_id,
team, user, model_group) must not be overwritten by user-controlled
request_tags or metadata, even when listed in cost_tag_keys.
"""
reserved = [
"env",
"service",
"host",
"pod_name",
"provider",
"model",
"model_id",
"team",
"user",
"model_group",
]
logger = DatadogCostManagementLogger(cost_tag_keys=reserved)
metadata_attack = {k: f"attacker-meta-{k}" for k in reserved}
metadata_attack["user_api_key_alias"] = "trusted-user"
metadata_attack["user_api_key_team_alias"] = "trusted-team"
metadata_attack["model_group"] = "trusted-group"
metadata_attack["spend_logs_metadata"] = {
k: f"attacker-spend-{k}" for k in reserved
}
metadata_attack["requester_metadata"] = {k: f"attacker-req-{k}" for k in reserved}
log = StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4",
model_id="router-id-123",
response_cost=0.01,
startTime=time.time(),
request_tags=[f"{k}:attacker-rt-{k}" for k in reserved],
metadata=metadata_attack,
)
tags = logger._extract_tags(log)
# Canonical FOCUS dims keep their trusted (top-level payload) values.
assert tags["provider"] == "openai"
assert tags["model"] == "gpt-4"
assert tags["model_id"] == "router-id-123"
# Backwards-compat trusted dims keep their proxy-controlled metadata values.
assert tags["user"] == "trusted-user"
assert tags["team"] == "trusted-team"
assert tags["model_group"] == "trusted-group"
# No reserved key carries an attacker-supplied prefix from any path.
for k in reserved:
assert not tags[k].startswith("attacker-"), (
f"reserved key {k!r} was overwritten by user-controlled input: "
f"{tags[k]!r}"
)

View file

@ -0,0 +1,397 @@
import os
import sys
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
sys.path.insert(0, os.path.abspath("../.."))
from litellm.integrations.galileo import GalileoObserve
from litellm.types.utils import (
Choices,
EmbeddingResponse,
ImageObject,
ImageResponse,
Message,
ModelResponse,
TextCompletionResponse,
)
@pytest.fixture
def galileo_v2_env(monkeypatch):
monkeypatch.setenv("GALILEO_API_KEY", "test-api-key")
monkeypatch.setenv("GALILEO_PROJECT_ID", "86ff8ebe-a297-4134-b167-748bdd8d2c20")
monkeypatch.setenv("GALILEO_LOG_STREAM_ID", "76c4ea50-8aa3-4771-a0d7-8567b112210f")
monkeypatch.setenv("GALILEO_BASE_URL", "https://api.galileo.ai")
@pytest.mark.asyncio
async def test_galileo_v2_ingest_url_and_headers(galileo_v2_env):
logger = GalileoObserve()
logger.in_memory_records = [
{
"latency_ms": 100,
"status_code": 200,
"input_text": "hi",
"output_text": "hello",
"node_type": "acompletion",
"model": "gpt-5.2",
"num_input_tokens": 1,
"num_output_tokens": 2,
"created_at": "2026-05-25T12:00:00",
}
]
url, payload = logger._get_ingest_request()
assert (
url
== "https://api.galileo.ai/v2/projects/86ff8ebe-a297-4134-b167-748bdd8d2c20/spans"
)
assert payload["log_stream_id"] == "76c4ea50-8aa3-4771-a0d7-8567b112210f"
assert payload["spans"][0]["type"] == "llm"
assert payload["spans"][0]["output"]["content"] == "hello"
assert await logger._ensure_headers() is True
assert logger.headers["Galileo-API-Key"] == "test-api-key"
def test_galileo_v2_span_preserves_message_roles(galileo_v2_env):
record = {
"latency_ms": 1,
"status_code": 200,
"input_text": "fallback",
"output_text": "ok",
"node_type": "acompletion",
"model": "gpt-5.2",
"num_input_tokens": 0,
"num_output_tokens": 0,
"created_at": "2026-05-25T12:00:00",
"messages": [
{"role": "system", "content": "be helpful"},
{"role": "user", "content": "hello"},
],
}
span = GalileoObserve._record_to_v2_span(record)
assert span["input"] == [
{"role": "system", "content": "be helpful"},
{"role": "user", "content": "hello"},
]
def test_galileo_output_text_from_model_response(galileo_v2_env):
logger = GalileoObserve()
response = ModelResponse(
choices=[
Choices(
message=Message(
content="assistant reply",
role="assistant",
annotations=[],
)
)
]
)
output = logger.get_output_str_from_response(response, {"call_type": "acompletion"})
assert output == "assistant reply"
@pytest.mark.asyncio
async def test_galileo_flush_swallows_http_errors(galileo_v2_env):
logger = GalileoObserve()
logger.in_memory_records = [
{
"latency_ms": 1,
"status_code": 200,
"input_text": "a",
"output_text": "b",
"node_type": "acompletion",
"model": "gpt-5.2",
"num_input_tokens": 0,
"num_output_tokens": 0,
"created_at": "2026-05-25T12:00:00",
}
]
with patch.object(
logger.async_httpx_handler, "post", new_callable=AsyncMock
) as mock_post:
mock_post.side_effect = Exception("404 Not Found")
await logger.flush_in_memory_records()
assert len(logger.in_memory_records) == 1
@pytest.mark.asyncio
async def test_galileo_flush_clears_records_on_201(galileo_v2_env):
logger = GalileoObserve()
logger.in_memory_records = [
{
"latency_ms": 1,
"status_code": 200,
"input_text": "a",
"output_text": "b",
"node_type": "acompletion",
"model": "gpt-5.2",
"num_input_tokens": 0,
"num_output_tokens": 0,
"created_at": "2026-05-25T12:00:00",
}
]
mock_response = AsyncMock()
mock_response.is_success = True
mock_response.status_code = 201
with patch.object(
logger.async_httpx_handler, "post", new_callable=AsyncMock
) as mock_post:
mock_post.return_value = mock_response
await logger.flush_in_memory_records()
assert logger.in_memory_records == []
def test_galileo_normalize_base_url_none(monkeypatch):
monkeypatch.delenv("GALILEO_API_KEY", raising=False)
monkeypatch.delenv("GALILEO_BASE_URL", raising=False)
monkeypatch.delenv("GALILEO_PROJECT_ID", raising=False)
logger = GalileoObserve()
assert logger.base_url is None
assert logger._normalize_base_url(None) is None
assert logger._normalize_base_url("https://x.example/") == "https://x.example"
def test_galileo_is_configured_branches(monkeypatch):
monkeypatch.delenv("GALILEO_API_KEY", raising=False)
monkeypatch.delenv("GALILEO_BASE_URL", raising=False)
monkeypatch.delenv("GALILEO_PROJECT_ID", raising=False)
monkeypatch.delenv("GALILEO_USERNAME", raising=False)
monkeypatch.delenv("GALILEO_PASSWORD", raising=False)
no_env = GalileoObserve()
assert no_env._is_configured() is False
monkeypatch.setenv("GALILEO_API_KEY", "k")
monkeypatch.setenv("GALILEO_PROJECT_ID", "p")
v2 = GalileoObserve()
assert v2._is_configured() is True
monkeypatch.delenv("GALILEO_API_KEY", raising=False)
monkeypatch.setenv("GALILEO_USERNAME", "u")
monkeypatch.setenv("GALILEO_PASSWORD", "pw")
monkeypatch.setenv("GALILEO_BASE_URL", "https://galileo.example")
legacy = GalileoObserve()
assert legacy._is_configured() is True
monkeypatch.delenv("GALILEO_PASSWORD", raising=False)
no_pw = GalileoObserve()
assert no_pw._is_configured() is False
def test_galileo_input_messages_fallbacks():
assert GalileoObserve._galileo_input_messages(None, "hi") == [
{"role": "user", "content": "hi"}
]
assert GalileoObserve._galileo_input_messages(
["not-a-dict", {"content": "no role"}], "fallback"
) == [{"role": "user", "content": "fallback"}]
def test_galileo_record_to_v2_span_with_tags_and_offset():
span = GalileoObserve._record_to_v2_span(
{
"latency_ms": 5,
"status_code": 200,
"input_text": "in",
"output_text": "out",
"node_type": "acompletion",
"model": "gpt-5.2",
"num_input_tokens": 1,
"num_output_tokens": 2,
"created_at": "2026-05-25T12:00:00",
"tags": ["t1"],
}
)
assert span["tags"] == ["t1"]
assert span["created_at"].endswith("Z")
offset = GalileoObserve._record_to_v2_span(
{"created_at": "2026-05-25T12:00:00-05:00"}
)
assert offset["created_at"] == "2026-05-25T12:00:00-05:00"
def test_galileo_get_output_str_variants(galileo_v2_env):
logger = GalileoObserve()
assert logger.get_output_str_from_response(None, {}) is None
assert (
logger.get_output_str_from_response(
EmbeddingResponse(), {"call_type": "embedding"}
)
is None
)
text_resp = TextCompletionResponse()
text_resp.choices = [MagicMock(text="text-completion-output")]
assert (
logger.get_output_str_from_response(text_resp, {"call_type": "text_completion"})
== "text-completion-output"
)
image_resp = ImageResponse(data=[ImageObject(url="https://x/y.png")])
assert "y.png" in logger.get_output_str_from_response(image_resp, {})
assert logger.get_output_str_from_response("not-a-supported-type", {}) is None
def test_galileo_get_ingest_request_unconfigured(monkeypatch):
monkeypatch.delenv("GALILEO_API_KEY", raising=False)
monkeypatch.delenv("GALILEO_BASE_URL", raising=False)
monkeypatch.delenv("GALILEO_PROJECT_ID", raising=False)
logger = GalileoObserve()
assert logger._get_ingest_request() is None
def test_galileo_get_ingest_request_legacy(monkeypatch):
monkeypatch.delenv("GALILEO_API_KEY", raising=False)
monkeypatch.setenv("GALILEO_USERNAME", "u")
monkeypatch.setenv("GALILEO_PASSWORD", "pw")
monkeypatch.setenv("GALILEO_BASE_URL", "https://galileo.example/")
monkeypatch.setenv("GALILEO_PROJECT_ID", "proj")
logger = GalileoObserve()
logger.in_memory_records = [{"foo": "bar"}]
url, payload = logger._get_ingest_request()
assert url == "https://galileo.example/projects/proj/observe/ingest"
assert payload == {"records": [{"foo": "bar"}]}
@pytest.mark.asyncio
async def test_galileo_ensure_headers_v2_missing_key(monkeypatch):
monkeypatch.delenv("GALILEO_API_KEY", raising=False)
monkeypatch.setenv("GALILEO_PROJECT_ID", "p")
monkeypatch.setenv("GALILEO_BASE_URL", "https://x")
logger = GalileoObserve()
logger.use_v2_api = True
logger.api_key = None
assert await logger._ensure_headers() is False
@pytest.mark.asyncio
async def test_galileo_ensure_headers_cached(galileo_v2_env):
logger = GalileoObserve()
logger.headers = {"Galileo-API-Key": "already-set"}
assert await logger._ensure_headers() is True
@pytest.mark.asyncio
async def test_galileo_ensure_headers_legacy_login(monkeypatch):
monkeypatch.delenv("GALILEO_API_KEY", raising=False)
monkeypatch.setenv("GALILEO_USERNAME", "u")
monkeypatch.setenv("GALILEO_PASSWORD", "pw")
monkeypatch.setenv("GALILEO_BASE_URL", "https://galileo.example")
monkeypatch.setenv("GALILEO_PROJECT_ID", "p")
logger = GalileoObserve()
login_resp = MagicMock()
login_resp.raise_for_status = MagicMock()
login_resp.json = MagicMock(return_value={"access_token": "tok"})
with patch.object(
logger.async_httpx_handler, "post", new_callable=AsyncMock
) as mock_post:
mock_post.return_value = login_resp
assert await logger._ensure_headers() is True
assert logger.headers["Authorization"] == "Bearer tok"
@pytest.mark.asyncio
async def test_galileo_ensure_headers_legacy_login_failure(monkeypatch):
monkeypatch.delenv("GALILEO_API_KEY", raising=False)
monkeypatch.setenv("GALILEO_USERNAME", "u")
monkeypatch.setenv("GALILEO_PASSWORD", "pw")
monkeypatch.setenv("GALILEO_BASE_URL", "https://galileo.example")
monkeypatch.setenv("GALILEO_PROJECT_ID", "p")
logger = GalileoObserve()
with patch.object(
logger.async_httpx_handler, "post", new_callable=AsyncMock
) as mock_post:
mock_post.side_effect = Exception("boom")
assert await logger._ensure_headers() is False
@pytest.mark.asyncio
async def test_galileo_flush_noop_when_unconfigured(monkeypatch):
monkeypatch.delenv("GALILEO_API_KEY", raising=False)
monkeypatch.delenv("GALILEO_BASE_URL", raising=False)
monkeypatch.delenv("GALILEO_PROJECT_ID", raising=False)
logger = GalileoObserve()
logger.in_memory_records = [{"foo": "bar"}]
await logger.flush_in_memory_records()
assert logger.in_memory_records == [{"foo": "bar"}]
@pytest.mark.asyncio
async def test_galileo_flush_resets_headers_on_401(monkeypatch):
monkeypatch.delenv("GALILEO_API_KEY", raising=False)
monkeypatch.setenv("GALILEO_USERNAME", "u")
monkeypatch.setenv("GALILEO_PASSWORD", "pw")
monkeypatch.setenv("GALILEO_BASE_URL", "https://galileo.example")
monkeypatch.setenv("GALILEO_PROJECT_ID", "p")
logger = GalileoObserve()
logger.headers = {"Authorization": "Bearer stale"}
logger.in_memory_records = [{"records": "x"}]
mock_response = MagicMock()
mock_response.is_success = False
mock_response.status_code = 401
mock_response.text = "unauthorized"
with patch.object(
logger.async_httpx_handler, "post", new_callable=AsyncMock
) as mock_post:
mock_post.return_value = mock_response
await logger.flush_in_memory_records()
assert logger.headers is None
assert logger.in_memory_records == [{"records": "x"}]
@pytest.mark.asyncio
async def test_galileo_async_log_success_appends_and_flushes(galileo_v2_env):
import datetime
logger = GalileoObserve()
response = ModelResponse(
choices=[
Choices(message=Message(content="reply", role="assistant", annotations=[]))
],
usage={"prompt_tokens": 1, "completion_tokens": 2},
)
flushed_url: dict = {}
mock_response = MagicMock()
mock_response.is_success = True
mock_response.status_code = 200
async def fake_post(**kwargs):
flushed_url["url"] = kwargs.get("url")
return mock_response
with patch.object(logger.async_httpx_handler, "post", side_effect=fake_post):
await logger.async_log_success_event(
kwargs={
"call_type": "acompletion",
"model": "gpt",
"messages": [{"role": "user", "content": "hi"}],
},
response_obj=response,
start_time=datetime.datetime(2026, 5, 25, 12, 0, 0),
end_time=datetime.datetime(2026, 5, 25, 12, 0, 1),
)
assert "/v2/projects/" in flushed_url["url"]
assert logger.in_memory_records == []

View file

@ -476,6 +476,50 @@ async def test_transcription_captured_in_backend_to_client():
assert logging_obj.model_call_details["messages"] == streaming.input_messages
@pytest.mark.asyncio
async def test_client_ack_caches_setup_to_prevent_duplicate_session_update_setup():
websocket = MagicMock()
backend_ws = MagicMock()
logging_obj = MagicMock()
logging_obj.pre_call = MagicMock()
# Two session.update messages arrive before setupComplete round-trip.
websocket.receive_text = AsyncMock(
side_effect=[
json.dumps({"type": "session.update", "session": {"tools": []}}),
json.dumps({"type": "session.update", "session": {"tools": []}}),
Exception("client done"),
]
)
provider_config = MagicMock()
def _transform(message: str, model: str, session_configuration_request=None):
if session_configuration_request is None:
return [json.dumps({"setup": {"model": "models/gemini-2.5-flash"}})]
return []
provider_config.transform_realtime_request = MagicMock(side_effect=_transform)
backend_ws.send = AsyncMock()
streaming = RealTimeStreaming(
websocket=websocket,
backend_ws=backend_ws,
logging_obj=logging_obj,
provider_config=provider_config,
model="gemini-2.5-flash",
)
await streaming.client_ack_messages()
# Setup should be forwarded exactly once even with repeated session.update.
assert backend_ws.send.await_count == 1
assert streaming.session_configuration_request is not None
sent_payload = json.loads(backend_ws.send.await_args_list[0].args[0])
assert "setup" in sent_payload
def test_collect_session_tools_from_session_update():
"""
Test that tools from session.update events are collected.
@ -879,6 +923,169 @@ async def test_realtime_text_input_guardrail_blocks_and_returns_error():
litellm.callbacks = [] # cleanup
@pytest.mark.asyncio
async def test_realtime_function_call_output_guardrail_blocks_and_returns_error():
"""
Test that a client-supplied function_call_output whose content triggers a
guardrail is blocked: it is not forwarded to the backend, and an error
event is sent to the client.
"""
from fastapi import HTTPException
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.types.guardrails import GuardrailEventHooks
class BlockingGuardrail(CustomGuardrail):
async def apply_guardrail(
self, inputs, request_data, input_type, logging_obj=None
):
texts = inputs.get("texts", [])
for text in texts:
if "@" in text:
raise HTTPException(
status_code=403,
detail={"error": "email address detected"},
)
return inputs
guardrail = BlockingGuardrail(
guardrail_name="email-blocker",
event_hook=GuardrailEventHooks.pre_call,
default_on=True,
)
litellm.callbacks = [guardrail]
client_ws = MagicMock()
client_ws.send_text = AsyncMock()
backend_ws = MagicMock()
backend_ws.send = AsyncMock()
backend_ws.recv = AsyncMock(side_effect=ConnectionClosed(None, None))
logging_obj = MagicMock()
logging_obj.pre_call = MagicMock()
streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj)
item_create_msg = json.dumps(
{
"type": "conversation.item.create",
"item": {
"type": "function_call_output",
"call_id": "call_123",
"output": "Tool says: my email is test@example.com",
},
}
)
client_ws.receive_text = AsyncMock(
side_effect=[
item_create_msg,
Exception("connection closed"),
]
)
await streaming.client_ack_messages()
sent_texts = [json.loads(c.args[0]) for c in client_ws.send_text.call_args_list]
error_events = [e for e in sent_texts if e.get("type") == "error"]
assert len(error_events) == 1, f"Expected one error event, got: {sent_texts}"
assert error_events[0]["error"]["type"] == "guardrail_violation"
sent_to_backend = [c.args[0] for c in backend_ws.send.call_args_list if c.args]
forwarded_tool_outputs = [
json.loads(m)
for m in sent_to_backend
if isinstance(m, str)
and json.loads(m).get("type") == "conversation.item.create"
and json.loads(m).get("item", {}).get("type") == "function_call_output"
]
# A sanitized placeholder must reach the backend so providers that pair
# every toolCall with a toolResponse (Gemini/Vertex Live) exit their
# pending-tool-call state instead of stalling. The placeholder must NOT
# contain any of the blocked content.
assert len(forwarded_tool_outputs) == 1, (
f"Sanitized function_call_output should be forwarded, got: "
f"{forwarded_tool_outputs}"
)
sanitized_item = forwarded_tool_outputs[0]["item"]
assert sanitized_item["call_id"] == "call_123"
assert "test@example.com" not in sanitized_item["output"]
litellm.callbacks = [] # cleanup
@pytest.mark.asyncio
async def test_realtime_function_call_output_guardrail_allows_clean_output():
"""
Test that a clean function_call_output passes through and reaches the backend
when guardrails are configured.
"""
import litellm
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.types.guardrails import GuardrailEventHooks
class BlockingGuardrail(CustomGuardrail):
async def apply_guardrail(
self, inputs, request_data, input_type, logging_obj=None
):
return inputs
guardrail = BlockingGuardrail(
guardrail_name="noop",
event_hook=GuardrailEventHooks.pre_call,
default_on=True,
)
litellm.callbacks = [guardrail]
client_ws = MagicMock()
client_ws.send_text = AsyncMock()
backend_ws = MagicMock()
backend_ws.send = AsyncMock()
backend_ws.recv = AsyncMock(side_effect=ConnectionClosed(None, None))
logging_obj = MagicMock()
logging_obj.pre_call = MagicMock()
streaming = RealTimeStreaming(client_ws, backend_ws, logging_obj)
item_create_msg = json.dumps(
{
"type": "conversation.item.create",
"item": {
"type": "function_call_output",
"call_id": "call_456",
"output": '{"temperature": 72, "unit": "F"}',
},
}
)
client_ws.receive_text = AsyncMock(
side_effect=[
item_create_msg,
Exception("connection closed"),
]
)
await streaming.client_ack_messages()
sent_to_backend = [c.args[0] for c in backend_ws.send.call_args_list if c.args]
forwarded = [
json.loads(m)
for m in sent_to_backend
if isinstance(m, str)
and json.loads(m).get("type") == "conversation.item.create"
and json.loads(m).get("item", {}).get("type") == "function_call_output"
]
assert (
len(forwarded) == 1
), f"Clean function_call_output should be forwarded, got: {forwarded}"
litellm.callbacks = [] # cleanup
@pytest.mark.asyncio
async def test_realtime_text_input_guardrail_uses_pre_call_mode():
"""
@ -1160,3 +1367,406 @@ async def test_on_violation_end_session_closes_on_first_fail():
assert streaming._violation_count == 1
litellm.callbacks = [] # cleanup
@pytest.mark.asyncio
async def test_provider_path_suppresses_duplicate_session_created_after_synthetic():
client_ws = MagicMock()
client_ws.send_text = AsyncMock()
backend_ws = MagicMock()
backend_ws.recv = AsyncMock(
side_effect=[b'{"setupComplete": {}}', ConnectionClosed(None, None)]
)
backend_ws.send = AsyncMock()
provider_config = MagicMock()
provider_config.transform_realtime_response = MagicMock(
return_value={
"response": [
{
"type": "session.created",
"event_id": "event_1",
"session": {"id": "sess_1", "modalities": ["audio"]},
}
],
"current_output_item_id": None,
"current_response_id": None,
"current_delta_chunks": [],
"current_conversation_id": None,
"current_item_chunks": [],
"current_delta_type": None,
"session_configuration_request": None,
}
)
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_1"
logging_obj.async_success_handler = AsyncMock()
logging_obj.success_handler = MagicMock()
streaming = RealTimeStreaming(
websocket=client_ws,
backend_ws=backend_ws,
logging_obj=logging_obj,
provider_config=provider_config,
model="gemini-2.5-flash",
)
# Simulate synthetic session.created already sent by llm_http_handler.
streaming._session_created_sent_to_client = True
await streaming.backend_to_client_send_messages()
sent_payloads = [json.loads(c.args[0]) for c in client_ws.send_text.call_args_list]
assert not any(
payload.get("type") == "session.created" for payload in sent_payloads
), f"Expected duplicate session.created to be suppressed, got: {sent_payloads}"
@pytest.mark.asyncio
async def test_duplicate_session_created_still_triggers_guardrail_turn_detection_update():
client_ws = MagicMock()
client_ws.send_text = AsyncMock()
backend_ws = MagicMock()
backend_ws.recv = AsyncMock(
side_effect=[b'{"setupComplete": {}}', ConnectionClosed(None, None)]
)
backend_ws.send = AsyncMock()
provider_config = MagicMock()
provider_config.transform_realtime_response = MagicMock(
return_value={
"response": [
{
"type": "session.created",
"event_id": "event_1",
"session": {"id": "sess_1", "modalities": ["audio"]},
}
],
"current_output_item_id": None,
"current_response_id": None,
"current_delta_chunks": [],
"current_conversation_id": None,
"current_item_chunks": [],
"current_delta_type": None,
"session_configuration_request": None,
}
)
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_1"
logging_obj.async_success_handler = AsyncMock()
logging_obj.success_handler = MagicMock()
streaming = RealTimeStreaming(
websocket=client_ws,
backend_ws=backend_ws,
logging_obj=logging_obj,
provider_config=provider_config,
model="gemini-2.5-flash",
)
# Synthetic session.created already sent by llm_http_handler.
streaming._session_created_sent_to_client = True
streaming._has_audio_transcription_guardrails = MagicMock(return_value=True) # type: ignore[method-assign]
streaming._send_to_backend = AsyncMock() # type: ignore[method-assign]
await streaming.backend_to_client_send_messages()
# Duplicate session.created should still cause the one-time guardrail
# turn_detection update to be sent to backend.
assert streaming._send_to_backend.await_count == 1
sent_update = json.loads(streaming._send_to_backend.await_args_list[0].args[0])
assert sent_update["type"] == "session.update"
injected_session = sent_update["session"]
assert injected_session["type"] == "realtime"
assert (
injected_session["audio"]["input"]["turn_detection"]["create_response"] is False
)
@pytest.mark.asyncio
async def test_guardrail_update_respects_idempotency_flag():
"""Verify guardrail turn-detection update uses idempotency flag correctly."""
client_ws = AsyncMock()
backend_ws = MagicMock()
backend_ws.send = AsyncMock()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_1"
logging_obj.async_success_handler = AsyncMock()
logging_obj.success_handler = MagicMock()
provider_config = MagicMock()
provider_config.transform_realtime_request = MagicMock(
side_effect=lambda msg, model, session_config: [msg]
)
streaming = RealTimeStreaming(
websocket=client_ws,
backend_ws=backend_ws,
logging_obj=logging_obj,
provider_config=provider_config,
model="gemini-2.5-flash",
)
streaming._has_audio_transcription_guardrails = MagicMock(return_value=True) # type: ignore[method-assign]
# First call should send the update
assert streaming._guardrail_turn_detection_update_sent is False
await streaming._maybe_send_guardrail_turn_detection_update()
assert streaming._guardrail_turn_detection_update_sent is True
assert backend_ws.send.await_count == 1
# Second call should be a no-op (idempotent)
await streaming._maybe_send_guardrail_turn_detection_update()
assert backend_ws.send.await_count == 1 # Still 1, not 2
@pytest.mark.asyncio
async def test_guardrail_turn_detection_injected_into_first_session_update_deferred_mode():
"""Verify turn_detection is injected into first session.update in deferred mode."""
client_ws = AsyncMock()
client_ws.receive_text = AsyncMock(
side_effect=[
json.dumps(
{
"type": "session.update",
"session": {
"modalities": ["text", "audio"],
"tools": [{"type": "function", "name": "get_weather"}],
},
}
),
ConnectionClosed(None, None),
]
)
backend_ws = MagicMock()
backend_ws.send = AsyncMock()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_1"
logging_obj.async_success_handler = AsyncMock()
logging_obj.success_handler = MagicMock()
provider_config = MagicMock()
transformed_messages = []
def mock_transform(msg, model, session_config):
transformed_messages.append((msg, session_config))
return [msg] # Pass through for simplicity
provider_config.transform_realtime_request = MagicMock(side_effect=mock_transform)
streaming = RealTimeStreaming(
websocket=client_ws,
backend_ws=backend_ws,
logging_obj=logging_obj,
provider_config=provider_config,
model="gemini-2.5-flash",
)
streaming._has_audio_transcription_guardrails = MagicMock(return_value=True) # type: ignore[method-assign]
# Simulate first session.update in deferred mode
await streaming.client_ack_messages()
# Verify turn_detection was injected into the session.update. The
# injection runs before the GA remap, so the create_response flag ends
# up nested under audio.input.turn_detection in the GA-shaped payload.
assert len(transformed_messages) == 1
transformed_msg, session_config = transformed_messages[0]
msg_obj = json.loads(transformed_msg)
assert msg_obj["type"] == "session.update"
session_obj = msg_obj["session"]
injected_turn_detection = session_obj.get("turn_detection") or session_obj.get(
"audio", {}
).get("input", {}).get("turn_detection")
assert injected_turn_detection is not None
assert injected_turn_detection["create_response"] is False
assert streaming._guardrail_turn_detection_update_sent is True
@pytest.mark.asyncio
@pytest.mark.parametrize("existing_turn_detection", [None, "auto", 42, ["server_vad"]])
async def test_guardrail_turn_detection_injection_tolerates_non_dict_value(
existing_turn_detection,
):
"""Client-supplied non-dict turn_detection must not crash client_ack_messages."""
client_ws = AsyncMock()
client_ws.receive_text = AsyncMock(
side_effect=[
json.dumps(
{
"type": "session.update",
"session": {
"modalities": ["text", "audio"],
"turn_detection": existing_turn_detection,
},
}
),
ConnectionClosed(None, None),
]
)
backend_ws = MagicMock()
backend_ws.send = AsyncMock()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_1"
logging_obj.async_success_handler = AsyncMock()
logging_obj.success_handler = MagicMock()
provider_config = MagicMock()
transformed_messages = []
def mock_transform(msg, model, session_config):
transformed_messages.append((msg, session_config))
return [msg]
provider_config.transform_realtime_request = MagicMock(side_effect=mock_transform)
streaming = RealTimeStreaming(
websocket=client_ws,
backend_ws=backend_ws,
logging_obj=logging_obj,
provider_config=provider_config,
model="gemini-2.5-flash",
)
streaming._has_audio_transcription_guardrails = MagicMock(return_value=True) # type: ignore[method-assign]
await streaming.client_ack_messages()
assert len(transformed_messages) == 1
transformed_msg, _ = transformed_messages[0]
msg_obj = json.loads(transformed_msg)
session_obj = msg_obj["session"]
injected_turn_detection = session_obj.get("turn_detection") or session_obj.get(
"audio", {}
).get("input", {}).get("turn_detection")
assert isinstance(injected_turn_detection, dict)
assert injected_turn_detection["create_response"] is False
assert streaming._guardrail_turn_detection_update_sent is True
@pytest.mark.asyncio
@pytest.mark.parametrize(
"client_session",
[
{"turn_detection": {"type": "server_vad", "create_response": True}},
{
"audio": {
"input": {
"turn_detection": {"type": "server_vad", "create_response": True}
}
}
},
],
)
async def test_subsequent_session_update_cannot_reenable_vad_when_guardrails_active(
client_session,
):
"""A subsequent client session.update must not be allowed to flip
``create_response`` back to True once audio transcription guardrails have
disabled VAD auto-response. Covers both the flat beta shape and the
nested GA ``audio.input.turn_detection`` shape.
"""
client_ws = AsyncMock()
client_ws.receive_text = AsyncMock(
side_effect=[
json.dumps({"type": "session.update", "session": client_session}),
ConnectionClosed(None, None),
]
)
backend_ws = MagicMock()
backend_ws.send = AsyncMock()
logging_obj = MagicMock()
logging_obj.litellm_trace_id = "trace_1"
logging_obj.async_success_handler = AsyncMock()
logging_obj.success_handler = MagicMock()
provider_config = MagicMock()
transformed_messages = []
def mock_transform(msg, model, session_config):
transformed_messages.append((msg, session_config))
return [msg]
provider_config.transform_realtime_request = MagicMock(side_effect=mock_transform)
streaming = RealTimeStreaming(
websocket=client_ws,
backend_ws=backend_ws,
logging_obj=logging_obj,
provider_config=provider_config,
model="gemini-2.5-flash",
)
streaming._has_audio_transcription_guardrails = MagicMock(return_value=True) # type: ignore[method-assign]
# Simulate that initial setup + guardrail disable have already happened.
streaming.session_configuration_request = json.dumps({"setup": {"model": "x"}})
streaming._guardrail_turn_detection_update_sent = True
await streaming.client_ack_messages()
assert len(transformed_messages) == 1
forwarded_msg, _ = transformed_messages[0]
msg_obj = json.loads(forwarded_msg)
session_obj = msg_obj["session"]
forwarded_turn_detection = session_obj.get("turn_detection") or session_obj.get(
"audio", {}
).get("input", {}).get("turn_detection")
assert isinstance(forwarded_turn_detection, dict)
assert forwarded_turn_detection["create_response"] is False
@pytest.mark.asyncio
async def test_follow_up_setup_updates_cached_session_configuration_request():
"""A follow-up setup produced by a subsequent session.update must replace
the cached ``session_configuration_request`` so downstream readers
(e.g. modality lookup in ``response.created``) see the latest config."""
client_ws = AsyncMock()
client_ws.receive_text = AsyncMock(
side_effect=[
json.dumps({"type": "session.update", "session": {"tools": []}}),
ConnectionClosed(None, None),
]
)
backend_ws = MagicMock()
backend_ws.send = AsyncMock()
logging_obj = MagicMock()
logging_obj.async_success_handler = AsyncMock()
logging_obj.success_handler = MagicMock()
provider_config = MagicMock()
follow_up_setup = json.dumps(
{
"setup": {
"model": "models/gemini-2.5-flash",
"generationConfig": {"responseModalities": ["TEXT"]},
"tools": [{"function_declarations": []}],
}
}
)
provider_config.transform_realtime_request = MagicMock(
return_value=[follow_up_setup]
)
streaming = RealTimeStreaming(
websocket=client_ws,
backend_ws=backend_ws,
logging_obj=logging_obj,
provider_config=provider_config,
model="gemini-2.5-flash",
)
# Simulate that the original auto-setup was already cached.
streaming.session_configuration_request = json.dumps(
{
"setup": {
"model": "models/gemini-2.5-flash",
"generationConfig": {"responseModalities": ["AUDIO"]},
}
}
)
await streaming.client_ack_messages()
assert streaming.session_configuration_request == follow_up_setup

View file

@ -0,0 +1,508 @@
"""
Tests for CustomStreamWrapper per-chunk behavior across Anthropic,
Bedrock Invoke, and Bedrock Converse: text passthrough, usage stripping,
hidden_params propagation, finish_reason, sync/async parity, and the
per-stream caches (_GCHUNK_FIELDS, _post_streaming_hooks).
"""
import asyncio
import time
from typing import List, Optional
from unittest.mock import MagicMock, patch
import litellm
from litellm.litellm_core_utils.streaming_handler import (
CustomStreamWrapper,
_GCHUNK_FIELDS,
generic_chunk_has_all_required_fields,
)
from litellm.types.utils import (
Delta,
GenericStreamingChunk as GChunk,
ModelResponseStream,
StreamingChoices,
Usage,
)
# ---------------------------------------------------------------------------
# Shared helpers
# ---------------------------------------------------------------------------
def _make_logging_obj(provider: str = "anthropic") -> MagicMock:
logging_obj = MagicMock()
logging_obj.model_call_details = {
"custom_llm_provider": provider,
"litellm_params": {},
}
logging_obj.call_type = "completion"
logging_obj.stream_options = None
logging_obj.messages = [{"role": "user", "content": "hi"}]
logging_obj.completion_start_time = None
logging_obj._llm_caching_handler = None
return logging_obj
def _make_generic_chunk(
text: str,
is_finished: bool = False,
finish_reason: str = "",
usage: Optional[dict] = None,
) -> GChunk:
return GChunk(
text=text,
is_finished=is_finished,
finish_reason=finish_reason,
usage=usage,
index=0,
tool_use=None,
)
def _make_bedrock_converse_chunk(
text: str = "",
finish_reason: str = "",
usage: Optional[Usage] = None,
) -> ModelResponseStream:
"""Simulate what AWSEventStreamDecoder.converse_chunk_parser returns."""
return ModelResponseStream(
choices=[
StreamingChoices(
finish_reason=finish_reason or None,
index=0,
delta=Delta(content=text, role="assistant"),
)
],
id="msg-test",
model="anthropic.claude-3-5-sonnet",
usage=usage,
)
async def _async_iter(chunks: list):
"""Wrap a list as a proper async iterator for use in __anext__ async branch."""
for chunk in chunks:
yield chunk
def _make_wrapper(
chunks: list,
provider: str = "anthropic",
async_stream: bool = False,
) -> CustomStreamWrapper:
logging_obj = _make_logging_obj(provider)
stream = _async_iter(chunks) if async_stream else iter(chunks)
wrapper = CustomStreamWrapper(
completion_stream=stream,
model="claude-3-5-sonnet",
logging_obj=logging_obj,
custom_llm_provider=provider,
)
return wrapper
def _drain_sync(wrapper: CustomStreamWrapper) -> List[ModelResponseStream]:
results = []
for chunk in wrapper:
results.append(chunk)
return results
async def _drain_async(wrapper: CustomStreamWrapper) -> List[ModelResponseStream]:
results = []
async for chunk in wrapper:
results.append(chunk)
return results
# ---------------------------------------------------------------------------
# 1. Module-level _GCHUNK_FIELDS constant
# ---------------------------------------------------------------------------
def test_gchunk_fields_is_frozenset():
"""_GCHUNK_FIELDS must be a frozenset built from GChunk.__annotations__."""
assert isinstance(_GCHUNK_FIELDS, frozenset)
assert _GCHUNK_FIELDS == frozenset(GChunk.__annotations__)
def test_generic_chunk_has_all_required_fields_uses_module_constant(monkeypatch):
"""generic_chunk_has_all_required_fields must use _GCHUNK_FIELDS, not __annotations__.
The check semantics: every key in `chunk` must be a known GChunk field.
This identifies GChunk-shaped dicts (all keys are valid GChunk fields).
"""
valid_chunk = _make_generic_chunk("hello")
assert generic_chunk_has_all_required_fields(valid_chunk) is True
# A dict with an extra unknown key should return False — the unknown key
# is not a GChunk field, so the chunk is not a pure GChunk.
extra_key_chunk = dict(valid_chunk)
extra_key_chunk["unknown_extra_key"] = "value"
assert generic_chunk_has_all_required_fields(extra_key_chunk) is False
# A dict with only known GChunk fields but fewer keys still passes because
# all its keys are valid (subset of GChunk fields).
partial_chunk = {"text": "hi", "is_finished": False}
assert generic_chunk_has_all_required_fields(partial_chunk) is True
# ---------------------------------------------------------------------------
# 2. Cached model name and provider at init time
# ---------------------------------------------------------------------------
def test_cached_model_name_simple():
"""For non-openai providers the cached model name must match the model arg."""
wrapper = _make_wrapper([], provider="anthropic")
assert wrapper._cached_model_name == "claude-3-5-sonnet"
assert wrapper._cached_logging_llm_provider == "anthropic"
def test_cached_model_name_openai_prefix():
"""For openai provider when logging provider differs, model name is prefixed."""
logging_obj = _make_logging_obj(provider="azure")
wrapper = CustomStreamWrapper(
completion_stream=iter([]),
model="gpt-4o",
logging_obj=logging_obj,
custom_llm_provider="openai",
)
assert wrapper._cached_model_name == "azure/gpt-4o"
assert wrapper._cached_logging_llm_provider == "azure"
def test_base_hidden_params_precomputed():
"""_base_hidden_params must be pre-built from _hidden_params at init."""
wrapper = _make_wrapper([], provider="anthropic")
assert "response_cost" in wrapper._base_hidden_params
assert wrapper._base_hidden_params["response_cost"] is None
# Must include all keys from _hidden_params
for k in wrapper._hidden_params:
assert k in wrapper._base_hidden_params
# ---------------------------------------------------------------------------
# 3. Sync path: model_dump() is NOT called on non-usage chunks
# ---------------------------------------------------------------------------
def test_sync_path_no_model_dump_on_text_chunks():
"""
The sync __next__ must NOT call model_dump() on chunks that have no usage.
ModelResponseStream declares `usage` as a field, so a `hasattr` check
would always succeed and trigger the model_dump()+recreate path on every
chunk. The wrapper must check `is not None` instead.
"""
chunks = [
_make_generic_chunk("Hello"),
_make_generic_chunk(" world"),
_make_generic_chunk("", is_finished=True, finish_reason="stop"),
]
wrapper = _make_wrapper(chunks)
model_dump_call_count = 0
original_model_dump = ModelResponseStream.model_dump
def counting_model_dump(self, **kwargs):
nonlocal model_dump_call_count
model_dump_call_count += 1
return original_model_dump(self, **kwargs)
with patch.object(ModelResponseStream, "model_dump", counting_model_dump):
results = _drain_sync(wrapper)
text_chunks = [r for r in results if r.choices and r.choices[0].delta.content]
assert len(text_chunks) >= 2, "Expected at least 2 text chunks"
assert model_dump_call_count <= 1, (
f"model_dump() called {model_dump_call_count} times — "
"usage check is firing on every chunk"
)
# ---------------------------------------------------------------------------
# 4. Sync path: usage chunk is stripped from body but preserved in hidden_params
# ---------------------------------------------------------------------------
def test_sync_path_usage_stripped_from_body_preserved_in_hidden_params():
"""Usage data must be removed from the returned chunk but added to _hidden_params."""
usage_dict = {"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}
chunks = [
_make_generic_chunk("Hello"),
_make_generic_chunk(
"", is_finished=True, finish_reason="stop", usage=usage_dict
),
]
wrapper = _make_wrapper(chunks)
results = _drain_sync(wrapper)
# The usage chunk must be returned (not silently dropped)
finish_chunks = [
r for r in results if r.choices and r.choices[0].finish_reason == "stop"
]
assert finish_chunks, "Finish-reason chunk was not returned"
# The final chunk must carry usage in _hidden_params
final = results[-1]
assert "usage" in final._hidden_params, "usage missing from _hidden_params"
hidden_usage = final._hidden_params["usage"]
assert hidden_usage is not None
# ---------------------------------------------------------------------------
# 5. Async path: usage chunk is stripped from body but preserved in hidden_params
# ---------------------------------------------------------------------------
def test_async_path_usage_stripped_from_body_preserved_in_hidden_params():
"""Async path mirrors sync path for usage handling."""
usage_dict = {"prompt_tokens": 5, "completion_tokens": 15, "total_tokens": 20}
chunks = [
_make_generic_chunk("Hi"),
_make_generic_chunk(
"", is_finished=True, finish_reason="stop", usage=usage_dict
),
]
async def _run():
# async_stream=True forces the real async-for branch of __anext__
wrapper = _make_wrapper(chunks, async_stream=True)
return await _drain_async(wrapper)
results = asyncio.run(_run())
final = results[-1]
assert "usage" in final._hidden_params
assert final._hidden_params["usage"] is not None
# ---------------------------------------------------------------------------
# 6. Bedrock Converse: ModelResponseStream chunks pass through correctly
# ---------------------------------------------------------------------------
def test_bedrock_converse_text_chunks_pass_through():
"""
Bedrock Converse returns ModelResponseStream objects directly.
They should pass through chunk_creator and appear in output unchanged.
"""
chunks = [
_make_bedrock_converse_chunk("Hello"),
_make_bedrock_converse_chunk(" world"),
_make_bedrock_converse_chunk("", finish_reason="end_turn"),
]
wrapper = _make_wrapper(chunks, provider="bedrock")
results = _drain_sync(wrapper)
texts = [
r.choices[0].delta.content
for r in results
if r.choices and r.choices[0].delta.content
]
assert "Hello" in texts or any("Hello" in (t or "") for t in texts)
def test_bedrock_converse_usage_chunk_stripped_and_in_hidden_params():
"""Usage in a Bedrock Converse ModelResponseStream chunk is handled correctly."""
usage = Usage(prompt_tokens=8, completion_tokens=12, total_tokens=20)
chunks = [
_make_bedrock_converse_chunk("Hi"),
_make_bedrock_converse_chunk("", finish_reason="end_turn", usage=usage),
]
wrapper = _make_wrapper(chunks, provider="bedrock")
results = _drain_sync(wrapper)
final = results[-1]
assert "usage" in final._hidden_params
assert final._hidden_params["usage"] is not None
# ---------------------------------------------------------------------------
# 7. Anthropic generic chunk (GChunk) path
# ---------------------------------------------------------------------------
def test_anthropic_generic_chunks_text_pass_through():
"""GChunk text chunks must arrive in the output with correct content."""
chunks = [
_make_generic_chunk("The"),
_make_generic_chunk(" answer"),
_make_generic_chunk("", is_finished=True, finish_reason="stop"),
]
wrapper = _make_wrapper(chunks, provider="anthropic")
results = _drain_sync(wrapper)
texts = [
r.choices[0].delta.content
for r in results
if r.choices and r.choices[0].delta.content
]
assert len(texts) >= 2
def test_anthropic_finish_reason_propagated():
"""finish_reason must be set on the final streaming chunk."""
chunks = [
_make_generic_chunk("Hi"),
_make_generic_chunk("", is_finished=True, finish_reason="stop"),
]
wrapper = _make_wrapper(chunks, provider="anthropic")
results = _drain_sync(wrapper)
finish_reasons = [
r.choices[0].finish_reason
for r in results
if r.choices and r.choices[0].finish_reason
]
assert "stop" in finish_reasons
# ---------------------------------------------------------------------------
# 8. Callback caching: _post_streaming_hooks resolved once per stream
# ---------------------------------------------------------------------------
def test_post_streaming_hooks_cached_after_first_call():
"""
_post_streaming_hooks must be None before the first hook call and a list after.
The same list object must be reused on subsequent calls (not re-built).
"""
wrapper = _make_wrapper([], provider="anthropic")
assert wrapper._post_streaming_hooks is None, "Must be None before first call"
async def _run():
# Simulate hook resolution with an empty callback list
with patch.object(litellm, "callbacks", []):
await wrapper._call_post_streaming_deployment_hook(
MagicMock(spec=ModelResponseStream)
)
first_list = wrapper._post_streaming_hooks
assert isinstance(first_list, list)
# Second call must reuse the same list object
with patch.object(litellm, "callbacks", []):
await wrapper._call_post_streaming_deployment_hook(
MagicMock(spec=ModelResponseStream)
)
assert (
wrapper._post_streaming_hooks is first_list
), "_post_streaming_hooks was rebuilt on second call — caching broken"
asyncio.run(_run())
def test_post_streaming_hooks_filters_correctly():
"""
Only CustomLogger instances must be included; plain callables are excluded.
Note: CustomLogger's base class already defines
async_post_call_streaming_deployment_hook, so ALL CustomLogger subclasses
pass the hasattr() check regardless of whether they override the method.
The filter therefore keeps any CustomLogger instance and drops anything else.
"""
from litellm.integrations.custom_logger import CustomLogger
class MyLogger(CustomLogger):
pass
plain_callable = MagicMock()
wrapper = _make_wrapper([], provider="anthropic")
async def _run():
with patch.object(litellm, "callbacks", [MyLogger(), plain_callable]):
await wrapper._call_post_streaming_deployment_hook(
MagicMock(spec=ModelResponseStream)
)
# plain_callable must be excluded; MyLogger (CustomLogger subclass) included
assert len(wrapper._post_streaming_hooks) == 1
assert isinstance(wrapper._post_streaming_hooks[0], MyLogger)
asyncio.run(_run())
# ---------------------------------------------------------------------------
# 9. model_response_creator: hidden_params built correctly
# ---------------------------------------------------------------------------
def test_model_response_creator_hidden_params_no_chunk():
"""model_response_creator() with no args must include all _base_hidden_params."""
wrapper = _make_wrapper([], provider="anthropic")
response = wrapper.model_response_creator()
assert response._hidden_params.get("response_cost") is None
assert response._hidden_params.get("custom_llm_provider") == "anthropic"
assert "created_at" in response._hidden_params
def test_model_response_creator_hidden_params_caller_merged():
"""When hidden_params are passed by caller, they must be included in result."""
wrapper = _make_wrapper([], provider="anthropic")
caller_params = {"some_key": "some_value"}
response = wrapper.model_response_creator(hidden_params=caller_params)
assert response._hidden_params.get("some_key") == "some_value"
assert response._hidden_params.get("response_cost") is None
def test_model_response_creator_stream_key_stripped():
"""The 'stream' key must be removed from chunk before constructing ModelResponseStream."""
wrapper = _make_wrapper([], provider="anthropic")
chunk = {"stream": True, "choices": []}
# Should not raise even if 'stream' would be an invalid ModelResponseStream field
response = wrapper.model_response_creator(chunk=chunk)
assert response is not None
# ---------------------------------------------------------------------------
# 10. Per-chunk overhead regression: sync path must not regress
# ---------------------------------------------------------------------------
def test_sync_streaming_overhead_not_regressed():
"""
Micro-benchmark: the sync hot path must process 200 text chunks in < 2 s.
This test acts as a canary for gross per-chunk overhead regressions.
It is intentionally generous (2 s) to avoid flakiness on slow CI runners.
"""
n_chunks = 200
chunks = [_make_generic_chunk(f"token-{i}") for i in range(n_chunks)]
chunks.append(_make_generic_chunk("", is_finished=True, finish_reason="stop"))
wrapper = _make_wrapper(chunks, provider="anthropic")
start = time.monotonic()
results = _drain_sync(wrapper)
elapsed = time.monotonic() - start
assert len(results) > 0, "No chunks returned"
assert elapsed < 2.0, (
f"Sync streaming of {n_chunks} chunks took {elapsed:.3f}s — "
"per-chunk overhead regression detected"
)
def test_async_streaming_overhead_not_regressed():
"""
Micro-benchmark for the async path: 200 text chunks in < 2 s.
"""
n_chunks = 200
chunks = [_make_generic_chunk(f"token-{i}") for i in range(n_chunks)]
chunks.append(_make_generic_chunk("", is_finished=True, finish_reason="stop"))
async def _run():
wrapper = _make_wrapper(chunks, provider="anthropic")
start = time.monotonic()
results = await _drain_async(wrapper)
return results, time.monotonic() - start
results, elapsed = asyncio.run(_run())
assert len(results) > 0
assert elapsed < 2.0, (
f"Async streaming of {n_chunks} chunks took {elapsed:.3f}s — "
"per-chunk overhead regression detected"
)

View file

@ -1622,6 +1622,29 @@ def test_effort_output_config_preservation():
assert result["output_config"]["effort"] == "medium"
def test_output_config_format_preservation_and_beta_header():
"""Test that output_config.format is preserved and treated as structured output."""
config = AnthropicConfig()
output_format = {
"type": "json_schema",
"schema": {"type": "object", "properties": {"answer": {"type": "string"}}},
}
optional_params = {"output_config": {"format": output_format, "effort": "xhigh"}}
result = config.transform_request(
model="claude-opus-4-7",
messages=[{"role": "user", "content": "Test"}],
optional_params=optional_params,
litellm_params={},
headers={},
)
headers = config.update_headers_with_optional_anthropic_beta({}, optional_params)
assert result["output_config"]["format"] == output_format
assert result["output_config"]["effort"] == "xhigh"
assert "structured-outputs-2025-11-13" in headers["anthropic-beta"]
def test_effort_beta_header_injection():
"""Test that effort beta header is automatically added when output_config is detected."""
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
@ -1648,7 +1671,7 @@ def test_effort_validation():
messages = [{"role": "user", "content": "Test"}]
# Valid values should work
# Valid values should work (xhigh is Opus 4.7+ only, not 4.5)
for effort in ["high", "medium", "low"]:
optional_params = {"output_config": {"effort": effort}}
result = config.transform_request(
@ -2513,14 +2536,14 @@ def test_reasoning_effort_accepts_dict_shape_for_adaptive_model(reasoning_effort
)
# thinking must be set (adaptive for 4.6+)
assert "thinking" in result, (
f"thinking missing for reasoning_effort={reasoning_effort_value!r}"
)
assert (
"thinking" in result
), f"thinking missing for reasoning_effort={reasoning_effort_value!r}"
assert result["thinking"]["type"] == "adaptive"
# output_config must carry the mapped effort
assert "output_config" in result, (
f"output_config missing for reasoning_effort={reasoning_effort_value!r}"
)
assert (
"output_config" in result
), f"output_config missing for reasoning_effort={reasoning_effort_value!r}"
assert result["output_config"]["effort"] == "low"
@ -2532,7 +2555,9 @@ def test_reasoning_effort_accepts_dict_shape_for_adaptive_model(reasoning_effort
{"effort": "low", "summary": "concise"},
],
)
def test_reasoning_effort_accepts_dict_shape_for_non_adaptive_model(reasoning_effort_value):
def test_reasoning_effort_accepts_dict_shape_for_non_adaptive_model(
reasoning_effort_value,
):
"""
Non-adaptive (pre-4.6) branch: dict-shape reasoning_effort must still map
to ``thinking.type='enabled'`` + ``budget_tokens``. ``output_config`` must
@ -2547,9 +2572,9 @@ def test_reasoning_effort_accepts_dict_shape_for_non_adaptive_model(reasoning_ef
drop_params=False,
)
assert "thinking" in result, (
f"thinking missing for reasoning_effort={reasoning_effort_value!r}"
)
assert (
"thinking" in result
), f"thinking missing for reasoning_effort={reasoning_effort_value!r}"
assert result["thinking"]["type"] == "enabled"
assert "budget_tokens" in result["thinking"]
assert result["thinking"]["budget_tokens"] > 0
@ -2582,12 +2607,12 @@ def test_reasoning_effort_unparseable_dict_is_dropped(bad_value):
model="claude-sonnet-4-6-20260219",
drop_params=False,
)
assert "thinking" not in result, (
f"thinking should not be set for bad value {bad_value!r}"
)
assert "output_config" not in result, (
f"output_config should not be set for bad value {bad_value!r}"
)
assert (
"thinking" not in result
), f"thinking should not be set for bad value {bad_value!r}"
assert (
"output_config" not in result
), f"output_config should not be set for bad value {bad_value!r}"
@pytest.mark.parametrize(

View file

@ -48,6 +48,39 @@ def test_output_format_supported_and_transforms_correctly():
assert "structured-outputs-2025-11-13" in headers["anthropic-beta"]
def test_output_config_format_supported_and_transforms_correctly():
"""Test that output_config.format is preserved and adds the structured-output beta."""
config = AnthropicMessagesConfig()
supported_params = config.get_supported_anthropic_messages_params("claude-opus-4-7")
assert "output_config" in supported_params
output_format = {
"type": "json_schema",
"schema": {"type": "object", "properties": {"result": {"type": "string"}}},
}
optional_params = {
"max_tokens": 1024,
"output_config": {"format": output_format, "effort": "xhigh"},
}
headers = {}
result = config.transform_anthropic_messages_request(
model="claude-opus-4-7",
messages=[{"role": "user", "content": "test"}],
anthropic_messages_optional_request_params=optional_params.copy(),
litellm_params={},
headers=headers,
)
headers = config._update_headers_with_anthropic_beta(headers, optional_params)
assert result["output_config"]["format"] == output_format
assert result["output_config"]["effort"] == "xhigh"
assert "anthropic-beta" in headers
assert "structured-outputs-2025-11-13" in headers["anthropic-beta"]
def test_output_format_works_with_bedrock_and_azure():
"""Test that output_format works with Bedrock and Azure Foundry models."""
config = AnthropicMessagesConfig()

View file

@ -6,6 +6,9 @@ from litellm.llms.anthropic.common_utils import AnthropicError
from litellm.llms.anthropic.experimental_pass_through.messages.transformation import (
AnthropicMessagesConfig,
)
from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import (
AmazonAnthropicClaudeMessagesConfig,
)
@pytest.mark.parametrize(
@ -102,7 +105,6 @@ def test_invalid_reasoning_effort_raises_400(bad_effort):
"model,bad_effort",
[
("claude-opus-4-6", "xhigh"),
("bedrock/invoke/us.anthropic.claude-opus-4-6-v1", "xhigh"),
("claude-sonnet-4-6", "xhigh"),
],
)
@ -123,6 +125,56 @@ def test_reasoning_effort_unsupported_tier_raises_400_messages(model, bad_effort
assert "not supported by this model" in str(exc_info.value)
@pytest.mark.parametrize(
"model,effort,expected_effort",
[
("invoke/us.anthropic.claude-opus-4-6-v1", "xhigh", "max"),
("invoke/us.anthropic.claude-opus-4-6-v1", "max", "max"),
("invoke/us.anthropic.claude-opus-4-6-v1", "high", "high"),
("invoke/us.anthropic.claude-opus-4-7", "xhigh", "xhigh"),
],
)
def test_bedrock_invoke_messages_clamps_effort_to_ceiling(
model, effort, expected_effort
):
"""Bedrock Invoke /v1/messages degrades effort to the model's ceiling.
Claude Code "goal mode" sends ``xhigh``; Opus 4.6 must clamp to ``max``
instead of raising, while Opus 4.7 (ceiling ``xhigh``) keeps ``xhigh``.
"""
config = AmazonAnthropicClaudeMessagesConfig()
optional_params = {"max_tokens": 1024, "reasoning_effort": effort}
result = config.transform_anthropic_messages_request(
model=model,
messages=[{"role": "user", "content": "Hello"}],
anthropic_messages_optional_request_params=optional_params,
litellm_params={},
headers={},
)
assert result["output_config"]["effort"] == expected_effort
assert result["thinking"]["type"] == "adaptive"
def test_bedrock_invoke_messages_rejects_xhigh_without_ceiling():
"""Sonnet 4.6 on Bedrock has no effort ceiling, so xhigh is still rejected."""
config = AmazonAnthropicClaudeMessagesConfig()
optional_params = {"max_tokens": 1024, "reasoning_effort": "xhigh"}
with pytest.raises(AnthropicError) as exc_info:
config.transform_anthropic_messages_request(
model="invoke/us.anthropic.claude-sonnet-4-6",
messages=[{"role": "user", "content": "Hello"}],
anthropic_messages_optional_request_params=optional_params,
litellm_params={},
headers={},
)
assert exc_info.value.status_code == 400
assert "not supported by this model" in str(exc_info.value)
@pytest.mark.parametrize(
"model",
[

View file

@ -63,7 +63,13 @@ class TestGetModelInfoReasoningEffortFields:
class TestModelRegistryReasoningEffortFields:
"""Verify specific models have the expected reasoning effort capability
values in the JSON registry file."""
values in the JSON registry file.
Claude models intentionally OMIT ``supports_minimal_reasoning_effort``:
``minimal`` is not a real Anthropic effort level (the API accepts only
low/medium/high/xhigh/max), so LiteLLM degrades ``minimal`` to ``low``
regardless of the flag. These tests guard against the flag being
re-added to the Claude fleet."""
@pytest.fixture(autouse=True)
def _load_registry(self):
@ -77,41 +83,41 @@ class TestModelRegistryReasoningEffortFields:
entry = self.registry["claude-opus-4-6"]
assert entry.get("supports_max_reasoning_effort") is True
def test_opus_4_7_supports_minimal(self):
def test_opus_4_7_omits_minimal(self):
entry = self.registry["claude-opus-4-7"]
assert entry.get("supports_minimal_reasoning_effort") is True
assert "supports_minimal_reasoning_effort" not in entry
def test_opus_4_6_supports_minimal(self):
def test_opus_4_6_omits_minimal(self):
entry = self.registry["claude-opus-4-6"]
assert entry.get("supports_minimal_reasoning_effort") is True
assert "supports_minimal_reasoning_effort" not in entry
def test_sonnet_4_6_supports_minimal(self):
def test_sonnet_4_6_omits_minimal(self):
entry = self.registry["anthropic.claude-sonnet-4-6"]
assert entry.get("supports_minimal_reasoning_effort") is True
assert "supports_minimal_reasoning_effort" not in entry
def test_bedrock_opus_4_7_supports_max(self):
entry = self.registry["anthropic.claude-opus-4-7"]
assert entry.get("supports_max_reasoning_effort") is True
assert entry.get("supports_minimal_reasoning_effort") is True
assert "supports_minimal_reasoning_effort" not in entry
def test_vertex_opus_4_7_supports_max(self):
entry = self.registry["vertex_ai/claude-opus-4-7"]
assert entry.get("supports_max_reasoning_effort") is True
assert entry.get("supports_minimal_reasoning_effort") is True
assert "supports_minimal_reasoning_effort" not in entry
def test_vertex_opus_4_6_supports_max(self):
entry = self.registry["vertex_ai/claude-opus-4-6"]
assert entry.get("supports_max_reasoning_effort") is True
assert entry.get("supports_minimal_reasoning_effort") is True
assert "supports_minimal_reasoning_effort" not in entry
def test_azure_ai_opus_4_6_supports_minimal(self):
def test_azure_ai_opus_4_6_omits_minimal(self):
entry = self.registry["azure_ai/claude-opus-4-6"]
assert entry.get("supports_minimal_reasoning_effort") is True
assert "supports_minimal_reasoning_effort" not in entry
def test_azure_ai_opus_4_7_supports_max(self):
entry = self.registry["azure_ai/claude-opus-4-7"]
assert entry.get("supports_max_reasoning_effort") is True
assert entry.get("supports_minimal_reasoning_effort") is True
assert "supports_minimal_reasoning_effort" not in entry
# ---------------------------------------------------------------------------

View file

@ -430,6 +430,61 @@ def test_output_config_forwarded_for_bedrock_chat_invoke_request():
assert result["max_tokens"] == 100
def test_output_config_format_converted_for_bedrock_chat_invoke_request():
"""Bedrock Invoke chat path consumes ``output_config.format`` before forwarding."""
config = AmazonAnthropicClaudeConfig()
schema = {
"type": "object",
"properties": {"answer": {"type": "string"}},
}
result = config.transform_request(
model="anthropic.claude-opus-4-7",
messages=[{"role": "user", "content": "test"}],
optional_params={
"max_tokens": 100,
"output_config": {
"effort": "xhigh",
"format": {"type": "json_schema", "schema": schema},
},
},
litellm_params={},
headers={},
)
assert result.get("output_config") == {"effort": "xhigh"}
last_content = result["messages"][0]["content"]
assert json.loads(last_content[-1]["text"]) == schema
@pytest.mark.parametrize(
"model,expected_effort",
[
("anthropic.claude-opus-4-5-20251101-v1:0", "high"),
("anthropic.claude-opus-4-6-v1", "max"),
("anthropic.claude-opus-4-7", "xhigh"),
],
)
def test_output_config_effort_normalized_for_bedrock_chat_invoke_request(
model, expected_effort
):
"""Bedrock Invoke chat path accepts ``xhigh`` and forwards the provider-safe effort."""
config = AmazonAnthropicClaudeConfig()
result = config.transform_request(
model=model,
messages=[{"role": "user", "content": "test"}],
optional_params={
"max_tokens": 100,
"output_config": {"effort": "xhigh"},
},
litellm_params={},
headers={},
)
assert result.get("output_config") == {"effort": expected_effort}
def test_bedrock_chat_invoke_checks_output_config_support_with_bedrock_provider():
config = AmazonAnthropicClaudeConfig()
messages = [{"role": "user", "content": "test"}]

View file

@ -318,6 +318,7 @@ def test_reasoning_effort_none_omits_thinking_for_anthropic_converse(model):
("bedrock/converse/us.anthropic.claude-opus-4-7", "high", "high"),
("bedrock/converse/us.anthropic.claude-opus-4-7", "xhigh", "xhigh"),
("bedrock/converse/us.anthropic.claude-opus-4-7", "max", "max"),
("bedrock/converse/us.anthropic.claude-opus-4-6-v1", "xhigh", "max"),
("bedrock/converse/us.anthropic.claude-opus-4-6-v1", "max", "max"),
("bedrock/converse/us.anthropic.claude-sonnet-4-6", "high", "high"),
("bedrock/converse/us.anthropic.claude-sonnet-4-6", "minimal", "low"),
@ -369,6 +370,132 @@ def test_output_config_effort_forwarded_into_additional_request_fields(model):
assert additional.get("output_config") == {"effort": "high"}
def test_output_config_format_translated_to_native_output_config_converse():
"""``output_config.format`` becomes Bedrock ``outputConfig`` and is not forwarded raw."""
config = AmazonConverseConfig()
schema = {
"type": "object",
"properties": {"answer": {"type": "string"}},
}
result = config._transform_request(
model="bedrock/converse/us.anthropic.claude-opus-4-7",
messages=[{"role": "user", "content": "hi"}],
optional_params={
"maxTokens": 256,
"thinking": {"type": "adaptive"},
"output_config": {
"effort": "xhigh",
"format": {"type": "json_schema", "schema": schema},
},
},
litellm_params={},
headers={},
)
additional = result.get("additionalModelRequestFields", {})
assert additional.get("output_config") == {"effort": "xhigh"}
assert "format" not in additional["output_config"]
assert result["outputConfig"]["textFormat"]["type"] == "json_schema"
parsed_schema = json.loads(
result["outputConfig"]["textFormat"]["structure"]["jsonSchema"]["schema"]
)
assert parsed_schema == {**schema, "additionalProperties": False}
def test_output_config_format_dropped_on_unsupported_converse_model_warns(caplog):
"""When Converse model lacks native structured-output support, the silently
dropped ``output_config.format`` must surface as a warning so callers can
diagnose plain-text responses."""
from unittest.mock import patch
config = AmazonConverseConfig()
schema = {
"type": "object",
"properties": {"answer": {"type": "string"}},
}
with patch.object(
AmazonConverseConfig,
"_supports_native_structured_outputs",
return_value=False,
):
with caplog.at_level("WARNING"):
result = config._transform_request(
model="bedrock/converse/us.anthropic.claude-3-haiku-20240307-v1:0",
messages=[{"role": "user", "content": "hi"}],
optional_params={
"maxTokens": 256,
"output_config": {
"format": {"type": "json_schema", "schema": schema},
},
},
litellm_params={},
headers={},
)
assert "outputConfig" not in result
assert any(
"dropping `output_config.format`" in record.getMessage()
for record in caplog.records
)
def test_output_config_normalized_marker_does_not_leak_into_optional_params():
"""The internal ``_output_config_normalized`` marker set by
``_handle_reasoning_effort_parameter`` must be consumed during request
preparation so it does not linger on the caller's ``optional_params``."""
config = AmazonConverseConfig()
optional_params = config.map_openai_params(
non_default_params={"reasoning_effort": "xhigh"},
optional_params={},
model="bedrock/converse/us.anthropic.claude-opus-4-6-v1",
drop_params=False,
)
assert optional_params.get("_output_config_normalized") is True
config._transform_request(
model="bedrock/converse/us.anthropic.claude-opus-4-6-v1",
messages=[{"role": "user", "content": "hi"}],
optional_params=optional_params,
litellm_params={},
headers={},
)
assert "_output_config_normalized" not in optional_params
@pytest.mark.parametrize(
"model,expected_effort",
[
("bedrock/converse/us.anthropic.claude-opus-4-5-20251101-v1:0", "high"),
("bedrock/converse/us.anthropic.claude-opus-4-6-v1", "max"),
("bedrock/converse/us.anthropic.claude-opus-4-7", "xhigh"),
],
)
def test_output_config_effort_normalized_for_bedrock_converse_opus(
model, expected_effort
):
"""Bedrock Converse accepts ``xhigh`` and forwards the provider-safe effort."""
config = AmazonConverseConfig()
result = config._transform_request(
model=model,
messages=[{"role": "user", "content": "hi"}],
optional_params={
"maxTokens": 256,
"thinking": {"type": "adaptive"},
"output_config": {"effort": "xhigh"},
},
litellm_params={},
headers={},
)
additional = result.get("additionalModelRequestFields", {})
assert additional.get("output_config") == {"effort": expected_effort}
@pytest.mark.parametrize(
"effort",
["disabled", "invalid", ""],

View file

@ -767,6 +767,163 @@ def test_bedrock_messages_forwards_output_config_with_output_format():
assert "output_format" not in result
def test_bedrock_messages_converts_output_config_format_to_inline_schema():
"""``output_config.format`` is consumed so Bedrock does not see an unknown nested key."""
from unittest.mock import patch
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
messages = [{"role": "user", "content": [{"type": "text", "text": "Hello"}]}]
schema = {
"type": "object",
"properties": {"answer": {"type": "string"}},
}
optional_params = {
"max_tokens": 4096,
"output_config": {
"effort": "xhigh",
"format": {"type": "json_schema", "schema": schema},
},
}
with patch(
"litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation._supports_factory",
return_value=True,
):
result = cfg.transform_anthropic_messages_request(
model="anthropic.claude-opus-4-7",
messages=messages,
anthropic_messages_optional_request_params=optional_params,
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert result.get("output_config") == {"effort": "xhigh"}
assert "output_format" not in result
last_content = result["messages"][0]["content"]
assert json.loads(last_content[-1]["text"]) == schema
@pytest.mark.parametrize(
"model,expected_effort",
[
("anthropic.claude-opus-4-5-20251101-v1:0", "high"),
("anthropic.claude-opus-4-6-v1", "max"),
("anthropic.claude-opus-4-7", "xhigh"),
],
)
def test_bedrock_messages_normalizes_output_config_effort_for_opus(
model, expected_effort
):
"""Bedrock /v1/messages accepts ``xhigh`` and forwards the provider-safe effort."""
from unittest.mock import patch
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
with patch(
"litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation._supports_factory",
return_value=True,
):
result = cfg.transform_anthropic_messages_request(
model=model,
messages=[{"role": "user", "content": [{"type": "text", "text": "Hello"}]}],
anthropic_messages_optional_request_params={
"max_tokens": 4096,
"output_config": {"effort": "xhigh"},
},
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert result.get("output_config") == {"effort": expected_effort}
def test_bedrock_messages_does_not_mutate_callers_messages_when_embedding_schema():
"""Inline-schema embedding must not mutate the caller's ``messages`` list,
message dicts, or content list."""
from unittest.mock import patch
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
caller_content = [{"type": "text", "text": "Hello"}]
caller_message = {"role": "user", "content": caller_content}
caller_messages = [caller_message]
schema = {"type": "object", "properties": {"answer": {"type": "string"}}}
optional_params = {
"max_tokens": 4096,
"output_config": {
"effort": "xhigh",
"format": {"type": "json_schema", "schema": schema},
},
}
with patch(
"litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation._supports_factory",
return_value=True,
):
result = cfg.transform_anthropic_messages_request(
model="anthropic.claude-opus-4-7",
messages=caller_messages,
anthropic_messages_optional_request_params=optional_params,
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert caller_messages == [
{"role": "user", "content": [{"type": "text", "text": "Hello"}]}
]
assert caller_message == {
"role": "user",
"content": [{"type": "text", "text": "Hello"}],
}
assert caller_content == [{"type": "text", "text": "Hello"}]
last_content = result["messages"][-1]["content"]
assert json.loads(last_content[-1]["text"]) == schema
def test_bedrock_messages_does_not_mutate_callers_output_config():
"""`pop_bedrock_invoke_output_config_format` / effort normalization must not
leak into the caller's ``optional_params`` dict."""
from unittest.mock import patch
from litellm.types.router import GenericLiteLLMParams
cfg = AmazonAnthropicClaudeMessagesConfig()
schema = {
"type": "object",
"properties": {"answer": {"type": "string"}},
}
caller_output_config = {
"effort": "xhigh",
"format": {"type": "json_schema", "schema": schema},
}
optional_params = {
"max_tokens": 4096,
"output_config": caller_output_config,
}
with patch(
"litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation._supports_factory",
return_value=True,
):
cfg.transform_anthropic_messages_request(
model="anthropic.claude-opus-4-5-20251101-v1:0",
messages=[{"role": "user", "content": [{"type": "text", "text": "Hello"}]}],
anthropic_messages_optional_request_params=optional_params,
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert caller_output_config == {
"effort": "xhigh",
"format": {"type": "json_schema", "schema": schema},
}
def test_bedrock_messages_strips_output_config_with_output_format():
"""
When both output_config and output_format are present, output_format
@ -1071,9 +1228,7 @@ def test_bedrock_messages_preserves_compact_context_management_and_adds_beta():
messages = [{"role": "user", "content": [{"type": "text", "text": "Hi"}]}]
optional_params = {
"max_tokens": 4096,
"context_management": {
"edits": [{"type": "compact_20260112"}]
},
"context_management": {"edits": [{"type": "compact_20260112"}]},
}
result = cfg.transform_anthropic_messages_request(
@ -1084,9 +1239,7 @@ def test_bedrock_messages_preserves_compact_context_management_and_adds_beta():
headers={},
)
assert result.get("context_management") == {
"edits": [{"type": "compact_20260112"}]
}
assert result.get("context_management") == {"edits": [{"type": "compact_20260112"}]}
assert "compact-2026-01-12" in result.get("anthropic_beta", [])
assert result["max_tokens"] == 4096
@ -1118,9 +1271,7 @@ def test_bedrock_messages_filters_unsupported_context_management_edits():
headers={},
)
assert result.get("context_management") == {
"edits": [{"type": "compact_20260112"}]
}
assert result.get("context_management") == {"edits": [{"type": "compact_20260112"}]}
assert "compact-2026-01-12" in result.get("anthropic_beta", [])

View file

@ -1,9 +1,7 @@
import json
import os
import sys
import pytest
from fastapi.testclient import TestClient
sys.path.insert(
0, os.path.abspath("../../../..")
@ -12,7 +10,6 @@ sys.path.insert(
from litellm.llms.bedrock.common_utils import BedrockModelInfo
# --------------------------------------------------------------------------- #
# get_bedrock_response_stream_shape lazy-load tests #
# --------------------------------------------------------------------------- #
@ -24,8 +21,10 @@ def _reset_bedrock_response_stream_shape_cache():
import litellm.llms.bedrock.common_utils as mod
mod.get_bedrock_response_stream_shape.cache_clear()
mod._get_local_model_cost_map.cache_clear()
yield
mod.get_bedrock_response_stream_shape.cache_clear()
mod._get_local_model_cost_map.cache_clear()
def test_bedrock_response_stream_shape_lazy_loads_once():
@ -222,3 +221,45 @@ def test_context_window_suffix_stripped_for_cost_lookup():
get_bedrock_base_model("anthropic.claude-3-5-sonnet-20241022-v2:0:51k")
== "anthropic.claude-3-5-sonnet-20241022-v2:0"
)
def test_output_config_effort_normalization_uses_model_info_ceiling(monkeypatch):
import litellm.llms.bedrock.common_utils as mod
calls = []
def fake_get_model_info(model, custom_llm_provider=None):
calls.append((model, custom_llm_provider))
return {"bedrock_output_config_effort_ceiling": "max"}
monkeypatch.setattr(mod, "_get_model_info", fake_get_model_info)
output_config = {"effort": "xhigh"}
mod.normalize_bedrock_opus_output_config_effort(
model="custom-bedrock-alias-without-opus-pattern",
output_config=output_config,
)
assert output_config == {"effort": "max"}
assert calls == [("custom-bedrock-alias-without-opus-pattern", "bedrock")]
@pytest.mark.parametrize(
"model,expected_ceiling",
[
("anthropic.claude-opus-4-5-20251101-v1:0", "high"),
("anthropic.claude-opus-4-6-v1", "max"),
("anthropic.claude-opus-4-7", "xhigh"),
("us.anthropic.claude-opus-4-5-20251101-v1:0", "high"),
("us.anthropic.claude-opus-4-6-v1", "max"),
("us.anthropic.claude-opus-4-7", "xhigh"),
],
)
def test_bundled_bedrock_opus_model_info_declares_output_config_effort_ceiling(
model, expected_ceiling
):
from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap
model_info = GetModelCostMap.load_local_model_cost_map()[model]
assert model_info["bedrock_output_config_effort_ceiling"] == expected_ceiling

View file

@ -19,6 +19,7 @@ import websockets.exceptions # registers websockets.exceptions on the websocket
sys.path.insert(0, os.path.abspath("../../../../.."))
import litellm
from litellm.llms.vertex_ai.realtime.transformation import VertexAIRealtimeConfig
# ---------------------------------------------------------------------------
@ -82,6 +83,85 @@ def test_session_configuration_request_model_format():
)
def test_vertex_requires_session_configuration_feature_flag(monkeypatch):
cfg = VertexAIRealtimeConfig(
access_token="tok", project="my-proj", location="us-central1"
)
# Default remains backwards-compatible (auto setup on connect)
monkeypatch.setattr(litellm, "gemini_live_defer_setup", False, raising=False)
assert cfg.requires_session_configuration() is True
# Opt-in deferred setup for tool-injection flow
monkeypatch.setattr(litellm, "gemini_live_defer_setup", True, raising=False)
assert cfg.requires_session_configuration() is False
def test_vertex_session_update_defaults_to_audio_modality():
cfg = VertexAIRealtimeConfig(
access_token="tok", project="my-proj", location="us-central1"
)
session_update = {
"type": "session.update",
"session": {
"instructions": "You are a helpful assistant.",
# No modalities provided on purpose
},
}
messages = cfg.transform_realtime_request(
json.dumps(session_update),
"gemini-live-2.5-flash-native-audio",
session_configuration_request=None,
)
assert len(messages) == 1
setup_payload = json.loads(messages[0])["setup"]
assert setup_payload["generationConfig"]["responseModalities"] == ["AUDIO"]
def test_vertex_session_update_normalizes_ga_remapped_fields():
"""GA-format clients send ``output_modalities`` and nested
``audio.input.transcription`` / ``audio.input.turn_detection``. These must
be normalised back to the flat beta keys before ``map_openai_params``
runs so client preferences aren't silently dropped.
"""
cfg = VertexAIRealtimeConfig(
access_token="tok", project="my-proj", location="us-central1"
)
session_update = {
"type": "session.update",
"session": {
"instructions": "Be concise.",
"output_modalities": ["text"],
"audio": {
"input": {
"transcription": {},
"turn_detection": {"silence_duration_ms": 1500},
},
},
},
}
messages = cfg.transform_realtime_request(
json.dumps(session_update),
"gemini-live-2.5-flash-native-audio",
session_configuration_request=None,
)
assert len(messages) == 1
setup_payload = json.loads(messages[0])["setup"]
assert setup_payload["generationConfig"]["responseModalities"] == ["TEXT"]
assert setup_payload["inputAudioTranscription"] == {}
assert (
setup_payload["realtimeInputConfig"]["automaticActivityDetection"][
"silenceDurationMs"
]
== 1500
)
# ---------------------------------------------------------------------------
# Round-trip test: text-in / text-out via RealTimeStreaming
# ---------------------------------------------------------------------------
@ -208,3 +288,61 @@ async def test_vertex_realtime_text_in_text_out():
# response.done should have been forwarded
done_msgs = [m for m in sent_to_client if '"response.done"' in m]
assert done_msgs, "Expected response.done to be sent to client"
def test_vertex_warns_when_dropping_guardrail_turn_detection_update(caplog):
"""A subsequent session.update carrying the guardrail's
``create_response: False`` cannot be forwarded as a follow-up setup on
Vertex AI (1007). Surface a warning so operators know the auto-response
suppression is being silently dropped."""
import logging
cfg = VertexAIRealtimeConfig(
access_token="tok", project="my-proj", location="us-central1"
)
session_update = {
"type": "session.update",
"session": {"turn_detection": {"create_response": False}},
}
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
result = cfg.transform_realtime_request(
json.dumps(session_update),
"gemini-live-2.5-flash-native-audio",
session_configuration_request=json.dumps({"setup": {"model": "x"}}),
)
assert result == []
assert any(
"Vertex AI Realtime" in record.message
and "create_response=False" in record.message
for record in caplog.records
)
def test_vertex_does_not_warn_when_dropping_non_guardrail_session_update(caplog):
"""A subsequent session.update without ``create_response: False`` is a
routine drop and should stay at debug level (no warning)."""
import logging
cfg = VertexAIRealtimeConfig(
access_token="tok", project="my-proj", location="us-central1"
)
session_update = {
"type": "session.update",
"session": {"instructions": "Be concise."},
}
with caplog.at_level(logging.WARNING, logger="LiteLLM"):
cfg.transform_realtime_request(
json.dumps(session_update),
"gemini-live-2.5-flash-native-audio",
session_configuration_request=json.dumps({"setup": {"model": "x"}}),
)
assert not any(
"Vertex AI Realtime" in record.message and "session.update" in record.message
for record in caplog.records
)

View file

@ -456,6 +456,127 @@ class TestVertexAIVideoConfig:
raw_response=mock_response, logging_obj=self.mock_logging_obj
)
def test_get_video_edit_prefetch_params(self):
"""Test that prefetch params returns the fetchPredictOperation URL and body."""
operation_name = "projects/test-project/locations/us-central1/publishers/google/models/veo-3.1-generate-001/operations/op-123"
api_base = "https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models"
fetch_url, fetch_body = self.config.get_video_edit_prefetch_params(
video_id=operation_name,
api_base=api_base,
litellm_params=GenericLiteLLMParams(),
headers={},
)
assert "fetchPredictOperation" in fetch_url
assert "veo-3.1-generate-001" in fetch_url
assert fetch_body == {"operationName": operation_name}
def test_transform_video_edit_request_with_bytes(self):
"""Test video edit request builds predictLongRunning body from pre-fetched bytes."""
operation_name = "projects/test-project/locations/us-central1/publishers/google/models/veo-3.1-generate-001/operations/op-123"
api_base = "https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models"
fake_bytes = base64.b64encode(b"fake_video").decode()
prefetched = {
"done": True,
"response": {
"videos": [{"bytesBase64Encoded": fake_bytes, "mimeType": "video/mp4"}]
},
}
url, data = self.config.transform_video_edit_request(
prompt="Make it brighter",
video_id=operation_name,
api_base=api_base,
litellm_params=GenericLiteLLMParams(),
headers={"Authorization": "Bearer token"},
prefetched_source_data=prefetched,
)
assert url.endswith(":predictLongRunning")
assert "veo-3.1-generate-001" in url
instance = data["instances"][0]
assert instance["prompt"] == "Make it brighter"
assert instance["video"]["bytesBase64Encoded"] == fake_bytes
assert instance["video"]["mimeType"] == "video/mp4"
def test_transform_video_edit_request_with_gcs_uri(self):
"""Test that gcsUri is used when present in source video."""
operation_name = "projects/test-project/locations/us-central1/publishers/google/models/veo-3.1-generate-001/operations/op-456"
api_base = "https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models"
prefetched = {
"done": True,
"response": {
"videos": [{"gcsUri": "gs://bucket/video.mp4", "mimeType": "video/mp4"}]
},
}
_, data = self.config.transform_video_edit_request(
prompt="Make it darker",
video_id=operation_name,
api_base=api_base,
litellm_params=GenericLiteLLMParams(),
headers={},
prefetched_source_data=prefetched,
)
assert data["instances"][0]["video"] == {"gcsUri": "gs://bucket/video.mp4"}
def test_transform_video_edit_request_source_not_done_raises(self):
"""Test that editing an in-progress video raises a clear error."""
operation_name = "projects/test-project/locations/us-central1/publishers/google/models/veo-3.1-generate-001/operations/op-789"
api_base = "https://us-central1-aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models"
with pytest.raises(ValueError, match="not complete yet"):
self.config.transform_video_edit_request(
prompt="Make it brighter",
video_id=operation_name,
api_base=api_base,
litellm_params=GenericLiteLLMParams(),
headers={},
prefetched_source_data={"done": False},
)
def test_transform_video_edit_response(self):
"""Test that edit response returns a processing VideoObject with encoded ID."""
operation_name = "projects/test-project/locations/us-central1/publishers/google/models/veo-3.1-generate-001/operations/new-op-123"
mock_response = Mock(spec=httpx.Response)
mock_response.json.return_value = {"name": operation_name}
video_obj = self.config.transform_video_edit_response(
raw_response=mock_response,
logging_obj=self.mock_logging_obj,
custom_llm_provider="vertex_ai",
)
assert isinstance(video_obj, VideoObject)
assert video_obj.status == "processing"
assert video_obj.id
assert video_obj.model == "veo-3.1-generate-001"
def test_transform_video_edit_response_includes_usage_for_cost(self):
"""Edit responses include duration/resolution usage for spend accounting."""
operation_name = "projects/test-project/locations/us-central1/publishers/google/models/veo-3.1-generate-001/operations/new-op-123"
mock_response = Mock(spec=httpx.Response)
mock_response.json.return_value = {"name": operation_name}
request_data = {
"instances": [{"prompt": "Make it brighter", "video": {}}],
"parameters": {"durationSeconds": 8, "resolution": "1080p"},
}
video_obj = self.config.transform_video_edit_response(
raw_response=mock_response,
logging_obj=self.mock_logging_obj,
custom_llm_provider="vertex_ai",
request_data=request_data,
)
assert video_obj.usage is not None
assert video_obj.usage["duration_seconds"] == 8.0
assert video_obj.usage["video_resolution"] == "1080p"
def test_transform_video_remix_request_not_supported(self):
"""Test that video remix raises NotImplementedError."""
with pytest.raises(NotImplementedError, match="Video remix is not supported"):

View file

@ -2444,13 +2444,14 @@ async def test_get_team_object_permission_with_core_auth_auto_loading():
@pytest.mark.asyncio
async def test_get_allowed_mcp_servers_for_team_uses_helper():
"""
Test that _get_allowed_mcp_servers_for_team properly uses _get_team_object_permission
helper which handles both loaded and unloaded object_permission cases.
Test that _get_allowed_mcp_servers_for_team resolves both legacy
object_permission fields (mcp_servers, mcp_access_groups) and the unified
team.access_group_ids → access_mcp_server_ids path.
"""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._types import LiteLLM_ObjectPermissionTable
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable
from litellm.types.mcp import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
@ -2464,53 +2465,51 @@ async def test_get_allowed_mcp_servers_for_team_uses_helper():
transport=MCPTransport.http,
)
try:
# Create mock object permission with servers and access groups
mock_object_permission = LiteLLM_ObjectPermissionTable(
object_permission_id="perm-789",
mcp_servers=["direct-server1", "direct-server2"],
mcp_access_groups=["dev-group"],
vector_stores=[],
)
mock_team = LiteLLM_TeamTable(
team_id="team-789",
access_group_ids=[],
object_permission_id="perm-789",
)
mock_team.object_permission = mock_object_permission
# Create mock user auth
mock_user_auth = UserAPIKeyAuth(
api_key="test-key",
user_id="test-user",
team_id="team-789",
)
# Mock the helper methods
with patch.object(
MCPRequestHandler, "_get_team_object_permission"
) as mock_get_team_perm:
with patch.object(
MCPRequestHandler, "_get_mcp_servers_from_access_groups"
) as mock_get_access_group_servers:
# Configure mocks
mock_get_team_perm.return_value = mock_object_permission
mock_get_access_group_servers.return_value = [
"group-server1",
"group-server2",
]
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
new_callable=AsyncMock,
return_value=mock_team,
),
patch.object(
MCPRequestHandler,
"_get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=["group-server1", "group-server2"],
) as mock_get_access_group_servers,
):
result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(
mock_user_auth
)
# Call the method
result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(
mock_user_auth
)
assert set(result) == {
"direct-server1",
"direct-server2",
"group-server1",
"group-server2",
}
# Assert the result contains both direct and access group servers
assert set(result) == {
"direct-server1",
"direct-server2",
"group-server1",
"group-server2",
}
# Verify _get_team_object_permission was called (the helper we fixed)
mock_get_team_perm.assert_called_once_with(mock_user_auth)
# Verify access groups were resolved
mock_get_access_group_servers.assert_called_once_with(["dev-group"])
mock_get_access_group_servers.assert_called_once_with(["dev-group"])
finally:
for sid in ("direct-server1", "direct-server2"):
global_mcp_server_manager.registry.pop(sid, None)
@ -2520,32 +2519,36 @@ async def test_get_allowed_mcp_servers_for_team_uses_helper():
async def test_get_allowed_mcp_servers_for_team_with_no_object_permission():
"""
Test that _get_allowed_mcp_servers_for_team returns empty list when
team has no object_permission.
the team has no object_permission and no access_group_ids.
"""
# Create mock user auth
from litellm.proxy._types import LiteLLM_TeamTable
mock_team = LiteLLM_TeamTable(
team_id="team-no-perm",
access_group_ids=[],
object_permission_id=None,
)
mock_user_auth = UserAPIKeyAuth(
api_key="test-key",
user_id="test-user",
team_id="team-no-perm",
)
# Mock the helper to return None (no object permission)
with patch.object(
MCPRequestHandler, "_get_team_object_permission"
) as mock_get_team_perm:
mock_get_team_perm.return_value = None
# Call the method
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
new_callable=AsyncMock,
return_value=mock_team,
),
):
result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(
mock_user_auth
)
# Assert empty list is returned
assert result == []
# Verify the helper was called
mock_get_team_perm.assert_called_once_with(mock_user_auth)
@pytest.mark.asyncio
async def test_get_allowed_mcp_servers_for_team_without_user_auth_returns_empty():
@ -3456,3 +3459,185 @@ async def test_get_allowed_mcp_servers_no_union_when_no_authorized_extras():
# key ∩ team = {} (no overlap), extras = [] → final = []
result = await MCPRequestHandler.get_allowed_mcp_servers(auth)
assert result == []
# ---------------------------------------------------------------------------
# Issue #27657: team unified access_group_ids resolve to MCP servers
# ---------------------------------------------------------------------------
@pytest.mark.asyncio
async def test_team_access_group_ids_resolve_to_mcp_servers():
"""A virtual key with empty access_group_ids inherits MCP servers from
its team's access_group_ids (mirror of the model-side resolution).
Reproduction of https://github.com/BerriAI/litellm/issues/27657:
the runtime used to ignore team.access_group_ids when computing the
MCP scope, so virtual keys saw empty server lists even when their
team had an MCP-granting access group attached.
"""
from litellm.proxy._types import LiteLLM_TeamTable
mock_team = LiteLLM_TeamTable(
team_id="team-a",
access_group_ids=["mcp-premium"],
object_permission_id=None,
)
auth = UserAPIKeyAuth(
token="test-token-hash",
api_key="sk-test",
team_id="team-a",
access_group_ids=[],
)
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
new_callable=AsyncMock,
return_value=mock_team,
),
patch(
"litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups",
new_callable=AsyncMock,
return_value=["srv-stripe"],
) as mock_resolver,
):
result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(auth)
assert result == ["srv-stripe"]
mock_resolver.assert_called_once()
assert mock_resolver.call_args.kwargs["access_group_ids"] == ["mcp-premium"]
@pytest.mark.asyncio
async def test_team_access_group_ids_union_with_object_permission():
"""When both legacy object_permission and unified team.access_group_ids
grant MCP servers, the final list is their union."""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import (
global_mcp_server_manager,
)
from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable
from litellm.types.mcp import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
for sid in ("srv-direct",):
global_mcp_server_manager.registry[sid] = MCPServer(
server_id=sid,
name=sid,
server_name=sid,
url=f"https://{sid}.example.com",
transport=MCPTransport.http,
)
try:
mock_object_permission = LiteLLM_ObjectPermissionTable(
object_permission_id="perm-1",
mcp_servers=["srv-direct"],
mcp_access_groups=[],
vector_stores=[],
)
mock_team = LiteLLM_TeamTable(
team_id="team-a",
access_group_ids=["mcp-premium"],
object_permission_id="perm-1",
)
mock_team.object_permission = mock_object_permission
auth = UserAPIKeyAuth(
token="test-token-hash",
api_key="sk-test",
team_id="team-a",
)
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
new_callable=AsyncMock,
return_value=mock_team,
),
patch(
"litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups",
new_callable=AsyncMock,
return_value=["srv-stripe"],
),
):
result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(auth)
assert set(result) == {"srv-direct", "srv-stripe"}
finally:
global_mcp_server_manager.registry.pop("srv-direct", None)
@pytest.mark.asyncio
async def test_team_access_group_ids_empty_returns_no_extras():
"""Empty team.access_group_ids → resolver called with [], short-circuits
without DB access, no extras added."""
from litellm.proxy._types import LiteLLM_TeamTable
mock_team = LiteLLM_TeamTable(
team_id="team-a",
access_group_ids=[],
object_permission_id=None,
)
auth = UserAPIKeyAuth(
token="test-token-hash",
api_key="sk-test",
team_id="team-a",
)
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
new_callable=AsyncMock,
return_value=mock_team,
),
patch(
"litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups",
new_callable=AsyncMock,
return_value=[],
) as mock_resolver,
):
result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(auth)
assert result == []
mock_resolver.assert_called_once()
assert mock_resolver.call_args.kwargs["access_group_ids"] == []
@pytest.mark.asyncio
async def test_get_allowed_mcp_servers_includes_team_access_group_extras_end_to_end():
"""End-to-end: virtual key has nothing of its own, team has an MCP
access group → key sees the granted server through get_allowed_mcp_servers."""
auth = UserAPIKeyAuth(
token="test-token",
api_key="sk-test",
team_id="team-a",
access_group_ids=[],
)
with (
patch.object(
MCPRequestHandler,
"_get_allowed_mcp_servers_for_key",
new_callable=AsyncMock,
return_value=[],
),
patch.object(
MCPRequestHandler,
"_get_allowed_mcp_servers_for_team",
new_callable=AsyncMock,
return_value=["srv-stripe"],
),
patch.object(
MCPRequestHandler,
"_get_key_access_group_mcp_server_extras",
new_callable=AsyncMock,
return_value=[],
),
):
result = await MCPRequestHandler.get_allowed_mcp_servers(auth)
assert result == ["srv-stripe"]

View file

@ -462,6 +462,9 @@ async def test_e2e_jwt_team_mcp_key_intersection(monkeypatch):
monkeypatch.setattr(
"litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object
)
monkeypatch.setattr(
"litellm.proxy.auth.auth_checks.get_team_object", mock_get_team_object
)
jwt_handler = JWTHandler()
jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_ids_jwt_field="groups")
@ -495,28 +498,25 @@ async def test_e2e_jwt_team_mcp_key_intersection(monkeypatch):
object_permission=key_object_permission, # Key has its own permissions
)
# Mock the helper methods to return our test data
with patch.object(
MCPRequestHandler, "_get_team_object_permission"
) as mock_team_perm:
mock_team_perm.return_value = team_object_permission
with (
patch.object(
MCPRequestHandler,
"_get_key_object_permission",
return_value=key_object_permission,
),
patch.object(
MCPRequestHandler,
"_get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
),
):
allowed_servers = await MCPRequestHandler.get_allowed_mcp_servers(
user_api_key_auth
)
with patch.object(
MCPRequestHandler, "_get_key_object_permission"
) as mock_key_perm:
mock_key_perm.return_value = key_object_permission
with patch.object(
MCPRequestHandler, "_get_mcp_servers_from_access_groups"
) as mock_access_groups:
mock_access_groups.return_value = []
allowed_servers = await MCPRequestHandler.get_allowed_mcp_servers(
user_api_key_auth
)
# Should be intersection: only server-2 is in both
expected = ["server-2"]
assert sorted(allowed_servers) == sorted(
expected
), f"Expected intersection {expected}, got {allowed_servers}"
# Should be intersection: only server-2 is in both
expected = ["server-2"]
assert sorted(allowed_servers) == sorted(
expected
), f"Expected intersection {expected}, got {allowed_servers}"

View file

@ -41,37 +41,44 @@ async def test_simple_jwt_mcp_permissions_enforced():
object_permission_id="perm-123",
mcp_servers=team_mcp_servers,
)
team_obj = LiteLLM_TeamTable(
team_id="my-team",
access_group_ids=[],
object_permission_id="perm-123",
)
team_obj.object_permission = team_object_permission
# 3. Mock the team permission lookup
with patch.object(
MCPRequestHandler, "_get_team_object_permission", new_callable=AsyncMock
) as mock_team_perm:
mock_team_perm.return_value = team_object_permission
# 3. Mock the team object lookup (object_permission attached) and prisma_client
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
new_callable=AsyncMock,
return_value=team_obj,
) as mock_get_team,
patch.object(
MCPRequestHandler,
"_get_key_object_permission",
new_callable=AsyncMock,
return_value=None,
),
patch.object(
MCPRequestHandler,
"_get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
return_value=[],
),
):
# 4. Call get_allowed_mcp_servers - this is what MCP routes use
allowed = await MCPRequestHandler.get_allowed_mcp_servers(user_auth)
# Mock key permissions (empty - user has no key-level MCP permissions)
with patch.object(
MCPRequestHandler, "_get_key_object_permission", new_callable=AsyncMock
) as mock_key_perm:
mock_key_perm.return_value = None
# 5. Verify only team's MCP servers are returned
assert sorted(allowed) == sorted(
team_mcp_servers
), f"Expected {team_mcp_servers}, got {allowed}"
# Mock access groups (empty)
with patch.object(
MCPRequestHandler,
"_get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
) as mock_access_groups:
mock_access_groups.return_value = []
# 4. Call get_allowed_mcp_servers - this is what MCP routes use
allowed = await MCPRequestHandler.get_allowed_mcp_servers(user_auth)
# 5. Verify only team's MCP servers are returned
assert sorted(allowed) == sorted(
team_mcp_servers
), f"Expected {team_mcp_servers}, got {allowed}"
# Verify team permission was looked up
mock_team_perm.assert_called_once_with(user_auth)
# Verify team was looked up
mock_get_team.assert_called()
@pytest.mark.asyncio
@ -120,25 +127,33 @@ async def test_simple_jwt_team_id_required_for_mcp_permissions():
object_permission_id="perm-1",
mcp_servers=team_mcp_servers,
)
team_obj = LiteLLM_TeamTable(
team_id="team-abc",
access_group_ids=[],
object_permission_id="perm-1",
)
team_obj.object_permission = team_perm
with patch.object(
MCPRequestHandler, "_get_team_object_permission", new_callable=AsyncMock
) as mock_perm:
mock_perm.return_value = team_perm
with patch.object(
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
patch(
"litellm.proxy.auth.auth_checks.get_team_object",
new_callable=AsyncMock,
return_value=team_obj,
) as mock_get_team,
patch.object(
MCPRequestHandler,
"_get_mcp_servers_from_access_groups",
new_callable=AsyncMock,
) as mock_groups:
mock_groups.return_value = []
return_value=[],
),
):
result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(
user_with_team
)
result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(
user_with_team
)
assert sorted(result) == sorted(team_mcp_servers)
mock_perm.assert_called_once() # Permission WAS checked
assert sorted(result) == sorted(team_mcp_servers)
mock_get_team.assert_called() # Team WAS looked up
# Case 2: team_id is None -> team permissions NOT checked
user_without_team = UserAPIKeyAuth(

View file

@ -1625,6 +1625,50 @@ async def test_reject_clientside_metadata_tags_non_llm_route():
assert result is True
@pytest.mark.asyncio
async def test_reject_clientside_metadata_tags_allows_key_tags_without_client_tags():
"""Key metadata.tags are injected after the reject check; requests without
client metadata.tags must not be blocked when reject_clientside_metadata_tags is on."""
from fastapi import Request
from litellm.proxy.auth.auth_checks import common_checks
request_body = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "test"}],
}
general_settings = {"reject_clientside_metadata_tags": True}
mock_request = MagicMock(spec=Request)
valid_token = UserAPIKeyAuth(
token="test-token",
models=["gpt-3.5-turbo"],
metadata={"tags": ["engineering"]},
)
with patch(
"litellm.proxy.auth.auth_checks.get_tag_objects_batch",
new_callable=AsyncMock,
return_value={},
):
result = await common_checks(
request_body=request_body,
team_object=None,
user_object=None,
end_user_object=None,
global_proxy_spend=None,
general_settings=general_settings,
route="/chat/completions",
llm_router=None,
proxy_logging_obj=MagicMock(),
valid_token=valid_token,
request=mock_request,
)
assert result is True
assert request_body["metadata"]["tags"] == ["engineering"]
@pytest.mark.asyncio
async def test_virtual_key_soft_budget_check_with_user_obj():
"""Test _virtual_key_soft_budget_check includes user_email when user_obj is provided"""

View file

@ -1149,6 +1149,13 @@ async def test_apply_guardrail_not_found(mocker):
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", mock_registry
)
mock_proxy_logging = mocker.Mock()
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mocker.patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging)
mocker.patch("litellm.proxy.proxy_server.general_settings", {})
mocker.patch("litellm.proxy.proxy_server.proxy_config", mocker.Mock())
mocker.patch("litellm.proxy.proxy_server.version", "test")
# Create request
request = ApplyGuardrailRequest(
guardrail_name="non-existent-guardrail", text="Test input text"
@ -1159,7 +1166,11 @@ async def test_apply_guardrail_not_found(mocker):
# Call endpoint and expect ProxyException
with pytest.raises(ProxyException) as exc_info:
await apply_guardrail(request=request, user_api_key_dict=mock_user_auth)
await apply_guardrail(
fastapi_request=mocker.Mock(),
request=request,
user_api_key_dict=mock_user_auth,
)
# Verify error details
assert str(exc_info.value.code) == "404"
@ -1186,6 +1197,25 @@ async def test_apply_guardrail_execution_error(mocker):
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", mock_registry
)
mock_logging_obj = mocker.Mock()
mock_logging_obj.async_failure_handler = AsyncMock()
mock_logging_obj.model_call_details = {}
mock_processor = mocker.Mock()
mock_processor.common_processing_pre_call_logic = AsyncMock(
return_value=({"guardrail_name": "test-guardrail"}, mock_logging_obj)
)
mocker.patch(
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing",
return_value=mock_processor,
)
mock_proxy_logging = mocker.Mock()
mock_proxy_logging.post_call_failure_hook = AsyncMock()
mocker.patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging)
mocker.patch("litellm.proxy.proxy_server.general_settings", {})
mocker.patch("litellm.proxy.proxy_server.proxy_config", mocker.Mock())
mocker.patch("litellm.proxy.proxy_server.version", "test")
mocker.patch("litellm.litellm_core_utils.thread_pool_executor.executor")
# Create request
request = ApplyGuardrailRequest(
guardrail_name="test-guardrail", text="Test input text with forbidden content"
@ -1196,12 +1226,70 @@ async def test_apply_guardrail_execution_error(mocker):
# Call endpoint and expect ProxyException
with pytest.raises(ProxyException) as exc_info:
await apply_guardrail(request=request, user_api_key_dict=mock_user_auth)
await apply_guardrail(
fastapi_request=mocker.Mock(),
request=request,
user_api_key_dict=mock_user_auth,
)
# Verify error is properly handled
assert "Bedrock guardrail failed" in str(exc_info.value.message)
@pytest.mark.asyncio
async def test_apply_guardrail_invokes_logging_pipeline(mocker):
mock_guardrail = mocker.Mock()
mock_guardrail.apply_guardrail = AsyncMock(return_value={"texts": ["masked"]})
mock_registry = mocker.Mock()
mock_registry.get_initialized_guardrail_callback.return_value = mock_guardrail
mocker.patch(
"litellm.proxy.guardrails.guardrail_endpoints.GUARDRAIL_REGISTRY", mock_registry
)
mock_logging_obj = mocker.Mock()
mock_logging_obj.async_success_handler = AsyncMock()
mock_logging_obj.model_call_details = {}
mock_processor = mocker.Mock()
mock_processor.common_processing_pre_call_logic = AsyncMock(
return_value=({"guardrail_name": "test-guardrail"}, mock_logging_obj)
)
mocker.patch(
"litellm.proxy.common_request_processing.ProxyBaseLLMRequestProcessing",
return_value=mock_processor,
)
mock_proxy_logging = mocker.Mock()
mock_proxy_logging.post_call_success_hook = AsyncMock()
mocker.patch("litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging)
mocker.patch("litellm.proxy.proxy_server.general_settings", {})
mocker.patch("litellm.proxy.proxy_server.proxy_config", mocker.Mock())
mocker.patch("litellm.proxy.proxy_server.version", "test")
mock_executor = mocker.Mock()
mocker.patch(
"litellm.litellm_core_utils.thread_pool_executor.executor", mock_executor
)
request = ApplyGuardrailRequest(
guardrail_name="test-guardrail", text="hello@example.com"
)
response = await apply_guardrail(
fastapi_request=mocker.Mock(),
request=request,
user_api_key_dict=UserAPIKeyAuth(),
)
assert response.response_text == "masked"
mock_processor.common_processing_pre_call_logic.assert_awaited_once()
mock_proxy_logging.post_call_success_hook.assert_awaited_once()
mock_logging_obj.async_success_handler.assert_awaited_once()
assert mock_logging_obj.call_type == "pass_through_endpoint"
mock_executor.submit.assert_called_once()
assert mock_logging_obj.async_success_handler.await_args.kwargs["result"] == {
"response": {"response_text": "masked"}
}
@pytest.mark.asyncio
async def test_get_guardrail_info_endpoint_config_guardrail(mocker):
"""

View file

@ -2832,6 +2832,7 @@ async def test_list_team_v2_security_check_non_admin_user_own_teams():
]
mock_db.litellm_teamtable.find_many = AsyncMock(return_value=mock_teams)
mock_db.litellm_teamtable.count = AsyncMock(return_value=2)
mock_db.litellm_verificationtoken.group_by = AsyncMock(return_value=[])
with patch(
"litellm.proxy.management_endpoints.team_endpoints.get_user_object",
@ -2888,6 +2889,7 @@ async def test_list_team_v2_security_check_admin_user():
]
mock_db.litellm_teamtable.find_many = AsyncMock(return_value=mock_teams)
mock_db.litellm_teamtable.count = AsyncMock(return_value=2)
mock_db.litellm_verificationtoken.group_by = AsyncMock(return_value=[])
# Should NOT raise an exception
result = await list_team_v2(
@ -3036,6 +3038,7 @@ async def test_list_team_v2_org_admin_sees_org_teams():
}
mock_db.litellm_teamtable.find_many = AsyncMock(return_value=[mock_team])
mock_db.litellm_teamtable.count = AsyncMock(return_value=1)
mock_db.litellm_verificationtoken.group_by = AsyncMock(return_value=[])
result = await list_team_v2(
http_request=mock_request,
@ -3211,6 +3214,7 @@ async def test_list_team_v2_org_admin_with_user_id_returns_user_teams():
}
mock_db.litellm_teamtable.find_many = AsyncMock(return_value=[mock_team])
mock_db.litellm_teamtable.count = AsyncMock(return_value=1)
mock_db.litellm_verificationtoken.group_by = AsyncMock(return_value=[])
result = await list_team_v2(
http_request=mock_request,
@ -3390,6 +3394,163 @@ async def test_list_team_v2_search_composes_with_user_id_filter():
assert where["team_id"] == {"in": ["team_a", "team_b"]}
@pytest.mark.asyncio
async def test_list_team_v2_populates_keys_count():
"""
Test that list_team_v2 returns a keys_count per team derived from a single
batched group_by against LiteLLM_VerificationToken.
"""
from unittest.mock import AsyncMock, Mock, patch
from fastapi import Request
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.management_endpoints.team_endpoints import list_team_v2
mock_request = Mock(spec=Request)
mock_user_api_key_dict_admin = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
user_id="admin_user_123",
)
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client:
mock_db = Mock()
mock_prisma_client.db = mock_db
team_a = Mock()
team_a.team_id = "team_a"
team_a.model_dump = lambda: {
"team_id": "team_a",
"team_alias": "Team A",
"members_with_roles": [{"user_id": "u1", "role": "user"}],
}
team_b = Mock()
team_b.team_id = "team_b"
team_b.model_dump = lambda: {
"team_id": "team_b",
"team_alias": "Team B",
"members_with_roles": [],
}
mock_db.litellm_teamtable.find_many = AsyncMock(return_value=[team_a, team_b])
mock_db.litellm_teamtable.count = AsyncMock(return_value=2)
mock_db.litellm_verificationtoken.group_by = AsyncMock(
return_value=[
{"team_id": "team_a", "_count": {"team_id": 3}},
# team_b intentionally absent → expect 0
]
)
result = await list_team_v2(
http_request=mock_request,
user_id=None,
user_api_key_dict=mock_user_api_key_dict_admin,
page=1,
page_size=10,
status=None,
)
assert result["total"] == 2
by_id = {t.team_id: t for t in result["teams"]}
assert by_id["team_a"].keys_count == 3
assert by_id["team_b"].keys_count == 0
# The aggregate is one batched query, filtered by the page's team IDs.
group_by_kwargs = mock_db.litellm_verificationtoken.group_by.call_args.kwargs
assert group_by_kwargs["by"] == ["team_id"]
assert group_by_kwargs["where"] == {"team_id": {"in": ["team_a", "team_b"]}}
assert group_by_kwargs["count"] == {"team_id": True}
@pytest.mark.asyncio
async def test_list_team_v2_keys_count_skipped_for_empty_page():
"""
When the page has no teams, the keys-count group_by must not be issued.
"""
from unittest.mock import AsyncMock, Mock, patch
from fastapi import Request
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.management_endpoints.team_endpoints import list_team_v2
mock_request = Mock(spec=Request)
mock_user_api_key_dict_admin = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
user_id="admin_user_123",
)
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client:
mock_db = Mock()
mock_prisma_client.db = mock_db
mock_db.litellm_teamtable.find_many = AsyncMock(return_value=[])
mock_db.litellm_teamtable.count = AsyncMock(return_value=0)
mock_db.litellm_verificationtoken.group_by = AsyncMock(return_value=[])
result = await list_team_v2(
http_request=mock_request,
user_id=None,
user_api_key_dict=mock_user_api_key_dict_admin,
page=1,
page_size=10,
status=None,
)
assert result["total"] == 0
assert result["teams"] == []
mock_db.litellm_verificationtoken.group_by.assert_not_called()
@pytest.mark.asyncio
async def test_list_team_v2_keys_count_skipped_for_deleted_status():
"""
The deleted-table branch returns LiteLLM_DeletedTeamTable items, which do
not carry keys_count — group_by must not be issued.
"""
from unittest.mock import AsyncMock, Mock, patch
from fastapi import Request
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.management_endpoints.team_endpoints import list_team_v2
mock_request = Mock(spec=Request)
mock_user_api_key_dict_admin = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
user_id="admin_user_123",
)
with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma_client:
mock_db = Mock()
mock_prisma_client.db = mock_db
mock_deleted = Mock()
mock_deleted.team_id = "team_d"
mock_deleted.model_dump = lambda: {
"team_id": "team_d",
"team_alias": "Deleted Team",
}
mock_db.litellm_deletedteamtable.find_many = AsyncMock(
return_value=[mock_deleted]
)
mock_db.litellm_deletedteamtable.count = AsyncMock(return_value=1)
mock_db.litellm_verificationtoken.group_by = AsyncMock(return_value=[])
result = await list_team_v2(
http_request=mock_request,
user_id=None,
user_api_key_dict=mock_user_api_key_dict_admin,
page=1,
page_size=10,
status="deleted",
)
assert result["total"] == 1
mock_db.litellm_verificationtoken.group_by.assert_not_called()
@pytest.mark.asyncio
async def test_team_member_delete_cleans_membership(mock_db_client, mock_admin_auth):
"""

View file

@ -515,6 +515,59 @@ async def test_add_litellm_data_to_request_body_snapshot_excludes_secret_fields(
)
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_body_snapshot_excludes_proxy_server_request():
"""Regression: the body snapshot used to include the proxy_server_request
key itself, producing the path
``proxy_server_request.body.proxy_server_request.body == body``. Custom
loggers and audit consumers must not see the self-referencing structure
(independent of redaction — fires on every successful call).
"""
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
request_mock = MagicMock(spec=Request)
request_mock.url.path = "/v1/chat/completions"
request_mock.url = MagicMock()
request_mock.url.__str__.return_value = "http://localhost/v1/chat/completions"
request_mock.method = "POST"
request_mock.query_params = {}
request_mock.headers = {"Content-Type": "application/json"}
request_mock.client = MagicMock()
request_mock.client.host = "127.0.0.1"
data = {
"model": "gpt-3.5-turbo",
"messages": [{"role": "user", "content": "hello"}],
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
user_id="test-user",
metadata={},
team_metadata={},
spend=0.0,
max_budget=100.0,
model_max_budget={},
team_spend=0.0,
team_max_budget=200.0,
)
updated = await add_litellm_data_to_request(
data=data,
request=request_mock,
user_api_key_dict=user_api_key_dict,
proxy_config=MagicMock(),
general_settings={},
version="test-version",
)
snapshot_body = updated["proxy_server_request"]["body"]
assert "proxy_server_request" not in snapshot_body, (
"proxy_server_request must be excluded from its own body snapshot "
"to prevent the body from self-referencing"
)
@pytest.mark.asyncio
async def test_add_litellm_data_to_request_strips_string_encoded_admin_injection():
"""Regression: metadata arriving as a JSON string (multipart/form-data or
@ -4182,6 +4235,209 @@ class TestApplyClientTagPolicyPreAuth:
assert exc_info.value.max_budget == 0.10
class TestApplyKeyTagsPreAuth:
def test_merges_key_tags_into_metadata(self):
data = {"model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={"tags": ["engineering", "production"]},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
request_data=data,
user_api_key_dict=user_api_key_dict,
)
assert data["metadata"]["tags"] == ["engineering", "production"]
def test_unions_key_tags_with_existing_request_tags(self):
data = {
"model": "gpt-3.5-turbo",
"metadata": {"tags": ["request-tag"]},
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={"tags": ["key-tag", "request-tag"]},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
request_data=data,
user_api_key_dict=user_api_key_dict,
)
# request-tag deduplicated; key-tag appended
assert data["metadata"]["tags"] == ["request-tag", "key-tag"]
def test_no_key_tags_no_mutation(self):
data = {"model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
request_data=data,
user_api_key_dict=user_api_key_dict,
)
assert "metadata" not in data or "tags" not in data.get("metadata", {})
def test_empty_key_metadata_no_mutation(self):
data = {"model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
request_data=data,
user_api_key_dict=user_api_key_dict,
)
assert "metadata" not in data
def test_uses_litellm_metadata_when_present(self):
data = {
"model": "gpt-3.5-turbo",
"litellm_metadata": {"foo": "bar"},
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={"tags": ["key-tag"]},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
request_data=data,
user_api_key_dict=user_api_key_dict,
)
assert data["litellm_metadata"]["tags"] == ["key-tag"]
assert "tags" not in data.get("metadata", {})
def test_string_metadata_parsed_before_merge(self):
data = {
"model": "gpt-3.5-turbo",
"metadata": '{"tags": ["existing"]}',
}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={"tags": ["key-tag"]},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
request_data=data,
user_api_key_dict=user_api_key_dict,
)
assert isinstance(data["metadata"], dict)
assert data["metadata"]["tags"] == ["existing", "key-tag"]
@pytest.mark.asyncio
async def test_key_tags_visible_to_tag_max_budget_check(self):
from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TagTable
from litellm.proxy.auth.auth_checks import _tag_max_budget_check
from litellm.proxy.utils import ProxyLogging
data = {"model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={"tags": ["engineering"]},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
request_data=data,
user_api_key_dict=user_api_key_dict,
)
tag_object = LiteLLM_TagTable(
tag_name="engineering",
spend=0.0,
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.10),
)
async def mock_get_current_spend(counter_key, fallback_spend):
if counter_key == "spend:tag:engineering":
return 0.50
return fallback_spend
with (
patch(
"litellm.proxy.proxy_server.get_current_spend",
mock_get_current_spend,
),
patch(
"litellm.proxy.auth.auth_checks.get_tag_objects_batch",
new_callable=AsyncMock,
return_value={"engineering": tag_object},
),
):
with pytest.raises(litellm.BudgetExceededError) as exc_info:
await _tag_max_budget_check(
request_body=data,
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
proxy_logging_obj=ProxyLogging(user_api_key_cache=None),
valid_token=UserAPIKeyAuth(token="test-token"),
)
assert exc_info.value.current_cost == 0.50
assert exc_info.value.max_budget == 0.10
@pytest.mark.asyncio
async def test_key_tags_within_budget_passes_check(self):
from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TagTable
from litellm.proxy.auth.auth_checks import _tag_max_budget_check
from litellm.proxy.utils import ProxyLogging
data = {"model": "gpt-3.5-turbo"}
user_api_key_dict = UserAPIKeyAuth(
api_key="hashed-key",
metadata={"tags": ["engineering"]},
team_metadata={},
)
LiteLLMProxyRequestSetup.apply_key_tags_pre_auth(
request_data=data,
user_api_key_dict=user_api_key_dict,
)
tag_object = LiteLLM_TagTable(
tag_name="engineering",
spend=0.05,
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.10),
)
async def mock_get_current_spend(counter_key, fallback_spend):
if counter_key == "spend:tag:engineering":
return 0.05
return fallback_spend
with (
patch(
"litellm.proxy.proxy_server.get_current_spend",
mock_get_current_spend,
),
patch(
"litellm.proxy.auth.auth_checks.get_tag_objects_batch",
new_callable=AsyncMock,
return_value={"engineering": tag_object},
),
):
await _tag_max_budget_check(
request_body=data,
prisma_client=MagicMock(),
user_api_key_cache=MagicMock(),
proxy_logging_obj=ProxyLogging(user_api_key_cache=None),
valid_token=UserAPIKeyAuth(token="test-token"),
)
# ============================================================================
# Tests for #27516: provider hint resolution from deployment when the
# user-facing model name has no provider prefix.

View file

@ -42,11 +42,6 @@ def test_bedrock_haiku_4_5_configuration():
model_info.get("supports_vision") is True
), f"{model} should support vision"
# Verify tool use system prompt tokens
assert (
model_info.get("tool_use_system_prompt_tokens") == 346
), f"{model} should have tool_use_system_prompt_tokens set to 346"
# Verify core capabilities
assert model_info.get("supports_computer_use") is True
assert model_info.get("supports_function_calling") is True
@ -96,7 +91,6 @@ def test_bedrock_haiku_4_5_matches_sonnet_capabilities():
"supports_pdf_input",
"supports_assistant_prefill",
"supports_reasoning",
"tool_use_system_prompt_tokens",
]
for capability in shared_capabilities:

View file

@ -82,31 +82,26 @@ def test_opus_4_6_model_pricing_and_capabilities():
"claude-opus-4-6": {
"provider": "anthropic",
"has_long_context_pricing": False,
"tool_use_system_prompt_tokens": 346,
"max_input_tokens": 1000000,
},
"claude-opus-4-6-20260205": {
"provider": "anthropic",
"has_long_context_pricing": False,
"tool_use_system_prompt_tokens": 346,
"max_input_tokens": 1000000,
},
"anthropic.claude-opus-4-6-v1": {
"provider": "bedrock_converse",
"has_long_context_pricing": False,
"tool_use_system_prompt_tokens": 346,
"max_input_tokens": 1000000,
},
"vertex_ai/claude-opus-4-6": {
"provider": "vertex_ai-anthropic_models",
"has_long_context_pricing": False,
"tool_use_system_prompt_tokens": 346,
"max_input_tokens": 1000000,
},
"azure_ai/claude-opus-4-6": {
"provider": "azure_ai",
"has_long_context_pricing": False,
"tool_use_system_prompt_tokens": 159,
"max_input_tokens": 200000,
},
}
@ -143,10 +138,6 @@ def test_opus_4_6_model_pricing_and_capabilities():
assert info["supports_reasoning"] is True
assert info["supports_tool_choice"] is True
assert info["supports_vision"] is True
assert (
info["tool_use_system_prompt_tokens"]
== config["tool_use_system_prompt_tokens"]
)
def test_opus_4_6_bedrock_regional_model_pricing():
@ -191,7 +182,6 @@ def test_opus_4_6_bedrock_regional_model_pricing():
assert info["max_output_tokens"] == 128000
assert info["max_tokens"] == 128000
assert info["supports_assistant_prefill"] is False
assert info["tool_use_system_prompt_tokens"] == 346
assert "input_cost_per_token_above_200k_tokens" not in info
assert "output_cost_per_token_above_200k_tokens" not in info
assert "cache_creation_input_token_cost_above_200k_tokens" not in info
@ -220,7 +210,6 @@ def test_opus_4_6_alias_and_dated_metadata_match():
"cache_creation_input_token_cost_above_1hr",
"cache_read_input_token_cost",
"supports_assistant_prefill",
"tool_use_system_prompt_tokens",
]
for key in keys_to_match:
assert alias[key] == dated[key], f"Mismatch for {key}"

View file

@ -0,0 +1,184 @@
"""
Validate Claude Opus 4.8 model configuration entries.
Regression coverage for the wildcard-routing failure where a bare model name
(``claude-opus-4-8``) could not match an ``anthropic/*`` deployment because
LiteLLM could not infer its provider — the model was simply missing from the
model cost map, so ``get_llm_provider`` raised and the router returned
"no healthy deployments for this model". The fix is the cost-map entries added
for Anthropic, Bedrock, Vertex AI, and Azure AI; those entries are what populate
``litellm.anthropic_models`` at import time, which is what the bare-name lookup
in ``get_llm_provider`` consumes.
"""
import json
import os
import pytest
import litellm
from litellm.constants import BEDROCK_CONVERSE_MODELS
from litellm.litellm_core_utils.get_model_cost_map import GetModelCostMap
REPO_ROOT = os.path.join(os.path.dirname(__file__), "../..")
def _load_root_cost_map() -> dict:
json_path = os.path.join(REPO_ROOT, "model_prices_and_context_window.json")
with open(json_path) as f:
return json.load(f)
@pytest.fixture
def local_model_cost_map(monkeypatch):
"""Force the bundled backup cost map so assertions don't depend on the
network-fetched ``main`` copy (which lags this branch until merge)."""
original_model_cost = litellm.model_cost
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
litellm.model_cost = litellm.get_model_cost_map(url="")
litellm.get_model_info.cache_clear()
try:
yield
finally:
litellm.model_cost = original_model_cost
litellm.get_model_info.cache_clear()
def test_opus_4_8_model_pricing_and_capabilities():
model_data = _load_root_cost_map()
expected_models = {
"claude-opus-4-8": {
"provider": "anthropic",
"max_input_tokens": 1000000,
},
"anthropic.claude-opus-4-8": {
"provider": "bedrock_converse",
"max_input_tokens": 1000000,
},
"vertex_ai/claude-opus-4-8": {
"provider": "vertex_ai-anthropic_models",
"max_input_tokens": 1000000,
},
# Microsoft Foundry / Azure caps Opus 4.8 at a 200k context window.
"azure_ai/claude-opus-4-8": {
"provider": "azure_ai",
"max_input_tokens": 200000,
},
}
for model_name, config in expected_models.items():
assert model_name in model_data, f"Missing model entry: {model_name}"
info = model_data[model_name]
assert info["litellm_provider"] == config["provider"]
assert info["mode"] == "chat"
assert info["max_input_tokens"] == config["max_input_tokens"]
assert info["max_output_tokens"] == 128000
assert info["max_tokens"] == 128000
# Base pricing matches Opus 4.7: $5 / $25 per MTok, with the standard
# 1.25x cache-write and 0.1x cache-read multipliers.
assert info["input_cost_per_token"] == 5e-06
assert info["output_cost_per_token"] == 2.5e-05
assert info["cache_creation_input_token_cost"] == 6.25e-06
assert info["cache_read_input_token_cost"] == 5e-07
# Opus 4.x flagships are flat-rate across the full context window.
assert "input_cost_per_token_above_200k_tokens" not in info
assert "output_cost_per_token_above_200k_tokens" not in info
assert info["supports_assistant_prefill"] is False
assert info["supports_function_calling"] is True
assert info["supports_prompt_caching"] is True
assert info["supports_reasoning"] is True
assert info["supports_tool_choice"] is True
assert info["supports_vision"] is True
def test_opus_4_8_bedrock_regional_model_pricing():
model_data = _load_root_cost_map()
# Global endpoints use base pricing; regional endpoints carry a 10% premium.
expected_models = {
"global.anthropic.claude-opus-4-8": {
"input_cost_per_token": 5e-06,
"output_cost_per_token": 2.5e-05,
"cache_creation_input_token_cost": 6.25e-06,
"cache_read_input_token_cost": 5e-07,
},
"us.anthropic.claude-opus-4-8": {
"input_cost_per_token": 5.5e-06,
"output_cost_per_token": 2.75e-05,
"cache_creation_input_token_cost": 6.875e-06,
"cache_read_input_token_cost": 5.5e-07,
},
"eu.anthropic.claude-opus-4-8": {
"input_cost_per_token": 5.5e-06,
"output_cost_per_token": 2.75e-05,
"cache_creation_input_token_cost": 6.875e-06,
"cache_read_input_token_cost": 5.5e-07,
},
"au.anthropic.claude-opus-4-8": {
"input_cost_per_token": 5.5e-06,
"output_cost_per_token": 2.75e-05,
"cache_creation_input_token_cost": 6.875e-06,
"cache_read_input_token_cost": 5.5e-07,
},
}
for model_name, expected in expected_models.items():
assert model_name in model_data, f"Missing model entry: {model_name}"
info = model_data[model_name]
assert info["litellm_provider"] == "bedrock_converse"
assert info["max_input_tokens"] == 1000000
assert info["max_output_tokens"] == 128000
assert info["bedrock_output_config_effort_ceiling"] == "xhigh"
for key, value in expected.items():
assert info[key] == value
def test_opus_4_8_fast_mode_multiplier():
"""Opus 4.8 dropped fast-mode pricing to 2x base ($10/$50 per MTok);
Opus 4.7 was 6x ($30/$150)."""
model_data = _load_root_cost_map()
entry = model_data["claude-opus-4-8"]["provider_specific_entry"]
assert entry["us"] == 1.1
assert entry["fast"] == 2.0
def test_opus_4_8_present_in_bundled_backup():
"""The bundled backup is the runtime fallback (and what tests load with
``LITELLM_LOCAL_MODEL_COST_MAP=True``) — it must carry the same entries as
the root cost map, otherwise the model resolves on one path but not the
other."""
backup = GetModelCostMap.load_local_model_cost_map()
for model_name in (
"claude-opus-4-8",
"anthropic.claude-opus-4-8",
"global.anthropic.claude-opus-4-8",
"us.anthropic.claude-opus-4-8",
"eu.anthropic.claude-opus-4-8",
"au.anthropic.claude-opus-4-8",
"vertex_ai/claude-opus-4-8",
"vertex_ai/claude-opus-4-8@default",
"azure_ai/claude-opus-4-8",
):
assert model_name in backup, f"Missing from backup cost map: {model_name}"
def test_opus_4_8_registered_for_bedrock_converse():
assert "anthropic.claude-opus-4-8" in BEDROCK_CONVERSE_MODELS
def test_opus_4_8_provider_resolves_via_model_info(local_model_cost_map):
"""Regression: ``claude-opus-4-8`` must resolve to provider ``anthropic``.
Before the cost-map entry existed, the model was unknown to LiteLLM, so it
could not be tied to the ``anthropic`` provider and an ``anthropic/*``
wildcard deployment would not match it.
"""
info = litellm.get_model_info(model="claude-opus-4-8")
assert info["litellm_provider"] == "anthropic"
assert info["max_input_tokens"] == 1000000
assert info["max_output_tokens"] == 128000

View file

@ -50,7 +50,6 @@ def test_bedrock_sonnet_4_6_region_prefixes():
assert model_info.get("supports_pdf_input") is True
assert model_info.get("supports_assistant_prefill") is True
assert model_info.get("supports_reasoning") is True
assert model_info.get("tool_use_system_prompt_tokens") == 346
def test_bedrock_sonnet_4_6_jp_matches_other_regional_pricing():

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