diff options
| author | sigoden <sigoden@gmail.com> | 2024-06-01 17:47:49 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-06-01 17:47:49 +0800 |
| commit | 571d1022f628cb7d2a3125664bd3293bac4471b5 (patch) | |
| tree | 67034aaf0a656ae9dce0cd0d3aae3605d2ea1767 /src/client/vertexai.rs | |
| parent | 259583f4f750e4ece7aed07858a9589110d7cf5d (diff) | |
| download | aichat-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.rs | 49 |
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, |
