diff --git a/server/src/cursor/prompting/derived_state.rs b/server/src/cursor/prompting/derived_state.rs index bad3696..642e04a 100644 --- a/server/src/cursor/prompting/derived_state.rs +++ b/server/src/cursor/prompting/derived_state.rs @@ -27,7 +27,9 @@ pub fn fold_derived_state(messages: &[CanonicalMessage]) -> DerivedState { continue; }; match normalize(&name).as_str() { - "todowrite" | "updatetodos" => state.todos = Some(input), + "todowrite" | "updatetodos" => { + state.todos = Some(apply_todo_write(state.todos.take(), input)); + } "createplan" | "updateplan" | "writeplan" => state.plan = Some(input), _ => {} } @@ -38,6 +40,39 @@ pub fn fold_derived_state(messages: &[CanonicalMessage]) -> DerivedState { state } +fn apply_todo_write(current: Option, mut input: Value) -> Value { + if !input.get("merge").and_then(Value::as_bool).unwrap_or(false) { + return input; + } + let mut todos = current + .as_ref() + .and_then(|value| value.get("todos")) + .and_then(Value::as_array) + .cloned() + .unwrap_or_default(); + let patches = input + .get("todos") + .and_then(Value::as_array) + .cloned() + .unwrap_or_default(); + for patch in patches { + let existing = patch.get("id").and_then(Value::as_str).and_then(|id| { + todos + .iter_mut() + .find(|todo| todo.get("id").and_then(Value::as_str) == Some(id)) + }); + match (existing, patch) { + (Some(Value::Object(todo)), Value::Object(patch)) => todo.extend(patch), + (_, patch) => todos.push(patch), + } + } + if let Some(object) = input.as_object_mut() { + object.insert("merge".into(), Value::Bool(false)); + object.insert("todos".into(), Value::Array(todos)); + } + input +} + fn normalize(value: &str) -> String { value .chars() @@ -45,3 +80,87 @@ fn normalize(value: &str) -> String { .flat_map(char::to_lowercase) .collect() } + +#[cfg(test)] +mod tests { + use super::*; + use crate::model::{Origin, Role, ToolCallContent, ToolResultContent}; + + #[test] + fn todo_write_merge_materializes_complete_existing_items_and_appends_new_ids() { + let messages = vec![ + assistant_call( + "create", + serde_json::json!({ + "merge": false, + "todos": [ + {"id": "first", "content": "First", "status": "in_progress"}, + {"id": "second", "content": "Second", "status": "pending"} + ] + }), + ), + successful_result("create"), + assistant_call( + "merge", + serde_json::json!({ + "merge": true, + "todos": [ + {"id": "first", "status": "completed"}, + {"id": "second", "content": "Second updated"}, + {"id": "third", "content": "Third", "status": "cancelled"} + ] + }), + ), + successful_result("merge"), + ]; + + let state = fold_derived_state(&messages); + + assert_eq!( + state.todos.unwrap()["todos"], + serde_json::json!([ + {"id": "first", "content": "First", "status": "completed"}, + {"id": "second", "content": "Second updated", "status": "pending"}, + {"id": "third", "content": "Third", "status": "cancelled"} + ]) + ); + } + + fn assistant_call(call_id: &str, arguments: Value) -> CanonicalMessage { + CanonicalMessage { + message_id: format!("assistant-{call_id}"), + role: Role::Assistant, + origin: Origin::Assistant, + content: MessageContent::Assistant { + text: String::new(), + thinking: String::new(), + tool_round_id: Some(format!("round-{call_id}").into()), + replay_state: None, + tool_calls: vec![ToolCallContent { + index: 0, + call_id: call_id.into(), + name: "TodoWrite".into(), + arguments, + }], + }, + runtime_event_id: None, + } + } + + fn successful_result(call_id: &str) -> CanonicalMessage { + CanonicalMessage { + message_id: format!("result-{call_id}"), + role: Role::Tool, + origin: Origin::Tool, + content: MessageContent::ToolResult(ToolResultContent { + call_id: call_id.into(), + name: "TodoWrite".into(), + content: "{}".into(), + is_error: false, + image: None, + provider_parts: Vec::new(), + }), + runtime_event_id: None, + } + } +} diff --git a/server/src/provider/openai_responses.rs b/server/src/provider/openai_responses.rs index d1b4578..f209328 100644 --- a/server/src/provider/openai_responses.rs +++ b/server/src/provider/openai_responses.rs @@ -119,7 +119,6 @@ impl Provider for OpenAiResponsesProvider { let mut reasoning_items = Vec::new(); let mut saw_tool = false; let mut saw_completed_item = false; - let mut saw_done_marker = false; let mut terminal = false; loop { let event = tokio::select! { @@ -128,7 +127,7 @@ impl Provider for OpenAiResponsesProvider { }; let Some(event) = event else { break }; let event = event.map_err(|error| Error::Provider(format!("OpenAI Responses SSE: {error}")))?; - if event.data == "[DONE]" { saw_done_marker = true; break; } + if event.data == "[DONE]" { break; } let value: Value = serde_json::from_str(&event.data)?; let kind = value.get("type").and_then(Value::as_str).unwrap_or(&event.event); match kind { @@ -163,19 +162,20 @@ impl Provider for OpenAiResponsesProvider { } "response.output_item.done" => { let item = value.get("item").unwrap_or(&Value::Null); - saw_completed_item = true; match item.get("type").and_then(Value::as_str) { Some("reasoning") => { if thinking_open { thinking_open = false; yield ModelEvent::ThinkingEnd; } reasoning_items.push(item.clone()); } Some("message") => { + saw_completed_item = true; if let Some(final_text) = response_item_text(item) { for event in reconcile_response_text(&mut text_open, &mut text, &final_text) { yield event; } } if text_open { text_open = false; yield ModelEvent::TextEnd; } } Some("function_call") => { + saw_completed_item = true; let index = required_u64(&value, "output_index")? as usize; saw_tool = true; let arguments = item @@ -196,11 +196,24 @@ impl Provider for OpenAiResponsesProvider { } "response.function_call_arguments.done" => { let index = required_u64(&value, "output_index")? as usize; - let arguments = value - .get("arguments") - .and_then(Value::as_str) - .map_or(ResponseToolArguments::None, ResponseToolArguments::Snapshot); - for event in update_response_tool(index, &Value::Null, arguments, true, &mut tools)? { yield event; } + match value.get("arguments").and_then(Value::as_str) { + Some("") => { + for event in update_response_tool( + index, + &Value::Null, + ResponseToolArguments::None, + false, + &mut tools, + )? { yield event; } + } + arguments => { + let arguments = arguments.map_or( + ResponseToolArguments::None, + ResponseToolArguments::Snapshot, + ); + for event in update_response_tool(index, &Value::Null, arguments, true, &mut tools)? { yield event; } + } + } } "response.completed" => { if let Some(usage) = value.pointer("/response/usage") { yield ModelEvent::Usage(responses_usage(usage)); } @@ -238,11 +251,11 @@ impl Provider for OpenAiResponsesProvider { _ => {} } } - if !terminal && saw_done_marker && saw_completed_item { + if !terminal && saw_completed_item { if thinking_open { yield ModelEvent::ThinkingEnd; } if text_open { yield ModelEvent::TextEnd; } if tools.values().any(|tool| !tool.ended) { - Err(Error::Provider("OpenAI Responses [DONE] arrived with an incomplete tool call".into()))?; + Err(Error::Provider("OpenAI Responses stream ended with an incomplete tool call".into()))?; } terminal = true; if !reasoning_items.is_empty() { diff --git a/server/tests/provider_stream.rs b/server/tests/provider_stream.rs index 24ebf3d..45e8513 100644 --- a/server/tests/provider_stream.rs +++ b/server/tests/provider_stream.rs @@ -485,6 +485,40 @@ async fn openai_responses_preserves_delta_that_repeats_the_streamed_suffix() { assert_eq!(result.unwrap().calls[0].arguments["block_until_ms"], 30000); } +#[tokio::test] +async fn openai_responses_accepts_empty_arguments_done_and_eof_after_completed_tool() { + let arguments = + r#"{"merge":false,"todos":[{"id":"first","content":"First","status":"pending"}]}"#; + let stream = format!( + concat!( + "data: {{\"type\":\"response.output_item.added\",\"output_index\":0,\"item\":{{\"type\":\"function_call\",\"call_id\":\"call-1\",\"name\":\"TodoWrite\"}}}}\n\n", + "data: {{\"type\":\"response.function_call_arguments.delta\",\"output_index\":0,\"delta\":{0:?}}}\n\n", + "data: {{\"type\":\"response.function_call_arguments.done\",\"output_index\":0,\"arguments\":\"\"}}\n\n", + "data: {{\"type\":\"response.output_item.done\",\"output_index\":0,\"item\":{{\"type\":\"function_call\",\"call_id\":\"call-1\",\"name\":\"TodoWrite\",\"arguments\":{0:?}}}}}\n\n", + ), + arguments, + ); + let stream = Box::leak(stream.into_boxed_str()); + let (base_url, _requests, server) = fixture_server("/v1/responses", stream).await; + let provider = OpenAiResponsesProvider::new( + reqwest::Client::new(), + config(ProviderKind::OpenAiResponses, base_url, None), + ); + let (sender, _receiver) = tokio::sync::mpsc::channel(32); + + let cycle = consume_model_cycle( + provider.stream(invocation(), CancellationToken::new()), + &sender, + &CancellationToken::new(), + ) + .await + .unwrap(); + server.abort(); + + assert_eq!(cycle.calls[0].name, "TodoWrite"); + assert_eq!(cycle.calls[0].arguments["todos"][0]["content"], "First"); +} + #[tokio::test] async fn openai_responses_completed_snapshot_does_not_reindex_streamed_tool() { let (base_url, _requests, server) = fixture_server(