From bf07cdf256e1e3008578972f5916a4d6763554f5 Mon Sep 17 00:00:00 2001 From: sigoden Date: Tue, 29 Oct 2024 14:15:02 +0800 Subject: refactor: improve code quality (#956) --- src/client/common.rs | 24 ++++++++++++++---------- src/client/stream.rs | 12 ++++++------ src/main.rs | 45 ++++++++------------------------------------- src/rag/mod.rs | 6 +++--- src/render/mod.rs | 6 +++--- src/render/stream.rs | 37 +++++++++++++------------------------ src/repl/mod.rs | 5 ++--- src/utils/abort_signal.rs | 32 ++++++++++++++++++++++++++++---- 8 files changed, 77 insertions(+), 90 deletions(-) (limited to 'src') 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 Result<(String, Vec)> { 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)> { 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, - abort: AbortSignal, + abort_signal: AbortSignal, buffer: String, tool_calls: Vec, } impl SseHandler { - pub fn new(sender: UnboundedSender, abort: AbortSignal) -> Self { + pub fn new(sender: UnboundedSender, 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] { diff --git a/src/main.rs b/src/main.rs index f95be6a..ee95045 100644 --- a/src/main.rs +++ b/src/main.rs @@ -13,14 +13,12 @@ mod utils; extern crate log; use crate::cli::Cli; -use crate::client::{ - call_chat_completions, call_chat_completions_streaming, list_chat_models, ChatCompletionsOutput, -}; +use crate::client::{call_chat_completions, call_chat_completions_streaming, list_chat_models}; use crate::config::{ ensure_parent_exists, list_agents, load_env_file, Config, GlobalConfig, Input, WorkingMode, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE, TEMP_SESSION_NAME, }; -use crate::function::{eval_tool_calls, need_send_tool_results}; +use crate::function::need_send_tool_results; use crate::render::render_error; use crate::repl::Repl; use crate::utils::*; @@ -175,28 +173,9 @@ async fn start_directive( let extract_code = !*IS_STDOUT_TERMINAL && code_mode; config.write().before_chat_completion(&input)?; let (output, tool_results) = if !input.stream() || extract_code { - let task = client.chat_completions(input.clone()); - let ret = run_with_spinner(task, "Generating").await; - match ret { - Ok(ret) => { - let ChatCompletionsOutput { - mut text, - tool_calls, - .. - } = ret; - if !text.is_empty() { - if extract_code && text.trim_start().starts_with("```") { - text = extract_block(&text); - } - config.read().print_markdown(&text)?; - } - (text, eval_tool_calls(config, tool_calls)?) - } - Err(err) => return Err(err), - } + call_chat_completions(&input, extract_code, client.as_ref()).await? } else { - call_chat_completions_streaming(&input, client.as_ref(), config, abort_signal.clone()) - .await? + call_chat_completions_streaming(&input, client.as_ref(), abort_signal.clone()).await? }; config .write() @@ -225,15 +204,7 @@ async fn start_interactive(config: &GlobalConfig) -> Result<()> { async fn shell_execute(config: &GlobalConfig, shell: &Shell, mut input: Input) -> Result<()> { let client = input.create_client()?; config.write().before_chat_completion(&input)?; - let ret = if *IS_STDOUT_TERMINAL { - let spinner = create_spinner("Generating").await; - let ret = client.chat_completions(input.clone()).await; - spinner.stop(); - ret - } else { - client.chat_completions(input.clone()).await - }; - let mut eval_str = ret?.text; + let (mut eval_str, _) = call_chat_completions(&input, false, client.as_ref()).await?; if let Ok(true) = CODE_BLOCK_RE.is_match(&eval_str) { eval_str = extract_block(&eval_str); } @@ -287,12 +258,12 @@ async fn shell_execute(config: &GlobalConfig, shell: &Shell, mut input: Input) - "d" => { let role = config.read().retrieve_role(EXPLAIN_SHELL_ROLE)?; let input = Input::from_str(config, &eval_str, Some(role)); - let abort = create_abort_signal(); + let abort_signal = create_abort_signal(); if input.stream() { - call_chat_completions_streaming(&input, client.as_ref(), config, abort) + call_chat_completions_streaming(&input, client.as_ref(), abort_signal) .await?; } else { - call_chat_completions(&input, client.as_ref(), config).await?; + call_chat_completions(&input, false, client.as_ref()).await?; } println!(); continue; diff --git a/src/rag/mod.rs b/src/rag/mod.rs index 48ace8e..fd734e2 100644 --- a/src/rag/mod.rs +++ b/src/rag/mod.rs @@ -93,7 +93,7 @@ impl Rag { spinner.stop(); ret?; } - _ = watch_abort_signal(abort_signal) => { + _ = wait_abort_signal(&abort_signal) => { spinner.stop(); bail!("Aborted!") }, @@ -142,7 +142,7 @@ impl Rag { spinner.stop(); ret?; } - _ = watch_abort_signal(abort_signal) => { + _ = wait_abort_signal(&abort_signal) => { spinner.stop(); bail!("Aborted!") }, @@ -320,7 +320,7 @@ impl Rag { ret = self.hybird_search(text, top_k, min_score_vector_search, min_score_keyword_search, rerank_model) => { ret } - _ = watch_abort_signal(abort_signal) => { + _ = wait_abort_signal(&abort_signal) => { bail!("Aborted!") }, }; diff --git a/src/render/mod.rs b/src/render/mod.rs index 97161f8..9d93203 100644 --- a/src/render/mod.rs +++ b/src/render/mod.rs @@ -13,14 +13,14 @@ use tokio::sync::mpsc::UnboundedReceiver; pub async fn render_stream( rx: UnboundedReceiver, config: &GlobalConfig, - abort: AbortSignal, + abort_signal: AbortSignal, ) -> Result<()> { let ret = if *IS_STDOUT_TERMINAL { let render_options = config.read().render_options()?; let mut render = MarkdownRender::init(render_options)?; - markdown_stream(rx, &mut render, &abort).await + markdown_stream(rx, &mut render, &abort_signal).await } else { - raw_stream(rx, &abort).await + raw_stream(rx, &abort_signal).await }; ret.map_err(|err| err.context("Failed to reader stream")) } diff --git a/src/render/stream.rs b/src/render/stream.rs index 2264c16..07a1a18 100644 --- a/src/render/stream.rs +++ b/src/render/stream.rs @@ -1,12 +1,10 @@ use super::{MarkdownRender, SseEvent}; -use crate::utils::{create_spinner, AbortSignal}; +use crate::utils::{create_spinner, poll_abort_signal, AbortSignal}; use anyhow::Result; use crossterm::{ - cursor, - event::{self, Event, KeyCode, KeyModifiers}, - queue, style, + cursor, queue, style, terminal::{self, disable_raw_mode, enable_raw_mode}, }; use std::{ @@ -19,12 +17,12 @@ use tokio::sync::mpsc::UnboundedReceiver; pub async fn markdown_stream( rx: UnboundedReceiver, render: &mut MarkdownRender, - abort: &AbortSignal, + abort_signal: &AbortSignal, ) -> Result<()> { enable_raw_mode()?; let mut stdout = io::stdout(); - let ret = markdown_stream_inner(rx, render, abort, &mut stdout).await; + let ret = markdown_stream_inner(rx, render, abort_signal, &mut stdout).await; disable_raw_mode()?; @@ -34,9 +32,12 @@ pub async fn markdown_stream( ret } -pub async fn raw_stream(mut rx: UnboundedReceiver, abort: &AbortSignal) -> Result<()> { +pub async fn raw_stream( + mut rx: UnboundedReceiver, + abort_signal: &AbortSignal, +) -> Result<()> { loop { - if abort.aborted() { + if abort_signal.aborted() { return Ok(()); } if let Some(evt) = rx.recv().await { @@ -57,7 +58,7 @@ pub async fn raw_stream(mut rx: UnboundedReceiver, abort: &AbortSignal async fn markdown_stream_inner( mut rx: UnboundedReceiver, render: &mut MarkdownRender, - abort: &AbortSignal, + abort_signal: &AbortSignal, writer: &mut Stdout, ) -> Result<()> { let mut buffer = String::new(); @@ -68,7 +69,7 @@ async fn markdown_stream_inner( let mut spinner = Some(create_spinner("Generating").await); 'outer: loop { - if abort.aborted() { + if abort_signal.aborted() { return Ok(()); } for reply_event in gather_events(&mut rx).await { @@ -141,20 +142,8 @@ async fn markdown_stream_inner( } } - if crossterm::event::poll(Duration::from_millis(25))? { - if let Event::Key(key) = event::read()? { - match key.code { - KeyCode::Char('c') if key.modifiers == KeyModifiers::CONTROL => { - abort.set_ctrlc(); - break; - } - KeyCode::Char('d') if key.modifiers == KeyModifiers::CONTROL => { - abort.set_ctrld(); - break; - } - _ => {} - } - } + if poll_abort_signal(abort_signal)? { + break; } } diff --git a/src/repl/mod.rs b/src/repl/mod.rs index e2181cc..b392eb4 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -658,10 +658,9 @@ async fn ask( let client = input.create_client()?; config.write().before_chat_completion(&input)?; let (output, tool_results) = if input.stream() { - call_chat_completions_streaming(&input, client.as_ref(), config, abort_signal.clone()) - .await? + call_chat_completions_streaming(&input, client.as_ref(), abort_signal.clone()).await? } else { - call_chat_completions(&input, client.as_ref(), config).await? + call_chat_completions(&input, false, client.as_ref()).await? }; config .write() diff --git a/src/utils/abort_signal.rs b/src/utils/abort_signal.rs index 5329a9e..a13e5bb 100644 --- a/src/utils/abort_signal.rs +++ b/src/utils/abort_signal.rs @@ -1,6 +1,11 @@ -use std::sync::{ - atomic::{AtomicBool, Ordering}, - Arc, +use anyhow::Result; +use crossterm::event::{self, Event, KeyCode, KeyModifiers}; +use std::{ + sync::{ + atomic::{AtomicBool, Ordering}, + Arc, + }, + time::Duration, }; pub type AbortSignal = Arc; @@ -54,7 +59,7 @@ impl AbortSignalInner { } } -pub async fn watch_abort_signal(abort_signal: AbortSignal) { +pub async fn wait_abort_signal(abort_signal: &AbortSignal) { loop { if abort_signal.aborted() { break; @@ -62,3 +67,22 @@ pub async fn watch_abort_signal(abort_signal: AbortSignal) { tokio::time::sleep(std::time::Duration::from_millis(25)).await; } } + +pub fn poll_abort_signal(abort_signal: &AbortSignal) -> Result { + if crossterm::event::poll(Duration::from_millis(25))? { + if let Event::Key(key) = event::read()? { + match key.code { + KeyCode::Char('c') if key.modifiers == KeyModifiers::CONTROL => { + abort_signal.set_ctrlc(); + return Ok(true); + } + KeyCode::Char('d') if key.modifiers == KeyModifiers::CONTROL => { + abort_signal.set_ctrld(); + return Ok(true); + } + _ => {} + } + } + } + Ok(false) +} -- cgit v1.2.3