From 638bf3276613265da761b4d74f269fa77c8f4bb1 Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 17 Jun 2024 09:14:54 +0800 Subject: refactor: improve code quatity (#604) --- src/main.rs | 25 +++++++++++++++---------- 1 file changed, 15 insertions(+), 10 deletions(-) (limited to 'src/main.rs') 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; } _ => {} -- cgit v1.2.3