From c19c4b9be62beb653f50b9f16971321a0dc4fbde Mon Sep 17 00:00:00 2001 From: zipg Date: Sat, 18 Jul 2026 16:38:01 +0800 Subject: [PATCH] fix(ai): finish compatible streams after finish reason --- crates/dbx-core/src/ai.rs | 119 ++++++++++++++++++++++++++++++++++++-- 1 file changed, 115 insertions(+), 4 deletions(-) diff --git a/crates/dbx-core/src/ai.rs b/crates/dbx-core/src/ai.rs index f619086c5..ea1028079 100644 --- a/crates/dbx-core/src/ai.rs +++ b/crates/dbx-core/src/ai.rs @@ -555,6 +555,12 @@ pub fn openai_stream_reasoning(event: &serde_json::Value) -> Option<&str> { .filter(|text| !text.is_empty()) } +fn openai_stream_has_finish_reason(event: &serde_json::Value) -> bool { + event["choices"].as_array().is_some_and(|choices| { + choices.iter().any(|choice| choice["finish_reason"].as_str().is_some_and(|reason| !reason.is_empty())) + }) +} + pub fn responses_stream_text(event: &serde_json::Value) -> Option<&str> { let event_type = event["type"].as_str().unwrap_or_default(); if !event_type.is_empty() && event_type != "response.output_text.delta" { @@ -1567,6 +1573,7 @@ async fn stream_openai( let mut byte_stream = res.bytes_stream(); let mut buf = Vec::new(); + let mut finish_reason_deadline = None; loop { tokio::select! { @@ -1600,12 +1607,22 @@ async fn stream_openai( done: false, }); } + if finish_reason_deadline.is_none() && openai_stream_has_finish_reason(&event) { + finish_reason_deadline = + Some(tokio::time::Instant::now() + std::time::Duration::from_secs(1)); + } } } if finished { break; } } _ = cancelled.notified() => { break; } + _ = async { + match finish_reason_deadline { + Some(deadline) => tokio::time::sleep_until(deadline).await, + None => std::future::pending().await, + } + } => { break; } } } @@ -2130,6 +2147,7 @@ async fn stream_openai_with_tools( let mut byte_stream = res.bytes_stream(); let mut buf = Vec::new(); let mut token_usage: Option = None; + let mut finish_reason_deadline = None; loop { tokio::select! { @@ -2192,6 +2210,10 @@ async fn stream_openai_with_tools( } } } + if finish_reason_deadline.is_none() && openai_stream_has_finish_reason(&event) { + finish_reason_deadline = + Some(tokio::time::Instant::now() + std::time::Duration::from_secs(1)); + } } } @@ -2200,6 +2222,12 @@ async fn stream_openai_with_tools( _ = cancelled.notified() => { return Err(AGENT_CANCELLED_ERROR.to_string()); } + _ = async { + match finish_reason_deadline { + Some(deadline) => tokio::time::sleep_until(deadline).await, + None => std::future::pending().await, + } + } => { break; } } } @@ -2600,6 +2628,7 @@ pub fn load_config(path: &Path) -> Result, String> { mod tests { use std::cell::RefCell; use std::collections::{HashMap, HashSet}; + use tokio::sync::Notify; use super::{ apply_chat_completion_thinking_toggle, build_ai_http_client, build_responses_input_with_tools, claude_headers, @@ -2607,10 +2636,10 @@ mod tests { maybe_bearer_headers, openai_response_text, openai_stream_reasoning, openai_stream_text, parse_model_list_response, resolve_endpoint, resolve_gemini_stream_endpoint, resolve_model_list_endpoint, responses_function_tool, responses_max_output_tokens, responses_stream_text, responses_text, - responses_token_usage, set_chat_completion_token_limit, stream_data_payload, uses_anthropic_messages_api, - validate_config, validate_model_list_config, AiApiStyle, AiAuthMethod, AiConfig, AiMessage, AiModelInfo, - AiProvider, AiReasoningLevel, StreamToolEvent, StreamingToolCallAccumulator, ToolCallRef, AUTHORIZATION, - CLAUDE_DEFAULT_SYSTEM, TEST_PROMPT, + responses_token_usage, set_chat_completion_token_limit, stream_data_payload, stream_openai_with_tools, + uses_anthropic_messages_api, validate_config, validate_model_list_config, AiApiStyle, AiAuthMethod, + AiCompletionRequest, AiConfig, AiMessage, AiModelInfo, AiProvider, AiReasoningLevel, StreamToolEvent, + StreamingToolCallAccumulator, ToolCallRef, AUTHORIZATION, CLAUDE_DEFAULT_SYSTEM, TEST_PROMPT, }; /// Reproduce the "Unknown tool:" bug: some OpenAI-compatible providers @@ -3467,6 +3496,88 @@ mod tests { assert!(body.get("reasoning_effort").is_none()); } + #[tokio::test] + async fn openai_tool_stream_finishes_without_done_marker_after_finish_reason() { + use tokio::io::AsyncWriteExt; + + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.unwrap(); + let payload = + concat!("data: {\"choices\":[{\"delta\":{\"content\":\"done\"},", "\"finish_reason\":\"stop\"}]}\n\n"); + let response = format!( + "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nTransfer-Encoding: chunked\r\nConnection: keep-alive\r\n\r\n{:X}\r\n{}\r\n", + payload.len(), + payload + ); + socket.write_all(response.as_bytes()).await.unwrap(); + socket.flush().await.unwrap(); + let keep_alive = ": keep-alive\n\n"; + for _ in 0..20 { + tokio::time::sleep(std::time::Duration::from_millis(250)).await; + let chunk = format!("{:X}\r\n{}\r\n", keep_alive.len(), keep_alive); + if socket.write_all(chunk.as_bytes()).await.is_err() { + break; + } + if socket.flush().await.is_err() { + break; + } + } + }); + + let config = AiConfig { + provider: AiProvider::OpenaiCompatible, + api_key: "lm-studio".to_string(), + auth_method: AiAuthMethod::Bearer, + endpoint: format!("http://{address}/v1"), + model: "local-model".to_string(), + models: Vec::new(), + api_style: AiApiStyle::Completions, + proxy_enabled: false, + proxy_url: String::new(), + enable_thinking: true, + reasoning_level: AiReasoningLevel::Default, + context_window: None, + codex_cli_path: None, + codex_cli_env: Default::default(), + claude_code_cli_path: None, + claude_code_cli_env: Default::default(), + }; + let request = AiCompletionRequest { + config: config.clone(), + system_prompt: "Use tools when needed.".to_string(), + messages: vec![AiMessage { + role: "tool".to_string(), + content: "query failed".to_string(), + tool_call_id: Some("call-1".to_string()), + tool_calls: Vec::new(), + }], + task_contract: None, + max_tokens: Some(64), + }; + let client = build_ai_http_client(&config, 10).unwrap(); + let events = std::sync::Arc::new(std::sync::Mutex::new(Vec::new())); + let captured = events.clone(); + + let result = tokio::time::timeout( + std::time::Duration::from_secs(3), + stream_openai_with_tools(&client, "lm-studio-test", &request, &[], &Notify::new(), &move |event| { + captured.lock().unwrap().push(event); + }), + ) + .await + .expect("stream should finish after the finish_reason grace period"); + + assert!(result.is_ok()); + assert!(events + .lock() + .unwrap() + .iter() + .any(|event| matches!(event, StreamToolEvent::Chunk(chunk) if chunk.delta == "done"))); + server.abort(); + } + #[test] fn uses_reasoning_effort_to_disable_ollama_thinking() { let config = AiConfig {