summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/client/common.rs28
-rw-r--r--src/main.rs33
-rw-r--r--src/repl/mod.rs2
-rw-r--r--src/utils/mod.rs28
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> {