diff options
| author | sigoden <sigoden@gmail.com> | 2024-06-17 09:14:54 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-06-17 09:14:54 +0800 |
| commit | 638bf3276613265da761b4d74f269fa77c8f4bb1 (patch) | |
| tree | 9667711b92bc3664ad72c179cfe1cf4bbf54d322 /src | |
| parent | 12872b3d2956a26fb7d1cd464b3672ebe40d94c1 (diff) | |
| download | aichat-638bf3276613265da761b4d74f269fa77c8f4bb1.tar.gz | |
refactor: improve code quatity (#604)
Diffstat (limited to 'src')
| -rw-r--r-- | src/client/common.rs | 6 | ||||
| -rw-r--r-- | src/config/input.rs | 8 | ||||
| -rw-r--r-- | src/config/mod.rs | 126 | ||||
| -rw-r--r-- | src/function.rs | 15 | ||||
| -rw-r--r-- | src/main.rs | 25 | ||||
| -rw-r--r-- | src/repl/mod.rs | 21 |
6 files changed, 108 insertions, 93 deletions
diff --git a/src/client/common.rs b/src/client/common.rs index 4ff476b..bc84fe8 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -2,7 +2,7 @@ use super::*; use crate::{ config::{GlobalConfig, Input}, - function::{eval_tool_calls, FunctionDeclaration, ToolCall, ToolCallResult}, + function::{eval_tool_calls, FunctionDeclaration, ToolCall, ToolResult}, render::{render_error, render_stream}, utils::{ prompt_input_integer, prompt_input_string, tokenize, watch_abort_signal, AbortSignal, @@ -505,12 +505,12 @@ pub fn create_openai_compatible_client_config(client: &str) -> Result<Option<(St } } -pub async fn send_stream( +pub async fn chat_completion_streaming( input: &Input, client: &dyn Client, config: &GlobalConfig, abort: AbortSignal, -) -> Result<(String, Vec<ToolCallResult>)> { +) -> Result<(String, Vec<ToolResult>)> { let (tx, rx) = unbounded_channel(); let mut handler = SseHandler::new(tx, abort.clone()); diff --git a/src/config/input.rs b/src/config/input.rs index 0c93c1f..48b0359 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -4,7 +4,7 @@ use crate::client::{ init_client, ChatCompletionsData, Client, ImageUrl, Message, MessageContent, MessageContentPart, MessageRole, Model, }; -use crate::function::{ToolCallResult, ToolResults}; +use crate::function::{ToolResult, ToolResults}; use crate::utils::{base64_encode, sha256, AbortSignal}; use anyhow::{bail, Context, Result}; @@ -154,11 +154,7 @@ impl Input { self.patched_text.take(); } - pub fn merge_tool_call( - mut self, - output: String, - tool_call_results: Vec<ToolCallResult>, - ) -> Self { + pub fn merge_tool_call(mut self, output: String, tool_call_results: Vec<ToolResult>) -> Self { match self.tool_call.as_mut() { Some(exist_tool_call_results) => { exist_tool_call_results.0.extend(tool_call_results); diff --git a/src/config/mod.rs b/src/config/mod.rs index a905ae8..efbf984 100644 --- a/src/config/mod.rs +++ b/src/config/mod.rs @@ -12,7 +12,7 @@ use crate::client::{ create_client_config, list_chat_models, list_client_types, ClientConfig, Model, OPENAI_COMPATIBLE_PLATFORMS, }; -use crate::function::{FunctionDeclaration, Functions, FunctionsFilter, ToolCallResult}; +use crate::function::{FunctionDeclaration, Functions, FunctionsFilter, ToolResult}; use crate::rag::Rag; use crate::render::{MarkdownRender, RenderOptions}; use crate::utils::*; @@ -229,60 +229,6 @@ impl Config { Ok(path) } - pub fn save_message( - &mut self, - input: &mut Input, - output: &str, - tool_call_results: &[ToolCallResult], - ) -> Result<()> { - input.clear_patch_text(); - self.last_message = Some((input.clone(), output.to_string(), self.bot.is_some())); - - if self.dry_run || output.is_empty() || !tool_call_results.is_empty() { - return Ok(()); - } - - if let Some(session) = input.session_mut(&mut self.session) { - session.add_message(input, output)?; - return Ok(()); - } - - if !self.save { - return Ok(()); - } - let mut file = self.open_message_file()?; - if output.is_empty() || !self.save { - return Ok(()); - } - let timestamp = now(); - let summary = input.summary(); - let input_markdown = input.render(); - let scope = if self.bot.is_none() { - let role_name = if input.role().is_derived() { - None - } else { - Some(input.role().name()) - }; - match (role_name, input.rag_name()) { - (Some(role), Some(rag_name)) => format!(" ({role}#{rag_name})"), - (Some(role), _) => format!(" ({role})"), - (None, Some(rag_name)) => format!(" (#{rag_name})"), - _ => String::new(), - } - } else { - String::new() - }; - let output = format!("# CHAT: {summary} [{timestamp}]{scope}\n{input_markdown}\n--------\n{output}\n--------\n\n",); - file.write_all(output.as_bytes()) - .with_context(|| "Failed to save message") - } - - pub fn maybe_copy(&self, text: &str) { - if self.auto_copy { - let _ = set_text(text); - } - } - pub fn config_file() -> Result<PathBuf> { match env::var(get_env_name("config_file")) { Ok(value) => Ok(PathBuf::from(value)), @@ -838,6 +784,7 @@ impl Config { if let Some(session) = self.session.as_mut() { session.set_compressing(false); } + self.last_message = None; } pub async fn use_rag( @@ -1246,6 +1193,75 @@ impl Config { output } + pub fn before_chat_completion(&mut self, input: &Input) -> Result<()> { + self.last_message = Some((input.clone(), String::new(), self.bot.is_some())); + Ok(()) + } + + pub fn after_chat_completion( + &mut self, + input: &mut Input, + output: &str, + tool_results: &[ToolResult], + ) -> Result<()> { + input.clear_patch_text(); + self.last_message = Some((input.clone(), output.to_string(), self.bot.is_some())); + self.save_message(input, output, tool_results)?; + self.maybe_copy(output); + Ok(()) + } + + fn save_message( + &mut self, + input: &mut Input, + output: &str, + tool_results: &[ToolResult], + ) -> Result<()> { + if self.dry_run || output.is_empty() || !tool_results.is_empty() { + return Ok(()); + } + + if let Some(session) = input.session_mut(&mut self.session) { + session.add_message(input, output)?; + return Ok(()); + } + + if !self.save { + return Ok(()); + } + let mut file = self.open_message_file()?; + if output.is_empty() || !self.save { + return Ok(()); + } + let timestamp = now(); + let summary = input.summary(); + let input_markdown = input.render(); + let scope = if self.bot.is_none() { + let role_name = if input.role().is_derived() { + None + } else { + Some(input.role().name()) + }; + match (role_name, input.rag_name()) { + (Some(role), Some(rag_name)) => format!(" ({role}#{rag_name})"), + (Some(role), _) => format!(" ({role})"), + (None, Some(rag_name)) => format!(" (#{rag_name})"), + _ => String::new(), + } + } else { + String::new() + }; + let output = format!("# CHAT: {summary} [{timestamp}]{scope}\n{input_markdown}\n--------\n{output}\n--------\n\n",); + file.write_all(output.as_bytes()) + .with_context(|| "Failed to save message") + } + + fn maybe_copy(&self, text: &str) { + if self.auto_copy { + let _ = set_text(text); + } + } + fn open_message_file(&self) -> Result<File> { let path = self.messages_file()?; ensure_parent_exists(&path)?; diff --git a/src/function.rs b/src/function.rs index b2833e8..bab8ca9 100644 --- a/src/function.rs +++ b/src/function.rs @@ -16,13 +16,10 @@ use std::{ }; pub const SELECTED_ALL_FUNCTIONS: &str = ".*"; -pub type ToolResults = (Vec<ToolCallResult>, String); +pub type ToolResults = (Vec<ToolResult>, String); pub type FunctionsFilter = String; -pub fn eval_tool_calls( - config: &GlobalConfig, - mut calls: Vec<ToolCall>, -) -> Result<Vec<ToolCallResult>> { +pub fn eval_tool_calls(config: &GlobalConfig, mut calls: Vec<ToolCall>) -> Result<Vec<ToolResult>> { let mut output = vec![]; if calls.is_empty() { return Ok(output); @@ -33,22 +30,22 @@ pub fn eval_tool_calls( } for call in calls { let result = call.eval(config)?; - output.push(ToolCallResult::new(call, result)); + output.push(ToolResult::new(call, result)); } Ok(output) } -pub fn need_send_call_results(arr: &[ToolCallResult]) -> bool { +pub fn need_send_tool_results(arr: &[ToolResult]) -> bool { arr.iter().any(|v| !v.output.is_null()) } #[derive(Debug, Clone, Deserialize, Serialize)] -pub struct ToolCallResult { +pub struct ToolResult { pub call: ToolCall, pub output: Value, } -impl ToolCallResult { +impl ToolResult { pub fn new(call: ToolCall, output: Value) -> Self { Self { call, output } } diff --git a/src/main.rs b/src/main.rs index 3f197eb..d9f258a 100644 --- a/src/main.rs +++ b/src/main.rs @@ -14,12 +14,12 @@ mod utils; extern crate log; use crate::cli::Cli; -use crate::client::{list_chat_models, send_stream, ChatCompletionsOutput}; +use crate::client::{chat_completion_streaming, list_chat_models, ChatCompletionsOutput}; use crate::config::{ list_bots, Config, GlobalConfig, Input, WorkingMode, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE, TEMP_SESSION_NAME, }; -use crate::function::{eval_tool_calls, need_send_call_results}; +use crate::function::{eval_tool_calls, need_send_tool_results}; use crate::render::{render_error, MarkdownRender}; use crate::repl::Repl; use crate::utils::*; @@ -169,7 +169,8 @@ async fn start_directive( ) -> Result<()> { let client = input.create_client()?; let extract_code = !*IS_STDOUT_TERMINAL && code_mode; - let (output, tool_call_results) = if no_stream || extract_code { + config.write().before_chat_completion(&input)?; + let (output, tool_results) = if no_stream || extract_code { let ChatCompletionsOutput { text, tool_calls, .. } = client.chat_completions(input.clone()).await?; @@ -191,16 +192,18 @@ async fn start_directive( (text, vec![]) } } else { - send_stream(&input, client.as_ref(), config, abort_signal.clone()).await? + chat_completion_streaming(&input, client.as_ref(), config, abort_signal.clone()).await? }; config .write() - .save_message(&mut input, &output, &tool_call_results)?; + .after_chat_completion(&mut input, &output, &tool_results)?; + config.write().exit_session()?; - if need_send_call_results(&tool_call_results) { + + if need_send_tool_results(&tool_results) { start_directive( config, - input.merge_tool_call(output, tool_call_results), + input.merge_tool_call(output, tool_results), no_stream, code_mode, abort_signal, @@ -219,6 +222,7 @@ async fn start_interactive(config: &GlobalConfig) -> Result<()> { #[async_recursion::async_recursion] 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 (stop_spinner_tx, _) = run_spinner("Generating").await; let ret = client.chat_completions(input.clone()).await; @@ -231,8 +235,9 @@ async fn shell_execute(config: &GlobalConfig, shell: &Shell, mut input: Input) - if let Ok(true) = CODE_BLOCK_RE.is_match(&eval_str) { eval_str = extract_block(&eval_str); } - config.write().save_message(&mut input, &eval_str, &[])?; - config.read().maybe_copy(&eval_str); + config + .write() + .after_chat_completion(&mut input, &eval_str, &[])?; let render_options = config.read().render_options()?; let mut markdown_render = MarkdownRender::init(render_options)?; if config.read().dry_run { @@ -265,7 +270,7 @@ async fn shell_execute(config: &GlobalConfig, shell: &Shell, mut input: Input) - let role = config.read().retrieve_role(EXPLAIN_SHELL_ROLE)?; let input = Input::from_str(config, &eval_str, Some(role)); let abort = create_abort_signal(); - send_stream(&input, client.as_ref(), config, abort).await?; + chat_completion_streaming(&input, client.as_ref(), config, abort).await?; continue; } _ => {} diff --git a/src/repl/mod.rs b/src/repl/mod.rs index d56b790..48d8224 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -6,9 +6,9 @@ use self::completer::ReplCompleter; use self::highlighter::ReplHighlighter; use self::prompt::ReplPrompt; -use crate::client::send_stream; +use crate::client::chat_completion_streaming; use crate::config::{AssertState, Config, GlobalConfig, Input, StateFlags}; -use crate::function::need_send_call_results; +use crate::function::need_send_tool_results; use crate::render::render_error; use crate::utils::{create_abort_signal, set_text, AbortSignal}; @@ -491,16 +491,17 @@ async fn ask( input.use_embeddings(abort_signal.clone()).await?; } while config.read().is_compressing_session() { - std::thread::sleep(std::time::Duration::from_millis(100)); + tokio::time::sleep(std::time::Duration::from_millis(100)).await; } - let client = input.create_client()?; - let (output, tool_call_results) = - send_stream(&input, client.as_ref(), config, abort_signal.clone()).await?; + let client = input.create_client()?; + config.write().before_chat_completion(&input)?; + let (output, tool_results) = + chat_completion_streaming(&input, client.as_ref(), config, abort_signal.clone()).await?; config .write() - .save_message(&mut input, &output, &tool_call_results)?; - config.read().maybe_copy(&output); + .after_chat_completion(&mut input, &output, &tool_results)?; + if config.write().should_compress_session() { let config = config.clone(); let color = if config.read().light_theme { @@ -521,11 +522,11 @@ async fn ask( config.write().end_compressing_session(); }); } - if need_send_call_results(&tool_call_results) { + if need_send_tool_results(&tool_results) { ask( config, abort_signal, - input.merge_tool_call(output, tool_call_results), + input.merge_tool_call(output, tool_results), false, ) .await |
