diff options
| -rw-r--r-- | src/client/common.rs | 28 | ||||
| -rw-r--r-- | src/main.rs | 33 | ||||
| -rw-r--r-- | src/repl/mod.rs | 2 | ||||
| -rw-r--r-- | src/utils/mod.rs | 28 |
4 files changed, 37 insertions, 54 deletions
diff --git a/src/client/common.rs b/src/client/common.rs index d8909ee..5ad27ae 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -14,7 +14,7 @@ use inquire::{required, Text}; use reqwest::{Client as ReqwestClient, RequestBuilder}; use serde::Deserialize; use serde_json::{json, Value}; -use std::{future::Future, time::Duration}; +use std::time::Duration; use tokio::sync::mpsc::unbounded_channel; const MODELS_YAML: &str = include_str!("../../models.yaml"); @@ -378,6 +378,7 @@ pub fn create_openai_compatible_client_config(client: &str) -> Result<Option<(St pub async fn call_chat_completions( input: &Input, + print: bool, extract_code: bool, client: &dyn Client, abort_signal: AbortSignal, @@ -397,10 +398,12 @@ pub async fn call_chat_completions( .. } = ret; if !text.is_empty() { - if extract_code && text.trim_start().starts_with("```") { - text = extract_block(&text); + if extract_code { + text = extract_code_block(&text).to_string(); + } + if print { + client.global_config().read().print_markdown(&text)?; } - client.global_config().read().print_markdown(&text)?; } Ok((text, eval_tool_calls(client.global_config(), tool_calls)?)) } @@ -444,23 +447,6 @@ pub async fn call_chat_completions_streaming( } } -#[allow(unused)] -pub async fn chat_completions_as_streaming<F, Fut>( - builder: RequestBuilder, - handler: &mut SseHandler, - f: F, -) -> Result<()> -where - F: FnOnce(RequestBuilder) -> Fut, - Fut: Future<Output = Result<String>>, -{ - let text = f(builder).await?; - handler.text(&text)?; - handler.done(); - - Ok(()) -} - pub fn noop_prepare_embeddings<T>(_client: &T, _data: &EmbeddingsData) -> Result<RequestData> { bail!("The client doesn't support embeddings api") } diff --git a/src/main.rs b/src/main.rs index f27f6e2..c7a1239 100644 --- a/src/main.rs +++ b/src/main.rs @@ -208,7 +208,14 @@ async fn start_directive( let extract_code = !*IS_STDOUT_TERMINAL && code_mode; config.write().before_chat_completion(&input)?; let (output, tool_results) = if !input.stream() || extract_code { - call_chat_completions(&input, extract_code, client.as_ref(), abort_signal.clone()).await? + call_chat_completions( + &input, + true, + extract_code, + client.as_ref(), + abort_signal.clone(), + ) + .await? } else { call_chat_completions_streaming(&input, client.as_ref(), abort_signal.clone()).await? }; @@ -244,17 +251,9 @@ async fn shell_execute( ) -> Result<()> { let client = input.create_client()?; config.write().before_chat_completion(&input)?; - let ret = abortable_run_with_spinner( - client.chat_completions(input.clone()), - "Generating", - abort_signal.clone(), - ) - .await; - let mut eval_str = ret?.text; - eval_str = strip_think_tag(&eval_str).to_string(); - if let Ok(true) = CODE_BLOCK_RE.is_match(&eval_str) { - eval_str = extract_block(&eval_str); - } + let (eval_str, _) = + call_chat_completions(&input, false, true, client.as_ref(), abort_signal.clone()).await?; + config .write() .after_chat_completion(&input, &eval_str, &[])?; @@ -314,8 +313,14 @@ async fn shell_execute( ) .await?; } else { - call_chat_completions(&input, false, client.as_ref(), abort_signal.clone()) - .await?; + call_chat_completions( + &input, + true, + false, + client.as_ref(), + abort_signal.clone(), + ) + .await?; } println!(); continue; diff --git a/src/repl/mod.rs b/src/repl/mod.rs index 19ab3dc..780023a 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -713,7 +713,7 @@ async fn ask( let (output, tool_results) = if input.stream() { call_chat_completions_streaming(&input, client.as_ref(), abort_signal.clone()).await? } else { - call_chat_completions(&input, false, client.as_ref(), abort_signal.clone()).await? + call_chat_completions(&input, true, false, client.as_ref(), abort_signal.clone()).await? }; config .write() diff --git a/src/utils/mod.rs b/src/utils/mod.rs index 049f361..38617a6 100644 --- a/src/utils/mod.rs +++ b/src/utils/mod.rs @@ -61,10 +61,6 @@ pub fn parse_bool(value: &str) -> Option<bool> { } } -pub fn strip_think_tag(text: &str) -> Cow<str> { - THINK_TAG_RE.replace_all(text, "") -} - pub fn estimate_token_length(text: &str) -> usize { let words: Vec<&str> = text.unicode_words().collect(); let mut output: f32 = 0.0; @@ -101,20 +97,16 @@ pub fn light_theme_from_colorfgbg(colorfgbg: &str) -> Option<bool> { Some(light) } -pub fn extract_block(input: &str) -> String { - let output: String = CODE_BLOCK_RE - .captures_iter(input) - .filter_map(|m| { - m.ok() - .and_then(|cap| cap.get(1)) - .map(|m| String::from(m.as_str())) - }) - .collect(); - if output.is_empty() { - input.trim().to_string() - } else { - output.trim().to_string() - } +pub fn strip_think_tag(text: &str) -> Cow<str> { + THINK_TAG_RE.replace_all(text, "") +} + +pub fn extract_code_block(text: &str) -> &str { + CODE_BLOCK_RE + .captures(text) + .ok() + .and_then(|v| v?.get(1).map(|v| v.as_str().trim())) + .unwrap_or(text) } pub fn convert_option_string(value: &str) -> Option<String> { |
