summaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorsigoden <sigoden@gmail.com>2025-02-03 09:07:37 +0800
committerGitHub <noreply@github.com>2025-02-03 09:07:37 +0800
commit458cfc4ad16c509e1aaec8ce2a7d555519861d67 (patch)
treeef8986fa6fc4749ef509c798066b470ceba80562
parent85e008e1b85ae2a843682e0a15d23b3c6a90dba6 (diff)
downloadaichat-458cfc4ad16c509e1aaec8ce2a7d555519861d67.tar.gz
feat: display reasoning tokens (#1139)
-rw-r--r--src/client/openai.rs28
1 files changed, 27 insertions, 1 deletions
diff --git a/src/client/openai.rs b/src/client/openai.rs
index ce00de7..6490236 100644
--- a/src/client/openai.rs
+++ b/src/client/openai.rs
@@ -104,6 +104,7 @@ pub async fn openai_chat_completions_streaming(
let mut function_name = String::new();
let mut function_arguments = String::new();
let mut function_id = String::new();
+ let mut reason_state = 0;
let handle = |message: SseMmessage| -> Result<bool> {
if message.data == "[DONE]" {
if !function_name.is_empty() {
@@ -124,6 +125,20 @@ pub async fn openai_chat_completions_streaming(
.as_str()
.filter(|v| !v.is_empty())
{
+ if reason_state == 1 {
+ handler.text("\n</think>\n\n")?;
+ reason_state = 0;
+ }
+ handler.text(text)?;
+ } else if let Some(text) = data["choices"][0]["delta"]["reasoning_content"]
+ .as_str()
+ .or_else(|| data["choices"][0]["delta"]["reasoning"].as_str())
+ .filter(|v| !v.is_empty())
+ {
+ if reason_state == 0 {
+ handler.text("<think>\n")?;
+ reason_state = 1;
+ }
handler.text(text)?;
} else if let (Some(function), index, id) = (
data["choices"][0]["delta"]["tool_calls"][0]["function"].as_object(),
@@ -314,6 +329,12 @@ pub fn openai_extract_chat_completions(data: &Value) -> Result<ChatCompletionsOu
.as_str()
.unwrap_or_default();
+ let reasoning = data["choices"][0]["message"]["reasoning_content"]
+ .as_str()
+ .or_else(|| data["choices"][0]["message"]["reasoning"].as_str())
+ .unwrap_or_default()
+ .trim();
+
let mut tool_calls = vec![];
if let Some(calls) = data["choices"][0]["message"]["tool_calls"].as_array() {
for call in calls {
@@ -337,8 +358,13 @@ pub fn openai_extract_chat_completions(data: &Value) -> Result<ChatCompletionsOu
if text.is_empty() && tool_calls.is_empty() {
bail!("Invalid response data: {data}");
}
+ let text = if !reasoning.is_empty() {
+ format!("<think>\n{reasoning}\n</think>\n\n{text}")
+ } else {
+ text.to_string()
+ };
let output = ChatCompletionsOutput {
- text: text.to_string(),
+ text,
tool_calls,
id: data["id"].as_str().map(|v| v.to_string()),
input_tokens: data["usage"]["prompt_tokens"].as_u64(),