summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2025-02-06 19:20:48 +0800
committerGitHub <noreply@github.com>2025-02-06 19:20:48 +0800
commita491820ab5ff9bb58e62fd25d88f31ab53a81db9 (patch)
tree86a84d7176a5fb16bce9b195808e952d181a8146
parent4812e446ee30467bd3ecce983a826f7649de7110 (diff)
downloadaichat-a491820ab5ff9bb58e62fd25d88f31ab53a81db9.tar.gz
feat: support model alias (#1150)
-rw-r--r--src/client/azure_openai.rs4
-rw-r--r--src/client/bedrock.rs4
-rw-r--r--src/client/claude.rs2
-rw-r--r--src/client/cohere.rs2
-rw-r--r--src/client/gemini.rs6
-rw-r--r--src/client/model.rs6
-rw-r--r--src/client/openai.rs4
-rw-r--r--src/client/openai_compatible.rs2
-rw-r--r--src/client/vertexai.rs10
9 files changed, 23 insertions, 17 deletions
diff --git a/src/client/azure_openai.rs b/src/client/azure_openai.rs
index fc856b2..61a54a2 100644
--- a/src/client/azure_openai.rs
+++ b/src/client/azure_openai.rs
@@ -50,7 +50,7 @@ fn prepare_chat_completions(
let url = format!(
"{}/openai/deployments/{}/chat/completions?api-version=2024-10-21",
&api_base,
- self_.model.name()
+ self_.model.real_name()
);
let body = openai_build_chat_completions_body(data, &self_.model);
@@ -69,7 +69,7 @@ fn prepare_embeddings(self_: &AzureOpenAIClient, data: &EmbeddingsData) -> Resul
let url = format!(
"{}/openai/deployments/{}/embeddings?api-version=2024-10-21",
&api_base,
- self_.model.name()
+ self_.model.real_name()
);
let body = openai_build_embeddings_body(data, &self_.model);
diff --git a/src/client/bedrock.rs b/src/client/bedrock.rs
index 435aa57..ebc0154 100644
--- a/src/client/bedrock.rs
+++ b/src/client/bedrock.rs
@@ -46,7 +46,7 @@ impl BedrockClient {
let region = self.get_region()?;
let host = format!("bedrock-runtime.{region}.amazonaws.com");
- let model_name = &self.model.name();
+ let model_name = &self.model.real_name();
let uri = if data.stream {
format!("/model/{model_name}/converse-stream")
@@ -95,7 +95,7 @@ impl BedrockClient {
let region = self.get_region()?;
let host = format!("bedrock-runtime.{region}.amazonaws.com");
- let uri = format!("/model/{}/invoke", self.model.name());
+ let uri = format!("/model/{}/invoke", self.model.real_name());
let input_type = match data.query {
true => "search_query",
diff --git a/src/client/claude.rs b/src/client/claude.rs
index 9a0e1f6..7167e80 100644
--- a/src/client/claude.rs
+++ b/src/client/claude.rs
@@ -248,7 +248,7 @@ pub fn claude_build_chat_completions_body(
}
let mut body = json!({
- "model": model.name(),
+ "model": model.real_name(),
"messages": messages,
});
if let Some(v) = system_message {
diff --git a/src/client/cohere.rs b/src/client/cohere.rs
index ae96977..a9c9216 100644
--- a/src/client/cohere.rs
+++ b/src/client/cohere.rs
@@ -76,7 +76,7 @@ fn prepare_embeddings(self_: &CohereClient, data: &EmbeddingsData) -> Result<Req
};
let body = json!({
- "model": self_.model.name(),
+ "model": self_.model.real_name(),
"texts": data.texts,
"input_type": input_type,
"embedding_types": ["float"],
diff --git a/src/client/gemini.rs b/src/client/gemini.rs
index 85917c3..77304e7 100644
--- a/src/client/gemini.rs
+++ b/src/client/gemini.rs
@@ -54,7 +54,7 @@ fn prepare_chat_completions(
let url = format!(
"{}/models/{}:{}?key={}",
api_base.trim_end_matches('/'),
- self_.model.name(),
+ self_.model.real_name(),
func,
api_key
);
@@ -75,11 +75,11 @@ fn prepare_embeddings(self_: &GeminiClient, data: &EmbeddingsData) -> Result<Req
let url = format!(
"{}/models/{}:batchEmbedContents?key={}",
api_base.trim_end_matches('/'),
- self_.model.name(),
+ self_.model.real_name(),
api_key
);
- let model_id = format!("models/{}", self_.model.name());
+ let model_id = format!("models/{}", self_.model.real_name());
let requests: Vec<_> = data
.texts
diff --git a/src/client/model.rs b/src/client/model.rs
index 8b8a290..d80a1d4 100644
--- a/src/client/model.rs
+++ b/src/client/model.rs
@@ -103,6 +103,10 @@ impl Model {
&self.data.name
}
+ pub fn real_name(&self) -> &str {
+ self.data.real_name.as_deref().unwrap_or(&self.data.name)
+ }
+
pub fn model_type(&self) -> ModelType {
if self.data.model_type.starts_with("embed") {
ModelType::Embedding
@@ -294,6 +298,8 @@ pub struct ModelData {
#[serde(default = "default_model_type", rename = "type")]
pub model_type: String,
#[serde(skip_serializing_if = "Option::is_none")]
+ pub real_name: Option<String>,
+ #[serde(skip_serializing_if = "Option::is_none")]
pub max_input_tokens: Option<usize>,
#[serde(skip_serializing_if = "Option::is_none")]
pub input_price: Option<f64>,
diff --git a/src/client/openai.rs b/src/client/openai.rs
index 7989c34..975afd7 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -294,7 +294,7 @@ pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Mod
.collect();
let mut body = json!({
- "model": &model.name(),
+ "model": &model.real_name(),
"messages": messages,
});
@@ -330,7 +330,7 @@ pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Mod
pub fn openai_build_embeddings_body(data: &EmbeddingsData, model: &Model) -> Value {
json!({
"input": data.texts,
- "model": model.name()
+ "model": model.real_name()
})
}
diff --git a/src/client/openai_compatible.rs b/src/client/openai_compatible.rs
index 30bc464..4f77e88 100644
--- a/src/client/openai_compatible.rs
+++ b/src/client/openai_compatible.rs
@@ -149,7 +149,7 @@ pub fn generic_build_rerank_body(data: &RerankData, model: &Model) -> Value {
} = data;
let mut body = json!({
- "model": model.name(),
+ "model": model.real_name(),
"query": query,
"documents": documents,
});
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index 19c8436..7cf5a1f 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -43,7 +43,7 @@ impl Client for VertexAIClient {
) -> Result<ChatCompletionsOutput> {
prepare_gcloud_access_token(client, self.name(), &self.config.adc_file).await?;
let model = self.model();
- let model_category = ModelCategory::from_str(model.name())?;
+ let model_category = ModelCategory::from_str(model.real_name())?;
let request_data = prepare_chat_completions(self, data, &model_category)?;
let builder = self.request_builder(client, request_data);
match model_category {
@@ -61,7 +61,7 @@ impl Client for VertexAIClient {
) -> Result<()> {
prepare_gcloud_access_token(client, self.name(), &self.config.adc_file).await?;
let model = self.model();
- let model_category = ModelCategory::from_str(model.name())?;
+ let model_category = ModelCategory::from_str(model.real_name())?;
let request_data = prepare_chat_completions(self, data, &model_category)?;
let builder = self.request_builder(client, request_data);
match model_category {
@@ -100,7 +100,7 @@ fn prepare_chat_completions(
let base_url = format!("https://{location}-aiplatform.googleapis.com/v1/projects/{project_id}/locations/{location}/publishers");
- let model_name = self_.model.name();
+ let model_name = self_.model.real_name();
let url = match model_category {
ModelCategory::Gemini => {
@@ -135,7 +135,7 @@ fn prepare_chat_completions(
ModelCategory::Mistral => {
let mut body = openai_build_chat_completions_body(data, &self_.model);
if let Some(body_obj) = body.as_object_mut() {
- body_obj["model"] = strip_model_version(self_.model.name()).into();
+ body_obj["model"] = strip_model_version(self_.model.real_name()).into();
}
body
}
@@ -154,7 +154,7 @@ fn prepare_embeddings(self_: &VertexAIClient, data: &EmbeddingsData) -> Result<R
let access_token = get_access_token(self_.name())?;
let base_url = format!("https://{location}-aiplatform.googleapis.com/v1/projects/{project_id}/locations/{location}/publishers");
- let url = format!("{base_url}/google/models/{}:predict", self_.model.name());
+ let url = format!("{base_url}/google/models/{}:predict", self_.model.real_name());
let instances: Vec<_> = data.texts.iter().map(|v| json!({"content": v})).collect();