summaryrefslogtreecommitdiffstats
path: root/src/client/ernie.rs
diff options
context:
space:
mode:
Diffstat (limited to 'src/client/ernie.rs')
-rw-r--r--src/client/ernie.rs47
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()),