summaryrefslogtreecommitdiffstats
path: root/src/client
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2025-02-03 22:35:01 +0800
committerGitHub <noreply@github.com>2025-02-03 22:35:01 +0800
commit700d8a3245f1133c37e035039e511a1e5dce1a5d (patch)
tree7b1b2cf2e816adb0af399c9ca34051d3fcbef47f /src/client
parent07e58aaacd8c7f486a0d2a7c8d35a195f27d487e (diff)
downloadaichat-700d8a3245f1133c37e035039e511a1e5dce1a5d.tar.gz
feat: strip reasoning contents (#1141)
Diffstat (limited to 'src/client')
-rw-r--r--src/client/message.rs4
-rw-r--r--src/client/model.rs14
-rw-r--r--src/client/openai.rs9
3 files changed, 23 insertions, 4 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 })]
}
})