mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-30 01:52:18 +00:00
Merge branch 'main' into litellm_mcp_continuous_tool_defaults
Co-Authored-By: bot_apk <apk@cognition.ai>
This commit is contained in:
commit
d816dc8fbc
86 changed files with 5042 additions and 1478 deletions
|
|
@ -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==",
|
||||
|
|
|
|||
|
|
@ -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==",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
11
litellm-rust/Cargo.lock
generated
11
litellm-rust/Cargo.lock
generated
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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" }
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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?;
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
});
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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`.
|
||||
|
|
|
|||
31
litellm-rust/crates/coroutine/AGENTS.md
Normal file
31
litellm-rust/crates/coroutine/AGENTS.md
Normal 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
|
||||
15
litellm-rust/crates/coroutine/Cargo.toml
Normal file
15
litellm-rust/crates/coroutine/Cargo.toml
Normal 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"] }
|
||||
42
litellm-rust/crates/coroutine/src/co.rs
Normal file
42
litellm-rust/crates/coroutine/src/co.rs
Normal 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
|
||||
}
|
||||
}
|
||||
94
litellm-rust/crates/coroutine/src/coroutine.rs
Normal file
94
litellm-rust/crates/coroutine/src/coroutine.rs
Normal 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() {}
|
||||
}
|
||||
}
|
||||
14
litellm-rust/crates/coroutine/src/error.rs
Normal file
14
litellm-rust/crates/coroutine/src/error.rs
Normal 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;
|
||||
12
litellm-rust/crates/coroutine/src/lib.rs
Normal file
12
litellm-rust/crates/coroutine/src/lib.rs
Normal 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};
|
||||
60
litellm-rust/crates/coroutine/src/reply.rs
Normal file
60
litellm-rust/crates/coroutine/src/reply.rs
Normal 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 })
|
||||
}
|
||||
256
litellm-rust/crates/coroutine/tests/coroutine.rs
Normal file
256
litellm-rust/crates/coroutine/tests/coroutine.rs
Normal 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) });
|
||||
}
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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<'_>);
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
);
|
||||
});
|
||||
}
|
||||
|
|
|
|||
241
litellm-rust/crates/host-python/src/file_reader.rs
Normal file
241
litellm-rust/crates/host-python/src/file_reader.rs
Normal 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");
|
||||
}
|
||||
}
|
||||
|
|
@ -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};
|
||||
|
|
|
|||
|
|
@ -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"] }
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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()))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
|
|||
137
litellm-rust/crates/host/src/machine/call_machine.rs
Normal file
137
litellm-rust/crates/host/src/machine/call_machine.rs
Normal 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()) })
|
||||
}
|
||||
}
|
||||
|
|
@ -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>;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()) })
|
||||
}
|
||||
}
|
||||
17
litellm-rust/crates/host/src/protocol.rs
Normal file
17
litellm-rust/crates/host/src/protocol.rs
Normal 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;
|
||||
}
|
||||
|
|
@ -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;
|
||||
}
|
||||
|
|
@ -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]);
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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))))
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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>> {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 = || {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"]);
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 ())})
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
364
litellm/types/litellm_params.py
Normal file
364
litellm/types/litellm_params.py
Normal 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)
|
||||
|
|
@ -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=())
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
655
tests/unit/types/test_litellm_params.py
Normal file
655
tests/unit/types/test_litellm_params.py
Normal 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
|
||||
|
|
@ -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 () => {
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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" });
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
4
uv.lock
generated
|
|
@ -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]]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue