feat: add rust vllm image edit

This commit is contained in:
Ishaan Jaff 2026-06-25 12:29:19 -07:00
parent bd2a1653bd
commit 62b85f23a2
No known key found for this signature in database
17 changed files with 1322 additions and 7 deletions

View file

@ -583,6 +583,7 @@ dependencies = [
name = "litellm-providers"
version = "0.1.0"
dependencies = [
"base64",
"futures-util",
"litellm-core",
"reqwest",
@ -636,6 +637,22 @@ dependencies = [
"autocfg",
]
[[package]]
name = "mime"
version = "0.3.17"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6877bb514081ee2a7ff5ef9de3281f14a4dd4bceac4c09388074a6b5df8a139a"
[[package]]
name = "mime_guess"
version = "2.0.5"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "f7c44f8e672c00fe5308fa235f821cb4198414e1c77935c1ab6948d3fd78550e"
dependencies = [
"mime",
"unicase",
]
[[package]]
name = "mio"
version = "1.2.1"
@ -918,6 +935,7 @@ dependencies = [
"base64",
"bytes",
"futures-core",
"futures-util",
"h2",
"http",
"http-body",
@ -927,6 +945,7 @@ dependencies = [
"hyper-util",
"js-sys",
"log",
"mime_guess",
"percent-encoding",
"pin-project-lite",
"quinn",
@ -1428,6 +1447,12 @@ version = "1.20.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "b6f5e870be6c3b371b77fe0ee0bafb859fa4964b4404c27de1d380043c4dda20"
[[package]]
name = "unicase"
version = "2.9.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "dbc4bc3a9f746d862c45cb89d705aa10f187bb96c76001afab07a0d35ce60142"
[[package]]
name = "unicode-ident"
version = "1.0.24"

View file

@ -16,10 +16,11 @@ litellm-core = { path = "crates/core" }
litellm-providers = { path = "crates/providers" }
pyo3 = "0.23.5"
pyo3-async-runtimes = { version = "0.23.0", features = ["tokio-runtime"] }
reqwest = { version = "0.12", default-features = false, features = ["json", "rustls-tls", "http2"] }
reqwest = { version = "0.12", default-features = false, features = ["json", "multipart", "rustls-tls", "http2"] }
serde = { version = "1.0", features = ["derive"] }
serde_json = "1.0"
thiserror = "2.0"
tokio = { version = "1", features = ["rt-multi-thread", "macros", "time"] }
tokio-tungstenite = { version = "0.24", default-features = false, features = ["connect", "rustls-tls-native-roots"] }
futures-util = { version = "0.3", default-features = false, features = ["sink", "std"] }
base64 = "0.22"

View file

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

View file

@ -0,0 +1,38 @@
use serde_json::{Map, Value};
use crate::CoreResult;
use super::types::{ImageEditInputFile, ImageEditRequestData, ImageEditResponseData};
/// Provider-specific image-edit transforms.
///
/// Implementations stay pure and non-blocking. The route layer owns async HTTP
/// I/O and multipart/JSON transport.
pub trait ImageEditProviderConfig: Send + Sync {
fn supported_image_edit_params(&self) -> &'static [&'static str];
fn map_image_edit_params(&self, non_default_params: &Map<String, Value>) -> Map<String, Value> {
let mut mapped_params = Map::new();
for (param, value) in non_default_params {
if self.supported_image_edit_params().contains(&param.as_str()) {
mapped_params.insert(param.clone(), value.clone());
}
}
mapped_params
}
fn transform_image_edit_request(
&self,
model: &str,
images: Vec<ImageEditInputFile>,
mask: Option<ImageEditInputFile>,
prompt: Option<&str>,
optional_params: Map<String, Value>,
) -> CoreResult<ImageEditRequestData>;
fn transform_image_edit_response(
&self,
model: &str,
response_json: Value,
) -> CoreResult<ImageEditResponseData>;
}

View file

@ -0,0 +1,41 @@
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ImageEditInputFile {
pub filename: String,
pub content_type: String,
pub data_base64: String,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ImageEditMultipartPart {
pub field_name: String,
pub filename: String,
pub content_type: String,
pub data_base64: String,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub enum ImageEditRequestFormat {
Multipart,
Json,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ImageEditRequestData {
pub data: Map<String, Value>,
pub files: Vec<ImageEditMultipartPart>,
pub format: ImageEditRequestFormat,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ImageEditResponseData {
pub data: Map<String, Value>,
}
impl ImageEditResponseData {
pub fn into_json(self) -> Value {
Value::Object(self.data)
}
}

View file

@ -1,4 +1,5 @@
pub mod error;
pub mod image_edit;
pub mod ocr;
pub mod providers;
pub mod realtime;

View file

@ -12,6 +12,7 @@ serde_json.workspace = true
tokio.workspace = true
tokio-tungstenite.workspace = true
futures-util.workspace = true
base64.workspace = true
[dev-dependencies]
serde_json.workspace = true

View file

@ -0,0 +1,300 @@
//! End-to-end image-edit orchestration.
//!
//! Owns vLLM image edit calls while the Python side stays a thin bridge:
//! resolve URL/auth, build multipart via pure transforms, POST it, and return
//! the OpenAI-compatible response JSON.
use std::str::FromStr;
use std::sync::OnceLock;
use std::time::Duration;
use base64::Engine;
use litellm_core::error::CoreError;
use litellm_core::image_edit::transformation::ImageEditProviderConfig;
use litellm_core::image_edit::types::{ImageEditInputFile, ImageEditRequestFormat};
use litellm_core::{CoreResult, LlmProvider};
use serde_json::{Map, Value};
use crate::vllm::image_edit::transformation as vllm;
use crate::vllm::image_edit::transformation::VLLM_IMAGE_EDIT_CONFIG;
const IMAGE_EDIT_TIMEOUT_SECS: u64 = 600;
const ERROR_BODY_MAX_CHARS: usize = 256;
fn http_client() -> &'static reqwest::Client {
static CLIENT: OnceLock<reqwest::Client> = OnceLock::new();
CLIENT.get_or_init(|| {
reqwest::Client::builder()
.timeout(Duration::from_secs(IMAGE_EDIT_TIMEOUT_SECS))
.build()
.expect("failed to build reqwest client")
})
}
fn truncate_error_body(body: &str) -> String {
if body.chars().count() <= ERROR_BODY_MAX_CHARS {
return body.to_string();
}
let truncated: String = body.chars().take(ERROR_BODY_MAX_CHARS).collect();
format!("{truncated}... (truncated)")
}
fn image_edit_config_for(provider: LlmProvider) -> Option<&'static dyn ImageEditProviderConfig> {
match provider {
LlmProvider::Vllm => Some(&VLLM_IMAGE_EDIT_CONFIG),
_ => None,
}
}
fn string_headers(extra_headers: Option<Map<String, Value>>) -> CoreResult<Vec<(String, String)>> {
extra_headers
.unwrap_or_default()
.into_iter()
.map(|(key, value)| {
value
.as_str()
.map(|value| (key.clone(), value.to_string()))
.ok_or_else(|| {
CoreError::InvalidRequest(format!(
"image_edit extra_headers.{key} must be a string, got {}",
litellm_core::error::json_type_name(&value)
))
})
})
.collect()
}
fn has_header(headers: &[(String, String)], expected: &str) -> bool {
headers
.iter()
.any(|(key, _)| key.eq_ignore_ascii_case(expected))
}
fn multipart_form(
data: Map<String, Value>,
files: Vec<litellm_core::image_edit::types::ImageEditMultipartPart>,
) -> CoreResult<reqwest::multipart::Form> {
let mut form = reqwest::multipart::Form::new();
for (key, value) in data {
let text = match value {
Value::String(value) => value,
Value::Number(value) => value.to_string(),
Value::Bool(value) => value.to_string(),
other => other.to_string(),
};
form = form.text(key, text);
}
for file in files {
let bytes = base64::engine::general_purpose::STANDARD
.decode(file.data_base64.as_bytes())
.map_err(|err| {
CoreError::InvalidRequest(format!(
"invalid base64 for multipart field {}: {err}",
file.field_name
))
})?;
let part = reqwest::multipart::Part::bytes(bytes)
.file_name(file.filename)
.mime_str(&file.content_type)
.map_err(|err| {
CoreError::InvalidRequest(format!(
"invalid content type for multipart field {}: {err}",
file.field_name
))
})?;
form = form.part(file.field_name, part);
}
Ok(form)
}
pub struct ImageEditRequest<'a> {
pub model: &'a str,
pub images: Vec<ImageEditInputFile>,
pub mask: Option<ImageEditInputFile>,
pub prompt: Option<&'a str>,
pub api_key: Option<&'a str>,
pub api_base: Option<&'a str>,
pub custom_llm_provider: &'a str,
pub extra_headers: Option<Map<String, Value>>,
pub optional_params: Map<String, Value>,
pub timeout: Option<Duration>,
}
pub async fn image_edit(request: ImageEditRequest<'_>) -> CoreResult<Value> {
let provider = LlmProvider::from_str(request.custom_llm_provider)?;
let config = image_edit_config_for(provider)
.ok_or_else(|| CoreError::InvalidProvider(provider.to_string()))?;
let api_key = vllm::resolve_api_key(request.api_key, &|key| std::env::var(key).ok());
let url = vllm::complete_url(request.api_base, &|key| std::env::var(key).ok())?;
let filtered_params = config.map_image_edit_params(&request.optional_params);
let transformed = config.transform_image_edit_request(
request.model,
request.images,
request.mask,
request.prompt,
filtered_params,
)?;
let headers = string_headers(request.extra_headers)?;
let mut request_builder = http_client().post(&url);
if let Some(api_key) = api_key {
if !has_header(&headers, "x-api-key") && !has_header(&headers, "authorization") {
request_builder = request_builder.header("x-api-key", api_key);
}
}
for (key, value) in headers {
request_builder = request_builder.header(&key, value);
}
if let Some(duration) = request.timeout {
request_builder = request_builder.timeout(duration);
}
request_builder = match transformed.format {
ImageEditRequestFormat::Multipart => {
request_builder.multipart(multipart_form(transformed.data, transformed.files)?)
}
ImageEditRequestFormat::Json => request_builder.json(&Value::Object(transformed.data)),
};
let response = request_builder
.send()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
let status = response.status();
let text = response
.text()
.await
.map_err(|err| CoreError::Network(err.to_string()))?;
if !status.is_success() {
return Err(CoreError::Http {
status: status.as_u16(),
body: truncate_error_body(&text),
});
}
let response_json: Value = serde_json::from_str(&text).map_err(|err| {
CoreError::InvalidResponse(format!("invalid image edit response JSON: {err}"))
})?;
Ok(config
.transform_image_edit_response(request.model, response_json)?
.into_json())
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
#[test]
fn image_edit_registry_supports_only_vllm() {
assert!(image_edit_config_for(LlmProvider::Vllm).is_some());
assert!(image_edit_config_for(LlmProvider::Openai).is_none());
}
#[test]
fn string_headers_rejects_non_string_values() {
let headers = json!({"x-retry-count": 3}).as_object().unwrap().clone();
let err = string_headers(Some(headers)).expect_err("non-string header rejected");
assert_eq!(
err,
CoreError::InvalidRequest(
"image_edit extra_headers.x-retry-count must be a string, got number".to_string()
)
);
}
#[test]
fn multipart_form_rejects_invalid_base64() {
let file = litellm_core::image_edit::types::ImageEditMultipartPart {
field_name: "image[]".to_string(),
filename: "image.png".to_string(),
content_type: "image/png".to_string(),
data_base64: "not base64".to_string(),
};
let err = multipart_form(Map::new(), vec![file]).expect_err("invalid base64 rejected");
assert!(err.to_string().contains("invalid base64"));
}
#[tokio::test]
async fn image_edit_posts_vllm_multipart_request() {
let listener = TcpListener::bind("127.0.0.1:0")
.await
.expect("test listener binds");
let addr = listener.local_addr().expect("listener has local addr");
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts one request");
let mut request = Vec::new();
let mut buffer = [0_u8; 4096];
loop {
let n = socket.read(&mut buffer).await.expect("reads request");
if n == 0 {
break;
}
request.extend_from_slice(&buffer[..n]);
if request.windows(5).any(|window| window == b"image") {
break;
}
}
let response_body = r#"{"created":1,"data":[{"b64_json":"ZmFrZQ=="}]}"#;
let response = format!(
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}",
response_body.len(),
response_body
);
socket
.write_all(response.as_bytes())
.await
.expect("writes response");
String::from_utf8_lossy(&request).to_string()
});
let response = image_edit(ImageEditRequest {
model: "qwen-image-edit",
images: vec![ImageEditInputFile {
filename: "image.png".to_string(),
content_type: "image/png".to_string(),
data_base64: "aW1hZ2U=".to_string(),
}],
mask: Some(ImageEditInputFile {
filename: "mask.png".to_string(),
content_type: "image/png".to_string(),
data_base64: "bWFzaw==".to_string(),
}),
prompt: Some("make it brighter"),
api_key: Some("sk-test"),
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: "vllm",
extra_headers: None,
optional_params: json!({"quality": "high"}).as_object().unwrap().clone(),
timeout: Some(Duration::from_secs(5)),
})
.await
.expect("image edit request succeeds");
assert_eq!(response["data"][0]["b64_json"], "ZmFrZQ==");
let request = server.await.expect("server task completes");
assert!(request.starts_with("POST /v1/images/edits "), "{request}");
assert!(request.contains("x-api-key: sk-test"), "{request}");
assert!(request.contains("name=\"model\""), "{request}");
assert!(request.contains("qwen-image-edit"), "{request}");
assert!(request.contains("name=\"prompt\""), "{request}");
assert!(request.contains("make it brighter"), "{request}");
assert!(request.contains("name=\"quality\""), "{request}");
assert!(request.contains("high"), "{request}");
assert!(request.contains("name=\"image[]\""), "{request}");
assert!(request.contains("filename=\"image.png\""), "{request}");
assert!(request.contains("name=\"mask\""), "{request}");
assert!(request.contains("filename=\"mask.png\""), "{request}");
}
}

View file

@ -1,4 +1,6 @@
pub mod image_edit;
pub mod mistral;
pub mod ocr;
pub mod openai;
pub mod realtime;
pub mod vllm;

View file

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

View file

@ -0,0 +1,254 @@
use litellm_core::error::{json_type_name, CoreError, CoreResult};
use litellm_core::image_edit::transformation::ImageEditProviderConfig;
use litellm_core::image_edit::types::{
ImageEditInputFile, ImageEditMultipartPart, ImageEditRequestData, ImageEditRequestFormat,
ImageEditResponseData,
};
use serde_json::{Map, Value};
const SUPPORTED_IMAGE_EDIT_PARAMS: &[&str] = &[
"background",
"input_fidelity",
"mask",
"n",
"quality",
"response_format",
"size",
"user",
"imageConfig",
];
pub const VLLM_API_BASE_ENV: &str = "VLLM_API_BASE";
pub const VLLM_API_KEY_ENV: &str = "VLLM_API_KEY";
pub const MISSING_API_BASE_MESSAGE: &str = "VLLM_API_BASE is not set. Please set the environment variable, to use VLLM's image edit endpoint.";
pub fn complete_url(
api_base: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
let base = api_base
.map(str::trim)
.filter(|base| !base.is_empty())
.map(str::to_string)
.or_else(|| env_lookup(VLLM_API_BASE_ENV).filter(|base| !base.trim().is_empty()))
.ok_or_else(|| CoreError::Auth(MISSING_API_BASE_MESSAGE.to_string()))?;
let base = base.trim_end_matches('/');
if base.ends_with("/v1") {
Ok(format!("{base}/images/edits"))
} else {
Ok(format!("{base}/v1/images/edits"))
}
}
pub fn resolve_api_key(
api_key: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Option<String> {
api_key
.map(str::trim)
.filter(|key| !key.is_empty())
.map(str::to_string)
.or_else(|| env_lookup(VLLM_API_KEY_ENV).filter(|key| !key.trim().is_empty()))
}
pub struct VllmImageEditConfig;
pub const VLLM_IMAGE_EDIT_CONFIG: VllmImageEditConfig = VllmImageEditConfig;
impl ImageEditProviderConfig for VllmImageEditConfig {
fn supported_image_edit_params(&self) -> &'static [&'static str] {
SUPPORTED_IMAGE_EDIT_PARAMS
}
fn transform_image_edit_request(
&self,
model: &str,
images: Vec<ImageEditInputFile>,
mask: Option<ImageEditInputFile>,
prompt: Option<&str>,
optional_params: Map<String, Value>,
) -> CoreResult<ImageEditRequestData> {
if images.is_empty() {
return Err(CoreError::MissingField("image"));
}
let mut data = Map::new();
data.insert("model".to_string(), Value::String(model.to_string()));
if let Some(prompt) = prompt {
data.insert("prompt".to_string(), Value::String(prompt.to_string()));
}
for (param, value) in optional_params {
if param != "mask" {
data.insert(param, value);
}
}
let mut files = Vec::new();
for image in images {
files.push(ImageEditMultipartPart {
field_name: "image[]".to_string(),
filename: image.filename,
content_type: image.content_type,
data_base64: image.data_base64,
});
}
if let Some(mask) = mask {
files.push(ImageEditMultipartPart {
field_name: "mask".to_string(),
filename: mask.filename,
content_type: mask.content_type,
data_base64: mask.data_base64,
});
}
Ok(ImageEditRequestData {
data,
files,
format: ImageEditRequestFormat::Multipart,
})
}
fn transform_image_edit_response(
&self,
_model: &str,
response_json: Value,
) -> CoreResult<ImageEditResponseData> {
let data = response_json
.as_object()
.cloned()
.ok_or_else(|| CoreError::InvalidType {
expected: "object",
actual: json_type_name(&response_json),
})?;
Ok(ImageEditResponseData { data })
}
}
pub fn supported_image_edit_params() -> &'static [&'static str] {
VLLM_IMAGE_EDIT_CONFIG.supported_image_edit_params()
}
pub fn map_image_edit_params(non_default_params: &Map<String, Value>) -> Map<String, Value> {
VLLM_IMAGE_EDIT_CONFIG.map_image_edit_params(non_default_params)
}
pub fn transform_image_edit_request(
model: &str,
images: Vec<ImageEditInputFile>,
mask: Option<ImageEditInputFile>,
prompt: Option<&str>,
optional_params: Map<String, Value>,
) -> CoreResult<ImageEditRequestData> {
VLLM_IMAGE_EDIT_CONFIG.transform_image_edit_request(
model,
images,
mask,
prompt,
optional_params,
)
}
pub fn transform_image_edit_response(
model: &str,
response_json: Value,
) -> CoreResult<ImageEditResponseData> {
VLLM_IMAGE_EDIT_CONFIG.transform_image_edit_response(model, response_json)
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn image_part() -> ImageEditInputFile {
ImageEditInputFile {
filename: "image.png".to_string(),
content_type: "image/png".to_string(),
data_base64: "aW1hZ2U=".to_string(),
}
}
#[test]
fn supported_params_match_openai_compatible_image_edit_subset() {
assert_eq!(
supported_image_edit_params(),
&[
"background",
"input_fidelity",
"mask",
"n",
"quality",
"response_format",
"size",
"user",
"imageConfig",
]
);
}
#[test]
fn map_image_edit_params_drops_unknown_params() {
let params = json!({
"quality": "high",
"unsupported_param": "value",
"size": "1024x1024"
});
let mapped = map_image_edit_params(params.as_object().unwrap());
assert_eq!(mapped.get("quality"), Some(&json!("high")));
assert_eq!(mapped.get("size"), Some(&json!("1024x1024")));
assert!(!mapped.contains_key("unsupported_param"));
}
#[test]
fn transform_request_builds_openai_compatible_multipart_body() {
let optional_params = json!({
"quality": "high",
"size": "1024x1024"
})
.as_object()
.unwrap()
.clone();
let request = transform_image_edit_request(
"qwen-image-edit",
vec![image_part()],
Some(ImageEditInputFile {
filename: "mask.png".to_string(),
content_type: "image/png".to_string(),
data_base64: "bWFzaw==".to_string(),
}),
Some("make it brighter"),
optional_params,
)
.expect("request transforms");
assert_eq!(request.format, ImageEditRequestFormat::Multipart);
assert_eq!(request.data["model"], json!("qwen-image-edit"));
assert_eq!(request.data["prompt"], json!("make it brighter"));
assert_eq!(request.data["quality"], json!("high"));
assert_eq!(request.files[0].field_name, "image[]");
assert_eq!(request.files[1].field_name, "mask");
}
#[test]
fn transform_request_requires_an_image() {
let err = transform_image_edit_request("model", Vec::new(), None, None, Map::new())
.expect_err("empty images rejected");
assert_eq!(err, CoreError::MissingField("image"));
}
#[test]
fn complete_url_uses_vllm_base_and_dedupes_v1() {
assert_eq!(
complete_url(Some("http://localhost:8000"), &|_| None).unwrap(),
"http://localhost:8000/v1/images/edits"
);
assert_eq!(
complete_url(Some("http://localhost:8000/v1/"), &|_| None).unwrap(),
"http://localhost:8000/v1/images/edits"
);
}
}

View file

@ -0,0 +1 @@
pub mod image_edit;

View file

@ -1,6 +1,7 @@
use std::time::Duration;
use litellm_core::error::CoreError;
use litellm_core::image_edit::types::ImageEditInputFile;
use pyo3::exceptions::{PyRuntimeError, PyValueError};
use pyo3::prelude::*;
use pyo3::types::{PyAny, PyDict};
@ -15,6 +16,14 @@ type MarshaledOcrInputs = (
Option<Duration>,
);
type MarshaledImageEditInputs = (
Vec<ImageEditInputFile>,
Option<ImageEditInputFile>,
Option<Map<String, Value>>,
Map<String, Value>,
Option<Duration>,
);
fn py_to_json(py: Python<'_>, value: &Bound<'_, PyAny>) -> PyResult<Value> {
let json = py.import("json")?;
let encoded: String = json.call_method1("dumps", (value,))?.extract()?;
@ -66,6 +75,22 @@ fn optional_timeout(timeout_seconds: Option<f64>) -> Option<Duration> {
})
}
fn value_to_image_part(name: &'static str, value: Value) -> PyResult<ImageEditInputFile> {
serde_json::from_value(value).map_err(|err| {
PyValueError::new_err(format!("{name} must be an image file part object: {err}"))
})
}
fn value_to_image_parts(name: &'static str, value: Value) -> PyResult<Vec<ImageEditInputFile>> {
match value {
Value::Array(items) => items
.into_iter()
.map(|item| value_to_image_part(name, item))
.collect(),
_ => Err(PyValueError::new_err(format!("{name} must be a list"))),
}
}
fn marshal_inputs(
py: Python<'_>,
document: Py<PyAny>,
@ -84,6 +109,29 @@ fn marshal_inputs(
Ok((document, extra_headers, optional_params, timeout))
}
fn marshal_image_edit_inputs(
py: Python<'_>,
images: Py<PyAny>,
mask: Option<Py<PyAny>>,
extra_headers: Option<Py<PyAny>>,
optional_params: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<MarshaledImageEditInputs> {
let images = value_to_image_parts("images", py_to_json(py, images.bind(py))?)?;
let mask = match mask {
Some(mask) => Some(value_to_image_part("mask", py_to_json(py, mask.bind(py))?)?),
None => None,
};
let extra_headers = match extra_headers {
Some(headers) => Some(optional_object_to_map(py, "extra_headers", Some(headers))?),
None => None,
};
let optional_params = optional_object_to_map(py, "optional_params", optional_params)?;
let timeout = optional_timeout(timeout_seconds);
Ok((images, mask, extra_headers, optional_params, timeout))
}
/// Perform a Mistral OCR call end to end and return the response as a dict.
#[pyfunction]
#[pyo3(signature = (model, document, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))]
@ -172,6 +220,107 @@ fn aocr(
})
}
/// Perform a vLLM image-edit call end to end and return the response as a dict.
#[pyfunction]
#[pyo3(signature = (model, images, prompt=None, mask=None, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
fn image_edit(
py: Python<'_>,
model: String,
images: Py<PyAny>,
prompt: Option<String>,
mask: Option<Py<PyAny>>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
extra_headers: Option<Py<PyAny>>,
optional_params: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Py<PyAny>> {
let custom_llm_provider = custom_llm_provider.unwrap_or_else(|| "vllm".to_string());
let (images, mask, extra_headers, optional_params, timeout) = marshal_image_edit_inputs(
py,
images,
mask,
extra_headers,
optional_params,
timeout_seconds,
)?;
let result = gil::release_gil(py, || {
pyo3_async_runtimes::tokio::get_runtime().block_on(
litellm_providers::image_edit::image_edit(
litellm_providers::image_edit::ImageEditRequest {
model: &model,
images,
mask,
prompt: prompt.as_deref(),
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: &custom_llm_provider,
extra_headers,
optional_params,
timeout,
},
),
)
});
match result {
Ok(value) => json_to_py(py, value),
Err(err) => Err(core_error_to_pyerr(err)),
}
}
/// Perform a vLLM image-edit call end to end and return an asyncio awaitable.
#[pyfunction]
#[pyo3(signature = (model, images, prompt=None, mask=None, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))]
#[allow(clippy::too_many_arguments)]
fn aimage_edit(
py: Python<'_>,
model: String,
images: Py<PyAny>,
prompt: Option<String>,
mask: Option<Py<PyAny>>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
extra_headers: Option<Py<PyAny>>,
optional_params: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Bound<'_, PyAny>> {
let custom_llm_provider = custom_llm_provider.unwrap_or_else(|| "vllm".to_string());
let (images, mask, extra_headers, optional_params, timeout) = marshal_image_edit_inputs(
py,
images,
mask,
extra_headers,
optional_params,
timeout_seconds,
)?;
pyo3_async_runtimes::tokio::future_into_py(py, async move {
let value = litellm_providers::image_edit::image_edit(
litellm_providers::image_edit::ImageEditRequest {
model: &model,
images,
mask,
prompt: prompt.as_deref(),
api_key: api_key.as_deref(),
api_base: api_base.as_deref(),
custom_llm_provider: &custom_llm_provider,
extra_headers,
optional_params,
timeout,
},
)
.await
.map_err(core_error_to_pyerr)?;
Python::with_gil(|py| json_to_py(py, value))
})
}
/// Bridge GIL accounting, e.g. `{"releases": 12}`. Lets the Python side observe
/// how often the sync bridge has dropped the GIL while awaiting Rust work.
#[pyfunction]
@ -185,6 +334,8 @@ fn gil_stats(py: Python<'_>) -> PyResult<Py<PyAny>> {
fn litellm_python_bridge(module: &Bound<'_, PyModule>) -> PyResult<()> {
module.add_function(wrap_pyfunction!(ocr, module)?)?;
module.add_function(wrap_pyfunction!(aocr, module)?)?;
module.add_function(wrap_pyfunction!(image_edit, module)?)?;
module.add_function(wrap_pyfunction!(aimage_edit, module)?)?;
module.add_function(wrap_pyfunction!(gil_stats, module)?)?;
Ok(())
}

View file

@ -1,7 +1,10 @@
import asyncio
import base64
import contextvars
import importlib
import os
from functools import partial
from io import BufferedReader, BytesIO
from typing import (
TYPE_CHECKING,
Any,
@ -70,6 +73,13 @@ from litellm.utils import (
get_optional_params_image_gen,
)
from litellm.ocr.rust_bridge import (
RustAimageEdit,
RustImageEdit,
load_rust_aimage_edit,
load_rust_image_edit,
)
# Cache for ImageEditRequestUtils to avoid repeated __getattr__ calls
_ImageEditRequestUtils_cache: Optional["ImageEditRequestUtils"] = None
@ -85,6 +95,194 @@ def _get_ImageEditRequestUtils() -> "ImageEditRequestUtils":
return _ImageEditRequestUtils_cache
def _timeout_to_seconds(
timeout: Optional[Union[float, httpx.Timeout]],
) -> Optional[float]:
if timeout is None:
return None
if isinstance(timeout, httpx.Timeout):
read_timeout = timeout.read
if read_timeout is None:
return None
return float(read_timeout)
return float(timeout)
def _read_file_content(file_value: Any) -> bytes:
if isinstance(file_value, bytes):
return file_value
if isinstance(file_value, bytearray):
return bytes(file_value)
if isinstance(file_value, (BufferedReader, BytesIO)):
current_pos = file_value.tell()
file_value.seek(0)
data = file_value.read()
file_value.seek(current_pos)
return data
if hasattr(file_value, "read"):
current_pos = getattr(file_value, "tell", lambda: 0)()
if hasattr(file_value, "seek"):
file_value.seek(0)
data = file_value.read()
if hasattr(file_value, "seek"):
file_value.seek(current_pos)
if isinstance(data, str):
return data.encode()
return bytes(data)
if isinstance(file_value, (str, os.PathLike)) and os.path.exists(file_value):
with open(file_value, "rb") as file:
return file.read()
if isinstance(file_value, str):
return file_value.encode()
raise TypeError(
f"Unsupported image file type for Rust image_edit: {type(file_value)}"
)
def _filename_for_file(file_value: Any, default: str) -> str:
name = getattr(file_value, "name", None)
if isinstance(name, str) and name:
return os.path.basename(name)
if isinstance(file_value, (str, os.PathLike)) and os.path.exists(file_value):
return os.path.basename(os.fspath(file_value))
return default
def _rust_image_file_part(file_value: Any, default_filename: str) -> Dict[str, object]:
filename = default_filename
content_type: Optional[str] = None
raw_file = file_value
if isinstance(file_value, tuple):
if len(file_value) >= 1 and file_value[0] is not None:
filename = str(file_value[0])
if len(file_value) >= 2:
raw_file = file_value[1]
if len(file_value) >= 3 and file_value[2] is not None:
content_type = str(file_value[2])
filename = filename or _filename_for_file(raw_file, default_filename)
content_type = content_type or _get_ImageEditRequestUtils().get_image_content_type(
raw_file
)
return {
"filename": filename,
"content_type": content_type,
"data_base64": base64.b64encode(_read_file_content(raw_file)).decode("ascii"),
}
def _rust_image_file_parts(images: List[FileTypes]) -> List[Dict[str, object]]:
return [
_rust_image_file_part(image, f"image-{index}.png")
for index, image in enumerate(images)
if image is not None
]
def _prepare_rust_image_edit_params(
image_edit_optional_params: ImageEditOptionalRequestParams,
non_default_params: Dict[str, Any],
) -> tuple[Dict[str, object], Optional[Dict[str, object]]]:
optional_params: Dict[str, object] = dict(image_edit_optional_params)
optional_params.update(non_default_params)
mask = optional_params.pop("mask", None)
mask_part = _rust_image_file_part(mask, "mask.png") if mask is not None else None
return optional_params, mask_part
def _run_rust_image_edit(
rust_image_edit: RustImageEdit,
*,
model: str,
images: List[FileTypes],
prompt: Optional[str],
api_key: Optional[str],
api_base: Optional[str],
custom_llm_provider: str,
extra_headers: Optional[Dict[str, Any]],
optional_params: Dict[str, object],
mask: Optional[Dict[str, object]],
timeout: Optional[Union[float, httpx.Timeout]],
logging_obj: LiteLLMLoggingObj,
) -> ImageResponse:
image_parts = _rust_image_file_parts(images)
logging_obj.pre_call(
input=prompt,
api_key="",
additional_args={
"complete_input_dict": {
"model": model,
"prompt": prompt,
**optional_params,
},
"api_base": api_base,
"headers": extra_headers or {},
},
)
return ImageResponse(
**rust_image_edit(
model=model,
images=image_parts,
prompt=prompt,
mask=mask,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=cast(Optional[dict[str, object]], extra_headers),
optional_params=optional_params,
timeout_seconds=_timeout_to_seconds(timeout),
)
)
async def _run_rust_aimage_edit(
rust_aimage_edit: RustAimageEdit,
*,
model: str,
images: List[FileTypes],
prompt: Optional[str],
api_key: Optional[str],
api_base: Optional[str],
custom_llm_provider: str,
extra_headers: Optional[Dict[str, Any]],
optional_params: Dict[str, object],
mask: Optional[Dict[str, object]],
timeout: Optional[Union[float, httpx.Timeout]],
logging_obj: LiteLLMLoggingObj,
) -> ImageResponse:
image_parts = _rust_image_file_parts(images)
logging_obj.pre_call(
input=prompt,
api_key="",
additional_args={
"complete_input_dict": {
"model": model,
"prompt": prompt,
**optional_params,
},
"api_base": api_base,
"headers": extra_headers or {},
},
)
return ImageResponse(
**(
await rust_aimage_edit(
model=model,
images=image_parts,
prompt=prompt,
mask=mask,
api_key=api_key,
api_base=api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=cast(Optional[dict[str, object]], extra_headers),
optional_params=optional_params,
timeout_seconds=_timeout_to_seconds(timeout),
)
)
)
##### Image Generation #######################
@client
async def aimage_generation(*args, **kwargs) -> ImageResponse:
@ -791,6 +989,8 @@ def image_edit(
model_info = kwargs.get("model_info", None)
metadata = kwargs.get("metadata", {})
_is_async = kwargs.pop("async_call", False) is True
api_key_param: Optional[str] = kwargs.get("api_key")
api_base_param: Optional[str] = kwargs.get("api_base")
# add images / or return a single image
images = (
@ -809,9 +1009,10 @@ def image_edit(
# get llm provider logic
litellm_params = GenericLiteLLMParams(**kwargs)
model, custom_llm_provider, _, _ = get_llm_provider(
model, custom_llm_provider, dynamic_api_key, api_base = get_llm_provider(
model=model or DEFAULT_IMAGE_ENDPOINT_MODEL,
custom_llm_provider=custom_llm_provider,
api_base=api_base_param,
)
# Check for custom provider
@ -867,6 +1068,81 @@ def image_edit(
client=custom_client,
)
local_vars.update(kwargs)
if custom_llm_provider == "vllm":
image_edit_optional_params = (
_get_ImageEditRequestUtils().get_requested_image_edit_optional_param(
local_vars
)
)
rust_optional_params, rust_mask = _prepare_rust_image_edit_params(
image_edit_optional_params=image_edit_optional_params,
non_default_params=non_default_params,
)
resolved_api_key = (
api_key_param or dynamic_api_key or get_secret_str("VLLM_API_KEY")
)
resolved_api_base = api_base or get_secret_str("VLLM_API_BASE")
if _is_async:
rust_aimage_edit = load_rust_aimage_edit()
if rust_aimage_edit is not None:
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
user=user,
optional_params=dict(rust_optional_params),
litellm_params={
**rust_optional_params,
"litellm_call_id": litellm_call_id,
"model_info": model_info,
},
custom_llm_provider=custom_llm_provider,
)
return _run_rust_aimage_edit(
rust_aimage_edit=rust_aimage_edit,
model=model,
images=images,
prompt=prompt,
api_key=resolved_api_key,
api_base=resolved_api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=rust_optional_params,
mask=rust_mask,
timeout=timeout,
logging_obj=litellm_logging_obj,
)
else:
rust_image_edit = load_rust_image_edit()
if rust_image_edit is not None:
litellm_logging_obj.update_from_kwargs(
kwargs=kwargs,
model=model,
user=user,
optional_params=dict(rust_optional_params),
litellm_params={
**rust_optional_params,
"litellm_call_id": litellm_call_id,
"model_info": model_info,
},
custom_llm_provider=custom_llm_provider,
)
return _run_rust_image_edit(
rust_image_edit=rust_image_edit,
model=model,
images=images,
prompt=prompt,
api_key=resolved_api_key,
api_base=resolved_api_base,
custom_llm_provider=custom_llm_provider,
extra_headers=extra_headers,
optional_params=rust_optional_params,
mask=rust_mask,
timeout=timeout,
logging_obj=litellm_logging_obj,
)
# get provider config
image_edit_provider_config: Optional[BaseImageEditConfig] = (
ProviderConfigManager.get_provider_image_edit_config(
@ -878,7 +1154,6 @@ def image_edit(
if image_edit_provider_config is None:
raise ValueError(f"image edit is not supported for {custom_llm_provider}")
local_vars.update(kwargs)
# Get ImageEditOptionalRequestParams with only valid parameters
image_edit_optional_params: (
ImageEditOptionalRequestParams

View file

@ -48,6 +48,44 @@ class RustAocr(Protocol):
raise NotImplementedError
class RustImageEdit(Protocol):
"""Signature of the compiled ``litellm_python_bridge.image_edit`` entrypoint."""
def __call__(
self,
model: str,
images: list[dict[str, object]],
prompt: str | None,
mask: dict[str, object] | None,
api_key: str | None,
api_base: str | None,
custom_llm_provider: str,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> dict[str, object]:
raise NotImplementedError
class RustAimageEdit(Protocol):
"""Signature of the compiled ``litellm_python_bridge.aimage_edit`` entrypoint."""
def __call__(
self,
model: str,
images: list[dict[str, object]],
prompt: str | None,
mask: dict[str, object] | None,
api_key: str | None,
api_base: str | None,
custom_llm_provider: str,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> Awaitable[dict[str, object]]:
raise NotImplementedError
class _Unset:
"""Sentinel type so ``ocr=None`` can clear a prior injection while omission preserves it."""
@ -57,6 +95,8 @@ _UNSET: Final[_Unset] = _Unset()
_rust_ocr_enabled = False
_rust_ocr_impl: RustOcr | None = None
_rust_aocr_impl: RustAocr | None = None
_rust_image_edit_impl: RustImageEdit | None = None
_rust_aimage_edit_impl: RustAimageEdit | None = None
def use_litellm_rust(
@ -64,19 +104,26 @@ def use_litellm_rust(
*,
ocr: RustOcr | None | _Unset = _UNSET,
aocr: RustAocr | None | _Unset = _UNSET,
image_edit: RustImageEdit | None | _Unset = _UNSET,
aimage_edit: RustAimageEdit | None | _Unset = _UNSET,
) -> None:
"""Route supported OCR calls through the Rust ``litellm_python_bridge`` extension.
"""Route supported calls through the Rust ``litellm_python_bridge`` extension.
``ocr`` and ``aocr`` inject bridge callables; when omitted the compiled
Bridge callables can be injected for tests; when omitted the compiled
extension is loaded on demand and any previously injected bridge is
preserved. Pass ``None`` explicitly to clear a prior injection.
"""
global _rust_ocr_enabled, _rust_ocr_impl, _rust_aocr_impl
global _rust_image_edit_impl, _rust_aimage_edit_impl
_rust_ocr_enabled = enabled
if not isinstance(ocr, _Unset):
_rust_ocr_impl = ocr
if not isinstance(aocr, _Unset):
_rust_aocr_impl = aocr
if not isinstance(image_edit, _Unset):
_rust_image_edit_impl = image_edit
if not isinstance(aimage_edit, _Unset):
_rust_aimage_edit_impl = aimage_edit
def rust_ocr_enabled() -> bool:
@ -84,6 +131,11 @@ def rust_ocr_enabled() -> bool:
return _rust_ocr_enabled
def rust_image_edit_enabled() -> bool:
"""Whether the Rust image-edit path has been turned on."""
return _rust_ocr_enabled
def load_rust_ocr() -> RustOcr | None:
"""Return the Rust OCR callable, or ``None`` when no bridge is available.
@ -109,3 +161,25 @@ def load_rust_aocr() -> RustAocr | None:
except ImportError:
return None
return cast(RustAocr, getattr(litellm_python_bridge, "aocr", None))
def load_rust_image_edit() -> RustImageEdit | None:
"""Return the Rust image-edit callable, or ``None`` when unavailable."""
if _rust_image_edit_impl is not None:
return _rust_image_edit_impl
try:
import litellm_python_bridge
except ImportError:
return None
return cast(RustImageEdit, getattr(litellm_python_bridge, "image_edit", None))
def load_rust_aimage_edit() -> RustAimageEdit | None:
"""Return the async Rust image-edit callable, or ``None`` when unavailable."""
if _rust_aimage_edit_impl is not None:
return _rust_aimage_edit_impl
try:
import litellm_python_bridge
except ImportError:
return None
return cast(RustAimageEdit, getattr(litellm_python_bridge, "aimage_edit", None))

View file

@ -0,0 +1,144 @@
"""Tests for vLLM image-edit routing through the optional Rust bridge."""
from __future__ import annotations
import importlib
import pytest
import litellm
rust_bridge = importlib.import_module("litellm.ocr.rust_bridge")
PNG_BYTES = b"\x89PNG\r\n\x1a\nfake-image"
class RecordingImageEditBridge:
def __init__(self) -> None:
self.calls: list[dict[str, object]] = []
def __call__(
self,
model: str,
images: list[dict[str, object]],
prompt: str | None,
mask: dict[str, object] | None,
api_key: str | None,
api_base: str | None,
custom_llm_provider: str,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> dict[str, object]:
self.calls.append(
{
"model": model,
"images": images,
"prompt": prompt,
"mask": mask,
"api_key": api_key,
"api_base": api_base,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"optional_params": optional_params,
"timeout_seconds": timeout_seconds,
}
)
return {"created": 1, "data": [{"b64_json": "ZmFrZS1pbWFnZQ=="}]}
class RecordingAsyncImageEditBridge:
def __init__(self) -> None:
self.calls: list[dict[str, object]] = []
async def __call__(
self,
model: str,
images: list[dict[str, object]],
prompt: str | None,
mask: dict[str, object] | None,
api_key: str | None,
api_base: str | None,
custom_llm_provider: str,
extra_headers: dict[str, object] | None,
optional_params: dict[str, object],
timeout_seconds: float | None,
) -> dict[str, object]:
self.calls.append(
{
"model": model,
"images": images,
"prompt": prompt,
"mask": mask,
"api_key": api_key,
"api_base": api_base,
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
"optional_params": optional_params,
"timeout_seconds": timeout_seconds,
}
)
return {"created": 1, "data": [{"b64_json": "YXN5bmMtaW1hZ2U="}]}
@pytest.fixture(autouse=True)
def reset_rust_bridge() -> None:
rust_bridge.use_litellm_rust(
False, ocr=None, aocr=None, image_edit=None, aimage_edit=None
)
yield
rust_bridge.use_litellm_rust(
False, ocr=None, aocr=None, image_edit=None, aimage_edit=None
)
def test_vllm_image_edit_routes_to_rust_bridge() -> None:
bridge = RecordingImageEditBridge()
litellm.use_litellm_rust(True, image_edit=bridge)
response = litellm.image_edit(
model="vllm/qwen-image-edit",
image=PNG_BYTES,
prompt="make it brighter",
quality="high",
size="1024x1024",
api_base="http://localhost:8000",
api_key="sk-test",
extra_headers={"x-trace-id": "trace-1"},
timeout=12.5,
)
assert response.data[0].b64_json == "ZmFrZS1pbWFnZQ=="
assert len(bridge.calls) == 1
call = bridge.calls[0]
assert call["model"] == "qwen-image-edit"
assert call["prompt"] == "make it brighter"
assert call["api_key"] == "sk-test"
assert call["api_base"] == "http://localhost:8000"
assert call["custom_llm_provider"] == "vllm"
assert call["extra_headers"] == {"x-trace-id": "trace-1"}
assert call["optional_params"] == {"quality": "high", "size": "1024x1024"}
assert call["timeout_seconds"] == 12.5
images = call["images"]
assert isinstance(images, list)
assert images[0]["filename"] == "image-0.png"
assert images[0]["content_type"] == "image/png"
assert images[0]["data_base64"] == "iVBORw0KGgpmYWtlLWltYWdl"
@pytest.mark.asyncio
async def test_vllm_aimage_edit_routes_to_async_rust_bridge() -> None:
bridge = RecordingAsyncImageEditBridge()
litellm.use_litellm_rust(True, aimage_edit=bridge)
response = await litellm.aimage_edit(
model="vllm/qwen-image-edit",
image=[PNG_BYTES],
prompt="make it darker",
api_base="http://localhost:8000",
)
assert response.data[0].b64_json == "YXN5bmMtaW1hZ2U="
assert len(bridge.calls) == 1
assert bridge.calls[0]["model"] == "qwen-image-edit"
assert bridge.calls[0]["prompt"] == "make it darker"

View file

@ -213,9 +213,13 @@ def build_prepared_request(
@pytest.fixture(autouse=True)
def _reset_rust_flag():
"""Keep the global toggle isolated between tests."""
rust_bridge.use_litellm_rust(False, ocr=None, aocr=None)
rust_bridge.use_litellm_rust(
False, ocr=None, aocr=None, image_edit=None, aimage_edit=None
)
yield
rust_bridge.use_litellm_rust(False, ocr=None, aocr=None)
rust_bridge.use_litellm_rust(
False, ocr=None, aocr=None, image_edit=None, aimage_edit=None
)
@pytest.fixture