summaryrefslogtreecommitdiffstats
path: root/src/main.rs
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/main.rs
parent12872b3d2956a26fb7d1cd464b3672ebe40d94c1 (diff)
downloadaichat-638bf3276613265da761b4d74f269fa77c8f4bb1.tar.gz
refactor: improve code quatity (#604)
Diffstat (limited to 'src/main.rs')
-rw-r--r--src/main.rs25
1 files changed, 15 insertions, 10 deletions
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;
}
_ => {}