summaryrefslogtreecommitdiffstats
path: root/src/client/vertexai.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-01 17:47:49 +0800
committerGitHub <noreply@github.com>2024-06-01 17:47:49 +0800
commit571d1022f628cb7d2a3125664bd3293bac4471b5 (patch)
tree67034aaf0a656ae9dce0cd0d3aae3605d2ea1767 /src/client/vertexai.rs
parent259583f4f750e4ece7aed07858a9589110d7cf5d (diff)
downloadaichat-571d1022f628cb7d2a3125664bd3293bac4471b5.tar.gz
refactor: rename some client structs and methods (#555)
* rename `Completeion*` to `ChatCompletions*` * rename `send_message*` to `chat_completions*` * rename `request_builder` to `chat_completions_builder` * rename `build_body` to `build_chat_completions_body` * rename `extract_completion` to `extract_chat_completions` * format * remove unused config fields
Diffstat (limited to 'src/client/vertexai.rs')
-rw-r--r--src/client/vertexai.rs49
1 files changed, 25 insertions, 24 deletions
diff --git a/src/client/vertexai.rs b/src/client/vertexai.rs
index 102b5cb..b40247d 100644
--- a/src/client/vertexai.rs
+++ b/src/client/vertexai.rs
@@ -1,7 +1,7 @@
use super::{
- access_token::*, catch_error, json_stream, message::*, patch_system_message, Client,
- CompletionData, CompletionOutput, ExtraConfig, Model, ModelData, ModelPatches, PromptAction,
- PromptKind, SseHandler, ToolCall, VertexAIClient,
+ access_token::*, catch_error, json_stream, message::*, patch_system_message,
+ ChatCompletionsData, ChatCompletionsOutput, Client, ExtraConfig, Model, ModelData,
+ ModelPatches, PromptAction, PromptKind, SseHandler, ToolCall, VertexAIClient,
};
use anyhow::{anyhow, bail, Context, Result};
@@ -18,8 +18,6 @@ pub struct VertexAIConfig {
pub project_id: Option<String>,
pub location: Option<String>,
pub adc_file: Option<String>,
- #[serde(rename = "safetySettings")]
- pub safety_settings: Option<Value>,
#[serde(default)]
pub models: Vec<ModelData>,
pub patches: Option<ModelPatches>,
@@ -35,10 +33,10 @@ impl VertexAIClient {
("location", "Location", true, PromptKind::String),
];
- fn request_builder(
+ fn chat_completions_builder(
&self,
client: &ReqwestClient,
- data: CompletionData,
+ data: ChatCompletionsData,
) -> Result<RequestBuilder> {
let project_id = self.get_project_id()?;
let location = self.get_location()?;
@@ -52,7 +50,7 @@ impl VertexAIClient {
};
let url = format!("{base_url}/google/models/{}:{func}", self.model.name());
- let mut body = gemini_build_body(data, &self.model)?;
+ let mut body = gemini_build_chat_completions_body(data, &self.model)?;
self.patch_request_body(&mut body);
debug!("VertexAI Request: {url} {body}");
@@ -67,29 +65,29 @@ impl VertexAIClient {
impl Client for VertexAIClient {
client_common_fns!();
- async fn send_message_inner(
+ async fn chat_completions_inner(
&self,
client: &ReqwestClient,
- data: CompletionData,
- ) -> Result<CompletionOutput> {
+ data: ChatCompletionsData,
+ ) -> Result<ChatCompletionsOutput> {
prepare_gcloud_access_token(client, self.name(), &self.config.adc_file).await?;
- let builder = self.request_builder(client, data)?;
- gemini_send_message(builder).await
+ let builder = self.chat_completions_builder(client, data)?;
+ gemini_chat_completions(builder).await
}
- async fn send_message_streaming_inner(
+ async fn chat_completions_streaming_inner(
&self,
client: &ReqwestClient,
handler: &mut SseHandler,
- data: CompletionData,
+ data: ChatCompletionsData,
) -> Result<()> {
prepare_gcloud_access_token(client, self.name(), &self.config.adc_file).await?;
- let builder = self.request_builder(client, data)?;
- gemini_send_message_streaming(builder, handler).await
+ let builder = self.chat_completions_builder(client, data)?;
+ gemini_chat_completions_streaming(builder, handler).await
}
}
-pub async fn gemini_send_message(builder: RequestBuilder) -> Result<CompletionOutput> {
+pub async fn gemini_chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutput> {
let res = builder.send().await?;
let status = res.status();
let data: Value = res.json().await?;
@@ -97,10 +95,10 @@ pub async fn gemini_send_message(builder: RequestBuilder) -> Result<CompletionOu
catch_error(&data, status.as_u16())?;
}
debug!("non-stream-data: {data}");
- gemini_extract_completion_text(&data)
+ gemini_extract_chat_completions_text(&data)
}
-pub async fn gemini_send_message_streaming(
+pub async fn gemini_chat_completions_streaming(
builder: RequestBuilder,
handler: &mut SseHandler,
) -> Result<()> {
@@ -140,7 +138,7 @@ pub async fn gemini_send_message_streaming(
Ok(())
}
-fn gemini_extract_completion_text(data: &Value) -> Result<CompletionOutput> {
+fn gemini_extract_chat_completions_text(data: &Value) -> Result<ChatCompletionsOutput> {
let text = data["candidates"][0]["content"]["parts"][0]["text"]
.as_str()
.unwrap_or_default();
@@ -171,7 +169,7 @@ fn gemini_extract_completion_text(data: &Value) -> Result<CompletionOutput> {
bail!("Invalid response data: {data}");
}
}
- let output = CompletionOutput {
+ let output = ChatCompletionsOutput {
text: text.to_string(),
tool_calls,
id: None,
@@ -181,8 +179,11 @@ fn gemini_extract_completion_text(data: &Value) -> Result<CompletionOutput> {
Ok(output)
}
-pub(crate) fn gemini_build_body(data: CompletionData, model: &Model) -> Result<Value> {
- let CompletionData {
+pub(crate) fn gemini_build_chat_completions_body(
+ data: ChatCompletionsData,
+ model: &Model,
+) -> Result<Value> {
+ let ChatCompletionsData {
mut messages,
temperature,
top_p,