summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-07-27 17:14:19 +0800
committerGitHub <noreply@github.com>2024-07-27 17:14:19 +0800
commit2eed63f014deae539b22360c44c40e9ab2fef9a0 (patch)
tree00d8dd93a01d59ec7f3a8b20cb7a9cfb3d2ee93e
parent2576c04f7d03747550a0c0f1b7ba62b9656a792c (diff)
downloadaichat-2eed63f014deae539b22360c44c40e9ab2fef9a0.tar.gz
feat: change model patch structure (#754)
-rw-r--r--config.example.yaml19
-rw-r--r--src/client/azure_openai.rs2
-rw-r--r--src/client/bedrock.rs2
-rw-r--r--src/client/claude.rs2
-rw-r--r--src/client/cloudflare.rs2
-rw-r--r--src/client/cohere.rs2
-rw-r--r--src/client/common.rs39
-rw-r--r--src/client/ernie.rs2
-rw-r--r--src/client/gemini.rs2
-rw-r--r--src/client/macros.rs4
-rw-r--r--src/client/ollama.rs2
-rw-r--r--src/client/openai.rs2
-rw-r--r--src/client/openai_compatible.rs2
-rw-r--r--src/client/qianwen.rs2
-rw-r--r--src/client/rag_dedicated.rs2
-rw-r--r--src/client/replicate.rs2
-rw-r--r--src/client/vertexai.rs2
17 files changed, 49 insertions, 41 deletions
diff --git a/config.example.yaml b/config.example.yaml
index 9e51c13..96b001c 100644
--- a/config.example.yaml
+++ b/config.example.yaml
@@ -95,9 +95,10 @@ clients:
# - name: xxxx # Reranker model
# type: reranker
# max_input_tokens: 2048
- # patches:
- # <regex>: # The regex to match model names, e.g. '.*' 'gpt-4o' 'gpt-4o|gpt-4-.*'
- # chat_completions_body: # The JSON to be merged with the chat completions request body.
+ # patch: # Patch api request
+ # chat_completions_body:
+ # <regex>: # The regex to match model names, e.g. '.*' 'gpt-4o' 'gpt-4o|gpt-4-.*'
+ # <json> # The JSON to be merged with the chat completions request body.
# extra:
# proxy: socks5://127.0.0.1:1080 # Set https/socks5 proxy. ENV: HTTPS_PROXY/https_proxy/ALL_PROXY/all_proxy
# connect_timeout: 10 # Set timeout in seconds for connect to api
@@ -121,9 +122,9 @@ clients:
# See https://ai.google.dev/docs
- type: gemini
api_key: xxx # ENV: {client}_API_KEY
- patches:
- '.*':
- chat_completions_body:
+ patch:
+ chat_completions_body:
+ '.*':
safetySettings:
- category: HARM_CATEGORY_HARASSMENT
threshold: BLOCK_NONE
@@ -187,9 +188,9 @@ clients:
# Run `gcloud auth application-default login` to init the adc file
# see https://cloud.google.com/docs/authentication/external/set-up-adc
adc_file: <path-to/gcloud/application_default_credentials.json>
- patches:
- 'gemini-.*':
- chat_completions_body:
+ patch:
+ chat_completions_body:
+ 'gemini-.*':
safetySettings:
- category: HARM_CATEGORY_HARASSMENT
threshold: BLOCK_ONLY_HIGH
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<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patches: Option<ModelPatches>,
+ pub patch: Option<ModelPatch>,
pub extra: Option<ExtraConfig>,
}
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<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patches: Option<ModelPatches>,
+ pub patch: Option<ModelPatch>,
pub extra: Option<ExtraConfig>,
}
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<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patches: Option<ModelPatches>,
+ pub patch: Option<ModelPatch>,
pub extra: Option<ExtraConfig>,
}
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<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patches: Option<ModelPatches>,
+ pub patch: Option<ModelPatch>,
pub extra: Option<ExtraConfig>,
}
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<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patches: Option<ModelPatches>,
+ pub patch: Option<ModelPatch>,
pub extra: Option<ExtraConfig>,
}
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<u64>,
}
-pub type ModelPatches = IndexMap<String, ModelPatch>;
-
-#[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<ModelPatches>, model: &Model) -> Option<ModelPatch> {
- 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<String, Value>;
+
+pub fn extract_chat_completions_body_patch(
+ patch: Option<ChatCompletionsBodyPatch>,
+ model: &Model,
+) -> Option<Value> {
+ 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<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patches: Option<ModelPatches>,
+ pub patch: Option<ModelPatch>,
pub extra: Option<ExtraConfig>,
}
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<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patches: Option<ModelPatches>,
+ pub patch: Option<ModelPatch>,
pub extra: Option<ExtraConfig>,
}
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<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patches: Option<ModelPatches>,
+ pub patch: Option<ModelPatch>,
pub extra: Option<ExtraConfig>,
}
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<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patches: Option<ModelPatches>,
+ pub patch: Option<ModelPatch>,
pub extra: Option<ExtraConfig>,
}
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<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patches: Option<ModelPatches>,
+ pub patch: Option<ModelPatch>,
pub extra: Option<ExtraConfig>,
}
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<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patches: Option<ModelPatches>,
+ pub patch: Option<ModelPatch>,
pub extra: Option<ExtraConfig>,
}
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<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patches: Option<ModelPatches>,
+ pub patch: Option<ModelPatch>,
pub extra: Option<ExtraConfig>,
}
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<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patches: Option<ModelPatches>,
+ pub patch: Option<ModelPatch>,
pub extra: Option<ExtraConfig>,
}
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<String>,
#[serde(default)]
pub models: Vec<ModelData>,
- pub patches: Option<ModelPatches>,
+ pub patch: Option<ModelPatch>,
pub extra: Option<ExtraConfig>,
}