mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
refactor(rust): type provider model namespaces
This commit is contained in:
parent
c80617c4a6
commit
f4a6f695c9
5 changed files with 258 additions and 16 deletions
|
|
@ -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, ¶ms)?;
|
||||
mapping::transform_ocr_request(provider_model(&request.model)?, document, ¶ms)?;
|
||||
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!(
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
220
litellm-rust/crates/core/src/providers/model.rs
Normal file
220
litellm-rust/crates/core/src/providers/model.rs
Normal 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());
|
||||
}
|
||||
}
|
||||
|
|
@ -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(), ¶ms).unwrap(),
|
||||
transform_ocr_request(
|
||||
provider_model("deepseek-ai/deepseek-ocr-maas").unwrap(),
|
||||
document(),
|
||||
¶ms,
|
||||
)
|
||||
.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(),
|
||||
)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue