summaryrefslogtreecommitdiffstats
path: root/src/client/cloudflare.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/cloudflare.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/cloudflare.rs')
-rw-r--r--src/client/cloudflare.rs17
1 files changed, 10 insertions, 7 deletions
diff --git a/src/client/cloudflare.rs b/src/client/cloudflare.rs
index 5a4bf8c..dfff009 100644
--- a/src/client/cloudflare.rs
+++ b/src/client/cloudflare.rs
@@ -1,5 +1,5 @@
use super::{
- catch_error, sse_stream, CloudflareClient, CompletionDetails, ExtraConfig, Model, ModelConfig,
+ catch_error, sse_stream, CloudflareClient, CompletionOutput, ExtraConfig, Model, ModelData,
PromptAction, PromptKind, SendData, SsMmessage, SseHandler,
};
@@ -16,7 +16,7 @@ pub struct CloudflareConfig {
pub account_id: Option<String>,
pub api_key: Option<String>,
#[serde(default)]
- pub models: Vec<ModelConfig>,
+ pub models: Vec<ModelData>,
pub extra: Option<ExtraConfig>,
}
@@ -37,7 +37,7 @@ impl CloudflareClient {
let url = format!(
"{API_BASE}/accounts/{account_id}/ai/run/{}",
- self.model.name
+ self.model.name()
);
debug!("Cloudflare Request: {url} {body}");
@@ -50,7 +50,7 @@ impl CloudflareClient {
impl_client_trait!(CloudflareClient, send_message, send_message_streaming);
-async fn send_message(builder: RequestBuilder) -> Result<(String, CompletionDetails)> {
+async fn send_message(builder: RequestBuilder) -> Result<CompletionOutput> {
let res = builder.send().await?;
let status = res.status();
let data: Value = res.json().await?;
@@ -58,6 +58,7 @@ async fn send_message(builder: RequestBuilder) -> Result<(String, CompletionDeta
catch_error(&data, status.as_u16())?;
}
+ debug!("non-stream-data: {data}");
extract_completion(&data)
}
@@ -67,6 +68,7 @@ async fn send_message_streaming(builder: RequestBuilder, handler: &mut SseHandle
return Ok(true);
}
let data: Value = serde_json::from_str(&message.data)?;
+ debug!("stream-data: {data}");
if let Some(text) = data["response"].as_str() {
handler.text(text)?;
}
@@ -80,11 +82,12 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
messages,
temperature,
top_p,
+ functions: _,
stream,
} = data;
let mut body = json!({
- "model": &model.name,
+ "model": &model.name(),
"messages": messages,
});
@@ -104,10 +107,10 @@ fn build_body(data: SendData, model: &Model) -> Result<Value> {
Ok(body)
}
-fn extract_completion(data: &Value) -> Result<(String, CompletionDetails)> {
+fn extract_completion(data: &Value) -> Result<CompletionOutput> {
let text = data["result"]["response"]
.as_str()
.ok_or_else(|| anyhow!("Invalid response data: {data}"))?;
- Ok((text.to_string(), CompletionDetails::default()))
+ Ok(CompletionOutput::new(text))
}