summaryrefslogtreecommitdiffstats
path: root/src
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
parent3c21fac3665698bec860a49a66f5967e4449dc95 (diff)
downloadaichat-9dfbdafe9f2c662c1d9dae96f3b850f4de5bca19.tar.gz
feat: add model field `patch` (#1169)
Diffstat (limited to 'src')
-rw-r--r--src/client/common.rs12
-rw-r--r--src/client/model.rs7
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")]