merge litellm_internal_staging into litellm_fireworks_e2e_test
Some checks failed
LiteLLM Rust / rustfmt, clippy, test (push) Has been cancelled

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Krrish Dholakia 2026-07-20 23:22:47 +00:00
commit 280090caa1
172 changed files with 9127 additions and 1052 deletions

View file

@ -1,61 +0,0 @@
name: Create Daily OSS Branch
on:
schedule:
- cron: "0 16 * * 1-5" # 9am PT during daylight saving time, weekdays.
workflow_dispatch:
inputs:
date:
description: "Branch date in YYYY_MM_DD format. Defaults to today's UTC date."
required: false
type: string
permissions:
contents: write
jobs:
create-oss-branch:
if: github.repository == 'BerriAI/litellm'
runs-on: ubuntu-latest
timeout-minutes: 10
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
fetch-depth: 0
persist-credentials: false
- name: Create dated OSS branch
env:
GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }}
REQUESTED_DATE: ${{ inputs.date }}
run: |
set -euo pipefail
if [ -n "${REQUESTED_DATE}" ]; then
if ! echo "${REQUESTED_DATE}" | grep -Eq '^[0-9]{4}_[0-9]{2}_[0-9]{2}$'; then
echo "::error::date must use YYYY_MM_DD format, got '${REQUESTED_DATE}'"
exit 1
fi
BRANCH_DATE="${REQUESTED_DATE}"
else
BRANCH_DATE="$(date -u +'%Y_%m_%d')"
fi
BRANCH_NAME="litellm_oss_daily_${BRANCH_DATE}"
echo "Creating branch: ${BRANCH_NAME}"
git config user.name "github-actions[bot]"
git config user.email "github-actions[bot]@users.noreply.github.com"
git fetch origin main "${BRANCH_NAME}" || true
if git show-ref --verify --quiet "refs/remotes/origin/${BRANCH_NAME}"; then
echo "Branch ${BRANCH_NAME} already exists. Skipping creation."
exit 0
fi
git checkout -b "${BRANCH_NAME}" origin/main
git push "https://x-access-token:${GITHUB_TOKEN}@github.com/${GITHUB_REPOSITORY}.git" "${BRANCH_NAME}"
echo "Successfully created and pushed branch: ${BRANCH_NAME}"

View file

@ -31,12 +31,12 @@ jobs:
echo "PR head repo: $HEAD_REPO"
echo "PR head branch: $HEAD_REF"
if [ "$HEAD_REPO" != "$BASE_REPO" ]; then
echo "::error::PRs to main must originate from the canonical repository ($BASE_REPO), not a fork ($HEAD_REPO). External contributors should open PRs against the current daily OSS branch (named litellm_oss_daily_YYYY_MM_DD; a fresh one is cut each weekday, so target the most recent) instead."
echo "::error::PRs to main must originate from the canonical repository ($BASE_REPO), not a fork ($HEAD_REPO). External contributors should open PRs against 'litellm_internal_staging' instead."
exit 1
fi
if [ "$HEAD_REF" = "litellm_internal_staging" ] || [[ "$HEAD_REF" == litellm_hotfix_?* ]]; then
echo "Allowed source branch."
exit 0
fi
echo "::error::PRs to main must originate from 'litellm_internal_staging' or a 'litellm_hotfix_*' branch. Got: '$HEAD_REF'. If this is a contribution, retarget the PR against the current daily OSS branch (named litellm_oss_daily_YYYY_MM_DD; a fresh one is cut each weekday, so target the most recent) instead."
echo "::error::PRs to main must originate from 'litellm_internal_staging' or a 'litellm_hotfix_*' branch. Got: '$HEAD_REF'. If this is a contribution, retarget the PR against 'litellm_internal_staging' instead."
exit 1

View file

@ -1,50 +0,0 @@
name: OSS Daily Guardrails
on:
push:
branches:
- "litellm_oss_daily_20*"
pull_request:
branches:
- "litellm_oss_daily_20*"
- litellm_internal_staging
permissions:
contents: read
concurrency:
group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }}
cancel-in-progress: true
jobs:
oss-safe-checks:
name: Run OSS daily safe checks
if: startsWith(github.ref_name, 'litellm_oss_daily_20') || startsWith(github.head_ref, 'litellm_oss_daily_20') || startsWith(github.base_ref, 'litellm_oss_daily_20')
runs-on: ubuntu-latest
timeout-minutes: 10
steps:
- name: Checkout repository
uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
with:
python-version: "3.12"
- name: Set up uv
uses: ./.github/actions/setup-uv-with-retries
with:
version: "0.10.9"
- name: Run secret scan test
run: |
uv run --frozen --with 'pytest==9.0.2' pytest tests/litellm/test_no_hardcoded_secrets.py -v
- name: Run Ruff
run: |
uv sync --frozen
cd litellm
uv run --no-sync ruff check .

View file

@ -19,7 +19,7 @@ Same thing for bug fixes. The tests should make it so that this specific bug can
End-to-end tests belong in `tests/e2e/` and must follow the harness conventions documented in that directory's `CLAUDE.md`
When creating PRs, don't set base to `main`. `litellm_internal_staging` serves that purpose for internal contributors; external / OSS contributions target the current daily OSS branch instead, named `litellm_oss_daily_YYYY_MM_DD` (a fresh one is cut each weekday, so use the most recent)
When creating PRs, don't set base to `main`. `litellm_internal_staging` is the default base branch and serves that purpose for both internal and external / OSS contributions
When writing a PR body, treat the comments and imperative instructions inside @.github/pull_request_template.md as rules to follow, not just layout. Agent harnesses may strip HTML comments from copies of that file injected into context, so read .github/pull_request_template.md from disk before writing a PR body to make sure you see every comment rule

View file

@ -322,7 +322,7 @@ npm run build
## Submitting Your PR
1. **Push your branch**: `git push origin your-feature-branch`
2. **Create a PR**: Go to GitHub and open a pull request against the current daily OSS branch, named `litellm_oss_daily_YYYY_MM_DD`. A fresh one is cut each weekday, so pick the most recent from the [branch list](https://github.com/BerriAI/litellm/branches/all?query=litellm_oss_daily). Do not target `main`.
2. **Create a PR**: Go to GitHub and open a pull request against [`litellm_internal_staging`](https://github.com/BerriAI/litellm/tree/litellm_internal_staging), which is the default base branch. Do not target `main`.
3. **Fill out the PR template**: Provide clear description of changes
4. **Wait for review**: Maintainers will review and provide feedback
5. **Address feedback**: Make requested changes and push updates

View file

@ -62,6 +62,13 @@ Not allowed in `core`:
Python owns rollout state and fallback while Rust is being introduced. Rust
paths must be off by default until parity tests prove equivalence with Python.
A new provider/route may instead be implemented rust-only with no Python
reference; then the Python interface is a thin dispatch that calls Rust with no
fallback, and you state the rust-only choice explicitly in the PR. Either way
the Python side stays minimal (it only marshals inputs and calls the Rust
interface), never add a per-route feature flag, and never push provider
dispatch into `litellm/main.py`; put it in a thin dispatch class under
`litellm/llms/<provider>/<route>/`.
## Production Bar

454
litellm-rust/Cargo.lock generated

File diff suppressed because it is too large Load diff

View file

@ -7,7 +7,8 @@ members = [
resolver = "2"
[workspace.package]
edition = "2021"
edition = "2024"
rust-version = "1.88"
license = "MIT"
repository = "https://github.com/BerriAI/litellm"

View file

@ -39,11 +39,17 @@ Rules for adding or changing an LLM provider/route in `litellm-rust`. OCR (`MIST
19. Every provider transform ships tests for: supported-param filtering, request body shape, response normalization, missing/null fields, bad input, and `*_match_python` fixture parity.
20. Lifecycle/hook tests cover hook order, success + failure callback payloads, pre-call guardrail blocking before any provider I/O, during-call body mutation, and provider-error mapping.
21. Rust paths stay off by default and behind Python parity tests (disabled / enabled-equals-Python / bridge-unavailable fallback) until parity is proven.
21. When a route has a Python reference implementation, the Rust path stays off by default and behind Python parity tests (disabled / enabled-equals-Python / bridge-unavailable fallback) until parity is proven. A new provider/route may instead be implemented rust-only with no Python reference; then the Python interface is a thin dispatch to Rust with no fallback, and tests cover the rust-backed path plus the unavailable-bridge error. State the rust-only choice explicitly in the PR.
## Python bridge (SDK side)
22. A Python -> Rust bridge keeps the Python side minimal: the Python interface only marshals inputs and calls the Rust interface, with no transform, handler, or business logic. Aim for well under 100 lines of interface code per route; if the Python grows past that, the logic belongs in Rust.
23. Do not bloat `litellm/main.py`. A route's provider dispatch lives in a thin dispatch class under `litellm/llms/<provider>/<route>/` that calls the Rust bridge; `main.py` only instantiates it and calls its sync/async method.
24. Do not add new feature flags unless explicitly requested. Reuse the existing litellm rust rollout mechanism (`use_litellm_rust`); never introduce a per-route env flag such as `LITELLM_USE_RUST_<ROUTE>`.
## Checks before push
22. Run, and keep green:
25. Run, and keep green:
```bash
cd litellm-rust
cargo fmt --check

View file

@ -14,7 +14,7 @@ path = "src/main.rs"
required-features = ["server"]
[dependencies]
litellm-core.workspace = true
litellm-core = { workspace = true, features = ["bedrock-auth"] }
# reqwest (rustls + json) is used by io/ocr and ships realtime logs to the
# Python proxy callbacks API.
reqwest.workspace = true

View file

@ -0,0 +1,48 @@
use std::collections::BTreeMap;
use litellm_core::CoreResult;
use litellm_core::audio_transcription::transformation::AudioTranscriptionProviderConfig;
use litellm_core::error::CoreError;
use litellm_core::providers::bedrock::audio_transcription::BEDROCK_AUDIO_TRANSCRIPTION_CONFIG;
use serde_json::{Map, Value};
pub(super) fn audio_transcription_provider_config(
provider: &str,
) -> Option<&'static dyn AudioTranscriptionProviderConfig> {
match provider {
"bedrock" => Some(&BEDROCK_AUDIO_TRANSCRIPTION_CONFIG),
_ => None,
}
}
pub(super) fn string_headers(
headers: Option<Map<String, Value>>,
) -> CoreResult<BTreeMap<String, String>> {
headers
.unwrap_or_default()
.into_iter()
.map(|(key, value)| {
value
.as_str()
.map(|value| (key.clone(), value.to_string()))
.ok_or_else(|| {
CoreError::InvalidRequest(format!(
"audio transcription extra_headers.{key} must be a string"
))
})
})
.collect()
}
pub(super) fn has_header(headers: &BTreeMap<String, String>, name: &str) -> bool {
headers.keys().any(|key| key.eq_ignore_ascii_case(name))
}
pub(super) fn truncate_error_body(body: &str) -> String {
let truncated: String = body.chars().take(256).collect();
if truncated.chars().count() == body.chars().count() {
truncated
} else {
format!("{truncated}... (truncated)")
}
}

View file

@ -0,0 +1,89 @@
use std::time::SystemTime;
use litellm_core::CoreResult;
use litellm_core::audio_transcription::transformation::AudioTranscriptionAuth;
use litellm_core::error::CoreError;
use litellm_core::providers::bedrock::audio_transcription::aws_auth_config;
use litellm_core::providers::bedrock::aws_base::{resolve_credentials, sign_bedrock_post};
use serde_json::Value;
use super::common_utils::truncate_error_body;
use super::types::ProviderAudioTranscriptionRequest;
use crate::client::http_client;
pub(crate) async fn execute_audio_transcription_provider_call(
request: ProviderAudioTranscriptionRequest,
) -> CoreResult<Value> {
let body = serde_json::to_vec(&request.body).map_err(|error| {
CoreError::InvalidRequest(format!("invalid audio request body: {error}"))
})?;
let mut request_builder = http_client().post(&request.url).body(body.clone());
for (key, value) in &request.upstream_headers {
request_builder = request_builder.header(key, value);
}
if let Some(duration) = request.timeout {
request_builder = request_builder.timeout(duration);
}
let response = request_builder
.send()
.await
.map_err(|error| CoreError::Network(error.to_string()))?;
let status = response.status();
let text = response
.text()
.await
.map_err(|error| CoreError::Network(error.to_string()))?;
if !status.is_success() {
return Err(CoreError::Http {
status: status.as_u16(),
body: truncate_error_body(&text),
});
}
let response_json: Value = serde_json::from_str(&text).map_err(|error| {
CoreError::InvalidResponse(format!("invalid audio response JSON: {error}"))
})?;
Ok(request
.config
.transform_transcription_response(&request.model, response_json)?
.into_json())
}
pub(crate) async fn sign_request(
request: &ProviderAudioTranscriptionRequest,
optional_params: &serde_json::Map<String, Value>,
) -> CoreResult<ProviderAudioTranscriptionRequest> {
let env_lookup = environment_lookup;
let auth = request
.config
.auth_strategy(&request.model, optional_params, &env_lookup)?;
let body = serde_json::to_vec(&request.body).map_err(|error| {
CoreError::InvalidRequest(format!("invalid audio request body: {error}"))
})?;
let mut headers = super::common_utils::string_headers(None)?;
headers.insert("Content-Type".to_string(), "application/json".to_string());
headers.extend(request.upstream_headers.iter().cloned());
match auth {
AudioTranscriptionAuth::Bearer => {}
AudioTranscriptionAuth::AwsSigV4 { region, .. } => {
let credentials =
resolve_credentials(aws_auth_config(optional_params, &env_lookup), &env_lookup)
.await?;
headers.extend(sign_bedrock_post(
&request.url,
&body,
&headers,
&region,
&credentials,
SystemTime::now(),
)?);
}
}
Ok(ProviderAudioTranscriptionRequest {
upstream_headers: headers.into_iter().collect(),
..request.clone()
})
}
pub(super) fn environment_lookup(key: &str) -> Option<String> {
std::env::var(key).ok()
}

View file

@ -0,0 +1,300 @@
use std::future::Future;
use std::pin::Pin;
use litellm_core::CoreResult;
use litellm_core::audio_transcription::transformation::AudioTranscriptionAuth;
use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming};
use litellm_core::error::CoreError;
use serde_json::{Map, Value, json};
use super::common_utils::{audio_transcription_provider_config, has_header, string_headers};
use super::handler::sign_request;
use super::types::{PreparedAudioTranscriptionRequest, ProviderAudioTranscriptionRequest};
use crate::integrations::custom_guardrail::{
CustomGuardrailRunner, GuardrailContext, GuardrailError, GuardrailRequest,
};
use crate::integrations::custom_logger::{
CallType, CallbackTiming, CallbackValue, CustomLoggerRunner, LoggingError, ModelCallDetails,
};
use crate::integrations::types::{
RequestMetadata, StandardLoggingMetadata, StandardLoggingPayload,
};
pub(crate) struct AudioTranscriptionLifecycleHooks {
logger_runner: CustomLoggerRunner,
guardrail_runner: CustomGuardrailRunner,
request_metadata: RequestMetadata,
}
type AudioFuture<'a, T> = Pin<Box<dyn Future<Output = CoreResult<T>> + Send + 'a>>;
type AudioLogFuture<'a> = Pin<Box<dyn Future<Output = ()> + Send + 'a>>;
impl AudioTranscriptionLifecycleHooks {
pub(crate) fn new(
logger_runner: CustomLoggerRunner,
guardrail_runner: CustomGuardrailRunner,
request_metadata: RequestMetadata,
) -> Self {
Self {
logger_runner,
guardrail_runner,
request_metadata,
}
}
async fn run_pre_call_guardrails(
&self,
request: PreparedAudioTranscriptionRequest,
) -> CoreResult<PreparedAudioTranscriptionRequest> {
if self.guardrail_runner.is_empty() {
return Ok(request);
}
let (guardrail_request, _) = self
.guardrail_runner
.run_pre_call(
&guardrail_context(&self.request_metadata),
GuardrailRequest::new(json!({
"model": request.model,
"custom_llm_provider": request.custom_llm_provider,
"audio": request.audio,
"optional_params": request.optional_params,
})),
)
.await
.map_err(guardrail_error_to_core_error)?;
let Value::Object(mut data) = guardrail_request.data else {
return Err(CoreError::InvalidRequest(
"audio transcription pre_call guardrail must return an object".to_string(),
));
};
let audio = data.remove("audio").ok_or_else(|| {
CoreError::InvalidRequest("audio transcription guardrail removed audio".to_string())
})?;
let optional_params = match data.remove("optional_params") {
Some(Value::Object(value)) => value,
Some(_) => {
return Err(CoreError::InvalidRequest(
"audio transcription optional_params must be an object".to_string(),
));
}
None => Map::new(),
};
Ok(PreparedAudioTranscriptionRequest {
audio,
optional_params,
..request
})
}
async fn prepare_provider_request(
&self,
request: PreparedAudioTranscriptionRequest,
) -> CoreResult<ProviderAudioTranscriptionRequest> {
let config = audio_transcription_provider_config(&request.custom_llm_provider)
.ok_or_else(|| CoreError::InvalidProvider(request.custom_llm_provider.clone()))?;
let env_lookup = super::handler::environment_lookup;
let headers = string_headers(request.extra_headers)?;
let url = config.complete_url(
request.api_base.as_deref(),
&request.model,
&request.optional_params,
&env_lookup,
)?;
let filtered_params = config.map_transcription_params(&request.optional_params);
let body = config.transform_transcription_request(
&request.model,
request.audio,
filtered_params,
)?;
let auth = config.auth_strategy(&request.model, &request.optional_params, &env_lookup)?;
let mut upstream_headers = headers.into_iter().collect::<Vec<_>>();
if matches!(auth, AudioTranscriptionAuth::Bearer)
&& !has_header(
&upstream_headers
.iter()
.cloned()
.collect::<std::collections::BTreeMap<_, _>>(),
"authorization",
)
&& let Some(api_key) = request.api_key.as_deref()
{
upstream_headers.push(("Authorization".to_string(), format!("Bearer {api_key}")));
}
let provider_request = ProviderAudioTranscriptionRequest {
model: request.model,
config,
url,
body: body.body,
upstream_headers,
timeout: request.timeout,
};
let provider_request = self.run_during_call_guardrails(provider_request).await?;
sign_request(&provider_request, &request.optional_params).await
}
async fn run_during_call_guardrails(
&self,
request: ProviderAudioTranscriptionRequest,
) -> CoreResult<ProviderAudioTranscriptionRequest> {
if self.guardrail_runner.is_empty() {
return Ok(request);
}
let (guardrail_request, _) = self
.guardrail_runner
.run_during_call(
&guardrail_context(&self.request_metadata),
GuardrailRequest::new(json!({
"model": request.model,
"custom_llm_provider": "bedrock",
"url": request.url,
"body": request.body,
})),
)
.await
.map_err(guardrail_error_to_core_error)?;
let Value::Object(mut data) = guardrail_request.data else {
return Err(CoreError::InvalidRequest(
"audio transcription during_call guardrail must return an object".to_string(),
));
};
let body = data.remove("body").ok_or_else(|| {
CoreError::InvalidRequest("audio transcription guardrail removed body".to_string())
})?;
Ok(ProviderAudioTranscriptionRequest { body, ..request })
}
fn logging_payload(
&self,
context: &CallLifecycleContext,
timing: &CallLifecycleTiming,
) -> StandardLoggingPayload {
StandardLoggingPayload {
id: context.litellm_call_id.clone(),
litellm_call_id: context.litellm_call_id.clone(),
call_type: context.call_type.clone(),
model: context.model.clone(),
custom_llm_provider: context.custom_llm_provider.clone(),
response_cost: 0.0,
prompt_tokens: 0,
completion_tokens: 0,
total_tokens: 0,
start_time: timing.start_time,
end_time: timing.end_time,
stream: false,
metadata: StandardLoggingMetadata {
user_api_key_hash: self.request_metadata.user_api_key_hash.clone(),
user_api_key_user_id: self.request_metadata.user_api_key_user_id.clone(),
user_api_key_team_id: self.request_metadata.user_api_key_team_id.clone(),
..Default::default()
},
messages: None,
}
}
}
impl CallLifecycleHooks<PreparedAudioTranscriptionRequest, ProviderAudioTranscriptionRequest, Value>
for AudioTranscriptionLifecycleHooks
{
type PreCallFuture<'a> = AudioFuture<'a, PreparedAudioTranscriptionRequest>;
type DuringCallFuture<'a> = AudioFuture<'a, ProviderAudioTranscriptionRequest>;
type SuccessFuture<'a> = AudioLogFuture<'a>;
type FailureFuture<'a> = AudioLogFuture<'a>;
fn async_pre_call_hook<'a>(
&'a self,
_context: &'a CallLifecycleContext,
request: PreparedAudioTranscriptionRequest,
) -> Self::PreCallFuture<'a> {
Box::pin(async move { self.run_pre_call_guardrails(request).await })
}
fn async_during_call_hook<'a>(
&'a self,
_context: &'a CallLifecycleContext,
request: PreparedAudioTranscriptionRequest,
) -> Self::DuringCallFuture<'a> {
Box::pin(async move { self.prepare_provider_request(request).await })
}
fn async_log_success_event<'a>(
&'a self,
context: &'a CallLifecycleContext,
response: &'a Value,
timing: &'a CallLifecycleTiming,
) -> Self::SuccessFuture<'a> {
Box::pin(async move {
if self.logger_runner.is_empty() {
return;
}
self.logger_runner
.async_log_success_event(
&ModelCallDetails::from_standard_logging_payload(
self.logging_payload(context, timing),
),
&CallbackValue::new("audio_transcription", response.clone()),
CallbackTiming::new(timing.start_time, timing.end_time),
)
.await;
})
}
fn async_log_failure_event<'a>(
&'a self,
context: &'a CallLifecycleContext,
error: &'a CoreError,
timing: &'a CallLifecycleTiming,
) -> Self::FailureFuture<'a> {
Box::pin(async move {
if self.logger_runner.is_empty() {
return;
}
let logging_error = LoggingError {
message: error.to_string(),
kind: core_error_kind(error).to_string(),
};
self.logger_runner
.async_log_failure_event(
&ModelCallDetails::from_standard_logging_payload(
self.logging_payload(context, timing),
)
.with_failure_error(logging_error.clone()),
Some(&CallbackValue::new(
"error",
json!({"message": logging_error.message, "kind": logging_error.kind}),
)),
CallbackTiming::new(timing.start_time, timing.end_time),
)
.await;
})
}
}
fn guardrail_context(metadata: &RequestMetadata) -> GuardrailContext {
GuardrailContext {
call_type: CallType::Other("audio_transcription".to_string()),
selected_guardrails: Vec::new(),
metadata: std::collections::HashMap::new(),
user_api_key_hash: metadata.user_api_key_hash.clone(),
user_api_key_user_id: metadata.user_api_key_user_id.clone(),
user_api_key_team_id: metadata.user_api_key_team_id.clone(),
trace_parent: None,
}
}
fn guardrail_error_to_core_error(error: GuardrailError) -> CoreError {
CoreError::InvalidRequest(format!("{}: {}", error.kind, error.message))
}
fn core_error_kind(error: &CoreError) -> &'static str {
match error {
CoreError::Auth(_) => "AuthError",
CoreError::InvalidProvider(_) => "InvalidProvider",
CoreError::InvalidRequest(_) => "InvalidRequest",
CoreError::InvalidType { .. } => "InvalidType",
CoreError::MissingField(_) => "MissingField",
CoreError::Http { .. } => "HttpError",
CoreError::InvalidResponse(_) => "InvalidResponse",
CoreError::Network(_) => "NetworkError",
CoreError::Routing(_) => "RoutingError",
}
}

View file

@ -0,0 +1,25 @@
use litellm_core::CoreResult;
use litellm_core::call_lifecycle::CallLifecycle;
use serde_json::Value;
mod common_utils;
mod handler;
mod hooks;
mod prepare;
mod types;
pub use types::AudioTranscriptionRequest;
use handler::execute_audio_transcription_provider_call;
use prepare::{PreparedAudioTranscriptionCall, prepare_audio_transcription_call};
pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> CoreResult<Value> {
let PreparedAudioTranscriptionCall { request, hooks } =
prepare_audio_transcription_call(request);
CallLifecycle::default()
.run_request(request, &hooks, execute_audio_transcription_provider_call)
.await
}
#[cfg(test)]
mod tests;

View file

@ -0,0 +1,55 @@
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use super::hooks::AudioTranscriptionLifecycleHooks;
use super::types::{AudioTranscriptionRequest, PreparedAudioTranscriptionRequest};
use crate::integrations::custom_guardrail::CustomGuardrailRunner;
use crate::integrations::custom_logger::CustomLoggerRunner;
pub(crate) struct PreparedAudioTranscriptionCall {
pub(crate) request: PreparedAudioTranscriptionRequest,
pub(crate) hooks: AudioTranscriptionLifecycleHooks,
}
pub(crate) fn prepare_audio_transcription_call(
request: AudioTranscriptionRequest<'_>,
) -> PreparedAudioTranscriptionCall {
let call_id = request
.litellm_call_id
.map(str::to_string)
.unwrap_or_else(new_audio_transcription_call_id);
let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider)
.unwrap_or(CustomLlmProvider {
model: request.model,
custom_llm_provider: "bedrock",
});
PreparedAudioTranscriptionCall {
request: PreparedAudioTranscriptionRequest {
model: provider_info.model.to_string(),
custom_llm_provider: provider_info.custom_llm_provider.to_string(),
litellm_call_id: call_id,
audio: request.audio,
api_key: request.api_key.map(str::to_string),
api_base: request.api_base.map(str::to_string),
extra_headers: request.extra_headers,
optional_params: request.optional_params,
timeout: request.timeout,
},
hooks: AudioTranscriptionLifecycleHooks::new(
CustomLoggerRunner::new(request.callbacks),
CustomGuardrailRunner::new(request.guardrails),
request.request_metadata,
),
}
}
fn new_audio_transcription_call_id() -> String {
static COUNTER: AtomicU64 = AtomicU64::new(1);
let sequence = COUNTER.fetch_add(1, Ordering::Relaxed);
let timestamp = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map_or(0, |duration| duration.as_nanos());
format!("audio-transcription-{timestamp}-{sequence}")
}

View file

@ -0,0 +1,53 @@
use std::io::{Read, Write};
use std::net::TcpListener;
use std::thread;
use serde_json::{Map, json};
use super::{AudioTranscriptionRequest, audio_transcription};
#[tokio::test]
async fn bedrock_request_is_signed_and_contains_audio() {
let listener = TcpListener::bind("127.0.0.1:0").expect("listener");
let address = listener.local_addr().expect("address");
let server = thread::spawn(move || {
let (mut stream, _) = listener.accept().expect("connection");
let mut request = Vec::new();
let mut buffer = [0_u8; 16_384];
let count = stream.read(&mut buffer).expect("request");
request.extend_from_slice(&buffer[..count]);
let request = String::from_utf8_lossy(&request);
assert!(request.contains("POST /model/mistral.voxtral-mini-3b-2507/converse"));
assert!(request.contains("authorization: AWS4-HMAC-SHA256"));
assert!(request.contains("x-amz-date:"));
assert!(request.contains("\"bytes\":\"AQI=\""));
assert!(request.contains("Transcribe the audio. Respond with only the transcript."));
let response = b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: 53\r\nConnection: close\r\n\r\n{\"output\":{\"message\":{\"content\":[{\"text\":\"hello\"}]}}}";
stream.write_all(response).expect("response");
});
let optional_params = Map::from_iter([
("aws_access_key_id".to_string(), json!("access-key")),
("aws_secret_access_key".to_string(), json!("secret-key")),
("aws_region_name".to_string(), json!("us-east-1")),
]);
let api_base = format!("http://{address}");
let response = audio_transcription(AudioTranscriptionRequest {
model: "mistral.voxtral-mini-3b-2507",
audio: json!({"data": "AQI=", "format": "wav", "filename": "audio.wav"}),
api_key: None,
api_base: Some(&api_base),
custom_llm_provider: Some("bedrock"),
extra_headers: None,
optional_params,
timeout: None,
callbacks: Vec::new(),
guardrails: Vec::new(),
request_metadata: Default::default(),
litellm_call_id: None,
})
.await
.expect("transcription");
assert_eq!(response, json!({"text": "hello"}));
server.join().expect("server");
}

View file

@ -0,0 +1,58 @@
use std::sync::Arc;
use std::time::Duration;
use litellm_core::audio_transcription::transformation::AudioTranscriptionProviderConfig;
use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleRequest};
use serde_json::{Map, Value};
use crate::integrations::custom_guardrail::CustomGuardrail;
use crate::integrations::custom_logger::CustomLogger;
use crate::integrations::types::RequestMetadata;
pub struct AudioTranscriptionRequest<'a> {
pub model: &'a str,
pub audio: Value,
pub api_key: Option<&'a str>,
pub api_base: Option<&'a str>,
pub custom_llm_provider: Option<&'a str>,
pub extra_headers: Option<Map<String, Value>>,
pub optional_params: Map<String, Value>,
pub timeout: Option<Duration>,
pub callbacks: Vec<Arc<dyn CustomLogger>>,
pub guardrails: Vec<Arc<dyn CustomGuardrail>>,
pub request_metadata: RequestMetadata,
pub litellm_call_id: Option<&'a str>,
}
pub(crate) struct PreparedAudioTranscriptionRequest {
pub(crate) model: String,
pub(crate) custom_llm_provider: String,
pub(crate) litellm_call_id: String,
pub(crate) audio: Value,
pub(crate) api_key: Option<String>,
pub(crate) api_base: Option<String>,
pub(crate) extra_headers: Option<Map<String, Value>>,
pub(crate) optional_params: Map<String, Value>,
pub(crate) timeout: Option<Duration>,
}
impl CallLifecycleRequest for PreparedAudioTranscriptionRequest {
fn lifecycle_context(&self) -> CallLifecycleContext {
CallLifecycleContext::new(
"audio_transcription",
self.model.clone(),
self.custom_llm_provider.clone(),
self.litellm_call_id.clone(),
)
}
}
#[derive(Clone)]
pub(crate) struct ProviderAudioTranscriptionRequest {
pub(crate) model: String,
pub(crate) config: &'static dyn AudioTranscriptionProviderConfig,
pub(crate) url: String,
pub(crate) body: Value,
pub(crate) upstream_headers: Vec<(String, String)>,
pub(crate) timeout: Option<Duration>,
}

View file

@ -9,9 +9,9 @@
//! runs during extraction, before the handler body. Routes never re-implement it.
use axum::extract::FromRequestParts;
use axum::http::StatusCode;
use axum::http::header::AUTHORIZATION;
use axum::http::request::Parts;
use axum::http::StatusCode;
use sha2::{Digest, Sha256};
use subtle::ConstantTimeEq;

View file

@ -1,13 +1,13 @@
use std::sync::OnceLock;
use std::time::Duration;
const OCR_TIMEOUT_SECS: u64 = 600;
const HTTP_CLIENT_TIMEOUT_SECS: u64 = 600;
pub(super) fn http_client() -> &'static reqwest::Client {
pub(crate) fn http_client() -> &'static reqwest::Client {
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
CLIENT.get_or_init(|| {
reqwest::Client::builder()
.timeout(Duration::from_secs(OCR_TIMEOUT_SECS))
.timeout(Duration::from_secs(HTTP_CLIENT_TIMEOUT_SECS))
.build()
.expect("failed to build reqwest client")
})

View file

@ -0,0 +1 @@
pub use crate::audio_transcription::{AudioTranscriptionRequest, audio_transcription};

View file

@ -1 +1 @@
pub use crate::messages::{messages, MessagesRequest};
pub use crate::messages::{MessagesRequest, messages};

View file

@ -1,3 +1,4 @@
pub mod audio_transcription;
pub mod messages;
pub mod ocr;
pub mod realtime;

View file

@ -1 +1 @@
pub use crate::ocr::{ocr, OcrRequest};
pub use crate::ocr::{OcrRequest, ocr};

View file

@ -15,16 +15,16 @@ use std::time::Duration;
use futures_util::stream::{SplitSink, SplitStream};
use futures_util::{Sink, SinkExt, Stream, StreamExt};
use litellm_core::CoreResult;
use litellm_core::error::CoreError;
use litellm_core::realtime::transformation::RealtimeProviderConfig;
use litellm_core::realtime::types::RealtimeEvent;
use litellm_core::CoreResult;
use tokio::net::TcpStream;
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tokio_tungstenite::tungstenite::http::header::AUTHORIZATION;
use tokio_tungstenite::tungstenite::http::HeaderValue;
use tokio_tungstenite::tungstenite::Message;
use tokio_tungstenite::{connect_async, MaybeTlsStream, WebSocketStream};
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tokio_tungstenite::tungstenite::http::HeaderValue;
use tokio_tungstenite::tungstenite::http::header::AUTHORIZATION;
use tokio_tungstenite::{MaybeTlsStream, WebSocketStream, connect_async};
use litellm_core::providers::openai::realtime::transformation::OPENAI_REALTIME_CONFIG;
@ -113,7 +113,7 @@ pub(crate) async fn read_event(upstream_rx: &mut UpstreamRx) -> CoreResult<Realt
Message::Close(_) => {
return Err(CoreError::Network(
"upstream closed before first event".to_string(),
))
));
}
_ => continue,
}

View file

@ -28,11 +28,11 @@ use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use futures_util::StreamExt;
use litellm_core::realtime::types::RealtimeEvent;
use litellm_core::CoreResult;
use litellm_core::realtime::types::RealtimeEvent;
use crate::io::realtime::{
dial_upstream, read_event, resolve_api_key, UpstreamRx, UpstreamTx, UpstreamWs,
UpstreamRx, UpstreamTx, UpstreamWs, dial_upstream, read_event, resolve_api_key,
};
/// Default target warm sockets per key when pooling is enabled.
@ -473,8 +473,8 @@ pub fn upstream_key(
/// unhealthy — we'd rather discard and fresh-dial than hand over a socket in an
/// unexpected state. `Pending` (the healthy case) returns `false`.
fn is_dead(rx: &mut UpstreamRx) -> bool {
use futures_util::task::noop_waker_ref;
use futures_util::Stream;
use futures_util::task::noop_waker_ref;
use std::pin::Pin;
use std::task::{Context, Poll};
@ -523,15 +523,15 @@ mod tests {
))
.await;
while let Some(Ok(msg)) = ws.next().await {
if let Message::Text(text) = msg {
if text.contains("response.create") {
for frame in [
r#"{"type":"response.created"}"#,
r#"{"type":"response.output_audio.delta","delta":"AAAA"}"#,
r#"{"type":"response.done"}"#,
] {
let _ = ws.send(Message::Text(frame.to_string())).await;
}
if let Message::Text(text) = msg
&& text.contains("response.create")
{
for frame in [
r#"{"type":"response.created"}"#,
r#"{"type":"response.output_audio.delta","delta":"AAAA"}"#,
r#"{"type":"response.done"}"#,
] {
let _ = ws.send(Message::Text(frame.to_string())).await;
}
}
}

View file

@ -10,19 +10,18 @@ use litellm_core::responses::websocket::ResponsesWebSocketProviderConfig;
use litellm_core::{CoreError, CoreResult};
use tokio::net::TcpStream;
use tokio::sync::Mutex;
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tokio_tungstenite::tungstenite::http::header::{HeaderName, AUTHORIZATION};
use tokio_tungstenite::tungstenite::http::HeaderValue;
use tokio_tungstenite::tungstenite::Message;
use tokio_tungstenite::{connect_async, MaybeTlsStream, WebSocketStream};
use tokio_tungstenite::tungstenite::client::IntoClientRequest;
use tokio_tungstenite::tungstenite::http::HeaderValue;
use tokio_tungstenite::tungstenite::http::header::{AUTHORIZATION, HeaderName};
use tokio_tungstenite::{MaybeTlsStream, WebSocketStream, connect_async};
use crate::constants::{
DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS, DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS,
};
const OPENAI_API_KEY_ENV: &str = "OPENAI_API_KEY";
const MISSING_KEY_MESSAGE: &str =
"Missing OpenAI API Key - a Responses WebSocket call is being made but no key was passed via params or the OPENAI_API_KEY environment variable";
const MISSING_KEY_MESSAGE: &str = "Missing OpenAI API Key - a Responses WebSocket call is being made but no key was passed via params or the OPENAI_API_KEY environment variable";
pub type ResponsesUpstreamWs = WebSocketStream<MaybeTlsStream<TcpStream>>;
type UpstreamTx = SplitSink<ResponsesUpstreamWs, Message>;
@ -83,8 +82,8 @@ impl ResponsesWebSocketConnection {
}
pub async fn recv_text(&self) -> CoreResult<Option<String>> {
let mut socket = self.socket.lock().await;
let Some(socket) = socket.as_mut() else {
let mut socket_guard = self.socket.lock().await;
let Some(socket) = socket_guard.as_mut() else {
return Ok(None);
};
match socket.next().await {
@ -456,9 +455,11 @@ mod tests {
assert_eq!(fourth.event_type, ResponsesWsEventType::ResponseCompleted);
let observed: Vec<_> = observed_rx.collect().await;
assert_eq!(observed.len(), 4);
assert!(observed
.iter()
.all(|event| event.event_type != ResponsesWsEventType::ResponseCreate));
assert!(
observed
.iter()
.all(|event| event.event_type != ResponsesWsEventType::ResponseCreate)
);
}
#[tokio::test]

View file

@ -11,6 +11,8 @@
//! binary turns on. The `python-config` feature additionally pulls in [`python`]
//! for the load-time config reader.
pub mod audio_transcription;
mod client;
pub mod io;
pub mod messages;
pub mod ocr;

View file

@ -11,7 +11,7 @@
use std::sync::Arc;
use litellm_ai_gateway::io::realtime_pool::{upstream_key, PoolConfig, RealtimePool};
use litellm_ai_gateway::io::realtime_pool::{PoolConfig, RealtimePool, upstream_key};
use litellm_ai_gateway::routes;
use litellm_ai_gateway::state::AppState;
use litellm_core::router::{Deployment, LiteLLMParams, Router};

View file

@ -1,8 +1,8 @@
use litellm_core::error::{json_type_name, CoreError};
use litellm_core::CoreResult;
use litellm_core::error::{CoreError, json_type_name};
use litellm_core::messages::transformation::AnthropicMessagesProviderConfig;
use litellm_core::providers::anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG;
use litellm_core::providers::azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG;
use litellm_core::CoreResult;
use serde_json::{Map, Value};
use crate::constants::MESSAGES_ERROR_BODY_MAX_CHARS;

View file

@ -1,5 +1,5 @@
use litellm_core::error::CoreError;
use litellm_core::CoreResult;
use litellm_core::error::CoreError;
use serde_json::Value;
use super::client::http_client;

View file

@ -1,7 +1,7 @@
use litellm_core::messages::transformation::MessagesAuthStrategy;
use litellm_core::routing_utils::provider::{get_custom_llm_provider, CustomLlmProvider};
use litellm_core::CoreError;
use litellm_core::CoreResult;
use litellm_core::messages::transformation::MessagesAuthStrategy;
use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use super::common_utils::{has_header, messages_provider_config, string_headers};
use super::types::{MessagesRequest, ProviderMessagesRequest};

View file

@ -1,14 +1,14 @@
use std::time::Duration;
use litellm_core::error::CoreError;
use serde_json::{json, Map, Value};
use serde_json::{Map, Value, json};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use super::common_utils::{
has_header, messages_provider_config, string_headers, truncate_error_body,
};
use super::{messages, MessagesRequest};
use super::{MessagesRequest, messages};
async fn read_http_request(socket: &mut TcpStream) -> String {
let mut request = Vec::new();

View file

@ -1,11 +1,11 @@
use std::net::IpAddr;
use std::time::{Duration, Instant};
use base64::engine::general_purpose::STANDARD as BASE64_STANDARD;
use base64::Engine;
use base64::engine::general_purpose::STANDARD as BASE64_STANDARD;
use litellm_core::CoreResult;
use litellm_core::error::CoreError;
use litellm_core::ocr::transformation::OcrProviderConfig;
use litellm_core::CoreResult;
use reqwest::Url;
use serde_json::{Map, Value};
@ -18,7 +18,7 @@ use litellm_core::providers::vertex_ai::ocr::transformation::{
VERTEX_AI_DEEPSEEK_OCR_CONFIG, VERTEX_AI_OCR_CONFIG,
};
use super::client::http_client;
use crate::client::http_client;
const ERROR_BODY_MAX_CHARS: usize = 256;
const AZURE_DOCUMENT_INTELLIGENCE_POLL_TIMEOUT_SECS: u64 = 120;

View file

@ -1,11 +1,11 @@
use litellm_core::CoreResult;
use litellm_core::error::CoreError;
use litellm_core::ocr::transformation::OcrResponseHandling;
use litellm_core::CoreResult;
use serde_json::Value;
use super::client::http_client;
use super::common_utils::{poll_document_intelligence, truncate_error_body};
use super::types::ProviderOcrRequest;
use crate::client::http_client;
pub(crate) async fn execute_ocr_provider_call(request: ProviderOcrRequest) -> CoreResult<Value> {
let mut request_builder = http_client().post(&request.url).json(&request.body);

View file

@ -1,11 +1,11 @@
use std::future::Future;
use std::pin::Pin;
use litellm_core::CoreResult;
use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming};
use litellm_core::error::CoreError;
use litellm_core::ocr::transformation::OcrAuthStrategy;
use litellm_core::CoreResult;
use serde_json::{json, Map, Value};
use serde_json::{Map, Value, json};
use super::common_utils::{
convert_document_url_to_data_uri, has_header, ocr_provider_config, string_headers,
@ -292,7 +292,7 @@ fn parse_ocr_pre_call_guardrail_request(
Some(_) => {
return Err(CoreError::InvalidRequest(
"OCR pre_call guardrail optional_params must be an object".to_string(),
))
));
}
None => Map::new(),
};

View file

@ -1,8 +1,7 @@
use litellm_core::call_lifecycle::CallLifecycle;
use litellm_core::CoreResult;
use litellm_core::call_lifecycle::CallLifecycle;
use serde_json::Value;
mod client;
mod common_utils;
mod handler;
mod hooks;
@ -12,7 +11,7 @@ mod types;
pub use types::OcrRequest;
use handler::execute_ocr_provider_call;
use prepare::{prepare_ocr_call, PreparedOcrCall};
use prepare::{PreparedOcrCall, prepare_ocr_call};
pub async fn ocr(request: OcrRequest<'_>) -> CoreResult<Value> {
let PreparedOcrCall { request, hooks } = prepare_ocr_call(request);

View file

@ -1,7 +1,7 @@
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
use litellm_core::routing_utils::provider::{get_custom_llm_provider, CustomLlmProvider};
use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use super::hooks::OcrLifecycleHooks;
use super::types::{OcrRequest, PreparedOcrRequest};

View file

@ -3,12 +3,12 @@ use std::time::Duration;
use litellm_core::error::CoreError;
use litellm_core::ocr::transformation::OcrResponseHandling;
use serde_json::{json, Map, Value};
use serde_json::{Map, Value, json};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::{TcpListener, TcpStream};
use super::common_utils::{has_header, ocr_provider_config, string_headers, truncate_error_body};
use super::{ocr, OcrRequest};
use super::{OcrRequest, ocr};
use crate::integrations::custom_guardrail::{
CustomGuardrail, GuardrailContext, GuardrailDecision, GuardrailError, GuardrailEventHook,
GuardrailFuture, GuardrailRequest,
@ -228,19 +228,23 @@ fn truncate_error_body_does_not_split_multibyte_chars() {
#[test]
fn ocr_dispatch_supports_migrated_providers() {
assert!(ocr_provider_config("mistral", "mistral-ocr-latest").is_some());
assert!(ocr_provider_config("azure_ai", "pixtral-12b-2409")
.expect("azure ai config resolves")
.requires_data_uri_document());
assert!(
ocr_provider_config("azure_ai", "pixtral-12b-2409")
.expect("azure ai config resolves")
.requires_data_uri_document()
);
assert_eq!(
ocr_provider_config("azure_ai", "doc-intelligence/prebuilt-read")
.expect("document intelligence config resolves")
.response_handling(),
OcrResponseHandling::AzureDocumentIntelligencePoll
);
assert!(ocr_provider_config("vertex_ai", "deepseek-ocr-maas")
.expect("vertex deepseek config resolves")
.supported_ocr_params()
.contains(&"temperature"));
assert!(
ocr_provider_config("vertex_ai", "deepseek-ocr-maas")
.expect("vertex deepseek config resolves")
.supported_ocr_params()
.contains(&"temperature")
);
assert!(ocr_provider_config("openai", "gpt-4o").is_none());
}

View file

@ -7,9 +7,9 @@
//!
//! Compiled only under the `python-config` feature.
use litellm_core::CoreResult;
use litellm_core::error::CoreError;
use litellm_core::router::{Deployment, Router};
use litellm_core::CoreResult;
use pyo3::prelude::*;
use crate::gil;

View file

@ -1,8 +1,8 @@
//! Health probes. Simple-route template: a `router()` plus its handlers, in one file.
use axum::Router;
use axum::http::StatusCode;
use axum::routing::get;
use axum::Router;
use crate::state::AppState;

View file

@ -2,13 +2,13 @@
mod service;
use axum::Router;
use axum::body::Body;
use axum::extract::{Json, State};
use axum::http::header::{HeaderMap, HeaderValue, CACHE_CONTROL, CONTENT_TYPE};
use axum::http::StatusCode;
use axum::http::header::{CACHE_CONTROL, CONTENT_TYPE, HeaderMap, HeaderValue};
use axum::response::{IntoResponse, Response};
use axum::routing::post;
use axum::Router;
use litellm_core::CoreError;
use serde_json::{Map, Value};
@ -125,9 +125,9 @@ mod tests {
use std::sync::Arc;
use axum::body::Body;
use axum::http::header::{CACHE_CONTROL, CONTENT_TYPE};
use axum::http::Request;
use axum::http::StatusCode;
use axum::http::header::{CACHE_CONTROL, CONTENT_TYPE};
use litellm_core::router::{Deployment, LiteLLMParams, Router as ModelRouter};
use serde_json::json;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
@ -439,8 +439,8 @@ mod tests {
.await
.expect("response body reads");
assert_eq!(
serde_json::from_slice::<serde_json::Value>(&response_body).expect("error is json")
["error"]["message"],
serde_json::from_slice::<serde_json::Value>(&response_body).expect("error is json")["error"]
["message"],
"messages provider request failed"
);
server.await.expect("upstream task completes");

View file

@ -5,7 +5,7 @@ use litellm_core::{CoreError, CoreResult};
use serde_json::{Map, Value};
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
use crate::messages::{execute_messages, MessagesRequest};
use crate::messages::{MessagesRequest, execute_messages};
pub(crate) enum MessagesResponse {
Json(Value),

View file

@ -6,17 +6,17 @@
mod service;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
use crate::io::realtime_pool::RealtimePool;
use axum::Router;
use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
use axum::extract::{Query, State};
use axum::http::StatusCode;
use axum::response::Response;
use axum::routing::get;
use axum::Router;
use futures_util::{SinkExt, StreamExt};
use litellm_core::realtime::types::RealtimeEvent;
use litellm_core::router::Router as ModelRouter;

View file

@ -9,12 +9,12 @@
use std::time::Duration;
use crate::io::realtime_pool::{upstream_key, RealtimePool};
use crate::io::realtime_pool::{RealtimePool, upstream_key};
use futures_util::{Sink, Stream};
use litellm_core::CoreResult;
use litellm_core::error::CoreError;
use litellm_core::realtime::types::RealtimeEvent;
use litellm_core::router::Router;
use litellm_core::CoreResult;
/// Select a deployment for `model` and splice the client stream to the provider.
///

View file

@ -1,15 +1,15 @@
mod service;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
use axum::Router;
use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
use axum::extract::{Query, State};
use axum::http::StatusCode;
use axum::response::Response;
use axum::routing::get;
use axum::Router;
use futures_util::{Sink, SinkExt, StreamExt};
use litellm_core::responses::types::{ResponsesErrorFrame, ResponsesWsEvent, ResponsesWsEventType};
use litellm_core::router::Router as ModelRouter;

View file

@ -0,0 +1,2 @@
pub mod transformation;
pub mod types;

View file

@ -0,0 +1,57 @@
use serde_json::{Map, Value};
use crate::CoreResult;
use super::types::{AudioTranscriptionRequestData, AudioTranscriptionResponseData};
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum AudioTranscriptionAuth {
Bearer,
AwsSigV4 {
region: String,
service: &'static str,
},
}
pub trait AudioTranscriptionProviderConfig: Sync {
fn supported_transcription_params(&self) -> &'static [&'static str];
fn map_transcription_params(&self, params: &Map<String, Value>) -> Map<String, Value> {
params
.iter()
.filter(|(key, _)| {
self.supported_transcription_params()
.contains(&key.as_str())
})
.map(|(key, value)| (key.clone(), value.clone()))
.collect()
}
fn transform_transcription_request(
&self,
model: &str,
audio: Value,
optional_params: Map<String, Value>,
) -> CoreResult<AudioTranscriptionRequestData>;
fn transform_transcription_response(
&self,
model: &str,
response_json: Value,
) -> CoreResult<AudioTranscriptionResponseData>;
fn complete_url(
&self,
api_base: Option<&str>,
model: &str,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String>;
fn auth_strategy(
&self,
model: &str,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<AudioTranscriptionAuth>;
}

View file

@ -0,0 +1,20 @@
use serde::{Deserialize, Serialize};
use serde_json::Value;
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct AudioTranscriptionRequestData {
pub body: Value,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct AudioTranscriptionResponseData {
pub text: String,
}
impl AudioTranscriptionResponseData {
pub fn into_json(self) -> Value {
serde_json::json!({
"text": self.text,
})
}
}

View file

@ -134,8 +134,8 @@ impl<V: Clone> InMemoryCache<V> {
#[cfg(test)]
mod tests {
use std::sync::{
atomic::{AtomicU64, Ordering},
Arc,
atomic::{AtomicU64, Ordering},
};
use super::InMemoryCache;

View file

@ -1,3 +1,4 @@
pub mod audio_transcription;
pub mod caching;
pub mod call_lifecycle;
pub mod constants;

View file

@ -5,7 +5,7 @@ use crate::messages::types::{
MessageContent, SystemPrompt,
};
use crate::providers::anthropic::messages::transformation::{
non_empty, AnthropicMessagesConfig, ANTHROPIC_MESSAGES_CONFIG,
ANTHROPIC_MESSAGES_CONFIG, AnthropicMessagesConfig, non_empty,
};
use serde_json::{Map, Value};

View file

@ -1,9 +1,9 @@
use std::collections::BTreeSet;
use crate::error::{json_type_name, CoreError, CoreResult};
use crate::error::{CoreError, CoreResult, json_type_name};
use crate::ocr::transformation::{OcrAuthStrategy, OcrProviderConfig, OcrResponseHandling};
use crate::ocr::types::{OcrRequestData, OcrResponseData};
use serde_json::{json, Map, Value};
use serde_json::{Map, Value, json};
use crate::providers::mistral::ocr::transformation::MISTRAL_OCR_CONFIG;
@ -206,11 +206,11 @@ pub fn complete_document_intelligence_url(
AZURE_DOCUMENT_INTELLIGENCE_API_VERSION
);
if let Some(pages) = optional_params.get("pages") {
if let Some(normalized) = normalize_pages_param(pages)? {
url.push_str("&pages=");
url.push_str(&normalized);
}
if let Some(pages) = optional_params.get("pages")
&& let Some(normalized) = normalize_pages_param(pages)?
{
url.push_str("&pages=");
url.push_str(&normalized);
}
Ok(url)
@ -231,7 +231,7 @@ fn document_url_from_mistral_document(document: &Value) -> CoreResult<&str> {
other => {
return Err(CoreError::InvalidRequest(format!(
"Invalid document type: {other}. Must be 'document_url' or 'image_url'"
)))
)));
}
};
object

View file

@ -0,0 +1,310 @@
use serde_json::{Map, Value, json};
use crate::audio_transcription::transformation::{
AudioTranscriptionAuth, AudioTranscriptionProviderConfig,
};
use crate::audio_transcription::types::{
AudioTranscriptionRequestData, AudioTranscriptionResponseData,
};
use crate::error::{CoreError, CoreResult, json_type_name};
use super::aws_base::AwsAuthConfig;
use super::constants::{
AWS_REGION, AWS_REGION_NAME, BEDROCK_RUNTIME_ENDPOINT_TEMPLATE, BEDROCK_SERVICE,
DEFAULT_BEDROCK_REGION,
};
const SUPPORTED_PARAMS: &[&str] = &["language", "prompt", "temperature", "response_format"];
pub static BEDROCK_AUDIO_TRANSCRIPTION_CONFIG: BedrockAudioTranscriptionConfig =
BedrockAudioTranscriptionConfig;
pub struct BedrockAudioTranscriptionConfig;
pub fn bedrock_model_id_and_region(model: &str) -> (String, Option<String>) {
let mut stripped = model;
for prefix in ["bedrock/converse/", "bedrock/", "converse/"] {
if let Some(value) = stripped.strip_prefix(prefix) {
stripped = value;
break;
}
}
let mut region = None;
if let Some((candidate, remainder)) = stripped.split_once('/')
&& is_bedrock_region(candidate)
{
region = Some(candidate.to_string());
stripped = remainder;
}
for prefix in ["nova-2/", "nova/"] {
if let Some(value) = stripped.strip_prefix(prefix) {
stripped = value;
break;
}
}
if region.is_none() {
region = stripped
.strip_prefix("arn:")
.and_then(|value| value.split(':').nth(3))
.filter(|value| !value.is_empty())
.map(str::to_string);
}
(stripped.to_string(), region)
}
fn is_bedrock_region(value: &str) -> bool {
value.len() > 3
&& value.contains('-')
&& value
.chars()
.all(|char| char.is_ascii_alphanumeric() || char == '-')
}
pub fn resolve_bedrock_region(
model_region: Option<&str>,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> String {
if let Some(region) = optional_params
.get("aws_region_name")
.and_then(Value::as_str)
{
return region.to_string();
}
if let Some(region) = model_region {
return region.to_string();
}
env_lookup(AWS_REGION_NAME)
.or_else(|| env_lookup(AWS_REGION))
.unwrap_or_else(|| DEFAULT_BEDROCK_REGION.to_string())
}
fn audio_fields(audio: Value) -> CoreResult<(String, String)> {
let object = audio.as_object().ok_or_else(|| CoreError::InvalidType {
expected: "object",
actual: json_type_name(&audio),
})?;
let data = object
.get("data")
.and_then(Value::as_str)
.filter(|value| !value.is_empty())
.ok_or(CoreError::MissingField("audio.data"))?;
let format = object
.get("format")
.and_then(Value::as_str)
.filter(|value| matches!(*value, "wav" | "mp3" | "flac" | "ogg"))
.ok_or_else(|| {
CoreError::InvalidRequest("audio.format must be wav, mp3, flac, or ogg".to_string())
})?;
Ok((data.to_string(), format.to_string()))
}
fn optional_string<'a>(params: &'a Map<String, Value>, key: &str) -> Option<&'a str> {
params
.get(key)
.and_then(Value::as_str)
.filter(|value| !value.is_empty())
}
impl AudioTranscriptionProviderConfig for BedrockAudioTranscriptionConfig {
fn supported_transcription_params(&self) -> &'static [&'static str] {
SUPPORTED_PARAMS
}
fn transform_transcription_request(
&self,
_model: &str,
audio: Value,
optional_params: Map<String, Value>,
) -> CoreResult<AudioTranscriptionRequestData> {
let (data, format) = audio_fields(audio)?;
let mut instruction = "Transcribe the audio. Respond with only the transcript.".to_string();
if let Some(language) = optional_string(&optional_params, "language") {
instruction.push_str(&format!(" The audio language is {language}."));
}
if let Some(prompt) = optional_string(&optional_params, "prompt") {
instruction.push_str(&format!(" Additional context: {prompt}"));
}
let mut inference_config = Map::from_iter([("maxTokens".to_string(), json!(4096))]);
if let Some(temperature) = optional_params.get("temperature") {
inference_config.insert("temperature".to_string(), temperature.clone());
}
Ok(AudioTranscriptionRequestData {
body: json!({
"messages": [{
"role": "user",
"content": [
{"audio": {"format": format, "source": {"bytes": data}}},
{"text": instruction}
]
}],
"system": [{"text": "You are a transcription assistant."}],
"inferenceConfig": inference_config,
}),
})
}
fn transform_transcription_response(
&self,
_model: &str,
response_json: Value,
) -> CoreResult<AudioTranscriptionResponseData> {
let content = response_json
.get("output")
.and_then(|value| value.get("message"))
.and_then(|value| value.get("content"))
.and_then(Value::as_array)
.ok_or_else(|| {
CoreError::InvalidResponse("Bedrock response has no output content".to_string())
})?;
let mut text = String::new();
for block in content {
if let Some(value) = block.get("text").and_then(Value::as_str) {
text.push_str(value);
}
}
Ok(AudioTranscriptionResponseData { text })
}
fn complete_url(
&self,
api_base: Option<&str>,
model: &str,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
let (model_id, model_region) = bedrock_model_id_and_region(model);
let region = resolve_bedrock_region(model_region.as_deref(), optional_params, env_lookup);
let endpoint = optional_params
.get("aws_bedrock_runtime_endpoint")
.and_then(Value::as_str)
.or(api_base)
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_string)
.unwrap_or_else(|| BEDROCK_RUNTIME_ENDPOINT_TEMPLATE.replace("{region}", &region));
Ok(format!(
"{}/model/{model_id}/converse",
endpoint.trim_end_matches('/')
))
}
fn auth_strategy(
&self,
model: &str,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<AudioTranscriptionAuth> {
let (_, model_region) = bedrock_model_id_and_region(model);
Ok(AudioTranscriptionAuth::AwsSigV4 {
region: resolve_bedrock_region(model_region.as_deref(), optional_params, env_lookup),
service: BEDROCK_SERVICE,
})
}
}
pub fn aws_auth_config(
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> AwsAuthConfig {
let value = |key: &str| {
optional_params
.get(key)
.and_then(Value::as_str)
.map(str::to_string)
};
let env = |key: &str| env_lookup(key);
AwsAuthConfig {
access_key_id: value("aws_access_key_id").or_else(|| env("AWS_ACCESS_KEY_ID")),
secret_access_key: value("aws_secret_access_key").or_else(|| env("AWS_SECRET_ACCESS_KEY")),
session_token: value("aws_session_token").or_else(|| env("AWS_SESSION_TOKEN")),
region_name: value("aws_region_name").or_else(|| env(AWS_REGION_NAME)),
session_name: value("aws_session_name").or_else(|| env("AWS_SESSION_NAME")),
profile_name: value("aws_profile_name").or_else(|| env("AWS_PROFILE_NAME")),
role_name: value("aws_role_name").or_else(|| env("AWS_ROLE_NAME")),
web_identity_token: value("aws_web_identity_token")
.or_else(|| env("AWS_WEB_IDENTITY_TOKEN")),
sts_endpoint: value("aws_sts_endpoint").or_else(|| env("AWS_STS_ENDPOINT")),
external_id: value("aws_external_id").or_else(|| env("AWS_EXTERNAL_ID")),
}
}
#[cfg(test)]
mod tests {
use super::*;
fn no_env(_: &str) -> Option<String> {
None
}
#[test]
fn request_matches_python_shape() {
let params = Map::from_iter([
("language".to_string(), json!("en")),
("prompt".to_string(), json!("Speaker names")),
("temperature".to_string(), json!(0)),
("timestamp_granularities".to_string(), json!(["word"])),
]);
let params = BEDROCK_AUDIO_TRANSCRIPTION_CONFIG.map_transcription_params(&params);
let result = BEDROCK_AUDIO_TRANSCRIPTION_CONFIG
.transform_transcription_request(
"mistral.voxtral-mini-3b-2507",
json!({"data": "AQI=", "format": "wav", "filename": "sample.wav"}),
params,
)
.expect("request");
assert_eq!(
result.body,
json!({
"messages": [{
"role": "user",
"content": [
{"audio": {"format": "wav", "source": {"bytes": "AQI="}}},
{"text": "Transcribe the audio. Respond with only the transcript. The audio language is en. Additional context: Speaker names"}
]
}],
"system": [{"text": "You are a transcription assistant."}],
"inferenceConfig": {"maxTokens": 4096, "temperature": 0}
})
);
}
#[test]
fn response_concatenates_content_blocks() {
let result = BEDROCK_AUDIO_TRANSCRIPTION_CONFIG
.transform_transcription_response(
"model",
json!({"output": {"message": {"content": [{"text": "hello "}, {"text": "world"}]}}}),
)
.expect("response");
assert_eq!(result.text, "hello world");
assert_eq!(result.into_json(), json!({"text": "hello world"}));
}
#[test]
fn invalid_audio_is_rejected() {
let result = BEDROCK_AUDIO_TRANSCRIPTION_CONFIG.transform_transcription_request(
"model",
json!({"data": "AQI="}),
Map::new(),
);
assert!(result.is_err());
}
#[test]
fn region_and_url_precedence_match_python() {
let params = Map::from_iter([("aws_region_name".to_string(), json!("eu-west-1"))]);
let url = BEDROCK_AUDIO_TRANSCRIPTION_CONFIG
.complete_url(
None,
"bedrock/us-east-1/mistral.voxtral-mini-3b-2507",
&params,
&no_env,
)
.expect("url");
assert_eq!(
url,
"https://bedrock-runtime.eu-west-1.amazonaws.com/model/mistral.voxtral-mini-3b-2507/converse"
);
}
}

View file

@ -5,10 +5,10 @@ use std::time::{SystemTime, UNIX_EPOCH};
use crate::caching::in_memory_cache::InMemoryCache;
use crate::error::{CoreError, CoreResult};
use aws_credential_types::provider::ProvideCredentials;
use aws_credential_types::Credentials;
use aws_credential_types::provider::ProvideCredentials;
use aws_sigv4::http_request::{
sign, SignableBody, SignableRequest, SigningParams, SigningSettings,
SignableBody, SignableRequest, SigningParams, SigningSettings, sign,
};
use aws_sigv4::sign::v4;
use aws_smithy_runtime_api::client::identity::Identity;
@ -52,7 +52,7 @@ pub struct AwsAuthConfig {
}
impl AwsAuthConfig {
fn with_environment(self, env_lookup: &dyn Fn(&str) -> Option<String>) -> Self {
fn with_environment(self, env_lookup: &(dyn Fn(&str) -> Option<String> + Sync)) -> Self {
Self {
access_key_id: self.access_key_id.or_else(|| env_lookup(AWS_ACCESS_KEY_ID)),
secret_access_key: self
@ -144,7 +144,7 @@ fn same_role_arns(target: &str, caller: &str) -> bool {
pub fn classify_auth(
config: AwsAuthConfig,
env_lookup: &dyn Fn(&str) -> Option<String>,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> AwsAuthFlow {
let config = config.with_environment(env_lookup);
if let (Some(token), Some(role), Some(session_name)) = (
@ -194,7 +194,7 @@ pub fn classify_auth(
pub async fn resolve_credentials(
config: AwsAuthConfig,
env_lookup: &dyn Fn(&str) -> Option<String>,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> CoreResult<Credentials> {
let resolved = config.clone().with_environment(env_lookup);
let flow = classify_auth(config, env_lookup);
@ -368,11 +368,11 @@ async fn is_already_running_as_role(role: &str, config: &AwsAuthConfig) -> CoreR
if let (Ok(current_role), Ok(token_file)) = (
std::env::var(AWS_ROLE_ARN),
std::env::var(AWS_WEB_IDENTITY_TOKEN_FILE),
) {
if !token_file.is_empty() {
return Ok(same_role_arns(role, &current_role));
}
) && !token_file.is_empty()
{
return Ok(same_role_arns(role, &current_role));
}
let mut loader = aws_config::defaults(aws_config::BehaviorVersion::latest());
if let Some(region) = config.region_name.clone() {
loader = loader.region(aws_types::region::Region::new(region));
@ -639,7 +639,9 @@ mod tests {
);
assert_eq!(
signed.get("Authorization").map(String::as_str),
Some("AWS4-HMAC-SHA256 Credential=AKIDEXAMPLE/20240102/us-east-1/bedrock/aws4_request, SignedHeaders=content-type;host;x-amz-date;x-amz-security-token, Signature=55c027ef47527d3ad63f1735f9d099efdbc99f296ff914bd94e727e24ec0e464")
Some(
"AWS4-HMAC-SHA256 Credential=AKIDEXAMPLE/20240102/us-east-1/bedrock/aws4_request, SignedHeaders=content-type;host;x-amz-date;x-amz-security-token, Signature=55c027ef47527d3ad63f1735f9d099efdbc99f296ff914bd94e727e24ec0e464"
)
);
}

View file

@ -2,6 +2,7 @@ pub const AWS_ACCESS_KEY_ID: &str = "AWS_ACCESS_KEY_ID";
pub const AWS_SECRET_ACCESS_KEY: &str = "AWS_SECRET_ACCESS_KEY";
pub const AWS_SESSION_TOKEN: &str = "AWS_SESSION_TOKEN";
pub const AWS_REGION_NAME: &str = "AWS_REGION_NAME";
pub const AWS_REGION: &str = "AWS_REGION";
pub const AWS_SESSION_NAME: &str = "AWS_SESSION_NAME";
pub const AWS_PROFILE_NAME: &str = "AWS_PROFILE_NAME";
pub const AWS_ROLE_NAME: &str = "AWS_ROLE_NAME";
@ -12,3 +13,6 @@ pub const AWS_STS_ENDPOINT: &str = "AWS_STS_ENDPOINT";
pub const AWS_EXTERNAL_ID: &str = "AWS_EXTERNAL_ID";
pub const BEDROCK_SERVICE: &str = "bedrock";
pub const DEFAULT_SESSION_NAME_PREFIX: &str = "litellm-session";
pub const DEFAULT_BEDROCK_REGION: &str = "us-west-2";
pub const BEDROCK_RUNTIME_ENDPOINT_TEMPLATE: &str =
"https://bedrock-runtime.{region}.amazonaws.com";

View file

@ -2,5 +2,7 @@
//! with Python's `BaseAWSLLM`; the broader core purity guidance is reconciled
//! separately.
#[cfg(feature = "bedrock-auth")]
pub mod audio_transcription;
pub mod aws_base;
mod constants;

View file

@ -1,4 +1,4 @@
use crate::error::{json_type_name, CoreError, CoreResult};
use crate::error::{CoreError, CoreResult, json_type_name};
use crate::ocr::transformation::OcrProviderConfig;
use crate::ocr::types::{OcrRequestData, OcrResponseData};
use serde_json::{Map, Value};

View file

@ -1,6 +1,6 @@
use crate::CoreResult;
use crate::realtime::transformation::RealtimeProviderConfig;
use crate::realtime::types::{RealtimeEvent, RealtimeTransformResult};
use crate::CoreResult;
/// Default OpenAI API base, used when the caller does not override `api_base`.
pub const OPENAI_REALTIME_DEFAULT_API_BASE: &str = "https://api.openai.com";

View file

@ -1,6 +1,6 @@
use crate::responses::types::{ResponsesWsEvent, ResponsesWsTransformResult};
use crate::responses::websocket::{enforce_model, ResponsesWebSocketProviderConfig};
use crate::CoreResult;
use crate::responses::types::{ResponsesWsEvent, ResponsesWsTransformResult};
use crate::responses::websocket::{ResponsesWebSocketProviderConfig, enforce_model};
pub struct OpenAIResponsesWsConfig;

View file

@ -1,7 +1,7 @@
use crate::error::{json_type_name, CoreError, CoreResult};
use crate::error::{CoreError, CoreResult, json_type_name};
use crate::ocr::transformation::OcrProviderConfig;
use crate::ocr::types::{OcrRequestData, OcrResponseData};
use serde_json::{json, Map, Value};
use serde_json::{Map, Value, json};
use crate::providers::mistral::ocr::transformation::MISTRAL_OCR_CONFIG;
@ -140,7 +140,7 @@ fn document_content_item(document: &Value) -> CoreResult<Value> {
other => {
return Err(CoreError::InvalidRequest(format!(
"Unsupported document type: {other}. Expected 'image_url' or 'document_url'"
)))
)));
}
};
let url = object

View file

@ -1,5 +1,5 @@
use crate::realtime::types::{RealtimeEvent, RealtimeTransformResult};
use crate::CoreResult;
use crate::realtime::types::{RealtimeEvent, RealtimeTransformResult};
pub trait RealtimeProviderConfig {
/// Build the upstream WebSocket URL (e.g. `wss://api.openai.com/v1/realtime?model=…`).

View file

@ -1,6 +1,6 @@
use crate::CoreResult;
use crate::constants::{OPENAI_RESPONSES_DEFAULT_API_BASE, OPENAI_RESPONSES_PATH};
use crate::responses::types::{ResponsesWsEvent, ResponsesWsEventType, ResponsesWsTransformResult};
use crate::CoreResult;
pub trait ResponsesWebSocketProviderConfig: Sync {
fn supports_native_websocket(&self) -> bool {

View file

@ -17,7 +17,13 @@ Python-compatible dictionaries.
- Provider dispatch belongs in Rust route modules such as
`litellm_providers::ocr`, not in this PyO3 crate.
- Python owns rollout state and fallback. Rust should return errors; Python
decides whether to raise or fall back.
decides whether to raise or fall back. For a rust-only provider/route (no
Python reference), the Python side is a thin dispatch that calls Rust and
raises when the bridge is unavailable, with no fallback.
- Keep the Python interface minimal (well under 100 lines per route): it only
marshals inputs and calls Rust. Do not add per-route feature flags, and do
not put provider dispatch in `litellm/main.py`; it lives in a thin dispatch
class under `litellm/llms/<provider>/<route>/`.
## Data Handling

View file

@ -10,7 +10,7 @@ name = "_native"
crate-type = ["cdylib"]
[dependencies]
litellm-core.workspace = true
litellm-core = { workspace = true, features = ["bedrock-auth"] }
litellm-ai-gateway = { workspace = true, default-features = false }
pyo3 = { workspace = true, features = ["extension-module"] }
pyo3-async-runtimes.workspace = true

View file

@ -1,8 +1,11 @@
use std::collections::HashMap;
use std::time::Duration;
use litellm_ai_gateway::io::messages::{messages as run_messages, MessagesRequest};
use litellm_ai_gateway::io::ocr::{ocr as run_ocr, OcrRequest};
use litellm_ai_gateway::io::audio_transcription::{
AudioTranscriptionRequest, audio_transcription as run_audio_transcription,
};
use litellm_ai_gateway::io::messages::{MessagesRequest, messages as run_messages};
use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr};
use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
use litellm_core::error::CoreError;
use pyo3::exceptions::{PyRuntimeError, PyValueError};
@ -244,6 +247,93 @@ fn aocr(
})
}
#[pyfunction]
#[pyo3(signature = (model, audio, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
fn transcription(
py: Python<'_>,
model: String,
audio: Py<PyAny>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
extra_headers: Option<Py<PyAny>>,
optional_params: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Py<PyAny>> {
let audio = py_to_json(py, audio.bind(py))?;
let extra_headers = match extra_headers {
Some(headers) => Some(optional_object_to_map(py, "extra_headers", Some(headers))?),
None => None,
};
let optional_params = optional_object_to_map(py, "optional_params", optional_params)?;
let timeout = optional_timeout(timeout_seconds);
let result = gil::release_gil(py, || {
pyo3_async_runtimes::tokio::get_runtime().block_on(run_audio_transcription(
AudioTranscriptionRequest {
model: &model,
audio,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
optional_params,
timeout,
callbacks: Vec::new(),
guardrails: Vec::new(),
request_metadata: Default::default(),
litellm_call_id: None,
},
))
});
match result {
Ok(value) => json_to_py(py, value),
Err(err) => Err(core_error_to_pyerr(err)),
}
}
#[pyfunction]
#[pyo3(signature = (model, audio, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
fn atranscription(
py: Python<'_>,
model: String,
audio: Py<PyAny>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
extra_headers: Option<Py<PyAny>>,
optional_params: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Bound<'_, PyAny>> {
let audio = py_to_json(py, audio.bind(py))?;
let extra_headers = match extra_headers {
Some(headers) => Some(optional_object_to_map(py, "extra_headers", Some(headers))?),
None => None,
};
let optional_params = optional_object_to_map(py, "optional_params", optional_params)?;
let timeout = optional_timeout(timeout_seconds);
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let value = run_audio_transcription(AudioTranscriptionRequest {
model: &model,
audio,
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: custom_llm_provider.as_deref(),
extra_headers,
optional_params,
timeout,
callbacks: Vec::new(),
guardrails: Vec::new(),
request_metadata: Default::default(),
litellm_call_id: None,
})
.await
.map_err(core_error_to_pyerr)?;
Python::attach(|py| json_to_py(py, value))
})
}
type MarshaledMessagesInputs = (Value, Option<Map<String, Value>>, Option<Duration>);
fn marshal_messages_inputs(
@ -341,6 +431,8 @@ fn gil_stats(py: Python<'_>) -> PyResult<Py<PyAny>> {
fn _native(module: &Bound<'_, PyModule>) -> PyResult<()> {
module.add_function(wrap_pyfunction!(ocr, module)?)?;
module.add_function(wrap_pyfunction!(aocr, module)?)?;
module.add_function(wrap_pyfunction!(transcription, module)?)?;
module.add_function(wrap_pyfunction!(atranscription, module)?)?;
module.add_function(wrap_pyfunction!(messages, module)?)?;
module.add_function(wrap_pyfunction!(amessages, module)?)?;
module.add_class::<ResponsesWebSocketConnection>()?;

View file

@ -58,12 +58,12 @@ class DiskCache(BaseCache):
return return_val
def increment_cache(self, key, value: int, **kwargs) -> int:
# get the value
cached_value = self.get_cache(key=key)
init_value = cached_value if isinstance(cached_value, int) else 0
value = init_value + value
self.set_cache(key, value, **kwargs)
return value
with self.disk_cache.transact():
cached_value = self.get_cache(key=key)
init_value = cached_value if isinstance(cached_value, int) else 0
new_value = init_value + value
self.set_cache(key, new_value, **kwargs)
return new_value
async def async_get_cache(self, key, **kwargs):
return self.get_cache(key=key, **kwargs)
@ -76,12 +76,7 @@ class DiskCache(BaseCache):
return return_val
async def async_increment(self, key, value: int, **kwargs) -> int:
# get the value
cached_value = await self.async_get_cache(key=key)
init_value = cached_value if isinstance(cached_value, int) else 0
value = init_value + value
await self.async_set_cache(key, value, **kwargs)
return value
return self.increment_cache(key=key, value=value, **kwargs)
def flush_cache(self):
self.disk_cache.clear()

View file

@ -12,6 +12,7 @@ import json
import sys
import time
import heapq
import threading
from typing import TYPE_CHECKING, Any, List, Optional
if TYPE_CHECKING:
@ -46,6 +47,7 @@ class InMemoryCache(BaseCache):
self.cache_dict: dict = {}
self.ttl_dict: dict = {}
self.expiration_heap: list[tuple[float, str]] = []
self._increment_lock = threading.Lock()
def check_value_size(self, value: Any):
"""
@ -223,12 +225,13 @@ class InMemoryCache(BaseCache):
return_val.append(val)
return return_val
def increment_cache(self, key, value: int, **kwargs) -> int:
# get the value
init_value = self.get_cache(key=key) or 0
value = init_value + value
self.set_cache(key, value, **kwargs)
return value
def increment_cache(self, key, value: float, **kwargs) -> float:
with self._increment_lock:
# keep read-modify-write atomic
init_value = self.get_cache(key=key) or 0
value = init_value + value
self.set_cache(key, value, **kwargs)
return value
async def async_get_cache(self, key, **kwargs):
return self.get_cache(key=key, **kwargs)
@ -241,11 +244,7 @@ class InMemoryCache(BaseCache):
return return_val
async def async_increment(self, key, value: float, **kwargs) -> float:
# get the value
init_value = await self.async_get_cache(key=key) or 0
value = init_value + value
await self.async_set_cache(key, value, **kwargs)
return value
return self.increment_cache(key=key, value=value, **kwargs)
async def async_increment_pipeline(
self, increment_list: List["RedisPipelineIncrementOperation"], **kwargs

View file

@ -9,6 +9,7 @@ from openai.types.chat import ChatCompletionToolParam
from openai.types.responses.function_tool_param import FunctionToolParam
from openai.types.shared_params.function_definition import FunctionDefinition
from litellm.types.llms.anthropic import AnthropicMessagesTool
from litellm.types.utils import ChatCompletionMessageToolCall
@ -75,6 +76,20 @@ def transform_mcp_tool_to_openai_responses_api_tool(
)
def transform_mcp_tool_to_anthropic_tool(mcp_tool: MCPTool) -> AnthropicMessagesTool:
"""Convert an MCP tool to an Anthropic Messages API tool."""
from litellm.litellm_core_utils.prompt_templates.common_utils import (
sanitize_input_schema_for_anthropic,
)
return AnthropicMessagesTool(
name=mcp_tool.name,
description=mcp_tool.description or "",
input_schema=sanitize_input_schema_for_anthropic(mcp_tool.inputSchema),
type="custom",
)
async def load_mcp_tools(
session: ClientSession, format: Literal["mcp", "openai"] = "mcp"
) -> Union[List[MCPTool], List[ChatCompletionToolParam]]:

View file

@ -91,7 +91,9 @@ class AnthropicCacheControlHook(CustomPromptManagement):
# Pass through non-message injection points for provider-specific handling
if remaining_points:
non_default_params["cache_control_injection_points"] = remaining_points
non_default_params["cache_control_injection_points"] = AnthropicCacheControlHook._stamped_as_judged(
remaining_points
)
return model, processed_messages, non_default_params
@ -310,6 +312,35 @@ class AnthropicCacheControlHook(CustomPromptManagement):
return ChatCompletionCachedContent(type="ephemeral", ttl=ttl)
return ChatCompletionCachedContent(type="ephemeral")
@staticmethod
def _stamped_as_judged(points: list[CacheControlInjectionPoint]) -> list[dict[str, object]]:
"""Mark written-back points as having passed the client cache_control judgment.
Builds copies because config-owned point dicts are shared across
requests; mutating them would leak the stamp into future requests.
"""
return [{**point, "_litellm_judged": True} for point in points]
@staticmethod
def _should_stand_down(
points: list[CacheControlInjectionPoint],
messages: list[AllMessageValues],
system: str | list | None,
tools: list | None,
) -> bool:
"""Whether configured injection points must yield to client-set cache_control.
Points that a prior pass over this request already judged and wrote
back carry the internal judged stamp; any re-entry (acompletion
re-entering completion, the async-to-sync /v1/messages dispatch,
interceptor sub-calls reusing the request kwargs) must not re-judge
them, because by then the messages carry litellm's own injected marks
and the judgment would misread those as client breakpoints.
"""
if all(point.get("_litellm_judged") for point in points):
return False
return AnthropicCacheControlHook._request_has_cache_control(messages, system, tools)
@staticmethod
def _request_has_cache_control(
messages: list[AllMessageValues],
@ -322,7 +353,9 @@ class AnthropicCacheControlHook(CustomPromptManagement):
stand down entirely rather than add more, per the auto-caching contract.
Tools count: they are a breakpoint the client can mark, they count toward
the provider's four-block limit, and caching only the tool definitions is
a common pattern, so injecting alongside them can exceed the cap.
a common pattern, so injecting alongside them can exceed the cap. Tools
carry the mark either at the top level (Anthropic shape) or nested under
``function`` (OpenAI shape); the Anthropic chat transform accepts both.
"""
if any(AnthropicCacheControlHook._count_cache_control_blocks(msg) for msg in messages):
return True
@ -330,7 +363,14 @@ class AnthropicCacheControlHook(CustomPromptManagement):
if any(isinstance(block, dict) and block.get("cache_control") is not None for block in system):
return True
if tools is not None:
return any(isinstance(tool, dict) and tool.get("cache_control") is not None for tool in tools)
return any(
isinstance(tool, dict)
and (
tool.get("cache_control") is not None
or (isinstance(tool.get("function"), dict) and tool["function"].get("cache_control") is not None)
)
for tool in tools
)
return False
@staticmethod
@ -392,13 +432,23 @@ class AnthropicCacheControlHook(CustomPromptManagement):
custom_llm_provider: str | None,
tools: list | None = None,
) -> None:
"""For /chat/completions: add default injection points to the request params.
"""For /chat/completions: resolve the injection points the request should carry.
No-op when injection points are already configured (explicit config wins).
Seeding the param lets the existing prompt-management gate and the
AnthropicCacheControlHook run unchanged.
Configured injection points win over the automatic defaults, but stand
down entirely when the client already marked its own cache_control
breakpoints (messages or tools): injecting alongside them clashes with
the client's caching strategy and can exceed the provider's four-block
limit. The judgment happens once per request; points a prior pass
wrote back carry the judged stamp and are never re-judged (see
``_should_stand_down``). Seeding the param lets the existing
prompt-management gate and the AnthropicCacheControlHook run
unchanged.
"""
if non_default_params.get("cache_control_injection_points"):
if AnthropicCacheControlHook._should_stand_down(
non_default_params["cache_control_injection_points"], messages, None, tools
):
non_default_params.pop("cache_control_injection_points")
return
points = AnthropicCacheControlHook.get_default_injection_points(
messages=messages,
@ -421,18 +471,26 @@ class AnthropicCacheControlHook(CustomPromptManagement):
) -> Tuple[List[Dict], str | list | None]:
"""Extract cache_control_injection_points from kwargs and apply if present.
When none are configured but ``litellm.enable_anthropic_prompt_caching``
is on, synthesize default breakpoints for the native /v1/messages path.
Pops the key from kwargs; if remaining (non-message) points exist they
are written back so downstream transforms can handle them.
Configured points stand down entirely when the client already marked
its own cache_control breakpoints anywhere in the request. The
judgment happens once per request; points a prior pass wrote back
carry the judged stamp and are never re-judged (see
``_should_stand_down``). When none are configured but
``litellm.enable_anthropic_prompt_caching`` is on, synthesize default
breakpoints for the native /v1/messages path. Pops the key from kwargs;
if remaining (non-message) points exist they are written back so
downstream transforms can handle them.
"""
typed_messages = cast(list[AllMessageValues], messages) # cast-ok: Anthropic-shaped dicts from v1/messages
configured = cast( # cast-ok: kwargs is untyped; this key only holds the documented injection-point list
list[CacheControlInjectionPoint] | None, kwargs.pop("cache_control_injection_points", None)
)
if configured and AnthropicCacheControlHook._should_stand_down(configured, typed_messages, system, tools):
return messages, system
injection_points: list[CacheControlInjectionPoint] = configured or []
if not injection_points and model is not None:
injection_points = AnthropicCacheControlHook.get_default_injection_points(
messages=cast(list[AllMessageValues], messages), # cast-ok: Anthropic-shaped dicts from v1/messages
messages=typed_messages,
system=system,
tools=tools,
model=model,
@ -447,7 +505,7 @@ class AnthropicCacheControlHook(CustomPromptManagement):
injection_points=injection_points,
)
if remaining:
kwargs["cache_control_injection_points"] = remaining
kwargs["cache_control_injection_points"] = AnthropicCacheControlHook._stamped_as_judged(remaining)
return messages, system
@property

View file

@ -42,6 +42,7 @@ from litellm.types.utils import (
)
if TYPE_CHECKING: # newer pattern to avoid importing pydantic objects on __init__.py
from litellm.types.llms.anthropic import AnthropicInputSchema
from litellm.types.llms.openai import ChatCompletionImageObject
DEFAULT_USER_CONTINUE_MESSAGE = ChatCompletionUserMessage(content="Please continue.", role="user")
@ -1046,6 +1047,31 @@ def unpack_legacy_defs(
return schema
def sanitize_input_schema_for_anthropic(input_schema: dict) -> "AnthropicInputSchema":
"""Coerce an arbitrary tool input_schema into the shape Anthropic accepts.
Anthropic requires ``type == "object"``, only recognises ``$defs`` (legacy
``definitions`` / OpenAPI ``components.schemas`` refs must be inlined first),
and rejects keys outside ``AnthropicInputSchema``. Both the chat
(``AnthropicConfig._map_tool_helper``) and Anthropic Messages MCP paths run
a schema through here so an external MCP schema cannot succeed on one route
and 400 on the other.
"""
from litellm.types.llms.anthropic import AnthropicInputSchema
normalized = dict(input_schema) if input_schema else {}
if normalized.get("type") != "object":
normalized["type"] = "object"
if "properties" not in normalized:
normalized["properties"] = {}
normalized = unpack_legacy_defs(normalized, copy=True)
allowed_keys = set(AnthropicInputSchema.__annotations__.keys())
filtered = {key: value for key, value in normalized.items() if key in allowed_keys}
return AnthropicInputSchema(**filtered)
def _get_image_mime_type_from_url(url: str) -> Optional[str]:
"""
Get mime type for common image URLs

View file

@ -29,7 +29,9 @@ from litellm.constants import (
RESPONSE_FORMAT_TOOL_NAME,
)
from litellm.litellm_core_utils.core_helpers import map_finish_reason
from litellm.litellm_core_utils.prompt_templates.common_utils import unpack_legacy_defs
from litellm.litellm_core_utils.prompt_templates.common_utils import (
sanitize_input_schema_for_anthropic,
)
from litellm.llms.base_llm.base_utils import type_to_response_format_param
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
from litellm.types.llms.anthropic import (
@ -634,7 +636,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
mcp_server: Optional[AnthropicMcpServerTool] = None
if tool["type"] == "function" or tool["type"] == "custom":
_input_schema: dict = tool["function"].get(
_input_schema = tool["function"].get(
"parameters",
{
"type": "object",
@ -642,28 +644,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
},
)
# Anthropic requires input_schema.type to be "object". Normalize
# schemas from external sources (MCP servers, OpenAI callers) that
# may omit the type field or use a non-object type.
if _input_schema.get("type") != "object":
litellm.verbose_logger.debug(
"_map_tool_helper: coercing input_schema type from %r to "
"'object' for Anthropic compatibility (tool: %s)",
_input_schema.get("type"),
tool["function"].get("name"),
)
_input_schema = dict(_input_schema) # avoid mutating caller's dict
_input_schema["type"] = "object"
if "properties" not in _input_schema:
_input_schema["properties"] = {}
# Inline legacy / OpenAPI $refs before the allow-list filter strips
# their backing def blocks (https://github.com/BerriAI/litellm/issues/26692).
_input_schema = unpack_legacy_defs(_input_schema, copy=True)
_allowed_properties = set(AnthropicInputSchema.__annotations__.keys())
input_schema_filtered = {k: v for k, v in _input_schema.items() if k in _allowed_properties}
input_anthropic_schema: AnthropicInputSchema = AnthropicInputSchema(**input_schema_filtered)
input_anthropic_schema = sanitize_input_schema_for_anthropic(_input_schema)
_tool = AnthropicMessagesTool(
name=tool["function"]["name"],

View file

@ -485,6 +485,41 @@ def anthropic_messages_handler(
mock_response=litellm_params.mock_response,
)
# Expand litellm_proxy MCP references through the MCP gateway before dispatch, so every
# downstream path (native passthrough and both bridges) gets real tools rather than a
# reference the provider cannot resolve. Popped from kwargs so it never reaches the provider.
skip_mcp_handler = kwargs.pop("_skip_mcp_handler", False)
if not skip_mcp_handler and tools:
from litellm.llms.anthropic.experimental_pass_through.messages.mcp_handler import (
anthropic_messages_with_mcp,
)
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
LiteLLM_Proxy_MCP_Handler,
)
if LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(tools=tools):
return anthropic_messages_with_mcp(
max_tokens=max_tokens,
messages=messages,
model=model,
metadata=metadata,
stop_sequences=stop_sequences,
stream=stream,
system=system,
temperature=temperature,
thinking=thinking,
tool_choice=tool_choice,
tools=tools,
top_k=top_k,
top_p=top_p,
container=container,
api_key=api_key,
api_base=api_base,
client=client,
custom_llm_provider=custom_llm_provider,
**kwargs,
)
anthropic_messages_provider_config: Optional[BaseAnthropicMessagesConfig] = None
if custom_llm_provider is not None and custom_llm_provider in [provider.value for provider in LlmProviders]:

View file

@ -0,0 +1,176 @@
"""
MCP gateway support for the Anthropic `/v1/messages` API.
Mirrors ``litellm.responses.mcp.chat_completions_handler`` but speaks the
Anthropic Messages shapes: tools carry an ``input_schema``, the model asks for a
tool through a ``tool_use`` content block, and results are fed back as
``tool_result`` blocks in a user message.
"""
from typing import Any, AsyncIterator, Mapping, Sequence, Union
from litellm._logging import verbose_logger
from litellm.responses.mcp.request_context import MCPRequestContext
from litellm.types.llms.anthropic import (
AnthropicMessagesTool,
AnthropicMessagesToolResultParam,
AnthropicMessagesUserMessageParam,
)
from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
MAX_MCP_TOOL_USE_ITERATIONS = 10
def _get_response_content(response: AnthropicMessagesResponse) -> Sequence[Mapping[str, Any]]:
content = response.get("content")
if not isinstance(content, list):
return ()
return tuple(block for block in content if isinstance(block, dict))
def _extract_tool_use_blocks(response: AnthropicMessagesResponse) -> Sequence[Mapping[str, Any]]:
"""Return the ``tool_use`` content blocks the model emitted."""
return tuple(block for block in _get_response_content(response) if block.get("type") == "tool_use")
def _get_stop_reason(response: AnthropicMessagesResponse) -> Union[str, None]:
stop_reason = response.get("stop_reason")
return stop_reason if isinstance(stop_reason, str) else None
def _build_tool_result_message(tool_results: Sequence[Mapping[str, Any]]) -> AnthropicMessagesUserMessageParam:
"""Turn executed tool results into the user message Anthropic expects."""
return AnthropicMessagesUserMessageParam(
role="user",
content=tuple(
AnthropicMessagesToolResultParam(
type="tool_result",
tool_use_id=str(result.get("tool_call_id") or ""),
content=str(result.get("result") or ""),
)
for result in tool_results
),
)
async def anthropic_messages_with_mcp(
max_tokens: int,
messages: Sequence[Mapping[str, Any]],
model: str,
tools: Union[Sequence[Mapping[str, Any]], None] = None,
**kwargs: Any, # kwargs-ok: forwarded verbatim to litellm.anthropic_messages, which owns the param contract
) -> Union[AnthropicMessagesResponse, AsyncIterator[Any]]:
"""
Expand litellm_proxy MCP references for `/v1/messages` and run the tool loop.
The MCP gateway owns the expansion so the reference resolves against the
caller's own credentials and access control, rather than being handed to the
upstream provider as a url it cannot reach.
"""
import litellm
from litellm.experimental_mcp_client.tools import (
transform_mcp_tool_to_anthropic_tool,
)
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
LiteLLM_Proxy_MCP_Handler,
)
mcp_references, other_tools = LiteLLM_Proxy_MCP_Handler._parse_mcp_tools(tools)
if not mcp_references:
return await litellm.anthropic_messages(
max_tokens=max_tokens,
messages=list(messages),
model=model,
tools=list(tools) if tools else None,
_skip_mcp_handler=True,
**kwargs,
)
context = MCPRequestContext.resolve(kwargs=dict(kwargs), tools=tools)
(
deduplicated_mcp_tools,
tool_server_map,
) = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform(
context.user_api_key_auth,
mcp_references,
litellm_trace_id=context.litellm_trace_id,
mcp_auth_header=context.mcp_auth_header,
mcp_server_auth_headers=context.mcp_server_auth_headers,
request_tags=list(context.request_tags) if context.request_tags else None,
)
anthropic_tools: Sequence[AnthropicMessagesTool] = tuple(
transform_mcp_tool_to_anthropic_tool(mcp_tool) for mcp_tool in deduplicated_mcp_tools
)
all_tools = [*anthropic_tools, *(other_tools or ())]
should_auto_execute = LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools(
mcp_tools_with_litellm_proxy=mcp_references
)
stream = bool(kwargs.pop("stream", False))
base_call_args: Mapping[str, Any] = {
"max_tokens": max_tokens,
"model": model,
"tools": all_tools or None,
"_skip_mcp_handler": True,
**kwargs,
}
if not should_auto_execute:
return await litellm.anthropic_messages(messages=list(messages), stream=stream, **base_call_args)
working_messages: Sequence[Mapping[str, Any]] = tuple(messages)
response: AnthropicMessagesResponse = await litellm.anthropic_messages(
messages=list(working_messages), stream=False, **base_call_args
)
for _ in range(MAX_MCP_TOOL_USE_ITERATIONS):
if _get_stop_reason(response) != "tool_use":
break
tool_use_blocks = _extract_tool_use_blocks(response)
if not tool_use_blocks:
break
tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls(
tool_server_map=tool_server_map,
tool_calls=list(tool_use_blocks),
user_api_key_auth=context.user_api_key_auth,
mcp_auth_header=context.mcp_auth_header,
mcp_server_auth_headers=context.mcp_server_auth_headers,
oauth2_headers=context.oauth2_headers,
raw_headers=context.raw_headers,
litellm_call_id=context.litellm_call_id,
litellm_trace_id=context.litellm_trace_id,
request_tags=list(context.request_tags) if context.request_tags else None,
)
# Every tool call was skipped, so there is nothing to feed back; a
# tool_result message with empty content is rejected by Anthropic.
if not tool_results:
break
working_messages = (
*working_messages,
{"role": "assistant", "content": list(_get_response_content(response))},
_build_tool_result_message(tool_results),
)
response = await litellm.anthropic_messages(messages=list(working_messages), stream=False, **base_call_args)
else:
verbose_logger.warning(
f"MCP tool loop hit its {MAX_MCP_TOOL_USE_ITERATIONS} iteration cap for model {model}; "
"returning the last response"
)
if stream:
from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import (
FakeAnthropicMessagesStreamIterator,
)
return FakeAnthropicMessagesStreamIterator(response)
return response

View file

@ -0,0 +1,84 @@
import base64
from typing import Union
import httpx
from litellm.litellm_core_utils.audio_utils.utils import process_audio_file
from litellm.rust_bridge import transcription as rust_transcription_bridge
from litellm.types.utils import FileTypes, TranscriptionResponse
class BedrockAudioTranscriptionRustDispatch:
@staticmethod
def _audio_payload(audio_file: FileTypes) -> dict[str, object]:
processed_audio = process_audio_file(audio_file)
formats = {
"audio/flac": "flac",
"audio/mpeg": "mp3",
"audio/mp3": "mp3",
"audio/ogg": "ogg",
"audio/wav": "wav",
"audio/x-wav": "wav",
}
audio_format = formats.get(processed_audio.content_type) or (
processed_audio.filename.rsplit(".", 1)[-1].lower() if "." in processed_audio.filename else ""
)
if audio_format not in {"wav", "mp3", "flac", "ogg"}:
raise ValueError(f"Unsupported Bedrock audio format for file {processed_audio.filename!r}")
return {
"data": base64.b64encode(processed_audio.file_content).decode("ascii"),
"format": audio_format,
"filename": processed_audio.filename,
}
def audio_transcriptions(
self,
*,
model: str,
audio_file: FileTypes,
api_key: str | None,
api_base: str | None,
custom_llm_provider: str,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: Union[float, httpx.Timeout] | None,
) -> TranscriptionResponse:
rust_response = rust_transcription_bridge.transcription(
model=model,
audio=self._audio_payload(audio_file),
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout=timeout,
)
if rust_response is None:
raise RuntimeError("Rust audio transcription bridge is unavailable")
return TranscriptionResponse(**rust_response)
async def async_audio_transcriptions(
self,
*,
model: str,
audio_file: FileTypes,
api_key: str | None,
api_base: str | None,
custom_llm_provider: str,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: Union[float, httpx.Timeout] | None,
) -> TranscriptionResponse:
rust_response = await rust_transcription_bridge.atranscription(
model=model,
audio=self._audio_payload(audio_file),
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout=timeout,
)
if rust_response is None:
raise RuntimeError("Rust audio transcription bridge is unavailable")
return TranscriptionResponse(**rust_response)

View file

@ -4,6 +4,7 @@ import time
from typing import Any, Dict, List, Literal, Optional, Union, cast
from httpx import Headers, Response
from pydantic import TypeAdapter, ValidationError
from litellm.litellm_core_utils.cloud_storage_security import (
BEDROCK_MANAGED_S3_BATCH_PREFIX,
@ -19,6 +20,7 @@ from litellm.types.llms.bedrock import (
BedrockOutputDataConfig,
BedrockS3InputDataConfig,
BedrockS3OutputDataConfig,
BedrockTag,
)
from litellm.types.llms.openai import (
AllMessageValues,
@ -38,6 +40,18 @@ _S3_BATCH_FILE_UUID_SUFFIX_PATTERN = re.compile(
r"-[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}\.jsonl$"
)
_BEDROCK_TAGS_ADAPTER: TypeAdapter[list[BedrockTag]] = TypeAdapter(list[BedrockTag])
def _validate_bedrock_tags(raw_tags: object) -> list[BedrockTag]:
try:
return _BEDROCK_TAGS_ADAPTER.validate_python(raw_tags, strict=True)
except ValidationError as e:
raise ValueError(
"Invalid 'bedrock_tags' value. Expected a list of {'key': <str>, 'value': <str>} dicts, "
f"e.g. [{{'key': 'team', 'value': 'genai'}}]. Got: {raw_tags!r}"
) from e
class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
"""
@ -201,6 +215,11 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
"roleArn": role_arn,
}
config_bedrock_tags = litellm_params.get("bedrock_tags")
bedrock_tags = config_bedrock_tags if config_bedrock_tags is not None else optional_params.get("bedrock_tags")
if bedrock_tags is not None:
bedrock_request["tags"] = _validate_bedrock_tags(bedrock_tags)
# Add optional parameters if provided
completion_window = create_batch_data.get("completion_window")
if completion_window:

View file

@ -133,6 +133,32 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig):
def get_config(cls):
return super().get_config()
def validate_environment(
self,
headers: dict,
model: str,
messages: List[AllMessageValues],
optional_params: dict,
litellm_params: dict,
api_key: str | None = None,
api_base: str | None = None,
) -> dict:
api_key = self._get_api_key(api_key)
if api_key is None:
raise ValueError("FIREWORKS_API_KEY is not set")
validated_headers = OpenAIGPTConfig.validate_environment(
self,
headers=headers,
model=model,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
api_key=api_key,
api_base=api_base,
)
return self._add_session_affinity_header(validated_headers, litellm_params)
def get_supported_openai_params(self, model: str):
# Base parameters supported by all models
supported_params = [

View file

@ -64,9 +64,16 @@ class FireworksAIMixin:
if api_key is None:
raise ValueError("FIREWORKS_API_KEY is not set")
validated_headers = {"Authorization": "Bearer {}".format(api_key), **headers}
if not any(key.lower() == "x-session-affinity" for key in validated_headers):
session_id = get_fireworks_session_id(litellm_params)
if session_id:
validated_headers["x-session-affinity"] = session_id
return validated_headers
auth_headers = {"Authorization": "Bearer {}".format(api_key), **headers}
content_type_header = (
{} if any(key.lower() == "content-type" for key in auth_headers) else {"Content-Type": "application/json"}
)
return self._add_session_affinity_header({**auth_headers, **content_type_header}, litellm_params)
def _add_session_affinity_header(self, headers: dict, litellm_params: dict) -> dict:
if any(key.lower() == "x-session-affinity" for key in headers):
return headers
session_id = get_fireworks_session_id(litellm_params)
if not session_id:
return headers
return {**headers, "x-session-affinity": session_id}

View file

@ -27,7 +27,6 @@ from typing import (
TYPE_CHECKING,
Any,
AsyncIterator,
Callable,
Coroutine,
Dict,
Iterable,
@ -81,22 +80,19 @@ from litellm.constants import (
from litellm.exceptions import LiteLLMUnknownProvider
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.asyncify import run_async_function
from litellm.litellm_core_utils.chat_completion_agentic_loop import (
maybe_run_chat_completion_agentic_loop,
)
from litellm.litellm_core_utils.audio_utils.utils import (
calculate_request_duration,
get_audio_file_for_health_check,
)
from litellm.litellm_core_utils.completion_timeout import CompletionTimeout
from litellm.litellm_core_utils.request_timeout_resolver import (
get_configured_request_timeout,
from litellm.litellm_core_utils.chat_completion_agentic_loop import (
maybe_run_chat_completion_agentic_loop,
)
from litellm.litellm_core_utils.completion_timeout import CompletionTimeout
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.litellm_core_utils.get_litellm_params import (
AWS_CREDENTIAL_KWARGS_KEYS,
OPTIONAL_KWARGS_KEYS,
)
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.litellm_core_utils.get_provider_specific_headers import (
ProviderSpecificHeaderUtils,
)
@ -112,6 +108,9 @@ from litellm.litellm_core_utils.mock_functions import (
from litellm.litellm_core_utils.prompt_templates.common_utils import (
get_content_from_model_response,
)
from litellm.litellm_core_utils.request_timeout_resolver import (
get_configured_request_timeout,
)
from litellm.llms.base_llm import BaseConfig, BaseImageGenerationConfig
from litellm.llms.base_llm.base_model_iterator import (
convert_model_response_to_streaming,
@ -213,7 +212,6 @@ from .llms.bedrock.embed.embedding import BedrockEmbedding
from .llms.bedrock.image_edit.handler import BedrockImageEdit
from .llms.bedrock.image_generation.image_handler import BedrockImageGeneration
from .llms.bytez.chat.transformation import BytezChatConfig
from .llms.gdc.chat.transformation import GDCGeminiConfig
from .llms.clarifai.chat.transformation import ClarifaiConfig
from .llms.codestral.completion.handler import CodestralTextCompletion
from .llms.cohere.embed import handler as cohere_embed
@ -222,24 +220,25 @@ from .llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from .llms.custom_llm import CustomLLM, custom_chat_llm_router
from .llms.databricks.embed.handler import DatabricksEmbeddingHandler
from .llms.deprecated_providers import aleph_alpha, palm
from .llms.gdc.chat.transformation import GDCGeminiConfig
from .llms.gemini.common_utils import get_api_key_from_env
from .llms.groq.chat.handler import GroqChatCompletion
from .llms.heroku.chat.transformation import HerokuChatConfig
from .llms.huggingface.embedding.handler import HuggingFaceEmbedding
from .llms.lemonade.chat.transformation import LemonadeChatConfig
from .llms.nlp_cloud.chat.handler import completion as nlp_cloud_chat_completion
from .llms.oci.chat.transformation import OCIChatConfig
from .llms.ollama.completion import handler as ollama
from .llms.oobabooga.chat import oobabooga
from .llms.openai.completion.handler import OpenAITextCompletion
from .llms.openai.image_variations.handler import OpenAIImageVariationsHandler
from .llms.openai.openai import OpenAIChatCompletion
from .llms.nvidia_riva.audio_transcription.handler import (
NvidiaRivaAudioTranscription,
)
from .llms.nvidia_riva.audio_transcription.transformation import (
NvidiaRivaAudioTranscriptionConfig,
)
from .llms.oci.chat.transformation import OCIChatConfig
from .llms.ollama.completion import handler as ollama
from .llms.oobabooga.chat import oobabooga
from .llms.openai.completion.handler import OpenAITextCompletion
from .llms.openai.image_variations.handler import OpenAIImageVariationsHandler
from .llms.openai.openai import OpenAIChatCompletion
from .llms.openai.transcriptions.handler import OpenAIAudioTranscription
from .llms.openai_like.chat.handler import OpenAILikeChatHandler
from .llms.openai_like.embedding.handler import OpenAILikeEmbeddingHandler
@ -7722,6 +7721,32 @@ def transcription(
headers=extra_headers,
provider_config=provider_config, # type: ignore[arg-type]
)
elif custom_llm_provider == "bedrock":
from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch
dispatch = BedrockAudioTranscriptionRustDispatch()
if atranscription:
response = dispatch.async_audio_transcriptions(
model=model,
audio_file=file,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout=timeout,
)
else:
response = dispatch.audio_transcriptions(
model=model,
audio_file=file,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout=timeout,
)
elif provider_config is not None:
response = base_llm_http_handler.audio_transcriptions(
model=model,

View file

@ -93,6 +93,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
AuthorizationCodeConfig,
ClientCredentialsConfig,
CredError,
IdJagConfig,
PassthroughConfig,
@ -2739,14 +2740,18 @@ class MCPServerManager:
)
if not conflicts:
return auth, extra_headers
if isinstance(spec.config, (TokenExchangeConfig, AuthorizationCodeConfig, IdJagConfig)):
# The resolver owns the per-user credential here (token_exchange's exchanged
# token, authorization_code's stored token, id_jag's minted assertion). It is
# authoritative: a guardrail such
# as MCPJWTSigner, static_headers, or any other injected Authorization must NOT
# shadow it (otherwise the upstream gets e.g. the signer's JWT instead of the
# exchanged token and rejects it). Drop the conflicting header so the resolved
# token reaches upstream.
if isinstance(
spec.config,
(TokenExchangeConfig, AuthorizationCodeConfig, IdJagConfig, ClientCredentialsConfig),
):
# The resolver owns the credential here (token_exchange's exchanged token,
# authorization_code's stored token, id_jag's minted assertion,
# client_credentials' gateway-minted M2M token). It is authoritative: a
# guardrail such as MCPJWTSigner, static_headers, or any other injected
# Authorization must NOT shadow it (otherwise the upstream gets e.g. the
# signer's JWT instead of the minted token and rejects it, and for M2M the
# one-shot 401 refetch is lost with it). Drop the conflicting header so the
# resolved token reaches upstream.
return auth, _without_authorization(extra_headers)
# Other modes: an Authorization already supplied via extra_headers (a forwarded caller
# header or static_headers) is intentional and wins; v1 applies those last.

View file

@ -22,6 +22,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
ApiKeyConfig,
AuthorizationCodeConfig,
ClientAuth,
ClientCredentialsConfig,
ClientSecretAuth,
CredError,
IdJagConfig,
@ -70,10 +71,10 @@ def to_server_spec(server: MCPServer) -> Optional[ServerSpec]:
an ``assert_never`` tail, so a newly added auth mode fails the type gate here until it is
explicitly mapped or explicitly deferred, rather than silently falling through to v1. Live
modes: ``none``, the static-header family (``api_key`` plus the Authorization schemes,
all shared-key), ``oauth2`` per-user tokens (``authorization_code``), ``oauth2_token_exchange``
(OBO), and the client-forwarded token modes ``true_passthrough`` / ``oauth_delegate``
(``PassthroughConfig``); client_credentials (M2M), delegated/passthrough oauth2, and SigV4
return None and stay on v1.
all shared-key), ``oauth2`` per-user tokens (``authorization_code``), ``oauth2`` M2M
(``client_credentials``), ``oauth2_token_exchange`` (OBO), and the client-forwarded token
modes ``true_passthrough`` / ``oauth_delegate`` (``PassthroughConfig``); delegated/passthrough
oauth2 and SigV4 return None and stay on v1.
"""
if server.is_byok:
return None # per-user BYOK source not migrated yet -> defer to v1 (any auth_type)
@ -95,14 +96,7 @@ def to_server_spec(server: MCPServer) -> Optional[ServerSpec]:
case MCPAuth.basic:
return _shared_key_spec(server, resource, "Authorization", "Basic", encode=True)
case MCPAuth.oauth2:
if server.needs_user_oauth_token and not server.delegate_auth_to_upstream:
return ServerSpec(
server_id=server.server_id,
resource=resource,
config=AuthorizationCodeConfig(),
)
# client_credentials (M2M) and delegate/passthrough oauth2 stay on v1
return None
return _oauth2_spec(server, resource)
case MCPAuth.oauth2_id_jag:
return _id_jag_spec(server, resource)
case MCPAuth.true_passthrough | MCPAuth.oauth_delegate:
@ -114,6 +108,47 @@ def to_server_spec(server: MCPServer) -> Optional[ServerSpec]:
assert_never(auth_type)
def _oauth2_spec(server: MCPServer, resource: str) -> ServerSpec | None:
"""Dispatch the oauth2 auth_type across its sub-modes: M2M, gateway-managed interactive, or v1.
``client_credentials`` (the explicit ``oauth2_flow`` opt-in) builds the M2M spec, per-user
``authorization_code`` without upstream delegation builds the interactive spec, and the
delegate/passthrough shapes defer to v1 (None).
"""
if server.has_client_credentials:
return _client_credentials_spec(server, resource)
if server.needs_user_oauth_token and not server.delegate_auth_to_upstream:
return ServerSpec(
server_id=server.server_id,
resource=resource,
config=AuthorizationCodeConfig(),
)
return None
def _client_credentials_spec(server: MCPServer, resource: str) -> ServerSpec:
"""Build a client_credentials (M2M) spec; the explicit ``oauth2_flow`` opt-in owns the server.
Missing grant fields (``client_id``/``client_secret``/``token_url``) are NOT a reason to defer:
v1 would connect unauthenticated and the upstream's 401 gets absorbed into an empty tool list,
so the arm fails closed with ``misconfigured`` instead, naming the missing fields (mirrors the
OBO ownership rule). ``audience`` is forwarded only when the operator set it; a missing one is
omitted, not derived, since a fabricated value risks the IdP rejecting the grant.
"""
return ServerSpec(
server_id=server.server_id,
resource=resource,
config=ClientCredentialsConfig(
client_id=server.client_id,
client_secret=SecretStr(server.client_secret) if server.client_secret else None,
token_url=server.token_url,
scopes=tuple(server.scopes or ()),
audience=server.audience,
token_endpoint_auth_method=server.token_endpoint_auth_method,
),
)
def _token_exchange_spec(server: MCPServer, resource: str) -> Optional[ServerSpec]:
"""Build a token_exchange (OBO) spec, or defer (None) when it is not OBO-configured.

View file

@ -0,0 +1,348 @@
"""The ``client_credentials`` (M2M) arm's token source and retrying bearer auth.
Implements the client-credentials behavior contract for the v2 resolver:
- **Acquisition**: POST ``grant_type=client_credentials`` to the configured token endpoint with
the configured scopes and (when set) the IdP's ``audience`` parameter, authenticating the
client per ``token_endpoint_auth_method`` (RFC 6749 section 2.3.1, shared helper).
- **Caching**: tokens are cached per ``(client identity, server)`` where the identity key hashes
``token_url`` / ``client_id`` / ``client_secret`` / auth method / scopes / audience — rotating
or re-scoping the credentials changes the key, so a stale token can never be served for the
new identity (the contract's rotation-invalidation clause).
- **Expiry**: the cache TTL respects ``expires_in`` minus a skew so an entry lapses before the
real token does; a response with no ``expires_in`` is cached briefly
(``default_ttl_seconds``), not assumed long-lived. No refresh_token is ever expected.
- **401 recovery**: ``ClientCredentialsBearerAuth`` retries an upstream request exactly once
after a 401 — discard the cached token, mint a fresh one, resend; a second failure surfaces
the upstream's own auth error unchanged.
- **No user context**: nothing here reads a ``Subject``; every caller shares the one client
identity.
The token-endpoint POST is injected (``M2MTokenEndpointPost``) so the grant orchestration is
testable without a live IdP; ``post_client_credentials_grant`` is the httpx edge and the one
place the untyped response boundary is contained. Failures are values: the source returns
``Result[OAuthToken, CredError]``; only the httpx edge touches exceptions.
"""
from __future__ import annotations
import asyncio
import hashlib
import time
from collections.abc import AsyncGenerator, Awaitable, Callable, Generator
from dataclasses import dataclass
from typing import Annotated, Literal
import httpx
from pydantic import BaseModel, ConfigDict, Field, SecretStr, TypeAdapter, ValidationError
from typing_extensions import assert_never
from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import (
InMemoryTokenCacheBackend,
OAuthToken,
TokenCacheBackend,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.result import (
Error,
Ok,
Result,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.types import (
ClientCredentialsConfig,
CredError,
)
class TokenEndpointSuccess(BaseModel):
"""The endpoint returned a JSON object; field validation is the caller's job."""
model_config = ConfigDict(frozen=True)
tag: Literal["success"] = "success"
body: dict[str, object]
class TokenEndpointDenied(BaseModel):
"""The endpoint answered but did not grant a token (an HTTP error or a non-JSON body)."""
model_config = ConfigDict(frozen=True)
tag: Literal["denied"] = "denied"
status_code: int
detail: str
class TokenEndpointUnreachable(BaseModel):
"""The endpoint could not be reached (DNS, TLS, connect/read failure)."""
model_config = ConfigDict(frozen=True)
tag: Literal["unreachable"] = "unreachable"
detail: str
TokenEndpointOutcome = Annotated[
TokenEndpointSuccess | TokenEndpointDenied | TokenEndpointUnreachable,
Field(discriminator="tag"),
]
M2MTokenEndpointPost = Callable[[str, "dict[str, str]", "dict[str, str]"], Awaitable[TokenEndpointOutcome]]
_TOKEN_BODY_ADAPTER: TypeAdapter[dict[str, object]] = TypeAdapter(dict[str, object])
async def post_client_credentials_grant(
url: str, form: dict[str, str], headers: dict[str, str]
) -> TokenEndpointOutcome:
"""POST the grant to the token endpoint and classify the transport outcome.
The httpx edge: litellm's handler is partially typed (and raises ``HTTPStatusError`` itself on
a 4xx/5xx), so the untyped boundary is contained here and every field the caller reads comes
out of a validated ``TokenEndpointOutcome``.
"""
from litellm.llms.custom_httpx.http_handler import ( # noqa: PLC0415 # defer heavy handler import to call time
get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # handler is partially typed
)
from litellm.types.llms.custom_http import httpxSpecialProvider # noqa: PLC0415 # deferred with the handler import
try:
client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check)
response = await client.post( # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # handler is partially typed
url, headers={"Accept": "application/json", **headers}, data=form
)
except httpx.HTTPStatusError as status_err:
status_code = status_err.response.status_code
return TokenEndpointDenied(status_code=status_code, detail=f"token endpoint returned HTTP {status_code}")
except Exception as exc: # noqa: BLE001 # any transport failure is the same outcome: unreachable
return TokenEndpointUnreachable(detail=str(exc))
if not isinstance(response, httpx.Response):
return TokenEndpointUnreachable(detail="token endpoint returned no response")
try:
body = _TOKEN_BODY_ADAPTER.validate_json(response.content)
except ValidationError:
return TokenEndpointDenied(
status_code=response.status_code, detail="token endpoint returned a non-JSON-object body"
)
return TokenEndpointSuccess(body=body)
def _parse_expires_in(raw: object) -> int | None:
if isinstance(raw, bool):
return None
if isinstance(raw, int):
return raw
if isinstance(raw, str):
try:
return int(raw)
except ValueError:
return None
return None
def _parse_granted_scopes(raw: object) -> tuple[str, ...] | None:
return tuple(raw.split()) if isinstance(raw, str) and raw else None
@dataclass(frozen=True, slots=True)
class _PreparedGrant:
"""A validated, ready-to-POST grant plus the identity key its token caches under."""
token_url: str
form: dict[str, str]
headers: dict[str, str]
identity_key: str
class ClientCredentialsTokenSource:
"""Cached M2M access tokens, one per ``(client identity, server)``.
``get`` serves from the cache while the entry's TTL (derived from ``expires_in`` minus
``expiry_skew_seconds``) holds, fetching under a per-server lock so concurrent misses
produce one grant. ``refetch`` is the 401-recovery path: it drops the failed token and
mints a fresh one, unless a concurrent caller already replaced it.
"""
def __init__(
self,
post: M2MTokenEndpointPost = post_client_credentials_grant,
*,
backend: TokenCacheBackend | None = None,
default_ttl_seconds: float = 300.0,
expiry_skew_seconds: float = 60.0,
min_cache_seconds: float = 10.0,
max_locks: int = 1024,
clock: Callable[[], float] = time.time,
) -> None:
self._post = post
self._backend: TokenCacheBackend = backend or InMemoryTokenCacheBackend(clock=clock)
self._default_ttl_seconds = default_ttl_seconds
self._expiry_skew_seconds = expiry_skew_seconds
self._min_cache_seconds = min_cache_seconds
self._max_locks = max_locks
self._clock = clock
self._locks: dict[str, asyncio.Lock] = {}
def _lock(self, server_id: str) -> asyncio.Lock:
"""Per-server single-flight lock, bounded so ephemeral server ids (e.g. the REST tools
preview mints a fresh id per call) cannot grow the dict for the life of the process.
Evicting the oldest entry while a task still holds it only means a concurrent caller for
that server may run its own grant — single-flight is an optimization, not correctness.
"""
if server_id not in self._locks and len(self._locks) >= self._max_locks:
self._locks.pop(next(iter(self._locks)), None)
return self._locks.setdefault(server_id, asyncio.Lock())
async def get(self, server_id: str, config: ClientCredentialsConfig) -> Result[OAuthToken, CredError]:
match _prepare_grant(config):
case Error(err):
return Error(err)
case Ok(grant):
cached = await self._backend.get(grant.identity_key, server_id)
if cached is not None:
return Ok(cached)
async with self._lock(server_id):
cached = await self._backend.get(grant.identity_key, server_id)
if cached is not None:
return Ok(cached)
return await self._fetch_and_cache(server_id, grant)
async def refetch(self, server_id: str, config: ClientCredentialsConfig, failed_access_token: str) -> str | None:
"""Replace a token the upstream just 401'd; returns the fresh bearer value or ``None``.
Runs under the same per-server lock as ``get``: if a concurrent caller already replaced
the failed token, that replacement is returned without another grant, so a burst of 401s
yields one fetch. A failed refetch returns ``None`` and the caller surfaces the
upstream's original auth error (the contract's retry-once-then-give-up clause).
"""
match _prepare_grant(config):
case Error(_):
return None
case Ok(grant):
async with self._lock(server_id):
cached = await self._backend.get(grant.identity_key, server_id)
if cached is not None and cached.access_token != failed_access_token:
return cached.access_token
await self._backend.delete(grant.identity_key, server_id)
match await self._fetch_and_cache(server_id, grant):
case Ok(token):
return token.access_token
case Error(_):
return None
async def _fetch_and_cache(self, server_id: str, grant: _PreparedGrant) -> Result[OAuthToken, CredError]:
outcome = await self._post(grant.token_url, grant.form, grant.headers)
match outcome:
case TokenEndpointUnreachable():
return Error(CredError.of_upstream_unavailable(f"OAuth2 token endpoint unreachable: {outcome.detail}"))
case TokenEndpointDenied():
if outcome.status_code >= 500:
return Error(CredError.of_upstream_unavailable(f"OAuth2 token endpoint failed: {outcome.detail}"))
return Error(CredError.of_misconfigured(f"OAuth2 client_credentials grant rejected: {outcome.detail}"))
case TokenEndpointSuccess():
return await self._cache_token(server_id, grant, outcome.body)
assert_never(outcome)
async def _cache_token(
self, server_id: str, grant: _PreparedGrant, body: dict[str, object]
) -> Result[OAuthToken, CredError]:
access_token = body.get("access_token")
if not isinstance(access_token, str) or not access_token:
return Error(CredError.of_misconfigured("OAuth2 token response is missing 'access_token'"))
expires_in = _parse_expires_in(body.get("expires_in"))
token = OAuthToken(
access_token=access_token,
expires_at=self._clock() + expires_in if expires_in is not None else None,
scopes=_parse_granted_scopes(body.get("scope")) or (),
)
# The min-cache floor is itself capped at the token's real lifetime, so a token whose
# expires_in is below the skew is never served past its actual expiry; a non-positive
# expires_in caches nothing (every request re-fetches, serialized by the per-server lock).
ttl = (
max(expires_in - self._expiry_skew_seconds, min(float(expires_in), self._min_cache_seconds), 0.0)
if expires_in is not None
else self._default_ttl_seconds
)
if ttl > 0:
await self._backend.set(grant.identity_key, server_id, token, ttl)
return Ok(token)
def _prepare_grant(config: ClientCredentialsConfig) -> Result[_PreparedGrant, CredError]:
if not config.client_id or not config.client_secret or not config.token_url:
missing = ", ".join(
name
for name, present in (
("client_id", bool(config.client_id)),
("client_secret", bool(config.client_secret)),
("token_url", bool(config.token_url)),
)
if not present
)
return Error(CredError.of_misconfigured(f"client_credentials config is missing: {missing}"))
from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import ( # noqa: PLC0415 # keep package v1-free at import time
build_token_endpoint_client_auth,
)
client_auth = build_token_endpoint_client_auth(
auth_method=config.token_endpoint_auth_method,
client_id=config.client_id,
client_secret=config.client_secret.get_secret_value(),
)
form = {
"grant_type": "client_credentials",
**client_auth.body,
**({"scope": " ".join(config.scopes)} if config.scopes else {}),
**({"audience": config.audience} if config.audience else {}),
}
return Ok(
_PreparedGrant(
token_url=config.token_url,
form=form,
headers=client_auth.headers,
identity_key=_identity_key(config),
)
)
def _identity_key(config: ClientCredentialsConfig) -> str:
"""Hash of everything that names the client identity; any rotation yields a new key."""
material = "\n".join(
(
config.token_url or "",
config.client_id or "",
config.client_secret.get_secret_value() if config.client_secret else "",
config.token_endpoint_auth_method or "",
" ".join(config.scopes),
config.audience or "",
)
)
return hashlib.sha256(material.encode("utf-8")).hexdigest()
class ClientCredentialsBearerAuth(httpx.Auth):
"""Bearer auth that retries an upstream 401 exactly once with a freshly minted token.
The initial token was already resolved (so config/IdP failures surfaced as typed errors
before any upstream request); ``refetch`` is the source's 401-recovery callback. If the
refetch fails, or the retried request 401s again, the upstream's response stands.
"""
def __init__(self, access_token: str, refetch: Callable[[str], Awaitable[str | None]]) -> None:
self.header_name = "Authorization"
self._access_token = SecretStr(access_token)
self._refetch = refetch
async def async_auth_flow(self, request: httpx.Request) -> AsyncGenerator[httpx.Request, httpx.Response]:
token = self._access_token.get_secret_value()
request.headers[self.header_name] = f"Bearer {token}"
response = yield request
if response.status_code != 401:
return
fresh = await self._refetch(token)
if fresh is None:
return
self._access_token = SecretStr(fresh)
request.headers[self.header_name] = f"Bearer {fresh}"
yield request
def sync_auth_flow(self, request: httpx.Request) -> Generator[httpx.Request, httpx.Response, None]:
raise RuntimeError("ClientCredentialsBearerAuth only supports async httpx clients")

View file

@ -9,18 +9,24 @@ at runtime instead of returning `None`.
`none`, `api_key` (shared-key source), and `passthrough` (forwards the caller's own inbound token)
are live, as is `authorization_code`, which reads the user's token from the injected
`OAuthTokenStore`, and `token_exchange`, which swaps the caller's inbound token through the
injected `TokenExchanger`. The remaining arms are `not_implemented` stubs that each land in a
follow-up PR with their seam. Pure v2: no imports from v1.
`OAuthTokenStore`, `token_exchange`, which swaps the caller's inbound token through the injected
`TokenExchanger`, and `client_credentials`, which mints and caches the gateway's M2M token through
the injected `ClientCredentialsTokenSource`. The remaining arms are `not_implemented` stubs that
each land in a follow-up PR with their seam. Pure v2: no imports from v1.
"""
from __future__ import annotations
import hashlib
from functools import partial
import httpx
from typing_extensions import assert_never
from litellm.proxy._experimental.mcp_server.outbound_credentials.client_credentials import (
ClientCredentialsBearerAuth,
ClientCredentialsTokenSource,
)
from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import (
NoOpAuth,
StaticHeaderAuth,
@ -104,11 +110,13 @@ class UpstreamCredentialProvider:
token_exchanger: TokenExchanger | None = None,
token_endpoint: TokenEndpointClient | None = None,
exchanged_tokens: ExchangedTokenCache | None = None,
client_credentials_source: ClientCredentialsTokenSource | None = None,
) -> None:
self._oauth_token_store: OAuthTokenStore = oauth_token_store or _NullOAuthTokenStore()
self._token_exchanger: TokenExchanger = token_exchanger or _NullTokenExchanger()
self._token_endpoint: TokenEndpointClient = token_endpoint or TokenEndpointClient()
self._exchanged_tokens: ExchangedTokenCache = exchanged_tokens or ExchangedTokenCache()
self._client_credentials_source = client_credentials_source or ClientCredentialsTokenSource()
async def resolve_credentials(self, subject: Subject, server: ServerSpec) -> Result[httpx.Auth, CredError]:
match server.config:
@ -118,8 +126,8 @@ class UpstreamCredentialProvider:
return self._api_key(config)
case PassthroughConfig():
return self._passthrough(subject)
case ClientCredentialsConfig():
return _not_implemented(AuthSpecKind.client_credentials)
case ClientCredentialsConfig() as config:
return await self._client_credentials(server.server_id, config)
case TokenExchangeConfig() as config:
return await self._token_exchange(subject, server, config)
case IdJagConfig() as config:
@ -215,6 +223,23 @@ class UpstreamCredentialProvider:
return Error(CredError.of_unauthorized("Authorization required: complete the OAuth flow for this server."))
return Ok(StaticHeaderAuth(f"Bearer {token.access_token}", header_name="Authorization"))
async def _client_credentials(
self, server_id: str, config: ClientCredentialsConfig
) -> Result[httpx.Auth, CredError]:
"""The M2M arm: resolve a cached (or freshly minted) gateway token; no user context.
The token is resolved here, before any upstream request, so a misconfigured grant or an
unreachable IdP surfaces as a typed ``CredError``. The returned auth carries the source's
``refetch``, so an upstream 401 is retried exactly once with a freshly minted token (the
contract's invalid-token recovery); a second 401 surfaces the upstream's own error.
"""
match await self._client_credentials_source.get(server_id, config):
case Ok(token):
refetch = partial(self._client_credentials_source.refetch, server_id, config)
return Ok(ClientCredentialsBearerAuth(token.access_token, refetch))
case Error(err):
return Error(err)
async def _token_exchange(
self, subject: Subject, server: ServerSpec, config: TokenExchangeConfig
) -> Result[StaticHeaderAuth, CredError]:
@ -245,7 +270,9 @@ class UpstreamCredentialProvider:
Used after an upstream rejects the injected credential, so the next resolve re-mints rather
than serving the same rejected token until TTL. `token_exchange` and `id_jag` hold a
re-mintable cached credential here; other modes are a no-op.
re-mintable cached credential here; `client_credentials` recovers inside its own auth flow
(`ClientCredentialsBearerAuth` retries the 401'd request once with a fresh token), and
other modes are a no-op.
"""
if subject.inbound_token is None:
return

View file

@ -0,0 +1,190 @@
"""Producer and consumer helpers for the gateway-level DCR session token.
The aggregate ``/mcp`` front door (``mcp_gateway_dcr``) issues the identity-only session
tokens defined in :mod:`.session_token`. The gateway token endpoint mints them (producer)
after SSO sign-in, and at the MCP admission edge the gateway derives the session signing
key from the proxy ``master_key``, opens the bearer, and admits the request under the
recovered litellm user (consumer), reloading the live user record and policy before
anything runs. This module is the pure surface for both sides; the token-endpoint and
admission wiring live in their respective call sites.
The signing key is derived with the same memory-hard scrypt construction as
:func:`~.bridge_credentials.envelope_keys_from_master_key` but under a distinct domain
label, so session tokens and bridge envelopes never share key material: a token of one
family is unverifiable in the other by key separation, on top of the distinct issuers,
prefixes, and claim shapes.
"""
import hashlib
from datetime import datetime
from functools import lru_cache
from typing import Literal, TypeAlias
from pydantic import BaseModel, ConfigDict, SecretStr
from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import (
OpenedSessionToken,
SessionExpired,
SessionKeys,
SessionPrincipal,
is_session_refresh_token,
is_session_token,
open_session_refresh_token,
open_session_token,
)
_SESSION_SIGNING_KEY_DOMAIN = b"litellm-mcp-gateway:session-signing:"
# scrypt work factors (RFC 7914), identical to the envelope KDF: memory-hard so a captured
# session token is not a cheap offline oracle for the master key.
_SCRYPT_N = 2**15
_SCRYPT_R = 8
_SCRYPT_P = 1
_SCRYPT_MAXMEM = 128 * _SCRYPT_N * _SCRYPT_R * _SCRYPT_P * 2
_DERIVED_KEY_BYTES = 32
@lru_cache(maxsize=8)
def session_keys_from_master_key(master_key: str) -> SessionKeys:
"""Derive the session signing key from the proxy master key.
A memory-hard scrypt KDF (RFC 7914) over a session-specific domain-label salt yields a
256-bit subkey from the one secret, so the producer (mint) and consumer (open) agree on
the key without persisting any. The domain label differs from both envelope labels in
:mod:`.bridge_credentials`, so compromise or misuse of one token family never crosses
into the other. The result is cached (the master key is fixed for a process); rotating
``master_key`` invalidates every outstanding session, which is the intended behavior
for a signing-key change.
"""
signing = hashlib.scrypt(
master_key.encode(),
salt=_SESSION_SIGNING_KEY_DOMAIN,
n=_SCRYPT_N,
r=_SCRYPT_R,
p=_SCRYPT_P,
maxmem=_SCRYPT_MAXMEM,
dklen=_DERIVED_KEY_BYTES,
).hex()
return SessionKeys(signing_key=SecretStr(signing))
class NotSessionBearer(BaseModel):
"""The bearer is not session-shaped; admission continues on its normal path."""
model_config = ConfigDict(frozen=True)
tag: Literal["not_session_bearer"] = "not_session_bearer"
class SessionBearerAdmitted(BaseModel):
"""A valid session access token: the principal to admit under after a live reload."""
model_config = ConfigDict(frozen=True)
tag: Literal["admitted"] = "admitted"
principal: SessionPrincipal
class SessionBearerInvalid(BaseModel):
"""The bearer is session-shaped but must not admit (expired, tampered, wrong key, or a
refresh token presented at the tool-call edge); admission fails closed with the
``invalid_token`` challenge rather than falling through to another arm. ``expired``
distinguishes a routine expiry (debug-log worthy) from a tampered or foreign token."""
model_config = ConfigDict(frozen=True)
tag: Literal["invalid"] = "invalid"
expired: bool = False
SessionBearerResult: TypeAlias = NotSessionBearer | SessionBearerAdmitted | SessionBearerInvalid
def _strip_bearer(value: str) -> str:
parts = value.split(None, 1)
if len(parts) == 2 and parts[0].lower() == "bearer":
return parts[1]
return value
def is_session_bearer_shaped(authorization_value: str) -> bool:
"""Cheap, keyless test that an ``Authorization`` value carries a session token of either
kind (optional ``Bearer`` scheme stripped). The admission edge engages the session arm
for an access token (to admit) and for a refresh token (to reject it explicitly, since
a refresh credential is never usable at the tool-call edge); anything else falls
through to normal admission."""
candidate = _strip_bearer(authorization_value)
return is_session_token(candidate) or is_session_refresh_token(candidate)
def resolve_session_bearer(
authorization_value: str,
keys: SessionKeys,
now: datetime,
) -> SessionBearerResult:
"""Classify an ``Authorization`` value presented at the aggregate MCP edge.
Strips an optional ``Bearer`` scheme, then returns ``NotSessionBearer`` for a
non-session bearer (normal admission continues), ``SessionBearerAdmitted`` with the
recovered principal for a valid access token, and ``SessionBearerInvalid`` for a
session-shaped bearer that must not admit. Never raises: total over hostile input via
:func:`~.session_token.open_session_token`.
A refresh token is ``SessionBearerInvalid`` here: it is a valid gateway credential but
only ever presented back to the token endpoint, so admission must fail it closed rather
than let it fall through to another arm.
"""
candidate = _strip_bearer(authorization_value)
if is_session_refresh_token(candidate):
return SessionBearerInvalid()
if not is_session_token(candidate):
return NotSessionBearer()
opened = open_session_token(candidate, keys, now)
if isinstance(opened, OpenedSessionToken):
return SessionBearerAdmitted(principal=opened.principal)
return SessionBearerInvalid(expired=isinstance(opened, SessionExpired))
class SessionRefreshOpened(BaseModel):
"""A valid session refresh token presented to the token endpoint: the principal to
re-validate and renew under."""
model_config = ConfigDict(frozen=True)
tag: Literal["opened"] = "opened"
principal: SessionPrincipal
class SessionRefreshInvalid(BaseModel):
"""The presented refresh grant is not a valid session refresh token for this client
(not refresh-shaped, will not open, or bound to a different ``client_id``); the token
endpoint fails the refresh closed."""
model_config = ConfigDict(frozen=True)
tag: Literal["invalid"] = "invalid"
SessionRefreshResult: TypeAlias = SessionRefreshOpened | SessionRefreshInvalid
def open_session_refresh_bearer(
refresh_value: str,
keys: SessionKeys,
now: datetime,
expected_client_id: str,
) -> SessionRefreshResult:
"""Open a session refresh token presented on a ``refresh_token`` grant.
The token-endpoint mirror of :func:`resolve_session_bearer`: strips an optional
``Bearer`` scheme, then returns ``SessionRefreshOpened`` with the recovered principal,
or ``SessionRefreshInvalid`` for anything that is not a valid session refresh token
issued to ``expected_client_id``. Never raises. The client binding (RFC 6749 section 6)
stops a refresh token stolen from one DCR client from being renewed through another;
``client_id`` is not a secret (the caller presents it), so a plain equality check is
sufficient and, unlike ``hmac.compare_digest`` on ``str``, does not raise on non-ASCII.
"""
candidate = _strip_bearer(refresh_value)
if not is_session_refresh_token(candidate):
return SessionRefreshInvalid()
opened = open_session_refresh_token(candidate, keys, now)
if not isinstance(opened, OpenedSessionToken):
return SessionRefreshInvalid()
if opened.principal.client_id != expected_client_id:
return SessionRefreshInvalid()
return SessionRefreshOpened(principal=opened.principal)

View file

@ -0,0 +1,361 @@
"""Identity-only session tokens for the gateway-level (aggregate ``/mcp``) DCR front door.
A DCR client that signs in through LiteLLM SSO holds ONE bearer that carries ONLY a
litellm identity; unlike the :mod:`.envelope` bridge bearer it seals no upstream
credential, because the custody model vaults every upstream token server-side in
``LiteLLM_MCPUserCredentials`` and egress resolves them by user at call time. The token
is therefore a stable REFERENCE, not an authorization: admission reloads the live user
record and policy on every request, so deactivating the user (or their team) kills
outstanding sessions immediately without a revocation store.
Wire shape: ``llm_session_`` (access) / ``llm_srefresh_`` (refresh) + an HS256 JWT,
the same signing approach as :mod:`.envelope`. Claims are ``iss``/``iat``/``exp``
plus ``jti`` (per-mint uniqueness, so two tokens minted in the same second never
collide and a future revocation list has a stable handle), ``kind``, ``user_id``, and
``client_id``; ``client_id`` binds the refresh token
to the DCR client it was issued to (RFC 6749 section 6) and is carried on the access
token for parity and audit. There is no encrypted payload: nothing in a session token
is secret beyond the signature, and reprs never print the signed value because minted
tokens are ``SecretStr``.
This module is pure and unwired: it imports nothing from endpoint or edge code, reads
no proxy globals, and takes all key material and the clock as explicit parameters.
Failures are values: :func:`open_session_token` and :func:`open_session_refresh_token`
are total over hostile, attacker-controlled input and return a
``SessionTokenOpenError`` variant rather than raising. PyJWT's ``iat``/``nbf``/``exp``
validators are disabled for the same reasons documented in :mod:`.envelope` (they
raise on hostile claim types and compare against the wall clock instead of the
injected ``now``); the strict pydantic claims model is the sole, total type gate.
"""
from __future__ import annotations
import secrets
from datetime import datetime, timedelta
from typing import Literal, TypeAlias
import jwt
from pydantic import BaseModel, ConfigDict, Field, SecretStr, ValidationError
SESSION_TOKEN_PREFIX = "llm_session_"
"""Marker prefix on every serialized session ACCESS token so the admission edge can cheaply
tell a gateway session from a litellm key, JWT, or bridge envelope before doing any
cryptography. Distinct from the ``llm_env_``/``llm_refresh_`` envelope prefixes."""
SESSION_REFRESH_PREFIX = "llm_srefresh_"
"""Marker prefix on every serialized session REFRESH token. A distinct prefix keeps the two
credentials routable without crypto and, together with the signed ``kind`` claim, stops one
from being presented where the other is expected: the refresh token is only ever presented
back to the token endpoint, never at the MCP edge."""
SESSION_ISSUER = "litellm-mcp-gateway"
"""``iss`` claim stamped into every session token and required back on open. Distinct from
the envelope issuer so a token of one family can never validate in the other even under a
hypothetical shared signing key."""
SESSION_TTL_SECONDS = 3600
"""Session ACCESS token lifetime (1h), matching the access-envelope and BYOK session bearer
windows: a client-held credential never outlives a bounded window, and each refresh
re-validates the live user before re-minting."""
SESSION_REFRESH_TTL_SECONDS = 1209600
"""Session REFRESH token lifetime (14 days), matching the refresh-envelope bound. Each
renewal re-validates the sealed user against the live record (deactivation gates it) and
rotates the refresh token, so the practical bound is idle time, not a fixed session."""
MAX_SESSION_TOKEN_BYTES = 4096
"""Size cap on the serialized token (prefix + JWT, in bytes) and on any candidate accepted
by the openers. Session claims are small; the only variable-length field is ``client_id``
(a sealed DCR client record), and 4096 leaves ample headroom under common 8-16KB header
limits while bounding hostile input before JWT parsing."""
_SESSION_JWT_ALGORITHM = "HS256"
SessionTokenKind = Literal["session", "session_refresh"]
"""Which credential a session token is. Stamped into the signed claims and required to match
on open, so a signature-valid token of one kind cannot be replayed as the other even if its
wire prefix is swapped (the prefix is not part of the signed payload; this claim is)."""
class SessionPrincipal(BaseModel):
"""The litellm user a session token identifies and the DCR client it was issued to.
``user_id`` is the SSO-established litellm user subject, never a credential: admission
reloads the live user record by it, so current role, team, and revocation state are
enforced at use time rather than frozen at mint time. ``client_id`` is the (stateless,
gateway-sealed) DCR client identifier the token was issued to; the token endpoint
requires it to match on the refresh grant.
"""
model_config = ConfigDict(frozen=True)
user_id: str = Field(min_length=1)
client_id: str = Field(min_length=1)
class SessionKeys(BaseModel):
"""Injected key material: the HS256 signing key.
``signing_key`` must be at least 32 bytes: HS256's HMAC-SHA256 has a 256-bit security
level, RFC 7518 requires a key of at least that size, and a shorter key makes PyJWT
emit ``InsecureKeyLengthWarning``.
"""
model_config = ConfigDict(frozen=True)
signing_key: SecretStr = Field(min_length=32)
class MintedSessionToken(BaseModel):
"""A minted session token: the client-held bearer value and when it expires."""
model_config = ConfigDict(frozen=True)
token: SecretStr
expires_at: datetime
class OpenedSessionToken(BaseModel):
"""A validated session token of either kind: the principal it was minted for."""
model_config = ConfigDict(frozen=True)
principal: SessionPrincipal
class SessionTokenTooLarge(BaseModel):
"""The serialized token exceeded ``MAX_SESSION_TOKEN_BYTES``; carries sizes only. Only
reachable through an oversized ``client_id``, which registration should have bounded."""
model_config = ConfigDict(frozen=True)
tag: Literal["session_token_too_large"] = "session_token_too_large"
size_bytes: int
max_bytes: int
SessionTokenMintError: TypeAlias = SessionTokenTooLarge
class NotASessionToken(BaseModel):
"""The candidate does not carry the expected session prefix."""
model_config = ConfigDict(frozen=True)
tag: Literal["not_a_session_token"] = "not_a_session_token"
class SessionBadSignature(BaseModel):
"""The JWT signature does not verify under the provided signing key."""
model_config = ConfigDict(frozen=True)
tag: Literal["session_bad_signature"] = "session_bad_signature"
class SessionExpired(BaseModel):
"""The token's ``exp`` is not in the future relative to the provided ``now``."""
model_config = ConfigDict(frozen=True)
tag: Literal["session_expired"] = "session_expired"
class SessionMalformed(BaseModel):
"""The token is not a well-formed session token: undecodable JWT, wrong issuer, wrong
``kind``, or missing/mistyped/extra claims."""
model_config = ConfigDict(frozen=True)
tag: Literal["session_malformed"] = "session_malformed"
SessionTokenOpenError: TypeAlias = NotASessionToken | SessionBadSignature | SessionExpired | SessionMalformed
class _SessionClaims(BaseModel):
"""Decoded-claims boundary that pins the exact shape the mints emit.
``user_id``/``client_id`` mirror the ``min_length`` constraints of
:class:`SessionPrincipal` so any claim set that validates here also constructs a
principal, keeping the openers raise-free: a correctly signed JWT with an empty
identity claim fails here and maps to ``SessionMalformed``. ``strict`` rejects coerced
types (``exp: "123"``) and ``extra="forbid"`` rejects any claim the gateway never
mints; PyJWT's own registered-claim validators are disabled at decode (see module
docstring), so this model is the sole, total type gate for every claim.
"""
model_config = ConfigDict(frozen=True, strict=True, extra="forbid")
iss: str
iat: int
exp: int
jti: str = Field(min_length=1)
kind: SessionTokenKind
user_id: str = Field(min_length=1)
client_id: str = Field(min_length=1)
def is_session_token(candidate: str) -> bool:
"""Cheap prefix check for a session ACCESS token so the admission edge can route gateway
sessions vs keys, JWTs, and envelopes without crypto."""
return candidate.startswith(SESSION_TOKEN_PREFIX)
def is_session_refresh_token(candidate: str) -> bool:
"""Cheap prefix check for a session REFRESH token so the token endpoint can route a
refresh grant without crypto."""
return candidate.startswith(SESSION_REFRESH_PREFIX)
def mint_session_token(
principal: SessionPrincipal,
keys: SessionKeys,
now: datetime,
) -> MintedSessionToken | SessionTokenMintError:
"""Mint the short-lived session ACCESS token for ``principal``.
``exp`` is ``SESSION_TTL_SECONDS`` from ``now``. Returns ``SessionTokenTooLarge`` when
the serialized token exceeds ``MAX_SESSION_TOKEN_BYTES``.
"""
return _mint(
kind="session",
prefix=SESSION_TOKEN_PREFIX,
principal=principal,
expires_at=now + timedelta(seconds=SESSION_TTL_SECONDS),
keys=keys,
now=now,
)
def mint_session_refresh_token(
principal: SessionPrincipal,
keys: SessionKeys,
now: datetime,
) -> MintedSessionToken | SessionTokenMintError:
"""Mint the long-lived session REFRESH token for ``principal``.
``exp`` is ``SESSION_REFRESH_TTL_SECONDS`` from ``now``. Minting a distinct
``kind="session_refresh"`` claim is what keeps a refresh token from ever opening as an
access credential at the MCP edge.
"""
return _mint(
kind="session_refresh",
prefix=SESSION_REFRESH_PREFIX,
principal=principal,
expires_at=now + timedelta(seconds=SESSION_REFRESH_TTL_SECONDS),
keys=keys,
now=now,
)
def open_session_token(
candidate: str,
keys: SessionKeys,
now: datetime,
) -> OpenedSessionToken | SessionTokenOpenError:
"""Validate a session ACCESS ``candidate`` and recover the principal.
Never raises for bad input: every invalid, expired, tampered, or wrong-kind candidate
maps to a distinct ``SessionTokenOpenError`` variant.
"""
return _open(candidate, prefix=SESSION_TOKEN_PREFIX, expected_kind="session", keys=keys, now=now)
def open_session_refresh_token(
candidate: str,
keys: SessionKeys,
now: datetime,
) -> OpenedSessionToken | SessionTokenOpenError:
"""Validate a session REFRESH ``candidate`` and recover the principal.
Total over hostile input exactly like :func:`open_session_token`. The
``kind="session_refresh"`` claim is required, so an access token re-prefixed as a
refresh one is rejected as ``SessionMalformed``.
"""
return _open(candidate, prefix=SESSION_REFRESH_PREFIX, expected_kind="session_refresh", keys=keys, now=now)
def _mint(
kind: SessionTokenKind,
prefix: str,
principal: SessionPrincipal,
expires_at: datetime,
keys: SessionKeys,
now: datetime,
) -> MintedSessionToken | SessionTokenTooLarge:
"""Sign the claims for either token kind and enforce the size cap. Shared by both mints
so the JWT shape, issuer, and size guard cannot drift between access and refresh."""
claims = _SessionClaims(
iss=SESSION_ISSUER,
iat=int(now.timestamp()),
exp=int(expires_at.timestamp()),
jti=secrets.token_urlsafe(16),
kind=kind,
user_id=principal.user_id,
client_id=principal.client_id,
)
token = prefix + jwt.encode(
claims.model_dump(), keys.signing_key.get_secret_value(), algorithm=_SESSION_JWT_ALGORITHM
)
size_bytes = len(token.encode("utf-8"))
if size_bytes > MAX_SESSION_TOKEN_BYTES:
return SessionTokenTooLarge(size_bytes=size_bytes, max_bytes=MAX_SESSION_TOKEN_BYTES)
return MintedSessionToken(token=SecretStr(token), expires_at=expires_at)
def _open(
candidate: str,
prefix: str,
expected_kind: SessionTokenKind,
keys: SessionKeys,
now: datetime,
) -> OpenedSessionToken | SessionTokenOpenError:
"""Prefix-route, size-bound, signature-verify, kind-check, and expiry-check an
attacker-controlled candidate, shared by both openers so the security gate is identical
for access and refresh. Returns the opened token or a distinct error; never raises."""
if not candidate.startswith(prefix):
return NotASessionToken()
# UTF-8 byte length is never below character length, so a character count already over
# the cap rejects an oversize candidate in O(1) without encoding it; the exact byte
# check then runs only on candidates already bounded to the cap in characters.
if len(candidate) > MAX_SESSION_TOKEN_BYTES:
return SessionMalformed()
if len(candidate.encode("utf-8", "surrogatepass")) > MAX_SESSION_TOKEN_BYTES:
return SessionMalformed()
claims = _decode_claims(candidate.removeprefix(prefix), keys.signing_key)
if not isinstance(claims, _SessionClaims):
return claims
if claims.kind != expected_kind:
return SessionMalformed()
if now.timestamp() >= claims.exp:
return SessionExpired()
return OpenedSessionToken(principal=SessionPrincipal(user_id=claims.user_id, client_id=claims.client_id))
def _decode_claims(
compact: str,
signing_key: SecretStr,
) -> _SessionClaims | SessionBadSignature | SessionMalformed:
"""Verify the HS256 signature and shape of an attacker-controlled compact JWT.
``compact`` is fully hostile and bounded to ``MAX_SESSION_TOKEN_BYTES`` by the caller.
PyJWT's ``iat``/``nbf``/``exp`` validators are disabled: they raise on hostile claim
types and, for ``iat``/``nbf``, compare against the wall clock rather than the injected
``now`` (``exp`` is checked by the caller against ``now``). Apart from a signature
mismatch, every decode failure is ``SessionMalformed``: a non-UTF-8 candidate surfaces
as ``UnicodeEncodeError`` (a ``ValueError``), a non-string registered claim as a
``TypeError`` from PyJWT's claim validators, and a wrong issuer or structurally invalid
token as an ``InvalidTokenError``. ``_SessionClaims`` is the total type gate.
"""
try:
payload = jwt.decode(
compact,
signing_key.get_secret_value(),
algorithms=[_SESSION_JWT_ALGORITHM],
issuer=SESSION_ISSUER,
options={
"verify_exp": False,
"verify_iat": False,
"verify_nbf": False,
"require": ["iss", "iat", "exp"],
},
)
except jwt.InvalidSignatureError:
return SessionBadSignature()
except (jwt.InvalidTokenError, ValueError, TypeError):
return SessionMalformed()
try:
return _SessionClaims.model_validate(payload)
except ValidationError:
return SessionMalformed()

View file

@ -184,7 +184,12 @@ class ClientCredentialsConfig(BaseModel):
Fields are optional so the config can be built incomplete: a value may be supplied at
runtime (`token_url` via RFC 8414 discovery, `client_id`/`secret` via DCR), and the
resolver arm raises `CredError.misconfigured` when a needed field is still absent.
resolver arm returns `CredError.misconfigured` when a needed field is still absent.
`audience` is the IdP-specific audience parameter some authorization servers require on
the client_credentials grant (sent as `audience` in the token request when set).
`token_endpoint_auth_method` selects how the client authenticates to the token endpoint
(RFC 6749 section 2.3.1); `None` defaults to `client_secret_post`.
"""
model_config = ConfigDict(frozen=True)
@ -193,6 +198,8 @@ class ClientCredentialsConfig(BaseModel):
client_secret: SecretStr | None = None
token_url: str | None = None
scopes: tuple[str, ...] = ()
audience: str | None = None
token_endpoint_auth_method: Literal["client_secret_post", "client_secret_basic"] | None = None
class TokenExchangeConfig(BaseModel):

View file

@ -273,6 +273,7 @@ _BANNED_REQUEST_BODY_PARAMS: Tuple[str, ...] = (
# re-route the request's retention and accounting to any project
# reachable with the deployment's shared AWS credentials.
"aws_bedrock_project_id",
"bedrock_tags",
# Provider-specific endpoint overrides that flow into the outbound
# request via ``optional_params``. Same threat as ``api_base``:
# ``s3_endpoint_url`` redirects Bedrock file uploads to attacker

View file

@ -1647,7 +1647,7 @@ async def ui_view_spend_logs(
description="Time till which to view key spend",
),
page: int = fastapi.Query(default=1, description="Page number for pagination", ge=1),
page_size: int = fastapi.Query(default=50, description="Number of items per page", ge=1, le=100),
page_size: int = fastapi.Query(default=50, description="Number of items per page", ge=1, le=1000),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
status_filter: str | None = fastapi.Query(
default=None, description="Filter logs by status (e.g., success, failure)"

View file

@ -12,7 +12,7 @@ from typing import (
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
LiteLLM_Proxy_MCP_Handler,
)
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.responses.mcp.request_context import MCPRequestContext
from litellm.types.utils import ModelResponse
from litellm.utils import CustomStreamWrapper
@ -114,20 +114,13 @@ async def acompletion_with_mcp(
**kwargs,
)
# Extract user_api_key_auth from metadata or kwargs
user_api_key_auth = kwargs.get("user_api_key_auth") or ((kwargs.get("metadata", {}) or {}).get("user_api_key_auth"))
request_tags = LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(kwargs)
# Extract MCP auth headers before fetching tools (needed for dynamic auth)
(
mcp_auth_header,
mcp_server_auth_headers,
oauth2_headers,
raw_headers,
) = ResponsesAPIRequestUtils.extract_mcp_headers_from_request(
secret_fields=kwargs.get("secret_fields"),
tools=tools,
)
context = MCPRequestContext.resolve(kwargs=kwargs, tools=tools)
user_api_key_auth = context.user_api_key_auth
request_tags = list(context.request_tags) if context.request_tags else None
mcp_auth_header = context.mcp_auth_header
mcp_server_auth_headers = context.mcp_server_auth_headers
oauth2_headers = context.oauth2_headers
raw_headers = context.raw_headers
# Process MCP tools (pass auth headers for dynamic auth)
(

View file

@ -479,17 +479,25 @@ class LiteLLM_Proxy_MCP_Handler:
) -> bool:
"""Check if we should auto-execute tool calls.
Only auto-execute tools if user passed a MCP tool with require_approval set to "never".
Auto-execution requires EVERY MCP reference to opt in with
``require_approval="never"``. A single reference that requires approval
("always", "manual", the object form, or an unset value) disables
auto-execution for the whole request. This fails closed: when an
approval-required reference shares a request with a "never" one, the
model's tool calls are returned to the caller instead of being run, so
an approval-gated tool can never be invoked without approval. Returns
False for an empty list.
"""
for tool in mcp_tools_with_litellm_proxy:
if isinstance(tool, dict):
if tool.get("require_approval") == "never":
return True
elif getattr(tool, "require_approval", None) == "never":
return True
return False
references = list(mcp_tools_with_litellm_proxy or [])
if not references:
return False
for tool in references:
approval = (
tool.get("require_approval") if isinstance(tool, dict) else getattr(tool, "require_approval", None)
)
if approval != "never":
return False
return True
@staticmethod
def _extract_tool_calls_from_response(response: ResponsesAPIResponse) -> List[Any]:
@ -542,7 +550,10 @@ class LiteLLM_Proxy_MCP_Handler:
tool_arguments = function_block.get("arguments")
else:
tool_name = tool_call.get("name")
# Anthropic tool_use blocks carry the arguments under `input`
tool_arguments = tool_call.get("arguments")
if tool_arguments is None:
tool_arguments = tool_call.get("input")
else:
tool_call_id = getattr(tool_call, "call_id", None) or getattr(tool_call, "id", None)
@ -553,6 +564,8 @@ class LiteLLM_Proxy_MCP_Handler:
else:
tool_name = getattr(tool_call, "name", None)
tool_arguments = getattr(tool_call, "arguments", None)
if tool_arguments is None:
tool_arguments = getattr(tool_call, "input", None)
return tool_name, tool_arguments, tool_call_id

View file

@ -0,0 +1,73 @@
"""
The per-request context an MCP gateway handler needs.
Listing and executing MCP tools both need the caller's identity, their MCP auth
headers, and the request's trace/tag identifiers. Every gateway surface resolves
the same set from its own kwargs, so resolving it in one place keeps a new
surface from silently dropping a field: omitting the auth headers, for instance,
still executes the tool, just with no credentials.
"""
from dataclasses import dataclass
from typing import Any, Iterable, Mapping, Sequence, Union
@dataclass(frozen=True, slots=True)
class MCPRequestContext:
"""Everything a gateway handler must forward to MCP tool listing and execution."""
user_api_key_auth: Any # any-ok: UserAPIKeyAuth is proxy-only; importing it here would create a cycle
mcp_auth_header: Union[str, None] = None
mcp_server_auth_headers: Union[Mapping[str, Mapping[str, str]], None] = None
oauth2_headers: Union[Mapping[str, str], None] = None
raw_headers: Union[Mapping[str, str], None] = None
request_tags: Union[Sequence[str], None] = None
litellm_trace_id: Union[str, None] = None
litellm_call_id: Union[str, None] = None
@classmethod
def resolve(
cls,
kwargs: Mapping[str, Any],
tools: Union[Iterable[Any], None],
) -> "MCPRequestContext":
"""
Build the context from a gateway handler's kwargs.
``user_api_key_auth`` is read from both metadata keys because routes differ:
LITELLM_METADATA_ROUTES (``/v1/messages``, ``/responses``) carry it in
``litellm_metadata`` while ``/chat/completions`` uses ``metadata``.
"""
from litellm.responses.mcp.litellm_proxy_mcp_handler import (
LiteLLM_Proxy_MCP_Handler,
)
from litellm.responses.utils import ResponsesAPIRequestUtils
litellm_metadata = kwargs.get("litellm_metadata") or {}
metadata = kwargs.get("metadata") or {}
user_api_key_auth = (
kwargs.get("user_api_key_auth")
or litellm_metadata.get("user_api_key_auth")
or metadata.get("user_api_key_auth")
)
(
mcp_auth_header,
mcp_server_auth_headers,
oauth2_headers,
raw_headers,
) = ResponsesAPIRequestUtils.extract_mcp_headers_from_request(
secret_fields=kwargs.get("secret_fields"),
tools=tools,
)
return cls(
user_api_key_auth=user_api_key_auth,
mcp_auth_header=mcp_auth_header,
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
request_tags=LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(dict(kwargs)),
litellm_trace_id=kwargs.get("litellm_trace_id"),
litellm_call_id=kwargs.get("litellm_call_id"),
)

View file

@ -72,17 +72,28 @@ def use_litellm_rust(
messages: RustMessages | None | _Unset = _UNSET,
amessages: RustAmessages | None | _Unset = _UNSET,
responses_websocket: Any | None | _Unset = _UNSET,
transcription: Any | None | _Unset = _UNSET,
atranscription: Any | None | _Unset = _UNSET,
) -> None:
global _rust_ocr_enabled, _rust_ocr_impl, _rust_aocr_impl
configuring_ocr = not isinstance(ocr, _Unset) or not isinstance(aocr, _Unset)
configuring_messages = not isinstance(messages, _Unset) or not isinstance(amessages, _Unset)
configuring_responses_websocket = not isinstance(responses_websocket, _Unset)
configuring_transcription = not isinstance(transcription, _Unset) or not isinstance(atranscription, _Unset)
if configuring_ocr or (not configuring_messages and not configuring_responses_websocket):
_rust_ocr_enabled = enabled
if not isinstance(ocr, _Unset):
_rust_ocr_impl = ocr
if not isinstance(aocr, _Unset):
_rust_aocr_impl = aocr
if configuring_transcription:
from litellm.rust_bridge.transcription import configure_rust_transcription
configure_rust_transcription(
enabled=enabled,
transcription=transcription,
atranscription=atranscription,
)
if not configuring_messages and not configuring_responses_websocket:
return
if configuring_messages:

View file

@ -0,0 +1,148 @@
from __future__ import annotations
from dataclasses import dataclass
from typing import Awaitable, Final, Protocol, Union, cast
import httpx
from litellm.rust_bridge.timeouts import timeout_to_seconds
class RustTranscription(Protocol):
def __call__(
self,
model: str,
audio: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> dict[str, object]:
raise NotImplementedError
class RustAtranscription(Protocol):
def __call__(
self,
model: str,
audio: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> Awaitable[dict[str, object]]:
raise NotImplementedError
class _Unset:
pass
_UNSET: Final[_Unset] = _Unset()
@dataclass
class _RustTranscriptionState:
transcription: RustTranscription | None = None
atranscription: RustAtranscription | None = None
_STATE = _RustTranscriptionState()
def configure_rust_transcription(
enabled: bool = True,
*,
transcription: RustTranscription | None | _Unset = _UNSET,
atranscription: RustAtranscription | None | _Unset = _UNSET,
) -> None:
if not isinstance(transcription, _Unset):
_STATE.transcription = transcription
if not isinstance(atranscription, _Unset):
_STATE.atranscription = atranscription
def load_rust_transcription() -> RustTranscription | None:
if _STATE.transcription is not None:
return _STATE.transcription
from litellm.rust_bridge import get_native_bridge
native_bridge = get_native_bridge()
return (
None
if native_bridge is None
else cast( # cast-ok: native extension protocol is runtime-defined
RustTranscription, getattr(native_bridge, "transcription", None)
)
)
def load_rust_atranscription() -> RustAtranscription | None:
if _STATE.atranscription is not None:
return _STATE.atranscription
from litellm.rust_bridge import get_native_bridge
native_bridge = get_native_bridge()
return (
None
if native_bridge is None
else cast( # cast-ok: native extension protocol is runtime-defined
RustAtranscription, getattr(native_bridge, "atranscription", None)
)
)
def transcription(
*,
model: str,
audio: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: Union[float, httpx.Timeout] | None,
) -> dict[str, object] | None:
rust_transcription = load_rust_transcription()
if rust_transcription is None:
return None
return rust_transcription(
model=model,
audio=audio,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout_seconds=timeout_to_seconds(timeout),
)
async def atranscription(
*,
model: str,
audio: dict[str, object],
api_key: str | None,
api_base: str | None,
custom_llm_provider: str | None,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout: Union[float, httpx.Timeout] | None,
) -> dict[str, object] | None:
rust_atranscription = load_rust_atranscription()
if rust_atranscription is None:
return None
return await rust_atranscription(
model=model,
audio=audio,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=optional_params,
timeout_seconds=timeout_to_seconds(timeout),
)

View file

@ -1,6 +1,6 @@
from typing import Literal, Optional, Union
from typing_extensions import TypedDict
from typing_extensions import NotRequired, TypedDict
from litellm.types.llms.openai import ChatCompletionCachedContent
@ -12,6 +12,7 @@ class CacheControlMessageInjectionPoint(TypedDict):
role: Optional[Literal["user", "system", "assistant"]] # Optional: target by role (user, system, assistant)
index: Optional[Union[int, str]] # Optional: target by specific index
control: Optional[ChatCompletionCachedContent]
_litellm_judged: NotRequired[bool] # Internal: written back by litellm once the client cache_control judgment ran
class CacheControlToolConfigInjectionPoint(TypedDict):
@ -19,6 +20,7 @@ class CacheControlToolConfigInjectionPoint(TypedDict):
location: Literal["tool_config"]
control: Optional[ChatCompletionCachedContent]
_litellm_judged: NotRequired[bool] # Internal: written back by litellm once the client cache_control judgment ran
CacheControlInjectionPoint = Union[

View file

@ -985,6 +985,11 @@ class BedrockOutputDataConfig(TypedDict):
s3OutputDataConfig: BedrockS3OutputDataConfig
class BedrockTag(TypedDict):
key: str
value: str
class BedrockCreateBatchRequest(TypedDict, total=False):
"""
Request structure for creating a Bedrock batch inference job.
@ -999,7 +1004,7 @@ class BedrockCreateBatchRequest(TypedDict, total=False):
outputDataConfig: BedrockOutputDataConfig
timeoutDurationInHours: Optional[int]
clientRequestToken: Optional[str]
tags: Optional[List[dict]]
tags: Optional[List[BedrockTag]]
BedrockBatchJobStatus = Literal["Submitted", "InProgress", "Completed", "Failed", "Stopping", "Stopped"]

View file

@ -23,12 +23,16 @@ from access_control_client import (
ROUTE_NOT_ALLOWED_MARKER,
)
from e2e_config import unique_marker
from e2e_http import Success, UnauthorizedError, UnknownApiError, unwrap
from lifecycle import ResourceManager
from models import ChatBody, ChatMessage, LiteLLMParamsBody
from proxy_client import ProxyClient
pytestmark = pytest.mark.e2e
ALLOWED_MODEL = "gemini-2.5-flash"
DISALLOWED_MODEL = "gpt-5.5"
VIRTUAL_KEY_BACKEND = "anthropic/claude-haiku-4-5-20251001"
def _is_json(body: str) -> bool:
@ -39,6 +43,7 @@ def _is_json(body: str) -> bool:
return False
class TestAccessControl:
def test_disallowed_model_is_denied_403(
self, client: AccessControlClient, resources: ResourceManager
@ -81,3 +86,59 @@ class TestAccessControl:
f"{result.status_code}: {result.body[:300]}"
)
assert _is_json(result.body), f"400 body must be valid JSON: {result.body[:300]}"
class TestVirtualKeyAuth:
"""Virtual-key auth the way OpenAI-compatible clients send it: a real key
must reach chat, a forged bearer must be rejected before the provider."""
@pytest.mark.covers(
"mgmt.virtual_key.valid_allows",
"mgmt.virtual_key.invalid_denied",
exercised_on=[],
)
def test_valid_key_allows_and_invalid_key_denied(
self, proxy: ProxyClient, resources: ResourceManager
) -> None:
model = f"e2e-auth-chat-{unique_marker()}"
model_id = proxy.create_model(
model,
LiteLLMParamsBody(model=VIRTUAL_KEY_BACKEND, api_key="os.environ/ANTHROPIC_API_KEY"),
)
resources.defer(lambda: proxy.delete_model(model_id))
key = resources.key()
ok = unwrap(
proxy.chat(
key,
ChatBody(
model=model,
messages=[
ChatMessage(
role="user",
content=f"Reply with one word. {unique_marker()}",
)
],
max_tokens=16,
),
)
)
assert ok.choices, f"valid key must complete chat: {ok}"
bad = proxy.chat(
"sk-e2e-forged-not-a-real-key",
ChatBody(
model=model,
messages=[ChatMessage(role="user", content="should not run")],
max_tokens=8,
),
)
match bad:
case UnauthorizedError():
return
case UnknownApiError(status_code=status) if status in (401, 403):
return
case Success():
pytest.fail("forged bearer must not reach a successful completion")
case _:
pytest.fail(f"forged bearer must be auth-denied, got {bad}")

View file

@ -16,13 +16,14 @@ misroute to the wrong provider fails the create.
from __future__ import annotations
import json
import os
import time
from datetime import datetime, timedelta, timezone
from typing import Callable
import pytest
from e2e_config import unique_marker
from e2e_config import require_env, unique_marker
from batch_client import (
BatchClient,
@ -39,7 +40,9 @@ from capabilities import (
FILE_ID_SHAPE,
OPENAI_BATCH_MODEL,
Capability,
batch_model_name,
coverage_cells_for_lifecycle,
is_managed_id,
matches_id_shape,
raw_id_matches_provider,
)
@ -53,7 +56,7 @@ from e2e_http import (
unwrap,
)
from lifecycle import ResourceManager
from models import KeyGenerateBody, SpendLogRow
from models import KeyGenerateBody, LiteLLMParamsBody, SpendLogRow
pytestmark = pytest.mark.e2e
@ -457,3 +460,328 @@ def test_rate_limited_batch_create_leaves_no_unattributed_spend_row(
"batch create on a rate-limited key left an unattributed spend row "
f"(LIT-3266); rows={[(r.request_id, r.call_type, r.model) for r in new_orphans]}"
)
OPENAI_FILE_CONTENT_BACKEND = "gpt-4o-mini"
class TestBatchFileContent:
"""GET /v1/files/{id}/content returns the uploaded batch JSONL bytes."""
@pytest.mark.covers(
"llm.files.openai.content.nonstream.works",
exercised_on=["files"],
)
def test_file_content_matches_upload(
self, client: BatchClient, resources: ResourceManager
) -> None:
proxy_name = f"e2e-file-content-{unique_marker()}"
model_id = client.create_model(
proxy_name,
LiteLLMParamsBody(
model=f"openai/{OPENAI_FILE_CONTENT_BACKEND}",
api_key="os.environ/OPENAI_API_KEY",
),
)
resources.defer(lambda: client.delete_model(model_id))
key = resources.key()
payload = render_jsonl(OPENAI_FILE_CONTENT_BACKEND)
file = unwrap(
client.upload_file(
content=payload,
form=FileUploadForm(purpose="batch", target_model_names=proxy_name),
key=key,
)
)
resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
assert file.id
downloaded = client.proxy.transport.download(
f"/v1/files/{file.id}/content",
headers=client.proxy.transport.bearer(key),
)
assert downloaded.status_code == 200, (
f"file content must be 200, got {downloaded.status_code}: {downloaded.body[:300]}"
)
expected = payload.decode().rstrip("\n")
got = downloaded.body.rstrip("\n")
assert got == expected, (
"downloaded file content must match the uploaded JSONL bytes"
)
BATCH_RL_REQUEST_LINES = 3
BATCH_RL_RPM_LIMIT = 2
def _multi_request_jsonl(model: str, n: int) -> bytes:
lines = tuple(
json.dumps(
{
"custom_id": f"req-{i}",
"method": "POST",
"url": "/v1/chat/completions",
"body": {
"model": model,
"messages": [{"role": "user", "content": "ping"}],
"max_tokens": 8,
},
}
)
for i in range(n)
)
return ("\n".join(lines) + "\n").encode()
class TestBatchRateLimitErrorMapping:
"""Batch create that exceeds a key's RPM maps to a structured 429.
The batch rate limiter reads the input file at submission time and rejects
the create when the file's request count would exceed the key's remaining
RPM. The product promise is not only the block itself but the
OpenAI-compatible shape: HTTP 429, a body that names the batch rate limit,
and pacing headers so clients can back off. Complements the LIT-3266 hygiene
check (no orphan spend rows) by asserting the error mapping when the limiter
actually fires.
"""
@pytest.mark.covers(
"quota_management.ratelimit.batch_rpm.blocks_over_limit",
exercised_on=["batches"],
)
def test_batch_create_over_rpm_returns_mapped_429(
self, client: BatchClient, resources: ResourceManager, batch_deployments: None
) -> None:
user_id = f"e2e-batch-rl-map-{unique_marker()}"
key = client.proxy.generate_key(
KeyGenerateBody(
models=[], rpm_limit=BATCH_RL_RPM_LIMIT, tpm_limit=1_000_000, user_id=user_id
)
)
resources.defer(lambda: client.proxy.delete_key(key))
file = unwrap(
client.upload_file(
content=_multi_request_jsonl("gpt-4o-mini", BATCH_RL_REQUEST_LINES),
form=FileUploadForm(purpose="batch"),
model=OPENAI_BATCH_MODEL,
key=key,
)
)
resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
assert created.status_code == 429, (
f"expected batch RPM 429 when file has {BATCH_RL_REQUEST_LINES} requests and "
f"rpm_limit={BATCH_RL_RPM_LIMIT}, got {created.status_code}: {created.body[:400]}"
)
body_lower = created.body.lower()
assert "batch rate limit exceeded" in body_lower, (
f"429 body must name the batch rate limit so clients can branch on it; "
f"got: {created.body[:400]}"
)
assert str(BATCH_RL_REQUEST_LINES) in created.body, (
f"429 body should report the batch request count ({BATCH_RL_REQUEST_LINES}); "
f"got: {created.body[:400]}"
)
assert "rpm" in body_lower or "requests remaining" in body_lower, (
f"429 body must describe the RPM budget remaining so clients can pace; "
f"got: {created.body[:400]}"
)
retry_after = created.headers.get("retry-after")
if retry_after is not None:
assert retry_after.isdigit() and int(retry_after) > 0, (
f"retry-after must be a positive integer when present, got {retry_after!r}"
)
ASSUME_ROLE_RAW_MODEL = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0"
def _assume_role_params(role_arn: str, session_name: str) -> LiteLLMParamsBody:
return LiteLLMParamsBody(
model=ASSUME_ROLE_RAW_MODEL,
aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID",
aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY",
aws_region_name="os.environ/AWS_REGION",
s3_region_name="os.environ/AWS_REGION",
s3_bucket_name="os.environ/AWS_BATCH_S3_BUCKET",
s3_access_key_id="os.environ/AWS_ACCESS_KEY_ID",
s3_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY",
aws_batch_role_arn="os.environ/AWS_BATCH_ROLE_ARN",
aws_role_name=role_arn,
aws_session_name=session_name,
)
class TestBedrockBatchAssumeRole:
"""Bedrock batch create under STS assume-role credentials.
Provisions a bedrock batch deployment whose litellm_params carry
aws_role_name / aws_session_name (the product path for role assumption) and
runs the unified file-upload + batch-create lifecycle. Success means the
proxy assumed the role and Bedrock accepted the job; a misconfigured role
fails create with an AWS auth error rather than silently falling back to the
ambient key.
"""
@pytest.mark.covers(
"llm.batches.bedrock.assume_role.nonstream.works",
"llm.files.bedrock.upload.nonstream.works",
exercised_on=["batches", "files"],
)
def test_unified_batch_create_with_assume_role(
self, client: BatchClient, resources: ResourceManager
) -> None:
(role_arn,) = require_env("AWS_ROLE_NAME")
require_env(
"AWS_ACCESS_KEY_ID",
"AWS_SECRET_ACCESS_KEY",
"AWS_REGION",
"AWS_BATCH_S3_BUCKET",
"AWS_BATCH_ROLE_ARN",
)
session_name = f"e2e-batch-sts-{unique_marker()}"[:64]
model_name = batch_model_name("bedrock-sts-batch")
model_id = client.create_model(model_name, _assume_role_params(role_arn, session_name))
resources.defer(lambda: client.delete_model(model_id))
key = resources.key()
file = unwrap(
client.upload_file(
content=render_jsonl(ASSUME_ROLE_RAW_MODEL),
form=FileUploadForm(purpose="batch", target_model_names=model_name),
key=key,
)
)
resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
assert_file_object(file, provider="bedrock")
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
require_successful_call(created)
batch = BatchObject.model_validate_json(created.body)
resources.defer(quietly(lambda: client.cancel_batch(batch.id, key=key)))
assert batch.id, f"assume-role create returned no batch id: {created.body[:200]}"
assert is_managed_id(batch.id), (
f"assume-role create via target_model_names must return a managed batch id, "
f"got {batch.id!r}"
)
assert batch.status in CREATED_BATCH_STATUSES, (
f"assume-role batch has non-transitional status {batch.status!r}"
)
assert_batch_object(batch)
fetched = unwrap(client.retrieve_batch(batch.id, key=key))
assert fetched.id == batch.id
GEMINI_FILES_RAW_MODEL = "gemini-2.5-flash"
class TestGeminiFiles:
"""Gemini Files API upload through the proxy (LIT-3382).
gemini is a first-class FileCreateProvider. The test registers a gemini
deployment, uploads a tiny batch-purpose JSONL with target_model_names
routing, and asserts a FileObject comes back. Batch create for pure gemini
(non-Vertex) is out of scope here; Vertex covers the Gemini batch job path in
the main lifecycle matrix.
"""
@pytest.mark.covers(
"llm.files.gemini.upload.nonstream.works",
exercised_on=["files"],
)
def test_gemini_file_upload(
self, client: BatchClient, resources: ResourceManager
) -> None:
model_name = batch_model_name("gemini-files")
model_id = client.create_model(
model_name,
LiteLLMParamsBody(
model=f"gemini/{GEMINI_FILES_RAW_MODEL}",
api_key="os.environ/GEMINI_API_KEY",
),
)
resources.defer(lambda: client.delete_model(model_id))
key = resources.key()
file = unwrap(
client.upload_file(
content=render_jsonl(GEMINI_FILES_RAW_MODEL),
form=FileUploadForm(purpose="batch", target_model_names=model_name),
key=key,
)
)
resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
assert_file_object(file, provider="gemini")
assert file.id, "gemini file upload returned no id"
def _vllm_params(api_base: str, api_key: str | None, model_id: str) -> LiteLLMParamsBody:
return LiteLLMParamsBody(
model=f"hosted_vllm/{model_id}",
api_base=api_base,
api_key=api_key,
)
class TestHostedVllmBatch:
"""hosted_vllm file upload + batch create (OpenAI-compatible path, LIT-3266).
hosted_vllm is in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS, so /v1/files
and /v1/batches route through the OpenAI handler against the deployment's
api_base. Skipped for now: it needs a live vLLM (or OpenAI-compatible) server
exposing the files/batches APIs (HOSTED_VLLM_API_BASE), which the e2e
environment does not currently provision.
"""
@pytest.mark.skip(
reason="hosted_vllm batch/files needs a live vLLM server (HOSTED_VLLM_API_BASE) "
"not provisioned in the e2e environment; re-enable when available (LIT-3266)"
)
@pytest.mark.covers(
"llm.batches.hosted_vllm.basic.nonstream.works",
"llm.files.hosted_vllm.upload.nonstream.works",
exercised_on=["batches", "files"],
)
def test_unified_file_and_batch_create(
self, client: BatchClient, resources: ResourceManager
) -> None:
(api_base,) = require_env("HOSTED_VLLM_API_BASE")
api_key = (os.environ.get("HOSTED_VLLM_API_KEY") or "").strip() or None
model_id = (
os.environ.get("HOSTED_VLLM_MODEL") or "meta-llama/Llama-3.2-3B-Instruct"
).strip()
proxy_name = batch_model_name("hosted-vllm-batch")
model_row_id = client.create_model(
proxy_name, _vllm_params(api_base, api_key, model_id)
)
resources.defer(lambda: client.delete_model(model_row_id))
key = resources.key()
file = unwrap(
client.upload_file(
content=render_jsonl(model_id),
form=FileUploadForm(purpose="batch", target_model_names=proxy_name),
key=key,
)
)
resources.defer(quietly(lambda: client.delete_file(file.id, key=key)))
assert_file_object(file, provider="hosted_vllm")
created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key)
require_successful_call(created)
batch = BatchObject.model_validate_json(created.body)
resources.defer(quietly(lambda: client.cancel_batch(batch.id, key=key)))
assert batch.id, f"hosted_vllm create returned no batch id: {created.body[:200]}"
assert batch.status in CREATED_BATCH_STATUSES, (
f"hosted_vllm batch has non-transitional status {batch.status!r}"
)
assert_batch_object(batch)

View file

@ -14,14 +14,14 @@ shared fixtures build on it.
"""
import functools
import sys
import os
from collections.abc import Iterator
from pathlib import Path
import pytest
import requests
from e2e_config import CONTROL_PLANE_BASE_URL, PROXY_BASE_URL
from e2e_db import RESET_OPT_IN_ENV, reset_spend_logs, run_spend_log_cleanup
from junit_properties import attach_result_properties
from lifecycle import ProxyClientProvider, ResourceManager
from proxy_client import ProxyClient, build_proxy_client
@ -107,26 +107,17 @@ def pytest_runtest_call(item: pytest.Item) -> None:
def pytest_sessionfinish(session: pytest.Session, exitstatus: int) -> None:
"""Once the whole e2e session is done (all suites), truncate the spend logs so
the DB doesn't accumulate test rows. Sessions where no e2e test body ran leave
the DB alone so a `DATABASE_URL` pointing at a shared instance is never wiped
without an e2e run. Best-effort: a cleanup failure (no DB reachable) must not
fail the run. The spend_tracking dir goes on sys.path only for this import and
is removed after, so a broader `pytest tests/` run is not left with a mutated
path."""
if not session.stash.get(_E2E_TEST_RAN, False):
return
spend_dir = str(Path(__file__).parent / "quota_management" / "spend_tracking")
sys.path.insert(0, spend_dir)
try:
from spend_e2e_client import reset_spend_logs # pyright: ignore
reset_spend_logs()
except Exception as exc: # noqa: BLE001 - cleanup is best-effort
print(f"spend-log cleanup best-effort failed: {exc}")
finally:
if spend_dir in sys.path:
sys.path.remove(spend_dir)
"""Once the whole e2e session is done (all suites), optionally truncate the
spend logs so the DB doesn't accumulate test rows. The truncate is destructive
and irreversible, so it runs only when the operator explicitly opts in
(`E2E_RESET_SPEND_LOGS=1`) and an e2e test body actually ran; otherwise a
`DATABASE_URL` pointing at a shared or staging instance is left untouched.
Best-effort: a cleanup failure (no DB reachable) must not fail the run."""
run_spend_log_cleanup(
opt_in=os.environ.get(RESET_OPT_IN_ENV),
e2e_test_ran=session.stash.get(_E2E_TEST_RAN, False),
truncate=reset_spend_logs,
)
@pytest.fixture(scope="session")

View file

@ -4,6 +4,10 @@
- {id: guardrail.presidio.post_call.masks, module: guardrail, tier: P0, hook_point: post_call, assertions: [masks], exercised_on: [chat_completions, messages], source: "guardrail_hooks/presidio.py", rationale: "Mask PII in model output"}
- {id: guardrail.presidio.logging_only.masks, module: guardrail, tier: P0, hook_point: logging_only, assertions: [masks], exercised_on: [chat_completions, messages], source: "guardrail_hooks/presidio.py", rationale: "Redact in logs without blocking"}
- {id: guardrail.bedrock.pre_call.blocks, module: guardrail, tier: P0, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/bedrock_guardrails.py", rationale: "AWS content guardrail blocks harmful input"}
- {id: guardrail.litellm_content_filter.pre_call.blocks, module: guardrail, tier: P0, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "test_team_disable_global_guardrail_e2e.py", rationale: "Local content-filter default-on blocks banned keyword pre-call"}
- {id: guardrail.litellm_content_filter.pre_call.allows, module: guardrail, tier: P0, hook_point: pre_call, assertions: [allows], exercised_on: [chat_completions], source: "test_team_disable_global_guardrail_e2e.py", rationale: "Team disable_global_guardrails bypasses default-on content filter"}
- {id: guardrail.litellm_content_filter.apply_endpoint.blocks, module: guardrail, tier: P0, hook_point: apply_endpoint, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_endpoints.py:apply_guardrail", rationale: "POST /guardrails/apply_guardrail blocks banned content for customers that call the apply surface directly"}
- {id: guardrail.litellm_content_filter.apply_endpoint.allows, module: guardrail, tier: P0, hook_point: apply_endpoint, assertions: [allows], exercised_on: [chat_completions], source: "guardrail_endpoints.py:apply_guardrail", rationale: "POST /guardrails/apply_guardrail returns clean text for allowed input"}
- {id: guardrail.bedrock.during.blocks, module: guardrail, tier: P0, hook_point: during, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/bedrock_guardrails.py", rationale: "During-call moderation for streaming"}
- {id: guardrail.bedrock.post_call.blocks, module: guardrail, tier: P0, hook_point: post_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/bedrock_guardrails.py", rationale: "Block harmful output"}
- {id: guardrail.lakera.pre_call.blocks, module: guardrail, tier: P0, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions, messages], source: "guardrail_hooks/lakera_ai_v2.py", rationale: "Prompt-injection block pre-execution"}

View file

@ -24,6 +24,11 @@
- {id: llm.chat_completions.bedrock_converse.prompt_cache_5m.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: bedrock_converse, capability: prompt_cache_5m, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Anthropic-on-Bedrock caching"}
- {id: llm.chat_completions.bedrock_converse.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: bedrock_converse, capability: thinking, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Anthropic thinking on Bedrock"}
- {id: llm.chat_completions.vertex.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "proxy_server.py:8455", rationale: "P0 route; Vertex AI"}
- {id: llm.chat_completions.gemini.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: gemini, capability: basic, streaming: nonstream, assertions: [works], source: "test_chat_completions_regression_e2e.py", rationale: "Gemini OpenAI-compatible chat translation"}
- {id: llm.chat_completions.gemini.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: chat_completions, route: gemini, capability: basic, streaming: nonstream, assertions: [works, cost_logged], source: "test_chat_completions_regression_e2e.py", rationale: "Gemini chat cost lands in SpendLogs"}
- {id: llm.chat_completions.hosted_vllm.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: hosted_vllm, capability: basic, streaming: nonstream, assertions: [works], source: "test_chat_completions_regression_e2e.py", rationale: "OpenAI-compatible hosted_vllm chat is a confirmed self-hosted backend path"}
- {id: llm.chat_completions.cohere.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: cohere, capability: basic, streaming: nonstream, assertions: [works], source: "test_chat_completions_regression_e2e.py", rationale: "Cohere chat via OpenAI-compatible /chat/completions"}
- {id: llm.chat_completions.vertex.basic.stream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: vertex, capability: basic, streaming: stream, assertions: [works], source: "proxy_server.py:8455", rationale: "Streaming over Vertex"}
- {id: llm.chat_completions.vertex.tool_use.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: vertex, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Vertex Gemini function_calling"}
- {id: llm.chat_completions.vertex.vision.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: vertex, capability: vision, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Gemini vision"}

View file

@ -18,6 +18,8 @@
- {id: llm.batches.azure_openai.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:98", rationale: "Azure batches all scenarios"}
- {id: llm.batches.vertex.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:98", rationale: "Vertex batches"}
- {id: llm.batches.bedrock.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:98", rationale: "Bedrock batches (encoded/unified only)"}
- {id: llm.batches.bedrock.assume_role.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: bedrock_converse, capability: assume_role, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Bedrock batch create under STS assume-role credentials"}
- {id: llm.batches.hosted_vllm.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: batches, route: hosted_vllm, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "hosted_vllm OpenAI-compatible batch create"}
- {id: llm.batches.openai.key_model_access_denied.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Key model restriction 403 on upload/create"}
- {id: llm.files.openai.upload.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "openai_files_endpoints/files_endpoints.py:46", rationale: "File upload returns OpenAIFileObject"}
- {id: llm.files.openai.retrieve.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "files_endpoints.py", rationale: "File retrieve by id"}
@ -26,7 +28,11 @@
- {id: llm.files.azure_openai.upload.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:45", rationale: "Azure file upload managed backend"}
- {id: llm.files.vertex.upload.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:52", rationale: "Vertex file upload to GCS"}
- {id: llm.files.bedrock.upload.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:59", rationale: "Bedrock file upload to S3"}
- {id: llm.files.gemini.upload.nonstream.works, module: llm, tier: P1, subject_endpoint: files, route: gemini, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Gemini Files API upload via proxy"}
- {id: llm.files.hosted_vllm.upload.nonstream.works, module: llm, tier: P1, subject_endpoint: files, route: hosted_vllm, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "hosted_vllm OpenAI-compatible file upload"}
- {id: llm.rerank.cohere.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: rerank, route: cohere, capability: basic, streaming: nonstream, assertions: [works], source: "test_rerank_e2e.py:29", rationale: "Cohere rerank, top_n + relevance_score"}
- {id: llm.files.openai.content.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "GET /v1/files/{id}/content returns uploaded batch JSONL bytes"}
- {id: llm.realtime.bedrock_converse.basic.stream.works, module: llm, tier: P0, subject_endpoint: realtime, route: bedrock_converse, capability: basic, streaming: stream, assertions: [works], source: "test_nova_sonic_realtime_e2e.py", rationale: "Nova Sonic realtime session emits response.done (LIT-2239)"}
- {id: llm.rerank.bedrock.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: rerank, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "llms/bedrock/rerank/handler.py", rationale: "Bedrock rerank"}
- {id: llm.rerank.together_ai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: rerank, route: together_ai, capability: basic, streaming: nonstream, assertions: [works], source: "llms/together_ai/rerank/handler.py", rationale: "Together rerank"}
- {id: llm.images_generations.openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: images_generations, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_image_generation_e2e.py:22", rationale: "OpenAI image gen, b64/url"}

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