mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
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:
commit
4cda2c5a19
102 changed files with 6282 additions and 1750 deletions
4
.github/pull_request_template.md
vendored
4
.github/pull_request_template.md
vendored
|
|
@ -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
307
AGENTS.md
|
|
@ -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
214
CLAUDE.md
|
|
@ -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
109
GEMINI.md
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 -}}
|
||||
|
||||
|
|
|
|||
|
|
@ -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 }}
|
||||
|
|
|
|||
|
|
@ -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 }}
|
||||
|
|
|
|||
|
|
@ -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 }}
|
||||
|
|
|
|||
|
|
@ -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 }}
|
||||
|
|
|
|||
|
|
@ -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 }}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
||||
|
|
|
|||
|
|
@ -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
1
litellm/proxy/_experimental/mcp_server/CLAUDE.md
Normal file
1
litellm/proxy/_experimental/mcp_server/CLAUDE.md
Normal 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
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -839,6 +839,7 @@ class ProxyBaseLLMRequestProcessing:
|
|||
"aget_run",
|
||||
"acancel_run",
|
||||
"adelete_run",
|
||||
"apply_guardrail",
|
||||
],
|
||||
version: Optional[str] = None,
|
||||
user_model: Optional[str] = None,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)}"
|
||||
)
|
||||
},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
122
pyproject.toml
122
pyproject.toml
|
|
@ -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",
|
||||
]
|
||||
|
|
|
|||
191
scripts/benchmark_model_response_creator.py
Normal file
191
scripts/benchmark_model_response_creator.py
Normal 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()
|
||||
369
scripts/benchmark_streaming_chunk_overhead.py
Normal file
369
scripts/benchmark_streaming_chunk_overhead.py
Normal 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()
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
|
|
|||
508
tests/test_litellm/litellm_core_utils/test_streaming_overhead.py
Normal file
508
tests/test_litellm/litellm_core_utils/test_streaming_overhead.py
Normal 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"
|
||||
)
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
|
|
|
|||
|
|
@ -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"}]
|
||||
|
|
|
|||
|
|
@ -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", ""],
|
||||
|
|
|
|||
|
|
@ -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", [])
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"):
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
184
tests/test_litellm/test_claude_opus_4_8_config.py
Normal file
184
tests/test_litellm/test_claude_opus_4_8_config.py
Normal 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
|
||||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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": {
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
1
ui/litellm-dashboard/CLAUDE.md
Normal file
1
ui/litellm-dashboard/CLAUDE.md
Normal 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
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 });
|
||||
}
|
||||
|
|
|
|||
30
ui/litellm-dashboard/e2e_tests/tests/auth/logout.spec.ts
Normal file
30
ui/litellm-dashboard/e2e_tests/tests/auth/logout.spec.ts
Normal 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 });
|
||||
});
|
||||
});
|
||||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
|
|
@ -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 });
|
||||
|
|
|
|||
|
|
@ -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 });
|
||||
});
|
||||
});
|
||||
|
|
@ -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");
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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[];
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue