Merge remote-tracking branch 'origin/litellm_internal_staging' into litellm_batch-rate-limiter-skip-file-fetch

# Conflicts:
#	tests/test_litellm/proxy/hooks/test_batch_file_validation.py
This commit is contained in:
mateo-berri 2026-05-30 03:05:25 +00:00
commit 4cda2c5a19
No known key found for this signature in database
102 changed files with 6282 additions and 1750 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?

307
AGENTS.md
View file

@ -1,306 +1 @@
# INSTRUCTIONS FOR LITELLM
This document provides comprehensive instructions for AI agents working in the LiteLLM repository.
## Confidentiality: Customer and Company Names in Code
The codebase is public. Before writing **any** third-party organization name into this repository — in source code, file or directory names, docstrings, comments, tests, fixtures, mock payloads, error messages, log lines, commit messages, or PR descriptions — pause and check:
**Already in the codebase** (OpenAI, Anthropic, Google, Azure, Bedrock, Fireworks, and other established LLM providers / integrations) — fine to use. Quick check: `git grep -i "<name>"` — if it returns hits in real code (not just your current diff), the name is established.
**Anything else** — customers, prospects, partners, new vendor integrations, observability tools, infra vendors, or any organization name that does not already appear in the repo. STOP and surface it to the user. Ask for explicit consent before writing the name into any file, commit message, or PR description. Do not write it speculatively and clean up later. Do not substitute a placeholder and proceed. Do not assume it is safe because it "looks like" a public company. The user must approve first.
**What to do instead of a customer-specific reference:**
- If you find yourself reaching for a customer name — real or fake — step back. The code shouldn't be customer-specific in the first place. Generalize the feature, or capture the customer motivation in internal docs (Notion / Linear / the internal staging PR description), never in the repo.
- Frame changes by the capability they add, not the customer who asked for it ("add per-team Bedrock guardrail routing", not "add routing for $CUSTOMER").
- Standard "fake value" markers (`example.com`, `localhost`, `127.0.0.1`, `test@example.com`) and abstract identifiers (`team_a`, `user_1`, `tenant_x`) are fine — those are not customer stand-ins.
## 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

214
CLAUDE.md
View file

@ -1,194 +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
## Confidentiality: Customer and Company Names in Code
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
The codebase is public. Before writing **any** third-party organization name into this repository — in source code, file or directory names, docstrings, comments, tests, fixtures, mock payloads, error messages, log lines, commit messages, or PR descriptions — pause and check:
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
**Already in the codebase** (OpenAI, Anthropic, Google, Azure, Bedrock, Fireworks, and other established LLM providers / integrations) — fine to use. Quick check: `git grep -i "<name>"` — if it returns hits in real code (not just your current diff), the name is established.
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)
**Anything else** — customers, prospects, partners, new vendor integrations, observability tools, infra vendors, or any organization name that does not already appear in the repo. STOP and surface it to the user. Ask for explicit consent before writing the name into any file, commit message, or PR description. Do not write it speculatively and clean up later. Do not substitute a placeholder and proceed. Do not assume it is safe because it "looks like" a public company. The user must approve first.
When creating PRs, don't set base to `main`. `litellm_internal_staging` serves that purpose
**What to do instead of a customer-specific reference:**
- If you find yourself reaching for a customer name — real or fake — step back. The code shouldn't be customer-specific in the first place. Generalize the feature, or capture the customer motivation in internal docs (Notion / Linear / the internal staging PR description), never in the repo.
- Frame changes by the capability they add, not the customer who asked for it ("add per-team Bedrock guardrail routing", not "add routing for $CUSTOMER").
- Standard "fake value" markers (`example.com`, `localhost`, `127.0.0.1`, `test@example.com`) and abstract identifiers (`team_a`, `user_1`, `tenant_x`) are fine — those are not customer stand-ins.
Always use @.github/pull_request_template.md as a guide for your PR body
## Documentation
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
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.
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
## Development Commands
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
### 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
Run tests, format your code, and lint your code before each commit
### 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
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)
### 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.
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
### 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'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
### Running Scripts
- `uv run python script.py` - Run Python scripts (use for non-test files)
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
### GitHub Issue & PR Templates
When contributing to the project, use the appropriate templates:
When working on a PR, keep the PR description in sync with new commits being made
**Bug Reports** (`.github/ISSUE_TEMPLATE/bug_report.yml`):
- Describe what happened vs. what you expected
- Include relevant log output
- Specify your LiteLLM version
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
**Feature Requests** (`.github/ISSUE_TEMPLATE/feature_request.yml`):
- Describe the feature clearly
- Explain the motivation and use case
Do not put names of customers or customer company names in code, PRs, and issues. The codebase is public
**Pull Requests** (`.github/pull_request_template.md`):
- Add at least 1 test in `tests/litellm/`
- Ensure `make test-unit` passes
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
## Architecture Overview
## Think Before Coding
LiteLLM is a unified interface for 100+ LLM providers with two main components:
**Don't assume. Don't hide confusion. Surface tradeoffs.**
### 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.)
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.
### 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)
## Simplicity First
## Key Patterns
**Minimum code that solves the problem. Nothing speculative.**
### 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
- 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.
### 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
- **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.
### 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.
### 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

@ -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

@ -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

@ -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

@ -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

@ -230,6 +230,8 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
"""Translate legacy ``thinking.type=enabled`` to adaptive for 4.6/4.7.
Caller-provided ``output_config.effort`` is never overridden.
"""
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
if not AnthropicModelInfo._is_adaptive_thinking_model(model):
return
thinking = optional_params.get("thinking")
@ -237,7 +239,7 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
return
budget = int(thinking.get("budget_tokens") or 0)
if budget >= 24000:
if budget >= 24000 and AnthropicConfig._supports_effort_level(model, "xhigh"):
effort = "xhigh"
elif budget >= 10000:
effort = "high"
@ -427,8 +429,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

@ -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

@ -6560,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:
@ -6642,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:
@ -6734,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)
@ -6805,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)
@ -6888,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)
@ -6945,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)
@ -7021,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)
@ -7031,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,
@ -7063,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)
@ -7093,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)
@ -7103,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,
@ -7135,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)
@ -7182,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)
@ -7256,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)
@ -7467,6 +7523,7 @@ class BaseLLMHTTPHandler:
api_key=api_key,
headers=extra_headers or {},
model="",
litellm_params=litellm_params,
)
if extra_headers:

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

@ -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

@ -925,39 +925,67 @@ class MCPRequestHandler:
async def _get_allowed_mcp_servers_for_key(
user_api_key_auth: Optional[UserAPIKeyAuth] = None,
) -> List[str]:
"""
Get allowed MCP servers for a key (the key's own scope).
Unions two sources:
- Legacy key.object_permission (mcp_servers, mcp_access_groups,
mcp_tool_permissions).
- Unified key.access_group_ids → access_group.access_mcp_server_ids.
Mirrors the ungated fallback in can_key_call_model — the group is
attached to the key itself, so it grants the key's own scope (no
assigned_key_ids re-check). The gated, team-ceiling-busting override
lives in _get_key_access_group_mcp_server_extras.
"""
if user_api_key_auth is None:
return []
try:
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_object_permission,
)
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
# Unified key.access_group_ids → MCP servers (ungated: the group is
# attached to the key, so it grants the key's own scope). Entries in
# access_mcp_server_ids may be server_ids OR names/aliases, so expand
# to ids here — matching the legacy object_permission path below.
key_access_group_servers = global_mcp_server_manager.expand_permission_list(
await _get_mcp_server_ids_from_access_groups(
access_group_ids=user_api_key_auth.access_group_ids or [],
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
)
# Get key object permission (already loaded in main auth flow, or fetch from DB)
key_object_permission = MCPRequestHandler._get_key_object_permission(
user_api_key_auth
)
if (
key_object_permission is None
and user_api_key_auth
and user_api_key_auth.object_permission_id
and prisma_client is not None
):
from litellm.proxy.auth.auth_checks import get_object_permission
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
key_object_permission = await get_object_permission(
object_permission_id=user_api_key_auth.object_permission_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 prisma_client is not None:
key_object_permission = await get_object_permission(
object_permission_id=user_api_key_auth.object_permission_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 key_object_permission is None:
return []
return list(set(key_access_group_servers))
# 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,
)
direct_mcp_servers = global_mcp_server_manager.expand_permission_list(
key_object_permission.mcp_servers or []
)
@ -977,7 +1005,12 @@ class MCPRequestHandler:
)
# Combine all lists
all_servers = direct_mcp_servers + access_group_servers + tool_perm_servers
all_servers = (
direct_mcp_servers
+ access_group_servers
+ tool_perm_servers
+ key_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,

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

@ -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

@ -535,9 +535,15 @@ class _PROXY_BatchRateLimiter(CustomLogger):
llm_model_list = llm_router.model_list if llm_router is not None else None
for model in models:
# body.model may be the provider id after replace_model_in_jsonl; map to proxy model_name for auth.
model_to_check = model
if llm_router is not None:
proxy_model_name = llm_router.resolve_model_name_from_model_id(model)
if proxy_model_name is not None:
model_to_check = proxy_model_name
try:
await can_key_call_model(
model=model,
model=model_to_check,
llm_model_list=llm_model_list,
valid_token=user_api_key_dict,
llm_router=llm_router,
@ -553,7 +559,7 @@ class _PROXY_BatchRateLimiter(CustomLogger):
detail={
"error": (
"Batch input file references a model the caller is "
f"not authorized to use: model={model}, reason={str(e)}"
f"not authorized to use: model={model_to_check}, reason={str(e)}"
)
},
)

View file

@ -28,6 +28,7 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting):
def __init__(self, dual_cache: DualCache):
self.dual_cache = dual_cache
self.redis_increment_operation_queue = []
self.deployment_budget_config = None
async def is_key_within_model_budget(
self,

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,

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

@ -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

@ -21,7 +21,6 @@ from litellm.proxy.spend_tracking.spend_tracking_utils import (
get_spend_by_team_and_customer,
)
from litellm.proxy.utils import handle_exception_on_proxy
from litellm.router_strategy.budget_limiter import RouterBudgetLimiting
if TYPE_CHECKING:
from litellm.proxy.proxy_server import PrismaClient
@ -3149,18 +3148,12 @@ async def provider_budgets() -> ProviderBudgetResponse:
"No provider budget config found. Please set a provider budget config in the router settings. https://docs.litellm.ai/docs/proxy/provider_budget_routing"
)
router_budget_logger = llm_router._get_router_deployment_budget_limiter()
if router_budget_logger is None:
raise ValueError("No router budget logger found")
provider_budget_response_dict: Dict[str, ProviderBudgetResponseObject] = {}
for _provider, _budget_info in provider_budget_config.items():
router_budget_logger = next(
(
cb
for cb in (llm_router.optional_callbacks or [])
if isinstance(cb, RouterBudgetLimiting)
),
None,
)
if router_budget_logger is None:
raise ValueError("No router budget logger found")
_provider_spend = (
await router_budget_logger._get_current_provider_spend(_provider) or 0.0
)

View file

@ -1753,11 +1753,14 @@ class Router:
if pre_call_check == "prompt_caching":
_callback = PromptCachingDeploymentCheck(cache=self.cache)
elif pre_call_check == "router_budget_limiting":
if self._get_router_deployment_budget_limiter() is not None:
continue
_callback = RouterBudgetLimiting(
dual_cache=self.cache,
provider_budget_config=self.provider_budget_config,
model_list=self.model_list,
)
self.router_budget_logger = _callback
elif pre_call_check == "enforce_model_rate_limits":
_callback = ModelRateLimitingCheck(dual_cache=self.cache)
@ -8529,6 +8532,7 @@ class Router:
model=_deployment, model_id=deployment.model_info.id
)
self.model_names.add(deployment.model_name)
self._sync_deployment_budget_config(deployment=deployment)
return deployment
def _update_deployment_indices_after_removal(
@ -8717,12 +8721,64 @@ class Router:
self._update_deployment_indices_after_removal(
model_id=id, removal_idx=deployment_idx
)
_budget_limiter = self._get_router_deployment_budget_limiter()
if _budget_limiter is not None:
_budget_limiter.unregister_deployment_budget(model_id=id)
return item
else:
return None
except Exception:
return None
def _get_router_deployment_budget_limiter(
self,
) -> Optional[RouterBudgetLimiting]:
"""
Return the router's deployment-budget callback.
Uses exact-type matching so proxy subclasses (e.g. virtual-key model budgets)
registered on litellm.callbacks are not mistaken for router deployment budgets.
"""
if self.router_budget_logger is not None:
return self.router_budget_logger
if self.optional_callbacks:
for _cb in self.optional_callbacks:
if type(_cb) is RouterBudgetLimiting:
self.router_budget_logger = _cb
return _cb
return None
def _deployment_has_budget_limits(self, deployment: Deployment) -> bool:
return (
deployment.litellm_params.get("max_budget") is not None
and deployment.litellm_params.get("budget_duration") is not None
and deployment.model_info.id is not None
)
def _sync_deployment_budget_config(self, deployment: Deployment) -> None:
model_id = deployment.model_info.id
if model_id is None:
return
_budget_limiter = self._get_router_deployment_budget_limiter()
if not self._deployment_has_budget_limits(deployment=deployment):
if _budget_limiter is not None:
_budget_limiter.unregister_deployment_budget(model_id=model_id)
return
if _budget_limiter is None:
self.add_optional_pre_call_checks(
optional_pre_call_checks=["router_budget_limiting"]
)
_budget_limiter = self._get_router_deployment_budget_limiter()
if _budget_limiter is not None:
_budget_limiter.register_deployment_budget(
deployment=deployment.to_json(exclude_none=True)
)
def get_deployment(self, model_id: str) -> Optional[Deployment]:
"""
Returns -> Deployment or None

View file

@ -96,9 +96,7 @@ class RouterBudgetLimiting(CustomLogger):
self,
dual_cache: DualCache,
provider_budget_config: Optional[dict],
model_list: Optional[
Union[List[DeploymentTypedDict], List[Dict[str, Any]]]
] = None,
model_list: Optional[List[Union[DeploymentTypedDict, Dict[str, Any]]]] = None,
):
self.dual_cache = dual_cache
self.redis_increment_operation_queue: List[RedisPipelineIncrementOperation] = []
@ -854,9 +852,7 @@ class RouterBudgetLimiting(CustomLogger):
def _init_deployment_budgets(
self,
model_list: Optional[
Union[List[DeploymentTypedDict], List[Dict[str, Any]]]
] = None,
model_list: Optional[List[Union[DeploymentTypedDict, Dict[str, Any]]]] = None,
):
if model_list is None:
return
@ -887,6 +883,22 @@ class RouterBudgetLimiting(CustomLogger):
f"Initialized Deployment Budget Config: {self.deployment_budget_config}"
)
def register_deployment_budget(
self,
deployment: Union[Dict[str, Any], DeploymentTypedDict],
) -> None:
"""
Register or refresh deployment-level budget config for a runtime-added deployment.
"""
self._init_deployment_budgets(model_list=[deployment])
def unregister_deployment_budget(self, model_id: str) -> None:
if self.deployment_budget_config is None:
return
self.deployment_budget_config.pop(model_id, None)
if len(self.deployment_budget_config) == 0:
self.deployment_budget_config = None
def _init_tag_budgets(self):
if litellm.tag_budget_config is None:
return

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

@ -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

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

@ -14,7 +14,7 @@ import litellm
from unittest.mock import patch, MagicMock, AsyncMock
from create_mock_standard_logging_payload import create_standard_logging_payload
from litellm.types.utils import StandardLoggingPayload
from litellm.types.router import Deployment, LiteLLM_Params
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
@pytest.fixture
@ -997,9 +997,7 @@ def test_filter_cooldown_deployments(model_list):
healthy_deployments=router._get_all_deployments(model_name="gpt-5-mini"), # type: ignore
cooldown_deployments=[],
)
assert len(deployments) == len(
router._get_all_deployments(model_name="gpt-5-mini")
)
assert len(deployments) == len(router._get_all_deployments(model_name="gpt-5-mini"))
def test_track_deployment_metrics(model_list):
@ -2379,3 +2377,123 @@ def test_get_router_model_info_with_deployment_object():
# Verify we got valid model info back
assert model_info is not None
assert isinstance(model_info, dict)
def test_deployment_has_budget_limits():
router = Router(model_list=[])
with_budget = Deployment(
model_name="budgeted-model",
litellm_params=LiteLLM_Params(
model="openai/gpt-4o-mini",
max_budget=0.001,
budget_duration="1d",
),
model_info=ModelInfo(id="budget-deployment-id"),
)
without_budget = Deployment(
model_name="unbudgeted-model",
litellm_params=LiteLLM_Params(model="openai/gpt-4o-mini"),
model_info=ModelInfo(id="no-budget-deployment-id"),
)
assert router._deployment_has_budget_limits(deployment=with_budget) is True
assert router._deployment_has_budget_limits(deployment=without_budget) is False
def test_sync_deployment_budget_config(monkeypatch):
import asyncio
monkeypatch.setattr(asyncio, "create_task", lambda coro: None)
router = Router(model_list=[], optional_pre_call_checks=[])
deployment = Deployment(
model_name="dynamic-budget-model",
litellm_params=LiteLLM_Params(
model="openai/gpt-4o-mini",
api_key="fake-key",
max_budget=0.000000000001,
budget_duration="1d",
),
model_info=ModelInfo(id="runtime-budget-deployment"),
)
router._sync_deployment_budget_config(deployment=deployment)
budget_limiter = router._get_router_deployment_budget_limiter()
assert budget_limiter is not None
config = budget_limiter._get_budget_config_for_deployment(
"runtime-budget-deployment"
)
assert config is not None
assert config.max_budget == 0.000000000001
def test_sync_deployment_budget_config_clears_removed_limits(monkeypatch):
import asyncio
monkeypatch.setattr(asyncio, "create_task", lambda coro: None)
router = Router(model_list=[], optional_pre_call_checks=[])
model_id = "runtime-budget-deployment"
budgeted = Deployment(
model_name="dynamic-budget-model",
litellm_params=LiteLLM_Params(
model="openai/gpt-4o-mini",
api_key="fake-key",
max_budget=0.000000000001,
budget_duration="1d",
),
model_info=ModelInfo(id=model_id),
)
unbudgeted = Deployment(
model_name="dynamic-budget-model",
litellm_params=LiteLLM_Params(
model="openai/gpt-4o-mini",
api_key="fake-key",
),
model_info=ModelInfo(id=model_id),
)
router._sync_deployment_budget_config(deployment=budgeted)
budget_limiter = router._get_router_deployment_budget_limiter()
assert budget_limiter is not None
assert budget_limiter._get_budget_config_for_deployment(model_id) is not None
router._sync_deployment_budget_config(deployment=unbudgeted)
assert budget_limiter._get_budget_config_for_deployment(model_id) is None
def test_upsert_deployment_clears_stale_budget_config(monkeypatch):
import asyncio
monkeypatch.setattr(asyncio, "create_task", lambda coro: None)
router = Router(model_list=[], optional_pre_call_checks=[])
model_id = "upsert-budget-deployment"
budgeted = Deployment(
model_name="dynamic-budget-model",
litellm_params=LiteLLM_Params(
model="openai/gpt-4o-mini",
api_key="fake-key",
max_budget=0.000000000001,
budget_duration="1d",
),
model_info=ModelInfo(id=model_id),
)
unbudgeted = Deployment(
model_name="dynamic-budget-model",
litellm_params=LiteLLM_Params(
model="openai/gpt-4o-mini",
api_key="fake-key",
),
model_info=ModelInfo(id=model_id),
)
router.upsert_deployment(deployment=budgeted)
budget_limiter = router._get_router_deployment_budget_limiter()
assert budget_limiter is not None
assert budget_limiter._get_budget_config_for_deployment(model_id) is not None
router.upsert_deployment(deployment=unbudgeted)
assert budget_limiter._get_budget_config_for_deployment(model_id) is None

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,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",
[
@ -191,3 +243,124 @@ def test_reasoning_effort_in_supported_params():
assert "reasoning_effort" in config.get_supported_anthropic_messages_params(
"claude-opus-4-7"
)
@pytest.mark.parametrize(
"model",
[
"claude-sonnet-4-6",
"bedrock/invoke/us.anthropic.claude-sonnet-4-6",
"vertex_ai/claude-sonnet-4-6",
"claude-opus-4-6",
"bedrock/invoke/us.anthropic.claude-opus-4-6",
"vertex_ai/claude-opus-4-6",
],
)
def test_legacy_thinking_high_budget_clamps_to_high_when_xhigh_unsupported(model):
"""Claude Code sends ``thinking.budget_tokens=31999``; Sonnet 4.6 and Opus 4.6
have no ``xhigh`` tier, so the translator must emit ``high`` rather than the
provider-invalid ``xhigh`` (regression for issue #29282)."""
config = AnthropicMessagesConfig()
optional_params = {
"max_tokens": 1024,
"thinking": {"type": "enabled", "budget_tokens": 31999},
}
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.get("thinking") == {"type": "adaptive"}
assert result.get("output_config") == {"effort": "high"}
def test_legacy_thinking_high_budget_keeps_xhigh_when_supported():
"""Opus 4.7 advertises an ``xhigh`` tier, so the high-budget bucket keeps it."""
config = AnthropicMessagesConfig()
optional_params = {
"max_tokens": 1024,
"thinking": {"type": "enabled", "budget_tokens": 31999},
}
result = config.transform_anthropic_messages_request(
model="claude-opus-4-7",
messages=[{"role": "user", "content": "Hello"}],
anthropic_messages_optional_request_params=optional_params,
litellm_params={},
headers={},
)
assert result.get("thinking") == {"type": "adaptive"}
assert result.get("output_config") == {"effort": "xhigh"}
@pytest.mark.parametrize(
"budget_tokens,expected_effort",
[
(31999, "high"),
(24000, "high"),
(10000, "high"),
(9999, "medium"),
(5000, "medium"),
(4999, "low"),
(1024, "low"),
],
)
def test_legacy_thinking_budget_buckets_on_sonnet_46(budget_tokens, expected_effort):
config = AnthropicMessagesConfig()
optional_params = {
"max_tokens": 1024,
"thinking": {"type": "enabled", "budget_tokens": budget_tokens},
}
result = config.transform_anthropic_messages_request(
model="claude-sonnet-4-6",
messages=[{"role": "user", "content": "Hello"}],
anthropic_messages_optional_request_params=optional_params,
litellm_params={},
headers={},
)
assert result.get("output_config") == {"effort": expected_effort}
def test_legacy_thinking_does_not_override_explicit_output_config():
config = AnthropicMessagesConfig()
optional_params = {
"max_tokens": 1024,
"thinking": {"type": "enabled", "budget_tokens": 31999},
"output_config": {"effort": "low"},
}
result = config.transform_anthropic_messages_request(
model="claude-sonnet-4-6",
messages=[{"role": "user", "content": "Hello"}],
anthropic_messages_optional_request_params=optional_params,
litellm_params={},
headers={},
)
assert result.get("output_config") == {"effort": "low"}
def test_legacy_thinking_left_untouched_on_non_adaptive_model():
config = AnthropicMessagesConfig()
optional_params = {
"max_tokens": 1024,
"thinking": {"type": "enabled", "budget_tokens": 31999},
}
result = config.transform_anthropic_messages_request(
model="claude-opus-4-5",
messages=[{"role": "user", "content": "Hello"}],
anthropic_messages_optional_request_params=optional_params,
litellm_params={},
headers={},
)
assert result.get("thinking") == {"type": "enabled", "budget_tokens": 31999}
assert "output_config" not in result

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

@ -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

@ -3641,3 +3641,150 @@ async def test_get_allowed_mcp_servers_includes_team_access_group_extras_end_to_
):
result = await MCPRequestHandler.get_allowed_mcp_servers(auth)
assert result == ["srv-stripe"]
@pytest.mark.asyncio
async def test_key_access_group_ids_resolves_mcp_servers_ungated():
"""A teamless key whose unified access_group_ids grant an MCP server sees it
even though the group lists the key in NEITHER assigned_key_ids NOR
assigned_team_ids — the group is attached to the key, so it grants the key's
own scope (ungated, mirroring can_key_call_model's fallback)."""
auth = UserAPIKeyAuth(
token="test-token-hash",
api_key="sk-test",
access_group_ids=["mcp-premium"],
)
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
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_key(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_key_access_group_ids_union_with_object_permission():
"""When both legacy key.object_permission and unified key.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
from litellm.types.mcp import MCPTransport
from litellm.types.mcp_server.mcp_server_manager import MCPServer
global_mcp_server_manager.registry["srv-direct"] = MCPServer(
server_id="srv-direct",
name="srv-direct",
server_name="srv-direct",
url="https://srv-direct.example.com",
transport=MCPTransport.http,
)
try:
perms = LiteLLM_ObjectPermissionTable(
object_permission_id="perm-1",
mcp_servers=["srv-direct"],
mcp_access_groups=[],
vector_stores=[],
)
auth = UserAPIKeyAuth(
token="test-token-hash",
api_key="sk-test",
access_group_ids=["mcp-premium"],
object_permission=perms,
)
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
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_key(auth)
assert set(result) == {"srv-direct", "srv-stripe"}
finally:
global_mcp_server_manager.registry.pop("srv-direct", None)
@pytest.mark.asyncio
async def test_key_access_group_ids_empty_returns_no_extras():
"""Empty key.access_group_ids and no object_permission → resolver called with
[], short-circuits without DB access, returns []."""
auth = UserAPIKeyAuth(
token="test-token-hash",
api_key="sk-test",
access_group_ids=[],
)
with (
patch("litellm.proxy.proxy_server.prisma_client", MagicMock()),
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_key(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_key_access_group_base_end_to_end():
"""End-to-end bug repro: a teamless key has an MCP-granting access group on
its access_group_ids, but the group lists the key in NEITHER assigned_key_ids
NOR assigned_team_ids. The gated extras path returns [] (no override), yet the
ungated base key path grants the server → the key sees it through
get_allowed_mcp_servers."""
auth = UserAPIKeyAuth(
token="test-token",
api_key="sk-test",
access_group_ids=["mcp-group"],
)
# Group grants the server but admits neither this key nor its (absent) team.
fake_ag = _fake_mcp_access_group(
access_group_id="mcp-group",
access_mcp_server_ids=["srv-deepwiki"],
assigned_team_ids=[],
assigned_key_ids=[],
)
patches = _patch_proxy_server_globals_for_mcp() + [
# Ungated base resolver used by _get_allowed_mcp_servers_for_key.
patch(
"litellm.proxy.auth.auth_checks._get_mcp_server_ids_from_access_groups",
new_callable=AsyncMock,
return_value=["srv-deepwiki"],
),
# Gated path (_get_key_access_group_mcp_server_extras) resolves the group
# via get_access_object; empty assigned_* → it contributes nothing.
patch(
"litellm.proxy.auth.auth_checks.get_access_object",
new_callable=AsyncMock,
return_value=fake_ag,
),
]
_start_patches(patches)
try:
# Sanity: the gated extras path alone denies (the old behavior).
extras = await MCPRequestHandler._get_key_access_group_mcp_server_extras(auth)
assert extras == []
# But the key now sees the server via the ungated base path.
result = await MCPRequestHandler.get_allowed_mcp_servers(auth)
assert result == ["srv-deepwiki"]
finally:
_stop_patches(patches)

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

@ -365,6 +365,47 @@ async def test_count_input_file_usage_decodes_model_embedded_file_id():
assert mock_afile_content.await_args.kwargs["custom_llm_provider"] == "hosted_vllm"
@pytest.mark.asyncio
async def test_pre_call_allows_stripped_provider_model_when_key_has_proxy_alias():
"""After replace_model_in_jsonl, body.model is the provider id (e.g. gpt-5.5).
Auth must check the proxy model_name the key was granted, not the stripped id."""
from litellm.proxy.hooks.batch_rate_limiter import _PROXY_BatchRateLimiter
rate_limiter = _PROXY_BatchRateLimiter(
internal_usage_cache=MagicMock(),
parallel_request_limiter=MagicMock(),
)
proxy_alias = "openai/openai/gpt-5.5-batch"
file_dict = [
{"body": {"model": "gpt-5.5", "messages": [{"role": "user", "content": "x"}]}}
]
user = UserAPIKeyAuth(
api_key="sk-ok",
user_id="alice",
models=[proxy_alias],
user_role=LitellmUserRoles.INTERNAL_USER.value,
)
mock_router = MagicMock()
mock_router.model_list = []
mock_router.resolve_model_name_from_model_id.return_value = proxy_alias
can_key_call_model = AsyncMock(return_value=True)
with (
patch(
"litellm.proxy.auth.auth_checks.can_key_call_model",
new=can_key_call_model,
),
patch("litellm.proxy.proxy_server.llm_router", mock_router),
):
await rate_limiter._enforce_batch_file_model_access(
user_api_key_dict=user,
file_content_as_dict=file_dict,
)
can_key_call_model.assert_awaited_once()
assert can_key_call_model.await_args.kwargs["model"] == proxy_alias
@pytest.mark.asyncio
async def test_pre_call_skips_check_when_no_models_present():
"""Files without any `body.model` (corrupt or empty) must not 500;

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

@ -4235,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

@ -234,3 +234,72 @@ async def test_get_llm_provider_for_deployment_matches_legacy_behavior(
legacy_provider = _legacy_provider_resolution(deployment)
assert current_provider == legacy_provider
def test_register_deployment_budget_for_runtime_added_deployment(
disable_budget_sync, monkeypatch
):
import asyncio
monkeypatch.setattr(asyncio, "create_task", lambda coro: None)
budget_limiter = RouterBudgetLimiting(
dual_cache=DualCache(),
provider_budget_config={},
)
model_id = "dynamic-deployment-id"
budget_limiter.register_deployment_budget(
deployment={
"model_name": "dynamic-budget-model",
"litellm_params": {
"model": "openai/gpt-4o-mini",
"max_budget": 0.000000000001,
"budget_duration": "1d",
},
"model_info": {"id": model_id},
}
)
config = budget_limiter._get_budget_config_for_deployment(model_id)
assert config is not None
assert config.max_budget == 0.000000000001
assert config.budget_duration == "1d"
budget_limiter.unregister_deployment_budget(model_id=model_id)
assert budget_limiter._get_budget_config_for_deployment(model_id) is None
def test_router_add_deployment_registers_deployment_budget(
disable_budget_sync, monkeypatch
):
import asyncio
from litellm import Router
from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo
monkeypatch.setattr(asyncio, "create_task", lambda coro: None)
router = Router(
model_list=[],
optional_pre_call_checks=[],
)
router.add_deployment(
deployment=Deployment(
model_name="dynamic-budget-model",
litellm_params=LiteLLM_Params(
model="openai/gpt-4o-mini",
api_key="fake-key",
max_budget=0.000000000001,
budget_duration="1d",
),
model_info=ModelInfo(id="runtime-budget-deployment"),
)
)
budget_limiter = router._get_router_deployment_budget_limiter()
assert budget_limiter is not None
config = budget_limiter._get_budget_config_for_deployment(
"runtime-budget-deployment"
)
assert config is not None
assert config.max_budget == 0.000000000001

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():

View file

@ -859,8 +859,11 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"supports_adaptive_thinking": {"type": "boolean"},
"supports_service_tier": {"type": "boolean"},
"supports_preset": {"type": "boolean"},
"supports_output_config": {"type": "boolean"},
"tool_use_system_prompt_tokens": {"type": "number"},
"supports_output_config": {"type": "boolean"},
"bedrock_output_config_effort_ceiling": {
"type": "string",
"enum": ["low", "medium", "high", "max", "xhigh"],
},
"tpm": {"type": "number"},
"provider_specific_entry": {"type": "object"},
"supported_endpoints": {

View file

@ -398,6 +398,34 @@ class TestVideoGeneration:
)
assert abs(cost - 0.8) < 0.001
def test_completion_cost_video_edit_uses_video_calculator(self):
"""video_edit is charged via the same video cost path as create_video."""
from litellm.cost_calculator import completion_cost
mock_response = MagicMock()
mock_response.usage = MagicMock()
mock_response.usage.duration_seconds = 10.0
type(mock_response)._hidden_params = {}
mock_logging_obj = MagicMock()
mock_logging_obj.litellm_params = {
"metadata": {
"model_info": {
"output_cost_per_video_per_second": 0.05,
}
}
}
cost = completion_cost(
completion_response=mock_response,
model="vertex_ai/veo-3.1-generate-001",
call_type="video_edit",
custom_llm_provider="vertex_ai",
custom_pricing=True,
litellm_logging_obj=mock_logging_obj,
)
assert cost == 0.5
def test_video_generation_with_files(self):
"""Test video generation with file uploads."""
config = OpenAIVideoConfig()

View file

@ -0,0 +1 @@
Never put LiteLLM tokens or API keys in `localStorage`. `localStorage` survives browser close. Prefer `httpOnly` cookies, or `sessionStorage` at most, understanding that any web storage is readable by injected scripts (XSS), and only httpOnly cookies are not

View file

@ -33,6 +33,7 @@ VALUES
('e2e-internal-viewer', 'viewer@test.local', 'internal_user_viewer', '{"e2e-team-crud"}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
('e2e-team-admin', 'teamadmin@test.local', 'internal_user', '{"e2e-team-crud","e2e-team-delete"}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
('e2e-invitable-user', 'invitable@test.local', 'internal_user', '{}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
('e2e-invitable-by-team-admin', 'invitable-team@test.local', 'internal_user', '{}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr'),
('e2e-removable-member', 'removable@test.local', 'internal_user', '{"e2e-team-crud"}', 'scrypt:MU5CcTAi6rVK1HfY1rVPEWq6r4sxg837eq9dG4n5Q6BhDJ44442+seC6LAhLEAYr');
-- 5. Teams (members_with_roles is required JSON)

View file

@ -23,3 +23,15 @@ export async function dismissFeedbackPopup(page: PlaywrightPage): Promise<void>
await expect(dismissButton).not.toBeVisible({ timeout: 2_000 }).catch(() => {});
}
}
/**
* Click on a team ID in the table. Team IDs are rendered differently depending
* on the component version — try button first (Tremor Button), fall back to
* clickable span (OldTeams Typography.Text).
*/
export async function clickTeamId(page: PlaywrightPage, teamId: string): Promise<void> {
const cell = page.locator("td").filter({ hasText: teamId }).first();
await expect(cell).toBeVisible({ timeout: 10_000 });
await cell.click();
await expect(page.getByText("Back to Teams")).toBeVisible({ timeout: 10_000 });
}

View file

@ -0,0 +1,30 @@
import { test, expect } from "@playwright/test";
import { ADMIN_STORAGE_PATH } from "../../constants";
test.describe("Logout", () => {
test.use({ storageState: ADMIN_STORAGE_PATH });
test("Clicking Logout clears the session and forces re-login on a protected page", async ({ page }) => {
await page.goto("/ui");
await expect(page.getByText("Virtual Keys")).toBeVisible({ timeout: 10_000 });
// Open the navbar User dropdown. The trigger button exposes an aria-label
// of "Account menu — <role> — signed in as <email>", and the antd Dropdown
// is declared with trigger={["click"]}, so a plain click opens the popup.
await page.getByRole("button", { name: /Account menu/i }).click();
const popup = page.locator(".ant-dropdown:visible").filter({
has: page.locator(".bg-white.rounded-lg.shadow-lg"),
}).first();
await expect(popup).toBeVisible({ timeout: 5_000 });
// Click Logout — the handler clears the auth cookie and navigates via
// window.location.href = PROXY_LOGOUT_URL (empty string in the e2e env).
await popup.getByText("Logout", { exact: true }).click();
// The cookie is now gone — visiting a protected page must redirect to /ui/login.
await page.goto("/ui?page=llm-playground", { waitUntil: "domcontentloaded" });
await expect(page).toHaveURL(/\/ui\/login/);
await expect(page.getByRole("heading", { name: "Login" })).toBeVisible({ timeout: 10_000 });
});
});

View file

@ -0,0 +1,55 @@
import { test, expect } from "@playwright/test";
import {
E2E_INTERNAL_USER_KEY_ALIAS,
E2E_TEAM_CRUD_ALIAS,
E2E_TEAM_CRUD_ID,
INTERNAL_USER_STORAGE_PATH,
} from "../../constants";
import { Page } from "../../fixtures/pages";
import { navigateToPage, clickTeamId } from "../../helpers/navigation";
test.describe("Internal User", () => {
test.use({ storageState: INTERNAL_USER_STORAGE_PATH });
test("Create Key modal shows the team dropdown populated with the user's teams", async ({ page }) => {
await navigateToPage(page, Page.ApiKeys);
await page.getByRole("button", { name: /Create New Key/i }).click();
await expect(page.getByText("Key Ownership")).toBeVisible({ timeout: 10_000 });
// Open the team dropdown — seeded internal user is a member of
// e2e-team-crud and e2e-team-org, so we expect at least the CRUD alias.
const teamSelect = page.locator(".ant-select", { hasText: "Search or select a team" });
await teamSelect.click();
await page.keyboard.type(E2E_TEAM_CRUD_ALIAS);
await expect(
page.locator(".ant-select-dropdown:visible").getByText(E2E_TEAM_CRUD_ALIAS).first(),
).toBeVisible({ timeout: 5_000 });
});
test("Team info page omits the Settings tab for non-admin members", async ({ page }) => {
await navigateToPage(page, Page.Teams);
await clickTeamId(page, E2E_TEAM_CRUD_ID);
// Overview / My User / Virtual Keys are always visible; Settings is gated
// on canEditTeam and must NOT render for a regular team member.
await expect(page.getByRole("tab", { name: "Overview" })).toBeVisible({ timeout: 5_000 });
await expect(page.getByRole("tab", { name: "Settings" })).not.toBeVisible();
await expect(page.getByRole("tab", { name: "Members" })).not.toBeVisible();
});
test("Virtual Keys page does not surface litellm-dashboard team keys", async ({ page }) => {
await navigateToPage(page, Page.ApiKeys);
// Anchor on the user's own seeded key so the absence check below cannot
// pass vacuously against an empty table.
await expect(
page.locator("table tbody").getByText(E2E_INTERNAL_USER_KEY_ALIAS).first(),
).toBeVisible({ timeout: 10_000 });
// The litellm-dashboard team is the proxy's internal bookkeeping team —
// its keys must never leak into an internal user's Virtual Keys table.
await expect(page.locator("table tbody").getByText("litellm-dashboard")).toHaveCount(0);
});
});

View file

@ -0,0 +1,90 @@
import { test, expect } from "@playwright/test";
import {
E2E_TEAM_CRUD_ID,
E2E_VIEWER_KEY_ALIAS,
INTERNAL_VIEWER_STORAGE_PATH,
} from "../../constants";
import { Page } from "../../fixtures/pages";
import { navigateToPage } from "../../helpers/navigation";
async function clickTeamId(page: import("@playwright/test").Page, teamId: string) {
const cell = page.locator("td").filter({ hasText: teamId }).first();
await expect(cell).toBeVisible({ timeout: 10_000 });
await cell.click();
await expect(page.getByText("Back to Teams")).toBeVisible({ timeout: 10_000 });
}
test.describe("Internal Viewer", () => {
test.use({ storageState: INTERNAL_VIEWER_STORAGE_PATH });
test("Nav shows only the allowed options for the Internal Viewer role", async ({ page }) => {
// Use navigateToPage so the networkidle wait lets the async role-gated nav
// settle before we assert — a bare page.goto races the permission fetch.
await navigateToPage(page, Page.ApiKeys);
// Scope to the sidebar and match items by their link role + accessible
// name. The sidebar is a `complementary` landmark (the `navigation` role
// is the top bar), and each item renders as a link inside it — far tighter
// than a CSS `nav, aside` selector or a getByText on stray text nodes.
const nav = page.getByRole("complementary");
// Items that must be visible per the manual-QA checklist
const expectedVisible = [
"Virtual Keys",
"MCP Servers",
"Guardrails",
"Usage",
"Logs",
"Teams",
"API Reference",
"AI Hub",
];
for (const label of expectedVisible) {
await expect(
nav.getByRole("link", { name: label, exact: true }).first(),
`expected nav item "${label}" to render for Internal Viewer`,
).toBeVisible({ timeout: 5_000 });
}
// Items that must NOT be visible (admin-only surface)
const expectedHidden = ["Internal Users", "Organizations", "Models + Endpoints"];
for (const label of expectedHidden) {
await expect(
nav.getByRole("link", { name: label, exact: true }),
`nav item "${label}" must not render for Internal Viewer`,
).toHaveCount(0);
}
});
test("Virtual Keys page hides Create / Regenerate / Reset / Delete controls", async ({ page }) => {
await navigateToPage(page, Page.ApiKeys);
// Create button is gated on rolesWithWriteAccess (Internal Viewer is not in it)
await expect(page.getByRole("button", { name: /Create New Key/i })).toHaveCount(0);
// Open the viewer's own key info page
const keyRow = page.locator("tr", { hasText: E2E_VIEWER_KEY_ALIAS });
await expect(keyRow).toBeVisible({ timeout: 10_000 });
await keyRow.locator("button").first().click();
await expect(page.getByText("Back to Keys")).toBeVisible({ timeout: 10_000 });
// None of the destructive / mutating actions should render
await expect(page.getByRole("button", { name: "Regenerate Key" })).toHaveCount(0);
await expect(page.getByRole("button", { name: /Reset Spend/i })).toHaveCount(0);
await expect(page.getByRole("button", { name: "Delete Key" })).toHaveCount(0);
});
test("Team info page omits Members and Settings tabs for an Internal Viewer", async ({ page }) => {
await navigateToPage(page, Page.Teams);
await clickTeamId(page, E2E_TEAM_CRUD_ID);
// Overview / Virtual Keys are always visible; Settings + Members are not.
// Tabs are conditionally rendered (getTeamInfoVisibleTabs filters the list),
// so assert absence from the DOM with toHaveCount(0) to match the nav block.
await expect(page.getByRole("tab", { name: "Overview" })).toBeVisible({ timeout: 5_000 });
await expect(page.getByRole("tab", { name: "Virtual Keys" })).toBeVisible({ timeout: 5_000 });
await expect(page.getByRole("tab", { name: "Settings" })).toHaveCount(0);
await expect(page.getByRole("tab", { name: "Members" })).toHaveCount(0);
});
});

View file

@ -7,19 +7,7 @@ import {
E2E_TEAM_ORG_ID,
} from "../../constants";
import { Page } from "../../fixtures/pages";
import { navigateToPage, dismissFeedbackPopup } from "../../helpers/navigation";
/**
* Click on a team ID in the table. Team IDs are rendered differently depending
* on the component version — try button first (Tremor Button), fall back to
* clickable span (OldTeams Typography.Text).
*/
async function clickTeamId(page: import("@playwright/test").Page, teamId: string) {
const cell = page.locator("td").filter({ hasText: teamId }).first();
await expect(cell).toBeVisible({ timeout: 10_000 });
await cell.click();
await expect(page.getByText("Back to Teams")).toBeVisible({ timeout: 10_000 });
}
import { navigateToPage, dismissFeedbackPopup, clickTeamId } from "../../helpers/navigation";
test.describe("Proxy Admin - Teams", () => {
test.use({ storageState: ADMIN_STORAGE_PATH });

View file

@ -0,0 +1,117 @@
import { test, expect } from "@playwright/test";
import {
E2E_INTERNAL_USER_KEY_ALIAS,
E2E_TEAM_CRUD_ALIAS,
E2E_TEAM_CRUD_ID,
TEAM_ADMIN_STORAGE_PATH,
} from "../../constants";
import { Page } from "../../fixtures/pages";
import { navigateToPage, dismissFeedbackPopup } from "../../helpers/navigation";
async function clickTeamId(page: import("@playwright/test").Page, teamId: string) {
const cell = page.locator("td").filter({ hasText: teamId }).first();
await expect(cell).toBeVisible({ timeout: 10_000 });
await cell.click();
await expect(page.getByText("Back to Teams")).toBeVisible({ timeout: 10_000 });
}
test.describe("Team Admin", () => {
test.use({ storageState: TEAM_ADMIN_STORAGE_PATH });
test("Team admin can see all team keys including internal user keys", async ({ page }) => {
// Step from the manual-QA checklist: navigate into the team info page,
// open the Virtual Keys tab, and confirm a key belonging to another
// team member (the seeded internal user) is visible.
await navigateToPage(page, Page.Teams);
await dismissFeedbackPopup(page);
await clickTeamId(page, E2E_TEAM_CRUD_ID);
await page.getByRole("tab", { name: "Virtual Keys" }).click();
await expect(page.getByText(E2E_INTERNAL_USER_KEY_ALIAS).first())
.toBeVisible({ timeout: 10_000 });
// And from the global Virtual Keys page, the same key should be visible.
await navigateToPage(page, Page.ApiKeys);
await expect(page.getByText(E2E_INTERNAL_USER_KEY_ALIAS).first())
.toBeVisible({ timeout: 10_000 });
});
test("Team admin can add a member to their team", async ({ page }) => {
await navigateToPage(page, Page.Teams);
await dismissFeedbackPopup(page);
await clickTeamId(page, E2E_TEAM_CRUD_ID);
await page.getByRole("tab", { name: "Members" }).click();
await page.getByRole("button", { name: /Add Member/i }).click();
const modal = page.locator(".ant-modal:visible");
await expect(modal).toBeVisible({ timeout: 5_000 });
// Use a dedicated invitee user so this doesn't race with the proxy-admin
// "Invite a user" test that adds invitable@test.local to the same team.
await modal.locator(".ant-select").first().click();
await page.keyboard.type("invitable-team@test.local");
const emailOption = page.getByRole("option", { name: "invitable-team@test.local" }).first();
await expect(emailOption).toBeAttached({ timeout: 10_000 });
await page.keyboard.press("Enter");
await modal.getByRole("button", { name: /Add Member/i }).click();
await expect(page.getByText("Team member added successfully").first())
.toBeVisible({ timeout: 10_000 });
});
test("Team admin can remove a member from their team", async ({ page }) => {
await navigateToPage(page, Page.Teams);
await dismissFeedbackPopup(page);
await clickTeamId(page, E2E_TEAM_CRUD_ID);
await page.getByRole("tab", { name: "Members" }).click();
// Seeded members appear in the roster by user_id (members_with_roles has no
// email), so match the row on the user_id rather than the email.
const row = page.locator("tr", { hasText: "e2e-removable-member" }).first();
await expect(row).toBeVisible({ timeout: 10_000 });
await row.getByTestId("delete-member").click();
const modal = page.locator(".ant-modal:visible");
await expect(modal).toBeVisible({ timeout: 5_000 });
await modal.getByRole("button", { name: /^Delete$/ }).click();
await expect(page.getByText("Team member removed successfully").first())
.toBeVisible({ timeout: 10_000 });
});
test("Team admin can create a team key with All Team Models", async ({ page }) => {
await navigateToPage(page, Page.ApiKeys);
await dismissFeedbackPopup(page);
await page.getByRole("button", { name: /Create New Key/i }).click();
await expect(page.getByText("Key Ownership")).toBeVisible({ timeout: 10_000 });
const keyName = `e2e-team-admin-key-${Date.now()}`;
await page.getByTestId("base-input").fill(keyName);
// Team selector — same locator pattern as the proxy-admin keys test.
const teamSelect = page.locator(".ant-select", { hasText: "Search or select a team" });
await teamSelect.click();
await page.keyboard.type(E2E_TEAM_CRUD_ALIAS);
await page.locator(".ant-select-dropdown:visible").getByText(E2E_TEAM_CRUD_ALIAS).first().click();
// Models — pick "All Team Models"
await page.locator(".ant-select-selection-overflow").click();
await page.locator(".ant-select-dropdown:visible").getByText("All Team Models").click();
await page.keyboard.press("Escape");
await page.getByRole("button", { name: "Create Key", exact: true }).click();
await expect(page.getByText("Save your Key")).toBeVisible({ timeout: 10_000 });
await page.keyboard.press("Escape");
await expect(page.getByText(keyName)).toBeVisible({ timeout: 10_000 });
});
});

View file

@ -1018,3 +1018,83 @@ describe("OldTeams - organization alias display", () => {
});
});
});
describe("OldTeams - Resources column keys badge", () => {
beforeEach(() => {
vi.clearAllMocks();
mockUseOrganizations.mockReturnValue({ data: [] });
});
it("renders keys_count from the v2 payload in the Resources badge", async () => {
const { container } = renderWithQueryClient(
<OldTeams
teams={[
{
team_id: "1",
team_alias: "Team With Keys",
organization_id: "org-123",
models: ["gpt-4"],
max_budget: 100,
budget_duration: "1d",
tpm_limit: 1000,
rpm_limit: 1000,
created_at: new Date().toISOString(),
keys: [],
keys_count: 3,
members_with_roles: [],
spend: 0,
} as any,
]}
searchParams={{}}
accessToken="test-token"
setTeams={vi.fn()}
userID="user-123"
userRole="Admin"
organizations={[]}
/>,
);
await waitFor(() => {
expect(screen.getByText("Team With Keys")).toBeInTheDocument();
});
const cyanTag = container.querySelector(".ant-tag-cyan");
expect(cyanTag).not.toBeNull();
expect(cyanTag?.textContent).toContain("3");
});
it("falls back to keys.length when keys_count is absent", async () => {
const { container } = renderWithQueryClient(
<OldTeams
teams={[
{
team_id: "2",
team_alias: "Legacy Team",
organization_id: "org-123",
models: ["gpt-4"],
max_budget: 100,
budget_duration: "1d",
tpm_limit: 1000,
rpm_limit: 1000,
created_at: new Date().toISOString(),
keys: [{ token: "t1" } as any, { token: "t2" } as any],
members_with_roles: [],
spend: 0,
} as any,
]}
searchParams={{}}
accessToken="test-token"
setTeams={vi.fn()}
userID="user-123"
userRole="Admin"
organizations={[]}
/>,
);
await waitFor(() => {
expect(screen.getByText("Legacy Team")).toBeInTheDocument();
});
const cyanTag = container.querySelector(".ant-tag-cyan");
expect(cyanTag).not.toBeNull();
expect(cyanTag?.textContent).toContain("2");
});
});

View file

@ -105,6 +105,7 @@ interface TeamInfo {
interface PerTeamInfo {
keys: KeyResponse[];
keys_count: number;
team_info: TeamInfo;
}
@ -364,6 +365,7 @@ const Teams: React.FC<TeamProps> = ({
(acc, team) => {
acc[team.team_id] = {
keys: team.keys || [],
keys_count: team.keys_count ?? team.keys?.length ?? 0,
team_info: {
members_with_roles: team.members_with_roles || [],
},
@ -745,7 +747,7 @@ const Teams: React.FC<TeamProps> = ({
render: (_: unknown, record: Team) => {
const memberCount = perTeamInfo?.[record.team_id]?.team_info?.members_with_roles?.length ?? 0;
const modelCount = record.models?.length ?? 0;
const keyCount = perTeamInfo?.[record.team_id]?.keys?.length ?? 0;
const keyCount = perTeamInfo?.[record.team_id]?.keys_count ?? 0;
return (
<Flex gap={12} align="center">
<Tooltip title={`${memberCount} Members`}>
@ -977,17 +979,23 @@ const Teams: React.FC<TeamProps> = ({
<DeleteResourceModal
isOpen={isDeleteModalOpen}
title="Delete Team?"
alertMessage={
teamToDelete?.keys?.length === 0
alertMessage={(() => {
const deleteKeyCount =
teamToDelete?.keys_count ?? teamToDelete?.keys?.length ?? 0;
return deleteKeyCount === 0
? undefined
: `Warning: This team has ${teamToDelete?.keys?.length} keys associated with it. Deleting the team will also delete all associated keys. This action is irreversible.`
}
: `Warning: This team has ${deleteKeyCount} keys associated with it. Deleting the team will also delete all associated keys. This action is irreversible.`;
})()}
message="Are you sure you want to delete this team and all its keys? This action cannot be undone."
resourceInformationTitle="Team Information"
resourceInformation={[
{ label: "Team ID", value: teamToDelete?.team_id, code: true },
{ label: "Team Name", value: teamToDelete?.team_alias },
{ label: "Keys", value: teamToDelete?.keys?.length },
{
label: "Keys",
value:
teamToDelete?.keys_count ?? teamToDelete?.keys?.length ?? 0,
},
{ label: "Members", value: teamToDelete?.members_with_roles?.length },
]}
requiredConfirmation={teamToDelete?.team_alias}

View file

@ -13,6 +13,7 @@ export interface Team {
organization_id: string;
created_at: string;
keys: KeyResponse[];
keys_count?: number;
members_with_roles: Member[];
spend: number;
access_group_ids?: string[];

View file

@ -439,6 +439,12 @@ const CreateKey: React.FC<CreateKeyProps> = ({ team, teams, data, addKey, autoOp
// Update the formValues with the final metadata
formValues.metadata = JSON.stringify(metadata);
// disable_global_guardrails is premium-gated server-side; only send it when enabled
// so non-premium key creation isn't blocked by that gate.
if (!formValues.disable_global_guardrails) {
delete formValues.disable_global_guardrails;
}
// Transform allowed_vector_store_ids and allowed_mcp_servers_and_groups into object_permission format
if (formValues.allowed_vector_store_ids && formValues.allowed_vector_store_ids.length > 0) {
formValues.object_permission = {

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