fix: stabilize todo state and responses streams

This commit is contained in:
leokun
2026-08-24 18:43:44 +08:00
parent 95fb9be967
commit eb26b17ba0
3 changed files with 177 additions and 11 deletions
+120 -1
View File
@@ -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<Value>, 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,
}
}
}
+22 -9
View File
@@ -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,12 +196,25 @@ 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);
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)); }
if thinking_open { thinking_open = false; yield ModelEvent::ThinkingEnd; }
@@ -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() {
+34
View File
@@ -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(