mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
Merge 8e20a8644f into 6cba73a2c1
This commit is contained in:
commit
d2fc8b2438
59 changed files with 3765 additions and 135 deletions
14
.github/workflows/test-rust.yml
vendored
14
.github/workflows/test-rust.yml
vendored
|
|
@ -10,6 +10,10 @@ on:
|
|||
- ".github/scripts/smoke_test_native_wheel.py"
|
||||
- ".github/scripts/verify_linux_native_wheel.py"
|
||||
- "tests/test_litellm/rust_bridge/native_route_wheel_test.py"
|
||||
- "tests/test_litellm/rust_bridge/sdk_callback_wheel_test.py"
|
||||
- "tests/test_litellm/ocr/**"
|
||||
- "litellm/ocr/**"
|
||||
- "litellm/rust_bridge/**"
|
||||
- ".github/workflows/test-rust.yml"
|
||||
pull_request:
|
||||
branches:
|
||||
|
|
@ -25,6 +29,10 @@ on:
|
|||
- ".github/scripts/smoke_test_native_wheel.py"
|
||||
- ".github/scripts/verify_linux_native_wheel.py"
|
||||
- "tests/test_litellm/rust_bridge/native_route_wheel_test.py"
|
||||
- "tests/test_litellm/rust_bridge/sdk_callback_wheel_test.py"
|
||||
- "tests/test_litellm/ocr/**"
|
||||
- "litellm/ocr/**"
|
||||
- "litellm/rust_bridge/**"
|
||||
- ".github/workflows/test-rust.yml"
|
||||
|
||||
permissions:
|
||||
|
|
@ -133,3 +141,9 @@ jobs:
|
|||
|
||||
- name: Test native route wheel
|
||||
run: python tests/test_litellm/rust_bridge/native_route_wheel_test.py dist/*.whl
|
||||
|
||||
- name: Test installed SDK callback parity and fallback
|
||||
run: |
|
||||
uv venv /tmp/ocr-callback-sdk --python python
|
||||
uv pip install --python /tmp/ocr-callback-sdk/bin/python dist/*.whl
|
||||
/tmp/ocr-callback-sdk/bin/python tests/test_litellm/rust_bridge/sdk_callback_wheel_test.py
|
||||
|
|
|
|||
2
litellm-rust/Cargo.lock
generated
2
litellm-rust/Cargo.lock
generated
|
|
@ -1482,10 +1482,12 @@ name = "litellm-python-interop"
|
|||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"pyo3",
|
||||
"pyo3-async-runtimes",
|
||||
"pythonize",
|
||||
"rstest",
|
||||
"serde",
|
||||
"serde_json",
|
||||
"tokio",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
|
|||
|
|
@ -28,7 +28,7 @@ coverage and production evidence.
|
|||
## Native request boundary
|
||||
|
||||
Native HTTP routes and Responses WebSocket connections accept
|
||||
`native(request, *, options, context)`. The request carries only endpoint payload.
|
||||
`native(request, *, options, context, callback_adapter=None)`. The request carries only endpoint payload.
|
||||
`NativeRequestOptions` carries credentials, typed provider configuration, routing,
|
||||
headers, query parameters, and timeout. `NativeRequestContext` carries call identity,
|
||||
attribution, and typed capability facts separately from the provider payload.
|
||||
|
|
|
|||
|
|
@ -1 +1 @@
|
|||
pub use crate::ocr::{OcrRequest, ocr, ocr_provider_supported};
|
||||
pub use crate::ocr::{OcrRequest, ocr, ocr_provider_supported, ocr_with_observer};
|
||||
|
|
|
|||
|
|
@ -1,20 +1,39 @@
|
|||
use litellm_core::call_lifecycle::CallLifecycleContext;
|
||||
use litellm_core::error::Error;
|
||||
use litellm_core::http_utils::http_request;
|
||||
use litellm_core::ocr::transformation::OcrResponseHandling;
|
||||
use litellm_core::provider_callbacks::ProviderAttemptObserver;
|
||||
use litellm_core::provider_callbacks::handler::{
|
||||
ProviderAttemptContext, ProviderRequest, send_provider_request,
|
||||
};
|
||||
use serde_json::Value;
|
||||
|
||||
use super::common_utils::{poll_document_intelligence, truncate_error_body};
|
||||
use super::common_utils::poll_document_intelligence;
|
||||
use super::hooks::OcrLifecycleHooks;
|
||||
use super::types::PreparedOcrRequest;
|
||||
use crate::client::http_client;
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub(crate) async fn execute_ocr_provider_call(
|
||||
pub(crate) async fn execute_ocr_provider_call<Observer>(
|
||||
request: PreparedOcrRequest,
|
||||
context: &CallLifecycleContext,
|
||||
hooks: &OcrLifecycleHooks,
|
||||
) -> Result<Value, Error> {
|
||||
observer: &mut Observer,
|
||||
) -> Result<Value, Error>
|
||||
where
|
||||
Observer: ProviderAttemptObserver,
|
||||
Observer::Error: std::fmt::Display,
|
||||
{
|
||||
let request = hooks.prepare_provider_request(request).await?;
|
||||
let mut request_builder = http_client().post(&request.url).json(&request.body);
|
||||
let provider_request = ProviderRequest {
|
||||
provider: request.custom_llm_provider.clone(),
|
||||
model: request.model.clone(),
|
||||
body: serde_json::from_value(request.body).map_err(|error| {
|
||||
Error::InvalidRequest(format!("OCR provider request must be an object: {error}"))
|
||||
})?,
|
||||
api_base: request.url.clone(),
|
||||
headers: request.upstream_headers.iter().cloned().collect(),
|
||||
};
|
||||
let mut request_builder = http_client().post(&request.url);
|
||||
for (key, value) in &request.upstream_headers {
|
||||
request_builder = request_builder.header(key, value);
|
||||
}
|
||||
|
|
@ -22,16 +41,24 @@ pub(crate) async fn execute_ocr_provider_call(
|
|||
request_builder = request_builder.timeout(duration);
|
||||
}
|
||||
|
||||
let response = http_request(request_builder)
|
||||
.await
|
||||
.map_err(|err| Error::Network(err.to_string()))?;
|
||||
let response = send_provider_request(
|
||||
request_builder,
|
||||
provider_request,
|
||||
ProviderAttemptContext {
|
||||
call_id: context.litellm_call_id.clone(),
|
||||
trace_id: None,
|
||||
attempt: 1,
|
||||
},
|
||||
observer,
|
||||
)
|
||||
.await?;
|
||||
|
||||
let status = response.status();
|
||||
let status = response.status;
|
||||
if request.config.response_handling() == OcrResponseHandling::AzureDocumentIntelligencePoll
|
||||
&& status.as_u16() == 202
|
||||
{
|
||||
let operation_url = response
|
||||
.headers()
|
||||
.headers
|
||||
.get("operation-location")
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.map(str::to_string)
|
||||
|
|
@ -58,19 +85,7 @@ pub(crate) async fn execute_ocr_provider_call(
|
|||
.into_json());
|
||||
}
|
||||
|
||||
let text = response
|
||||
.text()
|
||||
.await
|
||||
.map_err(|err| Error::Network(err.to_string()))?;
|
||||
|
||||
if !status.is_success() {
|
||||
return Err(Error::Http {
|
||||
status: status.as_u16(),
|
||||
body: truncate_error_body(&text),
|
||||
});
|
||||
}
|
||||
|
||||
let response_json: Value = serde_json::from_str(&text)
|
||||
let response_json: Value = serde_json::from_str(&response.body)
|
||||
.map_err(|err| Error::InvalidResponse(format!("invalid OCR response JSON: {err}")))?;
|
||||
|
||||
Ok(request
|
||||
|
|
|
|||
|
|
@ -127,6 +127,7 @@ impl OcrLifecycleHooks {
|
|||
};
|
||||
Ok(ProviderOcrRequest {
|
||||
model,
|
||||
custom_llm_provider,
|
||||
config,
|
||||
url,
|
||||
body,
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
use crate::integrations::types::RequestHooks;
|
||||
use litellm_core::Error;
|
||||
use litellm_core::call_lifecycle::CallLifecycle;
|
||||
use litellm_core::provider_callbacks::{NoopProviderAttemptObserver, ProviderAttemptObserver};
|
||||
use litellm_core::request_context::LiteLlmRequestContext;
|
||||
use litellm_core::request_options::RequestOptions;
|
||||
use serde_json::Value;
|
||||
|
|
@ -23,14 +24,41 @@ pub async fn ocr(
|
|||
context: &LiteLlmRequestContext,
|
||||
hooks: RequestHooks,
|
||||
) -> Result<Value, Error> {
|
||||
ocr_with_observer(
|
||||
request,
|
||||
options,
|
||||
context,
|
||||
hooks,
|
||||
&mut NoopProviderAttemptObserver,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "ocr",
|
||||
target = "litellm::function_trace",
|
||||
level = "trace",
|
||||
skip_all
|
||||
)]
|
||||
pub async fn ocr_with_observer<Observer>(
|
||||
request: OcrRequest<'_>,
|
||||
options: &RequestOptions,
|
||||
context: &LiteLlmRequestContext,
|
||||
hooks: RequestHooks,
|
||||
observer: &mut Observer,
|
||||
) -> Result<Value, Error>
|
||||
where
|
||||
Observer: ProviderAttemptObserver,
|
||||
Observer::Error: std::fmt::Display,
|
||||
{
|
||||
let PreparedOcrCall {
|
||||
request,
|
||||
context,
|
||||
context: lifecycle,
|
||||
hooks,
|
||||
} = prepare_ocr_call(request, options.clone(), context, hooks);
|
||||
CallLifecycle::default()
|
||||
.run(context, request, &hooks, |request| {
|
||||
execute_ocr_provider_call(request, &hooks)
|
||||
.run(lifecycle.clone(), request, &hooks, |request| {
|
||||
execute_ocr_provider_call(request, &lifecycle, &hooks, observer)
|
||||
})
|
||||
.await
|
||||
}
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ pub(crate) struct PreparedOcrRequest {
|
|||
|
||||
pub(crate) struct ProviderOcrRequest {
|
||||
pub(crate) model: String,
|
||||
pub(crate) custom_llm_provider: String,
|
||||
pub(crate) config: &'static dyn OcrProviderConfig,
|
||||
pub(crate) url: String,
|
||||
pub(crate) body: Value,
|
||||
|
|
|
|||
|
|
@ -54,19 +54,19 @@ pub async fn run(
|
|||
let request = MessagesRequest {
|
||||
model: provider_model,
|
||||
body,
|
||||
options: RequestOptions {
|
||||
api_key: (deployment.litellm_params.api_key.as_deref()).map(|value| value.to_string()),
|
||||
api_base: (deployment.litellm_params.api_base.as_deref())
|
||||
.map(|value| value.to_string()),
|
||||
custom_llm_provider: (custom_llm_provider).map(|value| value.to_string()),
|
||||
extra_headers,
|
||||
timeout: None,
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
let options = RequestOptions {
|
||||
api_key: (deployment.litellm_params.api_key.as_deref()).map(|value| value.to_string()),
|
||||
api_base: (deployment.litellm_params.api_base.as_deref()).map(|value| value.to_string()),
|
||||
custom_llm_provider: (custom_llm_provider).map(|value| value.to_string()),
|
||||
extra_headers,
|
||||
timeout: None,
|
||||
..Default::default()
|
||||
};
|
||||
if request.body.get("stream").and_then(Value::as_bool) == Some(true) {
|
||||
return messages_stream(
|
||||
request,
|
||||
&options,
|
||||
&LiteLlmRequestContext {
|
||||
..Default::default()
|
||||
},
|
||||
|
|
@ -77,6 +77,7 @@ pub async fn run(
|
|||
|
||||
let response = messages(
|
||||
request,
|
||||
&options,
|
||||
&LiteLlmRequestContext {
|
||||
..Default::default()
|
||||
},
|
||||
|
|
|
|||
|
|
@ -12,13 +12,260 @@ use litellm_ai_gateway::integrations::custom_guardrail::{
|
|||
use litellm_ai_gateway::integrations::custom_logger::{
|
||||
CallbackTiming, CallbackValue, CustomLogger, LogFuture, ModelCallDetails,
|
||||
};
|
||||
|
||||
use litellm_ai_gateway::ocr::{OcrRequest, ocr};
|
||||
use litellm_ai_gateway::ocr::{OcrRequest, ocr, ocr_with_observer};
|
||||
use litellm_core::error::Error;
|
||||
use litellm_core::provider_callbacks::{
|
||||
CallbackDecision, ProviderAttemptObserver, ProviderError, ProviderPostCall, ProviderPreCall,
|
||||
};
|
||||
use serde_json::{Map, Value, json};
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
use tokio::net::{TcpListener, TcpStream};
|
||||
|
||||
struct ProviderObserver {
|
||||
events: Arc<Mutex<Vec<&'static str>>>,
|
||||
raw_response: Option<String>,
|
||||
rejected_callback: Option<&'static str>,
|
||||
decision: Option<&'static str>,
|
||||
}
|
||||
|
||||
impl ProviderAttemptObserver for ProviderObserver {
|
||||
type Error = &'static str;
|
||||
|
||||
async fn pre_call(&mut self, input: &ProviderPreCall) -> Result<CallbackDecision, Self::Error> {
|
||||
assert_eq!(input.model, "mistral-ocr-4-1");
|
||||
assert_eq!(input.call_id, "observer-test");
|
||||
assert_eq!(
|
||||
input.request["document"]["document_url"],
|
||||
"https://example.com/document.pdf"
|
||||
);
|
||||
assert!(input.api_base.ends_with("/v1/ocr"));
|
||||
assert!(
|
||||
input
|
||||
.headers
|
||||
.values()
|
||||
.any(|value| value == "Bearer test-key")
|
||||
);
|
||||
self.events.lock().unwrap().push("pre");
|
||||
match (self.rejected_callback, self.decision) {
|
||||
(Some("pre"), _) => Err("observer failure"),
|
||||
(_, Some("replace_pre")) => Ok(CallbackDecision::Replace {
|
||||
payload: Value::Object(
|
||||
input
|
||||
.request
|
||||
.iter()
|
||||
.map(|(key, value)| (key.clone(), value.clone()))
|
||||
.chain(std::iter::once((
|
||||
"callback_replaced".to_string(),
|
||||
json!(true),
|
||||
)))
|
||||
.collect(),
|
||||
),
|
||||
}),
|
||||
(_, Some("reject_pre")) => Ok(CallbackDecision::Reject {
|
||||
message: "callback rejected request".to_string(),
|
||||
status_code: Some(400),
|
||||
}),
|
||||
_ => Ok(CallbackDecision::Unchanged),
|
||||
}
|
||||
}
|
||||
|
||||
async fn post_call(
|
||||
&mut self,
|
||||
input: &ProviderPostCall,
|
||||
) -> Result<CallbackDecision, Self::Error> {
|
||||
self.events.lock().unwrap().push("post");
|
||||
self.raw_response = input.response.as_str().map(str::to_string);
|
||||
match (self.rejected_callback, self.decision) {
|
||||
(Some("post"), _) => Err("observer failure"),
|
||||
(_, Some("replace_post")) => Ok(CallbackDecision::Replace {
|
||||
payload: json!({"pages":[{"index":0,"markdown":"masked"}]}),
|
||||
}),
|
||||
_ => Ok(CallbackDecision::Unchanged),
|
||||
}
|
||||
}
|
||||
|
||||
async fn error(&mut self, input: &ProviderError) -> Result<(), Self::Error> {
|
||||
assert!(input.committed);
|
||||
assert!(!input.message.is_empty());
|
||||
self.events.lock().unwrap().push("error");
|
||||
if self.rejected_callback == Some("error") {
|
||||
Err("observer failure")
|
||||
} else {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn observer_request() -> OcrRequest<'static> {
|
||||
OcrRequest {
|
||||
model: "mistral/mistral-ocr-4-1",
|
||||
document: json!({"type":"document_url","document_url":"https://example.com/document.pdf"}),
|
||||
optional_params: Map::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn observer_options(api_base: &str) -> RequestOptions {
|
||||
RequestOptions {
|
||||
api_key: Some("test-key".into()),
|
||||
api_base: Some(api_base.into()),
|
||||
custom_llm_provider: Some("mistral".into()),
|
||||
timeout: Some(Duration::from_secs(2)),
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
fn observer_context() -> LiteLlmRequestContext {
|
||||
LiteLlmRequestContext {
|
||||
litellm_call_id: Some("observer-test".into()),
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
|
||||
async fn observer_case(status: u16, body: &'static str, decision: Option<&'static str>) {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let url = format!("http://{}/v1", listener.local_addr().unwrap());
|
||||
let events = Arc::new(Mutex::new(Vec::new()));
|
||||
let provider_events = Arc::clone(&events);
|
||||
let server = tokio::spawn(async move {
|
||||
let (mut socket, _) = listener.accept().await.unwrap();
|
||||
let request = read_http_request(&mut socket).await;
|
||||
assert!(request.starts_with("POST /v1/ocr "));
|
||||
assert_eq!(
|
||||
request.contains(r#""callback_replaced":true"#),
|
||||
decision == Some("replace_pre")
|
||||
);
|
||||
provider_events.lock().unwrap().push("http");
|
||||
let response = format!(
|
||||
"HTTP/1.1 {status} Test\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{body}",
|
||||
body.len()
|
||||
);
|
||||
socket.write_all(response.as_bytes()).await.unwrap();
|
||||
});
|
||||
let mut observer = ProviderObserver {
|
||||
events: Arc::clone(&events),
|
||||
raw_response: None,
|
||||
rejected_callback: None,
|
||||
decision,
|
||||
};
|
||||
let result = ocr_with_observer(
|
||||
observer_request(),
|
||||
&observer_options(&url),
|
||||
&observer_context(),
|
||||
RequestHooks::default(),
|
||||
&mut observer,
|
||||
)
|
||||
.await;
|
||||
tokio::time::timeout(Duration::from_secs(2), server)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
if status != 200 {
|
||||
assert!(matches!(result, Err(Error::Http { status: actual, .. }) if actual == status));
|
||||
assert_eq!(*events.lock().unwrap(), ["pre", "http", "error"]);
|
||||
assert_eq!(observer.raw_response, None);
|
||||
} else {
|
||||
assert_eq!(*events.lock().unwrap(), ["pre", "http", "post"]);
|
||||
assert_eq!(observer.raw_response.as_deref(), Some(body));
|
||||
if body == "invalid-json" {
|
||||
assert!(matches!(result, Err(Error::InvalidResponse(_))));
|
||||
} else {
|
||||
assert_eq!(
|
||||
result.unwrap()["pages"][0]["markdown"],
|
||||
if decision == Some("replace_post") {
|
||||
"masked"
|
||||
} else {
|
||||
"ok"
|
||||
}
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn provider_observers_surround_http_and_can_replace_request_or_response() {
|
||||
observer_case(200, r#"{"pages":[{"index":0,"markdown":"ok"}]}"#, None).await;
|
||||
observer_case(200, "invalid-json", None).await;
|
||||
observer_case(401, r#"{"error":"rejected"}"#, None).await;
|
||||
observer_case(
|
||||
200,
|
||||
r#"{"pages":[{"index":0,"markdown":"ok"}]}"#,
|
||||
Some("replace_pre"),
|
||||
)
|
||||
.await;
|
||||
observer_case(
|
||||
200,
|
||||
r#"{"pages":[{"index":0,"markdown":"ok"}]}"#,
|
||||
Some("replace_post"),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn provider_callback_rejection_stops_before_http() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let url = format!("http://{}/v1", listener.local_addr().unwrap());
|
||||
let events = Arc::new(Mutex::new(Vec::new()));
|
||||
let mut observer = ProviderObserver {
|
||||
events: Arc::clone(&events),
|
||||
raw_response: None,
|
||||
rejected_callback: None,
|
||||
decision: Some("reject_pre"),
|
||||
};
|
||||
|
||||
let result = ocr_with_observer(
|
||||
observer_request(),
|
||||
&observer_options(&url),
|
||||
&observer_context(),
|
||||
RequestHooks::default(),
|
||||
&mut observer,
|
||||
)
|
||||
.await;
|
||||
|
||||
assert!(
|
||||
matches!(result, Err(Error::InvalidRequest(message)) if message == "callback rejected request")
|
||||
);
|
||||
assert_eq!(*events.lock().unwrap(), ["pre"]);
|
||||
assert!(
|
||||
tokio::time::timeout(Duration::from_millis(50), listener.accept())
|
||||
.await
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn invalid_ocr_preparation_does_not_call_observers_or_provider() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let url = format!("http://{}/v1", listener.local_addr().unwrap());
|
||||
let events = Arc::new(Mutex::new(Vec::new()));
|
||||
let mut observer = ProviderObserver {
|
||||
events: Arc::clone(&events),
|
||||
raw_response: None,
|
||||
rejected_callback: None,
|
||||
decision: None,
|
||||
};
|
||||
let request = OcrRequest {
|
||||
document: json!(42),
|
||||
..observer_request()
|
||||
};
|
||||
assert!(
|
||||
ocr_with_observer(
|
||||
request,
|
||||
&observer_options(&url),
|
||||
&observer_context(),
|
||||
RequestHooks::default(),
|
||||
&mut observer
|
||||
)
|
||||
.await
|
||||
.is_err()
|
||||
);
|
||||
assert!(events.lock().unwrap().is_empty());
|
||||
assert!(
|
||||
tokio::time::timeout(Duration::from_millis(50), listener.accept())
|
||||
.await
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
|
||||
async fn read_http_headers(socket: &mut TcpStream) -> String {
|
||||
let mut request = Vec::new();
|
||||
let mut buffer = [0_u8; 1024];
|
||||
|
|
|
|||
36
litellm-rust/crates/core/src/hook_contracts.rs
Normal file
36
litellm-rust/crates/core/src/hook_contracts.rs
Normal file
|
|
@ -0,0 +1,36 @@
|
|||
use serde::{Serialize, de::DeserializeOwned};
|
||||
|
||||
pub trait Hook {
|
||||
type Input: Serialize + Send + Sync;
|
||||
type Output: DeserializeOwned + Send;
|
||||
}
|
||||
|
||||
#[macro_export]
|
||||
macro_rules! define_hooks {
|
||||
(
|
||||
$visibility:vis trait $hooks:ident;
|
||||
{ $($method:ident: $marker:ident($input:ty) -> $output:ty = $mode:ident;)* }
|
||||
) => {
|
||||
$(
|
||||
$visibility struct $marker;
|
||||
|
||||
impl $crate::hook_contracts::Hook for $marker {
|
||||
type Input = $input;
|
||||
type Output = $output;
|
||||
}
|
||||
)*
|
||||
|
||||
$visibility trait $hooks: Send {
|
||||
type Error: Send;
|
||||
|
||||
$(
|
||||
fn $method<'a>(
|
||||
&'a mut self,
|
||||
input: &'a <$marker as $crate::hook_contracts::Hook>::Input,
|
||||
) -> impl ::std::future::Future<
|
||||
Output = Result<<$marker as $crate::hook_contracts::Hook>::Output, Self::Error>,
|
||||
> + Send + 'a;
|
||||
)*
|
||||
}
|
||||
};
|
||||
}
|
||||
|
|
@ -5,11 +5,13 @@ pub mod chat_completions;
|
|||
pub mod constants;
|
||||
pub mod eligibility;
|
||||
pub mod error;
|
||||
pub mod hook_contracts;
|
||||
pub mod http_utils;
|
||||
pub mod messages;
|
||||
#[cfg(any(feature = "observability", test))]
|
||||
pub mod observability;
|
||||
pub mod ocr;
|
||||
pub mod provider_callbacks;
|
||||
pub mod providers;
|
||||
pub mod realtime;
|
||||
pub mod responses;
|
||||
|
|
|
|||
322
litellm-rust/crates/core/src/provider_callbacks/handler.rs
Normal file
322
litellm-rust/crates/core/src/provider_callbacks/handler.rs
Normal file
|
|
@ -0,0 +1,322 @@
|
|||
use std::collections::BTreeMap;
|
||||
use std::time::{SystemTime, UNIX_EPOCH};
|
||||
|
||||
use reqwest::{RequestBuilder, StatusCode, header::HeaderMap};
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::Error;
|
||||
use crate::http_utils::{http_request, truncate_error_body};
|
||||
use crate::provider_callbacks::{
|
||||
CallbackDecision, ProviderAttemptObserver, ProviderError, ProviderPostCall, ProviderPreCall,
|
||||
};
|
||||
|
||||
pub struct ProviderHttpResponse {
|
||||
pub status: StatusCode,
|
||||
pub headers: HeaderMap,
|
||||
pub body: String,
|
||||
}
|
||||
|
||||
pub struct ProviderRequest {
|
||||
pub provider: String,
|
||||
pub model: String,
|
||||
pub body: BTreeMap<String, Value>,
|
||||
pub api_base: String,
|
||||
pub headers: BTreeMap<String, String>,
|
||||
}
|
||||
|
||||
pub struct ProviderAttemptContext {
|
||||
pub call_id: String,
|
||||
pub trace_id: Option<String>,
|
||||
pub attempt: u32,
|
||||
}
|
||||
|
||||
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
|
||||
pub async fn send_provider_request<Observer>(
|
||||
request: RequestBuilder,
|
||||
input: ProviderRequest,
|
||||
context: ProviderAttemptContext,
|
||||
observer: &mut Observer,
|
||||
) -> Result<ProviderHttpResponse, Error>
|
||||
where
|
||||
Observer: ProviderAttemptObserver,
|
||||
Observer::Error: std::fmt::Display,
|
||||
{
|
||||
let event = ProviderPreCall {
|
||||
provider: input.provider,
|
||||
model: input.model,
|
||||
call_id: context.call_id,
|
||||
trace_id: context.trace_id,
|
||||
attempt: context.attempt,
|
||||
started_at: epoch_seconds(),
|
||||
request: input.body,
|
||||
api_base: input.api_base,
|
||||
headers: input.headers,
|
||||
};
|
||||
let body = match observer.pre_call(&event).await.map_err(callback_error)? {
|
||||
CallbackDecision::Unchanged => Value::Object(event.request.clone().into_iter().collect()),
|
||||
CallbackDecision::Replace { payload } => payload,
|
||||
CallbackDecision::Reject { message, .. } => return Err(Error::InvalidRequest(message)),
|
||||
};
|
||||
let response = match http_request(request.json(&body)).await {
|
||||
Ok(response) => response,
|
||||
Err(error) => {
|
||||
let mapped = transport_error(error);
|
||||
notify_error(observer, &event, &mapped, "provider_request", true).await?;
|
||||
return Err(mapped);
|
||||
}
|
||||
};
|
||||
let status = response.status();
|
||||
let headers = response.headers().clone();
|
||||
let body = match response.text().await {
|
||||
Ok(body) => body,
|
||||
Err(error) => {
|
||||
let mapped = transport_error(error);
|
||||
notify_error(observer, &event, &mapped, "response_body", true).await?;
|
||||
return Err(mapped);
|
||||
}
|
||||
};
|
||||
if !status.is_success() {
|
||||
let error = Error::Http {
|
||||
status: status.as_u16(),
|
||||
body: truncate_error_body(&body),
|
||||
};
|
||||
notify_error(observer, &event, &error, "provider_response", true).await?;
|
||||
return Err(error);
|
||||
}
|
||||
let post_call = ProviderPostCall {
|
||||
provider: event.provider.clone(),
|
||||
model: event.model.clone(),
|
||||
call_id: event.call_id.clone(),
|
||||
trace_id: event.trace_id.clone(),
|
||||
attempt: event.attempt,
|
||||
started_at: event.started_at,
|
||||
response: Value::String(body.clone()),
|
||||
status_code: status.as_u16(),
|
||||
headers: header_values(&headers),
|
||||
ended_at: epoch_seconds(),
|
||||
};
|
||||
let body = match observer
|
||||
.post_call(&post_call)
|
||||
.await
|
||||
.map_err(callback_error)?
|
||||
{
|
||||
CallbackDecision::Unchanged => body,
|
||||
CallbackDecision::Replace {
|
||||
payload: Value::String(replacement),
|
||||
} => replacement,
|
||||
CallbackDecision::Replace { payload } => {
|
||||
serde_json::to_string(&payload).map_err(|error| {
|
||||
Error::InvalidResponse(format!("callback response is invalid: {error}"))
|
||||
})?
|
||||
}
|
||||
CallbackDecision::Reject { message, .. } => return Err(Error::InvalidResponse(message)),
|
||||
};
|
||||
Ok(ProviderHttpResponse {
|
||||
status,
|
||||
headers,
|
||||
body,
|
||||
})
|
||||
}
|
||||
|
||||
async fn notify_error<Observer>(
|
||||
observer: &mut Observer,
|
||||
context: &ProviderPreCall,
|
||||
error: &Error,
|
||||
stage: &'static str,
|
||||
committed: bool,
|
||||
) -> Result<(), Error>
|
||||
where
|
||||
Observer: ProviderAttemptObserver,
|
||||
Observer::Error: std::fmt::Display,
|
||||
{
|
||||
let event = ProviderError {
|
||||
provider: context.provider.clone(),
|
||||
model: context.model.clone(),
|
||||
call_id: context.call_id.clone(),
|
||||
trace_id: context.trace_id.clone(),
|
||||
attempt: context.attempt,
|
||||
started_at: context.started_at,
|
||||
message: error.to_string(),
|
||||
stage,
|
||||
committed,
|
||||
status_code: match error {
|
||||
Error::Http { status, .. } => Some(*status),
|
||||
_ => None,
|
||||
},
|
||||
ended_at: epoch_seconds(),
|
||||
};
|
||||
observer.error(&event).await.map_err(callback_error)
|
||||
}
|
||||
|
||||
fn header_values(headers: &HeaderMap) -> BTreeMap<String, String> {
|
||||
headers
|
||||
.iter()
|
||||
.filter_map(|(name, value)| {
|
||||
value
|
||||
.to_str()
|
||||
.ok()
|
||||
.map(|value| (name.to_string(), value.to_string()))
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn epoch_seconds() -> f64 {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.map(|duration| duration.as_secs_f64())
|
||||
.unwrap_or(0.0)
|
||||
}
|
||||
|
||||
fn callback_error(error: impl std::fmt::Display) -> Error {
|
||||
Error::InvalidResponse(format!("provider callback failed: {error}"))
|
||||
}
|
||||
|
||||
fn transport_error(error: reqwest::Error) -> Error {
|
||||
Error::Network(if error.is_timeout() {
|
||||
"Request timed out".into()
|
||||
} else {
|
||||
error.to_string()
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use serde_json::json;
|
||||
use std::time::Duration;
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
use tokio::net::TcpListener;
|
||||
|
||||
struct Observer {
|
||||
events: Vec<&'static str>,
|
||||
reject: bool,
|
||||
}
|
||||
|
||||
impl ProviderAttemptObserver for Observer {
|
||||
type Error = std::convert::Infallible;
|
||||
|
||||
async fn pre_call(
|
||||
&mut self,
|
||||
event: &ProviderPreCall,
|
||||
) -> Result<CallbackDecision, Self::Error> {
|
||||
assert_eq!(event.provider, "test-provider");
|
||||
assert_eq!(event.model, "test-model");
|
||||
assert_eq!(event.call_id, "call-1");
|
||||
assert_eq!(event.trace_id.as_deref(), Some("trace-1"));
|
||||
assert_eq!(event.attempt, 3);
|
||||
assert_eq!(event.request["input"], "private");
|
||||
self.events.push("pre");
|
||||
Ok(if self.reject {
|
||||
CallbackDecision::Reject {
|
||||
message: "blocked".into(),
|
||||
status_code: Some(400),
|
||||
}
|
||||
} else {
|
||||
CallbackDecision::Replace {
|
||||
payload: json!({"masked": true}),
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
async fn post_call(
|
||||
&mut self,
|
||||
event: &ProviderPostCall,
|
||||
) -> Result<CallbackDecision, Self::Error> {
|
||||
assert_eq!(event.response, json!("raw-response"));
|
||||
assert_eq!(event.status_code, 200);
|
||||
assert_eq!(event.attempt, 3);
|
||||
assert_eq!(event.trace_id.as_deref(), Some("trace-1"));
|
||||
assert!(event.ended_at >= event.started_at);
|
||||
self.events.push("post");
|
||||
Ok(CallbackDecision::Replace {
|
||||
payload: json!("redacted-response"),
|
||||
})
|
||||
}
|
||||
|
||||
async fn error(&mut self, event: &ProviderError) -> Result<(), Self::Error> {
|
||||
assert_eq!(event.status_code, Some(429));
|
||||
assert_eq!(event.stage, "provider_response");
|
||||
assert_eq!(event.attempt, 3);
|
||||
assert_eq!(event.trace_id.as_deref(), Some("trace-1"));
|
||||
assert!(event.committed);
|
||||
self.events.push("error");
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn shared_attempt_preserves_context_decisions_and_provider_errors() {
|
||||
let client = reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(2))
|
||||
.build()
|
||||
.unwrap();
|
||||
for (status, reject) in [(200, false), (429, false), (200, true)] {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let url = format!(
|
||||
"http://{}/provider-operation",
|
||||
listener.local_addr().unwrap()
|
||||
);
|
||||
let server = tokio::spawn(async move {
|
||||
if reject {
|
||||
assert!(
|
||||
tokio::time::timeout(Duration::from_millis(50), listener.accept())
|
||||
.await
|
||||
.is_err()
|
||||
);
|
||||
return;
|
||||
}
|
||||
let (mut socket, _) = listener.accept().await.unwrap();
|
||||
let mut request = Vec::new();
|
||||
let mut buffer = [0; 1024];
|
||||
while !request.ends_with(b"{\"masked\":true}") {
|
||||
let count = socket.read(&mut buffer).await.unwrap();
|
||||
assert!(count > 0);
|
||||
request.extend_from_slice(&buffer[..count]);
|
||||
}
|
||||
let request = String::from_utf8(request).unwrap();
|
||||
assert!(request.starts_with("POST /provider-operation "));
|
||||
assert!(!request.contains("private"));
|
||||
socket.write_all(format!(
|
||||
"HTTP/1.1 {status} Test\r\ncontent-length: 12\r\nconnection: close\r\n\r\nraw-response"
|
||||
).as_bytes()).await.unwrap();
|
||||
});
|
||||
let mut observer = Observer {
|
||||
events: Vec::new(),
|
||||
reject,
|
||||
};
|
||||
let result = send_provider_request(
|
||||
client.post(&url),
|
||||
ProviderRequest {
|
||||
provider: "test-provider".into(),
|
||||
model: "test-model".into(),
|
||||
body: BTreeMap::from([("input".into(), json!("private"))]),
|
||||
api_base: url,
|
||||
headers: BTreeMap::new(),
|
||||
},
|
||||
ProviderAttemptContext {
|
||||
call_id: "call-1".into(),
|
||||
trace_id: Some("trace-1".into()),
|
||||
attempt: 3,
|
||||
},
|
||||
&mut observer,
|
||||
)
|
||||
.await;
|
||||
tokio::time::timeout(Duration::from_secs(2), server)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
if reject {
|
||||
assert!(
|
||||
matches!(result, Err(Error::InvalidRequest(message)) if message == "blocked")
|
||||
);
|
||||
assert_eq!(observer.events, ["pre"]);
|
||||
} else if status == 429 {
|
||||
assert!(matches!(result, Err(Error::Http { status: 429, .. })));
|
||||
assert_eq!(observer.events, ["pre", "error"]);
|
||||
} else {
|
||||
assert_eq!(result.unwrap().body, "redacted-response");
|
||||
assert_eq!(observer.events, ["pre", "post"]);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
207
litellm-rust/crates/core/src/provider_callbacks/mod.rs
Normal file
207
litellm-rust/crates/core/src/provider_callbacks/mod.rs
Normal file
|
|
@ -0,0 +1,207 @@
|
|||
pub mod handler;
|
||||
|
||||
use std::collections::BTreeMap;
|
||||
use std::convert::Infallible;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
|
||||
#[derive(Debug, Deserialize, PartialEq)]
|
||||
#[serde(tag = "action", rename_all = "snake_case")]
|
||||
pub enum CallbackDecision {
|
||||
Unchanged,
|
||||
Replace {
|
||||
payload: Value,
|
||||
},
|
||||
Reject {
|
||||
message: String,
|
||||
status_code: Option<u16>,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Clone, Serialize)]
|
||||
pub struct ProviderPreCall {
|
||||
pub provider: String,
|
||||
pub model: String,
|
||||
pub call_id: String,
|
||||
pub trace_id: Option<String>,
|
||||
pub attempt: u32,
|
||||
pub started_at: f64,
|
||||
pub request: BTreeMap<String, Value>,
|
||||
pub api_base: String,
|
||||
pub headers: BTreeMap<String, String>,
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
pub struct ProviderPostCall {
|
||||
pub provider: String,
|
||||
pub model: String,
|
||||
pub call_id: String,
|
||||
pub trace_id: Option<String>,
|
||||
pub attempt: u32,
|
||||
pub started_at: f64,
|
||||
pub response: Value,
|
||||
pub status_code: u16,
|
||||
pub headers: BTreeMap<String, String>,
|
||||
pub ended_at: f64,
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
pub struct ProviderError {
|
||||
pub provider: String,
|
||||
pub model: String,
|
||||
pub call_id: String,
|
||||
pub trace_id: Option<String>,
|
||||
pub attempt: u32,
|
||||
pub started_at: f64,
|
||||
pub message: String,
|
||||
pub stage: &'static str,
|
||||
pub committed: bool,
|
||||
pub status_code: Option<u16>,
|
||||
pub ended_at: f64,
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
pub struct ProviderStreamEvent {
|
||||
pub provider: String,
|
||||
pub model: String,
|
||||
pub call_id: String,
|
||||
pub trace_id: Option<String>,
|
||||
pub attempt: u32,
|
||||
pub started_at: f64,
|
||||
pub event: Value,
|
||||
pub sequence: u64,
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
pub struct ProviderStreamClose {
|
||||
pub provider: String,
|
||||
pub model: String,
|
||||
pub call_id: String,
|
||||
pub trace_id: Option<String>,
|
||||
pub attempt: u32,
|
||||
pub started_at: f64,
|
||||
pub outcome: String,
|
||||
pub ended_at: f64,
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
pub struct SessionEvent {
|
||||
pub session_id: String,
|
||||
pub call_id: String,
|
||||
pub trace_id: Option<String>,
|
||||
pub event: Option<Value>,
|
||||
pub response_id: Option<String>,
|
||||
pub sequence: Option<u64>,
|
||||
pub message: Option<String>,
|
||||
}
|
||||
|
||||
#[macro_export]
|
||||
macro_rules! provider_attempt_observer_catalog {
|
||||
($consumer:path, $($options:tt)*) => {
|
||||
$consumer! {
|
||||
$($options)*
|
||||
{
|
||||
pre_call: PreCall($crate::provider_callbacks::ProviderPreCall) -> $crate::provider_callbacks::CallbackDecision = direct;
|
||||
post_call: PostCall($crate::provider_callbacks::ProviderPostCall) -> $crate::provider_callbacks::CallbackDecision = direct;
|
||||
error: Error($crate::provider_callbacks::ProviderError) -> () = direct;
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
#[macro_export]
|
||||
macro_rules! streaming_observer_catalog {
|
||||
($consumer:path, $($options:tt)*) => {
|
||||
$consumer! {
|
||||
$($options)*
|
||||
{
|
||||
pre_call: StreamingPreCall($crate::provider_callbacks::ProviderPreCall) -> $crate::provider_callbacks::CallbackDecision = direct;
|
||||
post_call: StreamingPostCall($crate::provider_callbacks::ProviderPostCall) -> $crate::provider_callbacks::CallbackDecision = direct;
|
||||
error: StreamingError($crate::provider_callbacks::ProviderError) -> () = direct;
|
||||
stream_event: StreamEvent($crate::provider_callbacks::ProviderStreamEvent) -> $crate::provider_callbacks::CallbackDecision = direct;
|
||||
stream_close: StreamClose($crate::provider_callbacks::ProviderStreamClose) -> () = direct;
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
#[macro_export]
|
||||
macro_rules! session_observer_catalog {
|
||||
($consumer:path, $($options:tt)*) => {
|
||||
$consumer! {
|
||||
$($options)*
|
||||
{
|
||||
before_connect: BeforeConnect($crate::provider_callbacks::SessionEvent) -> $crate::provider_callbacks::CallbackDecision = awaitable;
|
||||
connected: Connected($crate::provider_callbacks::SessionEvent) -> () = awaitable;
|
||||
before_send: BeforeSend($crate::provider_callbacks::SessionEvent) -> $crate::provider_callbacks::CallbackDecision = awaitable;
|
||||
after_receive: AfterReceive($crate::provider_callbacks::SessionEvent) -> $crate::provider_callbacks::CallbackDecision = awaitable;
|
||||
response_complete: ResponseComplete($crate::provider_callbacks::SessionEvent) -> () = awaitable;
|
||||
response_error: ResponseError($crate::provider_callbacks::SessionEvent) -> () = awaitable;
|
||||
error: SessionError($crate::provider_callbacks::SessionEvent) -> () = awaitable;
|
||||
close: Close($crate::provider_callbacks::SessionEvent) -> () = awaitable;
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
provider_attempt_observer_catalog!(crate::define_hooks, pub trait ProviderAttemptObserver;);
|
||||
streaming_observer_catalog!(crate::define_hooks, pub trait StreamingObserver;);
|
||||
session_observer_catalog!(crate::define_hooks, pub trait SessionObserver;);
|
||||
|
||||
pub struct NoopProviderAttemptObserver;
|
||||
|
||||
impl ProviderAttemptObserver for NoopProviderAttemptObserver {
|
||||
type Error = Infallible;
|
||||
|
||||
async fn pre_call(&mut self, _input: &ProviderPreCall) -> Result<CallbackDecision, Infallible> {
|
||||
Ok(CallbackDecision::Unchanged)
|
||||
}
|
||||
|
||||
async fn post_call(
|
||||
&mut self,
|
||||
_input: &ProviderPostCall,
|
||||
) -> Result<CallbackDecision, Infallible> {
|
||||
Ok(CallbackDecision::Unchanged)
|
||||
}
|
||||
|
||||
async fn error(&mut self, _input: &ProviderError) -> Result<(), Infallible> {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::CallbackDecision;
|
||||
|
||||
#[test]
|
||||
fn callback_decisions_have_a_tagged_wire_contract() {
|
||||
assert_eq!(
|
||||
serde_json::from_value::<CallbackDecision>(json!({"action": "unchanged"})).unwrap(),
|
||||
CallbackDecision::Unchanged
|
||||
);
|
||||
assert_eq!(
|
||||
serde_json::from_value::<CallbackDecision>(
|
||||
json!({"action": "replace", "payload": {"masked": true}})
|
||||
)
|
||||
.unwrap(),
|
||||
CallbackDecision::Replace {
|
||||
payload: json!({"masked": true})
|
||||
}
|
||||
);
|
||||
assert_eq!(
|
||||
serde_json::from_value::<CallbackDecision>(json!({
|
||||
"action": "reject",
|
||||
"message": "blocked",
|
||||
"status_code": 400
|
||||
}))
|
||||
.unwrap(),
|
||||
CallbackDecision::Reject {
|
||||
message: "blocked".to_string(),
|
||||
status_code: Some(400)
|
||||
}
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -603,6 +603,12 @@ impl OcrProviderConfig for AzureDocumentIntelligenceOcrConfig {
|
|||
transform_document_intelligence_response(model, response_json, false)
|
||||
}
|
||||
|
||||
#[tracing::instrument(
|
||||
name = "transform_ocr_response",
|
||||
target = "litellm::function_trace",
|
||||
level = "trace",
|
||||
skip_all
|
||||
)]
|
||||
fn transform_ocr_response_with_params(
|
||||
&self,
|
||||
model: &str,
|
||||
|
|
|
|||
113
litellm-rust/crates/python-bridge/src/callback_bindings.rs
Normal file
113
litellm-rust/crates/python-bridge/src/callback_bindings.rs
Normal file
|
|
@ -0,0 +1,113 @@
|
|||
use std::num::NonZeroUsize;
|
||||
|
||||
use litellm_core::provider_callbacks::{
|
||||
CallbackDecision, ProviderAttemptObserver, ProviderError, ProviderPostCall, ProviderPreCall,
|
||||
};
|
||||
use litellm_python_interop::callback_runtime::{AsyncContext, CallbackRuntime, SyncContext};
|
||||
use pyo3::prelude::*;
|
||||
|
||||
use crate::constants::OCR_CALLBACK_CAPACITY;
|
||||
use crate::execution::PythonCallContext;
|
||||
|
||||
litellm_core::provider_attempt_observer_catalog!(crate::bind_python_hooks,
|
||||
pub(crate) struct PythonProviderSession;
|
||||
trait ProviderAttemptObserver;
|
||||
);
|
||||
litellm_core::streaming_observer_catalog!(crate::bind_python_hooks,
|
||||
pub struct PythonStreamingSession;
|
||||
trait litellm_core::provider_callbacks::StreamingObserver;
|
||||
);
|
||||
litellm_core::session_observer_catalog!(crate::bind_python_hooks,
|
||||
pub struct PythonSession;
|
||||
trait litellm_core::provider_callbacks::SessionObserver;
|
||||
);
|
||||
|
||||
#[pyclass(frozen)]
|
||||
struct PythonCallbackRuntime(CallbackRuntime);
|
||||
|
||||
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
let _streaming_constructor = PythonStreamingSession::<SyncContext>::new;
|
||||
let _session_constructor = PythonSession::<AsyncContext>::new;
|
||||
let capacity = NonZeroUsize::new(OCR_CALLBACK_CAPACITY)
|
||||
.expect("Python callback capacity is a positive constant");
|
||||
module.add(
|
||||
"__python_callback_runtime__",
|
||||
PythonCallbackRuntime(CallbackRuntime::new(module, capacity)?),
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) enum PythonProviderObserver {
|
||||
Disabled,
|
||||
Sync(PythonProviderSession<SyncContext>),
|
||||
Async(PythonProviderSession<AsyncContext>),
|
||||
}
|
||||
|
||||
impl PythonProviderObserver {
|
||||
pub(crate) fn new(
|
||||
adapter: Option<Py<PyAny>>,
|
||||
context: PythonCallContext<'_>,
|
||||
) -> PyResult<Self> {
|
||||
let Some(adapter) = adapter else {
|
||||
return Ok(Self::Disabled);
|
||||
};
|
||||
let py = context.py;
|
||||
let module = py.import("litellm.rust_bridge._native")?;
|
||||
let runtime = module
|
||||
.getattr("__python_callback_runtime__")?
|
||||
.extract::<PyRef<'_, PythonCallbackRuntime>>()?
|
||||
.0
|
||||
.clone();
|
||||
if context.asynchronous {
|
||||
Ok(Self::Async(PythonProviderSession::new(
|
||||
adapter.bind(py),
|
||||
runtime.async_context(py)?,
|
||||
)?))
|
||||
} else {
|
||||
Ok(Self::Sync(PythonProviderSession::new(
|
||||
adapter.bind(py),
|
||||
runtime.sync_context(py)?,
|
||||
)?))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn python_async_session(
|
||||
adapter: Py<PyAny>,
|
||||
py: Python<'_>,
|
||||
) -> PyResult<PythonSession<AsyncContext>> {
|
||||
let module = py.import("litellm.rust_bridge._native")?;
|
||||
let runtime = module
|
||||
.getattr("__python_callback_runtime__")?
|
||||
.extract::<PyRef<'_, PythonCallbackRuntime>>()?
|
||||
.0
|
||||
.clone();
|
||||
PythonSession::new(adapter.bind(py), runtime.async_context(py)?)
|
||||
}
|
||||
|
||||
impl ProviderAttemptObserver for PythonProviderObserver {
|
||||
type Error = PyErr;
|
||||
|
||||
async fn pre_call(&mut self, input: &ProviderPreCall) -> PyResult<CallbackDecision> {
|
||||
match self {
|
||||
Self::Disabled => Ok(CallbackDecision::Unchanged),
|
||||
Self::Sync(session) => session.pre_call(input).await,
|
||||
Self::Async(session) => session.pre_call(input).await,
|
||||
}
|
||||
}
|
||||
|
||||
async fn post_call(&mut self, input: &ProviderPostCall) -> PyResult<CallbackDecision> {
|
||||
match self {
|
||||
Self::Disabled => Ok(CallbackDecision::Unchanged),
|
||||
Self::Sync(session) => session.post_call(input).await,
|
||||
Self::Async(session) => session.post_call(input).await,
|
||||
}
|
||||
}
|
||||
|
||||
async fn error(&mut self, input: &ProviderError) -> PyResult<()> {
|
||||
match self {
|
||||
Self::Disabled => Ok(()),
|
||||
Self::Sync(session) => session.error(input).await,
|
||||
Self::Async(session) => session.error(input).await,
|
||||
}
|
||||
}
|
||||
}
|
||||
1
litellm-rust/crates/python-bridge/src/constants.rs
Normal file
1
litellm-rust/crates/python-bridge/src/constants.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub(crate) const OCR_CALLBACK_CAPACITY: usize = 1024;
|
||||
|
|
@ -3,7 +3,6 @@ use std::panic::AssertUnwindSafe;
|
|||
use std::time::Duration;
|
||||
|
||||
use futures_util::FutureExt;
|
||||
use litellm_core::error::Error;
|
||||
use litellm_python_interop::{Pythonized, panic_to_pyerr, release_gil};
|
||||
use pyo3::exceptions::PyRuntimeError;
|
||||
use pyo3::prelude::*;
|
||||
|
|
@ -11,14 +10,19 @@ use serde::Serialize;
|
|||
use tokio::runtime::{Handle, Runtime};
|
||||
use tokio::time::{self, MissedTickBehavior};
|
||||
|
||||
pub(crate) fn run_sync<T, F>(
|
||||
pub(crate) struct PythonCallContext<'py> {
|
||||
pub(crate) py: Python<'py>,
|
||||
pub(crate) asynchronous: bool,
|
||||
}
|
||||
|
||||
pub(crate) fn run_sync<T, E>(
|
||||
py: Python<'_>,
|
||||
future: F,
|
||||
map_error: fn(Error) -> PyErr,
|
||||
future: impl Future<Output = Result<T, E>> + Send + 'static,
|
||||
map_error: fn(E) -> PyErr,
|
||||
) -> PyResult<Py<PyAny>>
|
||||
where
|
||||
T: Serialize + Send + 'static,
|
||||
F: Future<Output = Result<T, Error>> + Send + 'static,
|
||||
E: Send + 'static,
|
||||
{
|
||||
run_sync_on(
|
||||
py,
|
||||
|
|
@ -28,15 +32,15 @@ where
|
|||
)
|
||||
}
|
||||
|
||||
fn run_sync_on<T, F>(
|
||||
fn run_sync_on<T, E>(
|
||||
py: Python<'_>,
|
||||
runtime: &Runtime,
|
||||
future: F,
|
||||
map_error: fn(Error) -> PyErr,
|
||||
future: impl Future<Output = Result<T, E>> + Send + 'static,
|
||||
map_error: fn(E) -> PyErr,
|
||||
) -> PyResult<Py<PyAny>>
|
||||
where
|
||||
T: Serialize + Send + 'static,
|
||||
F: Future<Output = Result<T, Error>> + Send + 'static,
|
||||
E: Send + 'static,
|
||||
{
|
||||
if Handle::try_current().is_ok() {
|
||||
return Err(PyRuntimeError::new_err(
|
||||
|
|
@ -45,27 +49,27 @@ where
|
|||
}
|
||||
|
||||
let result = release_gil(py, move || runtime.block_on(wait_for_sync_result(future)))?;
|
||||
let result = map_core_result(result, map_error)?;
|
||||
let result = map_result(result, map_error)?;
|
||||
Pythonized(result).into_pyobject(py).map(Bound::unbind)
|
||||
}
|
||||
|
||||
pub(crate) fn run_async<T, F>(
|
||||
pub(crate) fn run_async<T, E>(
|
||||
py: Python<'_>,
|
||||
future: F,
|
||||
map_error: fn(Error) -> PyErr,
|
||||
future: impl Future<Output = Result<T, E>> + Send + 'static,
|
||||
map_error: fn(E) -> PyErr,
|
||||
) -> PyResult<Bound<'_, PyAny>>
|
||||
where
|
||||
T: Serialize + Send + 'static,
|
||||
F: Future<Output = Result<T, Error>> + Send + 'static,
|
||||
E: Send + 'static,
|
||||
{
|
||||
pyo3_async_runtimes::tokio::future_into_py(py, async move {
|
||||
let result = catch_future_panic(future).await?;
|
||||
let result = map_core_result(result, map_error)?;
|
||||
let result = map_result(result, map_error)?;
|
||||
Ok(Pythonized(result))
|
||||
})
|
||||
}
|
||||
|
||||
fn map_core_result<T>(result: Result<T, Error>, map_error: fn(Error) -> PyErr) -> PyResult<T> {
|
||||
fn map_result<T, E>(result: Result<T, E>, map_error: fn(E) -> PyErr) -> PyResult<T> {
|
||||
match result {
|
||||
Ok(value) => Ok(value),
|
||||
Err(error) => Err(
|
||||
|
|
@ -75,9 +79,9 @@ fn map_core_result<T>(result: Result<T, Error>, map_error: fn(Error) -> PyErr) -
|
|||
}
|
||||
}
|
||||
|
||||
async fn catch_future_panic<T, F>(future: F) -> PyResult<Result<T, Error>>
|
||||
async fn catch_future_panic<T, E, F>(future: F) -> PyResult<Result<T, E>>
|
||||
where
|
||||
F: Future<Output = Result<T, Error>>,
|
||||
F: Future<Output = Result<T, E>>,
|
||||
{
|
||||
AssertUnwindSafe(future)
|
||||
.catch_unwind()
|
||||
|
|
@ -85,9 +89,9 @@ where
|
|||
.map_err(panic_to_pyerr)
|
||||
}
|
||||
|
||||
async fn wait_for_sync_result<T, F>(future: F) -> PyResult<Result<T, Error>>
|
||||
async fn wait_for_sync_result<T, E, F>(future: F) -> PyResult<Result<T, E>>
|
||||
where
|
||||
F: Future<Output = Result<T, Error>>,
|
||||
F: Future<Output = Result<T, E>>,
|
||||
{
|
||||
let future = catch_future_panic(future);
|
||||
tokio::pin!(future);
|
||||
|
|
@ -108,12 +112,14 @@ where
|
|||
mod tests {
|
||||
use std::ffi::CString;
|
||||
use std::future::poll_fn;
|
||||
use std::process::Command;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
use std::sync::{Arc, mpsc};
|
||||
use std::task::Poll;
|
||||
use std::thread;
|
||||
use std::time::Instant;
|
||||
|
||||
use litellm_core::error::Error;
|
||||
use pyo3::panic::PanicException;
|
||||
use pyo3::types::{PyDict, PyModule};
|
||||
use serde::Serializer;
|
||||
|
|
@ -385,6 +391,31 @@ asyncio.run(exercise())
|
|||
|
||||
#[test]
|
||||
fn async_result_delivery_does_not_stall_tokio_workers() {
|
||||
const CHILD_PROCESS: &str = "LITELLM_ASYNC_RESULT_DELIVERY_TEST_CHILD";
|
||||
|
||||
if std::env::var_os(CHILD_PROCESS).is_none() {
|
||||
let output =
|
||||
Command::new(std::env::current_exe().expect("test executable should exist"))
|
||||
.arg("--exact")
|
||||
.arg(
|
||||
thread::current()
|
||||
.name()
|
||||
.expect("test thread should be named"),
|
||||
)
|
||||
.arg("--nocapture")
|
||||
.env(CHILD_PROCESS, "1")
|
||||
.env("TOKIO_WORKER_THREADS", "1")
|
||||
.output()
|
||||
.expect("isolated result-delivery test should start");
|
||||
assert!(
|
||||
output.status.success(),
|
||||
"isolated result-delivery test failed:\n{}\n{}",
|
||||
String::from_utf8_lossy(&output.stdout),
|
||||
String::from_utf8_lossy(&output.stderr),
|
||||
);
|
||||
return;
|
||||
}
|
||||
|
||||
Python::initialize();
|
||||
ASYNC_PROBE_COMPLETED.store(0, Ordering::SeqCst);
|
||||
Python::attach(|py| {
|
||||
|
|
|
|||
|
|
@ -1,12 +1,21 @@
|
|||
mod callback_bindings;
|
||||
#[cfg(test)]
|
||||
#[path = "../tests/callbacks/mod.rs"]
|
||||
mod callback_tests;
|
||||
mod constants;
|
||||
mod diagnostics;
|
||||
mod errors;
|
||||
mod execution;
|
||||
#[cfg(feature = "trace-parity")]
|
||||
mod function_trace;
|
||||
mod marshal;
|
||||
mod python_hook_bindings;
|
||||
mod routes;
|
||||
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
|
||||
use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
|
||||
use litellm_core::provider_callbacks::{CallbackDecision, SessionEvent, SessionObserver};
|
||||
use litellm_core::responses::types::ResponsesWebSocketRequest;
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::PyAny;
|
||||
|
|
@ -14,6 +23,8 @@ use pyo3::types::PyAny;
|
|||
use crate::errors::core_error_to_pyerr;
|
||||
use crate::marshal::{NativeRequestContext, NativeRequestOptions};
|
||||
|
||||
static NEXT_WEBSOCKET_SESSION_ID: AtomicU64 = AtomicU64::new(1);
|
||||
|
||||
#[derive(FromPyObject)]
|
||||
struct WebSocketConnectRequest {
|
||||
url: String,
|
||||
|
|
@ -27,13 +38,14 @@ struct ResponsesWebSocketConnection {
|
|||
#[pymethods]
|
||||
impl ResponsesWebSocketConnection {
|
||||
#[classmethod]
|
||||
#[pyo3(signature = (request, *, options, context))]
|
||||
#[pyo3(signature = (request, *, options, context, callback_adapter=None))]
|
||||
fn connect<'py>(
|
||||
_cls: &Bound<'py, pyo3::types::PyType>,
|
||||
py: Python<'py>,
|
||||
request: WebSocketConnectRequest,
|
||||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
callback_adapter: Option<Py<PyAny>>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let provider_supported = litellm_core::responses::websocket::native_websocket_supported(
|
||||
options.provider("openai"),
|
||||
|
|
@ -43,11 +55,54 @@ impl ResponsesWebSocketConnection {
|
|||
return Err(crate::errors::RustBridgeDeclined::new_err(reason));
|
||||
}
|
||||
let options: litellm_core::request_options::RequestOptions = options.into();
|
||||
let call_id = context.litellm_call_id.clone().unwrap_or_default();
|
||||
let session_id = format!(
|
||||
"responses-websocket-{}",
|
||||
NEXT_WEBSOCKET_SESSION_ID.fetch_add(1, Ordering::Relaxed)
|
||||
);
|
||||
let mut observer = callback_adapter
|
||||
.map(|adapter| crate::callback_bindings::python_async_session(adapter, py))
|
||||
.transpose()?;
|
||||
let request = ResponsesWebSocketRequest { url: request.url };
|
||||
pyo3_async_runtimes::tokio::future_into_py(py, async move {
|
||||
let inner = RustResponsesWebSocketConnection::connect(request, &options, &context)
|
||||
if let Some(observer) = observer.as_mut() {
|
||||
let decision = observer
|
||||
.before_connect(&session_event(&session_id, &call_id, None))
|
||||
.await?;
|
||||
match decision {
|
||||
CallbackDecision::Unchanged => {}
|
||||
CallbackDecision::Replace { .. } => {
|
||||
return Err(pyo3::exceptions::PyValueError::new_err(
|
||||
"before_connect cannot replace WebSocket setup",
|
||||
));
|
||||
}
|
||||
CallbackDecision::Reject { message, .. } => {
|
||||
return Err(pyo3::exceptions::PyValueError::new_err(message));
|
||||
}
|
||||
}
|
||||
}
|
||||
let inner = match RustResponsesWebSocketConnection::connect(request, &options, &context)
|
||||
.await
|
||||
.map_err(core_error_to_pyerr)?;
|
||||
{
|
||||
Ok(inner) => inner,
|
||||
Err(error) => {
|
||||
if let Some(observer) = observer.as_mut() {
|
||||
observer
|
||||
.error(&session_event(
|
||||
&session_id,
|
||||
&call_id,
|
||||
Some(error.to_string()),
|
||||
))
|
||||
.await?;
|
||||
}
|
||||
return Err(core_error_to_pyerr(error));
|
||||
}
|
||||
};
|
||||
if let Some(observer) = observer.as_mut() {
|
||||
observer
|
||||
.connected(&session_event(&session_id, &call_id, None))
|
||||
.await?;
|
||||
}
|
||||
Ok(ResponsesWebSocketConnection { inner })
|
||||
})
|
||||
}
|
||||
|
|
@ -74,6 +129,18 @@ impl ResponsesWebSocketConnection {
|
|||
}
|
||||
}
|
||||
|
||||
fn session_event(session_id: &str, call_id: &str, message: Option<String>) -> SessionEvent {
|
||||
SessionEvent {
|
||||
session_id: session_id.to_string(),
|
||||
call_id: call_id.to_string(),
|
||||
trace_id: None,
|
||||
event: None,
|
||||
response_id: None,
|
||||
sequence: None,
|
||||
message,
|
||||
}
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (_model, custom_llm_provider, *, context))]
|
||||
fn responses_websocket_decline(
|
||||
|
|
@ -95,6 +162,8 @@ mod _native {
|
|||
#[pymodule_init]
|
||||
fn init(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
super::errors::register(module)?;
|
||||
litellm_python_interop::callback_runtime::register(module)?;
|
||||
super::callback_bindings::register(module)?;
|
||||
super::routes::register(module)?;
|
||||
module.add_class::<super::ResponsesWebSocketConnection>()?;
|
||||
module.add_function(wrap_pyfunction!(
|
||||
|
|
@ -229,6 +298,16 @@ mod tests {
|
|||
let code = CString::new(
|
||||
r#"
|
||||
import asyncio
|
||||
import sys
|
||||
import types
|
||||
|
||||
litellm_module = types.ModuleType('litellm')
|
||||
rust_bridge_module = types.ModuleType('litellm.rust_bridge')
|
||||
litellm_module.rust_bridge = rust_bridge_module
|
||||
rust_bridge_module._native = native
|
||||
sys.modules['litellm'] = litellm_module
|
||||
sys.modules['litellm.rust_bridge'] = rust_bridge_module
|
||||
sys.modules['litellm.rust_bridge._native'] = native
|
||||
|
||||
async def exercise():
|
||||
for request, request_options, request_context, field in (
|
||||
|
|
@ -248,7 +327,56 @@ async def exercise():
|
|||
else:
|
||||
raise AssertionError('invalid WebSocket input reached execution')
|
||||
|
||||
connection = await native.ResponsesWebSocketConnection.connect(Request(url=url), options=options, context=context)
|
||||
events = []
|
||||
|
||||
class Adapter:
|
||||
async def before_connect(self, event):
|
||||
events.append(('before_connect', event))
|
||||
return {'action': 'unchanged'}
|
||||
|
||||
async def connected(self, event):
|
||||
events.append(('connected', event))
|
||||
|
||||
async def before_send(self, event):
|
||||
return {'action': 'unchanged'}
|
||||
|
||||
async def after_receive(self, event):
|
||||
return {'action': 'unchanged'}
|
||||
|
||||
async def response_complete(self, event):
|
||||
pass
|
||||
|
||||
async def response_error(self, event):
|
||||
pass
|
||||
|
||||
async def error(self, event):
|
||||
events.append(('error', event))
|
||||
|
||||
async def close(self, event):
|
||||
pass
|
||||
|
||||
class RejectingAdapter(Adapter):
|
||||
async def before_connect(self, event):
|
||||
return {'action': 'reject', 'message': 'blocked', 'status_code': 400}
|
||||
|
||||
try:
|
||||
await native.ResponsesWebSocketConnection.connect(
|
||||
Request(url=url), options=options, context=context, callback_adapter=RejectingAdapter()
|
||||
)
|
||||
except ValueError as error:
|
||||
assert str(error) == 'blocked'
|
||||
else:
|
||||
raise AssertionError('rejected WebSocket setup reached execution')
|
||||
|
||||
connection = await native.ResponsesWebSocketConnection.connect(
|
||||
Request(url=url),
|
||||
options=options,
|
||||
context=replace(context, litellm_call_id='call-1'),
|
||||
callback_adapter=Adapter(),
|
||||
)
|
||||
assert [name for name, _ in events] == ['before_connect', 'connected']
|
||||
assert events[0][1]['call_id'] == 'call-1'
|
||||
assert events[0][1]['session_id'] == events[1][1]['session_id']
|
||||
assert type(connection) is native.ResponsesWebSocketConnection
|
||||
await connection.send_text("from-python")
|
||||
assert await connection.recv_text() == "from-server"
|
||||
|
|
@ -256,6 +384,9 @@ async def exercise():
|
|||
assert await connection.recv_text() is None
|
||||
|
||||
asyncio.run(asyncio.wait_for(exercise(), timeout=5))
|
||||
sys.modules.pop('litellm.rust_bridge._native')
|
||||
sys.modules.pop('litellm.rust_bridge')
|
||||
sys.modules.pop('litellm')
|
||||
"#,
|
||||
)
|
||||
.expect("Python source should not contain null bytes");
|
||||
|
|
|
|||
|
|
@ -0,0 +1,64 @@
|
|||
#[macro_export]
|
||||
macro_rules! callback_return_mode {
|
||||
(direct) => {
|
||||
::litellm_python_interop::callback_runtime::Direct
|
||||
};
|
||||
(awaitable) => {
|
||||
::litellm_python_interop::callback_runtime::Awaitable
|
||||
};
|
||||
}
|
||||
|
||||
#[macro_export]
|
||||
macro_rules! bind_python_hooks {
|
||||
(
|
||||
$visibility:vis struct $session:ident;
|
||||
trait $hooks:path;
|
||||
{ $($method:ident: $marker:ident($input:ty) -> $output:ty = $mode:ident;)* }
|
||||
) => {
|
||||
$visibility struct $session<C> {
|
||||
context: C,
|
||||
$(
|
||||
$method: ::litellm_python_interop::callback_runtime::Callback<
|
||||
$input, $output, $crate::callback_return_mode!($mode),
|
||||
>,
|
||||
)*
|
||||
}
|
||||
|
||||
impl<C> $session<C>
|
||||
where
|
||||
$(C: ::litellm_python_interop::callback_runtime::CallbackContext<
|
||||
$crate::callback_return_mode!($mode),
|
||||
>,)*
|
||||
{
|
||||
$visibility fn new(
|
||||
adapter: &::pyo3::Bound<'_, ::pyo3::PyAny>,
|
||||
context: C,
|
||||
) -> ::pyo3::PyResult<Self> {
|
||||
use ::pyo3::types::PyAnyMethods as _;
|
||||
Ok(Self {
|
||||
context,
|
||||
$(
|
||||
$method: ::litellm_python_interop::callback_runtime::Callback::new(
|
||||
adapter.getattr(stringify!($method))?,
|
||||
)?,
|
||||
)*
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
impl<C> $hooks for $session<C>
|
||||
where
|
||||
$(C: ::litellm_python_interop::callback_runtime::CallbackContext<
|
||||
$crate::callback_return_mode!($mode),
|
||||
>,)*
|
||||
{
|
||||
type Error = ::pyo3::PyErr;
|
||||
|
||||
$(
|
||||
async fn $method(&mut self, input: &$input) -> ::pyo3::PyResult<$output> {
|
||||
self.$method.call(&mut self.context, input).await
|
||||
}
|
||||
)*
|
||||
}
|
||||
};
|
||||
}
|
||||
|
|
@ -21,6 +21,8 @@ fn prepare_transcription(
|
|||
input: AudioTranscriptionInputs,
|
||||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
_callback_adapter: Option<Py<PyAny>>,
|
||||
_python_context: crate::execution::PythonCallContext<'_>,
|
||||
) -> PyResult<impl Future<Output = Result<Value, Error>> + Send + 'static> {
|
||||
let provider_supported = litellm_core::audio_transcription::transcription_provider_supported(
|
||||
options.provider("bedrock"),
|
||||
|
|
|
|||
|
|
@ -23,6 +23,8 @@ fn prepare_chat_completions(
|
|||
input: ChatCompletionsInputs,
|
||||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
_callback_adapter: Option<Py<PyAny>>,
|
||||
_python_context: crate::execution::PythonCallContext<'_>,
|
||||
) -> PyResult<impl Future<Output = Result<ChatCompletionsResponse, Error>> + Send + 'static> {
|
||||
let context: LiteLlmRequestContext = context.into();
|
||||
let messages = required_value("messages", input.messages, Value::is_array, "list")?;
|
||||
|
|
|
|||
|
|
@ -12,26 +12,28 @@ macro_rules! bridge_route {
|
|||
$(, extra = [$($extra:ident),* $(,)?])? $(,)?
|
||||
) => {
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (request, *, options, context))]
|
||||
#[pyo3(signature = (request, *, options, context, callback_adapter=None))]
|
||||
fn $sync_name(
|
||||
py: pyo3::Python<'_>,
|
||||
request: $inputs,
|
||||
options: $crate::marshal::NativeRequestOptions,
|
||||
context: $crate::marshal::NativeRequestContext,
|
||||
callback_adapter: Option<pyo3::Py<pyo3::PyAny>>,
|
||||
) -> pyo3::PyResult<pyo3::Py<pyo3::PyAny>> {
|
||||
let future = $prepare(request, options, context)?;
|
||||
let future = $prepare(request, options, context, callback_adapter, $crate::execution::PythonCallContext {py, asynchronous: false})?;
|
||||
$crate::execution::run_sync(py, future, $map_error)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (request, *, options, context))]
|
||||
#[pyo3(signature = (request, *, options, context, callback_adapter=None))]
|
||||
fn $async_name(
|
||||
py: pyo3::Python<'_>,
|
||||
request: $inputs,
|
||||
options: $crate::marshal::NativeRequestOptions,
|
||||
context: $crate::marshal::NativeRequestContext,
|
||||
callback_adapter: Option<pyo3::Py<pyo3::PyAny>>,
|
||||
) -> pyo3::PyResult<pyo3::Bound<'_, pyo3::PyAny>> {
|
||||
let future = $prepare(request, options, context)?;
|
||||
let future = $prepare(request, options, context, callback_adapter, $crate::execution::PythonCallContext {py, asynchronous: true})?;
|
||||
$crate::execution::run_async(py, future, $map_error)
|
||||
}
|
||||
|
||||
|
|
@ -48,26 +50,28 @@ macro_rules! bridge_route {
|
|||
use super::{$inputs, $map_error, $prepare};
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (request, *, options, context))]
|
||||
#[pyo3(signature = (request, *, options, context, callback_adapter=None))]
|
||||
fn $sync_name(
|
||||
py: pyo3::Python<'_>,
|
||||
request: $inputs,
|
||||
options: $crate::marshal::NativeRequestOptions,
|
||||
context: $crate::marshal::NativeRequestContext,
|
||||
callback_adapter: Option<pyo3::Py<pyo3::PyAny>>,
|
||||
) -> pyo3::PyResult<pyo3::Py<pyo3::PyAny>> {
|
||||
let future = $prepare(request, options, context)?;
|
||||
let future = $prepare(request, options, context, callback_adapter, $crate::execution::PythonCallContext {py, asynchronous: false})?;
|
||||
$crate::execution::run_sync(py, $crate::function_trace::capture(future), $map_error)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
#[pyo3(signature = (request, *, options, context))]
|
||||
#[pyo3(signature = (request, *, options, context, callback_adapter=None))]
|
||||
fn $async_name(
|
||||
py: pyo3::Python<'_>,
|
||||
request: $inputs,
|
||||
options: $crate::marshal::NativeRequestOptions,
|
||||
context: $crate::marshal::NativeRequestContext,
|
||||
callback_adapter: Option<pyo3::Py<pyo3::PyAny>>,
|
||||
) -> pyo3::PyResult<pyo3::Bound<'_, pyo3::PyAny>> {
|
||||
let future = $prepare(request, options, context)?;
|
||||
let future = $prepare(request, options, context, callback_adapter, $crate::execution::PythonCallContext {py, asynchronous: true})?;
|
||||
$crate::execution::run_async(py, $crate::function_trace::capture(future), $map_error)
|
||||
}
|
||||
|
||||
|
|
@ -155,6 +159,8 @@ mod tests {
|
|||
inputs: EchoInputs,
|
||||
_options: crate::marshal::NativeRequestOptions,
|
||||
_context: crate::marshal::NativeRequestContext,
|
||||
_callback_adapter: Option<Py<PyAny>>,
|
||||
_python_context: crate::execution::PythonCallContext<'_>,
|
||||
) -> PyResult<impl Future<Output = Result<String, Error>> + Send + 'static> {
|
||||
FUTURE_DROPPED.store(false, Ordering::SeqCst);
|
||||
let drop_guard = (inputs.value == "pending").then_some(DropGuard);
|
||||
|
|
@ -195,17 +201,25 @@ mod tests {
|
|||
let module = PyModule::new(py, "routes").expect("module should be created");
|
||||
crate::routes::register(&module).expect("routes should register");
|
||||
let routes = [
|
||||
("ocr", "aocr", "(request, *, options, context)"),
|
||||
(
|
||||
"ocr",
|
||||
"aocr",
|
||||
"(request, *, options, context, callback_adapter=None)",
|
||||
),
|
||||
(
|
||||
"transcription",
|
||||
"atranscription",
|
||||
"(request, *, options, context)",
|
||||
"(request, *, options, context, callback_adapter=None)",
|
||||
),
|
||||
(
|
||||
"messages",
|
||||
"amessages",
|
||||
"(request, *, options, context, callback_adapter=None)",
|
||||
),
|
||||
("messages", "amessages", "(request, *, options, context)"),
|
||||
(
|
||||
"chat_completions",
|
||||
"achat_completions",
|
||||
"(request, *, options, context)",
|
||||
"(request, *, options, context, callback_adapter=None)",
|
||||
),
|
||||
];
|
||||
|
||||
|
|
|
|||
|
|
@ -19,6 +19,8 @@ fn prepare_messages(
|
|||
input: MessagesInputs,
|
||||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
_callback_adapter: Option<Py<PyAny>>,
|
||||
_python_context: crate::execution::PythonCallContext<'_>,
|
||||
) -> PyResult<impl Future<Output = Result<AnthropicMessagesResponse, Error>> + Send + 'static> {
|
||||
let provider_supported =
|
||||
litellm_core::messages::messages_provider_supported(options.provider("anthropic"));
|
||||
|
|
|
|||
|
|
@ -1,8 +1,9 @@
|
|||
use crate::callback_bindings::PythonProviderObserver;
|
||||
use crate::errors::ocr_error_to_pyerr;
|
||||
use crate::marshal::{NativeRequestContext, NativeRequestOptions};
|
||||
use litellm_ai_gateway::integrations::types::RequestHooks;
|
||||
use litellm_ai_gateway::io::ocr::OcrRequest;
|
||||
use litellm_ai_gateway::io::ocr::ocr as run_route;
|
||||
use litellm_ai_gateway::io::ocr::ocr_with_observer as run_route;
|
||||
use litellm_core::Error;
|
||||
use litellm_core::request_context::LiteLlmRequestContext;
|
||||
use pyo3::prelude::*;
|
||||
|
|
@ -22,6 +23,8 @@ fn prepare_ocr(
|
|||
input: OcrInputs,
|
||||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
callback_adapter: Option<Py<PyAny>>,
|
||||
python_context: crate::execution::PythonCallContext<'_>,
|
||||
) -> PyResult<impl Future<Output = Result<Value, Error>> + Send + 'static> {
|
||||
let context: LiteLlmRequestContext = context.into();
|
||||
let provider_supported = litellm_ai_gateway::io::ocr::ocr_provider_supported(
|
||||
|
|
@ -33,6 +36,7 @@ fn prepare_ocr(
|
|||
return Err(crate::errors::RustBridgeDeclined::new_err(reason));
|
||||
}
|
||||
let document = input.document;
|
||||
let mut observer = PythonProviderObserver::new(callback_adapter, python_context)?;
|
||||
Ok(async move {
|
||||
run_route(
|
||||
OcrRequest {
|
||||
|
|
@ -46,6 +50,7 @@ fn prepare_ocr(
|
|||
callbacks: Vec::new(),
|
||||
guardrails: Vec::new(),
|
||||
},
|
||||
&mut observer,
|
||||
)
|
||||
.await
|
||||
})
|
||||
|
|
|
|||
432
litellm-rust/crates/python-bridge/tests/callbacks/mod.rs
Normal file
432
litellm-rust/crates/python-bridge/tests/callbacks/mod.rs
Normal file
|
|
@ -0,0 +1,432 @@
|
|||
use std::collections::BTreeMap;
|
||||
use std::ffi::{CStr, CString};
|
||||
use std::num::NonZeroUsize;
|
||||
use std::sync::Arc;
|
||||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
use litellm_core::provider_callbacks::{
|
||||
ProviderError, ProviderPostCall, ProviderPreCall, ProviderStreamClose, ProviderStreamEvent,
|
||||
SessionEvent, SessionObserver, StreamingObserver,
|
||||
};
|
||||
use litellm_python_interop::callback_runtime::CallbackRuntime;
|
||||
use pyo3::exceptions::PyRuntimeError;
|
||||
use pyo3::prelude::*;
|
||||
|
||||
const PYTHON_TEST_SOURCE: &CStr = pyo3::ffi::c_str!(include_str!("test_callbacks.py"));
|
||||
|
||||
macro_rules! provider_catalog {
|
||||
($consumer:path, $($options:tt)*) => {
|
||||
$consumer! {
|
||||
$($options)*
|
||||
{
|
||||
pre_request: PreRequest(crate::callback_tests::domain::Request)
|
||||
-> crate::callback_tests::domain::Request = awaitable;
|
||||
pre_api_call: PreApiCall(crate::callback_tests::domain::BeforeSend)
|
||||
-> () = direct;
|
||||
post_response: PostResponse(crate::callback_tests::domain::Response)
|
||||
-> crate::callback_tests::domain::Response = awaitable;
|
||||
failure: Failure(crate::callback_tests::domain::ProviderFailure)
|
||||
-> () = direct;
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
macro_rules! direct_catalog {
|
||||
($consumer:path, $($options:tt)*) => {
|
||||
$consumer! {
|
||||
$($options)*
|
||||
{
|
||||
transform: Transform(crate::callback_tests::domain::Request)
|
||||
-> crate::callback_tests::domain::Request = direct;
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
mod domain {
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use super::*;
|
||||
|
||||
#[derive(Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct Request {
|
||||
pub text: String,
|
||||
}
|
||||
|
||||
#[derive(Serialize, Deserialize)]
|
||||
#[serde(deny_unknown_fields)]
|
||||
pub struct Response {
|
||||
pub text: String,
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
pub struct BeforeSend {
|
||||
pub body: Request,
|
||||
}
|
||||
|
||||
#[derive(Serialize)]
|
||||
pub struct ProviderFailure {
|
||||
pub message: String,
|
||||
}
|
||||
|
||||
provider_catalog!(litellm_core::define_hooks, pub trait ProviderHooks;);
|
||||
direct_catalog!(litellm_core::define_hooks, pub trait DirectHooks;);
|
||||
|
||||
pub enum CallError<E> {
|
||||
Hook(E),
|
||||
Provider {
|
||||
error: ProviderFailure,
|
||||
observer_error: Option<E>,
|
||||
},
|
||||
}
|
||||
|
||||
struct PreparedCall(BeforeSend);
|
||||
struct ReadyCall(Request);
|
||||
|
||||
impl PreparedCall {
|
||||
async fn finish_hooks<H: ProviderHooks>(
|
||||
self,
|
||||
hooks: &mut H,
|
||||
) -> Result<ReadyCall, H::Error> {
|
||||
hooks.pre_api_call(&self.0).await?;
|
||||
Ok(ReadyCall(self.0.body))
|
||||
}
|
||||
}
|
||||
|
||||
impl ReadyCall {
|
||||
fn send(self, calls: &AtomicUsize, fail: bool) -> Result<Response, ProviderFailure> {
|
||||
calls.fetch_add(1, Ordering::SeqCst);
|
||||
if fail {
|
||||
return Err(ProviderFailure {
|
||||
message: "provider failed".into(),
|
||||
});
|
||||
}
|
||||
Ok(Response {
|
||||
text: format!("processed:{}", self.0.text),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn execute<H: ProviderHooks>(
|
||||
hooks: &mut H,
|
||||
calls: &AtomicUsize,
|
||||
fail: bool,
|
||||
) -> Result<Response, CallError<H::Error>> {
|
||||
let request = hooks
|
||||
.pre_request(&Request {
|
||||
text: "input".into(),
|
||||
})
|
||||
.await
|
||||
.map_err(CallError::Hook)?;
|
||||
let ready = PreparedCall(BeforeSend { body: request })
|
||||
.finish_hooks(hooks)
|
||||
.await
|
||||
.map_err(CallError::Hook)?;
|
||||
let response = match ready.send(calls, fail) {
|
||||
Ok(response) => response,
|
||||
Err(error) => {
|
||||
let observer_error = hooks.failure(&error).await.err();
|
||||
return Err(CallError::Provider {
|
||||
error,
|
||||
observer_error,
|
||||
});
|
||||
}
|
||||
};
|
||||
hooks
|
||||
.post_response(&response)
|
||||
.await
|
||||
.map_err(CallError::Hook)
|
||||
}
|
||||
}
|
||||
|
||||
use domain::{DirectHooks, ProviderHooks};
|
||||
|
||||
provider_catalog!(crate::bind_python_hooks,
|
||||
struct PythonProviderSession;
|
||||
trait domain::ProviderHooks;
|
||||
);
|
||||
direct_catalog!(crate::bind_python_hooks,
|
||||
struct PythonDirectSession;
|
||||
trait domain::DirectHooks;
|
||||
);
|
||||
|
||||
fn map_error(error: domain::CallError<PyErr>) -> PyErr {
|
||||
match error {
|
||||
domain::CallError::Hook(error) => error,
|
||||
domain::CallError::Provider {
|
||||
error,
|
||||
observer_error,
|
||||
} => Python::attach(|py| {
|
||||
let exception = PyRuntimeError::new_err(error.message);
|
||||
if let Some(observer_error) = observer_error {
|
||||
exception
|
||||
.value(py)
|
||||
.setattr("observer_error", observer_error.value(py))
|
||||
.expect("exception should retain its observer error");
|
||||
}
|
||||
exception
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
#[pyclass(frozen)]
|
||||
struct Harness {
|
||||
runtime: CallbackRuntime,
|
||||
calls: Arc<AtomicUsize>,
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl Harness {
|
||||
#[pyo3(signature = (adapter, fail=false))]
|
||||
fn execute<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
adapter: &Bound<'py, PyAny>,
|
||||
fail: bool,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let mut session = PythonProviderSession::new(adapter, self.runtime.async_context(py)?)?;
|
||||
let calls = Arc::clone(&self.calls);
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move { domain::execute(&mut session, &calls, fail).await },
|
||||
map_error,
|
||||
)
|
||||
}
|
||||
|
||||
fn sync(&self, py: Python<'_>, adapter: &Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
|
||||
let mut session = PythonDirectSession::new(adapter, self.runtime.sync_context(py)?)?;
|
||||
crate::execution::run_sync(
|
||||
py,
|
||||
async move {
|
||||
session
|
||||
.transform(&domain::Request {
|
||||
text: "input".into(),
|
||||
})
|
||||
.await
|
||||
},
|
||||
std::convert::identity,
|
||||
)
|
||||
}
|
||||
|
||||
fn interrupt<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
adapter: &Bound<'py, PyAny>,
|
||||
stop: Bound<'py, PyAny>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let mut session = PythonProviderSession::new(adapter, self.runtime.async_context(py)?)?;
|
||||
let stop = pyo3_async_runtimes::into_future_with_locals(
|
||||
&pyo3_async_runtimes::tokio::get_current_locals(py)?,
|
||||
stop,
|
||||
)?;
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move {
|
||||
let request = domain::Request {
|
||||
text: "input".into(),
|
||||
};
|
||||
tokio::select! {
|
||||
result = session.pre_request(&request) => { result?; },
|
||||
result = stop => { result?; },
|
||||
}
|
||||
session.pre_request(&request).await
|
||||
},
|
||||
std::convert::identity,
|
||||
)
|
||||
}
|
||||
|
||||
fn calls(&self) -> usize {
|
||||
self.calls.load(Ordering::SeqCst)
|
||||
}
|
||||
|
||||
fn streaming(&self, py: Python<'_>, adapter: &Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
|
||||
let mut session = crate::callback_bindings::PythonStreamingSession::new(
|
||||
adapter,
|
||||
self.runtime.sync_context(py)?,
|
||||
)?;
|
||||
crate::execution::run_sync(
|
||||
py,
|
||||
async move {
|
||||
let pre = provider_pre_call();
|
||||
assert!(matches!(
|
||||
session.pre_call(&pre).await?,
|
||||
litellm_core::provider_callbacks::CallbackDecision::Unchanged
|
||||
));
|
||||
assert!(matches!(
|
||||
session.post_call(&provider_post_call()).await?,
|
||||
litellm_core::provider_callbacks::CallbackDecision::Unchanged
|
||||
));
|
||||
assert!(matches!(
|
||||
session.stream_event(&provider_stream_event()).await?,
|
||||
litellm_core::provider_callbacks::CallbackDecision::Unchanged
|
||||
));
|
||||
session.stream_close(&provider_stream_close()).await?;
|
||||
session.error(&provider_error()).await
|
||||
},
|
||||
std::convert::identity,
|
||||
)
|
||||
}
|
||||
|
||||
fn session<'py>(
|
||||
&self,
|
||||
py: Python<'py>,
|
||||
adapter: &Bound<'py, PyAny>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let mut session =
|
||||
crate::callback_bindings::PythonSession::new(adapter, self.runtime.async_context(py)?)?;
|
||||
crate::execution::run_async(
|
||||
py,
|
||||
async move {
|
||||
let event = session_event();
|
||||
assert!(matches!(
|
||||
session.before_connect(&event).await?,
|
||||
litellm_core::provider_callbacks::CallbackDecision::Unchanged
|
||||
));
|
||||
session.connected(&event).await?;
|
||||
assert!(matches!(
|
||||
session.before_send(&event).await?,
|
||||
litellm_core::provider_callbacks::CallbackDecision::Unchanged
|
||||
));
|
||||
assert!(matches!(
|
||||
session.after_receive(&event).await?,
|
||||
litellm_core::provider_callbacks::CallbackDecision::Unchanged
|
||||
));
|
||||
session.response_complete(&event).await?;
|
||||
session.response_error(&event).await?;
|
||||
session.error(&event).await?;
|
||||
session.close(&event).await
|
||||
},
|
||||
std::convert::identity,
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_pre_call() -> ProviderPreCall {
|
||||
ProviderPreCall {
|
||||
provider: "test".to_string(),
|
||||
model: "model".to_string(),
|
||||
call_id: "call".to_string(),
|
||||
trace_id: Some("trace".to_string()),
|
||||
attempt: 2,
|
||||
started_at: 1.0,
|
||||
request: BTreeMap::new(),
|
||||
api_base: "https://provider.test".to_string(),
|
||||
headers: BTreeMap::new(),
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_post_call() -> ProviderPostCall {
|
||||
ProviderPostCall {
|
||||
provider: "test".to_string(),
|
||||
model: "model".to_string(),
|
||||
call_id: "call".to_string(),
|
||||
trace_id: Some("trace".to_string()),
|
||||
attempt: 2,
|
||||
started_at: 1.0,
|
||||
response: serde_json::json!({}),
|
||||
status_code: 200,
|
||||
headers: BTreeMap::new(),
|
||||
ended_at: 2.0,
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_stream_event() -> ProviderStreamEvent {
|
||||
ProviderStreamEvent {
|
||||
provider: "test".to_string(),
|
||||
model: "model".to_string(),
|
||||
call_id: "call".to_string(),
|
||||
trace_id: Some("trace".to_string()),
|
||||
attempt: 2,
|
||||
started_at: 1.0,
|
||||
event: serde_json::json!({"type": "delta"}),
|
||||
sequence: 1,
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_stream_close() -> ProviderStreamClose {
|
||||
ProviderStreamClose {
|
||||
provider: "test".to_string(),
|
||||
model: "model".to_string(),
|
||||
call_id: "call".to_string(),
|
||||
trace_id: Some("trace".to_string()),
|
||||
attempt: 2,
|
||||
started_at: 1.0,
|
||||
outcome: "completed".to_string(),
|
||||
ended_at: 2.0,
|
||||
}
|
||||
}
|
||||
|
||||
fn provider_error() -> ProviderError {
|
||||
ProviderError {
|
||||
provider: "test".to_string(),
|
||||
model: "model".to_string(),
|
||||
call_id: "call".to_string(),
|
||||
trace_id: Some("trace".to_string()),
|
||||
attempt: 2,
|
||||
started_at: 1.0,
|
||||
message: "retrying".to_string(),
|
||||
stage: "provider_response",
|
||||
committed: true,
|
||||
status_code: Some(429),
|
||||
ended_at: 2.0,
|
||||
}
|
||||
}
|
||||
|
||||
fn session_event() -> SessionEvent {
|
||||
SessionEvent {
|
||||
session_id: "session".to_string(),
|
||||
call_id: "call".to_string(),
|
||||
trace_id: Some("trace".to_string()),
|
||||
event: Some(serde_json::json!({"type": "response.create"})),
|
||||
response_id: Some("response".to_string()),
|
||||
sequence: Some(1),
|
||||
message: None,
|
||||
}
|
||||
}
|
||||
|
||||
fn run_python_test(name: &str, capacity: usize) {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let module = PyModule::new(py, "callback_test").expect("test module should load");
|
||||
litellm_python_interop::callback_runtime::register(&module).expect("shim should register");
|
||||
let runtime = CallbackRuntime::new(&module, NonZeroUsize::new(capacity).unwrap())
|
||||
.expect("runtime should initialize");
|
||||
let harness = Harness {
|
||||
runtime,
|
||||
calls: Arc::new(AtomicUsize::new(0)),
|
||||
};
|
||||
let tests = PyModule::from_code(
|
||||
py,
|
||||
PYTHON_TEST_SOURCE,
|
||||
c"callback_tests.py",
|
||||
&CString::new(format!("callback_tests_{name}")).unwrap(),
|
||||
)
|
||||
.expect("test definitions should load");
|
||||
tests
|
||||
.call_method1("run", (name, harness))
|
||||
.expect("callback contract should hold");
|
||||
});
|
||||
}
|
||||
|
||||
macro_rules! python_tests {
|
||||
($($name:ident: $capacity:literal,)*) => {
|
||||
$(
|
||||
#[test]
|
||||
fn $name() { run_python_test(stringify!($name), $capacity); }
|
||||
)*
|
||||
};
|
||||
}
|
||||
|
||||
python_tests! {
|
||||
transforms_and_context: 64,
|
||||
callback_errors: 64,
|
||||
retained_callbacks: 64,
|
||||
registration_and_return_contracts: 64,
|
||||
provider_failure: 64,
|
||||
cancellation_and_admission: 1,
|
||||
interrupted_session: 1,
|
||||
synchronous_callbacks: 1,
|
||||
callback_catalogs: 64,
|
||||
}
|
||||
|
|
@ -0,0 +1,422 @@
|
|||
import asyncio
|
||||
import contextvars
|
||||
import gc
|
||||
import inspect
|
||||
import threading
|
||||
import weakref
|
||||
from collections.abc import Awaitable, Callable
|
||||
from dataclasses import dataclass, replace
|
||||
from typing import Final, Protocol
|
||||
|
||||
Payload = dict[str, str]
|
||||
BeforeSend = dict[str, Payload]
|
||||
|
||||
|
||||
class Harness(Protocol):
|
||||
def execute(self, adapter: object, fail: bool = False) -> Awaitable[Payload]: ...
|
||||
def sync(self, adapter: object) -> Payload: ...
|
||||
def interrupt(self, adapter: object, stop: Awaitable[object]) -> Awaitable[Payload]: ...
|
||||
def calls(self) -> int: ...
|
||||
def streaming(self, adapter: object) -> None: ...
|
||||
def session(self, adapter: object) -> Awaitable[None]: ...
|
||||
|
||||
|
||||
REQUEST_CONTEXT: Final = contextvars.ContextVar("request_context", default="default")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Adapter:
|
||||
pre_request: Callable[[Payload], object]
|
||||
pre_api_call: Callable[[BeforeSend], object]
|
||||
post_response: Callable[[Payload], object]
|
||||
failure: Callable[[Payload], object]
|
||||
|
||||
|
||||
class Recorder:
|
||||
def __init__(self) -> None:
|
||||
self.context = REQUEST_CONTEXT.get()
|
||||
self.loop = asyncio.get_running_loop()
|
||||
self.thread = threading.get_ident()
|
||||
self.events: tuple[str, ...] = ()
|
||||
|
||||
def record(self, event: str) -> None:
|
||||
assert REQUEST_CONTEXT.get() == self.context
|
||||
assert asyncio.get_running_loop() is self.loop
|
||||
assert threading.get_ident() == self.thread
|
||||
self.events += (event,)
|
||||
|
||||
async def pre_request(self, request: Payload) -> Payload:
|
||||
self.record("pre_request")
|
||||
await asyncio.sleep(0)
|
||||
return {"text": request["text"].upper()}
|
||||
|
||||
def pre_api_call(self, event: BeforeSend) -> None:
|
||||
self.record("pre_api_call")
|
||||
event["body"]["text"] = "observer mutation"
|
||||
|
||||
async def post_response(self, response: Payload) -> Payload:
|
||||
self.record("post_response")
|
||||
await asyncio.sleep(0)
|
||||
return {"text": response["text"] + "!"}
|
||||
|
||||
def failure(self, error: Payload) -> None:
|
||||
self.record("failure")
|
||||
assert error == {"message": "provider failed"}
|
||||
|
||||
def adapter(self) -> Adapter:
|
||||
return Adapter(self.pre_request, self.pre_api_call, self.post_response, self.failure)
|
||||
|
||||
|
||||
async def transforms_and_context(harness: Harness) -> None:
|
||||
async def run_one(index: int) -> None:
|
||||
token: Final = REQUEST_CONTEXT.set(f"request-{index}")
|
||||
try:
|
||||
recorder: Final = Recorder()
|
||||
result: Final = await harness.execute(recorder.adapter())
|
||||
assert result == {"text": "processed:INPUT!"}
|
||||
assert recorder.events == ("pre_request", "pre_api_call", "post_response")
|
||||
finally:
|
||||
REQUEST_CONTEXT.reset(token)
|
||||
|
||||
await asyncio.gather(*(run_one(index) for index in range(32)))
|
||||
assert harness.calls() == 32
|
||||
|
||||
|
||||
class CallbackAbort(BaseException):
|
||||
pass
|
||||
|
||||
|
||||
async def callback_errors(harness: Harness) -> None:
|
||||
async def check(original: BaseException) -> None:
|
||||
async def reject(request: Payload) -> Payload:
|
||||
raise original
|
||||
|
||||
try:
|
||||
await harness.execute(replace(Recorder().adapter(), pre_request=reject))
|
||||
except BaseException as error:
|
||||
assert error is original
|
||||
assert error.__traceback__ is not None
|
||||
frames: Final = inspect.getinnerframes(error.__traceback__)
|
||||
assert "reject" in tuple(frame.function for frame in frames)
|
||||
else:
|
||||
raise AssertionError("callback failure was lost")
|
||||
|
||||
await check(ValueError("rejected"))
|
||||
await check(CallbackAbort("abort"))
|
||||
await check(asyncio.CancelledError())
|
||||
assert harness.calls() == 0
|
||||
|
||||
|
||||
async def retained_callbacks(harness: Harness) -> None:
|
||||
def start() -> tuple[Awaitable[Payload], weakref.ReferenceType[Recorder]]:
|
||||
recorder: Final = Recorder()
|
||||
return harness.execute(recorder.adapter()), weakref.ref(recorder)
|
||||
|
||||
future, reference = start()
|
||||
gc.collect()
|
||||
assert reference() is not None
|
||||
assert await future == {"text": "processed:INPUT!"}
|
||||
gc.collect()
|
||||
assert reference() is None
|
||||
|
||||
|
||||
async def registration_and_return_contracts(harness: Harness) -> None:
|
||||
adapter: Final = Recorder().adapter()
|
||||
for invalid in (object(), replace(adapter, pre_api_call=None)):
|
||||
try:
|
||||
harness.execute(invalid)
|
||||
except (AttributeError, TypeError):
|
||||
pass
|
||||
else:
|
||||
raise AssertionError("invalid adapter accepted")
|
||||
|
||||
async def wrong_shape(request: Payload) -> int:
|
||||
return 42
|
||||
|
||||
async def extra_field(request: Payload) -> Payload:
|
||||
return {"text": "input", "unknown": "rejected"}
|
||||
|
||||
def not_awaitable(request: Payload) -> Payload:
|
||||
return request
|
||||
|
||||
async def unfinished() -> None:
|
||||
pass
|
||||
|
||||
coroutine: Final = unfinished()
|
||||
|
||||
def wrong_direct(event: BeforeSend) -> object:
|
||||
return coroutine
|
||||
|
||||
def wrong_observer(event: BeforeSend) -> int:
|
||||
return 42
|
||||
|
||||
for invalid_adapter, expected in (
|
||||
(replace(adapter, pre_request=wrong_shape), "typed contract"),
|
||||
(replace(adapter, pre_request=extra_field), "typed contract"),
|
||||
(replace(adapter, pre_request=not_awaitable), "non-awaitable"),
|
||||
(replace(adapter, pre_api_call=wrong_direct), "direct hook returned an awaitable"),
|
||||
(replace(adapter, pre_api_call=wrong_observer), "typed contract"),
|
||||
):
|
||||
try:
|
||||
await harness.execute(invalid_adapter)
|
||||
except TypeError as error:
|
||||
assert expected in str(error)
|
||||
else:
|
||||
raise AssertionError("invalid callback result accepted")
|
||||
assert inspect.getcoroutinestate(coroutine) == inspect.CORO_CLOSED
|
||||
assert harness.calls() == 0
|
||||
|
||||
def future_value(request: Payload) -> asyncio.Future[Payload]:
|
||||
result: Final[asyncio.Future[Payload]] = asyncio.get_running_loop().create_future()
|
||||
result.set_result(request)
|
||||
return result
|
||||
|
||||
assert await harness.execute(replace(adapter, pre_request=future_value)) == {"text": "processed:input!"}
|
||||
assert harness.calls() == 1
|
||||
|
||||
|
||||
async def provider_failure(harness: Harness) -> None:
|
||||
recorder: Final = Recorder()
|
||||
original: Final = LookupError("observer failed")
|
||||
|
||||
def fail_observer(error: Payload) -> None:
|
||||
recorder.failure(error)
|
||||
raise original
|
||||
|
||||
try:
|
||||
await harness.execute(replace(recorder.adapter(), failure=fail_observer), fail=True)
|
||||
except RuntimeError as error:
|
||||
assert str(error) == "provider failed"
|
||||
assert getattr(error, "observer_error") is original
|
||||
else:
|
||||
raise AssertionError("provider error was lost")
|
||||
assert recorder.events == ("pre_request", "pre_api_call", "failure")
|
||||
assert harness.calls() == 1
|
||||
|
||||
|
||||
async def interrupted_session(harness: Harness) -> None:
|
||||
started: Final = asyncio.Event()
|
||||
stopped: Final = asyncio.Event()
|
||||
cleaned: Final = asyncio.Event()
|
||||
|
||||
async def block(request: Payload) -> Payload:
|
||||
started.set()
|
||||
try:
|
||||
await asyncio.Future[None]()
|
||||
finally:
|
||||
cleaned.set()
|
||||
return request
|
||||
|
||||
task: Final = asyncio.ensure_future(
|
||||
harness.interrupt(replace(Recorder().adapter(), pre_request=block), stopped.wait())
|
||||
)
|
||||
await started.wait()
|
||||
stopped.set()
|
||||
try:
|
||||
await task
|
||||
except RuntimeError as error:
|
||||
assert str(error) == "callback session was cancelled"
|
||||
else:
|
||||
raise AssertionError("interrupted session was reused")
|
||||
await cleaned.wait()
|
||||
assert harness.calls() == 0
|
||||
|
||||
|
||||
async def cancellation_and_admission(harness: Harness) -> None:
|
||||
started: Final = asyncio.Event()
|
||||
cleaning: Final = asyncio.Event()
|
||||
release: Final = asyncio.Event()
|
||||
finished: Final = asyncio.Event()
|
||||
recorder: Final = Recorder()
|
||||
|
||||
async def block(request: Payload) -> Payload:
|
||||
started.set()
|
||||
try:
|
||||
await asyncio.Future[None]()
|
||||
finally:
|
||||
cleaning.set()
|
||||
await release.wait()
|
||||
finished.set()
|
||||
return request
|
||||
|
||||
async def assert_full() -> None:
|
||||
try:
|
||||
await harness.execute(Recorder().adapter())
|
||||
except RuntimeError as error:
|
||||
assert str(error) == "callback capacity exhausted"
|
||||
else:
|
||||
raise AssertionError("capacity released before callback completion")
|
||||
|
||||
task: Final = asyncio.ensure_future(harness.execute(replace(recorder.adapter(), pre_request=block)))
|
||||
await started.wait()
|
||||
await assert_full()
|
||||
task.cancel()
|
||||
try:
|
||||
await task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
else:
|
||||
raise AssertionError("request ignored cancellation")
|
||||
await cleaning.wait()
|
||||
await assert_full()
|
||||
assert harness.calls() == 0
|
||||
assert recorder.events == ()
|
||||
release.set()
|
||||
await finished.wait()
|
||||
assert await harness.execute(Recorder().adapter()) == {"text": "processed:INPUT!"}
|
||||
assert harness.calls() == 1
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class DirectAdapter:
|
||||
transform: Callable[[Payload], object]
|
||||
|
||||
|
||||
async def synchronous_callbacks(harness: Harness) -> None:
|
||||
def on_caller_thread() -> None:
|
||||
token: Final = REQUEST_CONTEXT.set("synchronous")
|
||||
caller: Final = threading.get_ident()
|
||||
|
||||
def transform(request: Payload) -> Payload:
|
||||
assert REQUEST_CONTEXT.get() == "synchronous"
|
||||
assert threading.get_ident() == caller
|
||||
try:
|
||||
asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
pass
|
||||
else:
|
||||
raise AssertionError("sync callback invented an event loop")
|
||||
return {"text": request["text"].upper()}
|
||||
|
||||
try:
|
||||
assert harness.sync(DirectAdapter(transform)) == {"text": "INPUT"}
|
||||
finally:
|
||||
REQUEST_CONTEXT.reset(token)
|
||||
|
||||
await asyncio.to_thread(on_caller_thread)
|
||||
|
||||
original: Final = ValueError("sync failure")
|
||||
|
||||
def reject(request: Payload) -> Payload:
|
||||
raise original
|
||||
|
||||
try:
|
||||
harness.sync(DirectAdapter(reject))
|
||||
except ValueError as error:
|
||||
assert error is original
|
||||
else:
|
||||
raise AssertionError("sync callback error was lost")
|
||||
|
||||
def reenter(request: Payload) -> Payload:
|
||||
return harness.sync(DirectAdapter(lambda value: value))
|
||||
|
||||
try:
|
||||
harness.sync(DirectAdapter(reenter))
|
||||
except RuntimeError as error:
|
||||
assert "Tokio context" in str(error)
|
||||
else:
|
||||
raise AssertionError("synchronous re-entry was accepted")
|
||||
|
||||
async def unfinished() -> None:
|
||||
pass
|
||||
|
||||
coroutine: Final = unfinished()
|
||||
try:
|
||||
harness.sync(DirectAdapter(lambda value: coroutine))
|
||||
except TypeError as error:
|
||||
assert str(error) == "direct hook returned an awaitable"
|
||||
else:
|
||||
raise AssertionError("sync callback accepted an awaitable")
|
||||
assert inspect.getcoroutinestate(coroutine) == inspect.CORO_CLOSED
|
||||
|
||||
|
||||
class CatalogAdapter:
|
||||
def __init__(self) -> None:
|
||||
self.events: tuple[str, ...] = ()
|
||||
|
||||
def _record(self, name: str) -> None:
|
||||
self.events += (name,)
|
||||
|
||||
def pre_call(self, payload: object) -> dict[str, str]:
|
||||
self._record("pre_call")
|
||||
return {"action": "unchanged"}
|
||||
|
||||
def post_call(self, payload: object) -> dict[str, str]:
|
||||
self._record("post_call")
|
||||
return {"action": "unchanged"}
|
||||
|
||||
def error(self, payload: object) -> None:
|
||||
self._record("error")
|
||||
|
||||
def stream_event(self, payload: object) -> dict[str, str]:
|
||||
self._record("stream_event")
|
||||
return {"action": "unchanged"}
|
||||
|
||||
def stream_close(self, payload: object) -> None:
|
||||
self._record("stream_close")
|
||||
|
||||
async def before_connect(self, payload: object) -> dict[str, str]:
|
||||
self._record("before_connect")
|
||||
return {"action": "unchanged"}
|
||||
|
||||
async def connected(self, payload: object) -> None:
|
||||
self._record("connected")
|
||||
|
||||
async def before_send(self, payload: object) -> dict[str, str]:
|
||||
self._record("before_send")
|
||||
return {"action": "unchanged"}
|
||||
|
||||
async def after_receive(self, payload: object) -> dict[str, str]:
|
||||
self._record("after_receive")
|
||||
return {"action": "unchanged"}
|
||||
|
||||
async def response_complete(self, payload: object) -> None:
|
||||
self._record("response_complete")
|
||||
|
||||
async def response_error(self, payload: object) -> None:
|
||||
self._record("response_error")
|
||||
|
||||
async def close(self, payload: object) -> None:
|
||||
self._record("close")
|
||||
|
||||
|
||||
class SessionCatalogAdapter(CatalogAdapter):
|
||||
async def error(self, payload: object) -> None:
|
||||
self._record("error")
|
||||
|
||||
|
||||
async def callback_catalogs(harness: Harness) -> None:
|
||||
streaming: Final = CatalogAdapter()
|
||||
harness.streaming(streaming)
|
||||
assert streaming.events == ("pre_call", "post_call", "stream_event", "stream_close", "error")
|
||||
|
||||
session: Final = SessionCatalogAdapter()
|
||||
await harness.session(session)
|
||||
assert session.events == (
|
||||
"before_connect",
|
||||
"connected",
|
||||
"before_send",
|
||||
"after_receive",
|
||||
"response_complete",
|
||||
"response_error",
|
||||
"error",
|
||||
"close",
|
||||
)
|
||||
|
||||
|
||||
TESTS: Final = (
|
||||
transforms_and_context,
|
||||
callback_errors,
|
||||
retained_callbacks,
|
||||
registration_and_return_contracts,
|
||||
provider_failure,
|
||||
cancellation_and_admission,
|
||||
interrupted_session,
|
||||
synchronous_callbacks,
|
||||
callback_catalogs,
|
||||
)
|
||||
|
||||
|
||||
def run(name: str, harness: Harness) -> None:
|
||||
test: Final = next(test for test in TESTS if test.__name__ == name)
|
||||
asyncio.run(asyncio.wait_for(test(harness), timeout=10))
|
||||
|
|
@ -7,8 +7,10 @@ repository.workspace = true
|
|||
|
||||
[dependencies]
|
||||
pyo3.workspace = true
|
||||
pyo3-async-runtimes.workspace = true
|
||||
pythonize.workspace = true
|
||||
serde.workspace = true
|
||||
tokio = { workspace = true, features = ["sync"] }
|
||||
|
||||
[dev-dependencies]
|
||||
rstest.workspace = true
|
||||
|
|
|
|||
|
|
@ -0,0 +1,49 @@
|
|||
import asyncio
|
||||
import inspect
|
||||
from collections.abc import Awaitable, Callable
|
||||
from typing import Final, cast
|
||||
|
||||
|
||||
def invoke_direct(callback: Callable[[object], object], payload: object) -> object:
|
||||
result: Final = callback(payload)
|
||||
if inspect.isawaitable(result):
|
||||
if inspect.iscoroutine(result):
|
||||
result.close()
|
||||
raise TypeError("direct hook returned an awaitable")
|
||||
return result
|
||||
|
||||
|
||||
class Invocation:
|
||||
def __init__(
|
||||
self,
|
||||
callback: Callable[[object], object],
|
||||
payload: object,
|
||||
returns_awaitable: bool,
|
||||
admission: object,
|
||||
) -> None:
|
||||
self.callback = callback
|
||||
self.payload = payload
|
||||
self.returns_awaitable = returns_awaitable
|
||||
self.admission: object | None = admission
|
||||
self.task: asyncio.Task[object] | None = None
|
||||
self.cancelled = False
|
||||
|
||||
async def run(self) -> object:
|
||||
self.task = asyncio.current_task()
|
||||
try:
|
||||
if self.cancelled:
|
||||
raise asyncio.CancelledError()
|
||||
if not self.returns_awaitable:
|
||||
return invoke_direct(self.callback, self.payload)
|
||||
result: Final = self.callback(self.payload)
|
||||
if not inspect.isawaitable(result):
|
||||
raise TypeError("awaitable hook returned a non-awaitable")
|
||||
return await cast(Awaitable[object], result)
|
||||
finally:
|
||||
self.admission = None
|
||||
self.task = None
|
||||
|
||||
def cancel(self) -> None:
|
||||
self.cancelled = True
|
||||
if self.task is not None:
|
||||
self.task.cancel()
|
||||
229
litellm-rust/crates/python-interop/src/callback_runtime/mod.rs
Normal file
229
litellm-rust/crates/python-interop/src/callback_runtime/mod.rs
Normal file
|
|
@ -0,0 +1,229 @@
|
|||
use std::ffi::CStr;
|
||||
use std::marker::PhantomData;
|
||||
use std::num::NonZeroUsize;
|
||||
use std::sync::Arc;
|
||||
use std::thread::{self, ThreadId};
|
||||
|
||||
use pyo3::exceptions::{PyRuntimeError, PyTypeError, PyValueError};
|
||||
use pyo3::prelude::*;
|
||||
use pyo3_async_runtimes::TaskLocals;
|
||||
use serde::{Serialize, de::DeserializeOwned};
|
||||
use tokio::sync::{OwnedSemaphorePermit, Semaphore};
|
||||
|
||||
use crate::{Pythonized, from_py};
|
||||
|
||||
const PYTHON_RUNTIME_SOURCE: &CStr = pyo3::ffi::c_str!(include_str!("invoke.py"));
|
||||
|
||||
pub struct Direct;
|
||||
pub struct Awaitable;
|
||||
|
||||
pub trait ReturnMode: Send {
|
||||
const AWAITABLE: bool;
|
||||
}
|
||||
|
||||
impl ReturnMode for Direct {
|
||||
const AWAITABLE: bool = false;
|
||||
}
|
||||
|
||||
impl ReturnMode for Awaitable {
|
||||
const AWAITABLE: bool = true;
|
||||
}
|
||||
|
||||
pub trait CallbackContext<M: ReturnMode>: Send {
|
||||
fn invoke(
|
||||
&mut self,
|
||||
callable: Py<PyAny>,
|
||||
payload: Py<PyAny>,
|
||||
) -> impl Future<Output = PyResult<Py<PyAny>>> + Send;
|
||||
}
|
||||
|
||||
pub struct Callback<I, O, M> {
|
||||
callable: Py<PyAny>,
|
||||
signature: PhantomData<fn(I) -> (O, M)>,
|
||||
}
|
||||
|
||||
impl<I, O, M> Callback<I, O, M>
|
||||
where
|
||||
I: Serialize + Sync,
|
||||
O: DeserializeOwned + Send,
|
||||
M: ReturnMode,
|
||||
{
|
||||
pub fn new(callable: Bound<'_, PyAny>) -> PyResult<Self> {
|
||||
if !callable.is_callable() {
|
||||
return Err(PyTypeError::new_err("hook binding must be callable"));
|
||||
}
|
||||
Ok(Self {
|
||||
callable: callable.unbind(),
|
||||
signature: PhantomData,
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn call<C: CallbackContext<M>>(&mut self, context: &mut C, input: &I) -> PyResult<O> {
|
||||
let (callable, payload) = Python::attach(|py| {
|
||||
Ok::<_, PyErr>((
|
||||
self.callable.clone_ref(py),
|
||||
Pythonized(input).into_pyobject(py)?.unbind(),
|
||||
))
|
||||
})?;
|
||||
let result = context.invoke(callable, payload).await?;
|
||||
Python::attach(|py| {
|
||||
from_py(result.bind(py))
|
||||
.map_err(|_| PyTypeError::new_err("hook result does not match its typed contract"))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
let shim = PyModule::from_code(
|
||||
module.py(),
|
||||
PYTHON_RUNTIME_SOURCE,
|
||||
c"litellm_callbacks.py",
|
||||
c"_litellm_callbacks",
|
||||
)?;
|
||||
module.add("__callback_runtime__", shim)
|
||||
}
|
||||
|
||||
struct RuntimeState {
|
||||
invocation: Py<PyAny>,
|
||||
direct: Py<PyAny>,
|
||||
capacity: Arc<Semaphore>,
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct CallbackRuntime(Arc<RuntimeState>);
|
||||
|
||||
impl CallbackRuntime {
|
||||
pub fn new(module: &Bound<'_, PyModule>, max_in_flight: NonZeroUsize) -> PyResult<Self> {
|
||||
if max_in_flight.get() > Semaphore::MAX_PERMITS {
|
||||
return Err(PyValueError::new_err(
|
||||
"callback capacity exceeds the runtime limit",
|
||||
));
|
||||
}
|
||||
let shim = module.getattr("__callback_runtime__")?;
|
||||
Ok(Self(Arc::new(RuntimeState {
|
||||
invocation: shim.getattr("Invocation")?.unbind(),
|
||||
direct: shim.getattr("invoke_direct")?.unbind(),
|
||||
capacity: Arc::new(Semaphore::new(max_in_flight.get())),
|
||||
})))
|
||||
}
|
||||
|
||||
pub fn async_context(&self, py: Python<'_>) -> PyResult<AsyncContext> {
|
||||
Ok(AsyncContext {
|
||||
runtime: self.clone(),
|
||||
locals: pyo3_async_runtimes::tokio::get_current_locals(py)?,
|
||||
interrupted: false,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn sync_context(&self, py: Python<'_>) -> PyResult<SyncContext> {
|
||||
Ok(SyncContext {
|
||||
runtime: self.clone(),
|
||||
context: py
|
||||
.import("contextvars")?
|
||||
.call_method0("copy_context")?
|
||||
.unbind(),
|
||||
caller: thread::current().id(),
|
||||
})
|
||||
}
|
||||
|
||||
fn admit(&self) -> PyResult<OwnedSemaphorePermit> {
|
||||
Arc::clone(&self.0.capacity)
|
||||
.try_acquire_owned()
|
||||
.map_err(|_| PyRuntimeError::new_err("callback capacity exhausted"))
|
||||
}
|
||||
}
|
||||
|
||||
pub struct SyncContext {
|
||||
runtime: CallbackRuntime,
|
||||
context: Py<PyAny>,
|
||||
caller: ThreadId,
|
||||
}
|
||||
|
||||
impl CallbackContext<Direct> for SyncContext {
|
||||
async fn invoke(&mut self, callable: Py<PyAny>, payload: Py<PyAny>) -> PyResult<Py<PyAny>> {
|
||||
if thread::current().id() != self.caller {
|
||||
return Err(PyRuntimeError::new_err(
|
||||
"synchronous callbacks must run on the caller thread",
|
||||
));
|
||||
}
|
||||
let _permit = self.runtime.admit()?;
|
||||
Python::attach(|py| {
|
||||
self.context
|
||||
.call_method1(py, "run", (&self.runtime.0.direct, callable, payload))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub struct AsyncContext {
|
||||
runtime: CallbackRuntime,
|
||||
locals: TaskLocals,
|
||||
interrupted: bool,
|
||||
}
|
||||
|
||||
#[pyclass(frozen)]
|
||||
struct Admission {
|
||||
_permit: OwnedSemaphorePermit,
|
||||
}
|
||||
|
||||
impl<M: ReturnMode> CallbackContext<M> for AsyncContext {
|
||||
async fn invoke(&mut self, callable: Py<PyAny>, payload: Py<PyAny>) -> PyResult<Py<PyAny>> {
|
||||
if self.interrupted {
|
||||
return Err(PyRuntimeError::new_err("callback session was cancelled"));
|
||||
}
|
||||
let permit = self.runtime.admit()?;
|
||||
let (mut cancellation, future) = Python::attach(|py| {
|
||||
let invocation = self.runtime.0.invocation.call1(
|
||||
py,
|
||||
(
|
||||
callable,
|
||||
payload,
|
||||
M::AWAITABLE,
|
||||
Admission { _permit: permit },
|
||||
),
|
||||
)?;
|
||||
let cancellation = CancelOnDrop {
|
||||
event_loop: self.locals.event_loop(py).unbind(),
|
||||
invocation: Some(invocation.clone_ref(py)),
|
||||
};
|
||||
let coroutine = invocation.call_method0(py, "run")?;
|
||||
let future = pyo3_async_runtimes::into_future_with_locals(
|
||||
&self.locals,
|
||||
coroutine.clone_ref(py).into_bound(py),
|
||||
);
|
||||
match future {
|
||||
Ok(future) => Ok((cancellation, future)),
|
||||
Err(error) => {
|
||||
coroutine.call_method0(py, "close")?;
|
||||
Err(error)
|
||||
}
|
||||
}
|
||||
})?;
|
||||
self.interrupted = true;
|
||||
let result = future.await;
|
||||
cancellation.invocation = None;
|
||||
self.interrupted = false;
|
||||
result
|
||||
}
|
||||
}
|
||||
|
||||
struct CancelOnDrop {
|
||||
event_loop: Py<PyAny>,
|
||||
invocation: Option<Py<PyAny>>,
|
||||
}
|
||||
|
||||
impl Drop for CancelOnDrop {
|
||||
fn drop(&mut self) {
|
||||
let Some(invocation) = self.invocation.take() else {
|
||||
return;
|
||||
};
|
||||
Python::attach(|py| {
|
||||
let result = invocation.getattr(py, "cancel").and_then(|cancel| {
|
||||
self.event_loop
|
||||
.call_method1(py, "call_soon_threadsafe", (cancel,))
|
||||
});
|
||||
if let Err(error) = result {
|
||||
error.write_unraisable(py, Some(invocation.bind(py)));
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
@ -1,3 +1,4 @@
|
|||
pub mod callback_runtime;
|
||||
mod gil;
|
||||
mod marshal;
|
||||
|
||||
|
|
|
|||
|
|
@ -30,6 +30,7 @@ from litellm.llms.base_llm.ocr.transformation import (
|
|||
)
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.rust_bridge import ocr as rust_ocr_bridge
|
||||
from litellm.rust_bridge.callback_adapters import ProviderLoggingAdapter
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeRequestCapabilities,
|
||||
NativeRequestOptions,
|
||||
|
|
@ -240,27 +241,8 @@ def _prepare_rust_ocr_call(
|
|||
api_base=prepared_request.api_base,
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
)
|
||||
resolved_complete_url: Final = provider_config.get_complete_url(
|
||||
api_base=prepared_request.api_base,
|
||||
model=prepared_request.model,
|
||||
optional_params=prepared_request.optional_params,
|
||||
litellm_params=prepared_request.litellm_params,
|
||||
)
|
||||
rust_api_base: Final = _rust_bridge_api_base(prepared_request, resolve_api_key)
|
||||
rust_optional_params: Final = _rust_bridge_optional_params(prepared_request, resolve_api_key)
|
||||
prepared_request.litellm_logging_obj.pre_call(
|
||||
input="OCR document processing",
|
||||
api_key=resolved_api_key,
|
||||
additional_args={
|
||||
"complete_input_dict": {
|
||||
"model": prepared_request.model,
|
||||
"document": prepared_request.document,
|
||||
**rust_optional_params,
|
||||
},
|
||||
"api_base": resolved_complete_url,
|
||||
"headers": resolved_headers,
|
||||
},
|
||||
)
|
||||
return PreparedNativeCall(
|
||||
request=rust_ocr_bridge.NativeOCRRequest(
|
||||
model=prepared_request.model,
|
||||
|
|
@ -292,6 +274,11 @@ def _prepare_rust_ocr_call(
|
|||
native_response_format=(prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native"),
|
||||
),
|
||||
),
|
||||
callback_adapter=ProviderLoggingAdapter(
|
||||
prepared_request.litellm_logging_obj,
|
||||
"OCR document processing",
|
||||
resolved_api_key,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ from collections.abc import Callable
|
|||
from types import ModuleType
|
||||
from typing import Final, Generic, TypeVar, cast # noqa: TID251 # PyO3 module boundary
|
||||
|
||||
from litellm.rust_bridge.loader import get_native_bridge
|
||||
from litellm.rust_bridge.loader import get_native_bridge, module_route_ready
|
||||
from litellm.rust_bridge.protocols import NativeModule
|
||||
|
||||
BindingT = TypeVar("BindingT")
|
||||
|
|
@ -31,8 +31,12 @@ class NativeBinding(Generic[BindingT]):
|
|||
self,
|
||||
select: Callable[[NativeModule], BindingT],
|
||||
*,
|
||||
route: str = "",
|
||||
required_capabilities: frozenset[str] = frozenset({"callbacks"}),
|
||||
module_loader: Callable[[], ModuleType | None] | None = None,
|
||||
) -> None:
|
||||
self._route: Final = route
|
||||
self._required_capabilities: Final = required_capabilities
|
||||
self._select: Final = select
|
||||
self._module_loader: Final = module_loader
|
||||
self._override: BindingT | None | _Unset = _UNSET
|
||||
|
|
@ -48,7 +52,11 @@ class NativeBinding(Generic[BindingT]):
|
|||
value: Final = self._select(module)
|
||||
except AttributeError:
|
||||
return None
|
||||
return value if callable(value) else None
|
||||
if not callable(value):
|
||||
return None
|
||||
if self._route and not module_route_ready(native, self._route, self._required_capabilities):
|
||||
return None
|
||||
return value
|
||||
|
||||
def override(self, value: BindingT | None) -> None:
|
||||
self._override = value
|
||||
|
|
@ -56,6 +64,9 @@ class NativeBinding(Generic[BindingT]):
|
|||
def reset(self) -> None:
|
||||
self._override = _UNSET
|
||||
|
||||
def is_overridden(self) -> bool:
|
||||
return not isinstance(self._override, _Unset)
|
||||
|
||||
|
||||
_DECLINED: Final = NativeBinding(lambda native: native.RustBridgeDeclined)
|
||||
_UPSTREAM: Final = NativeBinding(lambda native: native.RustUpstreamError)
|
||||
|
|
|
|||
147
litellm/rust_bridge/callback_adapters.py
Normal file
147
litellm/rust_bridge/callback_adapters.py
Normal file
|
|
@ -0,0 +1,147 @@
|
|||
import json
|
||||
from collections.abc import Mapping, MutableMapping
|
||||
from dataclasses import dataclass
|
||||
from typing import Final, Protocol
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, JsonValue
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
from .callbacks import CallbackDecision, CallbackUnchanged, SessionCallbackHandle
|
||||
|
||||
|
||||
class PreCallArguments(TypedDict):
|
||||
complete_input_dict: ReadOnly[Mapping[str, JsonValue]]
|
||||
api_base: ReadOnly[str]
|
||||
headers: ReadOnly[Mapping[str, str]]
|
||||
|
||||
|
||||
class ProviderLogging(Protocol):
|
||||
@property
|
||||
def model_call_details(self) -> MutableMapping[str, object]: ... # mutable-ok: legacy logger stores provider events
|
||||
|
||||
def pre_call(self, *, input: object, api_key: str | None, additional_args: PreCallArguments) -> None: ...
|
||||
|
||||
def post_call(self, *, original_response: str, input: object, api_key: str | None) -> None: ...
|
||||
|
||||
|
||||
class ProviderEvent(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
provider: str
|
||||
model: str
|
||||
call_id: str
|
||||
trace_id: str | None = None
|
||||
attempt: int
|
||||
started_at: float
|
||||
|
||||
|
||||
class ProviderPreCall(ProviderEvent):
|
||||
request: Mapping[str, JsonValue]
|
||||
api_base: str
|
||||
headers: Mapping[str, str]
|
||||
|
||||
|
||||
class ProviderPostCall(ProviderEvent):
|
||||
response: JsonValue
|
||||
status_code: int
|
||||
headers: Mapping[str, str]
|
||||
ended_at: float
|
||||
|
||||
|
||||
class ProviderError(ProviderEvent):
|
||||
message: str
|
||||
stage: str
|
||||
committed: bool
|
||||
status_code: int | None
|
||||
ended_at: float
|
||||
|
||||
|
||||
class StreamEvent(ProviderEvent):
|
||||
event: JsonValue
|
||||
sequence: int
|
||||
|
||||
|
||||
class StreamClose(ProviderEvent):
|
||||
outcome: str
|
||||
ended_at: float
|
||||
|
||||
|
||||
class SessionEvent(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
session_id: str
|
||||
call_id: str
|
||||
trace_id: str | None = None
|
||||
event: JsonValue | None = None
|
||||
response_id: str | None = None
|
||||
sequence: int | None = None
|
||||
message: str | None = None
|
||||
|
||||
|
||||
def _unchanged() -> CallbackUnchanged:
|
||||
return {"action": "unchanged"} # mutable-ok: callback protocol requires a concrete decision payload
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ProviderLoggingAdapter:
|
||||
logging_obj: ProviderLogging
|
||||
input: object
|
||||
api_key: str | None
|
||||
|
||||
def pre_call(self, payload: object, /) -> CallbackDecision:
|
||||
event: Final = ProviderPreCall.model_validate(payload)
|
||||
additional_args: Final[PreCallArguments] = {
|
||||
"complete_input_dict": event.request,
|
||||
"api_base": event.api_base,
|
||||
"headers": event.headers,
|
||||
}
|
||||
self.logging_obj.pre_call(input=self.input, api_key=self.api_key, additional_args=additional_args)
|
||||
return _unchanged()
|
||||
|
||||
def post_call(self, payload: object, /) -> CallbackDecision:
|
||||
event: Final = ProviderPostCall.model_validate(payload)
|
||||
response: Final = event.response if isinstance(event.response, str) else json.dumps(event.response)
|
||||
self.logging_obj.post_call(original_response=response, input=self.input, api_key=self.api_key)
|
||||
return _unchanged()
|
||||
|
||||
def error(self, payload: object, /) -> None:
|
||||
event: Final = ProviderError.model_validate(payload)
|
||||
self.logging_obj.model_call_details["provider_error"] = event.model_dump()
|
||||
|
||||
def stream_event(self, payload: object, /) -> CallbackDecision:
|
||||
event: Final = StreamEvent.model_validate(payload)
|
||||
self.logging_obj.model_call_details["provider_stream_event"] = event.model_dump()
|
||||
return _unchanged()
|
||||
|
||||
def stream_close(self, payload: object, /) -> None:
|
||||
event: Final = StreamClose.model_validate(payload)
|
||||
self.logging_obj.model_call_details["provider_stream_close"] = event.model_dump()
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class SessionCallbackAdapter:
|
||||
callback: SessionCallbackHandle
|
||||
|
||||
def before_connect(self, payload: object, /) -> CallbackDecision:
|
||||
return self.callback.before_connect(SessionEvent.model_validate(payload).model_dump())
|
||||
|
||||
def connected(self, payload: object, /) -> None:
|
||||
self.callback.connected(SessionEvent.model_validate(payload).model_dump())
|
||||
|
||||
def before_send(self, payload: object, /) -> CallbackDecision:
|
||||
return self.callback.before_send(SessionEvent.model_validate(payload).model_dump())
|
||||
|
||||
def after_receive(self, payload: object, /) -> CallbackDecision:
|
||||
return self.callback.after_receive(SessionEvent.model_validate(payload).model_dump())
|
||||
|
||||
def response_complete(self, payload: object, /) -> None:
|
||||
self.callback.response_complete(SessionEvent.model_validate(payload).model_dump())
|
||||
|
||||
def response_error(self, payload: object, /) -> None:
|
||||
self.callback.response_error(SessionEvent.model_validate(payload).model_dump())
|
||||
|
||||
def error(self, payload: object, /) -> None:
|
||||
self.callback.error(SessionEvent.model_validate(payload).model_dump())
|
||||
|
||||
def close(self, payload: object, /) -> None:
|
||||
self.callback.close(SessionEvent.model_validate(payload).model_dump())
|
||||
64
litellm/rust_bridge/callbacks.py
Normal file
64
litellm/rust_bridge/callbacks.py
Normal file
|
|
@ -0,0 +1,64 @@
|
|||
from typing import Literal, Protocol, TypeAlias
|
||||
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
|
||||
class CallbackUnchanged(TypedDict):
|
||||
action: ReadOnly[Literal["unchanged"]]
|
||||
|
||||
|
||||
class CallbackReplace(TypedDict):
|
||||
action: ReadOnly[Literal["replace"]]
|
||||
payload: ReadOnly[object]
|
||||
|
||||
|
||||
class CallbackReject(TypedDict):
|
||||
action: ReadOnly[Literal["reject"]]
|
||||
message: ReadOnly[str]
|
||||
status_code: ReadOnly[int | None]
|
||||
|
||||
|
||||
CallbackDecision: TypeAlias = CallbackUnchanged | CallbackReplace | CallbackReject
|
||||
|
||||
|
||||
class ProviderAttemptCallbackHandle(Protocol):
|
||||
"""Observe the provider operation inside one native call.
|
||||
|
||||
Successful operations receive ``pre_call`` and ``post_call`` once. Failed
|
||||
operations receive ``pre_call`` and ``error`` once. Outer SDK success and
|
||||
failure callbacks remain owned by Python after endpoint dispatch completes.
|
||||
"""
|
||||
|
||||
def pre_call(self, payload: object, /) -> CallbackDecision: ...
|
||||
|
||||
def post_call(self, payload: object, /) -> CallbackDecision: ...
|
||||
|
||||
def error(self, payload: object, /) -> None: ...
|
||||
|
||||
|
||||
class OneShotCallbackHandle(ProviderAttemptCallbackHandle, Protocol):
|
||||
pass
|
||||
|
||||
|
||||
class StreamingCallbackHandle(ProviderAttemptCallbackHandle, Protocol):
|
||||
def stream_event(self, payload: object, /) -> CallbackDecision: ...
|
||||
|
||||
def stream_close(self, payload: object, /) -> None: ...
|
||||
|
||||
|
||||
class SessionCallbackHandle(Protocol):
|
||||
def before_connect(self, payload: object, /) -> CallbackDecision: ...
|
||||
|
||||
def connected(self, payload: object, /) -> None: ...
|
||||
|
||||
def before_send(self, payload: object, /) -> CallbackDecision: ...
|
||||
|
||||
def after_receive(self, payload: object, /) -> CallbackDecision: ...
|
||||
|
||||
def response_complete(self, payload: object, /) -> None: ...
|
||||
|
||||
def response_error(self, payload: object, /) -> None: ...
|
||||
|
||||
def error(self, payload: object, /) -> None: ...
|
||||
|
||||
def close(self, payload: object, /) -> None: ...
|
||||
|
|
@ -2,8 +2,9 @@
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Mapping, Set
|
||||
from types import ModuleType
|
||||
from typing import Final
|
||||
from typing import Final, cast # noqa: TID251 # runtime typing constructs
|
||||
|
||||
_BRIDGE_SENTINEL: Final = object()
|
||||
_cached_bridge: ModuleType | None | object = _BRIDGE_SENTINEL
|
||||
|
|
@ -33,3 +34,21 @@ def reset_native_bridge_cache() -> None:
|
|||
def native_bridge_available() -> bool:
|
||||
"""Whether the packaged Rust extension is importable."""
|
||||
return get_native_bridge() is not None
|
||||
|
||||
|
||||
def native_route_ready(route: str, required_capabilities: frozenset[str] = frozenset()) -> bool:
|
||||
native: Final = get_native_bridge()
|
||||
if native is None:
|
||||
return False
|
||||
return module_route_ready(native, route, required_capabilities)
|
||||
|
||||
|
||||
def module_route_ready(native: ModuleType, route: str, required_capabilities: frozenset[str]) -> bool:
|
||||
ready_endpoints: Final = getattr(native, "ready_endpoints", None)
|
||||
if not isinstance(ready_endpoints, Mapping):
|
||||
return False
|
||||
typed_endpoints: Final = cast( # cast-ok: native metadata was validated as a Mapping
|
||||
Mapping[object, object], ready_endpoints
|
||||
)
|
||||
capabilities: Final = typed_endpoints.get(route)
|
||||
return isinstance(capabilities, Set) and required_capabilities.issubset(capabilities)
|
||||
|
|
|
|||
|
|
@ -41,6 +41,21 @@ def load_rust_aocr() -> RustAocr | None:
|
|||
return _OCR.asynchronous.load()
|
||||
|
||||
|
||||
def supports_callback_adapter(*, asynchronous: bool = False) -> bool:
|
||||
binding = _OCR.asynchronous if asynchronous else _OCR.sync
|
||||
if binding.is_overridden():
|
||||
return False
|
||||
from litellm.rust_bridge import get_native_bridge
|
||||
from litellm.rust_bridge.loader import native_route_ready
|
||||
|
||||
native: Final = get_native_bridge()
|
||||
return (
|
||||
native is not None
|
||||
and hasattr(native, "__python_callback_runtime__")
|
||||
and native_route_ready("ocr", frozenset({"callbacks"}))
|
||||
)
|
||||
|
||||
|
||||
def dispatch_ocr(
|
||||
*,
|
||||
prepare: Callable[[], PreparedNativeCall[NativeOCRRequest]],
|
||||
|
|
|
|||
3
litellm/rust_bridge/ocr_callbacks.py
Normal file
3
litellm/rust_bridge/ocr_callbacks.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
from .callback_adapters import PreCallArguments, ProviderLoggingAdapter
|
||||
|
||||
__all__ = ("PreCallArguments", "ProviderLoggingAdapter")
|
||||
|
|
@ -3,6 +3,7 @@ from __future__ import annotations
|
|||
from collections.abc import Awaitable, Mapping, Sequence
|
||||
from typing import Protocol
|
||||
|
||||
from .callbacks import SessionCallbackHandle
|
||||
from .request import (
|
||||
NativeChatCompletionsRequest,
|
||||
NativeFunction,
|
||||
|
|
@ -53,6 +54,7 @@ class RustResponsesWebSocketConnection(Protocol):
|
|||
*,
|
||||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
callback_adapter: SessionCallbackHandle | None = None,
|
||||
) -> RustResponsesWebSocket: ...
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -5,6 +5,8 @@ from dataclasses import dataclass, replace
|
|||
from types import MappingProxyType
|
||||
from typing import Generic, Protocol, TypeVar
|
||||
|
||||
from .callbacks import OneShotCallbackHandle
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class NativeBedrockOptions:
|
||||
|
|
@ -73,7 +75,7 @@ def vertex_options(params: Mapping[str, object]) -> NativeVertexOptions:
|
|||
)
|
||||
|
||||
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
from typing_extensions import ReadOnly, TypedDict, TypeVar
|
||||
|
||||
|
||||
class NativePreCallDetails(TypedDict):
|
||||
|
|
@ -164,27 +166,45 @@ def with_capabilities(
|
|||
RequestT = TypeVar("RequestT")
|
||||
RequestContraT = TypeVar("RequestContraT", contravariant=True)
|
||||
ResultT = TypeVar("ResultT", covariant=True)
|
||||
CallbackT = TypeVar("CallbackT", default=OneShotCallbackHandle)
|
||||
CallbackContraT = TypeVar("CallbackContraT", contravariant=True, default=OneShotCallbackHandle)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PreparedNativeCall(Generic[RequestT]):
|
||||
class PreparedNativeCall(Generic[RequestT, CallbackT]):
|
||||
request: RequestT
|
||||
options: NativeRequestOptions = NativeRequestOptions()
|
||||
context: NativeRequestContext = NativeRequestContext()
|
||||
callback_adapter: CallbackT | None = None
|
||||
|
||||
|
||||
class NativeFunction(Protocol[RequestContraT, ResultT]):
|
||||
class NativeFunction(Protocol[RequestContraT, ResultT, CallbackContraT]):
|
||||
def __call__(
|
||||
self,
|
||||
request: RequestContraT,
|
||||
*,
|
||||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
callback_adapter: CallbackContraT | None = None,
|
||||
) -> ResultT: ...
|
||||
|
||||
|
||||
def call_native(native: NativeFunction[RequestT, ResultT], prepared: PreparedNativeCall[RequestT]) -> ResultT:
|
||||
return native(prepared.request, options=prepared.options, context=prepared.context)
|
||||
def call_native(
|
||||
native: NativeFunction[RequestT, ResultT, CallbackT],
|
||||
prepared: PreparedNativeCall[RequestT, CallbackT],
|
||||
) -> ResultT:
|
||||
if prepared.callback_adapter is None:
|
||||
return native(
|
||||
prepared.request,
|
||||
options=prepared.options,
|
||||
context=prepared.context,
|
||||
)
|
||||
return native(
|
||||
prepared.request,
|
||||
options=prepared.options,
|
||||
context=prepared.context,
|
||||
callback_adapter=prepared.callback_adapter,
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ import httpx
|
|||
from websockets.exceptions import ConnectionClosedOK
|
||||
|
||||
from litellm.rust_bridge.bindings import UNCHANGED, Unchanged
|
||||
from litellm.rust_bridge.callbacks import SessionCallbackHandle
|
||||
from litellm.rust_bridge.configuration import rust_enabled
|
||||
from litellm.rust_bridge.protocols import (
|
||||
RustResponsesWebSocket,
|
||||
|
|
@ -94,11 +95,12 @@ async def connect(
|
|||
timeout: float | httpx.Timeout | None,
|
||||
model: str = "responses websocket",
|
||||
provider: str = "openai",
|
||||
callback_adapter: SessionCallbackHandle | None = None,
|
||||
fallback: Callable[[], Awaitable[Connection | None]] = async_none,
|
||||
context: NativeRequestContext | None = None,
|
||||
) -> Connection | None:
|
||||
return await _RESPONSES_WEBSOCKET.ainvoke(
|
||||
prepare=lambda: PreparedNativeCall(
|
||||
prepare=lambda: PreparedNativeCall[NativeResponsesWebSocketRequest, SessionCallbackHandle](
|
||||
NativeResponsesWebSocketRequest(
|
||||
url=url,
|
||||
),
|
||||
|
|
@ -115,6 +117,7 @@ async def connect(
|
|||
requires_connection=True,
|
||||
),
|
||||
),
|
||||
callback_adapter=callback_adapter,
|
||||
),
|
||||
call=lambda connection_type, request: call_native(connection_type.connect, request),
|
||||
preflight=lambda: assess_route(_PREFLIGHT, model, provider),
|
||||
|
|
@ -132,6 +135,7 @@ async def open_connection(
|
|||
timeout: float | httpx.Timeout | None,
|
||||
model: str,
|
||||
provider: str,
|
||||
callback_adapter: SessionCallbackHandle | None = None,
|
||||
fallback: Callable[[], AbstractAsyncContextManager[Connection]],
|
||||
context: NativeRequestContext | None = None,
|
||||
) -> AsyncGenerator[Connection]:
|
||||
|
|
@ -146,6 +150,7 @@ async def open_connection(
|
|||
timeout=timeout,
|
||||
model=model,
|
||||
provider=provider,
|
||||
callback_adapter=callback_adapter,
|
||||
fallback=python_connection,
|
||||
context=context,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -90,7 +90,7 @@ class EndpointBinding(Generic[BindingT]):
|
|||
select: Callable[[NativeModule], SelectedT],
|
||||
enabled: RustEnablement,
|
||||
) -> EndpointBinding[SelectedT]:
|
||||
binding: Final = NativeBinding(select)
|
||||
binding: Final = NativeBinding(select, route=route)
|
||||
return EndpointBinding(
|
||||
route=route,
|
||||
load=binding.load,
|
||||
|
|
@ -108,6 +108,9 @@ class EndpointBinding(Generic[BindingT]):
|
|||
raise RuntimeError("only native Rust bridges support binding resets")
|
||||
self._native_binding.reset()
|
||||
|
||||
def is_overridden(self) -> bool:
|
||||
return self._native_binding is not None and self._native_binding.is_overridden()
|
||||
|
||||
def _attempt(
|
||||
self,
|
||||
*,
|
||||
|
|
|
|||
|
|
@ -45,7 +45,7 @@ from litellm.secret_managers.main import get_secret_str
|
|||
from litellm.types.utils import FileTypes, TranscriptionResponse
|
||||
|
||||
_TRANSCRIPTION: Final[EndpointDispatch[RustTranscription, RustAtranscription]] = EndpointDispatch.native(
|
||||
route="audio transcription",
|
||||
route="transcription",
|
||||
sync=lambda native: native.transcription,
|
||||
asynchronous=lambda native: native.atranscription,
|
||||
enabled=always_enabled,
|
||||
|
|
|
|||
|
|
@ -2,6 +2,7 @@ from __future__ import annotations
|
|||
|
||||
import asyncio
|
||||
from collections.abc import Awaitable
|
||||
from contextlib import ExitStack
|
||||
from pathlib import Path
|
||||
from typing import Final, Protocol, cast
|
||||
|
||||
|
|
@ -49,14 +50,40 @@ def _entrypoint(spec: RouteSpec, engine: Engine, *, asynchronous: bool) -> SdkCa
|
|||
|
||||
|
||||
def _collect(
|
||||
function: SdkCall, kwargs: dict[str, object], engine: Engine, *, asynchronous: bool
|
||||
function: SdkCall, kwargs: dict[str, object], engine: Engine, *, route: str, asynchronous: bool
|
||||
) -> tuple[FunctionTraceEvent, ...]:
|
||||
if engine == "rust":
|
||||
return native_trace_events(_invoke(function, kwargs, asynchronous=asynchronous))
|
||||
import litellm
|
||||
|
||||
with profile_python(Path(litellm.__file__).parent, threads=True) as profiler:
|
||||
_invoke(function, kwargs, asynchronous=asynchronous)
|
||||
with ExitStack() as stack:
|
||||
if route == "transcription":
|
||||
from litellm.rust_bridge import get_native_bridge
|
||||
from litellm.rust_bridge.transcription import (
|
||||
RustAtranscription,
|
||||
RustRouteDecline,
|
||||
RustTranscription,
|
||||
configure_rust_transcription,
|
||||
)
|
||||
|
||||
native: Final = get_native_bridge()
|
||||
if native is None:
|
||||
raise RuntimeError("native transcription is required for diagnostic trace parity")
|
||||
# This required-native SDK route is injected only for diagnostic comparison.
|
||||
# Normal callback readiness remains empty before and after this scope.
|
||||
configure_rust_transcription(
|
||||
transcription=cast(RustTranscription, native.transcription),
|
||||
atranscription=cast(RustAtranscription, native.atranscription),
|
||||
decline=cast(RustRouteDecline, native.transcription_decline),
|
||||
)
|
||||
stack.callback(
|
||||
configure_rust_transcription,
|
||||
transcription=None,
|
||||
atranscription=None,
|
||||
decline=None,
|
||||
)
|
||||
with profile_python(Path(litellm.__file__).parent, threads=True) as profiler:
|
||||
_invoke(function, kwargs, asynchronous=asynchronous)
|
||||
return tuple(profiler.events)
|
||||
|
||||
|
||||
|
|
@ -141,6 +168,7 @@ def collect_trace(
|
|||
function,
|
||||
_native_kwargs(spec.route, kwargs) if engine == "rust" else kwargs,
|
||||
engine,
|
||||
route=spec.route,
|
||||
asynchronous=asynchronous,
|
||||
)
|
||||
provider.take_requests(len(fixture.provider_responses))
|
||||
|
|
|
|||
|
|
@ -16,6 +16,8 @@ COMMON_MAPPINGS: Final = (
|
|||
mapping(rust_span="validate_environment", python_frame=r"(?<!_)validate_environment$"),
|
||||
mapping(rust_span="complete_url", python_frame=r"get_complete_url$"),
|
||||
mapping(rust_span="transform_ocr_request", python_frame=r"(?<!async_)transform_ocr_request$"),
|
||||
# Rust wraps the HTTP attempt to invoke observers; Python has no matching frame.
|
||||
mapping(rust_span="send_provider_request"),
|
||||
mapping(rust_span="http_request", python_frame=r"AsyncHTTPHandler\.post$|HTTPHandler\.post$"),
|
||||
)
|
||||
|
||||
|
|
@ -48,7 +50,7 @@ AZURE_COMMON_MAPPINGS: Final = (
|
|||
r"|MistralOCRConfig\.transform_ocr_request$"
|
||||
),
|
||||
),
|
||||
COMMON_MAPPINGS[-1],
|
||||
*COMMON_MAPPINGS[-2:],
|
||||
)
|
||||
AZURE_SYNC_MAPPINGS: Final = (
|
||||
*AZURE_COMMON_MAPPINGS,
|
||||
|
|
@ -201,7 +203,7 @@ VERTEX_COMMON_MAPPINGS: Final = (
|
|||
r"|MistralOCRConfig\.transform_ocr_request$"
|
||||
),
|
||||
),
|
||||
COMMON_MAPPINGS[-1],
|
||||
*COMMON_MAPPINGS[-2:],
|
||||
)
|
||||
VERTEX_SYNC_MAPPINGS: Final = (
|
||||
*VERTEX_COMMON_MAPPINGS,
|
||||
|
|
@ -228,6 +230,8 @@ DEEPSEEK_COMMON_MAPPINGS: Final = (
|
|||
rust_span="transform_ocr_request",
|
||||
python_frame=r"VertexAIDeepSeekOCRConfig\.transform_ocr_request$",
|
||||
),
|
||||
# Rust wraps the HTTP attempt to invoke observers; Python has no matching frame.
|
||||
mapping(rust_span="send_provider_request"),
|
||||
mapping(rust_span="http_request", python_frame=r"AsyncHTTPHandler\.post$|HTTPHandler\.post$"),
|
||||
mapping(
|
||||
rust_span="transform_ocr_response",
|
||||
|
|
@ -264,6 +268,8 @@ DOCUMENT_INTELLIGENCE_COMMON_MAPPINGS: Final = (
|
|||
mapping(
|
||||
rust_span="transform_ocr_request", python_frame=r"AzureDocumentIntelligenceOCRConfig\.transform_ocr_request$"
|
||||
),
|
||||
# Rust wraps the HTTP attempt to invoke observers; Python has no matching frame.
|
||||
mapping(rust_span="send_provider_request"),
|
||||
mapping(rust_span="http_request", python_frame=r"AsyncHTTPHandler\.post$|HTTPHandler\.post$"),
|
||||
mapping(
|
||||
rust_span="poll_document_intelligence",
|
||||
|
|
|
|||
|
|
@ -23,7 +23,7 @@ MAPPINGS: Final = (
|
|||
mapping(rust_span="transform_transcription_request"),
|
||||
mapping(
|
||||
rust_span="execute_audio_transcription_provider_call",
|
||||
python_frame=r"BedrockAudioTranscriptionRustDispatch\.(?:async_)?audio_transcriptions$",
|
||||
python_frame=r"rust_bridge/request\.py:\d+ call_native$",
|
||||
),
|
||||
mapping(rust_span="transform_transcription_response"),
|
||||
mapping(rust_span="http_request"),
|
||||
|
|
|
|||
|
|
@ -85,7 +85,12 @@ class ExplodingAsyncMessages:
|
|||
self.calls = 0
|
||||
|
||||
async def __call__(
|
||||
self, request: NativeMessagesRequest, *, options: object, context: NativeRequestContext
|
||||
self,
|
||||
request: NativeMessagesRequest,
|
||||
*,
|
||||
options: object,
|
||||
context: NativeRequestContext,
|
||||
callback_adapter: object | None = None,
|
||||
) -> dict[str, object]:
|
||||
self.calls += 1
|
||||
raise AssertionError("bridge must not be called")
|
||||
|
|
@ -96,7 +101,12 @@ class RaisingAsyncMessages:
|
|||
self.calls = 0
|
||||
|
||||
async def __call__(
|
||||
self, request: NativeMessagesRequest, *, options: object, context: NativeRequestContext
|
||||
self,
|
||||
request: NativeMessagesRequest,
|
||||
*,
|
||||
options: object,
|
||||
context: NativeRequestContext,
|
||||
callback_adapter: object | None = None,
|
||||
) -> dict[str, object]:
|
||||
self.calls += 1
|
||||
raise RuntimeError("upstream request failed with status 400: bad request")
|
||||
|
|
|
|||
201
tests/test_litellm/ocr/callback_support.py
Normal file
201
tests/test_litellm/ocr/callback_support.py
Normal file
|
|
@ -0,0 +1,201 @@
|
|||
import asyncio
|
||||
import contextvars
|
||||
import inspect
|
||||
import json
|
||||
import threading
|
||||
from collections.abc import Generator, Mapping
|
||||
from contextlib import contextmanager
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from pathlib import Path
|
||||
from typing import Final, TypedDict
|
||||
|
||||
from typing_extensions import ReadOnly
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
|
||||
MODEL: Final = "mistral/mistral-ocr-4-1"
|
||||
DOCUMENT: Final = {"type": "document_url", "document_url": "https://example.com/document.pdf"}
|
||||
RESPONSE: Final = (
|
||||
'{"pages":[{"index":0,"markdown":"callback-test"}],"model":"mistral-ocr-4-1","usage_info":{"pages_processed":1}}'
|
||||
)
|
||||
REQUEST_CONTEXT: Final = contextvars.ContextVar("ocr_callback_test_context", default="unset")
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class CallbackEvent:
|
||||
name: str
|
||||
model: object
|
||||
call_id: object
|
||||
original_response: object
|
||||
input: object
|
||||
additional_args: str
|
||||
response_type: str
|
||||
start_time: datetime | None
|
||||
end_time: datetime | None
|
||||
context: str
|
||||
native_provider_hook: bool
|
||||
|
||||
|
||||
class CallbackRecorder(CustomLogger):
|
||||
def __init__(self, asynchronous: bool, raises: str = "", name: str = "recorder", expected_calls: int = 1) -> None:
|
||||
super().__init__() # pyright: ignore[reportUnknownMemberType] # CustomLogger exposes untyped keyword arguments
|
||||
self.asynchronous = asynchronous
|
||||
self.name = name
|
||||
self.expected_calls = expected_calls
|
||||
self.raises = raises
|
||||
self.events: tuple[CallbackEvent, ...] = ()
|
||||
self.done = threading.Event()
|
||||
self.lock = threading.Lock()
|
||||
|
||||
def record(
|
||||
self,
|
||||
name: str,
|
||||
kwargs: Mapping[str, object],
|
||||
response: object = None,
|
||||
start_time: datetime | None = None,
|
||||
end_time: datetime | None = None,
|
||||
) -> None:
|
||||
event: Final = CallbackEvent(
|
||||
name,
|
||||
kwargs.get("model"),
|
||||
kwargs.get("litellm_call_id"),
|
||||
kwargs.get("original_response"),
|
||||
kwargs.get("input"),
|
||||
json.dumps(kwargs.get("additional_args"), sort_keys=True, default=str),
|
||||
type(response).__name__,
|
||||
start_time,
|
||||
end_time,
|
||||
REQUEST_CONTEXT.get(),
|
||||
any(frame.filename == "litellm_callbacks.py" for frame in inspect.stack(context=0)),
|
||||
)
|
||||
with self.lock:
|
||||
self.events += (event,)
|
||||
terminals: Final = tuple(
|
||||
item.name for item in self.events if "success" in item.name or "failure" in item.name
|
||||
)
|
||||
if len(terminals) >= self.expected_calls * (2 if self.asynchronous and "failure" in name else 1):
|
||||
self.done.set()
|
||||
if name == self.raises:
|
||||
raise ValueError("intentional observer failure")
|
||||
|
||||
def log_pre_api_call(self, model: str, messages: object, kwargs: Mapping[str, object]) -> None:
|
||||
assert model == kwargs.get("model")
|
||||
assert isinstance(messages, list)
|
||||
self.record("pre", kwargs)
|
||||
|
||||
def log_post_api_call(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime | None
|
||||
) -> None:
|
||||
self.record("post", kwargs, response_obj, start_time, end_time)
|
||||
|
||||
def log_success_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
|
||||
) -> None:
|
||||
self.record("success", kwargs, response_obj, start_time, end_time)
|
||||
|
||||
def log_failure_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
|
||||
) -> None:
|
||||
self.record("failure", kwargs, response_obj, start_time, end_time)
|
||||
|
||||
async def async_log_success_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
|
||||
) -> None:
|
||||
self.record("async_success", kwargs, response_obj, start_time, end_time)
|
||||
|
||||
async def async_log_failure_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
|
||||
) -> None:
|
||||
self.record("async_failure", kwargs, response_obj, start_time, end_time)
|
||||
|
||||
async def wait(self) -> tuple[CallbackEvent, ...]:
|
||||
assert await asyncio.to_thread(self.done.wait, 5), tuple(event.name for event in self.events)
|
||||
return self.events
|
||||
|
||||
|
||||
class OcrUpstream(ThreadingHTTPServer):
|
||||
daemon_threads = True
|
||||
|
||||
def __init__(self, status: int, body: str, stall: bool = False) -> None:
|
||||
super().__init__(("127.0.0.1", 0), OcrHandler)
|
||||
self.status = status
|
||||
self.body = body.encode()
|
||||
self.stall = stall
|
||||
self.release = threading.Event()
|
||||
self.started = threading.Event()
|
||||
self.requests: tuple[tuple[str, str], ...] = ()
|
||||
self.lock = threading.Lock()
|
||||
|
||||
@property
|
||||
def api_base(self) -> str:
|
||||
return f"http://127.0.0.1:{self.server_port}/v1"
|
||||
|
||||
|
||||
class OcrHandler(BaseHTTPRequestHandler):
|
||||
def do_POST(self) -> None:
|
||||
assert isinstance(self.server, OcrUpstream)
|
||||
body: Final = self.rfile.read(int(self.headers["content-length"])).decode()
|
||||
with self.server.lock:
|
||||
self.server.requests += ((self.path, body),)
|
||||
self.server.started.set()
|
||||
if self.server.stall:
|
||||
self.server.release.wait(5)
|
||||
return
|
||||
self.send_response(self.server.status)
|
||||
self.send_header("Content-Type", "application/json")
|
||||
self.send_header("Content-Length", str(len(self.server.body)))
|
||||
self.end_headers()
|
||||
self.wfile.write(self.server.body)
|
||||
|
||||
def log_message(self, format: str, *args: object) -> None:
|
||||
pass
|
||||
|
||||
|
||||
@contextmanager
|
||||
def ocr_upstream(status: int = 200, body: str = RESPONSE, stall: bool = False) -> Generator[OcrUpstream, None, None]:
|
||||
with OcrUpstream(status, body, stall) as server:
|
||||
thread: Final = threading.Thread(target=server.serve_forever, daemon=True)
|
||||
thread.start()
|
||||
try:
|
||||
yield server
|
||||
finally:
|
||||
server.release.set()
|
||||
server.shutdown()
|
||||
thread.join(timeout=5)
|
||||
assert not thread.is_alive()
|
||||
|
||||
|
||||
def assert_provider_request(server: OcrUpstream) -> None:
|
||||
assert len(server.requests) == 1
|
||||
path, body = server.requests[0]
|
||||
assert path == "/v1/ocr"
|
||||
assert json.loads(body) == {"model": "mistral-ocr-4-1", "document": DOCUMENT}
|
||||
|
||||
|
||||
class OcrArguments(TypedDict):
|
||||
model: ReadOnly[str]
|
||||
document: ReadOnly[dict[str, str]]
|
||||
api_key: ReadOnly[str]
|
||||
api_base: ReadOnly[str]
|
||||
callbacks: ReadOnly[list[CustomLogger]]
|
||||
num_retries: ReadOnly[int]
|
||||
timeout: ReadOnly[float]
|
||||
|
||||
|
||||
def verify_installed_package() -> None:
|
||||
root: Final = Path(litellm.__file__).resolve()
|
||||
assert "site-packages" in root.parts, f"SDK must come from the installed wheel: {root}"
|
||||
|
||||
|
||||
async def call_ocr(arguments: OcrArguments, asynchronous: bool) -> tuple[OCRResponse | None, Exception | None]:
|
||||
try:
|
||||
response: Final = (
|
||||
await litellm.aocr(**arguments) if asynchronous else await asyncio.to_thread(litellm.ocr, **arguments)
|
||||
)
|
||||
return OCRResponse.model_validate(response), None
|
||||
except Exception as error:
|
||||
return None, error
|
||||
66
tests/test_litellm/ocr/live_callback_smoke.py
Normal file
66
tests/test_litellm/ocr/live_callback_smoke.py
Normal file
|
|
@ -0,0 +1,66 @@
|
|||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
from typing import Final
|
||||
from unittest.mock import patch
|
||||
|
||||
from callback_support import CallbackRecorder, OcrArguments, call_ocr, verify_installed_package
|
||||
|
||||
import litellm
|
||||
from litellm.rust_bridge import get_native_bridge
|
||||
|
||||
|
||||
async def smoke(*, baseline: bool) -> None:
|
||||
verify_installed_package()
|
||||
litellm.rust(True)
|
||||
for asynchronous, rejected in ((False, False), (True, False), (True, True)):
|
||||
await smoke_case(asynchronous, rejected, baseline=baseline)
|
||||
|
||||
|
||||
async def smoke_case(asynchronous: bool, rejected: bool, *, baseline: bool) -> None:
|
||||
recorder: Final = CallbackRecorder(asynchronous, name=f"live-{asynchronous}-{rejected}")
|
||||
arguments: Final[OcrArguments] = {
|
||||
"model": "mistral/mistral-ocr-latest",
|
||||
"document": {
|
||||
"type": "document_url",
|
||||
"document_url": "https://www.w3.org/WAI/ER/tests/xhtml/testfiles/resources/pdf/dummy.pdf",
|
||||
},
|
||||
"api_key": "invalid-smoke-key" if rejected else os.environ["MISTRAL_API_KEY"],
|
||||
"callbacks": [recorder],
|
||||
"api_base": "https://api.mistral.ai/v1",
|
||||
"num_retries": 0,
|
||||
"timeout": 60,
|
||||
}
|
||||
response, error = await call_ocr(arguments, asynchronous)
|
||||
assert (error is not None) is rejected
|
||||
if error is not None and not baseline:
|
||||
assert isinstance(error, litellm.AuthenticationError) and error.status_code == 401
|
||||
if response is not None:
|
||||
assert len(response.pages) == 1
|
||||
outcome: Final = (
|
||||
f"pages={len(response.pages)}"
|
||||
if response is not None
|
||||
else f"{type(error).__name__} status={getattr(error, 'status_code', None)}"
|
||||
)
|
||||
events: Final = await recorder.wait()
|
||||
names: Final = tuple(event.name for event in events)
|
||||
if rejected:
|
||||
assert names[0] == "pre" and sorted(names[1:]) == ["async_failure", "failure"]
|
||||
else:
|
||||
assert names == ("pre", *(() if baseline else ("post",)), "async_success" if asynchronous else "success")
|
||||
if not baseline:
|
||||
assert all(event.native_provider_hook for event in events if event.name in ("pre", "post"))
|
||||
sys.stdout.write(f"async={asynchronous} rejected={rejected} {outcome} events={names}\n")
|
||||
sys.stdout.flush()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
if "--run-live" not in sys.argv:
|
||||
raise SystemExit("Pass --run-live and set MISTRAL_API_KEY to make real, billable provider calls")
|
||||
if "--baseline" in sys.argv:
|
||||
asyncio.run(smoke(baseline=True))
|
||||
else:
|
||||
native: Final = get_native_bridge()
|
||||
assert native is not None
|
||||
with patch.object(native, "ready_endpoints", {"ocr": frozenset({"callbacks"})}, create=True):
|
||||
asyncio.run(smoke(baseline=False))
|
||||
280
tests/test_litellm/ocr/sdk_callback_contract.py
Normal file
280
tests/test_litellm/ocr/sdk_callback_contract.py
Normal file
|
|
@ -0,0 +1,280 @@
|
|||
import asyncio
|
||||
import json
|
||||
import sys
|
||||
from collections import Counter
|
||||
from collections.abc import Mapping
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from typing import Final
|
||||
from unittest.mock import patch
|
||||
|
||||
from callback_support import (
|
||||
DOCUMENT,
|
||||
MODEL,
|
||||
REQUEST_CONTEXT,
|
||||
RESPONSE,
|
||||
CallbackRecorder,
|
||||
OcrArguments,
|
||||
assert_provider_request,
|
||||
call_ocr,
|
||||
ocr_upstream,
|
||||
verify_installed_package,
|
||||
)
|
||||
from pydantic import TypeAdapter
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
from litellm.rust_bridge import get_native_bridge
|
||||
from litellm.rust_bridge.bindings import NativeBinding
|
||||
from litellm.rust_bridge.ocr import supports_callback_adapter
|
||||
from litellm.rust_bridge.callback_adapters import PreCallArguments
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class Outcome:
|
||||
response: str | None
|
||||
error_type: str | None
|
||||
status: int | None
|
||||
events: tuple[str, ...]
|
||||
response_types: tuple[str, ...]
|
||||
provider_body: str
|
||||
|
||||
|
||||
async def exercise(asynchronous: bool, rust: bool, case: str, *, native_expected: bool | None = None) -> Outcome:
|
||||
litellm.rust(rust)
|
||||
litellm.logging_callback_manager._reset_all_callbacks() # pyright: ignore[reportPrivateUsage] # isolate SDK cases in this dedicated test process
|
||||
native: Final = rust if native_expected is None else native_expected
|
||||
if native:
|
||||
assert supports_callback_adapter(asynchronous=asynchronous), "native callback bridge must be installed"
|
||||
status: Final = int(case) if case.isdigit() else 500 if case in ("raise_failure", "raise_async_failure") else 200
|
||||
body: Final = "invalid-json" if case == "malformed" else RESPONSE if status == 200 else '{"error":"rejected"}'
|
||||
successful: Final = status == 200 and case not in ("malformed", "timeout")
|
||||
raises: Final = case.removeprefix("raise_") if case.startswith("raise_") else ""
|
||||
recorder: Final = CallbackRecorder(asynchronous, raises=raises, name=f"first-{asynchronous}-{rust}-{case}")
|
||||
follower: Final = CallbackRecorder(asynchronous, name=f"second-{asynchronous}-{rust}-{case}")
|
||||
context: Final = f"request-{asynchronous}-{rust}-{case}"
|
||||
token: Final = REQUEST_CONTEXT.set(context)
|
||||
try:
|
||||
with ocr_upstream(status, body, stall=case == "timeout") as upstream:
|
||||
params: Final[OcrArguments] = {
|
||||
"model": MODEL,
|
||||
"document": DOCUMENT,
|
||||
"api_key": "test-key",
|
||||
"api_base": upstream.api_base,
|
||||
"callbacks": [recorder, follower],
|
||||
"num_retries": 0,
|
||||
"timeout": 0.2 if case == "timeout" else 5,
|
||||
}
|
||||
result, error = await call_ocr(params, asynchronous)
|
||||
assert (error is None) is successful
|
||||
if error is not None:
|
||||
if case.isdigit():
|
||||
assert getattr(error, "status_code", None) == status
|
||||
if case == "timeout":
|
||||
assert isinstance(error, litellm.Timeout)
|
||||
if result is not None:
|
||||
assert result.pages[0].markdown == "callback-test"
|
||||
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), 5)
|
||||
events: Final = tuple(event for event in await recorder.wait() if native or event.name != "post")
|
||||
following: Final = tuple(event for event in await follower.wait() if native or event.name != "post")
|
||||
names: Final = tuple(event.name for event in events)
|
||||
prefix: Final = ("pre", "post") if native and status == 200 and case != "timeout" else ("pre",)
|
||||
assert names[: len(prefix)] == prefix, (asynchronous, rust, case, native, names, prefix)
|
||||
terminal: Final = "success" if successful else "failure"
|
||||
expected: Final = (
|
||||
("async_success",)
|
||||
if asynchronous and successful
|
||||
else ((terminal, f"async_{terminal}") if asynchronous else (terminal,))
|
||||
)
|
||||
assert sorted(names[len(prefix) :]) == sorted(expected), (asynchronous, rust, case, native, names, expected)
|
||||
assert sorted(event.name for event in following) == sorted(names)
|
||||
assert len({event.call_id for event in events}) == 1
|
||||
assert isinstance(events[0].call_id, str) and events[0].call_id
|
||||
assert all(event.model == "mistral-ocr-4-1" for event in events)
|
||||
assert events[0].input == "OCR document processing"
|
||||
pre_arguments: Final = TypeAdapter(PreCallArguments).validate_json(events[0].additional_args)
|
||||
assert pre_arguments["complete_input_dict"] == {"model": "mistral-ocr-4-1", "document": DOCUMENT}
|
||||
assert pre_arguments["api_base"] == upstream.api_base + "/ocr"
|
||||
assert {key.lower(): value for key, value in pre_arguments["headers"].items()}[
|
||||
"authorization"
|
||||
] == "Bearer test-key"
|
||||
for event in events[: len(prefix)]:
|
||||
assert event.context == context
|
||||
assert event.native_provider_hook is native
|
||||
if "post" in prefix:
|
||||
assert events[1].original_response == body
|
||||
assert events[1].response_type == "NoneType"
|
||||
assert events[1].start_time is not None and events[1].end_time is None
|
||||
for event in events[len(prefix) :]:
|
||||
assert event.start_time is not None and event.end_time is not None
|
||||
assert event.end_time >= event.start_time
|
||||
assert_provider_request(upstream)
|
||||
return Outcome(
|
||||
result.model_dump_json() if result is not None else None,
|
||||
type(error).__name__ if error is not None else None,
|
||||
getattr(error, "status_code", None),
|
||||
tuple(sorted(event.name for event in events if event.name != "post")),
|
||||
tuple(sorted(event.response_type for event in events if event.name != "post")),
|
||||
json.dumps(json.loads(upstream.requests[0][1]), sort_keys=True),
|
||||
)
|
||||
finally:
|
||||
REQUEST_CONTEXT.reset(token)
|
||||
|
||||
|
||||
async def verify_parity() -> None:
|
||||
for asynchronous in (False, True):
|
||||
for case in (
|
||||
"success",
|
||||
"malformed",
|
||||
"401",
|
||||
"429",
|
||||
"500",
|
||||
"timeout",
|
||||
"raise_pre",
|
||||
"raise_post",
|
||||
"raise_success",
|
||||
"raise_async_success",
|
||||
"raise_failure",
|
||||
"raise_async_failure",
|
||||
):
|
||||
await verify_case(asynchronous, case)
|
||||
await verify_concurrency()
|
||||
for rust in (False, True):
|
||||
await verify_delayed_terminal(rust)
|
||||
for rust in (False, True):
|
||||
await verify_without_loggers(rust)
|
||||
|
||||
|
||||
async def verify_case(asynchronous: bool, case: str) -> None:
|
||||
python: Final = await exercise(asynchronous, False, case)
|
||||
native: Final = await exercise(asynchronous, True, case)
|
||||
assert python == native, (asynchronous, case, python, native)
|
||||
sys.stdout.write(f"PASS async={asynchronous} case={case}\n")
|
||||
sys.stdout.flush()
|
||||
|
||||
|
||||
async def verify_concurrency() -> None:
|
||||
litellm.rust(True)
|
||||
litellm.logging_callback_manager._reset_all_callbacks() # pyright: ignore[reportPrivateUsage] # isolate the concurrent SDK scenario
|
||||
recorder: Final = CallbackRecorder(True, name="concurrent", expected_calls=32)
|
||||
with ocr_upstream() as upstream:
|
||||
|
||||
async def one(index: int) -> None:
|
||||
token: Final = REQUEST_CONTEXT.set(f"concurrent-{index}")
|
||||
try:
|
||||
result: Final = await litellm.aocr(
|
||||
model=MODEL,
|
||||
document=DOCUMENT,
|
||||
api_key="test-key",
|
||||
api_base=upstream.api_base,
|
||||
callbacks=[recorder],
|
||||
num_retries=0,
|
||||
timeout=10,
|
||||
)
|
||||
assert result.pages[0].markdown == "callback-test"
|
||||
finally:
|
||||
REQUEST_CONTEXT.reset(token)
|
||||
|
||||
await asyncio.gather(*(one(index) for index in range(32)))
|
||||
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), 5)
|
||||
events: Final = await recorder.wait()
|
||||
assert len(upstream.requests) == 32
|
||||
assert Counter(event.name for event in events) == {"pre": 32, "post": 32, "async_success": 32}
|
||||
assert len({event.call_id for event in events}) == 32
|
||||
for index in range(32):
|
||||
assert tuple(event.name for event in events if event.context == f"concurrent-{index}") == (
|
||||
"pre",
|
||||
"post",
|
||||
"async_success",
|
||||
)
|
||||
assert all(event.native_provider_hook for event in events if event.name in ("pre", "post"))
|
||||
|
||||
|
||||
class DelayedRecorder(CallbackRecorder):
|
||||
def __init__(self) -> None:
|
||||
super().__init__(True, name="delayed")
|
||||
self.started = asyncio.Event()
|
||||
self.release = asyncio.Event()
|
||||
|
||||
async def async_log_success_event(
|
||||
self, kwargs: Mapping[str, object], response_obj: object, start_time: datetime, end_time: datetime
|
||||
) -> None:
|
||||
self.started.set()
|
||||
await self.release.wait()
|
||||
await super().async_log_success_event(kwargs, response_obj, start_time, end_time)
|
||||
|
||||
|
||||
async def verify_delayed_terminal(rust: bool) -> None:
|
||||
litellm.rust(rust)
|
||||
litellm.logging_callback_manager._reset_all_callbacks() # pyright: ignore[reportPrivateUsage] # isolate delayed delivery
|
||||
recorder: Final = DelayedRecorder()
|
||||
with ocr_upstream() as upstream:
|
||||
result: Final = await asyncio.wait_for(
|
||||
litellm.aocr(
|
||||
model=MODEL,
|
||||
document=DOCUMENT,
|
||||
api_key="test-key",
|
||||
api_base=upstream.api_base,
|
||||
callbacks=[recorder],
|
||||
num_retries=0,
|
||||
),
|
||||
5,
|
||||
)
|
||||
assert result.pages[0].markdown == "callback-test"
|
||||
await asyncio.wait_for(recorder.started.wait(), 5)
|
||||
prefix: Final = ("pre", "post") if rust else ("pre",)
|
||||
assert tuple(event.name for event in recorder.events if rust or event.name != "post") == prefix
|
||||
recorder.release.set()
|
||||
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), 5)
|
||||
events: Final = await recorder.wait()
|
||||
assert tuple(event.name for event in events if rust or event.name != "post") == (*prefix, "async_success"), (
|
||||
events
|
||||
)
|
||||
|
||||
|
||||
async def verify_without_loggers(rust: bool) -> None:
|
||||
litellm.rust(rust)
|
||||
litellm.logging_callback_manager._reset_all_callbacks() # pyright: ignore[reportPrivateUsage] # exercise the no-logger configuration
|
||||
with ocr_upstream() as upstream:
|
||||
result: Final = await litellm.aocr(
|
||||
model=MODEL, document=DOCUMENT, api_key="test-key", api_base=upstream.api_base, num_retries=0
|
||||
)
|
||||
assert result.pages[0].markdown == "callback-test"
|
||||
assert_provider_request(upstream)
|
||||
await asyncio.wait_for(GLOBAL_LOGGING_WORKER.flush(), 5)
|
||||
|
||||
|
||||
def assert_native_unavailable() -> None:
|
||||
assert not supports_callback_adapter()
|
||||
for binding in (
|
||||
NativeBinding(lambda native: native.ocr, route="ocr"),
|
||||
NativeBinding(lambda native: native.chat_completions, route="chat_completions"),
|
||||
NativeBinding(lambda native: native.messages, route="messages"),
|
||||
NativeBinding(lambda native: native.transcription, route="transcription"),
|
||||
NativeBinding(lambda native: native.ResponsesWebSocketConnection, route="responses_websocket"),
|
||||
):
|
||||
assert binding.load() is None
|
||||
|
||||
|
||||
async def verify_unavailable() -> None:
|
||||
assert_native_unavailable()
|
||||
for asynchronous in (False, True):
|
||||
for case in ("success", "401"):
|
||||
await exercise(asynchronous, True, case, native_expected=False)
|
||||
|
||||
|
||||
async def verify_foundation() -> None:
|
||||
assert_native_unavailable()
|
||||
native: Final = get_native_bridge()
|
||||
assert native is not None
|
||||
with patch.object(native, "ready_endpoints", {"ocr": frozenset({"callbacks"})}, create=True):
|
||||
await verify_parity()
|
||||
assert_native_unavailable()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
if "--installed" in sys.argv:
|
||||
verify_installed_package()
|
||||
asyncio.run(
|
||||
verify_unavailable() if "--without-native" in sys.argv or "--unready" in sys.argv else verify_foundation()
|
||||
)
|
||||
|
|
@ -11,6 +11,7 @@ import pytest
|
|||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.rust_bridge.callback_adapters import ProviderLoggingAdapter
|
||||
from litellm.rust_bridge import configuration
|
||||
from litellm.rust_bridge.request import (
|
||||
NativeOCRRequest,
|
||||
|
|
@ -52,6 +53,7 @@ class RecordingBridge:
|
|||
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[dict[str, object]] = []
|
||||
self.callback_adapter: object | None = None
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
|
|
@ -59,7 +61,9 @@ class RecordingBridge:
|
|||
*,
|
||||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
callback_adapter: object | None = None,
|
||||
) -> dict[str, object]:
|
||||
self.callback_adapter = callback_adapter
|
||||
self.calls.append(
|
||||
{
|
||||
"model": request.model,
|
||||
|
|
@ -81,6 +85,7 @@ class RecordingAsyncBridge:
|
|||
|
||||
def __init__(self) -> None:
|
||||
self.calls: list[dict[str, object]] = []
|
||||
self.callback_adapter: object | None = None
|
||||
|
||||
async def __call__(
|
||||
self,
|
||||
|
|
@ -88,7 +93,9 @@ class RecordingAsyncBridge:
|
|||
*,
|
||||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
callback_adapter: object | None = None,
|
||||
) -> dict[str, object]:
|
||||
self.callback_adapter = callback_adapter
|
||||
self.calls.append(
|
||||
{
|
||||
"model": request.model,
|
||||
|
|
@ -112,6 +119,7 @@ class RaisingBridge:
|
|||
*,
|
||||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
callback_adapter: object | None = None,
|
||||
) -> dict[str, object]:
|
||||
raise RuntimeError("bridge failed")
|
||||
|
||||
|
|
@ -123,6 +131,7 @@ class RaisingAsyncBridge:
|
|||
*,
|
||||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
callback_adapter: object | None = None,
|
||||
) -> dict[str, object]:
|
||||
raise RuntimeError("bridge failed")
|
||||
|
||||
|
|
@ -394,6 +403,7 @@ def test_load_rust_ocr_uses_compiled_extension(monkeypatch):
|
|||
fake_module = types.ModuleType("litellm.rust_bridge._native")
|
||||
fake_module.ocr = lambda **kwargs: dict(FAKE_OCR_RESPONSE) # type: ignore[attr-defined]
|
||||
fake_module.aocr = lambda **kwargs: dict(FAKE_OCR_RESPONSE) # type: ignore[attr-defined]
|
||||
fake_module.ready_endpoints = {"ocr": {"callbacks"}} # type: ignore[attr-defined]
|
||||
monkeypatch.setattr(
|
||||
importlib.import_module("litellm.rust_bridge.bindings"),
|
||||
"get_native_bridge",
|
||||
|
|
@ -686,7 +696,7 @@ def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint():
|
|||
assert bridge.calls[0]["api_base"] == "https://document-intelligence.example.com"
|
||||
|
||||
|
||||
def test_run_rust_ocr_runs_pre_call_logging():
|
||||
def test_run_rust_ocr_passes_provider_logging_adapter():
|
||||
logging_obj = RecordingLogging()
|
||||
bridge = RecordingBridge()
|
||||
litellm.rust(True)
|
||||
|
|
@ -704,17 +714,11 @@ def test_run_rust_ocr_runs_pre_call_logging():
|
|||
resolve_api_key=lambda _name: None,
|
||||
)
|
||||
|
||||
assert logging_obj.pre_call_kwargs is not None
|
||||
assert logging_obj.pre_call_kwargs["input"] == "OCR document processing"
|
||||
additional_args = logging_obj.pre_call_kwargs["additional_args"]
|
||||
complete_input = additional_args["complete_input_dict"]
|
||||
assert complete_input["document"] == DOCUMENT
|
||||
assert complete_input["include_image_base64"] is True
|
||||
assert additional_args["api_base"] == "https://api.mistral.ai/v1/ocr"
|
||||
assert additional_args["headers"] == {
|
||||
"Authorization": "Bearer sk-test",
|
||||
"x-trace-id": "trace-1",
|
||||
}
|
||||
adapter = bridge.callback_adapter
|
||||
assert isinstance(adapter, ProviderLoggingAdapter)
|
||||
assert adapter.logging_obj is logging_obj
|
||||
assert adapter.input == "OCR document processing"
|
||||
assert adapter.api_key == "sk-test"
|
||||
|
||||
|
||||
def test_ocr_routes_azure_ai_to_rust_when_enabled(fake_bridge):
|
||||
|
|
|
|||
|
|
@ -1,8 +1,11 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from typing import cast
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.rust_bridge import configuration, responses_websocket
|
||||
from litellm.rust_bridge.callbacks import SessionCallbackHandle
|
||||
from litellm.rust_bridge.request import NativeRequestContext, NativeResponsesWebSocketRequest
|
||||
|
||||
|
||||
|
|
@ -34,6 +37,7 @@ class _FakeNativeBridge:
|
|||
*,
|
||||
options: object,
|
||||
context: NativeRequestContext,
|
||||
callback_adapter: object | None = None,
|
||||
) -> _FakeNativeConnection:
|
||||
return _FakeNativeConnection()
|
||||
|
||||
|
|
@ -97,6 +101,38 @@ async def test_enabled_bridge_connects_and_adapts_socket(
|
|||
await connection.close()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_connection_forwards_session_callback_adapter() -> None:
|
||||
configuration.rust(True)
|
||||
received: list[object] = []
|
||||
callback_adapter = cast(SessionCallbackHandle, object())
|
||||
|
||||
class Native:
|
||||
@classmethod
|
||||
async def connect(
|
||||
cls,
|
||||
request: NativeResponsesWebSocketRequest,
|
||||
*,
|
||||
options: object,
|
||||
context: NativeRequestContext,
|
||||
callback_adapter: object | None = None,
|
||||
) -> _FakeNativeConnection:
|
||||
received.append(callback_adapter)
|
||||
return _FakeNativeConnection()
|
||||
|
||||
responses_websocket.set_rust_responses_websocket(connection=Native)
|
||||
|
||||
connection = await responses_websocket.connect(
|
||||
url="wss://example.test/responses",
|
||||
headers={},
|
||||
timeout=None,
|
||||
callback_adapter=callback_adapter,
|
||||
)
|
||||
|
||||
assert connection is not None
|
||||
assert received == [callback_adapter]
|
||||
|
||||
|
||||
class _FailingNativeBridge:
|
||||
@classmethod
|
||||
async def connect(
|
||||
|
|
@ -105,6 +141,7 @@ class _FailingNativeBridge:
|
|||
*,
|
||||
options: object,
|
||||
context: NativeRequestContext,
|
||||
callback_adapter: object | None = None,
|
||||
) -> _FakeNativeConnection:
|
||||
raise RuntimeError("connection failed")
|
||||
|
||||
|
|
@ -131,7 +168,7 @@ async def test_connection_dispatch_cleans_up_without_reconnecting(native, sessio
|
|||
|
||||
class Native:
|
||||
@classmethod
|
||||
async def connect(cls, request, *, options, context):
|
||||
async def connect(cls, request, *, options, context, callback_adapter=None):
|
||||
connections.append("native")
|
||||
assert options.custom_llm_provider == "azure"
|
||||
return native_socket
|
||||
|
|
|
|||
49
tests/test_litellm/rust_bridge/sdk_callback_wheel_test.py
Normal file
49
tests/test_litellm/rust_bridge/sdk_callback_wheel_test.py
Normal file
|
|
@ -0,0 +1,49 @@
|
|||
import importlib.util
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import litellm
|
||||
|
||||
|
||||
def main() -> None:
|
||||
assert "site-packages" in Path(litellm.__file__).resolve().parts, "install the reviewed wheel first"
|
||||
spec: Final = importlib.util.find_spec("litellm.rust_bridge._native")
|
||||
assert spec is not None and spec.origin is not None, "wheel must contain the native extension"
|
||||
native: Final = Path(spec.origin)
|
||||
hidden: Final = native.with_suffix(native.suffix + ".disabled")
|
||||
assert not hidden.exists()
|
||||
script: Final = Path(__file__).resolve().parents[1] / "ocr" / "sdk_callback_contract.py"
|
||||
environment: Final = {
|
||||
**{key: value for key, value in os.environ.items() if key not in ("PYTHONPATH", "PYTHONHOME")},
|
||||
"LITELLM_LOCAL_MODEL_COST_MAP": "True",
|
||||
"NO_PROXY": "127.0.0.1,localhost",
|
||||
}
|
||||
with tempfile.TemporaryDirectory(prefix="ocr-sdk-callbacks-") as directory:
|
||||
for flags in (("--unready",), ()):
|
||||
subprocess.run(
|
||||
[sys.executable, str(script), "--installed", *flags],
|
||||
cwd=directory,
|
||||
env=environment,
|
||||
check=True,
|
||||
timeout=180,
|
||||
)
|
||||
native.rename(hidden)
|
||||
try:
|
||||
subprocess.run(
|
||||
[sys.executable, str(script), "--installed", "--without-native"],
|
||||
cwd=directory,
|
||||
env=environment,
|
||||
check=True,
|
||||
timeout=60,
|
||||
)
|
||||
finally:
|
||||
hidden.rename(native)
|
||||
sys.stdout.write("Installed-wheel callback foundation, unready-route, and unavailable-extension checks passed\n")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
|
@ -111,3 +111,52 @@ def test_selectors_are_checked_by_type_checker(tmp_path: Path, expression: str,
|
|||
assert [(item["rule"], item["range"]["start"]["line"]) for item in diagnostics] == [
|
||||
(expected_rule, len(source.read_text().splitlines()) - 1)
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("export", (None, 3))
|
||||
def test_missing_execution_export_does_not_inspect_readiness(export):
|
||||
from types import ModuleType
|
||||
|
||||
native = ModuleType("test_native")
|
||||
lookups = []
|
||||
|
||||
def missing(name: str):
|
||||
lookups.append(name)
|
||||
raise AttributeError(name)
|
||||
|
||||
native.__getattr__ = missing
|
||||
if export is not None:
|
||||
native.messages = export
|
||||
binding = bindings.NativeBinding(lambda module: module.messages, route="messages", module_loader=lambda: native)
|
||||
assert binding.load() is None
|
||||
assert "ready_endpoints" not in lookups
|
||||
|
||||
|
||||
def test_discovery_reuses_one_module_and_does_not_cache_binding():
|
||||
from types import ModuleType
|
||||
|
||||
native = ModuleType("test_native")
|
||||
native.ready_endpoints = {"messages": frozenset({"callbacks"})}
|
||||
|
||||
def first():
|
||||
return "first"
|
||||
|
||||
def second():
|
||||
return "second"
|
||||
|
||||
native.messages = first
|
||||
loads = []
|
||||
|
||||
def load():
|
||||
loads.append(native)
|
||||
return native
|
||||
|
||||
binding = bindings.NativeBinding(lambda module: module.messages, route="messages", module_loader=load)
|
||||
assert binding.load() is first
|
||||
assert len(loads) == 1
|
||||
native.messages = second
|
||||
assert binding.load() is second
|
||||
assert len(loads) == 2
|
||||
native.ready_endpoints = {"messages": frozenset()}
|
||||
assert binding.load() is None
|
||||
assert len(loads) == 3
|
||||
|
|
|
|||
130
tests/test_litellm/rust_bridge/test_callback_adapters.py
Normal file
130
tests/test_litellm/rust_bridge/test_callback_adapters.py
Normal file
|
|
@ -0,0 +1,130 @@
|
|||
from dataclasses import dataclass, field
|
||||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.callback_adapters import ProviderLoggingAdapter, SessionCallbackAdapter
|
||||
from litellm.rust_bridge.callbacks import CallbackDecision
|
||||
|
||||
|
||||
@dataclass
|
||||
class RecordingLogging:
|
||||
model_call_details: dict[str, object] = field(default_factory=dict)
|
||||
calls: list[tuple[str, object]] = field(default_factory=list)
|
||||
|
||||
def pre_call(self, *, input: object, api_key: str | None, additional_args: object) -> None:
|
||||
self.calls.append(("pre", (input, api_key, additional_args)))
|
||||
|
||||
def post_call(self, *, original_response: str, input: object, api_key: str | None) -> None:
|
||||
self.calls.append(("post", (original_response, input, api_key)))
|
||||
|
||||
|
||||
def provider_event(**updates: object) -> dict[str, object]:
|
||||
return {
|
||||
"provider": "mistral",
|
||||
"model": "mistral-ocr-latest",
|
||||
"call_id": "call-1",
|
||||
"trace_id": "trace-1",
|
||||
"attempt": 1,
|
||||
"started_at": 10.0,
|
||||
**updates,
|
||||
}
|
||||
|
||||
|
||||
def test_provider_logging_adapter_preserves_provider_lifecycle() -> None:
|
||||
logging: Final = RecordingLogging()
|
||||
adapter: Final = ProviderLoggingAdapter(logging, "OCR document processing", "secret")
|
||||
|
||||
assert adapter.pre_call(
|
||||
provider_event(request={"model": "mistral-ocr-latest"}, api_base="https://provider.test", headers={})
|
||||
) == {"action": "unchanged"}
|
||||
assert adapter.post_call(provider_event(response={"pages": []}, status_code=200, headers={}, ended_at=11.0)) == {
|
||||
"action": "unchanged"
|
||||
}
|
||||
adapter.error(
|
||||
provider_event(
|
||||
message="retryable",
|
||||
stage="provider_response",
|
||||
committed=True,
|
||||
status_code=429,
|
||||
ended_at=11.0,
|
||||
)
|
||||
)
|
||||
assert adapter.stream_event(provider_event(event={"type": "delta"}, sequence=1)) == {"action": "unchanged"}
|
||||
adapter.stream_close(provider_event(outcome="completed", ended_at=12.0))
|
||||
|
||||
assert [name for name, _ in logging.calls] == ["pre", "post"]
|
||||
assert logging.model_call_details["provider_stream_event"] == {
|
||||
**provider_event(),
|
||||
"event": {"type": "delta"},
|
||||
"sequence": 1,
|
||||
}
|
||||
assert logging.model_call_details["provider_stream_close"] == {
|
||||
**provider_event(),
|
||||
"outcome": "completed",
|
||||
"ended_at": 12.0,
|
||||
}
|
||||
assert logging.model_call_details["provider_error"] == {
|
||||
**provider_event(),
|
||||
"message": "retryable",
|
||||
"stage": "provider_response",
|
||||
"committed": True,
|
||||
"status_code": 429,
|
||||
"ended_at": 11.0,
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class RecordingSession:
|
||||
events: list[str] = field(default_factory=list)
|
||||
|
||||
def before_connect(self, payload: object, /) -> CallbackDecision:
|
||||
self.events.append("before_connect")
|
||||
return {"action": "unchanged"}
|
||||
|
||||
def connected(self, payload: object, /) -> None:
|
||||
self.events.append("connected")
|
||||
|
||||
def before_send(self, payload: object, /) -> CallbackDecision:
|
||||
self.events.append("before_send")
|
||||
return {"action": "reject", "message": "drop frame", "status_code": None}
|
||||
|
||||
def after_receive(self, payload: object, /) -> CallbackDecision:
|
||||
self.events.append("after_receive")
|
||||
return {"action": "replace", "payload": {"type": "masked"}}
|
||||
|
||||
def response_complete(self, payload: object, /) -> None:
|
||||
self.events.append("response_complete")
|
||||
|
||||
def response_error(self, payload: object, /) -> None:
|
||||
self.events.append("response_error")
|
||||
|
||||
def error(self, payload: object, /) -> None:
|
||||
self.events.append("error")
|
||||
|
||||
def close(self, payload: object, /) -> None:
|
||||
self.events.append("close")
|
||||
|
||||
|
||||
def test_session_adapter_preserves_frame_decisions_and_order() -> None:
|
||||
callback: Final = RecordingSession()
|
||||
adapter: Final = SessionCallbackAdapter(callback)
|
||||
event: Final = {"session_id": "session-1", "call_id": "call-1", "event": {"type": "response.create"}}
|
||||
|
||||
assert adapter.before_connect(event) == {"action": "unchanged"}
|
||||
adapter.connected(event)
|
||||
assert adapter.before_send(event) == {"action": "reject", "message": "drop frame", "status_code": None}
|
||||
assert adapter.after_receive(event) == {"action": "replace", "payload": {"type": "masked"}}
|
||||
adapter.response_complete(event)
|
||||
adapter.response_error(event)
|
||||
adapter.error(event)
|
||||
adapter.close(event)
|
||||
|
||||
assert callback.events == [
|
||||
"before_connect",
|
||||
"connected",
|
||||
"before_send",
|
||||
"after_receive",
|
||||
"response_complete",
|
||||
"response_error",
|
||||
"error",
|
||||
"close",
|
||||
]
|
||||
|
|
@ -101,7 +101,7 @@ class _RecordingCall:
|
|||
self.error = error
|
||||
self.calls: list[dict] = []
|
||||
|
||||
def __call__(self, request, *, options, context):
|
||||
def __call__(self, request, *, options, context, callback_adapter=None):
|
||||
self.calls.append({"request": request, "options": options, "context": context})
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
|
|
@ -109,8 +109,14 @@ class _RecordingCall:
|
|||
|
||||
|
||||
class _RecordingAsyncCall(_RecordingCall):
|
||||
async def __call__(self, request, *, options, context):
|
||||
return _RecordingCall.__call__(self, request, options=options, context=context)
|
||||
async def __call__(self, request, *, options, context, callback_adapter=None):
|
||||
return _RecordingCall.__call__(
|
||||
self,
|
||||
request,
|
||||
options=options,
|
||||
context=context,
|
||||
callback_adapter=callback_adapter,
|
||||
)
|
||||
|
||||
|
||||
def _accepts(**overrides) -> bool:
|
||||
|
|
|
|||
47
tests/test_litellm/rust_bridge/test_loader.py
Normal file
47
tests/test_litellm/rust_bridge/test_loader.py
Normal file
|
|
@ -0,0 +1,47 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Generator
|
||||
from types import ModuleType
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.rust_bridge import loader
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def reset_loader_cache() -> Generator[None]:
|
||||
loader.reset_native_bridge_cache()
|
||||
yield
|
||||
loader.reset_native_bridge_cache()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("ready_endpoints", "expected"),
|
||||
(
|
||||
pytest.param(None, False, id="missing-registry"),
|
||||
pytest.param({"messages"}, False, id="mutable-registry"),
|
||||
pytest.param(frozenset(), False, id="unregistered"),
|
||||
pytest.param({"messages": frozenset()}, True, id="registered"),
|
||||
),
|
||||
)
|
||||
def test_native_route_requires_explicit_readiness_registry(
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
ready_endpoints: object,
|
||||
expected: bool,
|
||||
) -> None:
|
||||
native: Final = ModuleType("litellm.rust_bridge._native")
|
||||
if ready_endpoints is not None:
|
||||
native.ready_endpoints = ready_endpoints
|
||||
monkeypatch.setattr(loader, "get_native_bridge", lambda: native)
|
||||
|
||||
assert loader.native_route_ready("messages") is expected
|
||||
|
||||
|
||||
def test_native_route_requires_declared_capabilities(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
native: Final = ModuleType("litellm.rust_bridge._native")
|
||||
native.ready_endpoints = {"messages": frozenset({"callbacks"})}
|
||||
monkeypatch.setattr(loader, "get_native_bridge", lambda: native)
|
||||
|
||||
assert loader.native_route_ready("messages", frozenset({"callbacks"}))
|
||||
assert not loader.native_route_ready("messages", frozenset({"streaming_callbacks"}))
|
||||
|
|
@ -7,7 +7,7 @@ from typing import Final
|
|||
import pytest
|
||||
|
||||
from litellm.exceptions import APIError, AuthenticationError, InternalServerError, RateLimitError
|
||||
from litellm.rust_bridge import bindings, runtime
|
||||
from litellm.rust_bridge import bindings, loader, runtime
|
||||
|
||||
|
||||
class RustBridgeDeclined(Exception):
|
||||
|
|
@ -336,7 +336,11 @@ def test_native_endpoint_applies_partial_overrides_and_reset(monkeypatch: pytest
|
|||
monkeypatch.setattr(
|
||||
bindings,
|
||||
"get_native_bridge",
|
||||
lambda: SimpleNamespace(chat_completions=native_sync, achat_completions=native_async),
|
||||
lambda: SimpleNamespace(
|
||||
chat_completions=native_sync,
|
||||
achat_completions=native_async,
|
||||
ready_endpoints={"test": frozenset({"callbacks"})},
|
||||
),
|
||||
)
|
||||
endpoint: Final[runtime.EndpointDispatch[object, object]] = runtime.EndpointDispatch.native(
|
||||
route="test",
|
||||
|
|
@ -439,7 +443,9 @@ async def test_preflight_runs_after_binding_selection_before_preparation(
|
|||
assert events == (
|
||||
["load", "preflight", "prepare", "native"]
|
||||
if available and accepted
|
||||
else ["load", "preflight", "python"] if available else ["load", "python"]
|
||||
else ["load", "preflight", "python"]
|
||||
if available
|
||||
else ["load", "python"]
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -458,3 +464,50 @@ def test_preflight_failure_is_not_a_native_decline() -> None:
|
|||
error_context=context(),
|
||||
preflight=preflight,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"route",
|
||||
(
|
||||
"ocr",
|
||||
"chat_completions",
|
||||
"messages",
|
||||
"responses_websocket",
|
||||
"transcription",
|
||||
),
|
||||
)
|
||||
@pytest.mark.parametrize("capabilities", (None, frozenset(), frozenset({"streaming_callbacks"})))
|
||||
async def test_unready_routes_never_prepare_or_call_native(
|
||||
monkeypatch: pytest.MonkeyPatch, route: str, capabilities: frozenset[str] | None
|
||||
) -> None:
|
||||
def unexpected(*_args: object) -> object:
|
||||
pytest.fail("unready native route must not prepare or execute")
|
||||
|
||||
native: Final = SimpleNamespace(
|
||||
ready_endpoints={} if capabilities is None else {route: capabilities},
|
||||
chat_completions=unexpected,
|
||||
)
|
||||
monkeypatch.setattr(loader, "get_native_bridge", lambda: native)
|
||||
monkeypatch.setattr(bindings, "get_native_bridge", lambda: native)
|
||||
endpoint: Final = runtime.EndpointBinding.native(
|
||||
route=route, select=lambda native: native.chat_completions, enabled=runtime.always_enabled
|
||||
)
|
||||
arguments: Final = {
|
||||
"prepare": unexpected,
|
||||
"preflight": unexpected,
|
||||
"call": unexpected,
|
||||
"adapt": unexpected,
|
||||
"error_context": runtime.BridgeErrorContext(provider="test", model="test-model"),
|
||||
}
|
||||
assert not endpoint.can_attempt()
|
||||
assert endpoint.invoke(**arguments, fallback=lambda: "python") == "python"
|
||||
with pytest.raises(RuntimeError, match=f"native {route} endpoint is unavailable"):
|
||||
endpoint.require(**arguments)
|
||||
|
||||
async def fallback() -> str:
|
||||
return "python"
|
||||
|
||||
assert await endpoint.ainvoke(**arguments, fallback=fallback) == "python"
|
||||
with pytest.raises(RuntimeError, match=f"native {route} endpoint is unavailable"):
|
||||
await endpoint.arequire(**arguments)
|
||||
|
|
|
|||
|
|
@ -34,6 +34,7 @@ class SyncBridge:
|
|||
*,
|
||||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
callback_adapter: object | None = None,
|
||||
) -> dict[str, object]:
|
||||
self.calls.append(
|
||||
{
|
||||
|
|
@ -53,6 +54,7 @@ class AsyncBridge:
|
|||
*,
|
||||
options: NativeRequestOptions,
|
||||
context: NativeRequestContext,
|
||||
callback_adapter: object | None = None,
|
||||
) -> dict[str, object]:
|
||||
return {"text": "async"}
|
||||
|
||||
|
|
@ -135,7 +137,7 @@ async def test_dispatch_async_path_requires_bridge(monkeypatch: pytest.MonkeyPat
|
|||
|
||||
def test_bedrock_transcription_uses_rust_only_path() -> None:
|
||||
rust_bridge.configure_rust_transcription(
|
||||
transcription=lambda request, *, options, context: {"text": "rust"},
|
||||
transcription=lambda request, *, options, context, callback_adapter=None: {"text": "rust"},
|
||||
atranscription=None,
|
||||
)
|
||||
try:
|
||||
|
|
@ -152,7 +154,11 @@ def test_bedrock_transcription_uses_rust_only_path() -> None:
|
|||
@pytest.mark.asyncio
|
||||
async def test_bedrock_atranscription_uses_rust_only_path() -> None:
|
||||
async def rust_response(
|
||||
request: NativeTranscriptionRequest, *, options: object, context: NativeRequestContext
|
||||
request: NativeTranscriptionRequest,
|
||||
*,
|
||||
options: object,
|
||||
context: NativeRequestContext,
|
||||
callback_adapter: object | None = None,
|
||||
) -> dict[str, object]:
|
||||
return {"text": "rust"}
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue