diff options
| author | sigoden <sigoden@gmail.com> | 2024-10-29 14:15:02 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-10-29 14:15:02 +0800 |
| commit | bf07cdf256e1e3008578972f5916a4d6763554f5 (patch) | |
| tree | 3f48ead051403a7f270198578e86a92672bb4208 /src/client | |
| parent | bb542f6e92fd782b70c44fddb21d013c289d0a3c (diff) | |
| download | aichat-bf07cdf256e1e3008578972f5916a4d6763554f5.tar.gz | |
refactor: improve code quality (#956)
Diffstat (limited to 'src/client')
| -rw-r--r-- | src/client/common.rs | 24 | ||||
| -rw-r--r-- | src/client/stream.rs | 12 |
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] { |
