summaryrefslogtreecommitdiffstats
path: root/src/client/vertexai.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-12-04 21:03:59 +0800
committerGitHub <noreply@github.com>2024-12-04 21:03:59 +0800
commit3a3388375be05758d5f5574cce1d0f8eb7bdd604 (patch)
treecc8e19e218c20f28a04feca33850fce4bb2cbe98 /src/client/vertexai.rs
parent7d42fe9429f75d195f865b07cef10d040d5397f2 (diff)
downloadaichat-3a3388375be05758d5f5574cce1d0f8eb7bdd604.tar.gz
refactor: improve retrieve model (#1036)
- check the model type while retrieve model - select chat/reranker model even if it is missed in client models - find predefined-models for openai-compatible client with startsWith - remove client::ApiType
Diffstat (limited to 'src/client/vertexai.rs')
-rw-r--r--src/client/vertexai.rs6
1 files changed, 3 insertions, 3 deletions
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index 07ff2a4..7b73164 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -45,7 +45,7 @@ impl Client for VertexAIClient {
let model = self.model();
let model_category = ModelCategory::from_str(model.name())?;
let request_data = prepare_chat_completions(self, data, &model_category)?;
- let builder = self.request_builder(client, request_data, ApiType::ChatCompletions);
+ let builder = self.request_builder(client, request_data);
match model_category {
ModelCategory::Gemini => gemini_chat_completions(builder, model).await,
ModelCategory::Claude => claude_chat_completions(builder, model).await,
@@ -63,7 +63,7 @@ impl Client for VertexAIClient {
let model = self.model();
let model_category = ModelCategory::from_str(model.name())?;
let request_data = prepare_chat_completions(self, data, &model_category)?;
- let builder = self.request_builder(client, request_data, ApiType::ChatCompletions);
+ let builder = self.request_builder(client, request_data);
match model_category {
ModelCategory::Gemini => {
gemini_chat_completions_streaming(builder, handler, model).await
@@ -84,7 +84,7 @@ impl Client for VertexAIClient {
) -> Result<Vec<Vec<f32>>> {
prepare_gcloud_access_token(client, self.name(), &self.config.adc_file).await?;
let request_data = prepare_embeddings(self, data)?;
- let builder = self.request_builder(client, request_data, ApiType::Embeddings);
+ let builder = self.request_builder(client, request_data);
embeddings(builder, self.model()).await
}
}