summaryrefslogtreecommitdiffstats
path: root/src/client/vertexai.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-09-24 07:42:24 +0800
committerGitHub <noreply@github.com>2024-09-24 07:42:24 +0800
commit912773c25a113f49c5df63cd3a8086d38c75103e (patch)
tree7d653edf095f0949d1eeb2f83ec3e5616972c3b1 /src/client/vertexai.rs
parent00c4a6e421f01590ffc3fa5601e93d5ec755fca7 (diff)
downloadaichat-912773c25a113f49c5df63cd3a8086d38c75103e.tar.gz
refactor: embeddings/rerank fn accept ref data (#878)
Diffstat (limited to 'src/client/vertexai.rs')
-rw-r--r--src/client/vertexai.rs10
1 files changed, 3 insertions, 7 deletions
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index 5349eaf..3ce73e0 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -80,7 +80,7 @@ impl Client for VertexAIClient {
async fn embeddings_inner(
&self,
client: &ReqwestClient,
- data: EmbeddingsData,
+ data: &EmbeddingsData,
) -> Result<Vec<Vec<f32>>> {
prepare_gcloud_access_token(client, self.name(), &self.config.adc_file).await?;
let request_data = prepare_embeddings(self, data)?;
@@ -148,7 +148,7 @@ fn prepare_chat_completions(
Ok(request_data)
}
-fn prepare_embeddings(self_: &VertexAIClient, data: EmbeddingsData) -> Result<RequestData> {
+fn prepare_embeddings(self_: &VertexAIClient, data: &EmbeddingsData) -> Result<RequestData> {
let project_id = self_.get_project_id()?;
let location = self_.get_location()?;
let access_token = get_access_token(self_.name())?;
@@ -156,11 +156,7 @@ fn prepare_embeddings(self_: &VertexAIClient, data: EmbeddingsData) -> Result<Re
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 instances: Vec<_> = data
- .texts
- .into_iter()
- .map(|v| json!({"content": v}))
- .collect();
+ let instances: Vec<_> = data.texts.iter().map(|v| json!({"content": v})).collect();
let body = json!({
"instances": instances,