Merge branch 'main' into litellm_mcp_continuous_tool_defaults

Co-Authored-By: bot_apk <apk@cognition.ai>
This commit is contained in:
Devin AI 2026-09-25 04:41:54 +00:00
commit d816dc8fbc
86 changed files with 5042 additions and 1478 deletions

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-enterprise"
version = "0.1.70"
version = "0.1.71"
description = "Package for LiteLLM Enterprise features"
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.1.70"
version = "0.1.71"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-enterprise==",

View file

@ -1,6 +1,6 @@
[project]
name = "litellm-proxy-extras"
version = "0.4.101"
version = "0.4.102"
description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package."
readme = "README.md"
requires-python = ">=3.9"
@ -26,7 +26,7 @@ required-version = ">=0.10.9"
module-root = ""
[tool.commitizen]
version = "0.4.101"
version = "0.4.102"
version_files = [
"pyproject.toml:^version",
"../pyproject.toml:litellm-proxy-extras==",

View file

@ -8,3 +8,11 @@
- Split a mixed test file along that line instead of widening visibility to move it
- A test for another crate's item belongs in that crate, not in a downstream one
- Never set `autotests = false` or hand-list `[[test]]` targets; every file directly under `tests/` is discovered by cargo, and a shared helper goes in `tests/<name>/mod.rs` or `tests/<subject>/support.rs` so it is not picked up as a test crate of its own
## Error definitions
- A crate's errors live in `src/error.rs`, defined with `thiserror`, and re-exported from `lib.rs`
- Default to one top-level `Error` enum per crate, with one variant per failure mode and a `#[error(...)]` message on each
- Wrap a lower-level error as a variant with `#[from]` or `#[source]` instead of flattening it to a string
- Exception: split into separate types when different functions fail in disjoint ways, especially when different callers see them. A shared enum would force every caller to match variants its function can never return
- Name a split type after what went wrong (a unit struct is fine for a single failure mode), not after the function that returns it

View file

@ -3013,6 +3013,15 @@ dependencies = [
"url",
]
[[package]]
name = "litellm-coroutine"
version = "0.1.0"
dependencies = [
"rstest",
"thiserror 2.0.19",
"tokio",
]
[[package]]
name = "litellm-cost"
version = "0.1.0"
@ -3040,6 +3049,7 @@ name = "litellm-host"
version = "0.1.0"
dependencies = [
"litellm-auth",
"litellm-coroutine",
"rstest",
"serde_json",
"tokio",
@ -3049,6 +3059,7 @@ dependencies = [
name = "litellm-host-python"
version = "0.1.0"
dependencies = [
"bytes",
"futures-util",
"litellm-host",
"pyo3",

View file

@ -12,6 +12,7 @@ repository = "https://github.com/BerriAI/litellm"
litellm-tracing = { path = "crates/tracing" }
tracing = "0.1"
litellm-core = { path = "crates/core" }
litellm-coroutine = { path = "crates/coroutine" }
litellm-host = { path = "crates/host" }
litellm-callbacks-legacy-python = { path = "crates/callbacks-legacy-python" }
litellm-framing = { path = "crates/framer" }

View file

@ -3,8 +3,8 @@
//! lifetime. No other callback host has that obligation, which is why nothing outside
//! this crate holds them.
use litellm_host::{machine::Machine, route::Route};
use litellm_host_python::{RouteHost, lookup, run_call};
use litellm_host::{machine::Machine, protocol::Protocol};
use litellm_host_python::{ProtocolHost, lookup, run_call};
use pyo3::{
gc::{PyTraverseError, PyVisit},
prelude::*,
@ -63,25 +63,25 @@ impl PublicCall {
}
}
/// Runs one native call under the legacy `Logging` contract: the route host projects from
/// Runs one native call under the legacy `Logging` contract: the protocol host projects from
/// the keyword view the contract prepares, and the contract observes the call.
pub fn run_legacy_call<H, M>(
py: Python<'_>,
surface: LegacySurface,
call: PublicCall,
machine: M,
route: H,
host: H,
asynchronous: bool,
) -> PyResult<Py<PyAny>>
where
H: RouteHost + 'static,
M: Machine<Route = H::Route, Complete = <H::Route as Route>::Response> + 'static,
H: ProtocolHost + 'static,
M: Machine<Protocol = H::Protocol, Complete = <H::Protocol as Protocol>::Response> + 'static,
{
let arguments = call.kwargs.clone_ref(py);
run_call(
py,
machine,
route,
host,
Box::new(LegacyLogging::new(py, surface, call, asynchronous)),
arguments,
asynchronous,

View file

@ -1,4 +1,5 @@
use std::{
convert::Infallible,
sync::{Arc, Mutex},
time::Duration,
};
@ -9,8 +10,8 @@ use litellm_core_utils::get_llm_provider_logic::get_custom_llm_provider;
use litellm_host::{
event::{MachineEvent, RawResponse, RequestContext, WireRequest},
host::{Demand, Host},
machine::{HostChannel, MachineFault, RouteMachine},
route::Route,
machine::{CallMachine, HostChannel, MachineFault},
protocol::Protocol,
};
use litellm_secrets::source::SecretSource;
use litellm_types::{
@ -28,15 +29,6 @@ use super::{
};
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MessagesOp {
ProjectRequest,
}
pub enum MessagesOpResult {
Request(Box<MessagesCall>),
}
/// The caller's request as the host projects it.
pub struct MessagesCall {
pub model: String,
@ -64,11 +56,11 @@ pub enum MessagesOutput {
pub struct Messages;
impl Route for Messages {
impl Protocol for Messages {
type Response = MessagesOutput;
type Error = Error;
type Op = MessagesOp;
type OpResult = MessagesOpResult;
type Projection = MessagesCall;
type Op = Infallible;
type Chunk = Bytes;
type StreamHead = ();
}
@ -78,13 +70,12 @@ impl From<MachineFault> for Error {
Self::InvalidRequest(match fault {
MachineFault::Abandoned => "messages host driver was abandoned".into(),
MachineFault::Protocol(message) => format!("messages {message}"),
MachineFault::Mismatch => "invalid messages host operation result".into(),
})
}
}
pub type MessagesHost = HostChannel<Messages>;
pub type MessagesMachine = RouteMachine<Messages>;
pub type MessagesMachine = CallMachine<Messages>;
/// Whether this route serves the request, decided before any callback runs so a host
/// can still run its own path.
@ -114,30 +105,28 @@ impl LocalMessagesHost {
}
impl Host<Messages> for LocalMessagesHost {
async fn route(&self, op: MessagesOp) -> Result<MessagesOpResult, Error> {
match op {
MessagesOp::ProjectRequest => self
.call
.lock()
.unwrap_or_else(|error| error.into_inner())
.take()
.map(|call| MessagesOpResult::Request(Box::new(call)))
.ok_or_else(|| {
Error::InvalidRequest("messages request was already projected".into())
}),
}
async fn project(&self) -> Result<MessagesCall, Error> {
self.call
.lock()
.unwrap_or_else(|error| error.into_inner())
.take()
.ok_or_else(|| Error::InvalidRequest("messages request was already projected".into()))
}
async fn custom_op(&self, op: Infallible) -> Result<(), Error> {
match op {}
}
}
pub fn messages_machine(secrets: Arc<dyn SecretSource>) -> MessagesMachine {
RouteMachine::new(move |host| Box::pin(execute(host, secrets.clone())))
CallMachine::new(move |host| Box::pin(execute(host, secrets.clone())))
}
async fn execute(
host: MessagesHost,
secrets: Arc<dyn SecretSource>,
) -> Result<MessagesOutput, Error> {
let MessagesOpResult::Request(call) = host.route(MessagesOp::ProjectRequest).await?;
let call = host.project().await?;
let stream = call.streams();
let resolved = resolve_provider(&call.model, call.custom_llm_provider.as_deref())?;
let secrets = secrets.resolve(resolved.config.secret_names()).await?;

View file

@ -23,9 +23,6 @@ pub fn prepare_document(input: OcrDocumentInput) -> Result<OcrDocument, Error> {
file_name.as_deref(),
mime_type.as_deref(),
)?),
OcrDocumentInput::HostReader { .. } => Err(Error::InvalidRequest(
"OCR file reader was not read by the host".into(),
)),
}
}
@ -207,7 +204,7 @@ mod tests {
}
#[test]
fn byte_documents_are_encoded_and_host_readers_must_be_read_first() {
fn byte_documents_are_encoded() {
assert_eq!(
prepare_document(OcrDocumentInput::Bytes {
bytes: b"abc".as_slice().into(),
@ -217,7 +214,6 @@ mod tests {
.unwrap(),
document("data:application/pdf;base64,YWJj")
);
assert!(prepare_document(OcrDocumentInput::HostReader { mime_type: None }).is_err());
}
#[test]

View file

@ -3,102 +3,73 @@ use std::sync::{Arc, Mutex};
use litellm_auth::ResolvedCredential;
use litellm_host::{
event::{CallEvent, RequestContext, WireRequest},
machine::{HostChannel, HostTokenProvider, MachineFault, RouteMachine, TokenRoute},
route::Route,
host::Reply,
machine::{CallMachine, HostChannel, HostTokenProvider, TokenProtocol},
protocol::Protocol,
};
use litellm_llms::base_llm::ocr::{
error::Error, handler::OcrClient, transformation::LiteLLMOcrResponse,
};
use super::handler::perform_ocr_request;
use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput, OcrFileContent, ResolvedOcrRequest};
use crate::ocr::types::{LiteLLMOcrRequest, OcrDocumentInput, ResolvedOcrRequest};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum OcrOp {
ProjectRequest,
ReadDocument,
AcquireAzureAdToken,
AcquireAzureAdToken(Reply<ResolvedCredential>),
}
pub enum OcrOpResult {
Request {
request: Box<LiteLLMOcrRequest<OcrDocumentInput>>,
caller_token: bool,
},
Document(OcrFileContent),
AzureAdToken(ResolvedCredential),
/// The caller's request as the host projects it.
pub struct OcrProjection {
pub request: LiteLLMOcrRequest<OcrDocumentInput>,
/// The caller passed its own Azure AD token provider, which the host keeps.
pub caller_token: bool,
}
pub struct Ocr;
impl Route for Ocr {
impl Protocol for Ocr {
type Response = LiteLLMOcrResponse;
type Error = Error;
type Projection = OcrProjection;
type Op = OcrOp;
type OpResult = OcrOpResult;
type Chunk = std::convert::Infallible;
type StreamHead = std::convert::Infallible;
}
impl TokenRoute for Ocr {
fn acquire_token_op() -> OcrOp {
OcrOp::AcquireAzureAdToken
}
fn token_credential(result: OcrOpResult) -> Option<ResolvedCredential> {
match result {
OcrOpResult::AzureAdToken(credential) => Some(credential),
_ => None,
}
impl TokenProtocol for Ocr {
fn acquire_token_op(reply: Reply<ResolvedCredential>) -> OcrOp {
OcrOp::AcquireAzureAdToken(reply)
}
}
pub type OcrHost = HostChannel<Ocr>;
pub type OcrMachine = RouteMachine<Ocr>;
pub type OcrMachine = CallMachine<Ocr>;
/// The OCR call as a machine: projection, document reading and token acquisition are
/// host operations; everything else runs in Rust.
/// The OCR call as a machine: projection and token acquisition are host operations;
/// everything else runs in Rust.
pub fn ocr_machine(client: OcrClient) -> OcrMachine {
RouteMachine::new(move |host| Box::pin(execute(client, host)))
CallMachine::new(move |host| Box::pin(execute(client, host)))
}
async fn execute(client: OcrClient, host: OcrHost) -> Result<LiteLLMOcrResponse, Error> {
let OcrOpResult::Request {
let OcrProjection {
request,
caller_token,
} = host.route(OcrOp::ProjectRequest).await?
else {
return Err(MachineFault::Mismatch.into());
};
} = host.project().await?;
let request = LiteLLMOcrRequest {
azure_ad_token_provider: caller_token
.then(|| HostTokenProvider::handle(host.clone()))
.or(request.azure_ad_token_provider),
..*request
..request
};
let caller_document = matches!(request.document, OcrDocumentInput::Document(_));
let request = prepare_request_document(request, &host).await?;
let request = prepare_request_document(request).await?;
perform_ocr_request(&client, request, &host, caller_document).await
}
async fn prepare_request_document(
request: LiteLLMOcrRequest<OcrDocumentInput>,
host: &OcrHost,
) -> Result<ResolvedOcrRequest, Error> {
let request = match &request.document {
OcrDocumentInput::HostReader { mime_type } => {
let mime_type = mime_type.clone();
let OcrOpResult::Document(content) = host.route(OcrOp::ReadDocument).await? else {
return Err(MachineFault::Mismatch.into());
};
request.with_document(OcrDocumentInput::Bytes {
bytes: content.bytes,
file_name: content.file_name,
mime_type,
})
}
_ => request,
};
if let OcrDocumentInput::Document(_) = &request.document {
return request.map_document(super::document::prepare_document);
}
@ -107,7 +78,6 @@ async fn prepare_request_document(
.map_err(|error| Error::DocumentTask(Arc::new(error)))?
}
type Reader = Box<dyn Fn() -> Result<OcrFileContent, Error> + Send + Sync>;
type BeforeSend =
Box<dyn Fn(WireRequest, &RequestContext) -> Result<WireRequest, Error> + Send + Sync>;
type Observer = Box<dyn Fn(&CallEvent) + Send + Sync>;
@ -116,7 +86,6 @@ type Observer = Box<dyn Fn(&CallEvent) + Send + Sync>;
/// projection, and the optional observer sees and may rewrite the wire request.
pub struct LocalOcrHost {
request: Mutex<Option<LiteLLMOcrRequest<OcrDocumentInput>>>,
reader: Option<Reader>,
before_send: Option<BeforeSend>,
observer: Option<Observer>,
}
@ -125,22 +94,11 @@ impl LocalOcrHost {
pub fn new(request: LiteLLMOcrRequest<OcrDocumentInput>) -> Self {
Self {
request: Mutex::new(Some(request)),
reader: None,
before_send: None,
observer: None,
}
}
pub fn with_reader(
self,
reader: impl Fn() -> Result<OcrFileContent, Error> + Send + Sync + 'static,
) -> Self {
Self {
reader: Some(Box::new(reader)),
..self
}
}
pub fn with_before_send(
self,
before_send: impl Fn(WireRequest, &RequestContext) -> Result<WireRequest, Error>
@ -163,25 +121,21 @@ impl LocalOcrHost {
}
impl litellm_host::host::Host<Ocr> for LocalOcrHost {
async fn route(&self, op: OcrOp) -> Result<OcrOpResult, Error> {
async fn project(&self) -> Result<OcrProjection, Error> {
self.request
.lock()
.unwrap_or_else(|error| error.into_inner())
.take()
.map(|request| OcrProjection {
request,
caller_token: false,
})
.ok_or_else(|| Error::InvalidRequest("OCR request was already projected".into()))
}
async fn custom_op(&self, op: OcrOp) -> Result<(), Error> {
match op {
OcrOp::ProjectRequest => self
.request
.lock()
.unwrap_or_else(|error| error.into_inner())
.take()
.map(|request| OcrOpResult::Request {
request: Box::new(request),
caller_token: false,
})
.ok_or_else(|| Error::InvalidRequest("OCR request was already projected".into())),
OcrOp::ReadDocument => self
.reader
.as_ref()
.ok_or_else(|| Error::InvalidRequest("OCR host has no document reader".into()))
.and_then(|reader| reader())
.map(OcrOpResult::Document),
OcrOp::AcquireAzureAdToken => {
OcrOp::AcquireAzureAdToken(_) => {
Err(Error::Auth(litellm_auth::Error::AzureTokenAcquisition(
"OCR host has no Azure AD token provider".into(),
)))
@ -2757,7 +2711,7 @@ pub(crate) mod tests {
use litellm_auth_gcp::VertexAuth;
use litellm_host::{
event::{CallEvent, MachineEvent, WireRequest},
host::{Host, HostOp, HostResult},
host::{Host, HostOp},
machine::{HostFailure, Machine, MachineStep},
};
use litellm_http::{
@ -2776,7 +2730,7 @@ pub(crate) mod tests {
use rstest::rstest;
use serde_json::{Value, json};
use crate::ocr::route::{LocalOcrHost, OcrOp, OcrOpResult, ocr_machine};
use crate::ocr::route::{LocalOcrHost, OcrOp, OcrProjection, ocr_machine};
use crate::ocr::{
test_support::{
MockResponse, mock_server, ocr_client, perform_ocr, perform_ocr_with, wire_request,
@ -3212,42 +3166,42 @@ pub(crate) mod tests {
crate::ocr::route::OcrMachine,
) {
let mut machine = ocr_machine(client);
let mut result = None;
let mut ops = Vec::new();
let outcome = loop {
let op = match machine.resume(result.take()).await {
let op = match machine.resume().await {
Ok(MachineStep::Host(op)) => op,
Ok(MachineStep::Complete(response)) => break Ok(response),
Err(error) => break Err(error),
};
let answer = match op {
HostOp::Route(op) => {
ops.push(match op {
OcrOp::ProjectRequest => "ProjectRequest",
OcrOp::ReadDocument => "ReadDocument",
OcrOp::AcquireAzureAdToken => "AcquireAzureAdToken",
});
host.route(op)
HostOp::Project(reply) => {
ops.push("Project");
host.project()
.await
.map(HostResult::Route)
.map(|projection| reply.send(projection))
.map_err(HostFailure::Error)
}
HostOp::BeforeSend { wire, .. } => {
ops.push("BeforeSend");
intercept(*wire).map(|wire| HostResult::BeforeSend(Box::new(wire)))
HostOp::Custom(op) => {
ops.push(match op {
OcrOp::AcquireAzureAdToken(_) => "AcquireAzureAdToken",
});
host.custom_op(op).await.map_err(HostFailure::Error)
}
HostOp::Emit(event) => {
HostOp::BeforeSend { wire, reply, .. } => {
ops.push("BeforeSend");
intercept(*wire).map(|wire| reply.send(wire))
}
HostOp::Emit(event, reply) => {
let event = CallEvent::Machine(event);
ops.push(event_name(&event));
host.emit(&event)
.await
.map(|()| HostResult::Emitted)
.map(|()| reply.send(()))
.map_err(HostFailure::Error)
}
};
match answer {
Ok(answer) => result = Some(answer),
Err(failure) => break machine.interrupt(failure).await,
if let Err(failure) = answer {
break machine.interrupt(failure).await;
}
};
(outcome, ops, machine)
@ -3269,8 +3223,8 @@ pub(crate) mod tests {
assert!(
matches!(outcome, Err(OcrError::InvalidRequest(message)) if message == "before_send failed")
);
assert_eq!(ops, ["ProjectRequest", "BeforeSend"]);
assert!(machine.resume(None).await.is_err());
assert_eq!(ops, ["Project", "BeforeSend"]);
assert!(machine.resume().await.is_err());
}
#[tokio::test]
@ -3306,80 +3260,24 @@ pub(crate) mod tests {
server.await.unwrap();
assert_eq!(outcome.unwrap().pages[0].markdown, "native");
assert_eq!(seen.lock().unwrap().len(), 1);
assert_eq!(ops, ["ProjectRequest", "BeforeSend", "response"]);
assert_eq!(ops, ["Project", "BeforeSend", "response"]);
assert!(matches!(
machine.resume(None).await,
machine.resume().await,
Err(OcrError::InvalidRequest(_))
));
}
async fn drive_native_file_call(
request: crate::ocr::types::LiteLLMOcrRequest<crate::ocr::types::OcrDocumentInput>,
content: Result<crate::ocr::types::OcrFileContent, OcrError>,
) -> (Result<LiteLLMOcrResponse, OcrError>, usize) {
let reads = Arc::new(Mutex::new(0));
let counted = reads.clone();
let content = Mutex::new(Some(content));
let host = LocalOcrHost::new(request).with_reader(move || {
*counted.lock().unwrap() += 1;
content.lock().unwrap().take().unwrap()
});
let outcome = perform_ocr_with(host).await;
let reads = *reads.lock().unwrap();
(outcome, reads)
}
#[tokio::test]
async fn host_reader_documents_are_read_once_at_the_core_selected_point_and_encoded() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"pages":[{"index":0,"markdown":"file"}]
}))])
.await;
let request = wire_request("mistral/model", &base, json!({})).with_document(
crate::ocr::types::OcrDocumentInput::HostReader {
mime_type: Some("application/pdf".into()),
},
);
let (response, reads) = drive_native_file_call(
request,
Ok(crate::ocr::types::OcrFileContent {
bytes: b"abc".as_slice().into(),
file_name: Some("scan.png".into()),
}),
)
.await;
server.await.unwrap();
assert_eq!(response.unwrap().pages[0].markdown, "file");
assert_eq!(reads, 1);
assert!(seen.lock().unwrap()[0].contains("data:application/pdf;base64,YWJj"));
}
#[tokio::test]
async fn host_reader_failures_and_empty_files_fail_before_the_provider_is_called() {
async fn empty_byte_documents_fail_before_the_provider_is_called() {
let (base, seen, _server) = mock_server(vec![]).await;
let request = wire_request("mistral/model", &base, json!({}));
let failure = OcrError::InvalidRequest("reader exploded".into());
let (response, reads) = drive_native_file_call(
request
.with_document(crate::ocr::types::OcrDocumentInput::HostReader { mime_type: None }),
Err(failure.clone()),
)
.await;
assert!(
matches!(response.unwrap_err(), OcrError::InvalidRequest(message) if message == "reader exploded")
);
assert_eq!(reads, 1);
let request = wire_request("mistral/model", &base, json!({}));
let (response, _) = drive_native_file_call(
request
.with_document(crate::ocr::types::OcrDocumentInput::HostReader { mime_type: None }),
Ok(crate::ocr::types::OcrFileContent {
let request = wire_request("mistral/model", &base, json!({})).with_document(
crate::ocr::types::OcrDocumentInput::Bytes {
bytes: Default::default(),
file_name: None,
}),
)
.await;
mime_type: None,
},
);
let response = perform_ocr_with(LocalOcrHost::new(request)).await;
assert!(matches!(response.unwrap_err(), OcrError::EmptyFile));
assert!(seen.lock().unwrap().is_empty());
}
@ -3400,24 +3298,21 @@ pub(crate) mod tests {
mime_type: None,
},
);
let (response, reads) =
drive_native_file_call(request, Err(OcrError::InvalidRequest("unused".into()))).await;
let (response, ops, _) = drive_until(ocr_client(), &LocalOcrHost::new(request), Ok).await;
server.await.unwrap();
std::fs::remove_dir_all(&dir).unwrap();
assert_eq!(response.unwrap().pages[0].markdown, "path");
assert_eq!(reads, 0);
assert_eq!(ops, ["Project", "BeforeSend", "response"]);
assert!(seen.lock().unwrap()[0].contains("data:image/png;base64,YWJj"));
let (base, seen, _server) = mock_server(vec![]).await;
let request = wire_request("mistral/model", &base, json!({}));
let (response, _) = drive_native_file_call(
request.with_document(crate::ocr::types::OcrDocumentInput::Path {
let request = wire_request("mistral/model", &base, json!({})).with_document(
crate::ocr::types::OcrDocumentInput::Path {
path: path.clone(),
mime_type: None,
}),
Err(OcrError::InvalidRequest("unused".into())),
)
.await;
},
);
let response = perform_ocr_with(LocalOcrHost::new(request)).await;
assert!(matches!(
response.unwrap_err(),
OcrError::FileRead { path: failed, source } if failed == path && source.kind() == std::io::ErrorKind::NotFound
@ -3441,28 +3336,25 @@ pub(crate) mod tests {
assert!(
matches!(outcome, Err(OcrError::InvalidRequest(message)) if message == "cancelled")
);
assert_eq!(ops, ["ProjectRequest", "BeforeSend"]);
assert!(machine.resume(Some(HostResult::Emitted)).await.is_err());
assert_eq!(ops, ["Project", "BeforeSend"]);
assert!(machine.resume().await.is_err());
}
#[tokio::test]
async fn missing_host_result_preserves_pending_operation() {
async fn resuming_before_answering_preserves_pending_operation() {
let request = wire_request("mistral/model", "http://127.0.0.1:1", json!({}));
let mut machine = ocr_machine(ocr_client());
let Ok(MachineStep::Host(HostOp::Project(reply))) = machine.resume().await else {
panic!("expected the projection op first");
};
assert!(machine.resume().await.is_err());
reply.send(OcrProjection {
request,
caller_token: false,
});
assert!(matches!(
machine.resume(None).await.unwrap(),
MachineStep::Host(HostOp::Route(OcrOp::ProjectRequest))
));
assert!(machine.resume(None).await.is_err());
assert!(matches!(
machine
.resume(Some(HostResult::Route(OcrOpResult::Request {
request: Box::new(request),
caller_token: false,
})))
.await
.unwrap(),
MachineStep::Host(HostOp::BeforeSend { .. })
machine.resume().await,
Ok(MachineStep::Host(HostOp::BeforeSend { .. }))
));
}
@ -3623,20 +3515,18 @@ pub(crate) mod tests {
};
let host = LocalOcrHost::new(request);
let mut machine = ocr_machine(ocr_client());
let mut result = None;
tokio::time::timeout(std::time::Duration::from_secs(2), async {
loop {
tokio::select! {
_ = entered.notified() => break,
step = machine.resume(result.take()) => {
result = Some(match step.unwrap() {
MachineStep::Host(HostOp::Route(op)) => HostResult::Route(host.route(op).await.unwrap()),
MachineStep::Host(HostOp::BeforeSend { wire, .. }) => {
HostResult::BeforeSend(wire)
}
MachineStep::Host(HostOp::Emit(_)) => HostResult::Emitted,
step = machine.resume() => {
match step.unwrap() {
MachineStep::Host(HostOp::Project(reply)) => reply.send(host.project().await.unwrap()),
MachineStep::Host(HostOp::Custom(op)) => host.custom_op(op).await.unwrap(),
MachineStep::Host(HostOp::BeforeSend { wire, reply, .. }) => reply.send(*wire),
MachineStep::Host(HostOp::Emit(_, reply)) => reply.send(()),
MachineStep::Complete(_) => panic!("pending provider completed"),
});
}
}
}
}
@ -3661,24 +3551,23 @@ pub(crate) mod tests {
}
impl Host<crate::ocr::route::Ocr> for CallerTokenHost {
async fn route(&self, op: OcrOp) -> Result<OcrOpResult, OcrError> {
async fn project(&self) -> Result<OcrProjection, OcrError> {
self.trace.lock().unwrap().push("project".into());
Ok(OcrProjection {
request: self.request.lock().unwrap().take().unwrap(),
caller_token: true,
})
}
async fn custom_op(&self, op: OcrOp) -> Result<(), OcrError> {
match op {
OcrOp::ProjectRequest => {
self.trace.lock().unwrap().push("project".into());
Ok(OcrOpResult::Request {
request: Box::new(self.request.lock().unwrap().take().unwrap()),
caller_token: true,
})
}
OcrOp::AcquireAzureAdToken => {
OcrOp::AcquireAzureAdToken(reply) => {
self.trace.lock().unwrap().push("token".into());
Ok(OcrOpResult::AzureAdToken(
litellm_auth::ResolvedCredential::Static(litellm_auth::SecretValue::new(
"caller-token",
)),
))
reply.send(litellm_auth::ResolvedCredential::Static(
litellm_auth::SecretValue::new("caller-token"),
));
Ok(())
}
OcrOp::ReadDocument => Err(OcrError::InvalidRequest("no reader".into())),
}
}
@ -3761,18 +3650,18 @@ pub(crate) mod tests {
});
let host = LocalOcrHost::new(wire_request("mistral/model", &base, json!({})));
let mut machine = ocr_machine(ocr_client());
let mut result = None;
tokio::time::timeout(std::time::Duration::from_secs(2), async {
loop {
tokio::select! {
_ = received.notified() => break,
step = machine.resume(result.take()) => {
result = Some(match step.unwrap() {
MachineStep::Host(HostOp::Route(op)) => HostResult::Route(host.route(op).await.unwrap()),
MachineStep::Host(HostOp::BeforeSend { wire, .. }) => HostResult::BeforeSend(wire),
MachineStep::Host(HostOp::Emit(_)) => HostResult::Emitted,
step = machine.resume() => {
match step.unwrap() {
MachineStep::Host(HostOp::Project(reply)) => reply.send(host.project().await.unwrap()),
MachineStep::Host(HostOp::Custom(op)) => host.custom_op(op).await.unwrap(),
MachineStep::Host(HostOp::BeforeSend { wire, reply, .. }) => reply.send(*wire),
MachineStep::Host(HostOp::Emit(_, reply)) => reply.send(()),
MachineStep::Complete(_) => panic!("the stalled provider completed"),
});
}
}
}
}

View file

@ -25,9 +25,6 @@ pub enum OcrDocumentInput {
file_name: Option<String>,
mime_type: Option<String>,
},
HostReader {
mime_type: Option<String>,
},
}
impl From<OcrDocument> for OcrDocumentInput {
@ -45,12 +42,6 @@ impl From<PathBuf> for OcrDocumentInput {
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct OcrFileContent {
pub bytes: Bytes,
pub file_name: Option<String>,
}
/// Caller-supplied connection overrides for a [`LiteLLMOcrRequest`], in the
/// shape hosts receive them: JSON-ish headers, optional timeout, optional
/// credentials, and per-field provenance in `input_sources`.

View file

@ -0,0 +1,31 @@
# Requirements
Core must pause mid-call to ask the host for things it cannot do itself (Python callbacks, secret and token reads, `before_send` rewrites, stream demand), then continue where it stopped. Any change to this crate must keep every requirement below; the alternatives section says which one each rejected design breaks
- R1 Core never calls the host: it names an op and waits for the answer, so it stays free of PyO3 and of any other host runtime
- R2 Async host work is awaited by the host's own driver in the caller's asyncio task (`litellm/rust_bridge/lifecycle.py`), so `contextvars` writes reach the caller; a Rust-side `into_future` would run it in a copied context
- R3 The body awaits real I/O (HTTP, `spawn_blocking`, timers) between yields, so `resume` is itself a future driven by the caller's runtime
- R4 Route code stays straight-line async (`host.route(OcrOp::ReadDocument).await?`) instead of hand-written states
- R5 Each op fixes its answer type at compile time: a host cannot answer `ReadDocument` with a token, and core never matches a result variant it did not ask for
- R6 A yield the body makes while being resumed is returned by that same poll, so the host driver's inline first poll needs no extra event-loop turn per op
- R7 No task is spawned: `cancel`, or dropping the coroutine, drops the body, and nothing waits forever on an answer that cannot come
- R8 Several yields can be pending at once, since route code hands clones of its `Co` to token providers and hooks
- R9 Stable Rust
# Other implementations and why they do not fit
- Nightly `std::ops::Coroutine`: breaks R9, and its body cannot await futures between yields (R3)
- `genawaiter`: resumes async bodies only with a noop waker, so the body cannot await real I/O (R3)
- `simple_coro`: typestate `Coro` makes answering before resuming a compile-time rule, but its body cannot await arbitrary futures (R3) and its reply type `R` is fixed per coroutine (R5)
- `corosensei` and other stackful coroutines: sync bodies on their own stack, no async I/O inside (R3)
- A hand-written phase enum with an `advance` match (the old `HostPhase`): every await point becomes a state (R4)
- An injected host trait with `async fn`s: core would call the host itself (R1, R2)
- Sans-IO, where core does no I/O and HTTP becomes one more host op: keeps every requirement and makes `resume` a pure step function, but HTTP, streaming, retries and timeouts would move out of core into every bridge; the one real alternative, not taken
- Temporal's Rust workflow SDK (`WorkflowFuture`, `WfContext`) is the closest precedent: an `async fn` polled in place, commands sent over a channel with a oneshot to unblock them. Roles are inverted there (the language SDK owns the program, core answers), and its workflow body may not do real I/O
# Tradeoffs accepted
- A tokio `mpsc` channel plus a `oneshot` per yield instead of compiler-generated states
- Protocol mistakes (resuming before answering, resuming after the end) are runtime `ResumeError`s, not compile errors
- Pending yields come out one per `resume`, in the order they were made, and each reply goes back to the yield that made it (R8)
- An answer sent after its yield stopped waiting (for example the body timed out on it) is discarded, since the body already moved on

View file

@ -0,0 +1,15 @@
[package]
name = "litellm-coroutine"
version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
description = "Async coroutines on stable Rust whose every yield carries its own typed reply"
[dependencies]
thiserror.workspace = true
tokio = { workspace = true, features = ["sync"] }
[dev-dependencies]
rstest.workspace = true
tokio = { workspace = true, features = ["rt", "macros", "time"] }

View file

@ -0,0 +1,42 @@
use std::sync::Weak;
use tokio::sync::mpsc;
use crate::{Abandoned, Reply, reply};
pub(crate) struct Request<Y> {
pub(crate) value: Y,
pub(crate) outstanding: Weak<()>,
}
/// The body's handle for yielding, `genawaiter`'s `Co`.
pub struct Co<Y> {
yields: mpsc::UnboundedSender<Request<Y>>,
}
impl<Y> Clone for Co<Y> {
fn clone(&self) -> Self {
Self {
yields: self.yields.clone(),
}
}
}
impl<Y> Co<Y> {
pub(crate) fn new(yields: mpsc::UnboundedSender<Request<Y>>) -> Self {
Self { yields }
}
/// Yields the value `ask` builds around a fresh [`Reply`] and waits for its answer.
pub async fn yield_<A>(&self, ask: impl FnOnce(Reply<A>) -> Y) -> Result<A, Abandoned> {
let (reply, answer) = reply();
let outstanding = reply.outstanding();
self.yields
.send(Request {
value: ask(reply),
outstanding,
})
.map_err(|_| Abandoned)?;
answer.await
}
}

View file

@ -0,0 +1,94 @@
use std::{
future::{Future, poll_fn},
pin::Pin,
sync::Weak,
task::{Context, Poll},
};
use tokio::sync::mpsc;
use crate::{Co, ResumeError, co::Request};
/// What one `resume` produced, as in [`std::ops::CoroutineState`].
#[derive(Debug, PartialEq, Eq)]
pub enum CoroutineState<Y, C> {
Yielded(Y),
Complete(C),
}
type Body<C> = Pin<Box<dyn Future<Output = C> + Send>>;
enum Step<Y, C> {
Yielded(Request<Y>),
Complete(C),
}
fn queued<Y>(
yields: &mut mpsc::UnboundedReceiver<Request<Y>>,
context: &mut Context<'_>,
) -> Option<Request<Y>> {
match yields.poll_recv(context) {
Poll::Ready(request) => request,
Poll::Pending => None,
}
}
pub struct Coroutine<Y, C> {
body: Option<Body<C>>,
yields: mpsc::UnboundedReceiver<Request<Y>>,
outstanding: Weak<()>,
}
impl<Y, C> Coroutine<Y, C> {
/// Builds the body from `producer`. Nothing runs until the first `resume`.
pub fn new<F>(producer: impl FnOnce(Co<Y>) -> F) -> Self
where
F: Future<Output = C> + Send + 'static,
{
let (sender, yields) = mpsc::unbounded_channel();
Self {
body: Some(Box::pin(producer(Co::new(sender)))),
yields,
outstanding: Weak::new(),
}
}
pub async fn resume(&mut self) -> Result<CoroutineState<Y, C>, ResumeError> {
let Some(body) = self.body.as_mut() else {
return Err(ResumeError::Finished);
};
if self.outstanding.strong_count() > 0 {
return Err(ResumeError::Unanswered);
}
let yields = &mut self.yields;
let step = poll_fn(|context| {
if let Some(request) = queued(yields, context) {
return Poll::Ready(Step::Yielded(request));
}
if let Poll::Ready(output) = body.as_mut().poll(context) {
return Poll::Ready(Step::Complete(output));
}
queued(yields, context)
.map_or(Poll::Pending, |request| Poll::Ready(Step::Yielded(request)))
})
.await;
match step {
Step::Yielded(Request { value, outstanding }) => {
self.outstanding = outstanding;
Ok(CoroutineState::Yielded(value))
}
Step::Complete(output) => {
self.cancel();
Ok(CoroutineState::Complete(output))
}
}
}
/// Drops the body and fails every yield still waiting, or yet to be made, with
/// [`Abandoned`](crate::Abandoned).
pub fn cancel(&mut self) {
self.body = None;
self.yields.close();
while self.yields.try_recv().is_ok() {}
}
}

View file

@ -0,0 +1,14 @@
/// A `resume` the coroutine refused, leaving it as it was.
#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
pub enum ResumeError {
#[error("coroutine resumed after it finished")]
Finished,
#[error("coroutine resumed before the reply to its last yield was sent or dropped")]
Unanswered,
}
/// No answer will come to a yield: its [`Reply`](crate::Reply) was dropped unsent, or the
/// coroutine it was sent to is gone.
#[derive(Clone, Copy, Debug, PartialEq, Eq, thiserror::Error)]
#[error("the yield was abandoned before it was answered")]
pub struct Abandoned;

View file

@ -0,0 +1,12 @@
//! Async coroutines on stable Rust whose every yield carries its own typed [`Reply`].
//! See `AGENTS.md` for the requirement, the alternatives and the contracts.
mod co;
mod coroutine;
mod error;
mod reply;
pub use co::Co;
pub use coroutine::{Coroutine, CoroutineState};
pub use error::{Abandoned, ResumeError};
pub use reply::{Answer, Reply, reply};

View file

@ -0,0 +1,60 @@
use std::{
fmt,
future::Future,
pin::Pin,
sync::{Arc, Weak},
task::{Context, Poll},
};
use tokio::sync::oneshot;
use crate::Abandoned;
/// The one way to answer a yield. Sending or dropping it settles the yield.
pub struct Reply<A> {
slot: oneshot::Sender<A>,
outstanding: Arc<()>,
}
impl<A> Reply<A> {
/// An answer the yield no longer awaits is discarded.
pub fn send(self, answer: A) {
let _ = self.slot.send(answer);
}
/// Alive until this reply is sent or dropped.
pub(crate) fn outstanding(&self) -> Weak<()> {
Arc::downgrade(&self.outstanding)
}
}
impl<A> fmt::Debug for Reply<A> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str("Reply")
}
}
/// The waiting end of a [`Reply`].
pub struct Answer<A> {
slot: oneshot::Receiver<A>,
}
impl<A> Future for Answer<A> {
type Output = Result<A, Abandoned>;
fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
Pin::new(&mut self.slot)
.poll(context)
.map(|answer| answer.map_err(|_| Abandoned))
}
}
/// A reply outside any coroutine, for answering a host operation directly.
pub fn reply<A>() -> (Reply<A>, Answer<A>) {
let (slot, answer) = oneshot::channel();
let reply = Reply {
slot,
outstanding: Arc::new(()),
};
(reply, Answer { slot: answer })
}

View file

@ -0,0 +1,256 @@
use std::{
future::Future,
sync::{Arc, Mutex},
time::Duration,
};
use litellm_coroutine::{Abandoned, Co, Coroutine, CoroutineState, Reply, ResumeError, reply};
use rstest::rstest;
use tokio::time::timeout;
#[derive(Debug)]
enum Ask {
Name(Reply<&'static str>),
Count(Reply<u32>),
}
type Test<C> = Coroutine<Ask, C>;
fn yielded<C>(state: Result<CoroutineState<Ask, C>, ResumeError>) -> Ask {
match state {
Ok(CoroutineState::Yielded(ask)) => ask,
Ok(CoroutineState::Complete(_)) => panic!("expected a yield, the body returned"),
Err(error) => panic!("expected a yield, resume failed: {error}"),
}
}
fn complete<C>(state: Result<CoroutineState<Ask, C>, ResumeError>) -> C {
match state {
Ok(CoroutineState::Complete(output)) => output,
Ok(CoroutineState::Yielded(ask)) => panic!("expected completion, got {ask:?}"),
Err(error) => panic!("expected completion, resume failed: {error}"),
}
}
fn name(ask: Ask) -> Reply<&'static str> {
match ask {
Ask::Name(reply) => reply,
other => panic!("expected a name ask, got {other:?}"),
}
}
fn count(ask: Ask) -> Reply<u32> {
match ask {
Ask::Count(reply) => reply,
other => panic!("expected a count ask, got {other:?}"),
}
}
/// A body parked at one name ask, with nothing else going on.
fn suspended_once() -> Test<Result<&'static str, Abandoned>> {
Coroutine::new(|co| async move { co.yield_(Ask::Name).await })
}
#[tokio::test]
async fn each_typed_answer_resumes_the_yield_that_asked_for_it() {
let mut coroutine: Test<String> = Coroutine::new(|co| async move {
let first = co.yield_(Ask::Name).await.unwrap();
let second = co.yield_(Ask::Count).await.unwrap();
format!("{first}+{second}")
});
name(yielded(coroutine.resume().await)).send("a");
count(yielded(coroutine.resume().await)).send(2);
assert_eq!(complete(coroutine.resume().await), "a+2");
}
/// A driver that polls `resume` once, inline, sees every yield the body makes during
/// that poll instead of being sent back to its event loop.
#[test]
fn a_yield_made_while_resuming_is_returned_by_that_same_poll() {
let mut coroutine: Test<u32> = Coroutine::new(|co| async move {
let first = co.yield_(Ask::Count).await.unwrap();
let second = co.yield_(Ask::Count).await.unwrap();
first + second
});
let mut context = std::task::Context::from_waker(std::task::Waker::noop());
let mut poll_once =
|coroutine: &mut Test<u32>| match std::pin::pin!(coroutine.resume()).poll(&mut context) {
std::task::Poll::Ready(state) => state,
std::task::Poll::Pending => panic!("resume needed a second poll"),
};
count(yielded(poll_once(&mut coroutine))).send(1);
count(yielded(poll_once(&mut coroutine))).send(2);
assert_eq!(complete(poll_once(&mut coroutine)), 3);
}
#[tokio::test]
async fn the_body_awaits_real_futures_between_yields() {
let mut coroutine: Test<u32> = Coroutine::new(|co| async move {
tokio::time::sleep(Duration::from_millis(5)).await;
co.yield_(Ask::Count).await.unwrap()
});
count(yielded(coroutine.resume().await)).send(7);
assert_eq!(complete(coroutine.resume().await), 7);
}
#[tokio::test]
async fn concurrent_yields_come_out_in_order_and_are_answered_separately() {
let mut coroutine: Test<(&str, u32)> = Coroutine::new(|co| async move {
let (first, second) = tokio::join!(co.yield_(Ask::Name), co.yield_(Ask::Count));
(first.unwrap(), second.unwrap())
});
name(yielded(coroutine.resume().await)).send("one");
count(yielded(coroutine.resume().await)).send(2);
assert_eq!(complete(coroutine.resume().await), ("one", 2));
}
#[tokio::test]
async fn resuming_before_the_reply_is_settled_is_refused_and_keeps_the_yield_waiting() {
let mut coroutine = suspended_once();
let reply = name(yielded(coroutine.resume().await));
assert_eq!(
coroutine.resume().await.unwrap_err(),
ResumeError::Unanswered
);
reply.send("real");
assert_eq!(complete(coroutine.resume().await), Ok("real"));
}
#[tokio::test]
async fn a_dropped_reply_abandons_its_yield() {
let mut coroutine = suspended_once();
drop(yielded(coroutine.resume().await));
assert_eq!(complete(coroutine.resume().await), Err(Abandoned));
}
#[tokio::test]
async fn an_answer_the_yield_no_longer_awaits_is_discarded() {
let mut coroutine: Test<&str> = Coroutine::new(|co| async move {
tokio::select! {
biased;
_ = co.yield_(Ask::Name) => unreachable!("the answer comes after the body moved on"),
() = std::future::ready(()) => {}
}
co.yield_(Ask::Name).await.unwrap()
});
let stale = name(yielded(coroutine.resume().await));
stale.send("stale");
name(yielded(coroutine.resume().await)).send("fresh");
assert_eq!(complete(coroutine.resume().await), "fresh");
}
#[rstest]
#[case::returned(false)]
#[case::cancelled(true)]
#[tokio::test]
async fn a_finished_coroutine_refuses_to_resume(#[case] cancel: bool) {
let mut coroutine = suspended_once();
let reply = name(yielded(coroutine.resume().await));
if cancel {
coroutine.cancel();
} else {
reply.send("done");
complete(coroutine.resume().await).unwrap();
}
assert_eq!(coroutine.resume().await.unwrap_err(), ResumeError::Finished);
}
#[tokio::test]
async fn a_dropped_resume_leaves_the_coroutine_resumable() {
let mut coroutine: Test<u32> = Coroutine::new(|co| async move {
tokio::time::sleep(Duration::from_millis(20)).await;
co.yield_(Ask::Count).await.unwrap()
});
assert!(
timeout(Duration::from_millis(1), coroutine.resume())
.await
.is_err()
);
count(yielded(coroutine.resume().await)).send(3);
assert_eq!(complete(coroutine.resume().await), 3);
}
struct Dropped(Arc<Mutex<bool>>);
impl Drop for Dropped {
fn drop(&mut self) {
*self.0.lock().unwrap() = true;
}
}
#[tokio::test]
async fn cancel_drops_the_body() {
let dropped = Arc::new(Mutex::new(false));
let guard = Dropped(Arc::clone(&dropped));
let mut coroutine: Test<()> = Coroutine::new(|co| async move {
let _guard = guard;
co.yield_(Ask::Count).await.unwrap();
});
let _reply = yielded(coroutine.resume().await);
coroutine.cancel();
assert!(*dropped.lock().unwrap());
}
#[rstest]
#[case::cancelled(true)]
#[case::dropped(false)]
#[tokio::test]
async fn a_co_that_escaped_the_body_is_abandoned_once_the_coroutine_ends(#[case] cancel: bool) {
let escaped: Arc<Mutex<Option<Co<Ask>>>> = Arc::default();
let slot = Arc::clone(&escaped);
let mut coroutine: Test<()> = Coroutine::new(move |co| {
*slot.lock().unwrap() = Some(co.clone());
async move {
co.yield_(Ask::Count).await.unwrap();
}
});
let _reply = yielded(coroutine.resume().await);
let co = escaped.lock().unwrap().take().unwrap();
let waiting = tokio::spawn(async move { co.yield_(Ask::Name).await });
tokio::task::yield_now().await;
if cancel {
coroutine.cancel();
} else {
drop(coroutine);
}
let outcome = timeout(Duration::from_secs(1), waiting)
.await
.expect("an escaped yield waits forever")
.unwrap();
assert_eq!(outcome, Err(Abandoned));
}
#[rstest]
#[case::sent(true)]
#[case::dropped(false)]
#[tokio::test]
async fn a_detached_reply_settles_its_answer(#[case] send: bool) {
let (reply, answer) = reply::<u32>();
if send {
reply.send(5);
} else {
drop(reply);
}
assert_eq!(answer.await, if send { Ok(5) } else { Err(Abandoned) });
}

View file

@ -1,8 +1,8 @@
- Target invariants; implementation and runtime validation may lag these rules
- Keep this crate the CPython runtime adapter and nothing more: Serde marshalling, interpreter detachment, tokio/asyncio glue, the `Execution` handle, the call driver and the `PythonLifecycle`/`RouteHost` traits
- Keep this crate the CPython runtime adapter and nothing more: Serde marshalling, interpreter detachment, tokio/asyncio glue, the `Execution` handle, the call driver and the `PythonLifecycle`/`ProtocolHost` traits
- No LiteLLM domain dependencies beyond `litellm-host`: no route types, no `Logging` policy, no public API registration, no cdylib build features
- The driver emits `Succeeded` or `Failed` exactly once and never dispatches after a cancellation; which Python objects consume those events is the adapter's business
- `RouteHost::invoke` receives the keyword view the adapter's `begin` returned, not the caller's dict; a route host that projects from it inherits that adapter's rewrites (for the legacy adapter: setup, deployment hooks, credential inheritance)
- `ProtocolHost::project` receives the keyword view the adapter's `begin` returned, not the caller's dict; a protocol host that projects from it inherits that adapter's rewrites (for the legacy adapter: setup, deployment hooks, credential inheritance)
- A native failure, including one a host op returns as `InvokeError::Native`, is classified exactly once through the route's `classify`; a Python exception raised inside the call, and a failure in `begin` or `after_success`, is raised as is
- A failing `classify` is raised with the native error's text as its `__context__`, never swallowed
- Use standard PyO3 ownership and conversion APIs

View file

@ -6,6 +6,7 @@ license.workspace = true
repository.workspace = true
[dependencies]
bytes.workspace = true
futures-util.workspace = true
litellm-host.workspace = true
pyo3.workspace = true

View file

@ -1,5 +1,5 @@
use litellm_host::event::{FailureOrigin, MachineEvent, RequestContext, Timing, WireRequest};
use litellm_host::route::Route;
use litellm_host::protocol::Protocol;
use pyo3::exceptions::PyRuntimeError;
use pyo3::gc::{PyTraverseError, PyVisit};
use pyo3::prelude::*;
@ -85,7 +85,7 @@ pub trait PythonLifecycle: Send + Sync {
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>;
}
/// Why a route operation the host answered did not produce a result: the route's own code
/// Why a custom operation the host answered did not produce a result: the route's own code
/// rejected it, which the route classifies like any other native failure, or Python code
/// raised, which reaches the caller as it was raised.
#[derive(Debug)]
@ -100,45 +100,54 @@ impl<E> From<PyErr> for InvokeError<E> {
}
}
/// The Python side of one route: answers the route's own operations, builds the public
/// The Python side of one protocol: answers its custom operations, builds the public
/// response and classifies native failures into public exceptions.
pub trait RouteHost: Send + Sync {
type Route: Route<Error: std::fmt::Display>;
pub trait ProtocolHost: Send + Sync {
type Protocol: Protocol<Error: std::fmt::Display>;
/// The public exception a native failure maps to, kept as a value until the driver
/// raises it.
type Failure: Into<PyErr>;
/// `arguments` is the keyword view the lifecycle's `begin` produced, not the
/// caller's own dict. A route host that projects from it inherits whatever that
/// adapter rewrote.
fn invoke(
/// Projects the call's request. `arguments` is the keyword view the lifecycle's
/// `begin` produced, not the caller's own dict, so the projection inherits whatever
/// that adapter rewrote.
fn project(
&mut self,
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
op: <Self::Route as Route>::Op,
) -> Result<<Self::Route as Route>::OpResult, InvokeError<<Self::Route as Route>::Error>>;
) -> Result<
<Self::Protocol as Protocol>::Projection,
InvokeError<<Self::Protocol as Protocol>::Error>,
>;
/// Answers `op` through its reply.
fn invoke(
&mut self,
py: Python<'_>,
op: <Self::Protocol as Protocol>::Op,
) -> Result<(), InvokeError<<Self::Protocol as Protocol>::Error>>;
fn complete(
&mut self,
py: Python<'_>,
response: <Self::Route as Route>::Response,
response: <Self::Protocol as Protocol>::Response,
) -> PyResult<Py<PyAny>>;
/// One streamed chunk as the caller receives it.
fn chunk(
&mut self,
py: Python<'_>,
chunk: <Self::Route as Route>::Chunk,
chunk: <Self::Protocol as Protocol>::Chunk,
) -> PyResult<Py<PyAny>>;
fn classify(
&self,
py: Python<'_>,
error: <Self::Route as Route>::Error,
error: <Self::Protocol as Protocol>::Error,
) -> PyResult<Self::Failure>;
fn host_error(error: &PyErr) -> <Self::Route as Route>::Error;
fn host_error(error: &PyErr) -> <Self::Protocol as Protocol>::Error;
fn close(&mut self, py: Python<'_>);

View file

@ -2,10 +2,11 @@ use std::sync::Arc;
use std::task::Poll;
use futures_util::future::{AbortHandle, Abortable};
use litellm_host::event::WireRequest;
use litellm_host::event::{FailureOrigin, Timing, epoch_seconds};
use litellm_host::host::{Demand, HostOp, HostResult, HostStep};
use litellm_host::host::{Demand, HostOp, HostStep, Reply};
use litellm_host::machine::{HostFailure, Machine, MachineStep};
use litellm_host::route::Route;
use litellm_host::protocol::Protocol;
use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError};
use pyo3::gc::{PyTraverseError, PyVisit};
use pyo3::prelude::*;
@ -13,21 +14,21 @@ use pyo3::types::PyDict;
use tokio::sync::Mutex;
use crate::adapter::{
InvokeError, LifecycleEvent, LifecycleStep, PythonLifecycle, RouteHost, missing_state,
InvokeError, LifecycleEvent, LifecycleStep, ProtocolHost, PythonLifecycle, missing_state,
};
use crate::execution::{poll_async_value, run_async_value, run_sync_value};
use crate::handle::{Execution, ExecutionBody, ExecutionStep};
type RouteOf<H> = <H as RouteHost>::Route;
type ErrorOf<H> = <RouteOf<H> as Route>::Error;
type ResponseOf<H> = <RouteOf<H> as Route>::Response;
type NativeStep<H> = MachineStep<RouteOf<H>, ResponseOf<H>>;
type ProtocolOf<H> = <H as ProtocolHost>::Protocol;
type ErrorOf<H> = <ProtocolOf<H> as Protocol>::Error;
type ResponseOf<H> = <ProtocolOf<H> as Protocol>::Response;
type NativeStep<H> = MachineStep<ProtocolOf<H>, ResponseOf<H>>;
type NativeResult<H> = Result<NativeStep<H>, ErrorOf<H>>;
type NativeResume<H> = Option<Result<HostResult<RouteOf<H>>, HostFailure<ErrorOf<H>>>>;
type Interruption<H> = Option<HostFailure<ErrorOf<H>>>;
type MachineResult<M> = Result<
MachineStep<<M as Machine>::Route, <M as Machine>::Complete>,
<<M as Machine>::Route as Route>::Error,
MachineStep<<M as Machine>::Protocol, <M as Machine>::Complete>,
<<M as Machine>::Protocol as Protocol>::Error,
>;
struct MachineState<M: Machine> {
@ -44,12 +45,11 @@ enum Stage {
Failed(Py<PyBaseException>),
}
#[derive(Clone, Copy)]
enum Expect {
Started,
Arguments,
Wire,
Emitted,
Wire(Reply<WireRequest>),
Emitted(Reply<()>),
Response,
Terminal,
}
@ -58,20 +58,30 @@ enum Pending {
Native,
Adapter(Expect),
/// The stream handed to the caller waits for its next read or its close.
Consumer,
Consumer(Reply<Demand>),
}
enum Next<H: RouteHost> {
/// A route answer as the driver resumes on it: a Python exception interrupts the call as
/// raised, a native rejection resumes the machine with it.
fn answered<E>(answer: Result<(), InvokeError<E>>) -> PyResult<Result<(), E>> {
match answer {
Ok(()) => Ok(Ok(())),
Err(InvokeError::Native(error)) => Ok(Err(error)),
Err(InvokeError::Python(error)) => Err(error),
}
}
enum Next<H: ProtocolHost> {
Return(ExecutionStep),
Continue(HostStep<NativeResult<H>, Py<PyAny>>),
}
struct PythonDriver<H, M>
where
H: RouteHost,
M: Machine<Route = H::Route, Complete = ResponseOf<H>> + 'static,
H: ProtocolHost,
M: Machine<Protocol = H::Protocol, Complete = ResponseOf<H>> + 'static,
{
route: H,
host: H,
adapter: Box<dyn PythonLifecycle>,
machine: Option<Arc<Mutex<MachineState<M>>>>,
arguments: Option<Py<PyDict>>,
@ -89,17 +99,17 @@ where
pub fn run_call<H, M>(
py: Python<'_>,
machine: M,
route: H,
host: H,
adapter: Box<dyn PythonLifecycle>,
arguments: Py<PyDict>,
asynchronous: bool,
) -> PyResult<Py<PyAny>>
where
H: RouteHost + 'static,
M: Machine<Route = H::Route, Complete = ResponseOf<H>> + 'static,
H: ProtocolHost + 'static,
M: Machine<Protocol = H::Protocol, Complete = ResponseOf<H>> + 'static,
{
let mut driver = PythonDriver {
route,
host,
adapter,
machine: Some(Arc::new(Mutex::new(MachineState {
machine,
@ -141,8 +151,8 @@ fn is_cancellation(py: Python<'_>, error: &PyErr) -> bool {
impl<H, M> PythonDriver<H, M>
where
H: RouteHost,
M: Machine<Route = H::Route, Complete = ResponseOf<H>> + 'static,
H: ProtocolHost,
M: Machine<Protocol = H::Protocol, Complete = ResponseOf<H>> + 'static,
{
fn timing(&self) -> Timing {
Timing {
@ -172,13 +182,13 @@ where
self.run_steps(py, HostStep::Ready(result))
}
(Some(Pending::Native), Some(Err(error))) => self.interrupt(py, error),
(Some(Pending::Consumer), Some(read)) => {
let demand = if read.is_ok() {
(Some(Pending::Consumer(reply)), Some(read)) => {
reply.send(if read.is_ok() {
Demand::More
} else {
Demand::Detached
};
self.resume_machine(py, Some(Ok(HostResult::Demand(demand))))
});
self.resume_machine(py, None)
}
(Some(Pending::Adapter(expect)), Some(result)) => {
match self.adapter.resume(py, result) {
@ -196,22 +206,24 @@ where
step: LifecycleStep,
expect: Expect,
) -> PyResult<ExecutionStep> {
if let LifecycleStep::Await(awaitable) = step {
self.pending = Some(Pending::Adapter(expect));
return Ok(ExecutionStep::Await(awaitable));
}
match (expect, step) {
(_, LifecycleStep::Await(awaitable)) => {
self.pending = Some(Pending::Adapter(expect));
Ok(ExecutionStep::Await(awaitable))
}
(Expect::Started, LifecycleStep::Done) => self.begin(py),
(Expect::Arguments, LifecycleStep::Arguments(arguments)) => {
self.arguments = Some(arguments);
self.stage = Stage::Call;
self.resume_machine(py, None)
}
(Expect::Wire, LifecycleStep::Wire(wire)) => {
self.resume_machine(py, Some(Ok(HostResult::BeforeSend(wire))))
(Expect::Wire(reply), LifecycleStep::Wire(wire)) => {
reply.send(*wire);
self.resume_machine(py, None)
}
(Expect::Emitted, LifecycleStep::Done) => {
self.resume_machine(py, Some(Ok(HostResult::Emitted)))
(Expect::Emitted(reply), LifecycleStep::Done) => {
reply.send(());
self.resume_machine(py, None)
}
(Expect::Response, LifecycleStep::Response(response)) => self.succeeded(py, response),
(Expect::Terminal, LifecycleStep::Done) => match &self.stage {
@ -242,9 +254,9 @@ where
fn resume_machine(
&mut self,
py: Python<'_>,
result: NativeResume<H>,
interruption: Interruption<H>,
) -> PyResult<ExecutionStep> {
let step = self.resume_core(py, result)?;
let step = self.resume_core(py, interruption)?;
self.run_steps(py, step)
}
@ -277,53 +289,62 @@ where
}
Err(error) => return self.machine_failed(py, error).map(Next::Return),
};
let answer = match op {
HostOp::Route(op) => {
let answered = match op {
HostOp::Project(reply) => {
let arguments = self.arguments.as_ref().ok_or_else(missing_state)?;
match self.route.invoke(py, arguments.bind(py), op) {
Ok(result) => Ok(HostResult::Route(result)),
Err(InvokeError::Native(error)) => {
return self
.resume_core(py, Some(Err(HostFailure::Error(error))))
.map(Next::Continue);
}
Err(InvokeError::Python(error)) => Err(error),
}
let projected = self.host.project(py, arguments.bind(py));
answered(projected.map(|projection| reply.send(projection)))
}
HostOp::BeforeSend { wire, context } => {
match self.adapter.before_send(py, wire, &context) {
Ok(LifecycleStep::Wire(wire)) => Ok(HostResult::BeforeSend(wire)),
HostOp::Custom(op) => answered(self.host.invoke(py, op)),
HostOp::BeforeSend {
wire,
context,
reply,
} => match self.adapter.before_send(py, wire, &context) {
Ok(LifecycleStep::Wire(wire)) => {
reply.send(*wire);
Ok(Ok(()))
}
Ok(LifecycleStep::Await(awaitable)) => {
self.pending = Some(Pending::Adapter(Expect::Wire(reply)));
return Ok(Next::Return(ExecutionStep::Await(awaitable)));
}
Ok(_) => return Err(missing_state()),
Err(error) => Err(error),
},
HostOp::Open(_, reply) => return self.opened(py, reply).map(Next::Return),
HostOp::Deliver(chunk, reply) => {
return self.delivered(py, chunk, reply).map(Next::Return);
}
HostOp::Emit(event, reply) => {
match self.adapter.emit(py, LifecycleEvent::Machine(&event)) {
Ok(LifecycleStep::Done) => {
reply.send(());
Ok(Ok(()))
}
Ok(LifecycleStep::Await(awaitable)) => {
self.pending = Some(Pending::Adapter(Expect::Wire));
self.pending = Some(Pending::Adapter(Expect::Emitted(reply)));
return Ok(Next::Return(ExecutionStep::Await(awaitable)));
}
Ok(_) => return Err(missing_state()),
Err(error) => Err(error),
}
}
HostOp::Open(_) => return self.opened(py).map(Next::Return),
HostOp::Deliver(chunk) => return self.delivered(py, chunk).map(Next::Return),
HostOp::Emit(event) => match self.adapter.emit(py, LifecycleEvent::Machine(&event)) {
Ok(LifecycleStep::Done) => Ok(HostResult::Emitted),
Ok(LifecycleStep::Await(awaitable)) => {
self.pending = Some(Pending::Adapter(Expect::Emitted));
return Ok(Next::Return(ExecutionStep::Await(awaitable)));
}
Ok(_) => return Err(missing_state()),
Err(error) => Err(error),
},
};
match answer {
Ok(answer) => self.resume_core(py, Some(Ok(answer))).map(Next::Continue),
match answered {
Ok(Ok(())) => self.resume_core(py, None).map(Next::Continue),
Ok(Err(native)) => self
.resume_core(py, Some(HostFailure::Error(native)))
.map(Next::Continue),
Err(error) => self.interrupt(py, error).map(Next::Return),
}
}
fn opened(&mut self, py: Python<'_>) -> PyResult<ExecutionStep> {
fn opened(&mut self, py: Python<'_>, reply: Reply<Demand>) -> PyResult<ExecutionStep> {
self.stage = Stage::Streaming;
match self.adapter.opened(py) {
Ok(()) => {
self.pending = Some(Pending::Consumer);
self.pending = Some(Pending::Consumer(reply));
Ok(ExecutionStep::Open)
}
Err(error) => self.interrupt(py, error),
@ -333,15 +354,16 @@ where
fn delivered(
&mut self,
py: Python<'_>,
chunk: <RouteOf<H> as Route>::Chunk,
chunk: <ProtocolOf<H> as Protocol>::Chunk,
reply: Reply<Demand>,
) -> PyResult<ExecutionStep> {
let chunk = match self.route.chunk(py, chunk) {
let chunk = match self.host.chunk(py, chunk) {
Ok(chunk) => chunk,
Err(error) => return self.interrupt(py, error),
};
match self.adapter.delivered(py, &chunk) {
Ok(()) => {
self.pending = Some(Pending::Consumer);
self.pending = Some(Pending::Consumer(reply));
Ok(ExecutionStep::Yield(chunk))
}
Err(error) => self.interrupt(py, error),
@ -357,25 +379,24 @@ where
} else {
HostFailure::Error(native)
};
self.resume_machine(py, Some(Err(failure)))
self.resume_machine(py, Some(failure))
}
fn resume_core(
&mut self,
py: Python<'_>,
result: NativeResume<H>,
interruption: Interruption<H>,
) -> PyResult<HostStep<NativeResult<H>, Py<PyAny>>> {
let state = Arc::clone(self.machine.as_ref().ok_or_else(missing_state)?);
let future = async move {
let mut state = state.lock().await;
let result = match result {
Some(Err(failure)) => state
let result = match interruption {
Some(failure) => state
.machine
.interrupt(failure)
.await
.map(MachineStep::Complete),
Some(Ok(result)) => state.machine.resume(Some(result)).await,
None => state.machine.resume(None).await,
None => state.machine.resume().await,
};
state.result = Some(result);
Ok(())
@ -414,7 +435,7 @@ where
fn completed(&mut self, py: Python<'_>, response: ResponseOf<H>) -> PyResult<ExecutionStep> {
self.ended_at = Some(epoch_seconds());
let public = match self.route.complete(py, response) {
let public = match self.host.complete(py, response) {
Ok(public) => public,
Err(error) => return self.failure(py, error, FailureOrigin::Call),
};
@ -441,7 +462,7 @@ where
/// fails, that failure is raised with the native error's text as its `__context__`.
fn classified(&self, py: Python<'_>, error: ErrorOf<H>) -> PyErr {
let native = error.to_string();
let classifier_error = match self.route.classify(py, error) {
let classifier_error = match self.host.classify(py, error) {
Ok(failure) => return failure.into(),
Err(classifier_error) => classifier_error,
};
@ -486,7 +507,7 @@ where
if self.machine.take().is_some() {
Python::attach(|py| {
self.adapter.close(py);
self.route.close(py);
self.host.close(py);
});
}
}
@ -494,15 +515,15 @@ where
impl<H, M> ExecutionBody for PythonDriver<H, M>
where
H: RouteHost,
M: Machine<Route = H::Route, Complete = ResponseOf<H>> + 'static,
H: ProtocolHost,
M: Machine<Protocol = H::Protocol, Complete = ResponseOf<H>> + 'static,
{
fn resume(&mut self, result: Option<PyResult<Py<PyAny>>>) -> PyResult<ExecutionStep> {
Python::attach(|py| self.drive(py, result))
}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
self.route.traverse(visit)?;
self.host.traverse(visit)?;
self.adapter.traverse(visit)?;
visit.call(&self.arguments)?;
visit.call(&self.interrupted)?;
@ -516,8 +537,8 @@ where
impl<H, M> Drop for PythonDriver<H, M>
where
H: RouteHost,
M: Machine<Route = H::Route, Complete = ResponseOf<H>> + 'static,
H: ProtocolHost,
M: Machine<Protocol = H::Protocol, Complete = ResponseOf<H>> + 'static,
{
fn drop(&mut self) {
self.clear();
@ -528,8 +549,8 @@ where
mod tests {
use std::sync::{Arc, Mutex};
use litellm_host::event::{MachineEvent, RequestContext, WireRequest};
use litellm_host::machine::{Interrupted, Step};
use litellm_host::event::{MachineEvent, RawResponse, RequestContext};
use litellm_host::machine::{CallMachine, MachineFault};
use pyo3::exceptions::{PyBaseException, PyValueError};
use pyo3::types::PyDict;
@ -573,22 +594,21 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
}
}
struct Synthetic;
impl Route for Synthetic {
type Response = String;
type Error = Error;
type Op = &'static str;
type OpResult = String;
type Chunk = std::convert::Infallible;
type StreamHead = std::convert::Infallible;
impl From<MachineFault> for Error {
fn from(fault: MachineFault) -> Self {
Self(format!("{fault:?}"))
}
}
/// Yields the scripted ops in order, then completes or fails as scripted.
struct ScriptedMachine {
ops: Vec<HostOp<Synthetic>>,
outcome: Option<Result<String, Error>>,
answers: Vec<String>,
struct Synthetic;
impl Protocol for Synthetic {
type Response = String;
type Error = Error;
type Projection = String;
type Op = (&'static str, Reply<String>);
type Chunk = std::convert::Infallible;
type StreamHead = std::convert::Infallible;
}
fn wire() -> WireRequest {
@ -609,37 +629,6 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
}
}
impl Machine for ScriptedMachine {
type Route = Synthetic;
type Complete = String;
fn resume(&mut self, result: Option<HostResult<Synthetic>>) -> Step<'_, Self> {
Box::pin(async move {
if let Some(result) = result {
self.answers.push(match result {
HostResult::Route(value) => value,
HostResult::BeforeSend(wire) => wire.url,
HostResult::Emitted => "emitted".into(),
HostResult::Demand(demand) => format!("{demand:?}"),
});
}
if !self.ops.is_empty() {
return Ok(MachineStep::Host(self.ops.remove(0)));
}
self.outcome
.take()
.ok_or_else(|| Error("resumed after completion".into()))?
.map(MachineStep::Complete)
})
}
fn interrupt(&mut self, failure: HostFailure<Error>) -> Interrupted<'_, Self> {
self.ops.clear();
self.outcome = None;
Box::pin(async move { Err(failure.into_error()) })
}
}
#[derive(Default)]
struct Log(Arc<Mutex<Vec<String>>>);
@ -677,22 +666,37 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
}
}
impl RouteHost for SyntheticHost {
type Route = Synthetic;
impl SyntheticHost {
fn answer(&self, value: impl FnOnce() -> String) -> Result<String, InvokeError<Error>> {
match self.op {
OpScript::Answer => Ok(value()),
OpScript::RaisePython => Err(PyValueError::new_err("op failed").into()),
OpScript::RejectNatively => Err(InvokeError::Native(Error("op rejected".into()))),
}
}
}
impl ProtocolHost for SyntheticHost {
type Protocol = Synthetic;
type Failure = Classified;
fn project(
&mut self,
_: Python<'_>,
arguments: &Bound<'_, PyDict>,
) -> Result<String, InvokeError<Error>> {
self.log.push("project");
self.answer(|| format!("project:{}", arguments.len()))
}
fn invoke(
&mut self,
_: Python<'_>,
arguments: &Bound<'_, PyDict>,
op: &'static str,
) -> Result<String, InvokeError<Error>> {
self.log.push(format!("route:{op}"));
match self.op {
OpScript::Answer => Ok(format!("{op}:{}", arguments.len())),
OpScript::RaisePython => Err(PyValueError::new_err("op failed").into()),
OpScript::RejectNatively => Err(InvokeError::Native(Error("op rejected".into()))),
}
(op, reply): (&'static str, Reply<String>),
) -> Result<(), InvokeError<Error>> {
self.log.push(format!("op:{op}"));
self.answer(|| op.to_string())
.map(|answer| reply.send(answer))
}
fn chunk(&mut self, _: Python<'_>, chunk: std::convert::Infallible) -> PyResult<Py<PyAny>> {
@ -719,7 +723,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
}
fn close(&mut self, _: Python<'_>) {
self.log.push("route.close");
self.log.push("host.close");
}
fn traverse(&self, _: &PyVisit<'_>) -> Result<(), PyTraverseError> {
@ -828,7 +832,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
fn run_scripted(
py: Python<'_>,
machine: ScriptedMachine,
machine: CallMachine<Synthetic>,
op: OpScript,
script: AdapterScript,
asynchronous: bool,
@ -848,12 +852,12 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
fn run_hosted(
py: Python<'_>,
machine: ScriptedMachine,
route: SyntheticHost,
machine: CallMachine<Synthetic>,
host: SyntheticHost,
script: AdapterScript,
asynchronous: bool,
) -> (PyResult<Py<PyAny>>, Vec<String>) {
let log = Log(route.log.0.clone());
let log = Log(host.log.0.clone());
let adapter = SyntheticAdapter {
log: Log(log.0.clone()),
script,
@ -863,7 +867,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
let result = run_call(
py,
machine,
route,
host,
Box::new(adapter),
arguments.unbind(),
asynchronous,
@ -884,21 +888,21 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
(result, log.entries())
}
fn success_machine() -> ScriptedMachine {
ScriptedMachine {
ops: vec![
HostOp::Route("project"),
HostOp::BeforeSend {
wire: Box::new(wire()),
context: Box::new(context()),
},
HostOp::Emit(MachineEvent::ResponseReceived {
raw: litellm_host::event::RawResponse { body: "raw".into() },
}),
],
outcome: Some(Ok("done".into())),
answers: Vec::new(),
}
/// Answers to projection, to the route op and to `before_send` all reach the
/// response, so a driver that misroutes a reply changes what the call returns.
fn success_machine() -> CallMachine<Synthetic> {
CallMachine::new(|host| {
Box::pin(async move {
let projected = host.project().await?;
let signed = host.custom_op(|reply| ("sign", reply)).await?;
let wire = host.before_send(wire(), context()).await?;
host.emit(MachineEvent::ResponseReceived {
raw: RawResponse { body: "raw".into() },
})
.await?;
Ok(format!("{projected}|{signed}|{}", wire.url))
})
})
}
#[test]
@ -917,32 +921,37 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
AdapterScript::Plain,
asynchronous,
);
assert_eq!(result.unwrap().extract::<String>(py).unwrap(), "done");
assert_eq!(
result.unwrap().extract::<String>(py).unwrap(),
"project:1|sign|rewritten"
);
assert_eq!(
log,
[
"started",
"begin",
"route:project",
"project",
"op:sign",
"before_send",
"response:raw",
"complete",
"after_success",
"succeeded:done",
"succeeded:project:1|sign|rewritten",
"adapter.close",
"route.close",
"host.close",
]
);
}
});
}
fn failing_machine() -> ScriptedMachine {
ScriptedMachine {
ops: vec![HostOp::Route("project")],
outcome: Some(Err(Error("provider exploded".into()))),
answers: Vec::new(),
}
fn failing_machine() -> CallMachine<Synthetic> {
CallMachine::new(|host| {
Box::pin(async move {
host.project().await?;
Err(Error("provider exploded".into()))
})
})
}
#[test]
@ -969,11 +978,11 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
[
"started",
"begin",
"route:project",
"project",
"classify:provider exploded",
"failed:Call:classified: provider exploded",
"adapter.close",
"route.close",
"host.close",
]
);
}
@ -1003,11 +1012,11 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
[
"started",
"begin",
"route:project",
"project",
"classify:op rejected",
"failed:Call:classified: op rejected",
"adapter.close",
"route.close",
"host.close",
]
);
});
@ -1035,10 +1044,10 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
[
"started",
"begin",
"route:project",
"project",
"failed:Call:op failed",
"adapter.close",
"route.close",
"host.close",
]
);
});
@ -1073,11 +1082,11 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
[
"started",
"begin",
"route:project",
"project",
"classify:provider exploded",
"failed:Call:classifier failed",
"adapter.close",
"route.close",
"host.close",
]
);
});
@ -1106,7 +1115,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
"begin",
"failed:Host:begin failed",
"adapter.close",
"route.close"
"host.close"
]
);
});
@ -1130,7 +1139,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
);
assert_eq!(result.unwrap().extract::<String>(py).unwrap(), "replaced");
assert!(log.contains(&"succeeded:replaced".to_string()));
assert!(!log.contains(&"succeeded:done".to_string()));
assert!(!log.contains(&"succeeded:project:1|rewritten".to_string()));
}
});
}
@ -1159,7 +1168,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
"after_success",
"failed:Host:after_success failed",
"adapter.close",
"route.close"
"host.close"
]
);
assert!(!log.iter().any(|entry| entry.starts_with("succeeded")));
@ -1175,18 +1184,24 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
crate::initialize_python();
Python::attach(|py| {
struct Cancelling(Log);
impl RouteHost for Cancelling {
type Route = Synthetic;
impl ProtocolHost for Cancelling {
type Protocol = Synthetic;
type Failure = Classified;
fn invoke(
fn project(
&mut self,
_: Python<'_>,
_: &Bound<'_, PyDict>,
_: &'static str,
) -> Result<String, InvokeError<Error>> {
self.0.push("route");
self.0.push("project");
Err(pyo3::exceptions::asyncio::CancelledError::new_err(()).into())
}
fn invoke(
&mut self,
_: Python<'_>,
_: (&'static str, Reply<String>),
) -> Result<(), InvokeError<Error>> {
Err(missing_state().into())
}
fn chunk(
&mut self,
_: Python<'_>,
@ -1210,7 +1225,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
}
}
let log = Log::default();
let route = Cancelling(Log(log.0.clone()));
let host = Cancelling(Log(log.0.clone()));
let adapter = SyntheticAdapter {
log: Log(log.0.clone()),
script: AdapterScript::Plain,
@ -1218,7 +1233,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
let error = run_call(
py,
success_machine(),
route,
host,
Box::new(adapter),
PyDict::new(py).unbind(),
false,
@ -1227,7 +1242,7 @@ sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bri
assert!(!error.is_instance_of::<pyo3::exceptions::PyException>(py));
assert_eq!(
log.entries(),
["started", "begin", "route", "adapter.close"]
["started", "begin", "project", "adapter.close"]
);
});
}

View file

@ -0,0 +1,241 @@
//! A caller's file-like object: anything with a callable `read`, kept as a handle and read
//! once, on the host's thread, into bytes Rust owns.
use bytes::Bytes;
use pyo3::{
exceptions::PyTypeError,
gc::{PyTraverseError, PyVisit},
prelude::*,
pybacked::PyBackedBytes,
types::{PyBytes, PyString},
};
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct FileContent {
pub bytes: Bytes,
pub file_name: Option<String>,
}
#[derive(Debug)]
pub struct PythonFileReader {
reader: Py<PyAny>,
name: Option<String>,
}
impl PythonFileReader {
/// `None` when `file` has no callable `read`. The object's `name` is read now, its
/// contents only on [`read`](Self::read).
pub fn from_file_like(file: &Bound<'_, PyAny>) -> PyResult<Option<Self>> {
let reader = file
.getattr_opt("read")?
.filter(|value| value.is_callable());
let Some(reader) = reader else {
return Ok(None);
};
let name = file
.getattr_opt("name")?
.filter(|value| !value.is_none())
.map(|value| value.extract::<String>())
.transpose()?;
Ok(Some(Self {
reader: reader.unbind(),
name,
}))
}
pub fn read(&self, py: Python<'_>) -> PyResult<FileContent> {
let value = self.reader.bind(py).call0()?;
let bytes = if value.is_instance_of::<PyString>() {
Bytes::from(value.extract::<String>()?)
} else if value.is_instance_of::<PyBytes>() {
py_bytes(&value)?
} else {
return Err(PyTypeError::new_err(format!(
"file read must return bytes or str, got {}",
value.get_type(),
)));
};
Ok(FileContent {
bytes,
file_name: self.name.clone(),
})
}
pub fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.reader)
}
}
/// An exact `bytes` object is shared without copying and keeps the Python object alive;
/// a `bytes` subclass is copied.
pub fn py_bytes(value: &Bound<'_, PyAny>) -> PyResult<Bytes> {
if value.is_exact_instance_of::<PyBytes>() {
return Ok(Bytes::from_owner(value.extract::<PyBackedBytes>()?));
}
Ok(Bytes::copy_from_slice(
value.extract::<PyBackedBytes>()?.as_ref(),
))
}
#[cfg(test)]
mod tests {
use pyo3::{exceptions::PyTypeError, types::PyDict};
use super::*;
fn eval<'py>(py: Python<'py>, source: &std::ffi::CStr) -> Bound<'py, PyDict> {
let locals = PyDict::new(py);
py.run(source, Some(&locals), Some(&locals)).unwrap();
locals
}
fn reader<'py>(locals: &Bound<'py, PyDict>, name: &str) -> PythonFileReader {
PythonFileReader::from_file_like(&locals.get_item(name).unwrap().unwrap())
.unwrap()
.unwrap()
}
#[test]
fn objects_without_a_callable_read_are_not_readers() {
Python::initialize();
Python::attach(|py| {
let locals = eval(
py,
c"
class Attribute:
read = 'not callable'
plain = object()
attribute = Attribute()
",
);
for name in ["plain", "attribute"] {
let file = locals.get_item(name).unwrap().unwrap();
assert!(PythonFileReader::from_file_like(&file).unwrap().is_none());
}
});
}
#[test]
fn the_name_is_taken_up_front_and_the_contents_only_on_read() {
Python::initialize();
Python::attach(|py| {
let locals = eval(
py,
c"
class Reader:
name = 'scan.png'
def __init__(self):
self.reads = 0
def read(self):
self.reads += 1
return b'abc'
file = Reader()
",
);
let reads = || {
locals
.get_item("file")
.unwrap()
.unwrap()
.getattr("reads")
.unwrap()
.extract::<usize>()
.unwrap()
};
let file = reader(&locals, "file");
assert_eq!(reads(), 0);
let content = file.read(py).unwrap();
assert_eq!(reads(), 1);
assert_eq!(
content,
FileContent {
bytes: b"abc".as_slice().into(),
file_name: Some("scan.png".into()),
}
);
});
}
#[test]
fn read_results_are_normalized_and_exceptions_keep_their_identity() {
Python::initialize();
Python::attach(|py| {
let locals = eval(
py,
c"
failure = KeyError('reader failed')
class Raising:
def read(self):
raise failure
class Text:
def read(self):
return 'héllo'
class Wrong:
def read(self):
return 7
raising = Raising()
text = Text()
wrong = Wrong()
",
);
let error = reader(&locals, "raising").read(py).unwrap_err();
assert!(
error
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
assert_eq!(
reader(&locals, "text").read(py).unwrap().bytes.as_ref(),
"héllo".as_bytes()
);
let error = reader(&locals, "wrong").read(py).unwrap_err();
assert!(error.is_instance_of::<PyTypeError>(py));
assert!(error.to_string().contains("bytes or str"));
});
}
#[rstest::rstest]
#[case::read("read")]
#[case::name("name")]
fn attribute_failures_keep_their_identity(#[case] attribute: &str) {
Python::initialize();
Python::attach(|py| {
let locals = eval(
py,
c"
failure = LookupError('file property failed')
class File:
def __getattribute__(self, name):
if name == attribute:
raise failure
return super().__getattribute__(name)
name = 'scan.pdf'
def read(self):
return b'abc'
file = File()
",
);
locals.set_item("attribute", attribute).unwrap();
let error =
PythonFileReader::from_file_like(&locals.get_item("file").unwrap().unwrap())
.unwrap_err();
assert!(
error
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
});
}
#[test]
fn exact_python_bytes_transfer_without_copying_and_outlive_the_input() {
Python::initialize();
let (bytes, pointer) = Python::attach(|py| {
let value = PyBytes::new(py, b"document bytes");
let pointer = value.as_bytes().as_ptr() as usize;
(py_bytes(value.as_any()).unwrap(), pointer)
});
assert_eq!(bytes.as_ptr() as usize, pointer);
assert_eq!(bytes.as_ref(), b"document bytes");
}
}

View file

@ -1,6 +1,6 @@
//! The CPython runtime adapter: value marshalling, interpreter detachment, the tokio and
//! asyncio glue, and the driver that runs a native [`Machine`](litellm_host::machine::Machine)
//! against a Python route host and a Python lifecycle. Everything here is Python-specific by
//! against a Python protocol host and a Python lifecycle. Everything here is Python-specific by
//! construction; another host language gets its own crate of the same shape.
mod adapter;
@ -8,13 +8,14 @@ mod argument;
mod callable;
mod driver;
mod execution;
mod file_reader;
mod fork_gate;
mod gil;
mod handle;
mod marshal;
pub use adapter::{
InvokeError, LifecycleEvent, LifecycleStep, PythonLifecycle, RouteHost, missing_state,
InvokeError, LifecycleEvent, LifecycleStep, ProtocolHost, PythonLifecycle, missing_state,
};
pub use argument::lookup;
pub use callable::wrap_failure;
@ -24,6 +25,7 @@ pub use execution::{
reserve_process_for_forking, run_async, run_async_value, run_sync, run_sync_value,
runtime_started,
};
pub use file_reader::{FileContent, PythonFileReader, py_bytes};
pub use fork_gate::RuntimeAlreadyStarted;
pub use gil::{PythonContext, attach_blocking, release_count, release_gil};
pub use handle::{Execution, ExecutionBody, ExecutionStep};

View file

@ -7,6 +7,7 @@ repository.workspace = true
[dependencies]
litellm-auth.workspace = true
litellm-coroutine.workspace = true
serde_json.workspace = true
tokio = { workspace = true, features = ["sync"] }

View file

@ -1,28 +1,27 @@
use std::future::Future;
use crate::event::{CallEvent, MachineEvent, RequestContext, WireRequest};
use crate::route::Route;
pub use litellm_coroutine::{Abandoned, Answer, Reply, reply};
/// One suspension point of a native call, performed by the host.
pub enum HostOp<R: Route> {
Route(R::Op),
use crate::event::{CallEvent, MachineEvent, RequestContext, WireRequest};
use crate::protocol::Protocol;
/// One suspension point of a native call, performed by the host and answered through the
/// [`Reply`] it carries.
pub enum HostOp<R: Protocol> {
/// The first op of every call: the caller's request as the host projects it.
Project(Reply<R::Projection>),
Custom(R::Op),
BeforeSend {
wire: Box<WireRequest>,
context: Box<RequestContext>,
reply: Reply<WireRequest>,
},
Emit(MachineEvent),
Emit(MachineEvent, Reply<()>),
/// The response streams: the host hands the caller a stream and answers once the
/// caller asks for the first chunk or goes away.
Open(R::StreamHead),
Open(R::StreamHead, Reply<Demand>),
/// The next chunk of an open stream, answered once the caller asks for the one after.
Deliver(R::Chunk),
}
pub enum HostResult<R: Route> {
Route(R::OpResult),
BeforeSend(Box<WireRequest>),
Emitted,
Demand(Demand),
Deliver(R::Chunk, Reply<Demand>),
}
/// Whether the caller of a streamed call still reads it.
@ -39,10 +38,13 @@ pub enum HostStep<V, S> {
Suspend(S),
}
/// An in-process host: answers route operations and observes the call without leaving
/// An in-process host: answers custom operations and observes the call without leaving
/// the Rust runtime. Language hosts implement their own driver instead.
pub trait Host<R: Route>: Send + Sync {
fn route(&self, op: R::Op) -> impl Future<Output = Result<R::OpResult, R::Error>> + Send;
pub trait Host<R: Protocol>: Send + Sync {
fn project(&self) -> impl Future<Output = Result<R::Projection, R::Error>> + Send;
/// Answers `op` through its reply, or fails the call.
fn custom_op(&self, op: R::Op) -> impl Future<Output = Result<(), R::Error>> + Send;
fn before_send(
&self,

View file

@ -1,12 +1,13 @@
//! The contract between a native call and the host runtime that drives it.
//!
//! A host is whatever sits on the far side of the language boundary: CPython today,
//! another runtime later. Core runs each route on a [`machine::RouteMachine`] and never learns
//! another runtime later. Core runs each route on a [`machine::CallMachine`] and never learns
//! which host is on the other end. The machine yields [`host::HostOp`]s; a driver answers
//! them, observes [`event::CallEvent`]s and may rewrite the wire request before it is sent.
//! each through the typed [`host::Reply`] it carries, observes [`event::CallEvent`]s and
//! may rewrite the wire request before it is sent.
pub mod event;
pub mod host;
pub mod machine;
pub mod route;
pub mod protocol;
pub mod run;

View file

@ -1,22 +1,21 @@
use std::sync::Arc;
use super::{HostChannel, MachineFault};
use crate::route::Route;
use crate::{host::Reply, protocol::Protocol};
use litellm_auth::{Error, ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle};
/// A route whose host can mint credentials on the call's behalf.
pub trait TokenRoute: Route {
fn acquire_token_op() -> Self::Op;
fn token_credential(result: Self::OpResult) -> Option<ResolvedCredential>;
/// A protocol whose host can mint credentials on the call's behalf.
pub trait TokenProtocol: Protocol {
fn acquire_token_op(reply: Reply<ResolvedCredential>) -> Self::Op;
}
/// A [`TokenProvider`] that asks the host for each credential through the call's own
/// operation channel, so the host answers it on the caller's thread and context.
pub struct HostTokenProvider<R: Route> {
pub struct HostTokenProvider<R: Protocol> {
channel: HostChannel<R>,
}
impl<R: Route> std::fmt::Debug for HostTokenProvider<R> {
impl<R: Protocol> std::fmt::Debug for HostTokenProvider<R> {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("HostTokenProvider")
}
@ -24,7 +23,7 @@ impl<R: Route> std::fmt::Debug for HostTokenProvider<R> {
impl<R> HostTokenProvider<R>
where
R: TokenRoute,
R: TokenProtocol,
R::Error: From<MachineFault> + std::fmt::Display,
{
pub fn handle(channel: HostChannel<R>) -> TokenProviderHandle {
@ -34,19 +33,15 @@ where
impl<R> TokenProvider for HostTokenProvider<R>
where
R: TokenRoute,
R: TokenProtocol,
R::Error: From<MachineFault> + std::fmt::Display,
{
fn acquire(&self) -> TokenFuture<'_> {
Box::pin(async move {
let result = self
.channel
.route(R::acquire_token_op())
self.channel
.custom_op(R::acquire_token_op)
.await
.map_err(|error| Error::AzureTokenAcquisition(error.to_string()))?;
R::token_credential(result).ok_or_else(|| {
Error::AzureTokenAcquisition("invalid token provider host result".into())
})
.map_err(|error| Error::AzureTokenAcquisition(error.to_string()))
})
}
}

View file

@ -0,0 +1,137 @@
//! The one machine every route runs on: the route's provider future as a
//! [`Coroutine`] that yields [`HostOp`]s, each answered through its own typed reply. No
//! task is spawned; dropping the machine drops the in-flight call.
use std::{future::Future, pin::Pin};
use litellm_coroutine::{Co, Coroutine, CoroutineState, ResumeError};
use super::{HostFailure, Interrupted, Machine, MachineStep, Step};
use crate::{
event::{MachineEvent, RequestContext, WireRequest},
host::{Demand, HostOp, Reply},
protocol::Protocol,
};
/// The machine's own failures, distinct from anything the provider call reports.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MachineFault {
/// The host dropped an op's reply unanswered, or went away while the call waited.
Abandoned,
/// The host resumed the call out of turn.
Protocol(ResumeError),
}
pub type ExecuteFuture<R> =
Pin<Box<dyn Future<Output = Result<<R as Protocol>::Response, <R as Protocol>::Error>> + Send>>;
/// The provider side of the machine: how the in-flight call reaches its host.
pub struct HostChannel<R: Protocol> {
co: Co<HostOp<R>>,
}
impl<R: Protocol> Clone for HostChannel<R> {
fn clone(&self) -> Self {
Self {
co: self.co.clone(),
}
}
}
impl<R: Protocol> HostChannel<R>
where
R::Error: From<MachineFault>,
{
async fn yield_<A: Send>(
&self,
ask: impl FnOnce(Reply<A>) -> HostOp<R> + Send,
) -> Result<A, R::Error> {
self.co
.yield_(ask)
.await
.map_err(|_| MachineFault::Abandoned.into())
}
pub async fn project(&self) -> Result<R::Projection, R::Error> {
self.yield_(HostOp::Project).await
}
/// Asks the host to perform the custom operation `ask` builds around its reply, as in
/// `host.custom_op(OcrOp::AcquireAzureAdToken)`.
pub async fn custom_op<A: Send>(
&self,
ask: impl FnOnce(Reply<A>) -> R::Op + Send,
) -> Result<A, R::Error> {
self.yield_(|reply| HostOp::Custom(ask(reply))).await
}
pub async fn before_send(
&self,
wire: WireRequest,
context: RequestContext,
) -> Result<WireRequest, R::Error> {
self.yield_(|reply| HostOp::BeforeSend {
wire: Box::new(wire),
context: Box::new(context),
reply,
})
.await
}
pub async fn emit(&self, event: MachineEvent) -> Result<(), R::Error> {
self.yield_(|reply| HostOp::Emit(event, reply)).await
}
pub async fn open(&self, head: R::StreamHead) -> Result<Demand, R::Error> {
self.yield_(|reply| HostOp::Open(head, reply)).await
}
pub async fn deliver(&self, chunk: R::Chunk) -> Result<Demand, R::Error> {
self.yield_(|reply| HostOp::Deliver(chunk, reply)).await
}
}
type CallCoroutine<R> =
Coroutine<HostOp<R>, Result<<R as Protocol>::Response, <R as Protocol>::Error>>;
pub struct CallMachine<R: Protocol> {
coroutine: CallCoroutine<R>,
}
impl<R: Protocol> CallMachine<R>
where
R::Error: From<MachineFault>,
{
pub fn new(execute: impl FnOnce(HostChannel<R>) -> ExecuteFuture<R> + Send + 'static) -> Self {
Self {
coroutine: Coroutine::new(|co| execute(HostChannel { co })),
}
}
}
impl<R: Protocol> Machine for CallMachine<R>
where
R::Error: From<MachineFault>,
{
type Protocol = R;
type Complete = R::Response;
fn resume(&mut self) -> Step<'_, Self> {
Box::pin(async move {
match self
.coroutine
.resume()
.await
.map_err(MachineFault::Protocol)?
{
CoroutineState::Yielded(op) => Ok(MachineStep::Host(op)),
CoroutineState::Complete(outcome) => outcome.map(MachineStep::Complete),
}
})
}
fn interrupt(&mut self, failure: HostFailure<R::Error>) -> Interrupted<'_, Self> {
self.coroutine.cancel();
Box::pin(async move { Err(failure.into_error()) })
}
}

View file

@ -1,16 +1,16 @@
mod auth;
mod route_machine;
mod call_machine;
use std::future::Future;
use std::pin::Pin;
pub use auth::{HostTokenProvider, TokenRoute};
pub use route_machine::{ExecuteFuture, HostChannel, MachineFault, RouteMachine};
pub use auth::{HostTokenProvider, TokenProtocol};
pub use call_machine::{CallMachine, ExecuteFuture, HostChannel, MachineFault};
use crate::host::{HostOp, HostResult};
use crate::route::Route;
use crate::host::HostOp;
use crate::protocol::Protocol;
pub enum MachineStep<R: Route, C> {
pub enum MachineStep<R: Protocol, C> {
Host(HostOp<R>),
Complete(C),
}
@ -19,8 +19,8 @@ pub type Step<'a, M> = Pin<
Box<
dyn Future<
Output = Result<
MachineStep<<M as Machine>::Route, <M as Machine>::Complete>,
<<M as Machine>::Route as Route>::Error,
MachineStep<<M as Machine>::Protocol, <M as Machine>::Complete>,
<<M as Machine>::Protocol as Protocol>::Error,
>,
> + Send
+ 'a,
@ -30,7 +30,10 @@ pub type Step<'a, M> = Pin<
pub type Interrupted<'a, M> = Pin<
Box<
dyn Future<
Output = Result<<M as Machine>::Complete, <<M as Machine>::Route as Route>::Error>,
Output = Result<
<M as Machine>::Complete,
<<M as Machine>::Protocol as Protocol>::Error,
>,
> + Send
+ 'a,
>,
@ -51,19 +54,18 @@ impl<E> HostFailure<E> {
}
/// A resumable call. Core implements it per route; a host drives it. Every suspension
/// point is an op the host performs and answers with a result.
/// point is an op the host performs and answers through the op's own reply before it
/// resumes the call again.
pub trait Machine: Send {
type Route: Route;
type Protocol: Protocol;
type Complete: Send + 'static;
/// `None` on the first call and whenever the previous step completed without
/// yielding an op; otherwise the result of the op last yielded.
fn resume(&mut self, result: Option<HostResult<Self::Route>>) -> Step<'_, Self>;
fn resume(&mut self) -> Step<'_, Self>;
/// The host failed to perform the pending op, or the caller cancelled. The call
/// yields no further ops.
fn interrupt(
&mut self,
failure: HostFailure<<Self::Route as Route>::Error>,
failure: HostFailure<<Self::Protocol as Protocol>::Error>,
) -> Interrupted<'_, Self>;
}

View file

@ -1,199 +0,0 @@
//! The one machine every route runs on: it owns the route's provider future, polls it in
//! place, and turns the host operations that future requests into [`Machine`] steps. No
//! task is spawned; dropping the machine drops the in-flight call.
use std::{future::Future, pin::Pin};
use tokio::sync::{mpsc, oneshot};
use super::{HostFailure, Interrupted, Machine, MachineStep, Step};
use crate::{
event::{MachineEvent, RequestContext, WireRequest},
host::{Demand, HostOp, HostResult},
route::Route,
};
/// The machine's own failures, distinct from anything the provider call reports.
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MachineFault {
/// The host driver went away while the call was waiting on it.
Abandoned,
/// The host answered out of turn: a result with nothing pending, or nothing when a
/// result was pending.
Protocol(&'static str),
/// The host answered a route operation with the wrong result variant.
Mismatch,
}
pub type ExecuteFuture<R> =
Pin<Box<dyn Future<Output = Result<<R as Route>::Response, <R as Route>::Error>> + Send>>;
struct PendingOp<R: Route> {
op: HostOp<R>,
reply: oneshot::Sender<HostResult<R>>,
}
/// The provider side of the machine: how the in-flight call reaches its host.
pub struct HostChannel<R: Route> {
ops: mpsc::UnboundedSender<PendingOp<R>>,
}
impl<R: Route> Clone for HostChannel<R> {
fn clone(&self) -> Self {
Self {
ops: self.ops.clone(),
}
}
}
impl<R: Route> HostChannel<R>
where
R::Error: From<MachineFault>,
{
async fn invoke(&self, op: HostOp<R>) -> Result<HostResult<R>, R::Error> {
let (reply, answer) = oneshot::channel();
self.ops
.send(PendingOp { op, reply })
.map_err(|_| MachineFault::Abandoned)?;
answer.await.map_err(|_| MachineFault::Abandoned.into())
}
pub async fn route(&self, op: R::Op) -> Result<R::OpResult, R::Error> {
match self.invoke(HostOp::Route(op)).await? {
HostResult::Route(result) => Ok(result),
_ => Err(MachineFault::Mismatch.into()),
}
}
pub async fn before_send(
&self,
wire: WireRequest,
context: RequestContext,
) -> Result<WireRequest, R::Error> {
let op = HostOp::BeforeSend {
wire: Box::new(wire),
context: Box::new(context),
};
match self.invoke(op).await? {
HostResult::BeforeSend(wire) => Ok(*wire),
_ => Err(MachineFault::Mismatch.into()),
}
}
pub async fn emit(&self, event: MachineEvent) -> Result<(), R::Error> {
match self.invoke(HostOp::Emit(event)).await? {
HostResult::Emitted => Ok(()),
_ => Err(MachineFault::Mismatch.into()),
}
}
pub async fn open(&self, head: R::StreamHead) -> Result<Demand, R::Error> {
self.demand(HostOp::Open(head)).await
}
pub async fn deliver(&self, chunk: R::Chunk) -> Result<Demand, R::Error> {
self.demand(HostOp::Deliver(chunk)).await
}
async fn demand(&self, op: HostOp<R>) -> Result<Demand, R::Error> {
match self.invoke(op).await? {
HostResult::Demand(demand) => Ok(demand),
_ => Err(MachineFault::Mismatch.into()),
}
}
}
enum Execution<R: Route> {
Unstarted(Box<dyn FnOnce(HostChannel<R>) -> ExecuteFuture<R> + Send>),
Running(ExecuteFuture<R>),
Done,
}
pub struct RouteMachine<R: Route> {
execution: Execution<R>,
ops: mpsc::UnboundedReceiver<PendingOp<R>>,
channel: HostChannel<R>,
reply: Option<oneshot::Sender<HostResult<R>>>,
}
impl<R: Route> RouteMachine<R>
where
R::Error: From<MachineFault>,
{
pub fn new(execute: impl FnOnce(HostChannel<R>) -> ExecuteFuture<R> + Send + 'static) -> Self {
let (ops_tx, ops) = mpsc::unbounded_channel();
Self {
execution: Execution::Unstarted(Box::new(execute)),
ops,
channel: HostChannel { ops: ops_tx },
reply: None,
}
}
async fn step(
&mut self,
result: Option<HostResult<R>>,
) -> Result<MachineStep<R, R::Response>, R::Error> {
match (self.reply.take(), result) {
(Some(reply), Some(result)) => {
reply
.send(result)
.map_err(|_| MachineFault::Protocol("the call stopped waiting on the host"))?;
}
(None, None) if matches!(self.execution, Execution::Unstarted(_)) => {}
(Some(reply), None) => {
self.reply = Some(reply);
return Err(MachineFault::Protocol("host operation result is required").into());
}
(None, Some(_)) => {
return Err(MachineFault::Protocol("unexpected host operation result").into());
}
(None, None) => {
return Err(
MachineFault::Protocol("call cannot be resumed after completion").into(),
);
}
}
if let Execution::Unstarted(_) = self.execution {
let Execution::Unstarted(start) =
std::mem::replace(&mut self.execution, Execution::Done)
else {
unreachable!()
};
self.execution = Execution::Running(start(self.channel.clone()));
}
let Execution::Running(future) = &mut self.execution else {
return Err(MachineFault::Protocol("call cannot be resumed after completion").into());
};
tokio::select! {
biased;
pending = self.ops.recv() => {
let pending = pending.ok_or(MachineFault::Abandoned)?;
self.reply = Some(pending.reply);
Ok(MachineStep::Host(pending.op))
}
outcome = future => {
self.execution = Execution::Done;
outcome.map(MachineStep::Complete)
}
}
}
}
impl<R: Route> Machine for RouteMachine<R>
where
R::Error: From<MachineFault>,
{
type Route = R;
type Complete = R::Response;
fn resume(&mut self, result: Option<HostResult<R>>) -> Step<'_, Self> {
Box::pin(self.step(result))
}
fn interrupt(&mut self, failure: HostFailure<R::Error>) -> Interrupted<'_, Self> {
self.reply = None;
self.execution = Execution::Done;
Box::pin(async move { Err(failure.into_error()) })
}
}

View file

@ -0,0 +1,17 @@
/// One public call surface: what a completed call produces, how it fails, what the host
/// projects the caller's request into, and the protocol-specific operations only its host
/// can perform mid-call (token acquisition, for one).
pub trait Protocol: Send + Sync + 'static {
type Response: Send + 'static;
type Error: Clone + Send + Sync + 'static;
/// The caller's request as the host projects it, answered once before anything else.
type Projection: Send + 'static;
/// Each operation carries the [`Reply`](crate::host::Reply) its answer goes through.
/// A protocol with no operations of its own uses `Infallible`.
type Op: Send + 'static;
/// One piece of a streamed response, handed to the caller as it arrives. A protocol
/// that never streams uses `Infallible`.
type Chunk: Send + 'static;
/// What the call knows once a streamed response starts, before its first chunk.
type StreamHead: Send + 'static;
}

View file

@ -1,14 +0,0 @@
/// One public call surface: what a completed call produces, how it fails, and the
/// route-specific operations only its host can perform (request projection, file reads,
/// token acquisition).
pub trait Route: Send + Sync + 'static {
type Response: Send + 'static;
type Error: Clone + Send + Sync + 'static;
type Op: Send + 'static;
type OpResult: Send + 'static;
/// One piece of a streamed response, handed to the caller as it arrives. A route
/// that never streams uses `Infallible`.
type Chunk: Send + 'static;
/// What the route knows once a streamed response starts, before its first chunk.
type StreamHead: Send + 'static;
}

View file

@ -1,40 +1,28 @@
use crate::event::{CallEvent, FailureOrigin, Timing, epoch_seconds};
use crate::host::{Host, HostOp, HostResult};
use crate::host::{Host, HostOp};
use crate::machine::{HostFailure, Machine, MachineStep};
use crate::route::Route;
use crate::protocol::Protocol;
/// Drives a machine to completion against an in-process host and emits exactly one
/// terminal event.
pub async fn run<M, H>(mut machine: M, host: &H) -> Result<M::Complete, <M::Route as Route>::Error>
pub async fn run<M, H>(
mut machine: M,
host: &H,
) -> Result<M::Complete, <M::Protocol as Protocol>::Error>
where
M: Machine,
H: Host<M::Route>,
H: Host<M::Protocol>,
{
let start_time = epoch_seconds();
let _ = host.emit(&CallEvent::Started { start_time }).await;
let mut result = None;
let outcome = loop {
let step = match machine.resume(result.take()).await {
let op = match machine.resume().await {
Ok(MachineStep::Complete(complete)) => break Ok(complete),
Ok(MachineStep::Host(op)) => op,
Err(error) => break Err(error),
};
let answer = match step {
HostOp::Route(op) => host.route(op).await.map(HostResult::Route),
HostOp::BeforeSend { wire, context } => host
.before_send(*wire, &context)
.await
.map(|wire| HostResult::BeforeSend(Box::new(wire))),
HostOp::Emit(event) => host
.emit(&CallEvent::Machine(event))
.await
.map(|()| HostResult::Emitted),
HostOp::Open(head) => host.open(head).await.map(HostResult::Demand),
HostOp::Deliver(chunk) => host.deliver(chunk).await.map(HostResult::Demand),
};
match answer {
Ok(answer) => result = Some(answer),
Err(error) => break machine.interrupt(HostFailure::Error(error)).await,
if let Err(error) = perform(host, op).await {
break machine.interrupt(HostFailure::Error(error)).await;
}
};
let timing = Timing {
@ -52,44 +40,52 @@ where
outcome
}
async fn perform<R: Protocol, H: Host<R>>(host: &H, op: HostOp<R>) -> Result<(), R::Error> {
match op {
HostOp::Project(reply) => host
.project()
.await
.map(|projection| reply.send(projection)),
HostOp::Custom(op) => host.custom_op(op).await,
HostOp::BeforeSend {
wire,
context,
reply,
} => host
.before_send(*wire, &context)
.await
.map(|wire| reply.send(wire)),
HostOp::Emit(event, reply) => host
.emit(&CallEvent::Machine(event))
.await
.map(|()| reply.send(())),
HostOp::Open(head, reply) => host.open(head).await.map(|demand| reply.send(demand)),
HostOp::Deliver(chunk, reply) => host.deliver(chunk).await.map(|demand| reply.send(demand)),
}
}
#[cfg(test)]
mod tests {
use std::sync::Mutex;
use super::*;
use crate::machine::{Interrupted, Step};
use crate::host::Reply;
use crate::machine::{CallMachine, MachineFault};
struct Unit;
impl Route for Unit {
impl Protocol for Unit {
type Response = ();
type Error = &'static str;
type Op = &'static str;
type OpResult = ();
type Projection = ();
type Op = (&'static str, Reply<()>);
type Chunk = std::convert::Infallible;
type StreamHead = std::convert::Infallible;
}
struct Scripted {
ops: Vec<&'static str>,
outcome: Result<(), &'static str>,
}
impl Machine for Scripted {
type Route = Unit;
type Complete = ();
fn resume(&mut self, _: Option<HostResult<Unit>>) -> Step<'_, Self> {
Box::pin(async move {
if !self.ops.is_empty() {
return Ok(MachineStep::Host(HostOp::Route(self.ops.remove(0))));
}
self.outcome.map(MachineStep::Complete)
})
}
fn interrupt(&mut self, failure: HostFailure<&'static str>) -> Interrupted<'_, Self> {
Box::pin(async move { Err(failure.into_error()) })
impl From<MachineFault> for &'static str {
fn from(_: MachineFault) -> Self {
"machine fault"
}
}
@ -100,12 +96,21 @@ mod tests {
}
impl Host<Unit> for Recording {
async fn route(&self, op: &'static str) -> Result<(), &'static str> {
self.seen.lock().unwrap().push(format!("route:{op}"));
match self.fail {
Some(failing) if failing == op => Err("host failed"),
_ => Ok(()),
async fn project(&self) -> Result<(), &'static str> {
self.seen.lock().unwrap().push("project".into());
Ok(())
}
async fn custom_op(
&self,
(op, reply): (&'static str, Reply<()>),
) -> Result<(), &'static str> {
self.seen.lock().unwrap().push(format!("op:{op}"));
if self.fail == Some(op) {
return Err("host failed");
}
reply.send(());
Ok(())
}
async fn emit(&self, event: &CallEvent) -> Result<(), &'static str> {
@ -119,21 +124,29 @@ mod tests {
}
}
fn scripted(ops: &[&'static str], outcome: Result<(), &'static str>) -> Scripted {
Scripted {
ops: ops.to_vec(),
outcome,
}
fn scripted(
ops: &'static [&'static str],
outcome: Result<(), &'static str>,
) -> CallMachine<Unit> {
CallMachine::new(move |host| {
Box::pin(async move {
host.project().await?;
for op in ops {
host.custom_op(|reply| (*op, reply)).await?;
}
outcome
})
})
}
#[tokio::test]
async fn forwards_every_op_then_emits_one_succeeded() {
let host = Recording::default();
let outcome = run(scripted(&["project", "send"], Ok(())), &host).await;
let outcome = run(scripted(&["sign", "send"], Ok(())), &host).await;
assert_eq!(outcome, Ok(()));
assert_eq!(
*host.seen.lock().unwrap(),
["started", "route:project", "route:send", "succeeded"]
["started", "project", "op:sign", "op:send", "succeeded"]
);
}
@ -142,24 +155,32 @@ mod tests {
let host = Recording::default();
let outcome = run(scripted(&[], Err("boom")), &host).await;
assert_eq!(outcome, Err("boom"));
assert_eq!(*host.seen.lock().unwrap(), ["started", "failed"]);
assert_eq!(*host.seen.lock().unwrap(), ["started", "project", "failed"]);
let host = Recording {
fail: Some("send"),
..Recording::default()
};
let outcome = run(scripted(&["project", "send", "never"], Ok(())), &host).await;
let outcome = run(scripted(&["sign", "send", "never"], Ok(())), &host).await;
assert_eq!(outcome, Err("host failed"));
assert_eq!(
*host.seen.lock().unwrap(),
["started", "route:project", "route:send", "failed"]
["started", "project", "op:sign", "op:send", "failed"]
);
}
struct StartTimes(Mutex<Vec<f64>>);
impl Host<Unit> for StartTimes {
async fn route(&self, _: &'static str) -> Result<(), &'static str> {
async fn project(&self) -> Result<(), &'static str> {
Ok(())
}
async fn custom_op(
&self,
(_, reply): (&'static str, Reply<()>),
) -> Result<(), &'static str> {
reply.send(());
Ok(())
}
@ -178,7 +199,7 @@ mod tests {
#[tokio::test]
async fn started_opens_the_call_at_the_terminal_start_time_and_cannot_fail_it() {
let host = StartTimes(Mutex::default());
assert_eq!(run(scripted(&["project"], Ok(())), &host).await, Ok(()));
assert_eq!(run(scripted(&["send"], Ok(())), &host).await, Ok(()));
let times = host.0.lock().unwrap();
assert_eq!(times.len(), 2);
assert_eq!(times[0], times[1]);

View file

@ -118,7 +118,6 @@ impl From<litellm_host::machine::MachineFault> for Error {
Self::InvalidRequest(match fault {
MachineFault::Abandoned => "OCR host driver was abandoned".into(),
MachineFault::Protocol(message) => format!("OCR {message}"),
MachineFault::Mismatch => "invalid OCR host operation result".into(),
})
}
}

View file

@ -1,9 +1,8 @@
use std::sync::OnceLock;
use litellm_host::{
host::HostResult,
machine::{HostFailure, Interrupted, Machine, Step},
route::Route,
protocol::Protocol,
};
use litellm_tracing::Logger;
use pyo3::Python;
@ -23,17 +22,17 @@ impl<M> LoggedMachine<M> {
}
impl<M: Machine> Machine for LoggedMachine<M> {
type Route = M::Route;
type Protocol = M::Protocol;
type Complete = M::Complete;
fn resume(&mut self, result: Option<HostResult<Self::Route>>) -> Step<'_, Self> {
fn resume(&mut self) -> Step<'_, Self> {
let logger = self.logger.get_or_init(|| Python::attach(super::capture));
Box::pin(logger.instrument(logger.scope(|| self.machine.resume(result))))
Box::pin(logger.instrument(logger.scope(|| self.machine.resume())))
}
fn interrupt(
&mut self,
failure: HostFailure<<Self::Route as Route>::Error>,
failure: HostFailure<<Self::Protocol as Protocol>::Error>,
) -> Interrupted<'_, Self> {
let logger = self.logger.get_or_init(|| Python::attach(super::capture));
Box::pin(logger.instrument(logger.scope(|| self.machine.interrupt(failure))))

View file

@ -1,29 +1,28 @@
use std::{process::Command, task::Poll};
use litellm_host::{
host::HostResult,
machine::{HostFailure, Interrupted, Machine, MachineStep, Step},
route::Route,
protocol::Protocol,
};
use pyo3::{prelude::*, types::PyDict};
struct DiagnosticMachine;
impl Route for DiagnosticMachine {
impl Protocol for DiagnosticMachine {
type Response = ();
type Error = String;
type Projection = ();
type Op = ();
type OpResult = ();
type Chunk = ();
type StreamHead = ();
}
impl Machine for DiagnosticMachine {
type Route = Self;
type Protocol = Self;
type Complete = ();
fn resume(&mut self, _: Option<HostResult<Self>>) -> Step<'_, Self> {
fn resume(&mut self) -> Step<'_, Self> {
litellm_tracing::warn!("machine started");
Box::pin(async {
tokio::task::yield_now().await;
@ -45,7 +44,7 @@ fn machine_warning(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
let mut machine = super::LoggedMachine::new(DiagnosticMachine);
let mut future = Box::pin(async move {
machine
.resume(None)
.resume()
.await
.map_err(pyo3::exceptions::PyValueError::new_err)?;
machine

View file

@ -1,10 +1,12 @@
use std::convert::Infallible;
use bytes::Bytes;
use litellm_core::messages::{
Error,
route::{Messages, MessagesCall, MessagesOp, MessagesOpResult, MessagesOutput},
route::{Messages, MessagesCall, MessagesOutput},
types::MessagesShaping,
};
use litellm_host_python::{InvokeError, RouteHost, from_py, lookup, to_py};
use litellm_host_python::{InvokeError, ProtocolHost, from_py, lookup, to_py};
use litellm_http::transport::Error as TransportError;
use litellm_types::utils::ProviderSpecificHeaders;
use pyo3::{
@ -80,16 +82,16 @@ fn native_error(py: Python<'_>, error: Error) -> PyResult<PyErr> {
/// The Python side of the Messages route: projects the prepared arguments and builds the
/// public response, chunks and exceptions.
pub(super) struct MessagesRouteHost {
pub(super) struct MessagesPythonHost {
request: Py<PyAny>,
}
impl MessagesRouteHost {
impl MessagesPythonHost {
pub(super) fn new(request: Py<PyAny>) -> Self {
Self { request }
}
fn project(&self, py: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult<MessagesCall> {
fn projection(&self, py: Python<'_>, arguments: &Bound<'_, PyDict>) -> PyResult<MessagesCall> {
let request = self.request.bind(py);
let argument = |name: &str| -> PyResult<Option<Bound<'_, PyAny>>> {
Ok(lookup(arguments, request, name)?.filter(|value| !value.is_none()))
@ -208,22 +210,21 @@ impl MessagesRouteHost {
}
}
impl RouteHost for MessagesRouteHost {
type Route = Messages;
impl ProtocolHost for MessagesPythonHost {
type Protocol = Messages;
type Failure = PyErr;
fn invoke(
fn project(
&mut self,
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
op: MessagesOp,
) -> Result<MessagesOpResult, InvokeError<Error>> {
match op {
MessagesOp::ProjectRequest => self
.project(py, arguments)
.map(|call| MessagesOpResult::Request(Box::new(call)))
.map_err(|error| InvokeError::Python(self.map_failure(py, error))),
}
) -> Result<MessagesCall, InvokeError<Error>> {
self.projection(py, arguments)
.map_err(|error| InvokeError::Python(self.map_failure(py, error)))
}
fn invoke(&mut self, _: Python<'_>, op: Infallible) -> Result<(), InvokeError<Error>> {
match op {}
}
fn complete(&mut self, py: Python<'_>, response: MessagesOutput) -> PyResult<Py<PyAny>> {

View file

@ -1,6 +1,6 @@
mod host;
use host::MessagesRouteHost;
use host::MessagesPythonHost;
use litellm_callbacks_legacy_python::{
LegacySurface, PassThroughStream, PublicCall, run_legacy_call,
};
@ -45,7 +45,7 @@ fn run_messages(
SURFACE,
PublicCall::capture(&request, &args, &kwargs)?,
crate::logger::LoggedMachine::new(messages_machine(secrets)),
MessagesRouteHost::new(request.unbind()),
MessagesPythonHost::new(request.unbind()),
asynchronous,
)
}

View file

@ -1,58 +1,38 @@
use std::path::PathBuf;
use bytes::Bytes;
use litellm_core::ocr::types::{OcrDocumentInput, OcrFileContent};
use litellm_core::ocr::types::OcrDocumentInput;
use litellm_host_python::{PythonFileReader, py_bytes};
use pyo3::{
exceptions::{PyTypeError, PyValueError},
gc::{PyTraverseError, PyVisit},
exceptions::PyValueError,
prelude::*,
pybacked::PyBackedBytes,
sync::PyOnceLock,
types::{PyBytes, PyString, PyType},
};
#[derive(Debug)]
pub(super) struct PythonFileReader {
reader: Py<PyAny>,
name: Option<String>,
/// A `type='file'` document as projected: paths and bytes are typed inputs already; a
/// file-like object is a reader the projection consumes once every other field is read.
pub(super) enum FileDocumentInput {
Ready(OcrDocumentInput),
Deferred {
reader: PythonFileReader,
mime_type: Option<String>,
},
}
impl PythonFileReader {
pub(super) fn read(&self, py: Python<'_>) -> PyResult<OcrFileContent> {
let value = self.reader.bind(py).call0()?;
let bytes = if value.is_instance_of::<PyString>() {
Bytes::from(value.extract::<String>()?)
} else if value.is_instance_of::<PyBytes>() {
extract_bytes(&value)?
} else {
return Err(PyTypeError::new_err(format!(
"OCR file read must return bytes or str, got {}",
value.get_type(),
)));
};
Ok(OcrFileContent {
bytes,
file_name: self.name.clone(),
})
impl FileDocumentInput {
pub(super) fn resolve(self, py: Python<'_>) -> PyResult<OcrDocumentInput> {
match self {
Self::Ready(input) => Ok(input),
Self::Deferred { reader, mime_type } => {
let content = reader.read(py)?;
Ok(OcrDocumentInput::Bytes {
bytes: content.bytes,
file_name: content.file_name,
mime_type,
})
}
}
}
pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.reader)
}
}
fn extract_bytes(value: &Bound<'_, PyAny>) -> PyResult<Bytes> {
if value.is_exact_instance_of::<PyBytes>() {
return Ok(Bytes::from_owner(value.extract::<PyBackedBytes>()?));
}
Ok(Bytes::copy_from_slice(
value.extract::<PyBackedBytes>()?.as_ref(),
))
}
pub(super) struct FileDocumentInput {
pub input: OcrDocumentInput,
pub reader: Option<PythonFileReader>,
}
impl FromPyObject<'_, '_> for FileDocumentInput {
@ -87,51 +67,31 @@ impl FromPyObject<'_, '_> for FileDocumentInput {
}
static PATH_LIKE: PyOnceLock<Py<PyType>> = PyOnceLock::new();
if file.is_instance(PATH_LIKE.import(py, "os", "PathLike")?)? {
return Ok(Self {
input: OcrDocumentInput::Path {
path: file.extract::<PathBuf>()?,
mime_type,
},
reader: None,
});
return Ok(Self::Ready(OcrDocumentInput::Path {
path: file.extract::<PathBuf>()?,
mime_type,
}));
}
if file.is_instance_of::<PyBytes>() {
return Ok(Self {
input: OcrDocumentInput::Bytes {
bytes: extract_bytes(&file)?,
file_name: None,
mime_type,
},
reader: None,
});
return Ok(Self::Ready(OcrDocumentInput::Bytes {
bytes: py_bytes(&file)?,
file_name: None,
mime_type,
}));
}
let reader = file
.getattr_opt("read")?
.filter(|value| value.is_callable());
let Some(reader) = reader else {
return Err(PyValueError::new_err(format!(
match PythonFileReader::from_file_like(&file)? {
Some(reader) => Ok(Self::Deferred { reader, mime_type }),
None => Err(PyValueError::new_err(format!(
"Unsupported file input type: {}. Expected pathlib.Path, bytes, or a file-like object.",
file.get_type(),
)));
};
let name = file
.getattr_opt("name")?
.filter(|value| !value.is_none())
.map(|value| value.extract::<String>())
.transpose()?;
Ok(Self {
input: OcrDocumentInput::HostReader { mime_type },
reader: Some(PythonFileReader {
reader: reader.unbind(),
name,
}),
})
))),
}
}
}
#[cfg(test)]
mod tests {
use pyo3::types::PyDict;
use pyo3::{exceptions::PyTypeError, types::PyDict};
use super::*;
@ -141,6 +101,13 @@ mod tests {
locals
}
fn ready(input: FileDocumentInput) -> OcrDocumentInput {
match input {
FileDocumentInput::Ready(input) => input,
FileDocumentInput::Deferred { .. } => panic!("expected a ready document"),
}
}
#[test]
fn extraction_validates_required_file_and_optional_mime_type() {
Python::initialize();
@ -167,13 +134,19 @@ mod tests {
.unwrap();
assert!(error.is_instance_of::<PyValueError>(py));
assert!(error.to_string().contains("bare str"));
let error = py
.eval(c"{'file': object()}", None, None)
.unwrap()
.extract::<FileDocumentInput>()
.err()
.unwrap();
assert!(error.is_instance_of::<PyValueError>(py));
assert!(error.to_string().contains("Unsupported file input type"));
let document = py
.eval(c"{'file': b'abc', 'mime_type': 'image/png'}", None, None)
.unwrap();
let input: FileDocumentInput = document.extract().unwrap();
assert!(input.reader.is_none());
assert_eq!(
input.input,
ready(document.extract().unwrap()),
OcrDocumentInput::Bytes {
bytes: b"abc".as_slice().into(),
file_name: None,
@ -199,7 +172,7 @@ class Reader:
return b'abc'
reader = Reader()
document = {'file': reader, 'mime_type': 7}
reader_document = {'file': reader}
reader_document = {'file': reader, 'mime_type': 'application/pdf'}
path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf'), 'mime_type': 'image/png'}",
);
let document = locals.get_item("document").unwrap().unwrap();
@ -208,10 +181,6 @@ path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf'), 'mime_typ
let document = locals.get_item("reader_document").unwrap().unwrap();
let input: FileDocumentInput = document.extract().unwrap();
assert_eq!(
input.input,
OcrDocumentInput::HostReader { mime_type: None }
);
let reads = || {
locals
.get_item("reader")
@ -223,21 +192,20 @@ path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf'), 'mime_typ
.unwrap()
};
assert_eq!(reads(), 0);
let content = input.reader.unwrap().read(py).unwrap();
let resolved = input.resolve(py).unwrap();
assert_eq!(reads(), 1);
assert_eq!(
content,
OcrFileContent {
resolved,
OcrDocumentInput::Bytes {
bytes: b"abc".as_slice().into(),
file_name: Some("scan.png".into()),
mime_type: Some("application/pdf".into()),
}
);
let document = locals.get_item("path_document").unwrap().unwrap();
let input: FileDocumentInput = document.extract().unwrap();
assert!(input.reader.is_none());
assert_eq!(
input.input,
ready(document.extract().unwrap()),
OcrDocumentInput::Path {
path: PathBuf::from("/nonexistent/ocr-projection-test.pdf"),
mime_type: Some("image/png".into()),
@ -245,97 +213,4 @@ path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf'), 'mime_typ
);
});
}
#[test]
fn reader_results_are_normalized_and_exceptions_keep_their_identity() {
Python::initialize();
Python::attach(|py| {
let locals = eval(
py,
c"failure = KeyError('reader failed')
class Raising:
def read(self):
raise failure
class Text:
def read(self):
return 'héllo'
class Wrong:
def read(self):
return 7
raising = {'file': Raising()}
text = {'file': Text()}
wrong = {'file': Wrong()}",
);
let reader = |name: &str| {
locals
.get_item(name)
.unwrap()
.unwrap()
.extract::<FileDocumentInput>()
.unwrap()
.reader
.unwrap()
};
let error = reader("raising").read(py).unwrap_err();
assert!(
error
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
assert_eq!(
reader("text").read(py).unwrap().bytes.as_ref(),
"héllo".as_bytes()
);
let error = reader("wrong").read(py).unwrap_err();
assert!(error.is_instance_of::<PyTypeError>(py));
assert!(error.to_string().contains("bytes or str"));
});
}
#[rstest::rstest]
#[case::read("read")]
#[case::name("name")]
fn reader_attribute_failures_keep_their_identity(#[case] attribute: &str) {
Python::initialize();
Python::attach(|py| {
let locals = eval(
py,
c"failure = LookupError('file property failed')
class File:
def __getattribute__(self, name):
if name == attribute:
raise failure
return super().__getattribute__(name)
name = 'scan.pdf'
def read(self):
return b'abc'
document = {'file': File()}",
);
locals.set_item("attribute", attribute).unwrap();
let error = locals
.get_item("document")
.unwrap()
.unwrap()
.extract::<FileDocumentInput>()
.err()
.unwrap();
assert!(
error
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
});
}
#[test]
fn exact_python_bytes_transfer_without_copying_and_outlive_the_input() {
Python::initialize();
let (bytes, pointer) = Python::attach(|py| {
let value = PyBytes::new(py, b"document bytes");
let pointer = value.as_bytes().as_ptr() as usize;
(extract_bytes(value.as_any()).unwrap(), pointer)
});
assert_eq!(bytes.as_ptr() as usize, pointer);
assert_eq!(bytes.as_ref(), b"document bytes");
}
}

View file

@ -1,6 +1,6 @@
use litellm_auth::ResolvedCredential;
use litellm_core::ocr::route::{Ocr, OcrOp, OcrOpResult};
use litellm_host_python::{InvokeError, RouteHost, missing_state, to_py};
use litellm_core::ocr::route::{Ocr, OcrOp, OcrProjection};
use litellm_host_python::{InvokeError, ProtocolHost, missing_state, to_py};
use litellm_llms::base_llm::ocr::{error::Error, transformation::LiteLLMOcrResponse};
use pyo3::{
exceptions::{PyBaseException, PyException},
@ -20,14 +20,15 @@ enum OcrHostData {
Released,
}
/// The Python side of the OCR route: projects the prepared arguments, reads file-like
/// documents, acquires Azure AD tokens, and builds the public response and exception.
pub(super) struct OcrRouteHost {
/// The Python side of the OCR route: projects the prepared arguments (reading a file-like
/// document as it goes), acquires Azure AD tokens, and builds the public response and
/// exception.
pub(super) struct OcrPythonHost {
request: Py<PyAny>,
data: OcrHostData,
}
impl OcrRouteHost {
impl OcrPythonHost {
pub(super) fn new(request: Py<PyAny>) -> Self {
Self {
request,
@ -42,14 +43,6 @@ impl OcrRouteHost {
}
}
fn read_document(&self, py: Python<'_>) -> PyResult<litellm_core::ocr::types::OcrFileContent> {
self.handles()?
.reader
.as_ref()
.ok_or_else(missing_state)?
.read(py)
}
fn acquire_azure_ad_token(&self, py: Python<'_>) -> PyResult<ResolvedCredential> {
self.handles()?
.azure_ad_token_provider
@ -58,30 +51,21 @@ impl OcrRouteHost {
.acquire(py)
}
fn answer(
fn projection(
&mut self,
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
op: OcrOp,
) -> PyResult<OcrOpResult> {
match op {
OcrOp::ProjectRequest => {
let OcrHostData::Unprojected = self.data else {
return Err(missing_state());
};
let (request, handles) = project_request(self.request.bind(py), arguments)?;
let caller_token = handles.azure_ad_token_provider.is_some();
self.data = OcrHostData::Projected(Box::new(handles));
Ok(OcrOpResult::Request {
request: Box::new(request),
caller_token,
})
}
OcrOp::ReadDocument => self.read_document(py).map(OcrOpResult::Document),
OcrOp::AcquireAzureAdToken => self
.acquire_azure_ad_token(py)
.map(OcrOpResult::AzureAdToken),
}
) -> PyResult<OcrProjection> {
let OcrHostData::Unprojected = self.data else {
return Err(missing_state());
};
let (request, handles) = project_request(self.request.bind(py), arguments)?;
let caller_token = handles.azure_ad_token_provider.is_some();
self.data = OcrHostData::Projected(Box::new(handles));
Ok(OcrProjection {
request,
caller_token,
})
}
fn map_failure(&self, py: Python<'_>, error: PyErr) -> PyErr {
@ -104,20 +88,28 @@ impl OcrRouteHost {
}
}
impl RouteHost for OcrRouteHost {
type Route = Ocr;
impl ProtocolHost for OcrPythonHost {
type Protocol = Ocr;
type Failure = PyErr;
fn invoke(
fn project(
&mut self,
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
op: OcrOp,
) -> Result<OcrOpResult, InvokeError<Error>> {
self.answer(py, arguments, op)
) -> Result<OcrProjection, InvokeError<Error>> {
self.projection(py, arguments)
.map_err(|error| InvokeError::Python(self.map_failure(py, error)))
}
fn invoke(&mut self, py: Python<'_>, op: OcrOp) -> Result<(), InvokeError<Error>> {
match op {
OcrOp::AcquireAzureAdToken(reply) => self
.acquire_azure_ad_token(py)
.map(|token| reply.send(token))
.map_err(|error| InvokeError::Python(self.map_failure(py, error))),
}
}
fn complete(&mut self, py: Python<'_>, response: LiteLLMOcrResponse) -> PyResult<Py<PyAny>> {
py.import("litellm.rust_bridge.ocr.route_host")?
.getattr("response")?
@ -148,13 +140,10 @@ impl RouteHost for OcrRouteHost {
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.request)?;
if let OcrHostData::Projected(handles) = &self.data {
if let Some(reader) = &handles.reader {
reader.traverse(visit)?;
}
if let Some(provider) = &handles.azure_ad_token_provider {
provider.traverse(visit)?;
}
if let OcrHostData::Projected(handles) = &self.data
&& let Some(provider) = &handles.azure_ad_token_provider
{
provider.traverse(visit)?;
}
Ok(())
}
@ -205,20 +194,13 @@ del provider
.unwrap()
.cast_into::<PyDict>()
.unwrap();
let mut host = OcrRouteHost::new(py.None());
let projected = host.invoke(py, &kwargs, OcrOp::ProjectRequest).unwrap();
assert!(matches!(
projected,
OcrOpResult::Request {
caller_token: true,
..
}
));
let mut host = OcrPythonHost::new(py.None());
assert!(host.project(py, &kwargs).unwrap().caller_token);
locals.del_item("kwargs").unwrap();
drop(kwargs);
let (reply, _) = litellm_host::host::reply();
assert_eq!(
host.invoke(py, &PyDict::new(py), OcrOp::AcquireAzureAdToken)
.is_ok(),
host.invoke(py, OcrOp::AcquireAzureAdToken(reply)).is_ok(),
succeeds
);
let alive = || {

View file

@ -5,7 +5,7 @@ mod project;
use std::sync::LazyLock;
use host::OcrRouteHost;
use host::OcrPythonHost;
use litellm_auth_gcp::VertexAuth;
use litellm_callbacks_legacy_python::{LegacySurface, PublicCall, run_legacy_call};
use litellm_core::ocr::{provider_config, route::ocr_machine};
@ -69,7 +69,7 @@ fn run_ocr(
if asynchronous { ASYNC_SURFACE } else { SURFACE },
PublicCall::capture(&request, &args, &kwargs)?,
crate::logger::LoggedMachine::new(ocr_machine(client)),
OcrRouteHost::new(request.unbind()),
OcrPythonHost::new(request.unbind()),
asynchronous,
)
}

View file

@ -8,19 +8,15 @@ use litellm_llms::base_llm::ocr::error::Error;
use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict};
use serde_json::{Map, Value};
use super::{
document::{FileDocumentInput, PythonFileReader},
errors::to_pyerr as ocr_error_to_pyerr,
};
use super::{document::FileDocumentInput, errors::to_pyerr as ocr_error_to_pyerr};
use crate::{
credentials::{self, CallerTokenProvider},
marshal::{project_optional_fields, python_timeout_seconds, request_input_sources},
};
/// What the host keeps after projection: the caller's callables that answer the document
/// read and token operations, and the provider name the failure mapping reports.
/// What the host keeps after projection: the caller's token callable that answers the
/// token operation, and the provider name the failure mapping reports.
pub(super) struct OcrHostHandles {
pub reader: Option<PythonFileReader>,
pub azure_ad_token_provider: Option<CallerTokenProvider>,
pub provider: &'static str,
}
@ -104,13 +100,11 @@ impl ProjectedDocument {
Ok(Self::File(document.extract()?))
}
fn into_parts(self) -> PyResult<(OcrDocumentInput, Option<PythonFileReader>)> {
/// Reads a file-like document now, so it runs after every other argument was read.
fn resolve(self, py: Python<'_>) -> PyResult<OcrDocumentInput> {
match self {
Self::File(FileDocumentInput { input, reader }) => Ok((input, reader)),
Self::Other(wire) => Ok((
decode_document(wire).map_err(ocr_error_to_pyerr)?.into(),
None,
)),
Self::File(file) => file.resolve(py),
Self::Other(wire) => Ok(decode_document(wire).map_err(ocr_error_to_pyerr)?.into()),
}
}
}
@ -136,24 +130,25 @@ pub(super) fn project_request(
.chain(["api_key", "api_base", "extra_headers"]),
)?;
let azure_ad_token_provider = credentials::azure_ad_token_provider(kwargs)?;
let (document, reader) = document.into_parts()?;
let api_base = arguments.api_base()?;
let extra_headers = arguments.extra_headers()?;
let timeout_seconds = arguments.timeout_seconds()?;
let wire = OcrWireRequest {
model,
document,
document: document.resolve(request.py())?,
api_key,
api_base: arguments.api_base()?,
api_base,
custom_llm_provider,
extra_headers: arguments.extra_headers()?,
extra_headers,
optional_params,
input_sources,
timeout_seconds: arguments.timeout_seconds()?,
timeout_seconds,
};
let request = decode_request_input(wire).map_err(ocr_error_to_pyerr)?;
let provider = request.provider_name();
Ok((
request,
OcrHostHandles {
reader,
azure_ad_token_provider,
provider,
},
@ -180,10 +175,8 @@ mod tests {
OcrArguments { request, kwargs }
}
fn project_document(
document: &Bound<'_, PyAny>,
) -> PyResult<(OcrDocumentInput, Option<PythonFileReader>)> {
ProjectedDocument::project(document)?.into_parts()
fn project_document(document: &Bound<'_, PyAny>) -> PyResult<OcrDocumentInput> {
ProjectedDocument::project(document)?.resolve(document.py())
}
fn url_document(url: &str) -> OcrDocumentInput {
@ -342,8 +335,11 @@ kwargs = {}
});
}
/// A reader that rewrites the request while it runs shows which arguments projection
/// read before it and which after: every other argument is read first, and the read
/// happens exactly once.
#[test]
fn document_readers_are_not_consumed_during_projection() {
fn document_readers_are_read_once_after_every_other_argument() {
Python::initialize();
Python::attach(|py| {
stub_timeout_conversion(py);
@ -351,17 +347,24 @@ kwargs = {}
py,
c"
class Request:
api_base = 'original'
model = 'mistral/mistral-ocr-latest'
custom_llm_provider = None
api_key = None
api_base = 'https://original.example.com'
extra_headers = {'x-source': 'original'}
timeout = 1
@property
def document(self):
return document
class Reader:
reads = 0
def read(self):
Request.api_base = 'mutated'
Reader.reads += 1
Request.api_base = 'https://mutated.example.com'
Request.extra_headers = {'x-source': 'mutated'}
Request.timeout = 9
return b'abc'
document = {'type': 'file', 'file': Reader()}
document = {'type': 'file', 'file': Reader(), 'mime_type': 'application/pdf'}
request = Request()
kwargs = {}
",
@ -373,15 +376,38 @@ kwargs = {}
.unwrap()
.cast_into::<PyDict>()
.unwrap();
let arguments = arguments(&request, &kwargs);
let document = arguments.document().unwrap();
let (input, reader) = project_document(&document).unwrap();
assert_eq!(input, OcrDocumentInput::HostReader { mime_type: None });
assert_eq!(arguments.api_base().unwrap().as_deref(), Some("original"));
assert_eq!(arguments.timeout_seconds().unwrap(), Some(1.0));
reader.unwrap().read(py).unwrap();
assert_eq!(arguments.api_base().unwrap().as_deref(), Some("mutated"));
assert_eq!(arguments.timeout_seconds().unwrap(), Some(9.0));
let (projected, _) = project_request(&request, &kwargs).unwrap();
assert_eq!(
py.eval(c"Reader.reads", Some(&locals), Some(&locals))
.unwrap()
.extract::<usize>()
.unwrap(),
1
);
assert_eq!(
projected.document,
OcrDocumentInput::Bytes {
bytes: b"abc".as_slice().into(),
file_name: None,
mime_type: Some("application/pdf".into()),
}
);
assert_eq!(
projected
.credentials
.api_base
.as_ref()
.map(|base| base.value().as_str()),
Some("https://original.example.com")
);
assert_eq!(
projected.transport.extra_headers,
[("x-source".to_string(), "original".to_string())]
);
assert_eq!(
projected.transport.timeout,
Some(std::time::Duration::from_secs(1))
);
});
}
@ -396,16 +422,14 @@ kwargs = {}
None,
)
.unwrap();
let (input, reader) = project_document(&file).unwrap();
assert_eq!(
input,
project_document(&file).unwrap(),
OcrDocumentInput::Bytes {
bytes: b"%PDF-1.4".as_slice().into(),
file_name: None,
mime_type: Some("application/pdf".into()),
}
);
assert!(reader.is_none());
let original = py
.eval(
@ -414,8 +438,10 @@ kwargs = {}
None,
)
.unwrap();
let (input, _) = project_document(&original).unwrap();
assert_eq!(input, url_document("https://example.com/a.pdf"));
assert_eq!(
project_document(&original).unwrap(),
url_document("https://example.com/a.pdf")
);
});
}
@ -617,7 +643,7 @@ document = Document()
",
);
let document = locals.get_item("document").unwrap().unwrap();
let (input, _) = project_document(&document).unwrap();
let input = project_document(&document).unwrap();
assert!(matches!(input, OcrDocumentInput::Bytes { .. }));
let reads: Vec<String> = document.getattr("reads").unwrap().extract().unwrap();
assert_eq!(reads, ["type", "mime_type", "file"]);

View file

@ -167,6 +167,9 @@ MCP_OAUTH2_TOKEN_CACHE_MAX_SIZE: Final = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_M
MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL: Final = int(os.getenv("MCP_OAUTH2_TOKEN_CACHE_DEFAULT_TTL", "3600"))
MCP_SSO_ASSERTION_CACHE_TTL_SECONDS: Final = int(os.getenv("MCP_SSO_ASSERTION_CACHE_TTL_SECONDS", "60"))
# mcp_tool_permissions entry that grants every current and future tool on a server
MCP_ALL_TOOLS_WILDCARD: Final = "*"
# Default npm cache directory for STDIO MCP servers.
# npm/npx needs a writable cache dir; in containers the default (~/.npm)
# may not exist or be read-only. /tmp is always writable.

View file

@ -790,6 +790,16 @@ def _get_hidden_str_for_cost_calc(hidden_params: object, key: str) -> str | None
return value if isinstance(value, str) and value else None
_NON_TOKEN_RATE_FIELDS: Final = frozenset({"input_cost_per_second", "input_cost_per_query", "tiered_pricing"})
def _cost_map_entry_prices_anything(entry: Mapping[str, object]) -> bool:
return any(
value is not None and (field in _NON_TOKEN_RATE_FIELDS or ("cost_per" in field and "token" in field))
for field, value in entry.items()
)
def _select_model_name_for_cost_calc(
model: str | None,
completion_response: object | None,
@ -828,12 +838,7 @@ def _select_model_name_for_cost_calc(
if custom_pricing is True:
if router_model_id is not None and router_model_id in litellm.model_cost:
entry: Final = litellm.model_cost[router_model_id]
if (
entry.get("input_cost_per_token") is not None
or entry.get("input_cost_per_second") is not None
or entry.get("input_cost_per_query") is not None
or entry.get("tiered_pricing") is not None
):
if _cost_map_entry_prices_anything(entry):
return_model = router_model_id
else:
return_model = model
@ -1699,6 +1704,8 @@ def completion_cost(
litellm_model_name=model,
data_residency=data_residency,
litellm_logging_obj=litellm_logging_obj,
custom_pricing_model=selected_model if custom_pricing else None,
base_pricing_model=(selected_model if base_model is not None and not custom_pricing else None),
)
elif call_type == _MCP_CALL_TYPE:
from litellm.proxy._experimental.mcp_server.cost_calculator import (
@ -2870,14 +2877,20 @@ def _candidate_realtime_token_costs(
def _cost_map_entry_declares_pricing(model_name: str, custom_llm_provider: str) -> bool:
"""Whether the entry behind ``model_name`` sets any rate of its own, even a zero one.
The name is resolved the way ``get_model_info`` resolves it before the raw entry is read,
because a deployment-scoped name arrives here already carrying its provider prefix. Two raw
lookups cannot strip that prefix, so a zero-rated override read as declaring nothing, and a
session that should bill nothing fell through to the public rates instead.
"""
resolved: Final = _get_model_info_or_none(model_name, custom_llm_provider)
entries: Final = (
litellm.model_cost.get(resolved.get("key")) if resolved is not None else None,
litellm.model_cost.get(model_name),
litellm.model_cost.get(f"{custom_llm_provider}/{model_name}"),
)
return any(
entry is not None and any("cost_per" in field and value is not None for field, value in entry.items())
for entry in entries
)
return any(entry is not None and _cost_map_entry_prices_anything(entry) for entry in entries)
def _first_priced_realtime_token_costs(
@ -2917,6 +2930,8 @@ def handle_realtime_stream_cost_calculation(
litellm_model_name: str,
data_residency: str | None = None,
litellm_logging_obj: LitellmLoggingObject | None = None,
custom_pricing_model: str | None = None,
base_pricing_model: str | None = None,
) -> float:
"""
Handles the cost calculation for realtime stream responses.
@ -2925,9 +2940,13 @@ def handle_realtime_stream_cost_calculation(
Args:
results: A list of OpenAIRealtimeStreamBaseObject objects
custom_pricing_model: deployment-scoped pricing key from the deployment's
custom rates, tried ahead of the session-reported model
base_pricing_model: the deployment's resolved base_model, tried ahead of the
session-reported model but after custom rates
"""
received_model = None
potential_model_names: Final = []
potential_model_names: Final = [custom_pricing_model, base_pricing_model]
for result in results:
if result["type"] == "session.created":
received_model = cast(OpenAIRealtimeStreamSessionEvents, result)["session"].get("model", None)
@ -2945,6 +2964,7 @@ def handle_realtime_stream_cost_calculation(
results=results,
custom_llm_provider=custom_llm_provider,
litellm_model_name=litellm_model_name,
custom_pricing_model=custom_pricing_model,
)
if any(r.get("type") == _TRANSCRIPTION_COMPLETED_EVENT_TYPE for r in results)
else 0.0
@ -2968,6 +2988,7 @@ def handle_realtime_transcription_cost_calculation(
results: OpenAIRealtimeStreamList,
custom_llm_provider: str,
litellm_model_name: str,
custom_pricing_model: str | None = None,
) -> float:
"""
Cost for realtime transcription sessions (e.g. gpt-realtime-whisper).
@ -2985,15 +3006,15 @@ def handle_realtime_transcription_cost_calculation(
return 0.0
model_name: Final = _get_transcription_model_name_from_results(results) or litellm_model_name
try:
model_info = litellm.get_model_info(model=model_name, custom_llm_provider=custom_llm_provider)
except Exception:
model_info = None
model_info: Final = _get_model_info_or_none(model_name, custom_llm_provider)
override_info: Final = (
_get_model_info_or_none(custom_pricing_model, custom_llm_provider) if custom_pricing_model is not None else None
)
total_cost = 0.0
for event in completed_events:
usage = event.get("usage") or {}
total_cost += _transcription_usage_cost(usage, model_info)
total_cost += _transcription_usage_cost(usage, model_info, override_info)
return total_cost
@ -3018,23 +3039,57 @@ def _get_transcription_model_name_from_results(
return None
def _transcription_usage_cost(usage: dict, model_info: ModelInfo | None) -> float:
if model_info is None:
def _get_model_info_or_none(model: str, custom_llm_provider: str) -> ModelInfo | None:
try:
return litellm.get_model_info(model=model, custom_llm_provider=custom_llm_provider)
except Exception:
return None
def _declared_transcription_rate(info: ModelInfo | None, keys: tuple[str, ...]) -> float | None:
"""First of ``keys`` this entry prices, read off the raw ``litellm.model_cost`` entry
because ``get_model_info`` synthesizes zero token rates for entries that omit them."""
if info is None:
return None
declared: Final = litellm.model_cost.get(info.get("key"))
if declared is None:
return None
return next(
(float(value) for key in keys if declared.get(key) is not None and (value := info.get(key)) is not None),
None,
)
def _transcription_rate(keys: tuple[str, ...], override: ModelInfo | None, base: ModelInfo | None) -> float:
rates: Final = (_declared_transcription_rate(info, keys) for info in (override, base))
return next((rate for rate in rates if rate is not None), 0.0)
def _transcription_usage_cost(
usage: dict,
model_info: ModelInfo | None,
override_info: ModelInfo | None = None,
) -> float:
if model_info is None and override_info is None:
return 0.0
usage_type: Final = usage.get("type")
if usage_type == "duration":
seconds: Final = usage.get("seconds") or 0.0
per_second: Final = model_info.get("input_cost_per_second") or 0.0
return float(seconds) * float(per_second)
return float(seconds) * _transcription_rate(("input_cost_per_second",), override_info, model_info)
if usage_type == "tokens":
input_token_details: Final = usage.get("input_token_details") or {}
audio_tokens: Final = input_token_details.get("audio_tokens") or 0
text_tokens: Final = input_token_details.get("text_tokens") or 0
output_tokens: Final = usage.get("output_tokens") or 0
audio_cost: Final = float(audio_tokens) * float(
model_info.get("input_cost_per_audio_token") or model_info.get("input_cost_per_token") or 0.0
audio_cost: Final = float(audio_tokens) * _transcription_rate(
("input_cost_per_audio_token", "input_cost_per_token"), override_info, model_info
)
text_cost: Final = float(text_tokens) * _transcription_rate(
("input_cost_per_token",), override_info, model_info
)
output_cost: Final = float(output_tokens) * _transcription_rate(
("output_cost_per_token",), override_info, model_info
)
text_cost: Final = float(text_tokens) * float(model_info.get("input_cost_per_token") or 0.0)
output_cost: Final = float(output_tokens) * float(model_info.get("output_cost_per_token") or 0.0)
return audio_cost + text_cost + output_cost
return 0.0

View file

@ -25,6 +25,21 @@ from litellm.types.llms.vertex_ai import (
from litellm.types.utils import TokenCountResponse
from litellm.utils import supports_response_schema, supports_system_messages
VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS: Final = frozenset(
{
"audio",
"max_retries",
"modalities",
"prediction",
"prompt_cache_key",
"prompt_cache_retention",
"safety_identifier",
"service_tier",
"store",
"web_search_options",
}
)
class VertexAILyriaModelInfo(TypedDict):
vertex_ai_audio_api: ReadOnly[Literal["lyria_predict", "lyria_interactions"]]
@ -370,6 +385,27 @@ def get_vertex_base_model_name(model: str) -> str:
return model
def vertex_model_garden_model_id_in_json_body(model: str) -> bool:
"""
Vertex catalog / publisher models are addressed as publisher/model (e.g.
xai/grok-4.1-fast-reasoning) on the shared OpenAPI URL, with the id in the JSON body.
Deployed Model Garden endpoints are typically a single segment (often numeric)
and use .../endpoints/{ENDPOINT_ID}/chat/completions with an empty model field.
"""
return "/" in model
def is_vertex_self_deployed_openai_compatible_endpoint(model: str) -> bool:
local_model: Final = model.removeprefix("vertex_ai/")
route: Final = get_vertex_ai_model_route(local_model)
if route == VertexAIModelRoute.GEMMA:
return True
return route == VertexAIModelRoute.MODEL_GARDEN and not vertex_model_garden_model_id_in_json_body(
get_vertex_base_model_name(local_model)
)
def get_vertex_ai_fine_tuned_endpoint_id(model: str) -> str | None:
"""
Fine-tuned Gemini deployments are addressed by a numeric endpoint id,

View file

@ -18,7 +18,11 @@ from litellm.types.utils import (
Usage,
)
from ...common_utils import VertexAIError
from ...common_utils import (
VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS,
VertexAIError,
is_vertex_self_deployed_openai_compatible_endpoint,
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer
@ -66,13 +70,15 @@ class VertexAILlama3Config(OpenAIGPTConfig):
and v is not None
}
def get_supported_openai_params(self, model: str):
supported_params: Final = super().get_supported_openai_params(model=model)
try:
supported_params.remove("max_retries")
except KeyError:
pass
return supported_params
def get_supported_openai_params(self, model: str) -> list[str]:
unsupported_params: Final = (
VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS
if is_vertex_self_deployed_openai_compatible_endpoint(model)
else frozenset({"max_retries"})
)
return [ # mutable-ok: get_optional_params extends the returned list with allowed_openai_params
param for param in super().get_supported_openai_params(model=model) if param not in unsupported_params
]
def map_openai_params(
self,

View file

@ -21,6 +21,7 @@ from litellm.llms.custom_httpx.http_handler import (
_get_httpx_client,
)
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
from litellm.llms.vertex_ai.common_utils import VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS
from litellm.types.llms.openai import AllMessageValues
from litellm.types.llms.vertex_ai_gemma import VertexGemmaContainerError
from litellm.types.utils import ModelResponse
@ -49,6 +50,13 @@ class VertexGemmaConfig(OpenAIGPTConfig):
def __init__(self) -> None:
super().__init__()
def get_supported_openai_params(self, model: str) -> list[str]:
return [ # mutable-ok: get_optional_params extends the returned list with allowed_openai_params
param
for param in super().get_supported_openai_params(model=model)
if param not in VERTEX_SELF_DEPLOYED_ENDPOINT_UNSUPPORTED_PARAMS
]
def should_fake_stream(
self,
model: str | None,

View file

@ -24,21 +24,14 @@ import httpx
from litellm.llms.vertex_ai.common_utils import get_vertex_base_url
from litellm.utils import ModelResponse
from ..common_utils import VertexAIError, get_vertex_base_model_name
from ..common_utils import (
VertexAIError,
get_vertex_base_model_name,
vertex_model_garden_model_id_in_json_body,
)
from ..vertex_llm_base import VertexBase
def _vertex_model_garden_model_id_in_json_body(model: str) -> bool:
"""
Vertex catalog / publisher models are addressed as publisher/model (e.g.
xai/grok-4.1-fast-reasoning) on the shared OpenAPI URL, with the id in the JSON body.
Deployed Model Garden endpoints are typically a single segment (often numeric)
and use .../endpoints/{ENDPOINT_ID}/chat/completions with an empty model field.
"""
return "/" in model
def create_vertex_url(
vertex_location: str,
vertex_project: str,
@ -48,7 +41,7 @@ def create_vertex_url(
) -> str:
"""Return the api base for vertex model garden (without /chat/completions)."""
base_url: Final = get_vertex_base_url(vertex_location)
if _vertex_model_garden_model_id_in_json_body(model):
if vertex_model_garden_model_id_in_json_body(model):
return f"{base_url}/v1/projects/{vertex_project}/locations/{vertex_location}/endpoints/openapi"
return f"{base_url}/v1beta1/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}"
@ -124,7 +117,7 @@ class VertexAIModelGardenModels(VertexBase):
)
# Publisher/catalog models: model id must be sent in the JSON body (OpenAPI route).
# Single-segment endpoint ids: model is encoded in the URL path; body model stays empty.
if not _vertex_model_garden_model_id_in_json_body(model):
if not vertex_model_garden_model_id_in_json_body(model):
model = ""
return openai_like_chat_completions.completion(
model=model,

View file

@ -126,6 +126,7 @@ from litellm.types.completion import (
_CompletionDispatchContext,
_CompletionDispatchResult,
)
from litellm.types.litellm_params import RetryStrategy
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import (
CustomPricingLiteLLMParams,
@ -6026,9 +6027,7 @@ def completion_with_retries(*args, **kwargs):
# reset retries in .completion()
kwargs["max_retries"] = 0
kwargs["num_retries"] = 0
retry_strategy: Final[Literal["exponential_backoff_retry", "constant_retry"]] = kwargs.pop(
"retry_strategy", "constant_retry"
)
retry_strategy: Final[RetryStrategy] = kwargs.pop("retry_strategy", "constant_retry")
original_function: Final = kwargs.pop("original_function", completion)
if retry_strategy == "exponential_backoff_retry":
retryer = tenacity.Retrying(
@ -6054,7 +6053,7 @@ async def acompletion_with_retries(*args, **kwargs):
num_retries: Final = kwargs.pop("num_retries", 3)
kwargs["max_retries"] = 0
kwargs["num_retries"] = 0
retry_strategy: Final = kwargs.pop("retry_strategy", "constant_retry")
retry_strategy: Final[RetryStrategy] = kwargs.pop("retry_strategy", "constant_retry")
original_function: Final = kwargs.pop("original_function", completion)
if retry_strategy == "exponential_backoff_retry":
retryer = tenacity.AsyncRetrying(
@ -6082,9 +6081,7 @@ def responses_with_retries(*args, **kwargs):
# reset retries in .responses()
kwargs["max_retries"] = 0
kwargs["num_retries"] = 0
retry_strategy: Final[Literal["exponential_backoff_retry", "constant_retry"]] = kwargs.pop(
"retry_strategy", "constant_retry"
)
retry_strategy: Final[RetryStrategy] = kwargs.pop("retry_strategy", "constant_retry")
original_function: Final = kwargs.pop("original_function", responses)
if retry_strategy == "exponential_backoff_retry":
retryer = tenacity.Retrying(
@ -6111,7 +6108,7 @@ async def aresponses_with_retries(*args, **kwargs):
num_retries: Final = kwargs.pop("num_retries", 3)
kwargs["max_retries"] = 0
kwargs["num_retries"] = 0
retry_strategy: Final = kwargs.pop("retry_strategy", "constant_retry")
retry_strategy: Final[RetryStrategy] = kwargs.pop("retry_strategy", "constant_retry")
original_function: Final = kwargs.pop("original_function", aresponses)
if retry_strategy == "exponential_backoff_retry":
retryer = tenacity.AsyncRetrying(

View file

@ -3342,6 +3342,7 @@
"supports_vision": true
},
"azure/command-r-plus": {
"deprecation_date": "2025-06-30",
"input_cost_per_token": 3e-06,
"litellm_provider": "azure",
"max_input_tokens": 128000,
@ -3349,6 +3350,7 @@
"max_tokens": 4096,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models",
"supports_function_calling": true
},
"azure_ai/claude-haiku-4-5": {
@ -5194,6 +5196,7 @@
"output_cost_per_token": 2e-06
},
"azure/gpt-4": {
"deprecation_date": "2025-06-06",
"input_cost_per_token": 3e-05,
"litellm_provider": "azure",
"max_input_tokens": 8192,
@ -5201,6 +5204,7 @@
"max_tokens": 4096,
"mode": "chat",
"output_cost_per_token": 6e-05,
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models",
"supports_function_calling": true,
"supports_tool_choice": true
},
@ -69074,6 +69078,7 @@
"source": "https://api.together.ai/v1/models"
},
"together_ai/google/gemma-4-31B-it": {
"deprecation_date": "2026-09-14",
"input_cost_per_token": 3.9e-07,
"litellm_provider": "together_ai",
"max_input_tokens": 262144,
@ -69203,6 +69208,7 @@
"source": "https://api.together.ai/v1/models"
},
"together_ai/openai/gpt-oss-20b": {
"deprecation_date": "2026-09-14",
"input_cost_per_token": 5e-08,
"litellm_provider": "together_ai",
"max_input_tokens": 131072,

View file

@ -13,6 +13,7 @@ from typing_extensions import assert_never
import litellm
from litellm._logging import verbose_logger
from litellm.constants import MCP_ALL_TOOLS_WILDCARD
from litellm.proxy._experimental.mcp_server.oauth_utils import (
get_passthrough_resource_metadata_url,
get_passthrough_www_authenticate,
@ -92,7 +93,9 @@ def level_allowed_tools(
answer from this level and intersects with the rest.
1. A legacy ``mcp_tool_permissions`` entry for the server stays a closed
allowlist (``[]`` denies all): allowed = legacy ∪ toolset tools.
allowlist (``[]`` denies all): allowed = legacy ∪ toolset tools. An
entry containing ``MCP_ALL_TOOLS_WILDCARD`` grants every current and
future tool, so the level places no restriction at all.
2. An unconverted row (``mcp_permission_version`` falsy) keeps pre-overrides
behavior: unrestricted unless a toolset names the server.
3. A converted row that does not grant the server places no restriction.
@ -109,6 +112,8 @@ def level_allowed_tools(
toolset: Final[frozenset[str]] = frozenset(toolset_tools or ())
legacy: Final = global_mcp_server_manager.expand_tool_permissions(row.mcp_tool_permissions).get(server_id)
if legacy is not None:
if MCP_ALL_TOOLS_WILDCARD in legacy:
return None
return frozenset(legacy) | toolset
if not row.mcp_permission_version:
return frozenset(toolset) if toolset_tools is not None else None
@ -2196,7 +2201,11 @@ class MCPRequestHandler:
via_toolsets: Sequence[str] | None,
) -> Sequence[str] | None:
"""Union of one level's direct tool grants and its toolset-granted tools on one server,
``None`` when neither source restricts (allow-all from this level)."""
``None`` when neither source restricts (allow-all from this level). A direct grant
containing ``MCP_ALL_TOOLS_WILDCARD`` makes the level unrestricted, so it returns
``None`` whatever the toolsets name."""
if direct is not None and MCP_ALL_TOOLS_WILDCARD in direct:
return None
if direct is None and via_toolsets is None:
return None
return tuple({*(direct or ()), *(via_toolsets or ())})

View file

@ -27,7 +27,8 @@ from collections.abc import (
from contextlib import asynccontextmanager
from dataclasses import dataclass, replace
from functools import lru_cache
from itertools import chain
from itertools import chain, groupby
from operator import itemgetter
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Generic, Literal, TypeAlias, TypedDict, TypeVar, cast
from urllib.parse import ParseResult, urlparse
@ -6857,9 +6858,11 @@ class MCPServerManager:
"""
Rewrite an ``mcp_tool_permissions`` dict keyed by id/name/alias so
every key is a concrete server_id where possible. Tool lists from
keys that point at the same server are unioned, matching the
"duplicate names grant access to all matches" semantics of
``expand_permission_list``.
keys that point at the same server are unioned and deduplicated
first-seen, matching the "duplicate names grant access to all
matches" semantics of ``expand_permission_list``; the
``MCP_ALL_TOOLS_WILDCARD`` entry is preserved as an ordinary list
entry for the caller to interpret.
Required so name-based keys don't silently drop their tool
restrictions when the lookup uses the resolved server_id. Unresolved
@ -6868,11 +6871,15 @@ class MCPServerManager:
"""
if not tool_permissions:
return {}
result: Final[dict[str, list[str]]] = {}
for key, tools in tool_permissions.items():
for server_id in self.expand_permission_list((key,)):
result.setdefault(server_id, []).extend(tools or [])
return result
expanded: Final = tuple(
(server_id, tuple(tools or ()))
for key, tools in tool_permissions.items()
for server_id in self.expand_permission_list([key])
)
return {
server_id: list(dict.fromkeys(tool for _, tools in group for tool in tools))
for server_id, group in groupby(sorted(expanded, key=itemgetter(0)), key=itemgetter(0))
}
def expand_tool_overrides(
self,

View file

@ -89,6 +89,7 @@ from litellm.proxy.common_utils.http_parsing_utils import (
_safe_get_request_headers,
_safe_get_request_query_params,
)
from litellm.proxy.common_utils.model_listing_utils import alias_map
from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time
from litellm.proxy.common_utils.user_api_key_cache import (
END_USER_RESTRICTED_REGISTRY_OVERFLOW_SENTINEL,
@ -722,6 +723,7 @@ async def _run_project_checks(
model=_model,
project_object=project_object,
llm_router=llm_router,
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
)
if not skip_budget_checks:
@ -1018,6 +1020,7 @@ async def common_checks(
team_object=team_object,
llm_router=llm_router,
team_model_aliases=(valid_token.team_model_aliases if valid_token else None),
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
)
except ProxyException as team_denial:
if team_denial.type != ProxyErrorTypes.team_model_access_denied:
@ -1027,6 +1030,7 @@ async def common_checks(
valid_token=valid_token,
team_object=team_object,
llm_router=llm_router,
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
):
raise
@ -1043,6 +1047,7 @@ async def common_checks(
proxy_logging_obj=proxy_logging_obj,
team_membership=loaded_team_membership,
team_membership_loaded=team_membership_loaded,
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
)
# Require trace id for agent keys when agent has require_trace_id_on_calls_by_agent
@ -1081,6 +1086,7 @@ async def common_checks(
model=_model,
llm_router=llm_router,
user_object=user_object,
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
)
# 1.1 - 2.2 - 3.0.2 - 3.0.3: Project checks (blocked, model access, budget)
@ -4349,6 +4355,7 @@ def _can_object_call_model(
models: list[str],
team_model_aliases: dict[str, str] | None = None,
team_id: str | None = None,
key_model_aliases: Mapping[str, str] | None = None,
object_type: Literal["user", "team", "key", "org", "project", "agent"] = "user",
fallback_depth: int = 0,
) -> Literal[True]:
@ -4378,6 +4385,7 @@ def _can_object_call_model(
models=models,
team_model_aliases=team_model_aliases,
team_id=team_id,
key_model_aliases=key_model_aliases,
object_type=object_type,
fallback_depth=fallback_depth + 1,
)
@ -4386,13 +4394,32 @@ def _can_object_call_model(
from litellm.router_strategy.complexity_router.context_compaction import native_compaction_parent
compaction_parent: Final = native_compaction_parent(model)
potential_models: Final = [model, compaction_parent] if compaction_parent is not None else [model]
if model in litellm.model_alias_map:
potential_models.append(litellm.model_alias_map[model])
elif llm_router and model in llm_router.model_group_alias:
_model: Final = llm_router._get_model_from_alias(model)
if _model:
potential_models.append(_model)
global_or_router_alias_target: Final = (
litellm.model_alias_map[model]
if model in litellm.model_alias_map
else (
llm_router._get_model_from_alias(model)
if llm_router is not None and model in llm_router.model_group_alias
else None
)
)
after_team_alias: Final = team_model_aliases.get(model, model) if team_model_aliases else model
after_key_alias: Final = (
key_model_aliases.get(after_team_alias, after_team_alias) if key_model_aliases else after_team_alias
)
after_global_alias: Final = litellm.model_alias_map.get(after_key_alias, after_key_alias)
dispatched_model: Final = (
key_model_aliases.get(after_global_alias, after_global_alias) if key_model_aliases else after_global_alias
)
key_alias_applied: Final = after_key_alias != after_team_alias or dispatched_model != after_global_alias
potential_models: Final = (
(dispatched_model,)
if key_alias_applied
else (
*((model, compaction_parent) if compaction_parent is not None else (model,)),
*((global_or_router_alias_target,) if global_or_router_alias_target else ()),
)
)
## check model access for alias + underlying model - allow if either is in allowed models
for m in potential_models:
@ -4418,6 +4445,35 @@ def _can_object_call_model(
)
def _resolve_team_alias(
model: str | list[str],
team_model_aliases: dict[str, str] | None,
team_id: str | None,
llm_router: Router | None,
) -> str | list[str]:
if not team_model_aliases:
return model
if isinstance(model, str):
return _live_team_alias_target(model, team_model_aliases, team_id, llm_router)
return [ # mutable-ok: _can_object_call_model takes list[str]
_live_team_alias_target(name, team_model_aliases, team_id, llm_router) for name in model
]
def _live_team_alias_target(
model: str, team_model_aliases: dict[str, str], team_id: str | None, llm_router: Router | None
) -> str:
target: Final = team_model_aliases.get(model)
if target is None:
return model
deleted_team_deployment: Final = (
llm_router is not None
and target.startswith(f"model_name_{team_id}_")
and target not in llm_router.model_name_to_deployment_indices
)
return model if deleted_team_deployment else target
async def _check_agent_access_group_model_access(
model: str | list[str] | None, # mutable-ok: _can_object_call_model and the client message helper take list[str]
valid_token: UserAPIKeyAuth | None,
@ -4438,12 +4494,14 @@ async def _check_agent_access_group_model_access(
param="model",
code=status.HTTP_403_FORBIDDEN,
)
dispatched: Final = _resolve_team_alias(model, valid_token.team_model_aliases, valid_token.team_id, llm_router)
return _can_object_call_model(
model=model,
model=dispatched,
llm_router=llm_router,
models=sorted(ceiling.models),
team_id=valid_token.team_id,
object_type="agent",
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
)
@ -4471,12 +4529,14 @@ async def _check_agent_caller_model_access(
if caller_auth is None:
return
caller_team: Final = await load_team(valid_token)
caller_key_model_aliases: Final = key_model_aliases_for_auth_check(valid_token)
if caller_team is not None:
await can_team_access_model(
model=model,
team_object=caller_team,
llm_router=llm_router,
prisma_client=prisma_client,
key_model_aliases=caller_key_model_aliases,
)
await _check_team_member_model_access(
model=model,
@ -4486,12 +4546,18 @@ async def _check_agent_caller_model_access(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
key_model_aliases=caller_key_model_aliases,
)
return
caller_user: Final = await load_user(valid_token)
if caller_user is None:
return
await can_user_call_model(model=model, llm_router=llm_router, user_object=caller_user)
await can_user_call_model(
model=model,
llm_router=llm_router,
user_object=caller_user,
key_model_aliases=caller_key_model_aliases,
)
def _model_in_team_aliases(model: str, team_model_aliases: dict[str, str] | None = None) -> bool:
@ -4512,6 +4578,10 @@ def _model_in_team_aliases(model: str, team_model_aliases: dict[str, str] | None
return False
def key_model_aliases_for_auth_check(valid_token: UserAPIKeyAuth | None) -> Mapping[str, str] | None:
return alias_map(valid_token.aliases) if valid_token is not None and valid_token.aliases else None
def _resolve_key_models_for_auth_check(valid_token: UserAPIKeyAuth) -> list[str]:
"""
Expand key model sentinels before auth checks.
@ -4831,6 +4901,7 @@ async def can_key_call_model(
models=key_models,
team_model_aliases=valid_token.team_model_aliases,
team_id=valid_token.team_id,
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
object_type="key",
)
except ProxyException:
@ -4848,6 +4919,7 @@ async def can_key_call_model(
models=models_from_groups,
team_model_aliases=valid_token.team_model_aliases,
team_id=valid_token.team_id,
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
object_type="key",
)
raise
@ -4906,6 +4978,7 @@ async def can_key_call_resolved_model(
team_object=team_object,
llm_router=llm_router,
team_model_aliases=valid_token.team_model_aliases,
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
)
except ProxyException as team_denial:
if team_denial.type != ProxyErrorTypes.team_model_access_denied:
@ -4915,6 +4988,7 @@ async def can_key_call_resolved_model(
valid_token=valid_token,
team_object=team_object,
llm_router=llm_router,
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
):
raise
@ -4927,6 +5001,7 @@ async def can_key_call_resolved_model(
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
)
if valid_token.project_id is not None:
@ -4941,6 +5016,7 @@ async def can_key_call_resolved_model(
model=model,
project_object=project_object,
llm_router=llm_router,
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
)
@ -4968,6 +5044,7 @@ async def can_team_access_model(
team_object: LiteLLM_TeamTable | None,
llm_router: Router | None,
team_model_aliases: dict[str, str] | None = None,
key_model_aliases: Mapping[str, str] | None = None,
prisma_client: DatabaseClient | None = None,
) -> Literal[True]:
"""
@ -4983,6 +5060,7 @@ async def can_team_access_model(
models=team_object.models if team_object else [],
team_model_aliases=team_model_aliases,
team_id=team_object.team_id if team_object else None,
key_model_aliases=key_model_aliases,
object_type="team",
)
except ProxyException:
@ -5000,6 +5078,7 @@ async def can_team_access_model(
models=list(dict.fromkeys([*(team_object.models if team_object else []), *models_from_groups])),
team_model_aliases=team_model_aliases,
team_id=team_object.team_id if team_object else None,
key_model_aliases=key_model_aliases,
object_type="team",
)
raise
@ -5058,6 +5137,7 @@ async def _key_access_group_grants_model(
valid_token: UserAPIKeyAuth | None,
team_object: LiteLLM_TeamTable | None,
llm_router: Router | None,
key_model_aliases: Mapping[str, str] | None = None,
) -> bool:
"""
Returns True if the key's `access_group_ids` expand to models that grant
@ -5078,6 +5158,7 @@ async def _key_access_group_grants_model(
models=authorized_models,
team_model_aliases=valid_token.team_model_aliases if valid_token else None,
team_id=valid_token.team_id if valid_token else None,
key_model_aliases=key_model_aliases,
object_type="key",
)
return True
@ -5089,6 +5170,7 @@ def can_project_access_model(
model: str | list[str],
project_object: LiteLLM_ProjectTable,
llm_router: Router | None,
key_model_aliases: Mapping[str, str] | None = None,
) -> Literal[True]:
"""
Returns True if the project can access a specific model.
@ -5099,6 +5181,7 @@ def can_project_access_model(
model=model,
llm_router=llm_router,
models=project_object.models if project_object else [],
key_model_aliases=key_model_aliases,
object_type="project",
)
@ -5107,6 +5190,7 @@ async def can_user_call_model(
model: str | list[str],
llm_router: Router | None,
user_object: LiteLLM_UserTable | None,
key_model_aliases: Mapping[str, str] | None = None,
) -> Literal[True]:
if user_object is None:
return True
@ -5128,6 +5212,7 @@ async def can_user_call_model(
model=model,
llm_router=llm_router,
models=user_object.models,
key_model_aliases=key_model_aliases,
object_type="user",
)
@ -5682,6 +5767,7 @@ async def _check_team_member_model_access(
proxy_logging_obj: ProxyLogging,
team_membership: LiteLLM_TeamMembership | None = None,
team_membership_loaded: bool = False,
key_model_aliases: Mapping[str, str] | None = None,
) -> None:
"""
Check if a team member's per-member model scope allows access to the requested model.
@ -5717,6 +5803,7 @@ async def _check_team_member_model_access(
models=member_allowed_models,
object_type="team",
team_id=team_object.team_id,
key_model_aliases=key_model_aliases,
)
except ProxyException:
internal_message: Final = (

View file

@ -66,6 +66,7 @@ from litellm.proxy.auth.auth_checks import (
get_user_object,
is_valid_fallback_model,
jwt_key_mapping_cache_key,
key_model_aliases_for_auth_check,
resolve_and_validate_end_user_id,
resolve_default_end_user_budget,
)
@ -469,6 +470,7 @@ async def _check_key_model_budget_with_fallback(
models=valid_token.team_models,
team_model_aliases=valid_token.team_model_aliases,
team_id=valid_token.team_id,
key_model_aliases=key_model_aliases_for_auth_check(valid_token),
object_type="team",
)
except ProxyException:

View file

@ -180,7 +180,7 @@ def caller_alias_maps(
return CallerAliases((team_aliases, key_aliases), (team_aliases, key_aliases, litellm.model_alias_map, key_aliases))
def _alias_map(aliases: object) -> Mapping[str, str]:
def alias_map(aliases: object) -> Mapping[str, str]:
try:
entries: Final = _ALIAS_ENTRIES.validate_python(aliases, strict=True)
except ValidationError:
@ -204,7 +204,7 @@ def alias_target(model_id: str, aliases: CallerAliases, listed: Container[str] =
already `listed` keeps its own row, so it is never rewritten."""
if model_id in listed:
return None
return _rewrite(model_id, tuple(_alias_map(alias_map) for alias_map in aliases.rewrite))
return _rewrite(model_id, tuple(alias_map(raw) for raw in aliases.rewrite))
def alias_listing_entries(
@ -213,8 +213,8 @@ def alias_listing_entries(
) -> tuple[tuple[str, str], ...]:
"""`entries` plus one `(alias, lookup_id)` row per key or team alias whose target is
listed. An alias colliding with a listed id keeps the listed entry."""
maps: Final = tuple(_alias_map(alias_map) for alias_map in aliases.rewrite)
own: Final = tuple(_alias_map(alias_map) for alias_map in aliases.own)
maps: Final = tuple(alias_map(raw) for raw in aliases.rewrite)
own: Final = tuple(alias_map(raw) for raw in aliases.own)
lookup_by_response: Final = MappingProxyType(dict(entries))
lookup_ids: Final = frozenset(lookup_by_response.values())
targets: Final = MappingProxyType(

View file

@ -105,6 +105,8 @@ from litellm.proxy.route_llm_request import ProxyModelNotFoundError
from litellm.proxy.utils import normalize_route_for_root_path
from litellm.repositories.team_repository import TeamRepository
from litellm.secret_managers.main import get_secret_str
from litellm.types import utils as types_utils
from litellm.types.litellm_params import ProxyRequestState, wire_names
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
@ -133,6 +135,9 @@ router: Final = APIRouter()
pass_through_endpoint_logging: Final = PassThroughEndpointLogging()
_METADATA_KEYS: Final = frozenset(("litellm_metadata", "metadata"))
_KEPT_OUT_OF_LITELLM_PARAMS: Final = _METADATA_KEYS | frozenset(wire_names(ProxyRequestState))
# Global registry to track registered pass-through routes and prevent memory leaks
_registered_pass_through_routes: Final[dict[str, dict[str, str | bool | list[str] | Mapping[str, object]]]] = {}
@ -578,21 +583,21 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
"""
Filter out litellm params from the request body
"""
from litellm.types.utils import all_litellm_params
_parsed_body = _parsed_body or {}
litellm_params_in_body: Final = {}
for k in all_litellm_params:
if k in _parsed_body:
litellm_params_in_body[k] = _parsed_body.pop(k, None)
litellm_keys_in_body: Final = MappingProxyType(
{k: _parsed_body.pop(k) for k in types_utils.all_litellm_params if k in _parsed_body}
)
litellm_params_in_body: Final = MappingProxyType(
{k: v for k, v in litellm_keys_in_body.items() if k not in _KEPT_OUT_OF_LITELLM_PARAMS}
)
_metadata = dict(
LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict)
)
litellm_metadata: Final = litellm_params_in_body.pop("litellm_metadata", None)
metadata: Final = litellm_params_in_body.pop("metadata", None)
litellm_metadata: Final = litellm_keys_in_body.get("litellm_metadata")
metadata: Final = litellm_keys_in_body.get("metadata")
if litellm_metadata:
_metadata.update(litellm_metadata)
if metadata:

View file

@ -259,6 +259,7 @@ from litellm.router_utils.routing_groups import (
validate_routing_strategy,
)
from litellm.scheduler import FlowItem, Scheduler
from litellm.types.litellm_params import RoutingStrategyName
from litellm.types.llms.openai import (
AllMessageValues,
ChatCompletionToolParam,
@ -796,15 +797,7 @@ class Router:
allowed_fails_policy: AllowedFailsPolicy | None = None, # set custom allowed fails policy
cooldown_time: float | None = None, # (seconds) time to cooldown a deployment after failure
disable_cooldowns: bool | None = None,
routing_strategy: Literal[
"simple-shuffle",
"least-busy",
"usage-based-routing",
"latency-based-routing",
"cost-based-routing",
"usage-based-routing-v2",
"lar1",
] = "simple-shuffle",
routing_strategy: RoutingStrategyName = "simple-shuffle",
optional_pre_call_checks: OptionalPreCallChecks | None = None,
routing_strategy_args: dict = {}, # just for latency-based
routing_groups: list[RoutingGroup | dict] | None = None,
@ -2940,7 +2933,7 @@ class Router:
self._update_kwargs_before_fallbacks(model=model_group, kwargs=initial_kwargs)
fallback_response = await self.async_function_with_fallbacks_common_utils(
e=e,
disable_fallbacks=False,
disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs),
fallbacks=fallbacks,
context_window_fallbacks=context_window_fallbacks,
content_policy_fallbacks=content_policy_fallbacks,
@ -3384,7 +3377,7 @@ class Router:
)
fallback_response = await self.async_function_with_fallbacks_common_utils(
e=fallback_trigger,
disable_fallbacks=False,
disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs),
fallbacks=fallbacks,
context_window_fallbacks=context_window_fallbacks,
content_policy_fallbacks=content_policy_fallbacks,
@ -3475,8 +3468,9 @@ class Router:
for item in model_response:
yield item
except MidStreamFallbackError as e:
if not e.is_pre_first_chunk and (
e.generated_content or _stream_chunks_have_generated_content(model_response.chunks)
if fallbacks_disabled_for_request(initial_kwargs) or (
not e.is_pre_first_chunk
and (e.generated_content or _stream_chunks_have_generated_content(model_response.chunks))
):
if e.original_exception is not None:
raise e.original_exception from e
@ -5611,7 +5605,7 @@ class Router:
)
fallback_response = await self.async_function_with_fallbacks_common_utils( # rebind-ok: set on success
e=fallback_trigger,
disable_fallbacks=False,
disable_fallbacks=fallbacks_disabled_for_request(initial_kwargs),
fallbacks=fallbacks,
context_window_fallbacks=context_window_fallbacks,
content_policy_fallbacks=content_policy_fallbacks,

View file

@ -3,8 +3,10 @@ from typing import Any, Final
from pydantic import BaseModel, Field
CHAT_COMPLETION_AGENTIC_SURFACE: Final = "chat_completions"
RESPONSES_AGENTIC_SURFACE: Final = "responses"
from litellm.types.litellm_params import AgenticSurface
CHAT_COMPLETION_AGENTIC_SURFACE: Final[AgenticSurface] = "chat_completions"
RESPONSES_AGENTIC_SURFACE: Final[AgenticSurface] = "responses"
CODE_INTERPRETER_INTERCEPTION_PREFIX: Final = "_code_interpreter_interception"
HEADROOM_INTERCEPTION_PREFIX: Final = "_headroom_interception"
HEADROOM_CONVERTED_STREAM_KEY: Final = f"{HEADROOM_INTERCEPTION_PREFIX}_converted_stream"

View file

@ -0,0 +1,364 @@
"""LiteLLM-owned request kwargs declared as typed fields; types/utils.py splices these with the callback and pricing
models and KWARG_ARTIFACTS into all_litellm_params."""
from collections.abc import Callable, Iterator, Mapping, MutableMapping, Sequence
from dataclasses import dataclass, field, fields, is_dataclass
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Literal, TypeAlias
if TYPE_CHECKING:
import httpx
from aiohttp import ClientSession
from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.router_strategy.complexity_router.context_compaction import CompactionState
from litellm.router_utils.fallback_event_handlers import AttemptedFallbackTargets
from litellm.types.caching import DynamicCacheControl
from litellm.types.llms.openai import ChatCompletionAssistantMessage, ChatCompletionUserMessage
from litellm.types.proxy.litellm_pre_call_utils import SecretFields
from litellm.types.router import ConfigurableClientsideParamsCustomAuth, DeploymentTypedDict, RetryPolicy
from litellm.types.router_weights import RouterWeights
from litellm.types.utils import ModelResponse, ModelResponseStream, ProviderSpecificHeader
ProviderClient: TypeAlias = (
OpenAI
| AsyncOpenAI
| AzureOpenAI
| AsyncAzureOpenAI
| HTTPHandler
| AsyncHTTPHandler
| httpx.Client
| httpx.AsyncClient
)
MockResponse: TypeAlias = (
str | Exception | Mapping[str, object] | Sequence[float] | ModelResponse | ModelResponseStream
)
RetryStrategy: TypeAlias = Literal["constant_retry", "exponential_backoff_retry"]
AgenticSurface: TypeAlias = Literal["chat_completions", "responses"]
RoutingStrategyName: TypeAlias = Literal[
"simple-shuffle",
"least-busy",
"usage-based-routing",
"latency-based-routing",
"cost-based-routing",
"usage-based-routing-v2",
"lar1",
]
TRUSTED_CALLBACK_VARS_FIELD: Final = "litellm_trusted_callback_vars"
ADDRESSED_RESPONSE_ID_FIELD: Final = "_litellm_addressed_response_id"
WIRE_NAME: Final = "wire_name"
def wire(name: str) -> Mapping[str, str]:
return MappingProxyType({WIRE_NAME: name})
@dataclass(frozen=True, slots=True, kw_only=True)
class ProviderConnection:
api_key: str | None = None
api_base: str | None = None
api_version: str | None = None
region_name: str | None = None
headers: Mapping[str, str] | None = None
provider_specific_header: "ProviderSpecificHeader | Sequence[ProviderSpecificHeader] | None" = None
client: "ProviderClient | None" = None
shared_session: "ClientSession | None" = None
ssl_verify: bool | str | None = None
request_timeout: float | None = None
force_timeout: float | None = None
stream_timeout: float | str | None = None
max_retries: int | None = None
tenant_id: str | None = None
client_id: str | None = None
client_secret: str | None = None
azure_username: str | None = None
azure_password: str | None = None
azure_scope: str | None = None
azure_ad_token_provider: Callable[[], str] | None = None
litellm_credential_name: str | None = None
configurable_clientside_auth_params: "Sequence[str | ConfigurableClientsideParamsCustomAuth] | None" = None
use_xai_oauth: bool | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class BedrockBatchConnection:
# Bedrock rejects these names in request bodies, so register them as LiteLLM-owned
aws_batch_role_arn: str | None = None
s3_bucket_name: str | None = None
s3_region_name: str | None = None
s3_endpoint_url: str | None = None
s3_output_bucket_name: str | None = None
s3_bucket_owner: str | None = None
s3_access_key_id: str | None = None
s3_secret_access_key: str | None = None
s3_encryption_key_id: str | None = None
bedrock_tags: Sequence[Mapping[str, str]] | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class ConnectionSettings:
provider: ProviderConnection
bedrock_batch: BedrockBatchConnection
@dataclass(frozen=True, slots=True, kw_only=True)
class DispatchOptions:
custom_llm_provider: str | None = None
azure: bool | None = None
use_litellm_proxy: bool | None = None
use_chat_completions_api: bool | None = None
use_in_pass_through: bool | None = None
allowed_openai_params: Sequence[str] | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class RoutingOptions:
fallbacks: Sequence[str | Mapping[str, object]] | None = None
context_window_fallback_dict: Mapping[str, str] | None = None
num_retries: int | None = None
retry_policy: "RetryPolicy | Mapping[str, object] | None" = None
retry_strategy: RetryStrategy | None = None
routing_strategy: RoutingStrategyName | None = None
cooldown_time: float | None = None
allowed_model_region: str | None = None
enable_tag_filtering: bool | None = None
fastest_response: bool | None = None
provider_affinity_header: str | None = None
search_tool_name: str | None = None
model_list: "Sequence[DeploymentTypedDict] | None" = None
@dataclass(frozen=True, slots=True, kw_only=True)
class DeploymentOptions:
model_info: Mapping[str, object] | None = None
rpm: int | None = None
tpm: int | None = None
itpm: int | None = None
otpm: int | None = None
default_api_key_rpm_limit: int | None = None
default_api_key_tpm_limit: int | None = None
max_parallel_requests: int | None = None
weight: int | None = None
order: int | None = None
tag_regex: Sequence[str] | None = None
max_file_size_mb: float | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class SpecializedRouterOptions:
auto_router_config_path: str | None = None
auto_router_config: str | None = None
auto_router_default_model: str | None = None
auto_router_embedding_model: str | None = None
auto_router_max_input_chars: int | None = None
auto_router_routing_compression: str | None = None
auto_router_model_compression: str | None = None
complexity_router_config: Mapping[str, object] | None = None
complexity_router_default_model: str | None = None
adaptive_router_config: Mapping[str, object] | None = None
adaptive_router_default_model: str | None = None
quality_router_config: Mapping[str, object] | None = None
quality_router_default_model: str | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class CachingOptions:
caching: bool | None = None
cache: "DynamicCacheControl | None" = None
ttl: float | None = None
enable_prompt_caching: bool | None = None
caching_groups: Sequence[Sequence[str]] | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class CostOptions:
cost_per_query: float | None = None
base_model: str | None = None
max_budget: float | None = None
budget_duration: str | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class ObservabilityOptions:
id: str | None = None
metadata: MutableMapping[str, object] | None = None # mutable-ok: the router and logging write keys into it
litellm_metadata: MutableMapping[str, object] | None = None # mutable-ok: the proxy writes keys into it
tags: Sequence[str] | None = None
litellm_trace_id: str | None = None
litellm_session_id: str | None = None
litellm_request_debug: bool | None = None
logger_fn: Callable[[Mapping[str, object]], None] | None = None
verbose: bool | None = None
no_log: bool | None = field(default=None, metadata=wire("no-log"))
@dataclass(frozen=True, slots=True, kw_only=True)
class AgenticLoopOptions:
max_agentic_loops: int | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class GuardrailOptions:
guardrails: Sequence[str] | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class PromptOptions:
prompt_id: str | None = None
prompt_variables: Mapping[str, object] | None = None
prompt_version: str | None = None
prompt_environment: str | None = None
prompt_label: str | None = None
litellm_system_prompt: str | None = None
custom_prompt_dict: Mapping[str, object] | None = None
roles: Mapping[str, object] | None = None
final_prompt_value: str | None = None
bos_token: str | None = None
eos_token: str | None = None
hf_model_name: str | None = None
supports_system_message: bool | None = None
ensure_alternating_roles: bool | None = None
user_continue_message: "ChatCompletionUserMessage | None" = None
assistant_continue_message: "ChatCompletionAssistantMessage | None" = None
disable_add_transform_inline_image_block: bool | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class ResponseOptions:
merge_reasoning_content_in_choices: bool | None = None
enable_json_schema_validation: bool | None = None
complete_response: bool | None = None
stream_chunk_size: int | None = None
keepalive_seconds: float | None = None
allow_client_keepalive_override: bool | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class MockOptions:
mock_response: "MockResponse | None" = None
mock_timeout: bool | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class LiteLLMOptions:
dispatch: DispatchOptions
routing: RoutingOptions
deployment: DeploymentOptions
specialized_routers: SpecializedRouterOptions
caching: CachingOptions
cost: CostOptions
observability: ObservabilityOptions
agentic_loop: AgenticLoopOptions
guardrails: GuardrailOptions
prompt: PromptOptions
response: ResponseOptions
mock: MockOptions
@dataclass(frozen=True, slots=True, kw_only=True)
class CallState:
litellm_call_id: str | None = None
completion_call_id: str | None = None
model_alias_map: Mapping[str, str] | None = None
data_residency: str | None = None
litellm_logging_obj: "Logging | None" = None
preset_cache_key: str | None = None
cache_key: str | None = None
stream_response: "Mapping[str, ModelResponse] | None" = None
context_compaction_state: "CompactionState | None" = field(default=None, metadata=wire("_context_compaction_state"))
@dataclass(frozen=True, slots=True, kw_only=True)
class AgenticLoopState:
depth: int | None = field(default=None, metadata=wire("_agentic_loop_depth"))
fingerprints: Sequence[str] | None = field(default=None, metadata=wire("_agentic_loop_fingerprints"))
api_surface: Literal["chat_completions", "responses"] | None = field(
default=None, metadata=wire("_agentic_loop_api_surface")
)
code_interpreter_active: bool | None = field(default=None, metadata=wire("_code_interpreter_interception_active"))
code_interpreter_sandbox_key: str | None = field(
default=None, metadata=wire("_code_interpreter_interception_sandbox_key")
)
code_interpreter_session_scoped: bool | None = field(
default=None, metadata=wire("_code_interpreter_interception_session_scoped")
)
code_interpreter_converted_stream: bool | None = field(
default=None, metadata=wire("_code_interpreter_interception_converted_stream")
)
websearch_emit_native_blocks: bool | None = field(
default=None, metadata=wire("_websearch_interception_emit_native_blocks")
)
websearch_converted_stream: bool | None = field(
default=None, metadata=wire("_websearch_interception_converted_stream")
)
headroom_converted_stream: bool | None = field(
default=None, metadata=wire("_headroom_interception_converted_stream")
)
@dataclass(frozen=True, slots=True, kw_only=True)
class RouterState:
weights: "RouterWeights | None" = field(default=None, metadata=wire("_router_weights"))
fallback_depth: int | None = None
max_fallbacks: int | None = None
attempted_targets: "AttemptedFallbackTargets | None" = None
@dataclass(frozen=True, slots=True, kw_only=True)
class ProxyRequestState:
proxy_server_request: Mapping[str, object] | None = None
secret_fields: "SecretFields | None" = None
trusted_callback_vars: Mapping[str, str] | None = field(default=None, metadata=wire(TRUSTED_CALLBACK_VARS_FIELD))
addressed_response_id: str | None = field(default=None, metadata=wire(ADDRESSED_RESPONSE_ID_FIELD))
strip_stream_usage: bool | None = field(default=None, metadata=wire("_litellm_strip_stream_usage"))
client_side_timeout: bool | None = None
model_file_id_mapping: Mapping[str, Mapping[str, str]] | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class EntrypointState:
acompletion: bool | None = None
aembedding: bool | None = None
aimg_generation: bool | None = None
atext_completion: bool | None = None
text_completion: bool | None = None
allm_passthrough_route: bool | None = None
async_call: bool | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class InternalState:
call: CallState
agentic_loop: AgenticLoopState
router: RouterState
proxy: ProxyRequestState
entrypoint: EntrypointState
KWARG_ARTIFACTS: Final[tuple[str, ...]] = ("self", "use_client", "model_config", "rust")
LITELLM_OWNED_ROOTS: Final = (ConnectionSettings, LiteLLMOptions, InternalState)
def wire_names(owner: type) -> tuple[str, ...]:
return tuple(owned.metadata.get(WIRE_NAME, owned.name) for owned in fields(owner))
def owned_wire_names(root: type) -> tuple[str, ...]:
def names() -> Iterator[str]:
for leaf in fields(root):
if not is_dataclass(leaf.type):
raise TypeError(f"{root.__name__}.{leaf.name} is not a dataclass leaf")
yield from wire_names(leaf.type) # pyright: ignore[reportArgumentType] # Field.type admits str
return tuple(names())
OWNED_KWARG_NAMES: Final = tuple(name for root in LITELLM_OWNED_ROOTS for name in owned_wire_names(root))
AGENTIC_LOOP_KWARG_NAMES: Final = (*wire_names(AgenticLoopState), *wire_names(AgenticLoopOptions))
BEDROCK_BATCH_KWARG_NAMES: Final = wire_names(BedrockBatchConnection)

View file

@ -24,6 +24,7 @@ if TYPE_CHECKING:
from .completion import CompletionRequest
from .embedding import EmbeddingRequest
from .litellm_params import RoutingStrategyName
from .llms.bedrock import AwsSessionTag
from .llms.openai import OpenAIFileObject
from .search import SearchProvider
@ -104,12 +105,7 @@ class RouterConfig(BaseModel):
context_window_fallbacks: list | None = []
model_group_alias: dict[str, list[str]] | None = {}
retry_after: int | None = 0
routing_strategy: Literal[
"simple-shuffle",
"least-busy",
"usage-based-routing",
"latency-based-routing",
] = "simple-shuffle"
routing_strategy: RoutingStrategyName = "simple-shuffle"
routing_groups: list[RoutingGroup] | None = None
model_config = ConfigDict(protected_namespaces=())

View file

@ -56,8 +56,15 @@ from litellm.types.llms.base import (
from litellm.types.mcp import MCPServerCostInfo
from ..litellm_core_utils.core_helpers import map_finish_reason, process_response_headers
from . import litellm_params as _litellm_params
from .agents import LiteLLMSendMessageResponse
from .guardrails import GuardrailEventHooks
from .litellm_params import (
AGENTIC_LOOP_KWARG_NAMES,
BEDROCK_BATCH_KWARG_NAMES,
KWARG_ARTIFACTS,
OWNED_KWARG_NAMES,
)
from .llms.anthropic_messages.anthropic_response import AnthropicMessagesResponse
from .llms.base import HiddenParams
from .llms.openai import (
@ -3901,205 +3908,20 @@ def pricing_override_fields(*sources: Mapping[str, object]) -> tuple[str, ...]:
)
# Server-controlled fields that bound or drive an interceptor's agentic loop
# (depth, cycle fingerprints, ceiling, code-interpreter sandbox state). Listed
# in all_litellm_params so they are treated as LiteLLM-level and excluded from
# get_non_default_completion_params; otherwise the OpenAI param builder sweeps
# any unrecognized top-level key into extra_body and leaks them to the provider.
# This is what lets the loop carry state across rerun calls without a provider
# scrubber.
agentic_loop_internal_litellm_params: Final = [
"_agentic_loop_depth",
"_agentic_loop_fingerprints",
"_agentic_loop_api_surface",
"max_agentic_loops",
"_code_interpreter_interception_active",
"_code_interpreter_interception_sandbox_key",
"_code_interpreter_interception_session_scoped",
"_code_interpreter_interception_converted_stream",
"_websearch_interception_emit_native_blocks",
"_websearch_interception_converted_stream",
"_headroom_interception_converted_stream",
agentic_loop_internal_litellm_params: Final = list(AGENTIC_LOOP_KWARG_NAMES) # mutable-ok: public type stays a list
bedrock_batch_litellm_params: Final = BEDROCK_BATCH_KWARG_NAMES
TRUSTED_CALLBACK_VARS_FIELD: Final = _litellm_params.TRUSTED_CALLBACK_VARS_FIELD
ADDRESSED_RESPONSE_ID_FIELD: Final = _litellm_params.ADDRESSED_RESPONSE_ID_FIELD
all_litellm_params = [ # rebind-ok: two star imports in litellm/__init__.py re-bind it # mutable-ok: callers concat
*OWNED_KWARG_NAMES,
*KWARG_ARTIFACTS,
*StandardCallbackDynamicParams.__annotations__,
*CustomPricingLiteLLMParams.model_fields,
]
# Proxy-owned callback credentials, stamped from admin-configured team/key callback
# settings. Listed in all_litellm_params for the same reason as the agentic-loop
# fields above: an unrecognized top-level key is swept into extra_body and sent to
# the provider.
TRUSTED_CALLBACK_VARS_FIELD: Final = "litellm_trusted_callback_vars"
ADDRESSED_RESPONSE_ID_FIELD: Final = "_litellm_addressed_response_id"
# Bedrock managed-batch deployment config, read from litellm_params by the batch and
# files transformations. Listed for the same reason as the fields above: these sit on
# a deployment that also serves chat, so leaking them into extra_body makes Bedrock
# reject every non-batch request to that deployment.
bedrock_batch_litellm_params: Final = (
"aws_batch_role_arn",
"s3_bucket_name",
"s3_region_name",
"s3_endpoint_url",
"s3_output_bucket_name",
"s3_bucket_owner",
"s3_access_key_id",
"s3_secret_access_key",
"s3_encryption_key_id",
"bedrock_tags",
)
all_litellm_params = (
agentic_loop_internal_litellm_params
+ [TRUSTED_CALLBACK_VARS_FIELD, ADDRESSED_RESPONSE_ID_FIELD, *bedrock_batch_litellm_params]
+ [
"_context_compaction_state",
"metadata",
"litellm_metadata",
"keepalive_seconds",
"allow_client_keepalive_override",
"litellm_trace_id",
"litellm_request_debug",
"guardrails",
"tags",
"acompletion",
"aimg_generation",
"atext_completion",
"text_completion",
"caching",
"mock_response",
"mock_timeout",
"disable_add_transform_inline_image_block",
"api_key",
"api_version",
"prompt_id",
"prompt_variables",
"litellm_system_prompt",
"provider_specific_header",
"prompt_version",
"prompt_environment",
"api_base",
"force_timeout",
"logger_fn",
"verbose",
"custom_llm_provider",
"model_file_id_mapping",
"litellm_logging_obj",
"litellm_call_id",
"completion_call_id",
"model_alias_map",
"custom_prompt_dict",
"stream_response",
"cost_per_query",
"ssl_verify",
"data_residency",
"async_call",
"aembedding",
"allm_passthrough_route",
"_litellm_strip_stream_usage",
"use_client",
"id",
"fallbacks",
"routing_strategy",
"_router_weights",
"azure",
"headers",
"model_list",
"num_retries",
"context_window_fallback_dict",
"retry_policy",
"retry_strategy",
"roles",
"final_prompt_value",
"bos_token",
"eos_token",
"request_timeout",
"client_side_timeout",
"complete_response",
"self",
"client",
"rpm",
"tpm",
"default_api_key_rpm_limit",
"default_api_key_tpm_limit",
"itpm",
"otpm",
"max_parallel_requests",
"input_cost_per_token",
"output_cost_per_token",
"input_cost_per_second",
"output_cost_per_second",
"hf_model_name",
"model_info",
"proxy_server_request",
"secret_fields",
"preset_cache_key",
"caching_groups",
"ttl",
"cache",
"enable_prompt_caching",
"no-log",
"base_model",
"stream_timeout",
"stream_chunk_size",
"supports_system_message",
"region_name",
"allowed_model_region",
"model_config",
"fastest_response",
"cooldown_time",
"cache_key",
"max_retries",
"azure_ad_token_provider",
"tenant_id",
"client_id",
"azure_username",
"azure_password",
"azure_scope",
"client_secret",
"user_continue_message",
"configurable_clientside_auth_params",
"weight",
"ensure_alternating_roles",
"assistant_continue_message",
"user_continue_message",
"fallback_depth",
"max_fallbacks",
"attempted_targets",
"max_budget",
"budget_duration",
"use_in_pass_through",
"merge_reasoning_content_in_choices",
"litellm_credential_name",
"allowed_openai_params",
"litellm_session_id",
"provider_affinity_header",
"use_litellm_proxy",
"use_chat_completions_api",
"rust",
"prompt_label",
"shared_session",
"search_tool_name",
"order",
"enable_tag_filtering",
"enable_json_schema_validation",
"use_xai_oauth",
"auto_router_config_path",
"auto_router_config",
"auto_router_default_model",
"auto_router_embedding_model",
"auto_router_max_input_chars",
"auto_router_routing_compression",
"auto_router_model_compression",
"complexity_router_config",
"complexity_router_default_model",
"adaptive_router_config",
"adaptive_router_default_model",
"quality_router_config",
"quality_router_default_model",
]
+ list(StandardCallbackDynamicParams.__annotations__.keys())
+ list(CustomPricingLiteLLMParams.model_fields.keys())
)
class KeyGenerationConfig(TypedDict, total=False):
required_params: list[str] # specify params that must be present in the key generation request

View file

@ -3342,6 +3342,7 @@
"supports_vision": true
},
"azure/command-r-plus": {
"deprecation_date": "2025-06-30",
"input_cost_per_token": 3e-06,
"litellm_provider": "azure",
"max_input_tokens": 128000,
@ -3349,6 +3350,7 @@
"max_tokens": 4096,
"mode": "chat",
"output_cost_per_token": 1.5e-05,
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models",
"supports_function_calling": true
},
"azure_ai/claude-haiku-4-5": {
@ -5194,6 +5196,7 @@
"output_cost_per_token": 2e-06
},
"azure/gpt-4": {
"deprecation_date": "2025-06-06",
"input_cost_per_token": 3e-05,
"litellm_provider": "azure",
"max_input_tokens": 8192,
@ -5201,6 +5204,7 @@
"max_tokens": 4096,
"mode": "chat",
"output_cost_per_token": 6e-05,
"source": "https://learn.microsoft.com/en-us/azure/foundry/openai/concepts/retired-models",
"supports_function_calling": true,
"supports_tool_choice": true
},
@ -69074,6 +69078,7 @@
"source": "https://api.together.ai/v1/models"
},
"together_ai/google/gemma-4-31B-it": {
"deprecation_date": "2026-09-14",
"input_cost_per_token": 3.9e-07,
"litellm_provider": "together_ai",
"max_input_tokens": 262144,
@ -69203,6 +69208,7 @@
"source": "https://api.together.ai/v1/models"
},
"together_ai/openai/gpt-oss-20b": {
"deprecation_date": "2026-09-14",
"input_cost_per_token": 5e-08,
"litellm_provider": "together_ai",
"max_input_tokens": 131072,

View file

@ -75,8 +75,8 @@ proxy = [
"mcp>=2.2.0,<3",
"httpx2>=2.5.0,<3",
"pydantic>=2.12.0,<3",
"litellm-proxy-extras==0.4.101",
"litellm-enterprise==0.1.70",
"litellm-proxy-extras==0.4.102",
"litellm-enterprise==0.1.71",
"RestrictedPython>=8.5,<9.0",
"rich>=13.9.4,<14.0",
"InquirerPy>=0.3.4,<1.0",

View file

@ -0,0 +1,72 @@
import uuid
from typing import Final
import httpx
from tests.integration._support.client import Gateway, eventually, object_value, string_value
def _listed_model_ids(response: httpx.Response) -> frozenset[str]:
entries: Final = response.json()["data"]
assert isinstance(entries, list), response.text
return frozenset(string_value(object_value(entry)["id"]) for entry in entries)
def _listed_and_callable(gateway: Gateway, key: str, model: str, alias: str) -> None:
"""Every id /v1/models lists for this key must be callable by the same key."""
response: Final = eventually(
lambda: gateway.request("GET", "/v1/models", key=key),
lambda value: value.status_code == 200 and _listed_model_ids(value) == frozenset({model, alias}),
return_last_on_timeout=True,
)
assert response.status_code == 200, response.text
listed: Final = _listed_model_ids(response)
assert listed == frozenset({model, alias}), response.text
for model_id in sorted(listed):
called: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model_id, "messages": [{"role": "user", "content": "ping"}]},
key=key,
)
assert called.status_code == 200, f"listed id {model_id} is not callable: {called.status_code} {called.text}"
def test_key_alias_listed_by_v1_models_is_callable(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
alias: Final = f"integration-alias-{uuid.uuid4().hex}"
key: Final = scenario.key(models=[model], aliases={alias: model})
_listed_and_callable(gateway, key, model, alias)
def test_team_key_alias_listed_by_v1_models_is_callable(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
team_id: Final = scenario.team(models=[model])
alias: Final = f"integration-alias-{uuid.uuid4().hex}"
key: Final = scenario.key(team_id=team_id, aliases={alias: model})
_listed_and_callable(gateway, key, model, alias)
def test_key_alias_to_model_outside_key_allowlist_is_hidden_and_denied(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
allowed: Final = scenario.model()
hidden: Final = scenario.model()
alias: Final = f"integration-alias-{uuid.uuid4().hex}"
key: Final = scenario.key(models=[allowed], aliases={alias: hidden})
response: Final = eventually(
lambda: gateway.request("GET", "/v1/models", key=key),
lambda value: value.status_code == 200 and _listed_model_ids(value) == frozenset({allowed}),
return_last_on_timeout=True,
)
assert response.status_code == 200, response.text
assert _listed_model_ids(response) == frozenset({allowed}), response.text
called: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": alias, "messages": [{"role": "user", "content": "ping"}]},
key=key,
)
assert called.status_code == 403, called.text
assert "key_model_access_denied" in called.text, called.text

View file

@ -1,6 +1,9 @@
import asyncio
import json
import os
import uuid
from collections.abc import Iterator, Mapping
from hashlib import sha256
from pathlib import Path
from typing import Final
@ -12,6 +15,96 @@ from litellm import get_model_info
from tests.integration._support.client import Gateway, eventually, object_value, string_value
from tests.integration._support.database import read_rows
from tests.integration._support.process import owned_proxy
from tests.integration._support.upstream import delete_scenario, register_scenario
from tests.integration.cost_calculation.cost_tracking_case import RealtimeResponse
from tests.integration.pricing.test_realtime_cached_audio_pricing import one_realtime_turn
REALTIME_MODEL: Final = "gpt-realtime-2"
REALTIME_INPUT_TEXT_TOKENS: Final = 10
REALTIME_INPUT_AUDIO_TOKENS: Final = 20
REALTIME_OUTPUT_TEXT_TOKENS: Final = 5
REALTIME_OUTPUT_AUDIO_TOKENS: Final = 7
def _realtime_response_done() -> RealtimeResponse:
return RealtimeResponse(
content_type="application/x-realtime",
events=(
{
"type": "response.done",
"event_id": "evt_$REQUEST_ID",
"response": {
"id": "resp_$REQUEST_ID",
"object": "realtime.response",
"status": "completed",
"output": [],
"usage": {
"total_tokens": REALTIME_INPUT_TEXT_TOKENS
+ REALTIME_INPUT_AUDIO_TOKENS
+ REALTIME_OUTPUT_TEXT_TOKENS
+ REALTIME_OUTPUT_AUDIO_TOKENS,
"input_tokens": REALTIME_INPUT_TEXT_TOKENS + REALTIME_INPUT_AUDIO_TOKENS,
"output_tokens": REALTIME_OUTPUT_TEXT_TOKENS + REALTIME_OUTPUT_AUDIO_TOKENS,
"input_token_details": {
"text_tokens": REALTIME_INPUT_TEXT_TOKENS,
"audio_tokens": REALTIME_INPUT_AUDIO_TOKENS,
"cached_tokens": 0,
},
"output_token_details": {
"text_tokens": REALTIME_OUTPUT_TEXT_TOKENS,
"audio_tokens": REALTIME_OUTPUT_AUDIO_TOKENS,
},
},
},
},
),
)
@pytest.mark.parametrize(
("input_text_rate", "input_audio_rate", "output_text_rate", "output_audio_rate"),
((0.001, 0.002, 0.003, 0.004), (0.0, 0.0, 0.0, 0.0)),
ids=("custom_rates", "zero_rated"),
)
def test_realtime_session_is_charged_at_the_deployment_configured_rates(
gateway: Gateway,
input_text_rate: float,
input_audio_rate: float,
output_text_rate: float,
output_audio_rate: float,
) -> None:
with gateway.scenario() as scenario:
scenario_id: Final = f"realtime-configured-price-{uuid.uuid4().hex[:12]}"
handle: Final = register_scenario(scenario_id, _realtime_response_done())
scenario.cleanups.callback(delete_scenario, handle)
key: Final = scenario.key()
model: Final = scenario.model(
model=f"openai/{REALTIME_MODEL}",
api_key=scenario_id,
api_base=gateway.upstream_url.rstrip("/"),
input_cost_per_token=input_text_rate,
input_cost_per_audio_token=input_audio_rate,
output_cost_per_token=output_text_rate,
output_cost_per_audio_token=output_audio_rate,
)
session: Final = asyncio.run(one_realtime_turn(os.environ["INTEGRATION_PROXY_URL"].rstrip("/"), key, model))
assert session.get("type") == "session.created", session
rows: Final = eventually(
lambda: read_rows(
'SELECT spend, call_type FROM "LiteLLM_SpendLogs" WHERE api_key = %s',
(sha256(key.encode()).hexdigest(),),
),
lambda values: len(values) == 1,
seconds=70,
)
assert rows[0]["call_type"] == "_arealtime", rows
assert float(str(rows[0]["spend"])) == pytest.approx(
REALTIME_INPUT_TEXT_TOKENS * input_text_rate
+ REALTIME_INPUT_AUDIO_TOKENS * input_audio_rate
+ REALTIME_OUTPUT_TEXT_TOKENS * output_text_rate
+ REALTIME_OUTPUT_AUDIO_TOKENS * output_audio_rate,
abs=1e-9,
), rows
@pytest.mark.covers("quota_management.spend_tracking.custom_price.matches_input_rates")

View file

@ -92,7 +92,7 @@ def cached_audio_response_done() -> RealtimeResponse:
)
async def _one_realtime_turn(proxy_url: str, key: str, model: str) -> dict[str, JsonValue]:
async def one_realtime_turn(proxy_url: str, key: str, model: str) -> dict[str, JsonValue]:
async with websockets.connect(
f"{proxy_url.replace('http://', 'ws://').replace('https://', 'wss://')}/v1/realtime?model={model}",
additional_headers={"Authorization": f"Bearer {key}"},
@ -115,7 +115,7 @@ def test_realtime_cached_audio_tokens_bill_at_audio_cache_read_rate_not_full_aud
model: Final = scenario.model(
model=f"openai/{MODEL}", api_key=scenario_id, api_base=gateway.upstream_url.rstrip("/")
)
session: Final = asyncio.run(_one_realtime_turn(os.environ["INTEGRATION_PROXY_URL"].rstrip("/"), key, model))
session: Final = asyncio.run(one_realtime_turn(os.environ["INTEGRATION_PROXY_URL"].rstrip("/"), key, model))
assert session.get("type") == "session.created", session
rows: Final = eventually(
lambda: read_rows(

View file

@ -834,12 +834,12 @@ async def test_key_model_list(model_access, model_access_level, model_endpoint):
assert len(model_list["data"]) > 0
if model_access == "gpt-3.5-turbo":
if model_endpoint == "/v1/models":
assert (
len(model_list["data"]) == 1
), "model_access={}, model_access_level={}".format(
assert {entry["id"] for entry in model_list["data"]} == {
model_access,
"mistral-7b",
}, "generate_key sets alias mistral-7b -> gpt-3.5-turbo, so /v1/models lists both; model_access={}, model_access_level={}".format(
model_access, model_access_level
)
assert model_list["data"][0]["id"] == model_access
elif model_endpoint == "/model/info":
assert isinstance(model_list["data"], list)
assert len(model_list["data"]) == 1

View file

@ -1046,11 +1046,26 @@ def test_translate_anthropic_to_openai_sets_prompt_cache_key_for_azure(model: st
assert openai_request["prompt_cache_key"] == "session-abc"
@pytest.mark.parametrize(
"model",
[
"vertex_ai/meta/llama-4-maverick-17b-128e-instruct-maas",
"vertex_ai/moonshotai/kimi-k2-thinking-maas",
"vertex_ai/xai/grok-4.1-fast-non-reasoning",
],
)
def test_translate_anthropic_to_openai_sets_prompt_cache_key_for_vertex_maas_models(model: str):
openai_request = _translate_with_metadata(model, {"user_id": CLAUDE_CODE_USER_ID}, "vertex_ai")
assert openai_request["prompt_cache_key"] == "session-abc"
@pytest.mark.parametrize(
"model, custom_llm_provider",
[
("gemini/gemini-2.5-pro", "gemini"),
("vertex_ai/gemini-2.5-pro", "vertex_ai"),
("vertex_ai/gemma/gemma-2-2b-it", "vertex_ai"),
("vertex_ai/openai/mg-endpoint-lit8592", "vertex_ai"),
("anthropic/claude-sonnet-4-5", "anthropic"),
("bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0", "bedrock"),
("no-such-model-lit5875", "no-such-provider-lit5875"),

View file

@ -8,10 +8,10 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import litellm
from litellm.llms.vertex_ai.vertex_model_garden.main import (
_vertex_model_garden_model_id_in_json_body,
create_vertex_url,
from litellm.llms.vertex_ai.common_utils import (
vertex_model_garden_model_id_in_json_body,
)
from litellm.llms.vertex_ai.vertex_model_garden.main import create_vertex_url
@pytest.mark.parametrize(
@ -43,11 +43,8 @@ def test_create_vertex_url_openapi_vs_deployed_endpoint(
def test_model_id_in_json_body_heuristic() -> None:
assert (
_vertex_model_garden_model_id_in_json_body("xai/grok-4.1-fast-reasoning")
is True
)
assert _vertex_model_garden_model_id_in_json_body("5464397967697903616") is False
assert vertex_model_garden_model_id_in_json_body("xai/grok-4.1-fast-reasoning") is True
assert vertex_model_garden_model_id_in_json_body("5464397967697903616") is False
@pytest.fixture

View file

@ -346,6 +346,15 @@ class TestMCPRequestHandler:
mock_manager.discovered_inventory = MagicMock(return_value=inventory or {})
return mock_manager
def _real_manager_with_toolsets(self, toolset_perms):
"""A real MCPServerManager so the real expand_tool_permissions runs;
only the DB-backed toolset lookup is stubbed"""
from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager
manager = MCPServerManager()
manager.resolve_toolset_tool_permissions = AsyncMock(return_value=toolset_perms)
return manager
async def test_get_allowed_mcp_servers_for_key_includes_toolset_servers(self):
"""A key granted only mcp_toolsets must reach the toolset's servers on
every path (list, call, REST); regression for the list-ok/call-403 bug"""
@ -538,6 +547,141 @@ class TestMCPRequestHandler:
assert result is None
@pytest.mark.parametrize(
"direct,via_toolsets,expected",
[
(["*"], None, None),
(["*"], ["read_file"], None),
(None, None, None),
([], None, ()),
(None, ["read_file"], ("read_file",)),
],
)
def test_union_tool_grants_wildcard_and_union_cases(self, direct, via_toolsets, expected):
"""A direct ["*"] makes the level unrestricted even beside a toolset
list (regression: mapping ["*"] to None in expand_tool_permissions let
a same-level toolset list deny every other tool)"""
result = MCPRequestHandler._union_tool_grants(direct, via_toolsets)
if expected is None:
assert result is None
else:
assert result is not None
assert set(result) == set(expected)
def test_union_tool_grants_unions_two_concrete_lists(self):
result = MCPRequestHandler._union_tool_grants(["read_file"], ["write_file"])
assert result is not None
assert set(result) == {"read_file", "write_file"}
async def test_key_wildcard_allows_a_tool_never_enumerated(self):
"""End to end at the key level: object_permission sits on the auth
object already, no team named, so no patching is needed; the real
global manager expands ["*"] and the level reads unrestricted"""
user_api_key_auth = UserAPIKeyAuth(
api_key="test-key",
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="perm-1",
mcp_tool_permissions={"server-a": ["*"]},
),
)
allowed = await MCPRequestHandler.get_allowed_tools_for_server(
server_id="server-a", user_api_key_auth=user_api_key_auth
)
brand_new_tool_allowed = await MCPRequestHandler.is_tool_allowed_for_server(
tool_name="brand_new_tool", server_id="server-a", user_api_key_auth=user_api_key_auth
)
assert allowed is None
assert brand_new_tool_allowed is True
async def test_key_wildcard_stays_capped_by_team_allowlist(self):
"""A wildcard on the key must never widen a team's enumerated ceiling:
the intersection keeps only the team's named tools"""
user_api_key_auth = UserAPIKeyAuth(
api_key="test-key",
team_id="team-1",
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="perm-1",
mcp_tool_permissions={"server-a": ["*"]},
),
)
team_object_permission = self._toolset_only_object_permission([])
team_object_permission.mcp_tool_permissions = {"server-a": ["read_file"]}
manager = self._real_manager_with_toolsets({})
with (
patch.object( # test-quality-ok: stub the DB team loader to drive the real team-server resolution path
MCPRequestHandler, "_get_team_object_permission", AsyncMock(return_value=team_object_permission)
),
patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
manager,
),
):
allowed = await MCPRequestHandler.get_allowed_tools_for_server(
server_id="server-a", user_api_key_auth=user_api_key_auth
)
brand_new_tool_allowed = await MCPRequestHandler.is_tool_allowed_for_server(
tool_name="brand_new_tool", server_id="server-a", user_api_key_auth=user_api_key_auth
)
assert allowed == ["read_file"]
assert brand_new_tool_allowed is False
async def test_team_wildcard_stays_capped_by_key_allowlist(self):
"""A wildcard on the team leaves the key's enumerated list as the
effective ceiling"""
user_api_key_auth = UserAPIKeyAuth(
api_key="test-key",
team_id="team-1",
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="perm-1",
mcp_tool_permissions={"server-a": ["read_file"]},
),
)
team_object_permission = self._toolset_only_object_permission([])
team_object_permission.mcp_tool_permissions = {"server-a": ["*"]}
manager = self._real_manager_with_toolsets({})
with (
patch.object( # test-quality-ok: stub the DB team loader to drive the real team-server resolution path
MCPRequestHandler, "_get_team_object_permission", AsyncMock(return_value=team_object_permission)
),
patch( # test-quality-ok: isolate the MCP registry, same seam as the sibling tests
"litellm.proxy._experimental.mcp_server.mcp_server_manager.global_mcp_server_manager",
manager,
),
):
allowed = await MCPRequestHandler.get_allowed_tools_for_server(
server_id="server-a", user_api_key_auth=user_api_key_auth
)
assert allowed == ["read_file"]
async def test_key_empty_tool_list_stays_deny_all(self):
"""[] on the key is deny-all, distinct from the wildcard: it must not
be widened into allow-all"""
user_api_key_auth = UserAPIKeyAuth(
api_key="test-key",
object_permission=LiteLLM_ObjectPermissionTable(
object_permission_id="perm-1",
mcp_tool_permissions={"server-a": []},
),
)
allowed = await MCPRequestHandler.get_allowed_tools_for_server(
server_id="server-a", user_api_key_auth=user_api_key_auth
)
read_file_allowed = await MCPRequestHandler.is_tool_allowed_for_server(
tool_name="read_file", server_id="server-a", user_api_key_auth=user_api_key_auth
)
assert allowed == []
assert read_file_allowed is False
# ------------------------------------------------------------------
# LIT-5749: toolsets attached to a TEAM, ORG, or internal USER must be
# enforced exactly like inline tool allowlists, on both axes

View file

@ -8783,6 +8783,35 @@ class TestMCPServerManagerExpandToolPermissions:
result = manager.expand_tool_permissions({"uuid-a": ["read_file"], "alias-a": ["write_file"]})
assert sorted(result["uuid-a"]) == ["read_file", "write_file"]
def test_wildcard_survives_expansion_as_list_entry(self):
"""["*"] stays in the expanded list so the caller's wildcard check
(``_union_tool_grants``) can read it; this function only normalizes
keys and never maps grants to None."""
manager = MCPServerManager()
manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="alpha")
result = manager.expand_tool_permissions({"uuid-a": ["*"]})
assert result == {"uuid-a": ["*"]}
def test_wildcard_unions_with_concrete_names_across_keys_for_same_server(self):
"""An alias key carrying ["*"] unioned with an id key naming one tool
keeps both entries; interpretation of the wildcard belongs to the
caller, not the expansion."""
manager = MCPServerManager()
manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="alias-a", alias="alias-a")
result = manager.expand_tool_permissions({"uuid-a": ["read_file"], "alias-a": ["*"]})
assert sorted(result["uuid-a"]) == ["*", "read_file"]
def test_empty_list_stays_deny_all(self):
"""[] is deny-all, a distinct meaning from no entry (unrestricted);
the key must survive expansion rather than disappear."""
manager = MCPServerManager()
manager.config_mcp_servers["uuid-a"] = self._make_server("uuid-a", server_name="alpha")
result = manager.expand_tool_permissions({"uuid-a": []})
assert result == {"uuid-a": []}
class TestOAuthDiscoverySSRFGuard:
"""SSRF guard for the OAuth metadata discovery follow-up fetches.

View file

@ -1783,6 +1783,336 @@ def test_can_object_call_model_access_via_alias_only():
assert result is True
def test_can_object_call_model_key_alias_to_allowed_target_is_allowed():
"""A key alias whose target is on the key allowlist resolves like a team alias."""
from litellm.proxy.auth.auth_checks import _can_object_call_model
result = _can_object_call_model(
model="mistral-7b",
llm_router=None,
models=["gpt-4o-mini"],
key_model_aliases={"mistral-7b": "gpt-4o-mini"},
object_type="key",
fallback_depth=0,
)
assert result is True
def test_can_object_call_model_key_alias_to_disallowed_target_is_denied():
"""A key alias whose target is outside the key allowlist stays denied."""
from litellm.proxy._types import ProxyErrorTypes, ProxyException
from litellm.proxy.auth.auth_checks import _can_object_call_model
with pytest.raises(ProxyException) as exc_info:
_can_object_call_model(
model="mistral-7b",
llm_router=None,
models=["gpt-4o-mini"],
key_model_aliases={"mistral-7b": "gpt-4"},
object_type="key",
fallback_depth=0,
)
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied
assert exc_info.value.code == "403"
@pytest.mark.asyncio
async def test_can_team_access_model_honors_key_alias():
"""A key on a team can call a model through its own alias when the target is on the team allowlist."""
from litellm.proxy.auth.auth_checks import can_team_access_model
team_object = LiteLLM_TeamTable(
team_id="team-123",
models=["gpt-4o-mini"],
)
assert (
await can_team_access_model(
model="mistral-7b",
team_object=team_object,
llm_router=None,
key_model_aliases={"mistral-7b": "gpt-4o-mini"},
)
is True
)
with pytest.raises(ProxyException) as exc_info:
await can_team_access_model(
model="mistral-7b",
team_object=team_object,
llm_router=None,
)
assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied
@pytest.mark.asyncio
async def test_can_key_call_model_honors_key_alias():
"""The real key entry point resolves a key alias to its target before the allowlist check."""
from litellm.proxy.auth.auth_checks import can_key_call_model
allowed_token = UserAPIKeyAuth(
api_key="sk-test",
models=["gpt-4o-mini"],
aliases={"mistral-7b": "gpt-4o-mini"},
)
assert (
await can_key_call_model(
model="mistral-7b",
llm_model_list=None,
valid_token=allowed_token,
llm_router=None,
)
is True
)
denied_token = UserAPIKeyAuth(
api_key="sk-test",
models=["gpt-4o-mini"],
aliases={"mistral-7b": "gpt-4"},
)
with pytest.raises(ProxyException) as exc_info:
await can_key_call_model(
model="mistral-7b",
llm_model_list=None,
valid_token=denied_token,
llm_router=None,
)
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied
def test_can_object_call_model_key_alias_applies_before_global_alias(monkeypatch):
"""The key alias rewrite precedes the global one at dispatch, so the key target is authorized."""
from litellm.proxy.auth.auth_checks import _can_object_call_model
monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"})
assert (
_can_object_call_model(
model="foo",
llm_router=None,
models=["baz"],
key_model_aliases={"foo": "baz"},
object_type="key",
fallback_depth=0,
)
is True
)
with pytest.raises(ProxyException) as exc_info:
_can_object_call_model(
model="foo",
llm_router=None,
models=["bar"],
key_model_aliases={"foo": "baz"},
object_type="key",
fallback_depth=0,
)
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied
def test_can_object_call_model_key_alias_matches_global_rewritten_name(monkeypatch):
"""A key alias on the globally rewritten name resolves the same way the request chain does."""
from litellm.proxy.auth.auth_checks import _can_object_call_model
monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"})
assert (
_can_object_call_model(
model="foo",
llm_router=None,
models=["baz"],
key_model_aliases={"bar": "baz"},
object_type="key",
fallback_depth=0,
)
is True
)
def test_can_object_call_model_chained_alias_requires_final_target(monkeypatch):
"""When a key alias fires on the globally rewritten name, only the final target is dispatched."""
from litellm.proxy.auth.auth_checks import _can_object_call_model
monkeypatch.setattr(litellm, "model_alias_map", {"foo": "bar"})
with pytest.raises(ProxyException) as exc_info:
_can_object_call_model(
model="foo",
llm_router=None,
models=["bar"],
key_model_aliases={"bar": "baz"},
object_type="key",
fallback_depth=0,
)
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied
assert (
_can_object_call_model(
model="foo",
llm_router=None,
models=["baz"],
key_model_aliases={"bar": "baz"},
object_type="key",
fallback_depth=0,
)
is True
)
def test_can_object_call_model_key_alias_name_alone_is_not_enough():
"""A key that may call the alias name but not its target cannot call the alias."""
from litellm.proxy.auth.auth_checks import _can_object_call_model
with pytest.raises(ProxyException) as exc_info:
_can_object_call_model(
model="bar",
llm_router=None,
models=["bar"],
key_model_aliases={"bar": "baz"},
object_type="key",
fallback_depth=0,
)
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied
assert (
_can_object_call_model(
model="bar",
llm_router=None,
models=["baz"],
key_model_aliases={"bar": "baz"},
object_type="key",
fallback_depth=0,
)
is True
)
def test_can_object_call_model_team_alias_applies_before_key_alias():
"""A key alias on the raw name loses to the team alias that rewrites it first at dispatch."""
from litellm.proxy.auth.auth_checks import _can_object_call_model
assert (
_can_object_call_model(
model="foo",
llm_router=None,
models=["bar"],
team_model_aliases={"foo": "bar"},
key_model_aliases={"foo": "baz"},
object_type="key",
fallback_depth=0,
)
is True
)
def test_can_object_call_model_key_alias_on_team_alias_target():
"""A key alias on the team-rewritten name resolves like the dispatch chain does."""
from litellm.proxy.auth.auth_checks import _can_object_call_model
assert (
_can_object_call_model(
model="foo",
llm_router=None,
models=["baz"],
team_model_aliases={"foo": "bar"},
key_model_aliases={"bar": "baz"},
object_type="key",
fallback_depth=0,
)
is True
)
with pytest.raises(ProxyException) as exc_info:
_can_object_call_model(
model="foo",
llm_router=None,
models=["bar"],
team_model_aliases={"foo": "bar"},
key_model_aliases={"bar": "baz"},
object_type="key",
fallback_depth=0,
)
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied
@pytest.mark.asyncio
async def test_can_user_call_model_honors_key_alias():
"""A personal-scope key alias resolves to its target before the user allowlist check."""
from litellm.proxy.auth.auth_checks import can_user_call_model
user_object = LiteLLM_UserTable(user_id="test-user", models=["gpt-4o-mini"])
assert (
await can_user_call_model(
model="mistral-7b",
llm_router=None,
user_object=user_object,
key_model_aliases={"mistral-7b": "gpt-4o-mini"},
)
is True
)
with pytest.raises(ProxyException) as exc_info:
await can_user_call_model(
model="mistral-7b",
llm_router=None,
user_object=user_object,
)
assert exc_info.value.type == ProxyErrorTypes.user_model_access_denied
@pytest.mark.asyncio
async def test_check_team_member_model_access_honors_key_alias():
"""A key alias resolves against the member allowlist, not just the raw alias name."""
from litellm.proxy._types import LiteLLM_TeamMembership
from litellm.proxy.auth.auth_checks import _check_team_member_model_access
membership = LiteLLM_TeamMembership(
user_id="alice",
team_id="team-a",
litellm_budget_table=LiteLLM_BudgetTable(allowed_models=["gpt-4o-mini"]),
)
await _check_team_member_model_access(
model="mistral-7b",
team_object=LiteLLM_TeamTable(team_id="team-a"),
valid_token=UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a"),
llm_router=None,
prisma_client=None,
user_api_key_cache=UserApiKeyCache(),
proxy_logging_obj=MagicMock(),
team_membership=membership,
team_membership_loaded=True,
key_model_aliases={"mistral-7b": "gpt-4o-mini"},
)
with pytest.raises(ProxyException) as exc_info:
await _check_team_member_model_access(
model="mistral-7b",
team_object=LiteLLM_TeamTable(team_id="team-a"),
valid_token=UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a"),
llm_router=None,
prisma_client=None,
user_api_key_cache=UserApiKeyCache(),
proxy_logging_obj=MagicMock(),
team_membership=membership,
team_membership_loaded=True,
)
assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied
def test_can_object_call_model_access_via_underlying_model_only():
"""
Test that a key can access a model via underlying model even when using an alias.
@ -9139,6 +9469,50 @@ async def test_agent_access_groups_cap_models_even_when_key_allows_them():
assert asked == ["agent-1", "agent-1"]
@pytest.mark.asyncio
async def test_agent_access_group_ceiling_admits_the_key_alias_target():
agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=["gpt-5"])
agent_key.aliases = {"fast": "gpt-5"}
resolve, asked = _agent_model_ceiling_resolver(frozenset({"gpt-5"}))
assert await _check_agent_access_group_model_access("fast", agent_key, None, resolve) is True
assert asked == ["agent-1"]
@pytest.mark.asyncio
async def test_agent_access_group_ceiling_checks_the_team_alias_target():
agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", team_id="team-1", models=["gpt-5"])
agent_key.team_model_aliases = {"foo": "gpt-5"}
resolve, asked = _agent_model_ceiling_resolver(frozenset({"gpt-5"}))
assert await _check_agent_access_group_model_access("foo", agent_key, None, resolve) is True
assert asked == ["agent-1"]
@pytest.mark.asyncio
async def test_agent_access_group_ceiling_denies_a_team_alias_outside_the_ceiling():
agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", team_id="team-1", models=["gpt-5"])
agent_key.team_model_aliases = {"foo": "claude-sonnet-4-5"}
resolve, _ = _agent_model_ceiling_resolver(frozenset({"gpt-5"}))
with pytest.raises(ModelAccessDeniedProxyException) as exc:
await _check_agent_access_group_model_access("foo", agent_key, None, resolve)
assert exc.value.type == ProxyErrorTypes.agent_model_access_denied
@pytest.mark.asyncio
async def test_agent_access_group_ceiling_keeps_the_name_for_a_deleted_team_deployment():
from litellm.router import Router
agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", team_id="team-1", models=["foo"])
agent_key.team_model_aliases = {"foo": "model_name_team-1_deadbeef"}
router: Final = Router(model_list=[])
resolve, asked = _agent_model_ceiling_resolver(frozenset({"foo"}))
assert await _check_agent_access_group_model_access("foo", agent_key, router, resolve) is True
assert asked == ["agent-1"]
@pytest.mark.asyncio
async def test_agent_access_groups_naming_no_model_deny_every_model():
agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=[])
@ -9437,6 +9811,18 @@ async def test_agent_key_acting_for_a_teamless_user_is_capped_at_that_users_mode
assert asked == ["team:None", "user:alice", "team:None", "user:alice"]
@pytest.mark.asyncio
async def test_agent_key_alias_resolves_against_the_echoed_teams_models():
agent_key: Final = _agent_key_acting_for(user_id="alice", team_id="team-a")
agent_key.aliases = {"foo": "bar"}
load_team, load_user, asked = _caller_loaders(LiteLLM_TeamTable(team_id="team-a", models=["bar"]), None)
cache: Final = await _cache_with_membership("alice", "team-a", allowed_models=None)
await _check_caller_models(agent_key, "foo", load_team, load_user, cache)
assert asked == ["team:team-a"]
@pytest.mark.asyncio
async def test_agent_key_without_an_echoed_caller_keeps_its_own_models():
agent_key: Final = UserAPIKeyAuth(token="agent-token", agent_id="agent-1", models=["gpt-5", "claude-sonnet"])

View file

@ -5,10 +5,11 @@ import logging
import os
import sys
import zlib
from collections.abc import Callable
from collections.abc import Callable, Mapping
from contextlib import ExitStack, contextmanager
from dataclasses import dataclass
from io import BytesIO
from types import SimpleNamespace
from types import MappingProxyType, SimpleNamespace
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
@ -16,7 +17,7 @@ import httpx
import pytest
from fastapi import HTTPException, Request, Response, UploadFile
from fastapi.responses import StreamingResponse
from pydantic import ValidationError
from pydantic import TypeAdapter, ValidationError
from starlette.datastructures import FormData, Headers, QueryParams
from starlette.datastructures import UploadFile as StarletteUploadFile
@ -45,6 +46,7 @@ from litellm.proxy.pass_through_endpoints.success_handler import (
PassThroughEndpointLogging,
)
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
from litellm.types import utils as types_utils
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_DEPLOYMENT_MODEL_INFO_STATE_KEY,
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
@ -7305,6 +7307,156 @@ def test_passthrough_logs_the_resolved_deployment_model_info_over_the_request_bo
assert kwargs["litellm_params"]["metadata"]["model_info"] == {"id": "vertex-gemini-38-flash-dep"}
@dataclass(frozen=True, slots=True, kw_only=True)
class _PassThroughSplit:
litellm_params: Mapping[str, object]
forwarded_body: Mapping[str, object]
_LITELLM_PARAMS: Final = TypeAdapter(dict[str, object])
_PROXY_SERVER_REQUEST: Final = TypeAdapter(dict[str, object])
def _split_pass_through_body(body: str) -> _PassThroughSplit:
mock_request: Final = MagicMock(spec=Request)
mock_request.method = "POST"
mock_request.url = "http://0.0.0.0:4000/gemini/v1beta/models/gemini-2.5-flash:generateContent"
mock_request.headers = Headers()
mock_request.scope = MappingProxyType({})
init_kwargs_for_pass_through_endpoint: Final = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint # pyright: ignore[reportUnknownVariableType, reportUnknownMemberType] # untyped legacy helper
kwargs: Final = init_kwargs_for_pass_through_endpoint(
request=mock_request,
user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key"),
passthrough_logging_payload=MagicMock(),
logging_obj=MagicMock(),
_parsed_body=json.loads(body),
litellm_call_id="lit-owned-keys-call-id",
)
validate_litellm_params: Final = _LITELLM_PARAMS.validate_python # pyright: ignore[reportUnknownArgumentType] # untyped legacy helper
litellm_params: Final = validate_litellm_params(kwargs["litellm_params"])
return _PassThroughSplit(
litellm_params=MappingProxyType(litellm_params),
forwarded_body=MappingProxyType(
_LITELLM_PARAMS.validate_python(
_PROXY_SERVER_REQUEST.validate_python(litellm_params["proxy_server_request"])["body"]
)
),
)
GEMINI_BODY: Final = '{"contents": [{"parts": [{"text": "hi"}]}], "generationConfig": {"temperature": 0}}'
def _metadata_of(split: _PassThroughSplit) -> Mapping[str, object]:
return MappingProxyType(_LITELLM_PARAMS.validate_python(split.litellm_params["metadata"]))
def test_passthrough_moves_every_litellm_owned_key_from_the_forwarded_body_into_litellm_params() -> None:
split: Final = _split_pass_through_body(
'{"ttl": 30, "contents": [{"parts": [{"text": "hi"}]}], "num_retries": 2,'
' "generationConfig": {"temperature": 0}, "litellm_trace_id": "trace-a"}'
)
assert frozenset(split.litellm_params) == frozenset(
("ttl", "num_retries", "litellm_trace_id", "metadata", "proxy_server_request")
)
assert tuple(split.litellm_params[k] for k in ("ttl", "num_retries", "litellm_trace_id")) == (30, 2, "trace-a")
assert split.forwarded_body == json.loads(GEMINI_BODY)
PROXY_STAMPED_NAMES: Final = frozenset(
(
"proxy_server_request",
"secret_fields",
"litellm_trusted_callback_vars",
"_litellm_addressed_response_id",
"_litellm_strip_stream_usage",
"client_side_timeout",
"model_file_id_mapping",
)
)
@pytest.mark.parametrize(
"name",
sorted(frozenset(litellm.all_litellm_params) - frozenset(("metadata", "litellm_metadata")) - PROXY_STAMPED_NAMES),
)
def test_passthrough_keeps_each_registered_litellm_owned_name_out_of_the_forwarded_body(name: str) -> None:
split: Final = _split_pass_through_body(json.dumps({name: "owned", **json.loads(GEMINI_BODY)}))
assert frozenset(split.litellm_params) == frozenset((name, "metadata", "proxy_server_request"))
assert split.litellm_params[name] == "owned"
assert split.forwarded_body == json.loads(GEMINI_BODY)
@pytest.mark.parametrize("name", sorted(PROXY_STAMPED_NAMES))
def test_passthrough_drops_a_client_supplied_proxy_stamped_name(name: str) -> None:
split: Final = _split_pass_through_body(json.dumps({name: {"forged": "by-client"}, **json.loads(GEMINI_BODY)}))
assert frozenset(split.litellm_params) == frozenset(("metadata", "proxy_server_request"))
assert split.litellm_params["proxy_server_request"] != {"forged": "by-client"}, split.litellm_params
assert split.forwarded_body == json.loads(GEMINI_BODY)
def test_passthrough_merges_both_metadata_carriers_from_the_body_into_one_metadata_key() -> None:
split: Final = _split_pass_through_body(
'{"metadata": {"client_tag": "a"}, "contents": [{"parts": [{"text": "hi"}]}], "ttl": 30,'
' "litellm_metadata": {"lm": "b"}, "generationConfig": {"temperature": 0}}'
)
assert frozenset(split.litellm_params) == frozenset(("ttl", "metadata", "proxy_server_request"))
assert _metadata_of(split) == {**_metadata_of(_split_pass_through_body(GEMINI_BODY)), "client_tag": "a", "lm": "b"}
assert split.forwarded_body == json.loads(GEMINI_BODY)
def test_passthrough_lets_metadata_win_over_litellm_metadata_on_a_shared_key() -> None:
split: Final = _split_pass_through_body(
'{"litellm_metadata": {"shared": "from-litellm-metadata", "lm": "b"},'
' "metadata": {"shared": "from-metadata", "client_tag": "a"}, "contents": []}'
)
assert _metadata_of(split) == {
**_metadata_of(_split_pass_through_body('{"contents": []}')),
"shared": "from-metadata",
"lm": "b",
"client_tag": "a",
}
def test_passthrough_orders_extracted_litellm_params_by_the_registry() -> None:
body: Final = json.dumps({"ttl": 30, "tags": ["team-a"], "num_retries": 2, "contents": []})
split: Final = _split_pass_through_body(body)
body_keys: Final = frozenset(json.loads(body))
assert tuple(k for k in split.litellm_params if k in body_keys) == tuple(
k for k in types_utils.all_litellm_params if k in body_keys
)
LATE_REGISTERED_BODY: Final = '{"registered_later": 1, "contents": [{"parts": [{"text": "hi"}]}]}'
def test_passthrough_sees_a_name_appended_to_the_public_list_after_import() -> None:
litellm.all_litellm_params.append("registered_later")
try:
split: Final = _split_pass_through_body(LATE_REGISTERED_BODY)
finally:
litellm.all_litellm_params.remove("registered_later")
assert frozenset(split.litellm_params) == frozenset(("registered_later", "metadata", "proxy_server_request"))
assert split.forwarded_body == {"contents": [{"parts": [{"text": "hi"}]}]}
def test_passthrough_sees_the_public_list_rebound_after_import(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setattr(types_utils, "all_litellm_params", (*litellm.all_litellm_params, "registered_later"))
split: Final = _split_pass_through_body(LATE_REGISTERED_BODY)
assert frozenset(split.litellm_params) == frozenset(("registered_later", "metadata", "proxy_server_request"))
assert split.forwarded_body == {"contents": [{"parts": [{"text": "hi"}]}]}
@pytest.mark.asyncio
async def test_chat_completion_pass_through_endpoint_answers_an_openai_typed_error_for_an_unknown_model(
monkeypatch: pytest.MonkeyPatch,

View file

@ -309,6 +309,243 @@ def test_realtime_logging_object_does_not_validate_unknown_event_types():
assert len(dumped["results"]) == len(results)
def test_realtime_transcription_honors_deployment_pricing_override(monkeypatch: pytest.MonkeyPatch) -> None:
"""A deployment's pricing override must reach transcription events too.
Transcription is billed separately from response usage inside the same realtime
session, so a deployment registered at zero rates has to zero both. Resolving
transcription against the public ASR model instead billed a zero-rated
deployment for every .completed event.
"""
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
deployment_id = "deployment-hash-zero-rated-asr"
litellm.register_model(
model_cost={
deployment_id: {
"litellm_provider": "openai",
"mode": "realtime",
"input_cost_per_second": 0.0,
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
"input_cost_per_audio_token": 0.0,
}
}
)
results: OpenAIRealtimeStreamList = [
{
"type": "session.created",
"session": {
"type": "transcription",
"audio": {"input": {"transcription": {"model": "gpt-realtime-whisper"}}},
},
},
{
"type": "conversation.item.input_audio_transcription.completed",
"usage": {"type": "duration", "seconds": 120.0},
},
]
public_rate_cost = 120.0 * litellm.model_cost["gpt-realtime-whisper"]["input_cost_per_second"]
assert public_rate_cost > 0, "the public ASR rate must be non-zero for this test to mean anything"
without_override = handle_realtime_stream_cost_calculation(
results=results,
combined_usage_object=Usage(),
custom_llm_provider="openai",
litellm_model_name="gpt-realtime-whisper",
)
assert abs(without_override - public_rate_cost) < 1e-9
with_override = handle_realtime_stream_cost_calculation(
results=results,
combined_usage_object=Usage(),
custom_llm_provider="openai",
litellm_model_name="gpt-realtime-whisper",
custom_pricing_model=deployment_id,
)
assert with_override == 0.0, "the zero-rated deployment must not be billed for transcription"
def test_realtime_transcription_partial_override_keeps_unset_rates(monkeypatch: pytest.MonkeyPatch) -> None:
"""An override must not blank the rates it does not set.
A deployment that prices tokens but omits input_cost_per_second would otherwise
bill duration-based transcription at nothing, because the cost helpers read
`.get(key) or 0.0`. Only the fields the operator actually set may win.
"""
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
deployment_id = "deployment-hash-tokens-only"
litellm.register_model(
model_cost={
deployment_id: {
"litellm_provider": "openai",
"mode": "realtime",
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
}
}
)
results: OpenAIRealtimeStreamList = [
{
"type": "session.created",
"session": {
"type": "transcription",
"audio": {"input": {"transcription": {"model": "gpt-realtime-whisper"}}},
},
},
{
"type": "conversation.item.input_audio_transcription.completed",
"usage": {"type": "duration", "seconds": 120.0},
},
]
cost = handle_realtime_stream_cost_calculation(
results=results,
combined_usage_object=Usage(),
custom_llm_provider="openai",
litellm_model_name="gpt-realtime-whisper",
custom_pricing_model=deployment_id,
)
expected = 120.0 * litellm.model_cost["gpt-realtime-whisper"]["input_cost_per_second"]
assert expected > 0, "the public ASR per-second rate must be non-zero for this test to mean anything"
assert cost == pytest.approx(expected, rel=1e-9), (
"duration must keep the ASR per-second rate the override left unset"
)
@pytest.mark.parametrize(
"label,override,expected_audio_rate,expected_per_second",
[
("tokens only", {"input_cost_per_token": 0.0}, 0.0, 0.017 / 60),
("audio zeroed", {"input_cost_per_audio_token": 0.0}, 0.0, 0.017 / 60),
("per second only", {"input_cost_per_second": 0.001}, 6e-06, 0.001),
("empty override", {}, 6e-06, 0.017 / 60),
("no override", None, 6e-06, 0.017 / 60),
],
)
def test_transcription_rate_precedence(
monkeypatch: pytest.MonkeyPatch,
label: str,
override: dict[str, float] | None,
expected_audio_rate: float,
expected_per_second: float,
) -> None:
"""Rates resolve within one entry before moving to the next, and zero is a real value.
An override that prices only tokens must apply its own token rate to audio rather
than reaching past itself for the public audio rate, a deliberate zero must win
instead of being treated as unset, and a rate the override never mentions must keep
the base entry's value.
"""
from litellm.cost_calculator import handle_realtime_transcription_cost_calculation
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
base_model = "asr-precedence-base"
deployment_id = "asr-precedence-deployment"
litellm.register_model(
model_cost={
base_model: {
"litellm_provider": "openai",
"mode": "audio_transcription",
"input_cost_per_audio_token": 6e-06,
"input_cost_per_token": 2.5e-06,
"input_cost_per_second": 0.017 / 60,
}
}
)
if override is not None:
litellm.register_model(
model_cost={deployment_id: {"litellm_provider": "openai", "mode": "audio_transcription", **override}}
)
def cost_for(usage: dict[str, object]) -> float:
return handle_realtime_transcription_cost_calculation(
results=[
{"type": "transcription_session.created", "session": {"model": base_model}},
{"type": "conversation.item.input_audio_transcription.completed", "usage": usage},
],
custom_llm_provider="openai",
litellm_model_name=base_model,
custom_pricing_model=deployment_id if override is not None else None,
)
audio_cost = cost_for({"type": "tokens", "input_token_details": {"audio_tokens": 100}})
assert audio_cost == pytest.approx(100 * expected_audio_rate, rel=1e-9), f"{label}: audio rate"
per_second_cost = cost_for({"type": "duration", "seconds": 120.0})
assert per_second_cost == pytest.approx(120.0 * expected_per_second, rel=1e-9), (
f"{label}: an override must never blank a rate it does not set"
)
def test_realtime_transcription_per_second_override_keeps_public_token_rates(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""A per-second override must not zero the token rates ``get_model_info`` synthesizes.
``get_model_info`` defaults input_cost_per_token and output_cost_per_token to 0 for entries
that omit them, so a deployment priced only per second looked like it had declared token
rates of 0. Token-shaped transcription then billed nothing instead of falling through to the
public ASR rates, while the per-second rate the operator did set stayed in force.
"""
from litellm.cost_calculator import handle_realtime_transcription_cost_calculation
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
asr_model = "gpt-4o-transcribe"
per_second_rate = 0.001
deployment_id = "deployment-hash-per-second-only"
litellm.register_model(
model_cost={
deployment_id: {
"litellm_provider": "openai",
"mode": "audio_transcription",
"input_cost_per_second": per_second_rate,
}
}
)
public = litellm.model_cost[asr_model]
session_event = {"type": "transcription_session.created", "session": {"model": asr_model}}
def cost_for(usage: dict[str, object]) -> float:
return handle_realtime_transcription_cost_calculation(
results=[session_event, {"type": "conversation.item.input_audio_transcription.completed", "usage": usage}],
custom_llm_provider="openai",
litellm_model_name=asr_model,
custom_pricing_model=deployment_id,
)
token_cost = cost_for(
{
"type": "tokens",
"input_token_details": {"audio_tokens": 400, "text_tokens": 12},
"output_tokens": 30,
}
)
expected_token_cost = (
400 * public["input_cost_per_audio_token"]
+ 12 * public["input_cost_per_token"]
+ 30 * public["output_cost_per_token"]
)
assert expected_token_cost > 0, "the public ASR token rates must be non-zero for this test to mean anything"
assert token_cost == pytest.approx(expected_token_cost, rel=1e-9), (
"an override that prices only seconds must leave the public token rates in place"
)
assert cost_for({"type": "duration", "seconds": 120.0}) == pytest.approx(120.0 * per_second_rate, rel=1e-9)
def test_realtime_transcription_no_completed_events_is_zero(monkeypatch):
"""A realtime stream without transcription completed events adds no extra cost."""
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
@ -4635,6 +4872,391 @@ def test_gemini_live_native_audio_limits_and_capabilities_match_vendor_model_car
assert info["supports_pdf_input"] is False
def test_realtime_honours_deployment_custom_pricing(monkeypatch: pytest.MonkeyPatch) -> None:
"""Regression: a deployment's pricing override never reached realtime costing.
`model_info` overrides are registered under the deployment's own model_id, and
only `_select_model_name_for_cost_calc` knows to look there. The realtime branch
discarded that result and priced by the model the session reported, so a config
that zeroes a realtime deployment was billed at the public rate anyway. Audio is
the bulk of a voice call, so the gap was most of the cost.
"""
from litellm.types.utils import CompletionTokensDetailsWrapper
model = "gemini-3.1-flash-live-preview"
deployment_key = "deployment-id-for-a-zero-rated-realtime-group"
paid = litellm.model_cost[model]
monkeypatch.setitem(
litellm.model_cost,
deployment_key,
{
**paid,
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
"input_cost_per_audio_token": 0.0,
"output_cost_per_audio_token": 0.0,
"cache_read_input_token_cost": 0.0,
},
)
results: OpenAIRealtimeStreamList = [
{"type": "session.created", "session": {"model": model}},
{
"type": "response.done",
"response": {"usage": {"input_tokens": 10, "output_tokens": 200, "total_tokens": 210}},
},
]
usage = Usage(
prompt_tokens=10,
completion_tokens=200,
total_tokens=210,
prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=10, cached_tokens=0),
completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=20, audio_tokens=180),
)
paid_cost = handle_realtime_stream_cost_calculation(
results=results,
combined_usage_object=usage,
custom_llm_provider="gemini",
litellm_model_name=model,
)
expected_paid = (
10 * paid["input_cost_per_token"]
+ 20 * paid["output_cost_per_token"]
+ 180 * paid["output_cost_per_audio_token"]
)
assert paid_cost == pytest.approx(expected_paid, rel=1e-9)
assert paid_cost > 0
zero_rated_cost = handle_realtime_stream_cost_calculation(
results=results,
combined_usage_object=usage,
custom_llm_provider="gemini",
litellm_model_name=model,
custom_pricing_model=deployment_key,
)
assert zero_rated_cost == 0.0
def test_realtime_honours_a_provider_prefixed_zero_rated_deployment(monkeypatch: pytest.MonkeyPatch) -> None:
"""Regression: the override arrived provider-prefixed and was read as pricing nothing.
`_select_model_name_for_cost_calc` hands back `<provider>/<model_id>`, so the name reaching
the pricing guard carries a prefix the raw cost-map lookups cannot strip. The rates resolved
correctly through `get_model_info`, then the guard rejected them as undeclared and the session
billed the public rates. A zero-rated deployment must stay at zero however its name arrives.
"""
from litellm.types.utils import CompletionTokensDetailsWrapper
model = "gemini-live-2.5-flash-native-audio"
deployment_key = "deployment-id-for-a-prefixed-zero-rated-realtime-group"
paid = litellm.model_cost[model]
monkeypatch.setitem(
litellm.model_cost,
deployment_key,
{
**paid,
"input_cost_per_token": 0.0,
"output_cost_per_token": 0.0,
"input_cost_per_audio_token": 0.0,
"output_cost_per_audio_token": 0.0,
},
)
results: OpenAIRealtimeStreamList = [
{"type": "session.created", "session": {"model": model}},
{
"type": "response.done",
"response": {"usage": {"input_tokens": 219, "output_tokens": 81, "total_tokens": 300}},
},
]
usage = Usage(
prompt_tokens=219,
completion_tokens=81,
total_tokens=300,
prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=16, audio_tokens=203),
completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=23, audio_tokens=58),
)
paid_cost = handle_realtime_stream_cost_calculation(
results=results,
combined_usage_object=usage,
custom_llm_provider="vertex_ai",
litellm_model_name=model,
)
assert paid_cost > 0
zero_rated_cost = handle_realtime_stream_cost_calculation(
results=results,
combined_usage_object=usage,
custom_llm_provider="vertex_ai",
litellm_model_name=model,
custom_pricing_model=f"vertex_ai/{deployment_key}",
)
assert zero_rated_cost == 0.0
def test_unpriced_deployment_entry_still_falls_through_to_the_session_model(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""The guard's own purpose must survive: an entry that prices nothing is not an override.
Deployments are auto-registered under their model_id with no rates at all, and those must
keep billing at the session model's public rates rather than silently costing nothing.
"""
from litellm.types.utils import CompletionTokensDetailsWrapper
model = "gemini-live-2.5-flash-native-audio"
deployment_key = "deployment-id-with-no-declared-rates"
monkeypatch.setitem(
litellm.model_cost,
deployment_key,
{key: value for key, value in litellm.model_cost[model].items() if "cost_per" not in key},
)
results: OpenAIRealtimeStreamList = [
{"type": "session.created", "session": {"model": model}},
{
"type": "response.done",
"response": {"usage": {"input_tokens": 219, "output_tokens": 81, "total_tokens": 300}},
},
]
usage = Usage(
prompt_tokens=219,
completion_tokens=81,
total_tokens=300,
prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=16, audio_tokens=203),
completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=23, audio_tokens=58),
)
with_unpriced_override = handle_realtime_stream_cost_calculation(
results=results,
combined_usage_object=usage,
custom_llm_provider="vertex_ai",
litellm_model_name=model,
custom_pricing_model=f"vertex_ai/{deployment_key}",
)
without_override = handle_realtime_stream_cost_calculation(
results=results,
combined_usage_object=usage,
custom_llm_provider="vertex_ai",
litellm_model_name=model,
)
assert with_unpriced_override == pytest.approx(without_override, rel=1e-9)
assert with_unpriced_override > 0
def test_realtime_audio_only_override_bills_audio_at_the_deployment_rate(
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""Regression: an audio-only pricing override was never selected as the pricing key.
The deployment-selection guard recognised only text, per-second, per-query and
tiered rates, so a deployment that priced just the audio meters was passed over
and the session kept billing the public rates for the exact tokens it priced.
"""
from litellm.types.utils import CompletionTokensDetailsWrapper
model = "gemini-live-2.5-flash-native-audio"
deployment_key = "deployment-id-for-an-audio-only-realtime-group"
monkeypatch.setitem(
litellm.model_cost,
deployment_key,
{
"litellm_provider": "vertex_ai",
"mode": "realtime",
"input_cost_per_audio_token": 0.0,
"output_cost_per_audio_token": 0.0,
},
)
logging_object = LiteLLMRealtimeStreamLoggingObject(
usage=Usage(
prompt_tokens=203,
completion_tokens=58,
total_tokens=261,
prompt_tokens_details=PromptTokensDetailsWrapper(audio_tokens=203),
completion_tokens_details=CompletionTokensDetailsWrapper(audio_tokens=58),
),
results=[
{"type": "session.created", "session": {"model": model}},
{
"type": "response.done",
"response": {"usage": {"input_tokens": 203, "output_tokens": 58, "total_tokens": 261}},
},
],
)
public_cost = completion_cost(
completion_response=logging_object,
model=model,
call_type=CallTypes.arealtime.value,
custom_llm_provider="vertex_ai",
)
assert public_cost > 0
overridden_cost = completion_cost(
completion_response=logging_object,
model=model,
call_type=CallTypes.arealtime.value,
custom_llm_provider="vertex_ai",
custom_pricing=True,
router_model_id=deployment_key,
)
assert overridden_cost == pytest.approx(0.0)
def test_realtime_session_falls_back_to_base_model_pricing(monkeypatch: pytest.MonkeyPatch) -> None:
"""Regression: a priced base_model was discarded for realtime sessions.
The resolved base model only reached the realtime cost path when custom pricing
was on, so a session reporting an alias unmapped in the cost map recorded zero
instead of the base model's published price.
"""
from litellm.types.utils import CompletionTokensDetailsWrapper
base_model = "gemini-live-2.5-flash-native-audio"
logging_object = LiteLLMRealtimeStreamLoggingObject(
usage=Usage(
prompt_tokens=219,
completion_tokens=81,
total_tokens=300,
prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=16, audio_tokens=203),
completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=23, audio_tokens=58),
),
results=[
{
"type": "session.created",
"session": {"model": "my-voice-alias"},
},
{
"type": "response.done",
"response": {"usage": {"input_tokens": 219, "output_tokens": 81, "total_tokens": 300}},
},
],
)
aliased_cost = completion_cost(
completion_response=logging_object,
model="my-voice-alias",
call_type=CallTypes.arealtime.value,
custom_llm_provider="vertex_ai",
base_model=base_model,
)
base_cost = completion_cost(
completion_response=logging_object,
model=base_model,
call_type=CallTypes.arealtime.value,
custom_llm_provider="vertex_ai",
)
assert aliased_cost == pytest.approx(base_cost, rel=1e-9)
assert aliased_cost > 0
def test_base_model_does_not_override_transcription_rates(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
base_model = "gpt-realtime-2"
asr_model = "gpt-4o-transcribe"
logging_object = LiteLLMRealtimeStreamLoggingObject(
usage=Usage(),
results=[
{
"type": "session.created",
"session": {
"model": "my-voice-alias",
"audio": {"input": {"transcription": {"model": asr_model}}},
},
},
{
"type": "conversation.item.input_audio_transcription.completed",
"usage": {
"type": "tokens",
"input_token_details": {"audio_tokens": 400, "text_tokens": 12},
"output_tokens": 30,
},
},
],
)
with_base_model = completion_cost(
completion_response=logging_object,
model="my-voice-alias",
call_type=CallTypes.arealtime.value,
custom_llm_provider="openai",
base_model=base_model,
)
asr_priced = completion_cost(
completion_response=logging_object,
model="my-voice-alias",
call_type=CallTypes.arealtime.value,
custom_llm_provider="openai",
)
realtime_card = litellm.model_cost[base_model]
billed_at_realtime = (
400 * realtime_card["input_cost_per_audio_token"]
+ 12 * realtime_card["input_cost_per_token"]
+ 30 * realtime_card["output_cost_per_audio_token"]
)
assert billed_at_realtime != pytest.approx(asr_priced, rel=1e-9)
assert with_base_model == pytest.approx(asr_priced, rel=1e-9)
assert with_base_model > 0
def test_realtime_base_model_outranks_the_session_reported_model(monkeypatch: pytest.MonkeyPatch) -> None:
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url=""))
from litellm.types.utils import CompletionTokensDetailsWrapper
session_model = "gpt-realtime-mini"
base_model = "gpt-realtime-2"
def logging_object_for(session: str) -> LiteLLMRealtimeStreamLoggingObject:
return LiteLLMRealtimeStreamLoggingObject(
usage=Usage(
prompt_tokens=120,
completion_tokens=60,
total_tokens=180,
prompt_tokens_details=PromptTokensDetailsWrapper(text_tokens=20, audio_tokens=100),
completion_tokens_details=CompletionTokensDetailsWrapper(text_tokens=10, audio_tokens=50),
),
results=[
{
"type": "session.created",
"session": {"model": session},
},
{
"type": "response.done",
"response": {"usage": {"input_tokens": 120, "output_tokens": 60, "total_tokens": 180}},
},
],
)
with_base_model = completion_cost(
completion_response=logging_object_for(session_model),
model=session_model,
call_type=CallTypes.arealtime.value,
custom_llm_provider="openai",
base_model=base_model,
)
base_priced = completion_cost(
completion_response=logging_object_for(base_model),
model=base_model,
call_type=CallTypes.arealtime.value,
custom_llm_provider="openai",
)
session_priced = completion_cost(
completion_response=logging_object_for(session_model),
model=session_model,
call_type=CallTypes.arealtime.value,
custom_llm_provider="openai",
)
assert base_priced != pytest.approx(session_priced, rel=1e-9)
assert with_base_model == pytest.approx(base_priced, rel=1e-9)
def test_baseten_glm_5_3_fast_is_priced_from_registry(_local_model_cost_map: None) -> None:
model: Final = "baseten/zai-org/GLM-5.3-Fast"
prompt_tokens: Final = 1000

View file

@ -25,6 +25,7 @@ from litellm import Router
from litellm.caching.caching import DualCache
from litellm.caching.redis_cache import _redis_circuit_breaker_guard
from litellm.exceptions import MidStreamFallbackError
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
@ -49,6 +50,7 @@ from litellm.router import (
from litellm.router_strategy import simple_shuffle
from litellm.router_utils.client_initalization_utils import MaxParallelRequestsLimit
from litellm.router_utils.cooldown_handlers import _async_get_cooldown_deployments
from litellm.router_utils.fallback_event_handlers import DISABLE_FALLBACKS_METADATA_KEY
from litellm.router_utils.router_callbacks.track_deployment_metrics import get_deployment_successes_for_current_minute
from litellm.types.llms.openai import ChatCompletionRequest
from litellm.types.router import Deployment, DeploymentTypedDict, LiteLLM_Params, ModelInfo, PreRoutingHookResponse, RetryPolicy
@ -14392,6 +14394,193 @@ async def test_anthropic_messages_fallback_also_catches_raised_midstream_error()
assert mock_fallback.await_args.kwargs["e"] is raised_error
_MID_STREAM_OPT_OUT_SHAPES: Final = (
pytest.param({"disable_fallbacks": True}, id="raw-kwarg"),
pytest.param({"metadata": {DISABLE_FALLBACKS_METADATA_KEY: True}}, id="metadata-stamp"),
pytest.param({"litellm_metadata": {DISABLE_FALLBACKS_METADATA_KEY: True}}, id="litellm_metadata-stamp"),
)
def _mid_stream_opt_out_router() -> Router:
return Router(
model_list=[
{"model_name": "primary", "litellm_params": {"model": "openai/gpt-5.4", "api_key": "k1"}},
{"model_name": "fallback", "litellm_params": {"model": "openai/gpt-5.4-mini", "api_key": "k2"}},
],
fallbacks=[{"primary": ["fallback"]}],
)
def _mid_stream_opt_out_primary_error() -> litellm.InternalServerError:
return litellm.InternalServerError(message="primary failed at stream start", llm_provider="openai", model="primary")
def _mid_stream_opt_out_trigger(primary_error: Exception) -> MidStreamFallbackError:
return MidStreamFallbackError(
message=str(primary_error),
model="primary",
llm_provider="openai",
original_exception=primary_error,
is_pre_first_chunk=True,
)
class _MidStreamOptOutChatStream(CustomStreamWrapper):
"""A chat deployment stream, as the router sees one, that dies before its first chunk."""
def __init__(self, error: Exception, model: str = "primary") -> None:
super().__init__(completion_stream=object(), model=model, custom_llm_provider="openai", logging_obj=MagicMock())
self._error: Final = error
def __aiter__(self):
return self
async def __anext__(self) -> object:
raise self._error
def __iter__(self):
return self
def __next__(self) -> object:
raise self._error
@pytest.mark.asyncio
@pytest.mark.parametrize("opt_out", _MID_STREAM_OPT_OUT_SHAPES)
async def test_acompletion_streaming_iterator_honors_disable_fallbacks(opt_out):
"""A chat stream that fails before its first chunk on a request that opted out of fallbacks
surfaces the primary's own error and never tries the fallback deployment."""
router = _mid_stream_opt_out_router()
primary_error = _mid_stream_opt_out_primary_error()
source = _MidStreamOptOutChatStream(_mid_stream_opt_out_trigger(primary_error))
with patch.object(router, "_acompletion", new=AsyncMock(return_value=_AsyncList([]))) as fallback_attempt:
wrapped = await router._acompletion_streaming_iterator(
model_response=source,
messages=[{"role": "user", "content": "Hi"}],
initial_kwargs={"model": "primary", "stream": True, **copy.deepcopy(opt_out)},
)
with pytest.raises(litellm.InternalServerError) as raised:
[chunk async for chunk in wrapped]
assert raised.value is primary_error
fallback_attempt.assert_not_awaited()
@pytest.mark.parametrize("opt_out", _MID_STREAM_OPT_OUT_SHAPES)
def test_completion_streaming_iterator_honors_disable_fallbacks(opt_out):
"""Sync counterpart of test_acompletion_streaming_iterator_honors_disable_fallbacks."""
router = _mid_stream_opt_out_router()
primary_error = _mid_stream_opt_out_primary_error()
source = _MidStreamOptOutChatStream(_mid_stream_opt_out_trigger(primary_error))
with patch.object(router, "_completion", new=MagicMock(return_value=iter([]))) as fallback_attempt:
wrapped = router._completion_streaming_iterator(
model_response=source,
messages=[{"role": "user", "content": "Hi"}],
initial_kwargs={"model": "primary", "stream": True, **copy.deepcopy(opt_out)},
)
with pytest.raises(litellm.InternalServerError) as raised:
list(wrapped)
assert raised.value is primary_error
fallback_attempt.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.parametrize("opt_out", _MID_STREAM_OPT_OUT_SHAPES)
async def test_aresponses_streaming_iterator_honors_disable_fallbacks(opt_out):
"""Same opt-out contract on the Responses API mid-stream fallback path."""
router = _mid_stream_opt_out_router()
primary_error = _mid_stream_opt_out_primary_error()
source = _make_responses_iterator(error=_mid_stream_opt_out_trigger(primary_error), model="primary")
with patch.object(
router,
"_ageneric_api_call_with_fallbacks_responses_attempt",
new=AsyncMock(return_value=_AsyncList([])),
) as fallback_attempt:
wrapped = await router._aresponses_streaming_iterator(
response=source,
initial_kwargs={
"model": "primary",
"stream": True,
"input": "Hi",
"original_generic_function": litellm.aresponses,
**copy.deepcopy(opt_out),
},
)
with pytest.raises(litellm.InternalServerError) as raised:
[chunk async for chunk in wrapped]
assert raised.value is primary_error
fallback_attempt.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize("opt_out", _MID_STREAM_OPT_OUT_SHAPES)
async def test_anthropic_messages_streaming_iterator_honors_disable_fallbacks(opt_out):
"""Same opt-out contract on the Anthropic Messages mid-stream fallback path."""
router = _mid_stream_opt_out_router()
primary_error = _mid_stream_opt_out_primary_error()
source = _AnthropicMessagesRaisingByteStream([], _mid_stream_opt_out_trigger(primary_error))
with patch.object(
router,
"_ageneric_api_call_with_fallbacks_anthropic_messages_attempt",
new=AsyncMock(return_value=_AnthropicMessagesFallbackByteStream([])),
) as fallback_attempt:
wrapped = await router._aanthropic_messages_streaming_iterator(
response=source,
initial_kwargs={"model": "primary", "stream": True, **copy.deepcopy(opt_out)},
)
with pytest.raises(litellm.InternalServerError) as raised:
[chunk async for chunk in wrapped]
assert raised.value is primary_error
fallback_attempt.assert_not_awaited()
@pytest.mark.asyncio
async def test_acompletion_disable_fallbacks_reaches_the_mid_stream_hop():
"""`disable_fallbacks=True` sent to the public entrypoint survives the fallback wrapper's handoff
into the stream: the primary's own error surfaces and no fallback deployment is ever called."""
router = _mid_stream_opt_out_router()
primary_error = _mid_stream_opt_out_primary_error()
async def primary_stream(**kwargs):
return _MidStreamOptOutChatStream(_mid_stream_opt_out_trigger(primary_error), model=kwargs["model"])
with patch("litellm.acompletion", side_effect=primary_stream) as provider_calls:
response = await router.acompletion(
model="primary", messages=[{"role": "user", "content": "Hi"}], stream=True, disable_fallbacks=True
)
with pytest.raises(litellm.InternalServerError) as raised:
[chunk async for chunk in response]
assert raised.value is primary_error
assert [call.kwargs["metadata"]["model_group"] for call in provider_calls.call_args_list] == ["primary"]
def test_completion_disable_fallbacks_reaches_the_mid_stream_hop():
"""Sync counterpart of test_acompletion_disable_fallbacks_reaches_the_mid_stream_hop."""
router = _mid_stream_opt_out_router()
primary_error = _mid_stream_opt_out_primary_error()
def primary_stream(**kwargs):
return _MidStreamOptOutChatStream(_mid_stream_opt_out_trigger(primary_error), model=kwargs["model"])
with patch("litellm.completion", side_effect=primary_stream) as provider_calls:
response = router.completion(
model="primary", messages=[{"role": "user", "content": "Hi"}], stream=True, disable_fallbacks=True
)
with pytest.raises(litellm.InternalServerError) as raised:
list(response)
assert raised.value is primary_error
assert [call.kwargs["metadata"]["model_group"] for call in provider_calls.call_args_list] == ["primary"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"raised_error",

View file

@ -3730,6 +3730,23 @@ def test_scoped_weights_are_excluded_from_provider_params(filter_name: str) -> N
assert filtered == {"provider_option": "kept"}
@pytest.mark.parametrize(
"provider_filter",
[
litellm.utils.get_non_default_completion_params,
litellm.utils.get_non_default_transcription_params,
litellm.utils.filter_out_litellm_params,
],
)
@pytest.mark.parametrize("setting", [("tag_regex", ["^team-a$"]), ("max_file_size_mb", 5)])
def test_deployment_only_settings_copied_by_the_router_stay_out_of_provider_params(
provider_filter: Callable[[dict[str, object]], Mapping[str, object]], setting: tuple[str, object]
) -> None:
name, value = setting
filtered: Final = provider_filter({"provider_option": "kept", name: value})
assert filtered == {"provider_option": "kept"}, filtered
class TestGetOptionalParamsTencent:
"""Tests that tencent provider uses TencentChatConfig for parameter mapping."""

View file

@ -11,6 +11,39 @@ from litellm.llms.vertex_ai.vertex_ai_partner_models.llama3.transformation impor
)
OPENAI_PLATFORM_PARAMS = (
"prompt_cache_key",
"prompt_cache_retention",
"safety_identifier",
"service_tier",
"store",
"web_search_options",
"modalities",
"prediction",
"audio",
)
SELF_DEPLOYED_ENDPOINT_MODELS = (
"gemma/gemma-2-2b-it",
"vertex_ai/gemma/gemma-2-2b-it",
"openai/mg-endpoint-lit8592",
"vertex_ai/openai/mg-endpoint-lit8592",
"openai/5464397967697903616",
)
MAAS_MODELS = (
"meta/llama-4-maverick-17b-128e-instruct-maas",
"vertex_ai/meta/llama-4-maverick-17b-128e-instruct-maas",
"moonshotai/kimi-k2-thinking-maas",
"qwen/qwen3-next-80b-a3b-instruct-maas",
"google/gemma-4-26b-a4b-it-maas",
"xai/grok-4.1-fast-non-reasoning",
"openai/xai/grok-4.1-fast-reasoning",
"1984786713414729728",
"llama3",
)
class TestVertexAILlama3Config:
def test_transform_choices(self):
"""
@ -56,6 +89,52 @@ class TestVertexAILlama3Config:
assert response[0].message.tool_calls is not None
assert response[0].finish_reason == "tool_calls"
@pytest.mark.parametrize("model", SELF_DEPLOYED_ENDPOINT_MODELS)
@pytest.mark.parametrize("param", OPENAI_PLATFORM_PARAMS)
def test_get_supported_openai_params_omits_platform_params_for_self_deployed_endpoints(
self, model: str, param: str
):
assert param not in VertexAILlama3Config().get_supported_openai_params(model=model)
@pytest.mark.parametrize("model", MAAS_MODELS)
@pytest.mark.parametrize("param", OPENAI_PLATFORM_PARAMS)
def test_get_supported_openai_params_keeps_platform_params_for_maas_models(self, model: str, param: str):
assert param in VertexAILlama3Config().get_supported_openai_params(model=model)
@pytest.mark.parametrize("model", [*SELF_DEPLOYED_ENDPOINT_MODELS, *MAAS_MODELS])
def test_get_supported_openai_params_never_lists_max_retries(self, model: str):
assert "max_retries" not in VertexAILlama3Config().get_supported_openai_params(model=model)
@pytest.mark.parametrize("model", [*SELF_DEPLOYED_ENDPOINT_MODELS, *MAAS_MODELS])
@pytest.mark.parametrize(
"param",
["max_completion_tokens", "tools", "tool_choice", "response_format", "seed", "logprobs", "parallel_tool_calls"],
)
def test_get_supported_openai_params_keeps_params_every_vertex_openai_endpoint_accepts(
self, model: str, param: str
):
assert param in VertexAILlama3Config().get_supported_openai_params(model=model)
@pytest.mark.parametrize("model", SELF_DEPLOYED_ENDPOINT_MODELS)
def test_map_openai_params_drops_prompt_cache_key_for_self_deployed_endpoints(self, model: str):
mapped = VertexAILlama3Config().map_openai_params(
{"prompt_cache_key": "session-lit8592", "max_completion_tokens": 10},
{},
model,
drop_params=True,
)
assert mapped == {"max_tokens": 10}
@pytest.mark.parametrize("model", MAAS_MODELS)
def test_map_openai_params_forwards_prompt_cache_key_for_maas_models(self, model: str):
mapped = VertexAILlama3Config().map_openai_params(
{"prompt_cache_key": "session-lit8592", "max_completion_tokens": 10},
{},
model,
drop_params=True,
)
assert mapped == {"prompt_cache_key": "session-lit8592", "max_tokens": 10}
class TestVertexAILlama3StreamingHandler:
def test_first_chunk_has_role_assistant_when_missing(self):

View file

@ -694,6 +694,89 @@ class TestVertexGemmaCompletion:
assert instance["@requestFormat"] == "chatCompletions"
assert "messages" in instance
@pytest.mark.parametrize(
"param",
[
"prompt_cache_key",
"prompt_cache_retention",
"safety_identifier",
"service_tier",
"store",
"web_search_options",
"modalities",
"prediction",
"audio",
"max_retries",
],
)
def test_get_supported_openai_params_omits_params_the_predict_endpoint_rejects(self, param: str):
from litellm.llms.vertex_ai.vertex_gemma_models.transformation import (
VertexGemmaConfig,
)
assert param not in VertexGemmaConfig().get_supported_openai_params(model="gemma-2-2b-it")
@pytest.mark.asyncio
async def test_acompletion_drops_prompt_cache_key_when_drop_params_is_set(self):
with (
patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client") as mock_get_client,
patch(
"litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token",
return_value=("fake-access-token", "PROJECT_ID"),
),
):
mock_client = Mock()
mock_response = Mock()
mock_response.status_code = 200
mock_response.json.return_value = _make_gemma_vertex_response()
mock_client.post = AsyncMock(return_value=mock_response)
mock_get_client.return_value = mock_client
await litellm.acompletion(
model="vertex_ai/gemma/gemma-2-2b-it",
messages=[{"role": "user", "content": "Test"}],
prompt_cache_key="session-lit8592",
service_tier="default",
max_completion_tokens=16,
drop_params=True,
api_base="https://test.us-central1-project.prediction.vertexai.goog/v1/projects/PROJECT_ID/locations/us-central1/endpoints/ENDPOINT_ID:predict",
vertex_project="PROJECT_ID",
vertex_location="us-central1",
)
instance = mock_client.post.call_args.kwargs["json"]["instances"][0]
assert "prompt_cache_key" not in instance
assert "service_tier" not in instance
assert instance["max_tokens"] == 16
assert instance["messages"] == [{"role": "user", "content": "Test"}]
@pytest.mark.asyncio
async def test_acompletion_rejects_prompt_cache_key_before_calling_vertex(self, monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr(litellm, "drop_params", False)
with (
patch("litellm.llms.custom_httpx.http_handler.get_async_httpx_client") as mock_get_client,
patch(
"litellm.llms.vertex_ai.vertex_gemma_models.main.VertexAIGemmaModels._ensure_access_token",
return_value=("fake-access-token", "PROJECT_ID"),
),
):
mock_client = Mock()
mock_client.post = AsyncMock()
mock_get_client.return_value = mock_client
with pytest.raises(litellm.UnsupportedParamsError, match="prompt_cache_key"):
await litellm.acompletion(
model="vertex_ai/gemma/gemma-2-2b-it",
messages=[{"role": "user", "content": "Test"}],
prompt_cache_key="session-lit8592",
drop_params=False,
api_base="https://test.us-central1-project.prediction.vertexai.goog/v1/projects/PROJECT_ID/locations/us-central1/endpoints/ENDPOINT_ID:predict",
vertex_project="PROJECT_ID",
vertex_location="us-central1",
)
mock_client.post.assert_not_called()
def test_transform_request_strips_context_management(self):
"""
Direct unit test for VertexGemmaConfig.transform_request: verify that

View file

@ -0,0 +1,655 @@
import inspect
from collections.abc import Callable, Mapping, Sequence
from dataclasses import dataclass, field, fields
from operator import attrgetter
from types import MappingProxyType
from typing import Final, TypeAlias, cast, get_type_hints
import httpx
import pytest
from aiohttp import ClientSession
from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI
from pydantic import BaseModel, ConfigDict, TypeAdapter, ValidationError
import litellm
from litellm.caching.caching import Cache
from litellm.litellm_core_utils.get_litellm_params import (
get_litellm_params, # pyright: ignore[reportUnknownVariableType] # untyped legacy carrier
)
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.router_strategy.complexity_router.context_compaction import CompactionState
from litellm.router_utils.fallback_event_handlers import AttemptedFallbackTargets
from litellm.types import litellm_params
from litellm.types import utils as types_utils
from litellm.types.caching import DynamicCacheControl
from litellm.types.litellm_params import (
ADDRESSED_RESPONSE_ID_FIELD,
LITELLM_OWNED_ROOTS,
TRUSTED_CALLBACK_VARS_FIELD,
CachingOptions,
owned_wire_names,
wire,
wire_names,
)
from litellm.types.llms.openai import ChatCompletionAssistantMessage, ChatCompletionUserMessage
from litellm.types.proxy.litellm_pre_call_utils import SecretFields
from litellm.types.router import (
ConfigurableClientsideParamsCustomAuth,
CredentialLiteLLMParams,
DeploymentTypedDict,
RetryPolicy,
RouterConfig,
UpdateRouterConfig,
)
from litellm.types.router_weights import RouterWeights
from litellm.types.utils import (
CustomPricingLiteLLMParams,
ModelResponse,
ModelResponseStream,
ProviderSpecificHeader,
StandardCallbackDynamicParams,
agentic_loop_internal_litellm_params,
all_litellm_params,
bedrock_batch_litellm_params,
)
from litellm.utils import (
filter_out_litellm_params, # pyright: ignore[reportUnknownVariableType] # untyped legacy classifier
get_non_default_completion_params, # pyright: ignore[reportUnknownVariableType] # untyped legacy classifier
get_non_default_transcription_params, # pyright: ignore[reportUnknownVariableType] # untyped legacy classifier
)
PROVIDER_KNOB: Final = "registry_test_provider_only_knob"
CONNECTION_NAMES: Final = (
"api_key",
"api_base",
"api_version",
"region_name",
"headers",
"provider_specific_header",
"client",
"shared_session",
"ssl_verify",
"request_timeout",
"force_timeout",
"stream_timeout",
"max_retries",
"tenant_id",
"client_id",
"client_secret",
"azure_username",
"azure_password",
"azure_scope",
"azure_ad_token_provider",
"litellm_credential_name",
"configurable_clientside_auth_params",
"use_xai_oauth",
"aws_batch_role_arn",
"s3_bucket_name",
"s3_region_name",
"s3_endpoint_url",
"s3_output_bucket_name",
"s3_bucket_owner",
"s3_access_key_id",
"s3_secret_access_key",
"s3_encryption_key_id",
"bedrock_tags",
)
OPTION_NAMES: Final = (
"custom_llm_provider",
"azure",
"use_litellm_proxy",
"use_chat_completions_api",
"use_in_pass_through",
"allowed_openai_params",
"fallbacks",
"context_window_fallback_dict",
"num_retries",
"retry_policy",
"retry_strategy",
"routing_strategy",
"cooldown_time",
"allowed_model_region",
"enable_tag_filtering",
"fastest_response",
"provider_affinity_header",
"search_tool_name",
"model_list",
"model_info",
"rpm",
"tpm",
"itpm",
"otpm",
"default_api_key_rpm_limit",
"default_api_key_tpm_limit",
"max_parallel_requests",
"weight",
"order",
"tag_regex",
"max_file_size_mb",
"auto_router_config_path",
"auto_router_config",
"auto_router_default_model",
"auto_router_embedding_model",
"auto_router_max_input_chars",
"auto_router_routing_compression",
"auto_router_model_compression",
"complexity_router_config",
"complexity_router_default_model",
"adaptive_router_config",
"adaptive_router_default_model",
"quality_router_config",
"quality_router_default_model",
"caching",
"cache",
"ttl",
"enable_prompt_caching",
"caching_groups",
"cost_per_query",
"base_model",
"max_budget",
"budget_duration",
"id",
"metadata",
"litellm_metadata",
"tags",
"litellm_trace_id",
"litellm_session_id",
"litellm_request_debug",
"logger_fn",
"verbose",
"no-log",
"max_agentic_loops",
"guardrails",
"prompt_id",
"prompt_variables",
"prompt_version",
"prompt_environment",
"prompt_label",
"litellm_system_prompt",
"custom_prompt_dict",
"roles",
"final_prompt_value",
"bos_token",
"eos_token",
"hf_model_name",
"supports_system_message",
"ensure_alternating_roles",
"user_continue_message",
"assistant_continue_message",
"disable_add_transform_inline_image_block",
"merge_reasoning_content_in_choices",
"enable_json_schema_validation",
"complete_response",
"stream_chunk_size",
"keepalive_seconds",
"allow_client_keepalive_override",
"mock_response",
"mock_timeout",
)
AGENTIC_LOOP_STATE_NAMES: Final = (
"_agentic_loop_depth",
"_agentic_loop_fingerprints",
"_agentic_loop_api_surface",
"_code_interpreter_interception_active",
"_code_interpreter_interception_sandbox_key",
"_code_interpreter_interception_session_scoped",
"_code_interpreter_interception_converted_stream",
"_websearch_interception_emit_native_blocks",
"_websearch_interception_converted_stream",
"_headroom_interception_converted_stream",
)
INTERNAL_STATE_NAMES: Final = (
"litellm_call_id",
"completion_call_id",
"model_alias_map",
"data_residency",
"litellm_logging_obj",
"preset_cache_key",
"cache_key",
"stream_response",
"_context_compaction_state",
*AGENTIC_LOOP_STATE_NAMES,
"_router_weights",
"fallback_depth",
"max_fallbacks",
"attempted_targets",
"proxy_server_request",
"secret_fields",
"litellm_trusted_callback_vars",
"_litellm_addressed_response_id",
"_litellm_strip_stream_usage",
"client_side_timeout",
"model_file_id_mapping",
"acompletion",
"aembedding",
"aimg_generation",
"atext_completion",
"text_completion",
"allm_passthrough_route",
"async_call",
)
BEDROCK_BATCH_NAMES: Final = (
"aws_batch_role_arn",
"s3_bucket_name",
"s3_region_name",
"s3_endpoint_url",
"s3_output_bucket_name",
"s3_bucket_owner",
"s3_access_key_id",
"s3_secret_access_key",
"s3_encryption_key_id",
"bedrock_tags",
)
ARTIFACT_NAMES: Final = ("self", "use_client", "model_config", "rust")
CALLBACK_VAR_NAMES: Final = tuple(StandardCallbackDynamicParams.__annotations__)
PRICING_NAMES: Final = tuple(CustomPricingLiteLLMParams.model_fields)
OWNED_NAMES: Final = (
*CONNECTION_NAMES,
*OPTION_NAMES,
*INTERNAL_STATE_NAMES,
*ARTIFACT_NAMES,
*CALLBACK_VAR_NAMES,
*PRICING_NAMES,
)
Classifier: TypeAlias = Callable[[dict[str, object]], dict[str, object]] # mutable-ok: classifiers use dict
CLASSIFIERS: Final[Mapping[str, Classifier]] = MappingProxyType(
{ # pyright: ignore[reportUnknownArgumentType] # untyped legacy classifiers
"completion": get_non_default_completion_params,
"transcription": get_non_default_transcription_params,
"filter_out": filter_out_litellm_params,
}
)
@pytest.mark.parametrize("classifier_name", CLASSIFIERS)
@pytest.mark.parametrize("name", OWNED_NAMES)
def test_owned_name_is_kept_out_of_provider_params(name: str, classifier_name: str) -> None:
provider_value: Final = object()
classify: Final = CLASSIFIERS[classifier_name]
result: Final = classify({name: object(), PROVIDER_KNOB: provider_value}) # mutable-ok: classifiers take a dict
assert result == MappingProxyType({PROVIDER_KNOB: provider_value})
assert result[PROVIDER_KNOB] is provider_value
def test_a_name_no_object_declares_reaches_the_provider() -> None:
result: Final = CLASSIFIERS["completion"]({PROVIDER_KNOB: 1}) # mutable-ok: classifier input type
assert result == MappingProxyType({PROVIDER_KNOB: 1})
def _cache_key_for_model_group(cache: Cache, model_group: str, options: CachingOptions) -> str:
return cache.get_cache_key( # pyright: ignore[reportUnknownMemberType] # untyped legacy key builder
model=model_group,
messages=(MappingProxyType({"role": "user", "content": "shared prompt"}),),
metadata=MappingProxyType({"caching_groups": options.caching_groups, "model_group": model_group}),
)
def test_caching_groups_is_a_flat_sequence_of_model_groups_that_share_one_cache_key(
monkeypatch: pytest.MonkeyPatch,
) -> None:
for callback_list in ("input_callback", "success_callback", "_async_success_callback"):
monkeypatch.setattr(litellm, callback_list, []) # mutable-ok: Cache() appends "cache" to these lists
options: Final = CachingOptions(caching_groups=(("gpt-4", "gpt-4o"), ("claude-3",)))
cache: Final = Cache()
keys: Final = tuple(_cache_key_for_model_group(cache, group, options) for group in ("gpt-4", "gpt-4o", "claude-3"))
assert (keys[0] == keys[1], keys[0] == keys[2]) == (True, False)
def test_all_litellm_params_is_exactly_the_owned_inventory() -> None:
assert frozenset(all_litellm_params) == frozenset(OWNED_NAMES)
assert frozenset(ARTIFACT_NAMES).isdisjoint(DECLARED_NAMES)
def test_every_owned_name_has_exactly_one_owner() -> None:
duplicated: Final = tuple(name for name in dict.fromkeys(all_litellm_params) if all_litellm_params.count(name) > 1)
assert duplicated == ()
@pytest.mark.parametrize(
("exported", "declared"),
(
pytest.param(
types_utils.TRUSTED_CALLBACK_VARS_FIELD,
litellm_params.TRUSTED_CALLBACK_VARS_FIELD,
id="TRUSTED_CALLBACK_VARS_FIELD",
),
pytest.param(
types_utils.ADDRESSED_RESPONSE_ID_FIELD,
litellm_params.ADDRESSED_RESPONSE_ID_FIELD,
id="ADDRESSED_RESPONSE_ID_FIELD",
),
),
)
def test_types_utils_still_exports_the_field_constant(exported: str, declared: str) -> None:
assert exported == declared
@dataclass(frozen=True, slots=True, kw_only=True)
class _Leaf:
plain: int | None = None
renamed: int | None = field(default=None, metadata=wire("wire-name"))
@dataclass(frozen=True, slots=True, kw_only=True)
class _OtherLeaf:
plain: int | None = None
trailing: int | None = None
@dataclass(frozen=True, slots=True, kw_only=True)
class _Root:
first: _Leaf
second: _OtherLeaf
@dataclass(frozen=True, slots=True, kw_only=True)
class _RootDeclaringAKwargDirectly:
first: _Leaf
stray: int | None = None
def test_wire_names_are_the_field_names_in_declaration_order_unless_wire_renames_them() -> None:
assert wire_names(_Leaf) == ("plain", "wire-name")
def test_owned_wire_names_walk_leaves_in_declaration_order_and_keep_every_occurrence() -> None:
assert owned_wire_names(_Root) == ("plain", "wire-name", "plain", "trailing")
def test_owned_wire_names_refuse_a_root_that_declares_a_kwarg_outside_a_leaf() -> None:
with pytest.raises(TypeError):
owned_wire_names(_RootDeclaringAKwargDirectly)
def test_agentic_loop_names_concatenate_as_a_list() -> None:
extended: Final = agentic_loop_internal_litellm_params + ["caller_added"] # mutable-ok: list contract under test
assert (type(extended), len(extended), frozenset(extended)) == (
list,
len(AGENTIC_LOOP_STATE_NAMES) + 2,
frozenset((*AGENTIC_LOOP_STATE_NAMES, "max_agentic_loops", "caller_added")),
)
def test_bedrock_batch_names_concatenate_as_a_tuple() -> None:
extended: Final = bedrock_batch_litellm_params + ("caller_added",)
assert extended == (*BEDROCK_BATCH_NAMES, "caller_added")
def test_proxy_stamped_fields_keep_their_wire_names() -> None:
assert (TRUSTED_CALLBACK_VARS_FIELD, ADDRESSED_RESPONSE_ID_FIELD) == (
"litellm_trusted_callback_vars",
"_litellm_addressed_response_id",
)
def test_all_litellm_params_concatenates_with_a_list_like_the_completion_entrypoint_does() -> None:
extended: Final = ["aembedding", "extra_headers"] + all_litellm_params # mutable-ok: list contract under test
assert (type(extended), frozenset(extended)) == (list, frozenset(("aembedding", "extra_headers", *OWNED_NAMES)))
CARRIED_AND_FORWARDED: Final = frozenset(("drop_params", "hugging_face", "no_log", "replicate", "together_ai"))
CARRIER_SIGNATURE: Final = inspect.signature(get_litellm_params) # pyright: ignore[reportUnknownArgumentType] # legacy
CARRIED_PARAMS: Final = tuple(
name for name in CARRIER_SIGNATURE.parameters if name != "kwargs" and name not in CARRIED_AND_FORWARDED
)
@pytest.mark.parametrize("name", CARRIED_PARAMS)
def test_every_param_get_litellm_params_carries_is_kept_out_of_provider_params(name: str) -> None:
provider_value: Final = object()
result: Final = CLASSIFIERS["completion"](
{name: object(), PROVIDER_KNOB: provider_value} # mutable-ok: classifier input type
)
assert result == MappingProxyType({PROVIDER_KNOB: provider_value})
TYPED_CONFIG_MODELS: Final[Mapping[str, tuple[type[BaseModel], ...]]] = MappingProxyType(
{
"credentials": (CredentialLiteLLMParams,),
"router": (RouterConfig, UpdateRouterConfig),
}
)
DECLARED_NAMES: Final = frozenset(name for root in LITELLM_OWNED_ROOTS for name in owned_wire_names(root))
ProviderClient: TypeAlias = (
OpenAI
| AsyncOpenAI
| AzureOpenAI
| AsyncAzureOpenAI
| HTTPHandler
| AsyncHTTPHandler
| httpx.Client
| httpx.AsyncClient
)
MockResponse: TypeAlias = str | Exception | Mapping[str, object] | Sequence[float] | ModelResponse | ModelResponseStream
TYPE_HINT_NAMESPACE: Final[Mapping[str, object]] = {
"ProviderClient": ProviderClient,
"ProviderSpecificHeader": ProviderSpecificHeader,
"ClientSession": ClientSession,
"AsyncAzureOpenAI": AsyncAzureOpenAI,
"AsyncOpenAI": AsyncOpenAI,
"AzureOpenAI": AzureOpenAI,
"OpenAI": OpenAI,
"AsyncHTTPHandler": AsyncHTTPHandler,
"HTTPHandler": HTTPHandler,
"ConfigurableClientsideParamsCustomAuth": ConfigurableClientsideParamsCustomAuth,
"RetryPolicy": RetryPolicy,
"DeploymentTypedDict": DeploymentTypedDict,
"DynamicCacheControl": DynamicCacheControl,
"ChatCompletionUserMessage": ChatCompletionUserMessage,
"ChatCompletionAssistantMessage": ChatCompletionAssistantMessage,
"MockResponse": MockResponse,
"ModelResponse": ModelResponse,
"ModelResponseStream": ModelResponseStream,
"Logging": Logging,
"SecretFields": SecretFields,
"CompactionState": CompactionState,
"RouterWeights": RouterWeights,
"AttemptedFallbackTargets": AttemptedFallbackTargets,
}
LEAF_SAMPLES: Final[Mapping[type, Mapping[str, object]]] = {
litellm_params.ProviderConnection: {"api_key": "k", "request_timeout": 1.5},
litellm_params.BedrockBatchConnection: {"aws_batch_role_arn": "arn", "bedrock_tags": ({"k": "v"},)},
litellm_params.DispatchOptions: {"custom_llm_provider": "openai"},
litellm_params.RoutingOptions: {
"fallbacks": [{"model": "gpt-4o", "api_key": "k", "temperature": 0}],
"num_retries": 2,
"retry_strategy": "constant_retry",
"routing_strategy": "simple-shuffle",
},
litellm_params.DeploymentOptions: {"model_info": {"region": "us"}, "rpm": 2},
litellm_params.SpecializedRouterOptions: {"adaptive_router_default_model": "gpt-4o"},
litellm_params.CachingOptions: {"ttl": 30.0, "caching_groups": (("gpt-4o", "gpt-4o-mini"),)},
litellm_params.CostOptions: {"max_budget": 10.0},
litellm_params.ObservabilityOptions: {"metadata": {"request": "test"}, "no_log": True},
litellm_params.AgenticLoopOptions: {"max_agentic_loops": 2},
litellm_params.GuardrailOptions: {"guardrails": ("default",)},
litellm_params.PromptOptions: {"prompt_id": "prompt", "prompt_variables": {"name": "value"}},
litellm_params.ResponseOptions: {"stream_chunk_size": 64},
litellm_params.MockOptions: {"mock_timeout": True},
litellm_params.CallState: {
"completion_call_id": "call",
"model_alias_map": {"alias": "gpt-4o"},
"data_residency": "us",
},
litellm_params.AgenticLoopState: {"api_surface": "chat_completions", "depth": 1},
litellm_params.RouterState: {"fallback_depth": 1},
litellm_params.ProxyRequestState: {
"proxy_server_request": {"path": "/chat/completions"},
"trusted_callback_vars": {"dd_api_key": "k"},
},
litellm_params.EntrypointState: {"acompletion": True},
}
LEAF_BAD_SAMPLES: Final[Mapping[type, Mapping[str, object]]] = {
litellm_params.ProviderConnection: {"api_key": 1},
litellm_params.BedrockBatchConnection: {"aws_batch_role_arn": 1},
litellm_params.DispatchOptions: {"custom_llm_provider": 1},
litellm_params.RoutingOptions: {"num_retries": "2"},
litellm_params.DeploymentOptions: {"rpm": "2"},
litellm_params.SpecializedRouterOptions: {"auto_router_max_input_chars": "2"},
litellm_params.CachingOptions: {"ttl": "30"},
litellm_params.CostOptions: {"max_budget": "10"},
litellm_params.ObservabilityOptions: {"verbose": "true"},
litellm_params.AgenticLoopOptions: {"max_agentic_loops": "2"},
litellm_params.GuardrailOptions: {"guardrails": (1,)},
litellm_params.PromptOptions: {"prompt_id": 1},
litellm_params.ResponseOptions: {"stream_chunk_size": "64"},
litellm_params.MockOptions: {"mock_timeout": "true"},
litellm_params.CallState: {"completion_call_id": 1},
litellm_params.AgenticLoopState: {"depth": "1"},
litellm_params.RouterState: {"fallback_depth": "1"},
litellm_params.ProxyRequestState: {"proxy_server_request": "request"},
litellm_params.EntrypointState: {"acompletion": "true"},
}
INVALID_LITERAL_SAMPLES: Final[tuple[tuple[type, Mapping[str, object]], ...]] = (
(litellm_params.RoutingOptions, {"retry_strategy": "linear"}),
(litellm_params.RoutingOptions, {"routing_strategy": "random"}),
(litellm_params.AgenticLoopState, {"api_surface": "batches"}),
)
def _leaf_id(value: object) -> str:
return value.__name__ if isinstance(value, type) else ""
def _leaf_instance(leaf: type, sample: Mapping[str, object]) -> object:
constructor: Final = cast(Callable[..., object], leaf)
return constructor(**sample)
def _strict_leaf_validation(leaf: type, instance: object) -> object:
hints: Final[Mapping[str, object]] = cast(
Mapping[str, object], get_type_hints(type(instance), localns=TYPE_HINT_NAMESPACE)
)
for field_info in fields(leaf):
value = cast(Callable[[object], object], attrgetter(field_info.name))(instance)
field_adapter: TypeAdapter[object] = TypeAdapter[object](
hints[field_info.name],
config=ConfigDict(arbitrary_types_allowed=True),
)
field_adapter.validate_python(value, strict=True)
return instance
@pytest.mark.parametrize("leaf,sample", LEAF_SAMPLES.items(), ids=_leaf_id)
def test_every_owned_leaf_accepts_a_strict_reader_shaped_sample(leaf: type, sample: Mapping[str, object]) -> None:
instance: Final = _leaf_instance(leaf, sample)
result: Final = _strict_leaf_validation(leaf, instance)
assert result == instance
assert frozenset(sample) <= frozenset(field.name for field in fields(leaf))
@pytest.mark.parametrize("leaf,sample", LEAF_BAD_SAMPLES.items(), ids=_leaf_id)
def test_every_owned_leaf_rejects_a_strict_wrong_typed_sample(leaf: type, sample: Mapping[str, object]) -> None:
instance: Final = _leaf_instance(leaf, sample)
with pytest.raises(ValidationError):
_strict_leaf_validation(leaf, instance)
@pytest.mark.parametrize("leaf,sample", INVALID_LITERAL_SAMPLES, ids=_leaf_id)
def test_owned_leaf_literals_reject_unknown_values(leaf: type, sample: Mapping[str, object]) -> None:
instance: Final = _leaf_instance(leaf, sample)
with pytest.raises(ValidationError):
_strict_leaf_validation(leaf, instance)
@pytest.mark.parametrize(
"strategy",
[
"simple-shuffle",
"least-busy",
"usage-based-routing",
"latency-based-routing",
"cost-based-routing",
"usage-based-routing-v2",
"lar1",
],
)
def test_routing_options_accept_every_strategy_the_router_accepts(strategy: str) -> None:
instance: Final = _leaf_instance(litellm_params.RoutingOptions, {"routing_strategy": strategy})
assert _strict_leaf_validation(litellm_params.RoutingOptions, instance) is instance
NAMES_SHARED_WITH_TYPED_MODELS: Final[Mapping[str, tuple[str, ...]]] = MappingProxyType(
{
"credentials": (
"api_base",
"api_key",
"api_version",
"aws_batch_role_arn",
"azure_password",
"azure_scope",
"azure_username",
"bedrock_tags",
"client_id",
"client_secret",
"region_name",
"s3_access_key_id",
"s3_bucket_name",
"s3_bucket_owner",
"s3_encryption_key_id",
"s3_endpoint_url",
"s3_output_bucket_name",
"s3_region_name",
"s3_secret_access_key",
"tenant_id",
),
"router": (
"caching_groups",
"cooldown_time",
"enable_tag_filtering",
"fallbacks",
"max_retries",
"model_list",
"num_retries",
"retry_policy",
"routing_strategy",
),
}
)
@pytest.mark.parametrize("source", TYPED_CONFIG_MODELS)
def test_names_a_typed_config_model_shares_with_the_owned_inventory_are_exactly_these(source: str) -> None:
model_names: Final = frozenset(name for model in TYPED_CONFIG_MODELS[source] for name in model.model_fields)
assert DECLARED_NAMES & model_names == frozenset(NAMES_SHARED_WITH_TYPED_MODELS[source])
@pytest.mark.parametrize("name", PRICING_NAMES)
def test_pricing_name_is_owned_by_the_pricing_model_alone(name: str) -> None:
assert name not in DECLARED_NAMES

View file

@ -134,9 +134,10 @@ describe("MCPToolPermissions", () => {
const selectAllButton = screen.getByRole("button", { name: "Select All" });
await userEvent.click(selectAllButton);
// Verify onChange was called with all tools selected
// Selecting every displayed tool writes the wildcard, which also covers tools the
// server adds later.
expect(mockOnChange).toHaveBeenCalledWith({
[mockServerId]: ["read_wiki_structure", "read_wiki_contents", "ask_question"],
[mockServerId]: ["*"],
});
});
@ -191,6 +192,77 @@ describe("MCPToolPermissions", () => {
});
});
describe("wildcard all-tools grant", () => {
const wildcardServerId = "server-1";
const wildcardServer = { server_id: wildcardServerId, server_name: "Wildcard Server", alias: "Wildcard Server" };
const wildcardTools = [
{ name: "read_wiki_structure", description: "Get documentation topics" },
{ name: "read_wiki_contents", description: "View documentation" },
{ name: "ask_question", description: "Ask questions" },
];
beforeEach(() => {
vi.mocked(networking.fetchMCPServers).mockResolvedValue([wildcardServer]);
vi.mocked(networking.listMCPTools).mockResolvedValue({ tools: wildcardTools, error: false });
});
it("renders every tool checked with the future-tools note when the entry is the wildcard", async () => {
renderWithProviders(
<MCPToolPermissions
accessToken={mockAccessToken}
selectedServers={[wildcardServerId]}
toolPermissions={{ [wildcardServerId]: ["*"] }}
onChange={vi.fn()}
/>,
);
expect(await screen.findByText("Wildcard Server")).toBeInTheDocument();
expect(screen.getByText("All tools allowed, including tools added to this server later")).toBeInTheDocument();
await userEvent.click(screen.getByText("Flat List"));
for (const checkbox of screen.getAllByRole("checkbox")) {
expect(checkbox).toBeChecked();
}
});
it("writes the wildcard when Select All covers every displayed tool", async () => {
const mockOnChange = vi.fn();
renderWithProviders(
<MCPToolPermissions
accessToken={mockAccessToken}
selectedServers={[wildcardServerId]}
toolPermissions={{ [wildcardServerId]: ["read_wiki_structure"] }}
onChange={mockOnChange}
/>,
);
expect(await screen.findByText("read_wiki_structure")).toBeInTheDocument();
await userEvent.click(screen.getByRole("button", { name: "Select All" }));
expect(mockOnChange).toHaveBeenCalledWith({ [wildcardServerId]: ["*"] });
});
it("converts back to an enumerated list when one tool is unchecked from a wildcard grant", async () => {
const mockOnChange = vi.fn();
renderWithProviders(
<MCPToolPermissions
accessToken={mockAccessToken}
selectedServers={[wildcardServerId]}
toolPermissions={{ [wildcardServerId]: ["*"] }}
onChange={mockOnChange}
/>,
);
expect(await screen.findByText("read_wiki_structure")).toBeInTheDocument();
await userEvent.click(screen.getByText("Flat List"));
await userEvent.click(screen.getByRole("checkbox", { name: "ask_question" }));
expect(mockOnChange).toHaveBeenCalledWith({
[wildcardServerId]: ["read_wiki_structure", "read_wiki_contents"],
});
});
});
describe("servers reached indirectly", () => {
const groupServer = {
server_id: "srv-group-1",
@ -432,6 +504,8 @@ describe("MCPToolPermissions", () => {
expect(await screen.findByText("list_issues")).toBeInTheDocument();
await userEvent.click(screen.getByText("Select All"));
// A toolset-sourced server never writes the wildcard: that would create a standing direct
// grant outliving the toolset. The write keeps only the tools this level grants itself.
expect(mockOnChange).toHaveBeenCalledWith({ [toolsetServer.server_id]: ["delete_issue"] });
});
@ -853,7 +927,7 @@ describe("MCPToolPermissions", () => {
const written = mockOnChange.mock.calls.at(-1)?.[0] as Record<string, string[]>;
expect(written["github_mcp"]).toEqual(["list_issues"]);
expect(written[twin.server_id]).toEqual(["list_issues", "create_issue", "delete_issue"]);
expect(written[twin.server_id]).toEqual(["*"]);
});
it("says nothing about shared names when every key names one server", async () => {

View file

@ -8,7 +8,7 @@ import { useMCPAccessGroups } from "../../app/(dashboard)/hooks/mcpServers/useMC
import { useMCPToolsets } from "../../app/(dashboard)/hooks/mcpServers/useMCPToolsets";
import McpCrudPermissionPanel from "../mcp_tools/McpCrudPermissionPanel";
import { classifyToolOp } from "../../utils/mcpToolCrudClassification";
import { NO_MCP_SERVERS_SENTINEL } from "../mcp_tools/constants";
import { MCP_ALL_TOOLS_WILDCARD, NO_MCP_SERVERS_SENTINEL } from "../mcp_tools/constants";
import {
EffectiveMcpServer,
McpGrantSource,
@ -18,6 +18,8 @@ import {
applyToolPermissionWrite,
emptyMcpAccessGroups,
isConventionServer,
mcpAllowedToolsFor,
mcpGrantsAllTools,
mcpToolState,
resolveEffectiveMcpServers,
} from "./effectiveMcpServers";
@ -135,7 +137,12 @@ const MCPToolPermissions: React.FC<MCPToolPermissionsProps> = ({
}, [servers, accessToken, toolsetsLoading]);
const writeAllowedTools = (entry: EffectiveMcpServer, allowed: string[]) => {
onChange(applyToolPermissionWrite({ toolPermissions, entry, allowed }));
const names = (serverTools[entry.server.server_id] ?? []).map((t) => t.name);
const next =
entry.source.kind !== "toolset" && names.length > 0 && names.every((n) => allowed.includes(n))
? [MCP_ALL_TOOLS_WILDCARD]
: allowed;
onChange(applyToolPermissionWrite({ toolPermissions, entry, allowed: next }));
};
const isDelete = (tool: MCPTool) => classifyToolOp(tool.name, tool.description || "") === "delete";
@ -152,7 +159,9 @@ const MCPToolPermissions: React.FC<MCPToolPermissionsProps> = ({
onOverridesChange?.(applyToolOverrideWrite(write));
return;
}
const current = entry.allowedTools ?? (serverTools[entry.server.server_id] || []).map((t) => t.name);
const current = mcpGrantsAllTools(entry.keyedTools)
? (serverTools[entry.server.server_id] || []).map((t) => t.name)
: entry.allowedTools ?? (serverTools[entry.server.server_id] || []).map((t) => t.name);
writeAllowedTools(entry, checked ? [...current, tool.name] : current.filter((name) => name !== tool.name));
};
@ -238,9 +247,12 @@ const MCPToolPermissions: React.FC<MCPToolPermissionsProps> = ({
const serverId = server.server_id;
const serverName = server.server_name || server.alias || serverId;
const tools = serverTools[serverId] || [];
const grantsAll = mcpGrantsAllTools(entry.keyedTools);
const stateFor = (tool: MCPTool) =>
mcpToolState(entry, tool.name, classifyToolOp(tool.name, tool.description || "") === "delete");
const selectedTools = tools.filter((tool) => stateFor(tool).checked).map((tool) => tool.name);
const selectedTools = grantsAll
? tools.map((tool) => tool.name)
: tools.filter((tool) => stateFor(tool).checked).map((tool) => tool.name);
const isLoading = loadingTools[serverId];
const error = toolErrors[serverId];
const viewMode = viewModes[serverId] ?? "crud";
@ -265,6 +277,11 @@ const MCPToolPermissions: React.FC<MCPToolPermissionsProps> = ({
)}
</div>
{server.description && <p className="text-sm text-muted-foreground">{server.description}</p>}
{grantsAll && (
<p className="text-sm text-muted-foreground mt-1">
All tools allowed, including tools added to this server later
</p>
)}
{entry.ambiguousKeys.length > 0 && (
<p className="text-sm text-amber-700 mt-1">
{`Also granted by ${entry.ambiguousKeys.map((key) => `"${key}"`).join(", ")}, which names another server too. Those tools stay allowed here until the servers no longer share that name`}
@ -376,10 +393,10 @@ const MCPToolPermissions: React.FC<MCPToolPermissionsProps> = ({
<input
type="checkbox"
aria-label={tool.name}
checked={state.checked}
checked={grantsAll || state.checked}
onChange={() => {
if (disabled || state.locked) return;
writeToolToggle(entry, tool, !state.checked);
writeToolToggle(entry, tool, !(grantsAll || state.checked));
}}
disabled={disabled || state.locked}
className="mt-0.5"

View file

@ -7,6 +7,7 @@ import {
emptyMcpAccessGroups,
isConventionServer,
mcpAllowedToolsFor,
mcpGrantsAllTools,
mcpServersForIdentifier,
mcpToolOverridesFor,
mcpToolPermissionKeyFor,
@ -72,6 +73,16 @@ describe("mcpServersForIdentifier", () => {
});
});
describe("mcpGrantsAllTools", () => {
it("is true only when the union carries the wildcard, never for an absent grant", () => {
expect(mcpGrantsAllTools(["*"])).toBe(true);
expect(mcpGrantsAllTools(["read_file", "*"])).toBe(true);
expect(mcpGrantsAllTools(["read_file"])).toBe(false);
expect(mcpGrantsAllTools([])).toBe(false);
expect(mcpGrantsAllTools(undefined)).toBe(false);
});
});
describe("mcpToolPermissionKeyFor", () => {
const target = server({ server_id: "uuid-1", server_name: "github_mcp", alias: "GitHub" });

View file

@ -1,5 +1,6 @@
import { z } from "zod/v4";
import { MCPServer, MCPToolset } from "../mcp_tools/types";
import { MCP_ALL_TOOLS_WILDCARD } from "../mcp_tools/constants";
// Mirrors the backend resolver's union (direct + access_group + tool_perm + toolset), so the
// editor shows exactly the servers this permission level entitles.
@ -156,6 +157,12 @@ export const mcpToolOverridesFor = (
};
};
// An allowed-tools union carrying the wildcard grants every current and future tool on the
// server; `undefined` (no entry at all) is unrestricted for a different reason and is not a
// wildcard grant the editor should expand.
export const mcpGrantsAllTools = (allowed: readonly string[] | undefined): boolean =>
allowed !== undefined && allowed.includes(MCP_ALL_TOOLS_WILDCARD);
// Tool names the given toolsets grant on this server, `undefined` when they grant none.
const mcpToolsetToolsFor = (
server: MCPServer,

View file

@ -3,5 +3,8 @@ export const NO_MCP_SERVERS_SENTINEL = "no-mcp-servers";
export const ALL_PROXY_MCP_SERVERS_SENTINEL = "all-proxy-mcpservers";
// Must match the backend MCP_ALL_TOOLS_WILDCARD constant in litellm/constants.py.
export const MCP_ALL_TOOLS_WILDCARD = "*";
export const MCP_TOOLS_PREVIEW_FORBIDDEN_MESSAGE =
"Tool preview is not available for submissions. Tools will be verified by an admin during review.";

4
uv.lock generated
View file

@ -4959,12 +4959,12 @@ proxy-dev = [
[[package]]
name = "litellm-enterprise"
version = "0.1.70"
version = "0.1.71"
source = { editable = "enterprise" }
[[package]]
name = "litellm-proxy-extras"
version = "0.4.101"
version = "0.4.102"
source = { editable = "litellm-proxy-extras" }
[[package]]