summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
-rw-r--r--assets/arena.html22
-rw-r--r--assets/playground.html22
-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
9 files changed, 81 insertions, 10 deletions
diff --git a/assets/arena.html b/assets/arena.html
index 31973eb..2296de2 100644
--- a/assets/arena.html
+++ b/assets/arena.html
@@ -870,6 +870,7 @@
});
}
}
+ sanitizeMessages(messages);
const body = {
model: chat.model,
messages: messages,
@@ -991,6 +992,20 @@
});
}
+ function sanitizeMessages(messages) {
+ let messagesLen = messages.length;
+ for (let i = 0; i < messagesLen; i++) {
+ const message = messages[i];
+ if (typeof message.content === "string" && message.role === "assistant" && i !== messagesLen - 1) {
+ message.content = stripThinkTag(message.content);
+ }
+ }
+ }
+
+ function stripThinkTag(text) {
+ return text.replace(/^\s*<think>([\s\S]*?)<\/think>(\s*|$)/g, '')
+ }
+
function setupMarked() {
const renderer = {
code({ text, lang }) {
@@ -1013,7 +1028,12 @@
name: 'think',
level: 'block',
start(src) {
- return src.indexOf('<think>');
+ const match = /^(\s*)<think>/.exec(src);
+ if (match) {
+ return match[1].length
+ } else {
+ return -1;
+ }
},
tokenizer(src, tokens) {
const rule = /^\s*<think>([\s\S]*?)(<\/think>|$)/;
diff --git a/assets/playground.html b/assets/playground.html
index d038bd1..06e2e62 100644
--- a/assets/playground.html
+++ b/assets/playground.html
@@ -1298,6 +1298,7 @@
messages = [...promptMessages, ...messages];
}
}
+ sanitizeMessages(messages);
const body = {
model: this.settings.model,
messages: messages,
@@ -1458,6 +1459,20 @@
return { system: prompt, cases: [] }
}
+ function sanitizeMessages(messages) {
+ let messagesLen = messages.length;
+ for (let i = 0; i < messagesLen; i++) {
+ const message = messages[i];
+ if (typeof message.content === "string" && message.role === "assistant" && i !== messagesLen - 1) {
+ message.content = stripThinkTag(message.content);
+ }
+ }
+ }
+
+ function stripThinkTag(text) {
+ return text.replace(/^\s*<think>([\s\S]*?)<\/think>(\s*|$)/g, '')
+ }
+
function convertImageToDataURL(imageFile) {
return new Promise((resolve, reject) => {
if (!imageFile) {
@@ -1494,7 +1509,12 @@
name: 'think',
level: 'block',
start(src) {
- return src.indexOf('<think>');
+ const match = /^(\s*)<think>/.exec(src);
+ if (match) {
+ return match[1].length
+ } else {
+ return -1;
+ }
},
tokenizer(src, tokens) {
const rule = /^\s*<think>([\s\S]*?)(<\/think>|$)/;
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;