mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
refactor(rust-gateway): scope provider resolution to router::resolve_deployment_provider
Rename the prefix/custom_llm_provider helper to resolve_deployment_provider and move it under core::router so its narrow contract is explicit; it does not implement Python's full provider inference. Rebuild the OCR response object without mutating a Map when normalizing the model alias. Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com>
This commit is contained in:
parent
c19e6382d6
commit
f47752cee6
5 changed files with 145 additions and 110 deletions
|
|
@ -1,6 +1,5 @@
|
||||||
use litellm_core::error::CoreError;
|
use litellm_core::error::CoreError;
|
||||||
use litellm_core::get_llm_provider::get_llm_provider;
|
use litellm_core::router::{resolve_deployment_provider, Router};
|
||||||
use litellm_core::router::Router;
|
|
||||||
use litellm_core::CoreResult;
|
use litellm_core::CoreResult;
|
||||||
use serde_json::Value;
|
use serde_json::Value;
|
||||||
|
|
||||||
|
|
@ -23,7 +22,7 @@ pub async fn run_ocr(router: &Router, call: OcrCall) -> CoreResult<Value> {
|
||||||
})?;
|
})?;
|
||||||
let params = &deployment.litellm_params;
|
let params = &deployment.litellm_params;
|
||||||
let (provider_model, provider) =
|
let (provider_model, provider) =
|
||||||
get_llm_provider(¶ms.model, params.custom_llm_provider.as_deref())?;
|
resolve_deployment_provider(¶ms.model, params.custom_llm_provider.as_deref())?;
|
||||||
|
|
||||||
let OcrCall {
|
let OcrCall {
|
||||||
model: requested_model,
|
model: requested_model,
|
||||||
|
|
@ -32,7 +31,7 @@ pub async fn run_ocr(router: &Router, call: OcrCall) -> CoreResult<Value> {
|
||||||
timeout,
|
timeout,
|
||||||
} = call;
|
} = call;
|
||||||
|
|
||||||
let mut response = ocr(OcrRequest {
|
let response = ocr(OcrRequest {
|
||||||
model: &provider_model,
|
model: &provider_model,
|
||||||
document,
|
document,
|
||||||
api_key: present(params.api_key.as_deref()),
|
api_key: present(params.api_key.as_deref()),
|
||||||
|
|
@ -44,10 +43,22 @@ pub async fn run_ocr(router: &Router, call: OcrCall) -> CoreResult<Value> {
|
||||||
})
|
})
|
||||||
.await?;
|
.await?;
|
||||||
|
|
||||||
if let Value::Object(map) = &mut response {
|
Ok(normalize_response_model(response, requested_model))
|
||||||
map.insert("model".to_string(), Value::String(requested_model));
|
}
|
||||||
}
|
|
||||||
Ok(response)
|
fn normalize_response_model(response: Value, requested_model: String) -> Value {
|
||||||
|
let Value::Object(map) = response else {
|
||||||
|
return response;
|
||||||
|
};
|
||||||
|
Value::Object(
|
||||||
|
map.into_iter()
|
||||||
|
.filter(|(key, _)| key != "model")
|
||||||
|
.chain(std::iter::once((
|
||||||
|
"model".to_string(),
|
||||||
|
Value::String(requested_model),
|
||||||
|
)))
|
||||||
|
.collect(),
|
||||||
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
|
|
@ -62,4 +73,31 @@ mod tests {
|
||||||
assert_eq!(present(Some(" ")), None);
|
assert_eq!(present(Some(" ")), None);
|
||||||
assert_eq!(present(None), None);
|
assert_eq!(present(None), None);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn normalize_response_model_replaces_model_and_keeps_other_fields() {
|
||||||
|
let response = serde_json::json!({
|
||||||
|
"object": "ocr",
|
||||||
|
"model": "mistral-ocr-latest",
|
||||||
|
"pages": [{"markdown": "hi"}]
|
||||||
|
});
|
||||||
|
let normalized = normalize_response_model(response, "rust-ocr-mistral".to_string());
|
||||||
|
assert_eq!(normalized["model"], "rust-ocr-mistral");
|
||||||
|
assert_eq!(normalized["object"], "ocr");
|
||||||
|
assert_eq!(normalized["pages"][0]["markdown"], "hi");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn normalize_response_model_adds_model_when_absent() {
|
||||||
|
let response = serde_json::json!({"object": "ocr"});
|
||||||
|
let normalized = normalize_response_model(response, "rust-ocr-mistral".to_string());
|
||||||
|
assert_eq!(normalized["model"], "rust-ocr-mistral");
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn normalize_response_model_leaves_non_object_untouched() {
|
||||||
|
let response = serde_json::json!("not-an-object");
|
||||||
|
let normalized = normalize_response_model(response, "rust-ocr-mistral".to_string());
|
||||||
|
assert_eq!(normalized, serde_json::json!("not-an-object"));
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -1,100 +0,0 @@
|
||||||
use crate::error::CoreError;
|
|
||||||
use crate::CoreResult;
|
|
||||||
|
|
||||||
pub fn get_llm_provider(
|
|
||||||
model: &str,
|
|
||||||
custom_llm_provider: Option<&str>,
|
|
||||||
) -> CoreResult<(String, String)> {
|
|
||||||
let model = model.trim();
|
|
||||||
if model.is_empty() {
|
|
||||||
return Err(CoreError::InvalidProvider(
|
|
||||||
"deployment model is empty".to_string(),
|
|
||||||
));
|
|
||||||
}
|
|
||||||
|
|
||||||
if let Some(provider) = custom_llm_provider
|
|
||||||
.map(str::trim)
|
|
||||||
.filter(|provider| !provider.is_empty())
|
|
||||||
{
|
|
||||||
let resolved_model = model
|
|
||||||
.strip_prefix(&format!("{provider}/"))
|
|
||||||
.unwrap_or(model)
|
|
||||||
.to_string();
|
|
||||||
return Ok((resolved_model, provider.to_string()));
|
|
||||||
}
|
|
||||||
|
|
||||||
model
|
|
||||||
.split_once('/')
|
|
||||||
.filter(|(provider, rest)| !provider.is_empty() && !rest.is_empty())
|
|
||||||
.map(|(provider, rest)| (rest.to_string(), provider.to_string()))
|
|
||||||
.ok_or_else(|| {
|
|
||||||
CoreError::InvalidProvider(format!(
|
|
||||||
"deployment model '{model}' must be prefixed with a provider (e.g. \
|
|
||||||
'mistral/mistral-ocr-latest') or set 'custom_llm_provider'"
|
|
||||||
))
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
#[cfg(test)]
|
|
||||||
mod tests {
|
|
||||||
use super::*;
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn splits_provider_prefix_when_no_explicit_provider() {
|
|
||||||
assert_eq!(
|
|
||||||
get_llm_provider("mistral/mistral-ocr-latest", None).expect("splits"),
|
|
||||||
("mistral-ocr-latest".to_string(), "mistral".to_string())
|
|
||||||
);
|
|
||||||
assert_eq!(
|
|
||||||
get_llm_provider("azure_ai/doc-intelligence/prebuilt-layout", None).expect("splits"),
|
|
||||||
(
|
|
||||||
"doc-intelligence/prebuilt-layout".to_string(),
|
|
||||||
"azure_ai".to_string()
|
|
||||||
)
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn explicit_provider_wins_and_strips_matching_prefix() {
|
|
||||||
assert_eq!(
|
|
||||||
get_llm_provider("mistral/mistral-ocr-latest", Some("mistral")).expect("resolves"),
|
|
||||||
("mistral-ocr-latest".to_string(), "mistral".to_string())
|
|
||||||
);
|
|
||||||
assert_eq!(
|
|
||||||
get_llm_provider("mistral-ocr-latest", Some("mistral")).expect("resolves"),
|
|
||||||
("mistral-ocr-latest".to_string(), "mistral".to_string())
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn explicit_provider_keeps_unrelated_prefix() {
|
|
||||||
assert_eq!(
|
|
||||||
get_llm_provider("openai/some-model", Some("mistral")).expect("resolves"),
|
|
||||||
("openai/some-model".to_string(), "mistral".to_string())
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn blank_explicit_provider_falls_back_to_prefix() {
|
|
||||||
assert_eq!(
|
|
||||||
get_llm_provider("mistral/mistral-ocr-latest", Some(" ")).expect("splits"),
|
|
||||||
("mistral-ocr-latest".to_string(), "mistral".to_string())
|
|
||||||
);
|
|
||||||
}
|
|
||||||
|
|
||||||
#[test]
|
|
||||||
fn rejects_model_without_provider() {
|
|
||||||
assert!(matches!(
|
|
||||||
get_llm_provider("mistral-ocr-latest", None),
|
|
||||||
Err(CoreError::InvalidProvider(_))
|
|
||||||
));
|
|
||||||
assert!(matches!(
|
|
||||||
get_llm_provider("mistral/", None),
|
|
||||||
Err(CoreError::InvalidProvider(_))
|
|
||||||
));
|
|
||||||
assert!(matches!(
|
|
||||||
get_llm_provider(" ", Some("mistral")),
|
|
||||||
Err(CoreError::InvalidProvider(_))
|
|
||||||
));
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
@ -1,5 +1,4 @@
|
||||||
pub mod error;
|
pub mod error;
|
||||||
pub mod get_llm_provider;
|
|
||||||
pub mod ocr;
|
pub mod ocr;
|
||||||
pub mod providers;
|
pub mod providers;
|
||||||
pub mod realtime;
|
pub mod realtime;
|
||||||
|
|
|
||||||
|
|
@ -3,6 +3,9 @@
|
||||||
|
|
||||||
use serde::Deserialize;
|
use serde::Deserialize;
|
||||||
|
|
||||||
|
use crate::error::CoreError;
|
||||||
|
use crate::CoreResult;
|
||||||
|
|
||||||
/// Per-deployment call parameters, mirroring Python's `litellm_params`.
|
/// Per-deployment call parameters, mirroring Python's `litellm_params`.
|
||||||
#[derive(Clone, Debug, Deserialize)]
|
#[derive(Clone, Debug, Deserialize)]
|
||||||
pub struct LiteLLMParams {
|
pub struct LiteLLMParams {
|
||||||
|
|
@ -24,6 +27,40 @@ pub struct Deployment {
|
||||||
pub litellm_params: LiteLLMParams,
|
pub litellm_params: LiteLLMParams,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
pub fn resolve_deployment_provider(
|
||||||
|
model: &str,
|
||||||
|
custom_llm_provider: Option<&str>,
|
||||||
|
) -> CoreResult<(String, String)> {
|
||||||
|
let model = model.trim();
|
||||||
|
if model.is_empty() {
|
||||||
|
return Err(CoreError::InvalidProvider(
|
||||||
|
"deployment model is empty".to_string(),
|
||||||
|
));
|
||||||
|
}
|
||||||
|
|
||||||
|
if let Some(provider) = custom_llm_provider
|
||||||
|
.map(str::trim)
|
||||||
|
.filter(|provider| !provider.is_empty())
|
||||||
|
{
|
||||||
|
let resolved_model = model
|
||||||
|
.strip_prefix(&format!("{provider}/"))
|
||||||
|
.unwrap_or(model)
|
||||||
|
.to_string();
|
||||||
|
return Ok((resolved_model, provider.to_string()));
|
||||||
|
}
|
||||||
|
|
||||||
|
model
|
||||||
|
.split_once('/')
|
||||||
|
.filter(|(provider, rest)| !provider.is_empty() && !rest.is_empty())
|
||||||
|
.map(|(provider, rest)| (rest.to_string(), provider.to_string()))
|
||||||
|
.ok_or_else(|| {
|
||||||
|
CoreError::InvalidProvider(format!(
|
||||||
|
"deployment model '{model}' must be prefixed with a provider (e.g. \
|
||||||
|
'mistral/mistral-ocr-latest') or set 'custom_llm_provider'"
|
||||||
|
))
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
#[cfg(test)]
|
#[cfg(test)]
|
||||||
mod tests {
|
mod tests {
|
||||||
use super::*;
|
use super::*;
|
||||||
|
|
@ -57,4 +94,65 @@ mod tests {
|
||||||
Some("mistral")
|
Some("mistral")
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn resolves_provider_from_prefix_when_no_explicit_provider() {
|
||||||
|
assert_eq!(
|
||||||
|
resolve_deployment_provider("mistral/mistral-ocr-latest", None).expect("splits"),
|
||||||
|
("mistral-ocr-latest".to_string(), "mistral".to_string())
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
resolve_deployment_provider("azure_ai/doc-intelligence/prebuilt-layout", None)
|
||||||
|
.expect("splits"),
|
||||||
|
(
|
||||||
|
"doc-intelligence/prebuilt-layout".to_string(),
|
||||||
|
"azure_ai".to_string()
|
||||||
|
)
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn explicit_provider_wins_and_strips_matching_prefix() {
|
||||||
|
assert_eq!(
|
||||||
|
resolve_deployment_provider("mistral/mistral-ocr-latest", Some("mistral"))
|
||||||
|
.expect("resolves"),
|
||||||
|
("mistral-ocr-latest".to_string(), "mistral".to_string())
|
||||||
|
);
|
||||||
|
assert_eq!(
|
||||||
|
resolve_deployment_provider("mistral-ocr-latest", Some("mistral")).expect("resolves"),
|
||||||
|
("mistral-ocr-latest".to_string(), "mistral".to_string())
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn explicit_provider_keeps_unrelated_prefix() {
|
||||||
|
assert_eq!(
|
||||||
|
resolve_deployment_provider("openai/some-model", Some("mistral")).expect("resolves"),
|
||||||
|
("openai/some-model".to_string(), "mistral".to_string())
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn blank_explicit_provider_falls_back_to_prefix() {
|
||||||
|
assert_eq!(
|
||||||
|
resolve_deployment_provider("mistral/mistral-ocr-latest", Some(" ")).expect("splits"),
|
||||||
|
("mistral-ocr-latest".to_string(), "mistral".to_string())
|
||||||
|
);
|
||||||
|
}
|
||||||
|
|
||||||
|
#[test]
|
||||||
|
fn rejects_model_without_resolvable_provider() {
|
||||||
|
assert!(matches!(
|
||||||
|
resolve_deployment_provider("mistral-ocr-latest", None),
|
||||||
|
Err(CoreError::InvalidProvider(_))
|
||||||
|
));
|
||||||
|
assert!(matches!(
|
||||||
|
resolve_deployment_provider("mistral/", None),
|
||||||
|
Err(CoreError::InvalidProvider(_))
|
||||||
|
));
|
||||||
|
assert!(matches!(
|
||||||
|
resolve_deployment_provider(" ", Some("mistral")),
|
||||||
|
Err(CoreError::InvalidProvider(_))
|
||||||
|
));
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -14,7 +14,7 @@
|
||||||
mod deployment;
|
mod deployment;
|
||||||
mod strategy;
|
mod strategy;
|
||||||
|
|
||||||
pub use deployment::{Deployment, LiteLLMParams};
|
pub use deployment::{resolve_deployment_provider, Deployment, LiteLLMParams};
|
||||||
pub use strategy::RoutingStrategy;
|
pub use strategy::RoutingStrategy;
|
||||||
|
|
||||||
/// Load-balancing router over a `model_list`.
|
/// Load-balancing router over a `model_list`.
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue