summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2024-06-17 09:14:54 +0800
committerGitHub <noreply@github.com>2024-06-17 09:14:54 +0800
commit638bf3276613265da761b4d74f269fa77c8f4bb1 (patch)
tree9667711b92bc3664ad72c179cfe1cf4bbf54d322 /src
parent12872b3d2956a26fb7d1cd464b3672ebe40d94c1 (diff)
downloadaichat-638bf3276613265da761b4d74f269fa77c8f4bb1.tar.gz
refactor: improve code quatity (#604)
Diffstat (limited to 'src')
-rw-r--r--src/client/common.rs6
-rw-r--r--src/config/input.rs8
-rw-r--r--src/config/mod.rs126
-rw-r--r--src/function.rs15
-rw-r--r--src/main.rs25
-rw-r--r--src/repl/mod.rs21
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