From b4a40e3fedb438570770a224b890ea24f6e660a9 Mon Sep 17 00:00:00 2001 From: sigoden Date: Sat, 18 May 2024 19:06:21 +0800 Subject: 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 --- src/client/cloudflare.rs | 17 ++++++++++------- 1 file changed, 10 insertions(+), 7 deletions(-) (limited to 'src/client/cloudflare.rs') 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, pub api_key: Option, #[serde(default)] - pub models: Vec, + pub models: Vec, pub extra: Option, } @@ -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 { 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 { 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 { Ok(body) } -fn extract_completion(data: &Value) -> Result<(String, CompletionDetails)> { +fn extract_completion(data: &Value) -> Result { 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)) } -- cgit v1.2.3