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:
Devin AI 2026-07-16 23:44:26 +00:00
parent c19e6382d6
commit f47752cee6
5 changed files with 145 additions and 110 deletions

View file

@ -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(&params.model, params.custom_llm_provider.as_deref())?; resolve_deployment_provider(&params.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"));
}
} }

View file

@ -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(_))
));
}
}

View file

@ -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;

View file

@ -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(_))
));
}
} }

View file

@ -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`.