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/sse_handler.rs | 23 ++++++++++++++++++----- 1 file changed, 18 insertions(+), 5 deletions(-) (limited to 'src/client/sse_handler.rs') diff --git a/src/client/sse_handler.rs b/src/client/sse_handler.rs index dbf78c2..ddbdbcd 100644 --- a/src/client/sse_handler.rs +++ b/src/client/sse_handler.rs @@ -3,10 +3,13 @@ use crate::utils::AbortSignal; use anyhow::{Context, Result}; use tokio::sync::mpsc::UnboundedSender; +use super::ToolCall; + pub struct SseHandler { sender: UnboundedSender, - buffer: String, abort: AbortSignal, + buffer: String, + tool_calls: Vec, } impl SseHandler { @@ -15,11 +18,12 @@ impl SseHandler { sender, abort, buffer: String::new(), + tool_calls: Vec::new(), } } pub fn text(&mut self, text: &str) -> Result<()> { - // debug!("ReplyText: {}", text); + // debug!("HandleText: {}", text); if text.is_empty() { return Ok(()); } @@ -33,7 +37,7 @@ impl SseHandler { } pub fn done(&mut self) -> Result<()> { - // debug!("ReplyDone"); + // debug!("HandleDone"); let ret = self .sender .send(SseEvent::Done) @@ -42,14 +46,23 @@ impl SseHandler { Ok(()) } - pub fn get_buffer(&self) -> &str { - &self.buffer + pub fn tool_call(&mut self, call: ToolCall) -> Result<()> { + // debug!("HandleCall: {:?}", call); + self.tool_calls.push(call); + Ok(()) } pub fn get_abort(&self) -> AbortSignal { self.abort.clone() } + pub fn take(self) -> (String, Vec) { + let Self { + buffer, tool_calls, .. + } = self; + (buffer, tool_calls) + } + fn safe_ret(&self, ret: Result<()>) -> Result<()> { if ret.is_err() && self.abort.aborted() { return Ok(()); -- cgit v1.2.3