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/common.rs | 39 +++++++++++++++++++++++---------------- 1 file changed, 23 insertions(+), 16 deletions(-) (limited to 'src/client/common.rs') 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()) { -- cgit v1.2.3