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/ernie.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/ernie.rs')
| -rw-r--r-- | src/client/ernie.rs | 47 |
1 files changed, 25 insertions, 22 deletions
diff --git a/src/client/ernie.rs b/src/client/ernie.rs index 49e158b..097ee68 100644 --- a/src/client/ernie.rs +++ b/src/client/ernie.rs @@ -1,7 +1,7 @@ use super::{ - access_token::*, maybe_catch_error, patch_system_message, sse_stream, Client, CompletionData, - CompletionOutput, ErnieClient, ExtraConfig, Model, ModelData, ModelPatches, PromptAction, - PromptKind, SseHandler, SseMmessage, + access_token::*, maybe_catch_error, patch_system_message, sse_stream, ChatCompletionsData, + ChatCompletionsOutput, Client, ErnieClient, ExtraConfig, Model, ModelData, ModelPatches, + PromptAction, PromptKind, SseHandler, SseMmessage, }; use anyhow::{anyhow, Context, Result}; @@ -31,12 +31,12 @@ impl ErnieClient { ("secret_key", "Secret Key:", true, PromptKind::String), ]; - fn request_builder( + fn chat_completions_builder( &self, client: &ReqwestClient, - data: CompletionData, + data: ChatCompletionsData, ) -> Result<RequestBuilder> { - let mut body = build_body(data, &self.model); + let mut body = build_chat_completions_body(data, &self.model); self.patch_request_body(&mut body); let access_token = get_access_token(self.name())?; @@ -81,36 +81,39 @@ impl ErnieClient { impl Client for ErnieClient { client_common_fns!(); - async fn send_message_inner( + async fn chat_completions_inner( &self, client: &ReqwestClient, - data: CompletionData, - ) -> Result<CompletionOutput> { + data: ChatCompletionsData, + ) -> Result<ChatCompletionsOutput> { self.prepare_access_token().await?; - let builder = self.request_builder(client, data)?; - send_message(builder).await + let builder = self.chat_completions_builder(client, data)?; + 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<()> { self.prepare_access_token().await?; - let builder = self.request_builder(client, data)?; - send_message_streaming(builder, handler).await + let builder = self.chat_completions_builder(client, data)?; + chat_completions_streaming(builder, handler).await } } -async fn send_message(builder: RequestBuilder) -> Result<CompletionOutput> { +async fn chat_completions(builder: RequestBuilder) -> Result<ChatCompletionsOutput> { let data: Value = builder.send().await?.json().await?; maybe_catch_error(&data)?; debug!("non-stream-data: {data}"); - extract_completion_text(&data) + extract_chat_completions_text(&data) } -async fn send_message_streaming(builder: RequestBuilder, handler: &mut SseHandler) -> Result<()> { +async fn chat_completions_streaming( + builder: RequestBuilder, + handler: &mut SseHandler, +) -> Result<()> { let handle = |message: SseMmessage| -> Result<bool> { let data: Value = serde_json::from_str(&message.data)?; debug!("stream-data: {data}"); @@ -123,8 +126,8 @@ async fn send_message_streaming(builder: RequestBuilder, handler: &mut SseHandle sse_stream(builder, handle).await } -fn build_body(data: CompletionData, model: &Model) -> Value { - let CompletionData { +fn build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Value { + let ChatCompletionsData { mut messages, temperature, top_p, @@ -155,11 +158,11 @@ fn build_body(data: CompletionData, model: &Model) -> Value { body } -fn extract_completion_text(data: &Value) -> Result<CompletionOutput> { +fn extract_chat_completions_text(data: &Value) -> Result<ChatCompletionsOutput> { let text = data["result"] .as_str() .ok_or_else(|| anyhow!("Invalid response data: {data}"))?; - let output = CompletionOutput { + let output = ChatCompletionsOutput { text: text.to_string(), tool_calls: vec![], id: data["id"].as_str().map(|v| v.to_string()), |
