diff options
| author | sigoden <sigoden@gmail.com> | 2024-05-14 11:16:55 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-05-14 11:16:55 +0800 |
| commit | 5284a18248bb8e48eaa4a1e6ddcf73d944d23783 (patch) | |
| tree | fe60e7b0be7c4ea40051f40f88362e0487da6355 /src/main.rs | |
| parent | 154c1e0b4b7fad08094c601893b081581d2ee0c8 (diff) | |
| download | aichat-5284a18248bb8e48eaa4a1e6ddcf73d944d23783.tar.gz | |
refactor: config::Input (#503)
Diffstat (limited to 'src/main.rs')
| -rw-r--r-- | src/main.rs | 17 |
1 files changed, 8 insertions, 9 deletions
diff --git a/src/main.rs b/src/main.rs index 558e473..0511f09 100644 --- a/src/main.rs +++ b/src/main.rs @@ -12,7 +12,7 @@ mod utils; extern crate log; use crate::cli::Cli; -use crate::client::{ensure_model_capabilities, init_client, list_models, send_stream}; +use crate::client::{ensure_model_capabilities, list_models, send_stream}; use crate::config::{ Config, GlobalConfig, Input, InputContext, WorkingMode, CODE_ROLE, EXPLAIN_SHELL_ROLE, SHELL_ROLE, @@ -138,9 +138,9 @@ async fn start_directive( no_stream: bool, code_mode: bool, ) -> Result<()> { - let mut client = init_client(config)?; + let mut client = input.create_client()?; ensure_model_capabilities(client.as_mut(), input.required_capabilities())?; - config.read().maybe_print_send_tokens(&input); + input.maybe_print_input_tokens(); let is_terminal_stdout = stdout().is_terminal(); let extract_code = !is_terminal_stdout && code_mode; let output = if no_stream || extract_code { @@ -176,8 +176,8 @@ async fn start_interactive(config: &GlobalConfig) -> Result<()> { #[async_recursion::async_recursion] async fn execute(config: &GlobalConfig, mut input: Input) -> Result<()> { - let client = init_client(config)?; - config.read().maybe_print_send_tokens(&input); + let client = input.create_client()?; + input.maybe_print_input_tokens(); let is_terminal_stdout = stdout().is_terminal(); let ret = if is_terminal_stdout { let (spinner_tx, spinner_rx) = oneshot::channel(); @@ -223,7 +223,7 @@ async fn execute(config: &GlobalConfig, mut input: Input) -> Result<()> { } "📙 Explain" => { let role = config.read().retrieve_role(EXPLAIN_SHELL_ROLE)?; - let input = Input::from_str(&eval_str, InputContext::role(role)); + let input = Input::from_str(config, &eval_str, Some(InputContext::role(role))); let abort = create_abort_signal(); send_stream(&input, client.as_ref(), config, abort).await?; continue; @@ -254,11 +254,10 @@ fn aggregate_text(text: Option<String>) -> Result<Option<String>> { } fn create_input(config: &GlobalConfig, text: Option<String>, file: &[String]) -> Result<Input> { - let input_context = config.read().input_context(); let input = if file.is_empty() { - Input::from_str(&text.unwrap_or_default(), input_context) + Input::from_str(config, &text.unwrap_or_default(), None) } else { - Input::new(&text.unwrap_or_default(), file.to_vec(), input_context)? + Input::new(config, &text.unwrap_or_default(), file.to_vec(), None)? }; if input.is_empty() { bail!("No input"); |
