summaryrefslogtreecommitdiffstats
path: root/src/client/sse_handler.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/sse_handler.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/sse_handler.rs')
-rw-r--r--src/client/sse_handler.rs23
1 files changed, 18 insertions, 5 deletions
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<SseEvent>,
- buffer: String,
abort: AbortSignal,
+ buffer: String,
+ tool_calls: Vec<ToolCall>,
}
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<ToolCall>) {
+ 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(());