summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
Diffstat (limited to 'src/client')
-rw-r--r--src/client/common.rs24
-rw-r--r--src/client/stream.rs12
2 files changed, 20 insertions, 16 deletions
diff --git a/src/client/common.rs b/src/client/common.rs
index 5fd4f24..fe12593 100644
--- a/src/client/common.rs
+++ b/src/client/common.rs
@@ -87,7 +87,7 @@ pub trait Client: Sync + Send {
handler.done();
ret.with_context(|| "Failed to call chat-completions api")
}
- _ = watch_abort_signal(abort_signal) => {
+ _ = wait_abort_signal(&abort_signal) => {
handler.done();
Ok(())
},
@@ -401,20 +401,25 @@ pub fn create_openai_compatible_client_config(client: &str) -> Result<Option<(St
pub async fn call_chat_completions(
input: &Input,
+ extract_code: bool,
client: &dyn Client,
- config: &GlobalConfig,
) -> Result<(String, Vec<ToolResult>)> {
let task = client.chat_completions(input.clone());
let ret = run_with_spinner(task, "Generating").await;
match ret {
Ok(ret) => {
let ChatCompletionsOutput {
- text, tool_calls, ..
+ mut text,
+ tool_calls,
+ ..
} = ret;
if !text.is_empty() {
- config.read().print_markdown(&text)?;
+ if extract_code && text.trim_start().starts_with("```") {
+ text = extract_block(&text);
+ }
+ client.global_config().read().print_markdown(&text)?;
}
- Ok((text, eval_tool_calls(config, tool_calls)?))
+ Ok((text, eval_tool_calls(client.global_config(), tool_calls)?))
}
Err(err) => Err(err),
}
@@ -423,15 +428,14 @@ pub async fn call_chat_completions(
pub async fn call_chat_completions_streaming(
input: &Input,
client: &dyn Client,
- config: &GlobalConfig,
- abort: AbortSignal,
+ abort_signal: AbortSignal,
) -> Result<(String, Vec<ToolResult>)> {
let (tx, rx) = unbounded_channel();
- let mut handler = SseHandler::new(tx, abort.clone());
+ let mut handler = SseHandler::new(tx, abort_signal.clone());
let (send_ret, render_ret) = tokio::join!(
client.chat_completions_streaming(input, &mut handler),
- render_stream(rx, config, abort.clone()),
+ render_stream(rx, client.global_config(), abort_signal.clone()),
);
render_ret?;
@@ -442,7 +446,7 @@ pub async fn call_chat_completions_streaming(
if !text.is_empty() && !text.ends_with('\n') {
println!();
}
- Ok((text, eval_tool_calls(config, tool_calls)?))
+ Ok((text, eval_tool_calls(client.global_config(), tool_calls)?))
}
Err(err) => {
if !text.is_empty() {
diff --git a/src/client/stream.rs b/src/client/stream.rs
index 735ec93..e89e7b9 100644
--- a/src/client/stream.rs
+++ b/src/client/stream.rs
@@ -10,16 +10,16 @@ use tokio::sync::mpsc::UnboundedSender;
pub struct SseHandler {
sender: UnboundedSender<SseEvent>,
- abort: AbortSignal,
+ abort_signal: AbortSignal,
buffer: String,
tool_calls: Vec<ToolCall>,
}
impl SseHandler {
- pub fn new(sender: UnboundedSender<SseEvent>, abort: AbortSignal) -> Self {
+ pub fn new(sender: UnboundedSender<SseEvent>, abort_signal: AbortSignal) -> Self {
Self {
sender,
- abort,
+ abort_signal,
buffer: String::new(),
tool_calls: Vec::new(),
}
@@ -36,7 +36,7 @@ impl SseHandler {
.send(SseEvent::Text(text.to_string()))
.with_context(|| "Failed to send SseEvent:Text");
if let Err(err) = ret {
- if self.abort.aborted() {
+ if self.abort_signal.aborted() {
return Ok(());
}
return Err(err);
@@ -48,7 +48,7 @@ impl SseHandler {
// debug!("HandleDone");
let ret = self.sender.send(SseEvent::Done);
if ret.is_err() {
- if self.abort.aborted() {
+ if self.abort_signal.aborted() {
return;
}
warn!("Failed to send SseEvent:Done");
@@ -62,7 +62,7 @@ impl SseHandler {
}
pub fn abort(&self) -> AbortSignal {
- self.abort.clone()
+ self.abort_signal.clone()
}
pub fn tool_calls(&self) -> &[ToolCall] {