From 2eed63f014deae539b22360c44c40e9ab2fef9a0 Mon Sep 17 00:00:00 2001 From: sigoden Date: Sat, 27 Jul 2024 17:14:19 +0800 Subject: feat: change model patch structure (#754) --- src/client/azure_openai.rs | 2 +- src/client/bedrock.rs | 2 +- src/client/claude.rs | 2 +- src/client/cloudflare.rs | 2 +- src/client/cohere.rs | 2 +- src/client/common.rs | 39 +++++++++++++++++++++++---------------- src/client/ernie.rs | 2 +- src/client/gemini.rs | 2 +- src/client/macros.rs | 4 ++-- src/client/ollama.rs | 2 +- src/client/openai.rs | 2 +- src/client/openai_compatible.rs | 2 +- src/client/qianwen.rs | 2 +- src/client/rag_dedicated.rs | 2 +- src/client/replicate.rs | 2 +- src/client/vertexai.rs | 2 +- 16 files changed, 39 insertions(+), 32 deletions(-) (limited to 'src') diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs index e3e9f16..2c4df05 100644 --- a/src/client/azure_openai.rs +++ b/src/client/azure_openai.rs @@ -12,7 +12,7 @@ pub struct AzureOpenAIConfig { pub api_key: Option, #[serde(default)] pub models: Vec, - pub patches: Option, + pub patch: Option, pub extra: Option, } diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs index 5d59ffe..5ae186d 100644 --- a/src/client/bedrock.rs +++ b/src/client/bedrock.rs @@ -28,7 +28,7 @@ pub struct BedrockConfig { pub region: Option, #[serde(default)] pub models: Vec, - pub patches: Option, + pub patch: Option, pub extra: Option, } diff --git a/src/client/claude.rs b/src/client/claude.rs index e239fc4..df0c034 100644 --- a/src/client/claude.rs +++ b/src/client/claude.rs @@ -13,7 +13,7 @@ pub struct ClaudeConfig { pub api_key: Option, #[serde(default)] pub models: Vec, - pub patches: Option, + pub patch: Option, pub extra: Option, } diff --git a/src/client/cloudflare.rs b/src/client/cloudflare.rs index 05891cc..439abac 100644 --- a/src/client/cloudflare.rs +++ b/src/client/cloudflare.rs @@ -14,7 +14,7 @@ pub struct CloudflareConfig { pub api_key: Option, #[serde(default)] pub models: Vec, - pub patches: Option, + pub patch: Option, pub extra: Option, } diff --git a/src/client/cohere.rs b/src/client/cohere.rs index b2857e7..4ab6f38 100644 --- a/src/client/cohere.rs +++ b/src/client/cohere.rs @@ -16,7 +16,7 @@ pub struct CohereConfig { pub api_key: Option, #[serde(default)] pub models: Vec, - pub patches: Option, + pub patch: Option, pub extra: Option, } diff --git a/src/client/common.rs b/src/client/common.rs index bfa747f..ea7aee5 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -31,7 +31,7 @@ pub trait Client: Sync + Send { fn extra_config(&self) -> Option<&ExtraConfig>; - fn patches_config(&self) -> Option<&ModelPatches>; + fn patch_config(&self) -> Option<&ModelPatch>; fn name(&self) -> &str; @@ -111,9 +111,12 @@ pub trait Client: Sync + Send { } fn patch_chat_completions_body(&self, body: &mut Value) { - if let Some(patch_data) = select_model_patch(self.patches_config().cloned(), self.model()) { - if body.is_object() && patch_data.chat_completions_body.is_object() { - json_patch::merge(body, &patch_data.chat_completions_body) + if let Some(patch) = extract_chat_completions_body_patch( + self.patch_config().map(|v| v.chat_completions_body.clone()), + self.model(), + ) { + if body.is_object() && patch.is_object() { + json_patch::merge(body, &patch) } } } @@ -160,21 +163,25 @@ pub struct ExtraConfig { pub connect_timeout: Option, } -pub type ModelPatches = IndexMap; - -#[derive(Debug, Clone, Deserialize)] +#[derive(Debug, Clone, Deserialize, Default)] pub struct ModelPatch { - #[serde(default)] - pub chat_completions_body: Value, + pub chat_completions_body: ChatCompletionsBodyPatch, } -pub fn select_model_patch(patches: Option, model: &Model) -> Option { - let patches: ModelPatches = - std::env::var(get_env_name(&format!("{}_patches", model.client_name()))) - .ok() - .and_then(|v| serde_json::from_str(&v).ok()) - .or(patches)?; - for (key, patch_data) in patches { +pub type ChatCompletionsBodyPatch = IndexMap; + +pub fn extract_chat_completions_body_patch( + patch: Option, + model: &Model, +) -> Option { + let patch = std::env::var(get_env_name(&format!( + "{}_chat_completions_body_patch", + model.client_name() + ))) + .ok() + .and_then(|v| serde_json::from_str(&v).ok()) + .or(patch)?; + for (key, patch_data) in patch { let key = ESCAPE_SLASH_RE.replace_all(&key, r"\/"); if let Ok(regex) = Regex::new(&format!("^({key})$")) { if let Ok(true) = regex.is_match(model.name()) { diff --git a/src/client/ernie.rs b/src/client/ernie.rs index fd5b907..f1d432d 100644 --- a/src/client/ernie.rs +++ b/src/client/ernie.rs @@ -19,7 +19,7 @@ pub struct ErnieConfig { pub secret_key: Option, #[serde(default)] pub models: Vec, - pub patches: Option, + pub patch: Option, pub extra: Option, } diff --git a/src/client/gemini.rs b/src/client/gemini.rs index 6382fb8..37f7c12 100644 --- a/src/client/gemini.rs +++ b/src/client/gemini.rs @@ -14,7 +14,7 @@ pub struct GeminiConfig { pub api_key: Option, #[serde(default)] pub models: Vec, - pub patches: Option, + pub patch: Option, pub extra: Option, } diff --git a/src/client/macros.rs b/src/client/macros.rs index 344bc92..778daf4 100644 --- a/src/client/macros.rs +++ b/src/client/macros.rs @@ -141,8 +141,8 @@ macro_rules! client_common_fns { self.config.extra.as_ref() } - fn patches_config(&self) -> Option<&$crate::client::ModelPatches> { - self.config.patches.as_ref() + fn patch_config(&self) -> Option<&$crate::client::ModelPatch> { + self.config.patch.as_ref() } fn name(&self) -> &str { diff --git a/src/client/ollama.rs b/src/client/ollama.rs index 5a3d201..299cae7 100644 --- a/src/client/ollama.rs +++ b/src/client/ollama.rs @@ -12,7 +12,7 @@ pub struct OllamaConfig { pub api_auth: Option, #[serde(default)] pub models: Vec, - pub patches: Option, + pub patch: Option, pub extra: Option, } diff --git a/src/client/openai.rs b/src/client/openai.rs index bfe896f..a494056 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -15,7 +15,7 @@ pub struct OpenAIConfig { pub organization_id: Option, #[serde(default)] pub models: Vec, - pub patches: Option, + pub patch: Option, pub extra: Option, } diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs index 132b40a..c59eff6 100644 --- a/src/client/openai_compatible.rs +++ b/src/client/openai_compatible.rs @@ -14,7 +14,7 @@ pub struct OpenAICompatibleConfig { pub chat_endpoint: Option, #[serde(default)] pub models: Vec, - pub patches: Option, + pub patch: Option, pub extra: Option, } diff --git a/src/client/qianwen.rs b/src/client/qianwen.rs index 3e2763b..e3aea31 100644 --- a/src/client/qianwen.rs +++ b/src/client/qianwen.rs @@ -27,7 +27,7 @@ pub struct QianwenConfig { pub api_key: Option, #[serde(default)] pub models: Vec, - pub patches: Option, + pub patch: Option, pub extra: Option, } diff --git a/src/client/rag_dedicated.rs b/src/client/rag_dedicated.rs index a9a9f18..19f1626 100644 --- a/src/client/rag_dedicated.rs +++ b/src/client/rag_dedicated.rs @@ -16,7 +16,7 @@ pub struct RagDedicatedConfig { pub api_key: Option, #[serde(default)] pub models: Vec, - pub patches: Option, + pub patch: Option, pub extra: Option, } diff --git a/src/client/replicate.rs b/src/client/replicate.rs index fd78ace..53097e6 100644 --- a/src/client/replicate.rs +++ b/src/client/replicate.rs @@ -16,7 +16,7 @@ pub struct ReplicateConfig { pub api_key: Option, #[serde(default)] pub models: Vec, - pub patches: Option, + pub patch: Option, pub extra: Option, } diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs index d93ccbd..4fc2ee1 100644 --- a/src/client/vertexai.rs +++ b/src/client/vertexai.rs @@ -19,7 +19,7 @@ pub struct VertexAIConfig { pub adc_file: Option, #[serde(default)] pub models: Vec, - pub patches: Option, + pub patch: Option, pub extra: Option, } -- cgit v1.2.3