diff options
| author | sigoden <sigoden@gmail.com> | 2025-02-03 09:07:37 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2025-02-03 09:07:37 +0800 |
| commit | 458cfc4ad16c509e1aaec8ce2a7d555519861d67 (patch) | |
| tree | ef8986fa6fc4749ef509c798066b470ceba80562 /src | |
| parent | 85e008e1b85ae2a843682e0a15d23b3c6a90dba6 (diff) | |
| download | aichat-458cfc4ad16c509e1aaec8ce2a7d555519861d67.tar.gz | |
feat: display reasoning tokens (#1139)
Diffstat (limited to 'src')
| -rw-r--r-- | src/client/openai.rs | 28 |
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(), |
