diff options
| author | sigoden <sigoden@gmail.com> | 2025-02-13 19:08:36 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2025-02-13 19:08:36 +0800 |
| commit | 9dfbdafe9f2c662c1d9dae96f3b850f4de5bca19 (patch) | |
| tree | 786634906e4c7df7aafe3c2df2ca83a5b93522db | |
| parent | 3c21fac3665698bec860a49a66f5967e4449dc95 (diff) | |
| download | aichat-9dfbdafe9f2c662c1d9dae96f3b850f4de5bca19.tar.gz | |
feat: add model field `patch` (#1169)
| -rw-r--r-- | src/client/common.rs | 12 | ||||
| -rw-r--r-- | src/client/model.rs | 7 |
2 files changed, 16 insertions, 3 deletions
diff --git a/src/client/common.rs b/src/client/common.rs index 4db87f2..27f66cb 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -154,7 +154,11 @@ pub trait Client: Sync + Send { fn patch_request_data(&self, request_data: &mut RequestData) { let model_type = self.model().model_type(); - let map = std::env::var(get_env_name(&format!( + if let Some(patch) = self.model().patch() { + request_data.apply_patch(patch.clone()); + } + + let patch_map = std::env::var(get_env_name(&format!( "patch_{}_{}", self.model().client_name(), model_type.api_name(), @@ -166,11 +170,11 @@ pub trait Client: Sync + Send { .and_then(|v| model_type.extract_patch(v)) .cloned() }); - let map = match map { + let patch_map = match patch_map { Some(v) => v, _ => return, }; - for (key, patch) in map { + for (key, patch) in patch_map { let key = ESCAPE_SLASH_RE.replace_all(&key, r"\/"); if let Ok(regex) = Regex::new(&format!("^({key})$")) { if let Ok(true) = regex.is_match(self.model().name()) { @@ -260,6 +264,8 @@ impl RequestData { for (key, value) in patch_headers { if let Some(value) = value.as_str() { self.header(key, value) + } else if value.is_null() { + self.headers.swap_remove(key); } } } diff --git a/src/client/model.rs b/src/client/model.rs index 6681856..a2865ab 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -9,6 +9,7 @@ use crate::utils::{estimate_token_length, strip_think_tag}; use anyhow::{bail, Result}; use serde::{Deserialize, Serialize}; +use serde_json::Value; use std::fmt::Display; const PER_MESSAGES_TOKENS: usize = 5; @@ -178,6 +179,10 @@ impl Model { } } + pub fn patch(&self) -> Option<&Value> { + self.data.patch.as_ref() + } + pub fn max_input_tokens(&self) -> Option<usize> { self.data.max_input_tokens } @@ -313,6 +318,8 @@ pub struct ModelData { pub input_price: Option<f64>, #[serde(skip_serializing_if = "Option::is_none")] pub output_price: Option<f64>, + #[serde(skip_serializing_if = "Option::is_none")] + pub patch: Option<Value>, // chat-only properties #[serde(skip_serializing_if = "Option::is_none")] |
