From 700d8a3245f1133c37e035039e511a1e5dce1a5d Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 3 Feb 2025 22:35:01 +0800 Subject: feat: strip reasoning contents (#1141) --- assets/arena.html | 22 +++++++++++++++++++++- assets/playground.html | 22 +++++++++++++++++++++- src/client/message.rs | 4 ++++ src/client/model.rs | 14 +++++++++++--- src/client/openai.rs | 9 ++++++++- src/config/input.rs | 7 +++++++ src/config/mod.rs | 6 ++---- src/main.rs | 1 + src/utils/mod.rs | 6 ++++++ 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*([\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(''); + const match = /^(\s*)/.exec(src); + if (match) { + return match[1].length + } else { + return -1; + } }, tokenizer(src, tokens) { const rule = /^\s*([\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*([\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(''); + const match = /^(\s*)/.exec(src); + if (match) { + return match[1].length + } else { + return -1; + } }, tokenizer(src, tokens) { const rule = /^\s*([\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 = 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 { + 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*.*?(\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 { } } +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; -- cgit v1.2.3