summaryrefslogtreecommitdiffstats
path: root/src/client/common.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2025-02-13 19:08:36 +0800
committerGitHub <noreply@github.com>2025-02-13 19:08:36 +0800
commit9dfbdafe9f2c662c1d9dae96f3b850f4de5bca19 (patch)
tree786634906e4c7df7aafe3c2df2ca83a5b93522db /src/client/common.rs
parent3c21fac3665698bec860a49a66f5967e4449dc95 (diff)
downloadaichat-9dfbdafe9f2c662c1d9dae96f3b850f4de5bca19.tar.gz
feat: add model field `patch` (#1169)
Diffstat (limited to 'src/client/common.rs')
-rw-r--r--src/client/common.rs12
1 files changed, 9 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);
}
}
}