summaryrefslogtreecommitdiffstats
path: root/src/client/ernie.rs
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-05-18 19:06:21 +0800
committerGitHub <noreply@github.com>2024-05-18 19:06:21 +0800
commitb4a40e3fedb438570770a224b890ea24f6e660a9 (patch)
tree344b96102da7cbedf1034d023aa82599940388b1 /src/client/ernie.rs
parent1348a62e5f8bc140a7218fbfe1b73f990ab16101 (diff)
downloadaichat-b4a40e3fedb438570770a224b890ea24f6e660a9.tar.gz
feat: support function calling (#514)
* feat: support function calling * fix on Windows OS * implement multi-steps function calling * fix on Windows OS * add error for client not support function calling * refactor message data structure and make claude client supporting function calling * support reuse previous call results * improve error handling for function calling * use prefix `may_` as indicator for `execute` type fucntions
Diffstat (limited to 'src/client/ernie.rs')
-rw-r--r--src/client/ernie.rs23
1 files changed, 14 insertions, 9 deletions
diff --git a/src/client/ernie.rs b/src/client/ernie.rs
index 28cb857..1d79138 100644
--- a/src/client/ernie.rs
+++ b/src/client/ernie.rs
@@ -1,7 +1,7 @@
use super::access_token::*;
use super::{
- maybe_catch_error, patch_system_message, sse_stream, Client, CompletionDetails, ErnieClient,
- ExtraConfig, Model, ModelConfig, PromptAction, PromptKind, SendData, SsMmessage, SseHandler,
+ maybe_catch_error, patch_system_message, sse_stream, Client, CompletionOutput, ErnieClient,
+ ExtraConfig, Model, ModelData, PromptAction, PromptKind, SendData, SsMmessage, SseHandler,
};
use anyhow::{anyhow, Context, Result};
@@ -20,7 +20,7 @@ pub struct ErnieConfig {
pub api_key: Option<String>,
pub secret_key: Option<String>,
#[serde(default)]
- pub models: Vec<ModelConfig>,
+ pub models: Vec<ModelData>,
pub extra: Option<ExtraConfig>,
}
@@ -36,7 +36,7 @@ impl ErnieClient {
let url = format!(
"{API_BASE}/wenxinworkshop/chat/{}?access_token={access_token}",
- &self.model.name,
+ &self.model.name(),
);
debug!("Ernie Request: {url} {body}");
@@ -78,7 +78,7 @@ impl Client for ErnieClient {
&self,
client: &ReqwestClient,
data: SendData,
- ) -> Result<(String, CompletionDetails)> {
+ ) -> Result<CompletionOutput> {
self.prepare_access_token().await?;
let builder = self.request_builder(client, data)?;
send_message(builder).await
@@ -96,15 +96,17 @@ impl Client for ErnieClient {
}
}
-async fn send_message(builder: RequestBuilder) -> Result<(String, CompletionDetails)> {
+async fn send_message(builder: RequestBuilder) -> Result<CompletionOutput> {
let data: Value = builder.send().await?.json().await?;
maybe_catch_error(&data)?;
+ debug!("non-stream-data: {data}");
extract_completion_text(&data)
}
async fn send_message_streaming(builder: RequestBuilder, handler: &mut SseHandler) -> Result<()> {
let handle = |message: SsMmessage| -> Result<bool> {
let data: Value = serde_json::from_str(&message.data)?;
+ debug!("stream-data: {data}");
if let Some(text) = data["result"].as_str() {
handler.text(text)?;
}
@@ -119,6 +121,7 @@ fn build_body(data: SendData, model: &Model) -> Value {
mut messages,
temperature,
top_p,
+ functions: _,
stream,
} = data;
@@ -145,16 +148,18 @@ fn build_body(data: SendData, model: &Model) -> Value {
body
}
-fn extract_completion_text(data: &Value) -> Result<(String, CompletionDetails)> {
+fn extract_completion_text(data: &Value) -> Result<CompletionOutput> {
let text = data["result"]
.as_str()
.ok_or_else(|| anyhow!("Invalid response data: {data}"))?;
- let details = CompletionDetails {
+ let output = CompletionOutput {
+ text: text.to_string(),
+ tool_calls: vec![],
id: data["id"].as_str().map(|v| v.to_string()),
input_tokens: data["usage"]["prompt_tokens"].as_u64(),
output_tokens: data["usage"]["completion_tokens"].as_u64(),
};
- Ok((text.to_string(), details))
+ Ok(output)
}
async fn fetch_access_token(