summaryrefslogtreecommitdiffstats
path: root/src
diff options
context:
space:
mode:
Diffstat (limited to 'src')
-rw-r--r--src/client/message.rs4
-rw-r--r--src/client/model.rs14
-rw-r--r--src/client/openai.rs9
-rw-r--r--src/config/input.rs7
-rw-r--r--src/config/mod.rs6
-rw-r--r--src/main.rs1
-rw-r--r--src/utils/mod.rs6
7 files changed, 39 insertions, 8 deletions
diff --git a/src/client/message.rs b/src/client/message.rs
index 2cfc183..2f7517c 100644
--- a/src/client/message.rs
+++ b/src/client/message.rs
@@ -65,6 +65,10 @@ impl MessageRole {
pub fn is_user(&self) -> bool {
matches!(self, MessageRole::User)
}
+
+ pub fn is_assistant(&self) -> bool {
+ matches!(self, MessageRole::Assistant)
+ }
}
#[derive(Debug, Clone, Deserialize, Serialize)]
diff --git a/src/client/model.rs b/src/client/model.rs
index b562705..bcd474f 100644
--- a/src/client/model.rs
+++ b/src/client/model.rs
@@ -5,7 +5,7 @@ use super::{
};
use crate::config::Config;
-use crate::utils::estimate_token_length;
+use crate::utils::{estimate_token_length, strip_think_tag};
use anyhow::{bail, Result};
use serde::{Deserialize, Serialize};
@@ -220,10 +220,18 @@ impl Model {
}
pub fn messages_tokens(&self, messages: &[Message]) -> usize {
+ let messages_len = messages.len();
messages
.iter()
- .map(|v| match &v.content {
- MessageContent::Text(text) => estimate_token_length(text),
+ .enumerate()
+ .map(|(i, v)| match &v.content {
+ MessageContent::Text(text) => {
+ if v.role.is_assistant() && i != messages_len - 1 {
+ estimate_token_length(&strip_think_tag(text))
+ } else {
+ estimate_token_length(text)
+ }
+ }
MessageContent::Array(list) => list
.iter()
.map(|v| match v {
diff --git a/src/client/openai.rs b/src/client/openai.rs
index 6490236..928a600 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -1,5 +1,7 @@
use super::*;
+use crate::utils::strip_think_tag;
+
use anyhow::{bail, Context, Result};
use reqwest::RequestBuilder;
use serde::Deserialize;
@@ -219,9 +221,11 @@ pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Mod
stream,
} = data;
+ let messages_len = messages.len();
let messages: Vec<Value> = messages
.into_iter()
- .flat_map(|message| {
+ .enumerate()
+ .flat_map(|(i, message)| {
let Message { role, content } = message;
match content {
MessageContent::ToolCalls(MessageContentToolCalls {
@@ -281,6 +285,9 @@ pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Mod
}).collect()
}
},
+ MessageContent::Text(text) if role.is_assistant() && i != messages_len - 1 => vec![
+ json!({ "role": role, "content": strip_think_tag(&text) }
+ )],
_ => vec![json!({ "role": role, "content": content })]
}
})
diff --git a/src/config/input.rs b/src/config/input.rs
index e468f19..482b1d5 100644
--- a/src/config/input.rs
+++ b/src/config/input.rs
@@ -244,6 +244,13 @@ impl Input {
init_client(&self.config, Some(self.role().model().clone()))
}
+ pub async fn fetch_chat_text(&self) -> Result<String> {
+ let client = self.create_client()?;
+ let text = client.chat_completions(self.clone()).await?.text;
+ let text = strip_think_tag(&text).to_string();
+ Ok(text)
+ }
+
pub fn prepare_completion_data(
&self,
model: &Model,
diff --git a/src/config/mod.rs b/src/config/mod.rs
index c9efd72..38f2f79 100644
--- a/src/config/mod.rs
+++ b/src/config/mod.rs
@@ -1262,8 +1262,7 @@ impl Config {
.clone()
.unwrap_or_else(|| SUMMARIZE_PROMPT.into());
let input = Input::from_str(config, &prompt, None);
- let client = input.create_client()?;
- let summary = client.chat_completions(input).await?.text;
+ let summary = input.fetch_chat_text().await?;
let summary_prompt = config
.read()
.summary_prompt
@@ -1322,8 +1321,7 @@ impl Config {
};
let role = config.read().retrieve_role(CREATE_TITLE_ROLE)?;
let input = Input::from_str(config, &text, Some(role));
- let client = input.create_client()?;
- let text = client.chat_completions(input).await?.text;
+ let text = input.fetch_chat_text().await?;
if let Some(session) = config.write().session.as_mut() {
session.set_autoname(&text);
}
diff --git a/src/main.rs b/src/main.rs
index 5251bad..f27f6e2 100644
--- a/src/main.rs
+++ b/src/main.rs
@@ -251,6 +251,7 @@ async fn shell_execute(
)
.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);
}
diff --git a/src/utils/mod.rs b/src/utils/mod.rs
index 9dde751..049f361 100644
--- a/src/utils/mod.rs
+++ b/src/utils/mod.rs
@@ -26,11 +26,13 @@ use anyhow::{Context, Result};
use fancy_regex::Regex;
use fuzzy_matcher::{skim::SkimMatcherV2, FuzzyMatcher};
use is_terminal::IsTerminal;
+use std::borrow::Cow;
use std::{env, path::PathBuf, process};
use unicode_segmentation::UnicodeSegmentation;
lazy_static::lazy_static! {
pub static ref CODE_BLOCK_RE: Regex = Regex::new(r"(?ms)```\w*(.*)```").unwrap();
+ pub static ref THINK_TAG_RE: Regex = Regex::new(r"(?s)^\s*<think>.*?</think>(\s*|$)").unwrap();
pub static ref IS_STDOUT_TERMINAL: bool = std::io::stdout().is_terminal();
pub static ref NO_COLOR: bool = env::var("NO_COLOR").ok().and_then(|v| parse_bool(&v)).unwrap_or_default() || !*IS_STDOUT_TERMINAL;
}
@@ -59,6 +61,10 @@ 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;