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::get_llm_provider::get_llm_provider;
use litellm_core::router::Router;
use litellm_core::router::{resolve_deployment_provider, Router};
use litellm_core::CoreResult;
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 (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 {
model: requested_model,
@ -32,7 +31,7 @@ pub async fn run_ocr(router: &Router, call: OcrCall) -> CoreResult<Value> {
timeout,
} = call;
let mut response = ocr(OcrRequest {
let response = ocr(OcrRequest {
model: &provider_model,
document,
api_key: present(params.api_key.as_deref()),
@ -44,10 +43,22 @@ pub async fn run_ocr(router: &Router, call: OcrCall) -> CoreResult<Value> {
})
.await?;
if let Value::Object(map) = &mut response {
map.insert("model".to_string(), Value::String(requested_model));
}
Ok(response)
Ok(normalize_response_model(response, requested_model))
}
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)]
@ -62,4 +73,31 @@ mod tests {
assert_eq!(present(Some(" ")), 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 get_llm_provider;
pub mod ocr;
pub mod providers;
pub mod realtime;

View file

@ -3,6 +3,9 @@
use serde::Deserialize;
use crate::error::CoreError;
use crate::CoreResult;
/// Per-deployment call parameters, mirroring Python's `litellm_params`.
#[derive(Clone, Debug, Deserialize)]
pub struct LiteLLMParams {
@ -24,6 +27,40 @@ pub struct Deployment {
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)]
mod tests {
use super::*;
@ -57,4 +94,65 @@ mod tests {
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 strategy;
pub use deployment::{Deployment, LiteLLMParams};
pub use deployment::{resolve_deployment_provider, Deployment, LiteLLMParams};
pub use strategy::RoutingStrategy;
/// Load-balancing router over a `model_list`.