refactor(rust): type provider model namespaces

This commit is contained in:
Yujong Lee 2026-09-15 10:11:37 -07:00
parent c80617c4a6
commit f4a6f695c9
5 changed files with 258 additions and 16 deletions

View file

@ -2,6 +2,9 @@ mod types {
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use super::DeepSeekAi;
use crate::providers::model::ProviderModel;
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
pub(crate) struct DeepSeekOcrParams {
#[serde(skip_serializing_if = "Option::is_none")]
@ -27,7 +30,7 @@ mod types {
#[derive(Clone, Debug, Serialize, Deserialize)]
pub(crate) struct DeepSeekOcrRequest {
pub model: String,
pub model: ProviderModel<DeepSeekAi>,
pub messages: Vec<DeepSeekOcrMessage>,
#[serde(flatten)]
pub params: DeepSeekOcrParams,
@ -105,13 +108,15 @@ mod mapping {
use serde::de::IntoDeserializer;
use serde_json::{Value, json};
use super::DeepSeekAi;
use super::types::*;
use crate::ocr::error::{OcrRequestError, OcrResponseError};
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
use crate::providers::model::ProviderModel;
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(crate) fn transform_ocr_request(
provider_model: &str,
provider_model: ProviderModel<DeepSeekAi>,
document: OcrDocument,
params: &DeepSeekOcrParams,
) -> Result<DeepSeekOcrRequest, OcrRequestError> {
@ -123,7 +128,7 @@ mod mapping {
extra_fields: serde_json::Map::new(),
};
Ok(DeepSeekOcrRequest {
model: provider_model.to_string(),
model: provider_model,
messages: vec![DeepSeekOcrMessage {
role: UserRole::User,
content: vec![content],
@ -220,12 +225,20 @@ use crate::ocr::prepare::{
_prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body,
};
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
use crate::providers::model::{ModelNamespace, ProviderModel, RoutedModel};
use crate::url_utils::ApiUrl;
use litellm_auth_gcp::{self as vertex, VertexConfig};
const DEFAULT_API_BASE: &str = "https://aiplatform.googleapis.com";
const MODEL_NAMESPACE: &str = "deepseek-ai";
const DEFAULT_LOCATION: &str = "us-central1";
#[derive(Clone, Debug)]
pub(crate) struct DeepSeekAi;
impl ModelNamespace for DeepSeekAi {
const NAME: &'static str = MODEL_NAMESPACE;
}
#[derive(Clone, Debug)]
pub(crate) struct VertexAIDeepSeekOCRConfig;
@ -270,7 +283,7 @@ impl BaseOcrConfig for VertexAIDeepSeekOCRConfig {
)?;
let document = request.document.clone();
let body =
mapping::transform_ocr_request(&provider_model(&request.model), document, &params)?;
mapping::transform_ocr_request(provider_model(&request.model)?, document, &params)?;
transform_request_body(
client,
request,
@ -292,12 +305,12 @@ impl BaseOcrConfig for VertexAIDeepSeekOCRConfig {
}
}
fn provider_model(model: &str) -> String {
if model.starts_with(&format!("{MODEL_NAMESPACE}/")) {
model.to_string()
} else {
format!("{MODEL_NAMESPACE}/{model}")
}
pub(crate) fn provider_model(model: &str) -> Result<ProviderModel<DeepSeekAi>, OcrRequestError> {
RoutedModel::new(model)
.and_then(RoutedModel::into_provider::<DeepSeekAi>)
.map_err(|_| OcrRequestError::RequestField {
path: "model".into(),
})
}
fn get_complete_url(
@ -339,11 +352,13 @@ mod tests {
#[test]
fn config_owns_model_namespace_and_endpoint() {
assert_eq!(
provider_model("deepseek-ocr-maas"),
provider_model("deepseek-ocr-maas").unwrap().as_str(),
"deepseek-ai/deepseek-ocr-maas"
);
assert_eq!(
provider_model("deepseek-ai/deepseek-ocr-maas"),
provider_model("deepseek-ai/deepseek-ocr-maas")
.unwrap()
.as_str(),
"deepseek-ai/deepseek-ocr-maas"
);
assert_eq!(

View file

@ -8,7 +8,7 @@ use crate::llms::mistral::ocr::transformation::MistralOCRConfig;
use crate::llms::reducto::ocr::transformation::{ReductoParseLegacyConfig, ReductoParseV3Config};
use crate::llms::vertex_ai::ocr::deepseek_transformation::VertexAIDeepSeekOCRConfig;
use crate::llms::vertex_ai::ocr::transformation::VertexAIOCRConfig;
use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use crate::providers::custom_llm_provider::{CustomLlmProvider, get_custom_llm_provider};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum OcrConfigKind {

View file

@ -3,4 +3,5 @@ pub mod azure_ai;
#[cfg(feature = "bedrock-auth")]
pub mod bedrock;
pub mod custom_llm_provider;
pub(crate) mod model;
pub mod openai;

View file

@ -0,0 +1,220 @@
use std::marker::PhantomData;
use serde::{Deserialize, Deserializer, Serialize, Serializer, de::Error as _};
use thiserror::Error;
#[derive(Clone, Copy, Debug, Eq, PartialEq, Error)]
pub(crate) enum ModelNameError {
#[error("model name cannot be empty")]
EmptyModel,
#[error("model namespace must be one non-empty path segment: {0}")]
InvalidNamespace(&'static str),
}
pub(crate) trait ModelNamespace {
const NAME: &'static str;
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) struct RoutedModel<'a>(&'a str);
impl<'a> RoutedModel<'a> {
pub(crate) fn new(value: &'a str) -> Result<Self, ModelNameError> {
if value.is_empty() {
return Err(ModelNameError::EmptyModel);
}
Ok(Self(value))
}
pub(crate) fn into_provider<N: ModelNamespace>(
self,
) -> Result<ProviderModel<N>, ModelNameError> {
let namespace = N::NAME;
if namespace.is_empty() || namespace.contains('/') {
return Err(ModelNameError::InvalidNamespace(namespace));
}
let prefix = format!("{namespace}/");
let local_model = self.0.trim_start_matches(prefix.as_str());
if local_model.is_empty() {
return Err(ModelNameError::EmptyModel);
}
Ok(ProviderModel {
value: format!("{prefix}{local_model}"),
namespace: PhantomData,
})
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub(crate) struct ProviderModel<N> {
value: String,
namespace: PhantomData<N>,
}
impl<N> ProviderModel<N> {
#[cfg(test)]
pub(crate) fn as_str(&self) -> &str {
&self.value
}
}
impl<N> Serialize for ProviderModel<N> {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
self.value.serialize(serializer)
}
}
impl<'de, N: ModelNamespace> Deserialize<'de> for ProviderModel<N> {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let value = String::deserialize(deserializer)?;
RoutedModel::new(&value)
.and_then(RoutedModel::into_provider::<N>)
.map_err(D::Error::custom)
}
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::*;
#[derive(Clone, Debug, Eq, PartialEq)]
struct DeepSeekAi;
impl ModelNamespace for DeepSeekAi {
const NAME: &'static str = "deepseek-ai";
}
#[derive(Clone, Debug, Eq, PartialEq)]
struct FalAi;
impl ModelNamespace for FalAi {
const NAME: &'static str = "fal-ai";
}
#[test]
fn qualifies_a_bare_model() {
let model = RoutedModel::new("deepseek-ocr-maas")
.and_then(RoutedModel::into_provider::<DeepSeekAi>)
.unwrap();
assert_eq!(model.as_str(), "deepseek-ai/deepseek-ocr-maas");
}
#[test]
fn preserves_an_already_qualified_model() {
let model = RoutedModel::new("deepseek-ai/deepseek-ocr-maas")
.and_then(RoutedModel::into_provider::<DeepSeekAi>)
.unwrap();
assert_eq!(model.as_str(), "deepseek-ai/deepseek-ocr-maas");
}
#[test]
fn collapses_repeated_owned_namespaces() {
let model = RoutedModel::new("deepseek-ai/deepseek-ai/deepseek-ai/deepseek-ocr-maas")
.and_then(RoutedModel::into_provider::<DeepSeekAi>)
.unwrap();
assert_eq!(model.as_str(), "deepseek-ai/deepseek-ocr-maas");
}
#[test]
fn matches_the_namespace_as_a_complete_segment() {
let model = RoutedModel::new("deepseek-ai-v2/model")
.and_then(RoutedModel::into_provider::<DeepSeekAi>)
.unwrap();
assert_eq!(model.as_str(), "deepseek-ai/deepseek-ai-v2/model");
}
#[test]
fn preserves_nested_provider_model_paths() {
let model = RoutedModel::new("publishers/vendor/models/model-v1")
.and_then(RoutedModel::into_provider::<DeepSeekAi>)
.unwrap();
assert_eq!(
model.as_str(),
"deepseek-ai/publishers/vendor/models/model-v1"
);
}
#[test]
fn namespace_markers_select_different_wire_names() {
let routed = RoutedModel::new("model-v1").unwrap();
let deepseek = routed.into_provider::<DeepSeekAi>().unwrap();
let fal = routed.into_provider::<FalAi>().unwrap();
assert_eq!(deepseek.as_str(), "deepseek-ai/model-v1");
assert_eq!(fal.as_str(), "fal-ai/model-v1");
}
#[test]
fn rejects_empty_routed_models() {
assert_eq!(RoutedModel::new(""), Err(ModelNameError::EmptyModel));
}
#[test]
fn rejects_a_namespace_without_a_model() {
let result =
RoutedModel::new("deepseek-ai/").and_then(RoutedModel::into_provider::<DeepSeekAi>);
assert_eq!(result, Err(ModelNameError::EmptyModel));
}
#[test]
fn rejects_invalid_namespace_markers() {
struct Empty;
impl ModelNamespace for Empty {
const NAME: &'static str = "";
}
struct MultipleSegments;
impl ModelNamespace for MultipleSegments {
const NAME: &'static str = "one/two";
}
assert!(matches!(
RoutedModel::new("model").and_then(RoutedModel::into_provider::<Empty>),
Err(ModelNameError::InvalidNamespace(""))
));
assert!(matches!(
RoutedModel::new("model").and_then(RoutedModel::into_provider::<MultipleSegments>),
Err(ModelNameError::InvalidNamespace("one/two"))
));
}
#[test]
fn provider_models_serialize_as_plain_strings() {
let model = RoutedModel::new("deepseek-ocr-maas")
.and_then(RoutedModel::into_provider::<DeepSeekAi>)
.unwrap();
assert_eq!(
serde_json::to_value(model).unwrap(),
json!("deepseek-ai/deepseek-ocr-maas")
);
}
#[test]
fn deserialization_reestablishes_the_namespace_invariant() {
let model: ProviderModel<DeepSeekAi> =
serde_json::from_value(json!("deepseek-ai/deepseek-ai/model-v1")).unwrap();
assert_eq!(model.as_str(), "deepseek-ai/model-v1");
}
#[test]
fn deserialization_rejects_missing_model_names() {
let result = serde_json::from_value::<ProviderModel<DeepSeekAi>>(json!("deepseek-ai/"));
assert!(result.is_err());
}
}

View file

@ -2,7 +2,8 @@ use rstest::rstest;
use serde_json::{Value, json};
use crate::llms::vertex_ai::ocr::deepseek_transformation::{
DeepSeekOcrParams, DeepSeekOcrResponse, transform_ocr_request, transform_ocr_response,
DeepSeekOcrParams, DeepSeekOcrResponse, provider_model, transform_ocr_request,
transform_ocr_response,
};
use crate::ocr::types::OcrDocument;
@ -22,7 +23,12 @@ fn request_mapping_matches_python(#[case] name: &str, #[case] value: Value) {
let params: DeepSeekOcrParams =
serde_json::from_value(json!({name: value.clone(), "ignored": true})).unwrap();
let result = serde_json::to_value(
transform_ocr_request("deepseek-ai/deepseek-ocr-maas", document(), &params).unwrap(),
transform_ocr_request(
provider_model("deepseek-ai/deepseek-ocr-maas").unwrap(),
document(),
&params,
)
.unwrap(),
)
.unwrap();
assert_eq!(result["model"], "deepseek-ai/deepseek-ocr-maas");
@ -44,7 +50,7 @@ fn request_maps_both_document_types_to_image_content(#[case] document: Value) {
.unwrap()
.clone();
let request = transform_ocr_request(
"deepseek-ai/deepseek-ocr-maas",
provider_model("deepseek-ai/deepseek-ocr-maas").unwrap(),
serde_json::from_value(document).unwrap(),
&DeepSeekOcrParams::default(),
)