From e8aa5d12a366a1fe14192e56321d39a7c9e5d87f Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 3 Feb 2025 22:40:13 +0800 Subject: refactor: improve call_chat_completions function (#1143) --- src/client/common.rs | 28 +++++++--------------------- src/main.rs | 33 +++++++++++++++++++-------------- src/repl/mod.rs | 2 +- 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( - builder: RequestBuilder, - handler: &mut SseHandler, - f: F, -) -> Result<()> -where - F: FnOnce(RequestBuilder) -> Fut, - Fut: Future>, -{ - let text = f(builder).await?; - handler.text(&text)?; - handler.done(); - - Ok(()) -} - pub fn noop_prepare_embeddings(_client: &T, _data: &EmbeddingsData) -> Result { 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 { } } -pub fn strip_think_tag(text: &str) -> Cow { - 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 { 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 { + 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 { -- cgit v1.2.3