mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
refactor(rust): extract inference-transcription crate (#44809)
Co-authored-by: Yujong Lee <yujong@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
1a9b533e71
commit
63babf23e6
18 changed files with 209 additions and 30 deletions
21
litellm-rust/Cargo.lock
generated
21
litellm-rust/Cargo.lock
generated
|
|
@ -3836,6 +3836,7 @@ dependencies = [
|
|||
"litellm-host-http",
|
||||
"litellm-http",
|
||||
"litellm-inference",
|
||||
"litellm-inference-transcription",
|
||||
"litellm-llms",
|
||||
"litellm-llms-types",
|
||||
"litellm-router",
|
||||
|
|
@ -4039,6 +4040,25 @@ dependencies = [
|
|||
"wiremock",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-inference-transcription"
|
||||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"futures-util",
|
||||
"litellm-auth",
|
||||
"litellm-core-utils",
|
||||
"litellm-http",
|
||||
"litellm-inference",
|
||||
"litellm-llms",
|
||||
"litellm-secrets",
|
||||
"litellm-tracing",
|
||||
"rstest",
|
||||
"serde_json",
|
||||
"tokio",
|
||||
"tracing",
|
||||
"wiremock",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "litellm-llms"
|
||||
version = "0.1.0"
|
||||
|
|
@ -4122,6 +4142,7 @@ dependencies = [
|
|||
"litellm-host-python",
|
||||
"litellm-http",
|
||||
"litellm-inference",
|
||||
"litellm-inference-transcription",
|
||||
"litellm-llms",
|
||||
"litellm-llms-types",
|
||||
"litellm-secrets",
|
||||
|
|
|
|||
|
|
@ -17,6 +17,7 @@ litellm-traces-cache = { path = "crates/traces-cache" }
|
|||
litellm-traces-clickhouse = { path = "crates/traces-clickhouse" }
|
||||
litellm-storage-clickhouse = { path = "crates/storage-clickhouse" }
|
||||
litellm-inference = { path = "crates/inference" }
|
||||
litellm-inference-transcription = { path = "crates/inference-transcription" }
|
||||
litellm-gateway-mcp = { path = "crates/gateway-mcp" }
|
||||
litellm-gateway = { path = "crates/gateway" }
|
||||
litellm-gateway-inference = { path = "crates/gateway-inference" }
|
||||
|
|
|
|||
|
|
@ -13,6 +13,7 @@ bytes.workspace = true
|
|||
litellm-auth.workspace = true
|
||||
litellm-gateway-auth.workspace = true
|
||||
litellm-inference.workspace = true
|
||||
litellm-inference-transcription.workspace = true
|
||||
litellm-host-http.workspace = true
|
||||
litellm-host.workspace = true
|
||||
litellm-http.workspace = true
|
||||
|
|
|
|||
|
|
@ -3,7 +3,7 @@ use std::{path::Path, sync::Arc};
|
|||
|
||||
use axum::{Json, extract::State, response::IntoResponse};
|
||||
use base64::{Engine, engine::general_purpose::STANDARD};
|
||||
use litellm_inference::audio_transcription::types::AudioTranscriptionRequest;
|
||||
use litellm_inference_transcription::types::AudioTranscriptionRequest;
|
||||
use serde_json::{Value, json};
|
||||
|
||||
use crate::{
|
||||
|
|
|
|||
|
|
@ -17,9 +17,10 @@ use std::sync::Arc;
|
|||
use axum::{Router, routing::post};
|
||||
use litellm_http::{ClientVariant, HttpClientConfig, media::UrlPolicy};
|
||||
use litellm_inference::{
|
||||
audio_transcription::AudioTranscriptionRoute, chat_completions::ChatCompletionsRoute,
|
||||
messages::MessagesRoute, ocr::OcrRoute, resources::CoreResources, responses::ResponsesRoute,
|
||||
chat_completions::ChatCompletionsRoute, messages::MessagesRoute, ocr::OcrRoute,
|
||||
resources::CoreResources, responses::ResponsesRoute,
|
||||
};
|
||||
use litellm_inference_transcription::AudioTranscriptionRoute;
|
||||
use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
|
||||
|
|
|
|||
25
litellm-rust/crates/inference-transcription/Cargo.toml
Normal file
25
litellm-rust/crates/inference-transcription/Cargo.toml
Normal file
|
|
@ -0,0 +1,25 @@
|
|||
[package]
|
||||
name = "litellm-inference-transcription"
|
||||
version = "0.1.0"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
futures-util.workspace = true
|
||||
litellm-auth.workspace = true
|
||||
litellm-core-utils.workspace = true
|
||||
litellm-http.workspace = true
|
||||
litellm-inference.workspace = true
|
||||
litellm-llms.workspace = true
|
||||
litellm-secrets.workspace = true
|
||||
serde_json.workspace = true
|
||||
tracing.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
litellm-http.workspace = true
|
||||
litellm-inference = { workspace = true, features = ["test-support"] }
|
||||
litellm-tracing.workspace = true
|
||||
rstest.workspace = true
|
||||
tokio.workspace = true
|
||||
wiremock.workspace = true
|
||||
|
|
@ -0,0 +1 @@
|
|||
pub(crate) const AUDIO_TRANSCRIPTION_TIMEOUT_SECS: u64 = 600;
|
||||
|
|
@ -6,8 +6,7 @@ use serde_json::Value;
|
|||
|
||||
use super::Error;
|
||||
use crate::{
|
||||
audio_transcription::types::ProviderAudioTranscriptionRequest,
|
||||
constants::AUDIO_TRANSCRIPTION_TIMEOUT_SECS,
|
||||
constants::AUDIO_TRANSCRIPTION_TIMEOUT_SECS, types::ProviderAudioTranscriptionRequest,
|
||||
};
|
||||
|
||||
pub async fn execute_audio_transcription_provider_call(
|
||||
|
|
@ -17,7 +16,7 @@ pub async fn execute_audio_transcription_provider_call(
|
|||
) -> Result<Value, Error> {
|
||||
let env_lookup = |key: &str| request.secrets.get(key);
|
||||
let authenticated = resolve_auth(auth, request.environment.clone(), &env_lookup).await?;
|
||||
let outbound = crate::outbound::outbound_request(
|
||||
let outbound = litellm_inference::outbound::outbound_request(
|
||||
authenticated,
|
||||
request.url.clone(),
|
||||
&request.body,
|
||||
|
|
@ -27,7 +26,7 @@ pub async fn execute_audio_transcription_provider_call(
|
|||
.unwrap_or(Duration::from_secs(AUDIO_TRANSCRIPTION_TIMEOUT_SECS)),
|
||||
),
|
||||
)?;
|
||||
let response = crate::outbound::send(outbound, http)
|
||||
let response = litellm_inference::outbound::send(outbound, http)
|
||||
.await
|
||||
.map_err(|error| {
|
||||
Error::Transport(litellm_http::transport::Error::Network(error.to_string()))
|
||||
|
|
@ -1,5 +1,6 @@
|
|||
pub mod types;
|
||||
pub use crate::error::RouteError as Error;
|
||||
pub use litellm_inference::RouteError as Error;
|
||||
mod constants;
|
||||
mod handler;
|
||||
mod prepare;
|
||||
pub use handler::execute_audio_transcription_provider_call;
|
||||
|
|
@ -9,7 +10,7 @@ pub use prepare::prepare_audio_transcription_provider_call;
|
|||
use serde_json::Value;
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::audio_transcription::types::AudioTranscriptionRequest;
|
||||
use crate::types::AudioTranscriptionRequest;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct AudioTranscriptionRoute {
|
||||
|
|
@ -40,10 +41,10 @@ impl AudioTranscriptionRoute {
|
|||
outcome
|
||||
))]
|
||||
pub async fn execute(&self, request: AudioTranscriptionRequest<'_>) -> Result<Value, Error> {
|
||||
crate::diagnostic::unary(async {
|
||||
litellm_inference::diagnostic::unary(async {
|
||||
let request =
|
||||
prepare_audio_transcription_provider_call(request, self.secrets.as_ref()).await?;
|
||||
crate::diagnostic::provider(&request.model, &request.custom_llm_provider);
|
||||
litellm_inference::diagnostic::provider(&request.model, &request.custom_llm_provider);
|
||||
let execute: futures_util::future::BoxFuture<'_, Result<Value, Error>> = Box::pin(
|
||||
execute_audio_transcription_provider_call(&self.http, &self.auth, request),
|
||||
);
|
||||
|
|
@ -11,10 +11,8 @@ use litellm_llms::{
|
|||
use litellm_secrets::source::SecretSource;
|
||||
|
||||
use super::Error;
|
||||
use crate::audio_transcription::types::{
|
||||
AudioTranscriptionRequest, ProviderAudioTranscriptionRequest,
|
||||
};
|
||||
use crate::provider::resolve_llm_provider;
|
||||
use crate::types::{AudioTranscriptionRequest, ProviderAudioTranscriptionRequest};
|
||||
use litellm_inference::provider::resolve_llm_provider;
|
||||
|
||||
fn provider_config(provider: LlmProviders) -> Option<&'static dyn BaseAudioTranscriptionConfig> {
|
||||
match provider {
|
||||
|
|
@ -1,4 +1,4 @@
|
|||
use litellm_inference::audio_transcription::{Error, types::AudioTranscriptionRequest};
|
||||
use litellm_inference_transcription::{Error, types::AudioTranscriptionRequest};
|
||||
use rstest::{fixture, rstest};
|
||||
use serde_json::{Map, Value, json};
|
||||
use wiremock::ResponseTemplate;
|
||||
143
litellm-rust/crates/inference-transcription/tests/support/mod.rs
Normal file
143
litellm-rust/crates/inference-transcription/tests/support/mod.rs
Normal file
|
|
@ -0,0 +1,143 @@
|
|||
//! Shared fixtures for route integration tests: a scripted upstream and a recording
|
||||
//! secret source.
|
||||
|
||||
#![allow(dead_code)] // each test binary compiles this module on its own and uses a different subset
|
||||
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use litellm_inference::test_support::{http_config, no_secrets, provider_http, resources};
|
||||
use serde_json::Value;
|
||||
use wiremock::{Mock, MockServer, Request, ResponseTemplate, matchers::any};
|
||||
|
||||
/// A port nothing listens on, for calls that must fail before any request is sent.
|
||||
pub const UNREACHABLE_BASE: &str = "http://127.0.0.1:1";
|
||||
|
||||
pub fn audio_transcription_route() -> litellm_inference_transcription::AudioTranscriptionRoute {
|
||||
let resources = resources();
|
||||
litellm_inference_transcription::AudioTranscriptionRoute::new(
|
||||
provider_http(&resources, &http_config()),
|
||||
resources.auth,
|
||||
no_secrets(),
|
||||
)
|
||||
}
|
||||
|
||||
/// Starts an upstream that answers its n-th request with the n-th response and 404s after.
|
||||
pub async fn upstream(responses: impl IntoIterator<Item = ResponseTemplate>) -> MockServer {
|
||||
let server = MockServer::start().await;
|
||||
respond_in_order(&server, responses).await;
|
||||
server
|
||||
}
|
||||
|
||||
/// Scripts responses on a started server, for responses that need its address.
|
||||
pub async fn respond_in_order(
|
||||
server: &MockServer,
|
||||
responses: impl IntoIterator<Item = ResponseTemplate>,
|
||||
) {
|
||||
for response in responses {
|
||||
Mock::given(any())
|
||||
.respond_with(response)
|
||||
.up_to_n_times(1)
|
||||
.mount(server)
|
||||
.await;
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn received(server: &MockServer) -> Vec<Request> {
|
||||
server
|
||||
.received_requests()
|
||||
.await
|
||||
.expect("request recording is on")
|
||||
}
|
||||
|
||||
pub async fn only_request(server: &MockServer) -> Request {
|
||||
let [request] = <[Request; 1]>::try_from(received(server).await)
|
||||
.unwrap_or_else(|requests| panic!("expected one request, got {}", requests.len()));
|
||||
request
|
||||
}
|
||||
|
||||
pub fn json_response(body: Value) -> ResponseTemplate {
|
||||
ResponseTemplate::new(200).set_body_json(body)
|
||||
}
|
||||
|
||||
pub trait ReceivedRequest {
|
||||
fn header(&self, name: &str) -> Option<&str>;
|
||||
fn header_values(&self, name: &str) -> Vec<&str>;
|
||||
fn json(&self) -> Value;
|
||||
fn body_text(&self) -> String;
|
||||
/// The path and query, as the request line carried them.
|
||||
fn target(&self) -> String;
|
||||
fn query(&self, name: &str) -> Option<String>;
|
||||
}
|
||||
|
||||
impl ReceivedRequest for Request {
|
||||
fn header(&self, name: &str) -> Option<&str> {
|
||||
self.headers.get(name).and_then(|value| value.to_str().ok())
|
||||
}
|
||||
|
||||
fn header_values(&self, name: &str) -> Vec<&str> {
|
||||
self.headers
|
||||
.get_all(name)
|
||||
.iter()
|
||||
.filter_map(|value| value.to_str().ok())
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn json(&self) -> Value {
|
||||
serde_json::from_slice(&self.body).expect("request body is json")
|
||||
}
|
||||
|
||||
fn body_text(&self) -> String {
|
||||
String::from_utf8_lossy(&self.body).into_owned()
|
||||
}
|
||||
|
||||
fn target(&self) -> String {
|
||||
match self.url.query() {
|
||||
Some(query) => format!("{}?{query}", self.url.path()),
|
||||
None => self.url.path().to_string(),
|
||||
}
|
||||
}
|
||||
|
||||
fn query(&self, name: &str) -> Option<String> {
|
||||
self.url
|
||||
.query_pairs()
|
||||
.find_map(|(key, value)| (key == name).then(|| value.into_owned()))
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub struct TraceCapture(Arc<Mutex<Vec<Value>>>);
|
||||
|
||||
impl TraceCapture {
|
||||
pub fn logger(&self) -> litellm_tracing::Logger {
|
||||
litellm_tracing::Logger::new(self.clone())
|
||||
}
|
||||
|
||||
pub fn records(&self) -> Vec<Value> {
|
||||
self.0.lock().unwrap().clone()
|
||||
}
|
||||
|
||||
pub fn summaries(&self, name: &str) -> Vec<Value> {
|
||||
self.records()
|
||||
.into_iter()
|
||||
.filter(|record| record["span_name"] == name)
|
||||
.collect()
|
||||
}
|
||||
}
|
||||
|
||||
impl litellm_tracing::Sink for TraceCapture {
|
||||
fn enabled(&self, metadata: &litellm_tracing::Metadata<'_>) -> bool {
|
||||
metadata.target().starts_with("litellm_inference")
|
||||
}
|
||||
|
||||
fn emit(&self, record: &litellm_tracing::Record) {
|
||||
self.0
|
||||
.lock()
|
||||
.unwrap()
|
||||
.push(Value::Object(record.fields.clone()));
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest::fixture]
|
||||
pub fn traces() -> TraceCapture {
|
||||
TraceCapture::default()
|
||||
}
|
||||
|
|
@ -9,7 +9,5 @@ pub(crate) const MESSAGES_TIMEOUT_SECS: u64 = 600;
|
|||
/// seconds. Mirrors the Python chat completions default.
|
||||
pub(crate) const CHAT_COMPLETIONS_TIMEOUT_SECS: u64 = 600;
|
||||
|
||||
pub(crate) const AUDIO_TRANSCRIPTION_TIMEOUT_SECS: u64 = 600;
|
||||
|
||||
/// `object` field every non-streaming chat completion response carries.
|
||||
pub const CHAT_COMPLETION_OBJECT: &str = "chat.completion";
|
||||
|
|
|
|||
|
|
@ -1,7 +1,6 @@
|
|||
pub mod context;
|
||||
pub mod diagnostic;
|
||||
|
||||
pub mod audio_transcription;
|
||||
pub mod caching;
|
||||
pub mod chat_completions;
|
||||
pub mod constants;
|
||||
|
|
|
|||
|
|
@ -48,16 +48,6 @@ pub fn responses_route(
|
|||
)
|
||||
}
|
||||
|
||||
pub fn audio_transcription_route() -> litellm_inference::audio_transcription::AudioTranscriptionRoute
|
||||
{
|
||||
let resources = resources();
|
||||
litellm_inference::audio_transcription::AudioTranscriptionRoute::new(
|
||||
provider_http(&resources, &http_config()),
|
||||
resources.auth,
|
||||
no_secrets(),
|
||||
)
|
||||
}
|
||||
|
||||
pub fn build_ocr_route(
|
||||
resources: &litellm_inference::resources::CoreResources,
|
||||
config: &HttpClientConfig,
|
||||
|
|
|
|||
|
|
@ -37,6 +37,7 @@ litellm-auth.workspace = true
|
|||
litellm-auth-aws.workspace = true
|
||||
litellm-callbacks-legacy-python.workspace = true
|
||||
litellm-inference.workspace = true
|
||||
litellm-inference-transcription.workspace = true
|
||||
litellm-core-utils.workspace = true
|
||||
litellm-http.workspace = true
|
||||
litellm-llms.workspace = true
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
use crate::execution::{run_async, run_sync};
|
||||
use litellm_host_python::from_py_argument;
|
||||
use litellm_inference::audio_transcription::{
|
||||
use litellm_inference_transcription::{
|
||||
AudioTranscriptionRoute, Error, types::AudioTranscriptionRequest,
|
||||
};
|
||||
use pyo3::{prelude::*, types::PyDict};
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue