From 228de238e167296b1e738e0b0652a3b201fed4b8 Mon Sep 17 00:00:00 2001 From: Francisco Javier Arceo Date: Mon, 3 Aug 2026 10:12:04 -0400 Subject: [PATCH] fix: reject stale conversation turns Capture the persisted conversation version during rehydration and verify it under the existing storage lock before appending. Return conversation_locked consistently across HTTP, SSE, and WebSocket when the history changes, with SQLite and PostgreSQL race coverage. Signed-off-by: Francisco Javier Arceo --- Cargo.lock | 1 + .../src/executor/compaction.rs | 1 + .../agentic-server-core/src/executor/error.rs | 122 ++++- .../src/executor/gateway_accumulator.rs | 47 +- .../agentic-server-core/src/executor/mod.rs | 2 +- .../src/executor/modes/conversation.rs | 115 ++++- .../src/executor/modes/response.rs | 4 +- .../src/executor/persist.rs | 21 +- .../src/executor/rehydrate.rs | 116 ++++- .../src/executor/request.rs | 7 +- .../src/storage/conversation.rs | 83 +++- crates/agentic-server-core/src/storage/mod.rs | 5 +- .../src/storage/models/item.rs | 14 + .../src/storage/response.rs | 19 +- .../src/storage/types/conversation.rs | 28 ++ .../src/storage/types/errors.rs | 14 + .../src/storage/types/item.rs | 21 + .../src/storage/types/mod.rs | 2 +- .../tests/postgres_storage_integration.rs | 129 ++++- .../stateful_conversation_integration.rs | 3 +- .../tests/stateful_responses_integration.rs | 1 + .../tests/storage_integration.rs | 293 +++++++++++- .../tests/tool_normalization_test.rs | 1 + crates/agentic-server/Cargo.toml | 1 + crates/agentic-server/src/app.rs | 44 ++ .../src/handler/websocket/error.rs | 77 ++- .../src/handler/websocket/responses.rs | 14 +- crates/agentic-server/tests/responses_test.rs | 444 +++++++++++++++++- .../tests/responses_websocket_test.rs | 273 ++++++++++- docs/deploying/container.md | 2 +- 30 files changed, 1843 insertions(+), 61 deletions(-) diff --git a/Cargo.lock b/Cargo.lock index 7a7d66e6..1abef8b9 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -24,6 +24,7 @@ dependencies = [ "reqwest 0.12.28", "serde", "serde_json", + "sqlx", "thiserror", "tokio", "tokio-tungstenite", diff --git a/crates/agentic-server-core/src/executor/compaction.rs b/crates/agentic-server-core/src/executor/compaction.rs index e1a221ca..7d1eb91a 100644 --- a/crates/agentic-server-core/src/executor/compaction.rs +++ b/crates/agentic-server-core/src/executor/compaction.rs @@ -198,6 +198,7 @@ pub(crate) async fn compact_items( new_input_items: Vec::new(), response_id: uuid7_str("resp_"), conversation_id: None, + conversation_version: None, }; let response = fetch_blocking_payload(&ctx, exec_ctx, auth).await?; let summary = completed_summary_text(&response)?; diff --git a/crates/agentic-server-core/src/executor/error.rs b/crates/agentic-server-core/src/executor/error.rs index f2cf5e19..fde513e4 100644 --- a/crates/agentic-server-core/src/executor/error.rs +++ b/crates/agentic-server-core/src/executor/error.rs @@ -19,6 +19,16 @@ pub enum ExecutorError { #[error("failed to persist response")] Persistence(#[source] Box), + /// A persisted conversation changed after its history was read. + /// + /// The storage source is retained for internal diagnostics while the + /// display message remains safe to send to API clients. + #[error("conversation changed while the response was being generated; retry the request")] + ConversationLocked { + #[source] + source: StorageError, + }, + /// The LLM backend returned a non-2xx status or was unreachable. #[error("LLM request failed ({status}): {body}")] LLMRequest { status: StatusCode, body: String }, @@ -71,33 +81,74 @@ pub enum ExecutorError { } impl ExecutorError { + fn client_visible_error(&self) -> &Self { + match self { + Self::Persistence(source) if source.contains_conversation_locked() => source.client_visible_error(), + _ => self, + } + } + + fn contains_conversation_locked(&self) -> bool { + match self { + Self::ConversationLocked { .. } => true, + Self::Persistence(source) => source.contains_conversation_locked(), + _ => false, + } + } + /// HTTP status code that best represents this error to an API caller. #[must_use] pub fn http_status(&self) -> StatusCode { - match self { + match self.client_visible_error() { Self::Storage(e) if e.is_not_found() => StatusCode::NOT_FOUND, Self::LLMRequest { status, .. } => *status, - Self::Tool(ToolError::Config(_)) | Self::InvalidRequest(_) | Self::JsonError(_) => StatusCode::BAD_REQUEST, + Self::ConversationLocked { .. } + | Self::Tool(ToolError::Config(_)) + | Self::InvalidRequest(_) + | Self::JsonError(_) => StatusCode::BAD_REQUEST, Self::Tool(ToolError::Execution(_)) | Self::CompactionFailed { .. } => StatusCode::BAD_GATEWAY, Self::ParseError(_) => StatusCode::UNPROCESSABLE_ENTITY, _ => StatusCode::INTERNAL_SERVER_ERROR, } } - /// Short machine-readable error code for the API error envelope. + /// Machine-readable error type for the API error envelope. #[must_use] - pub fn error_code(&self) -> &'static str { - match self { + pub fn error_type(&self) -> &'static str { + match self.client_visible_error() { + Self::ConversationLocked { .. } + | Self::Tool(ToolError::Config(_)) + | Self::InvalidRequest(_) + | Self::ParseError(_) + | Self::JsonError(_) => "invalid_request_error", Self::Storage(e) if e.is_not_found() => "not_found", Self::LLMRequest { .. } | Self::CompactionFailed { .. } => "upstream_error", - Self::Tool(ToolError::Config(_)) | Self::InvalidRequest(_) | Self::ParseError(_) | Self::JsonError(_) => { - "invalid_request_error" - } Self::Tool(ToolError::Execution(_)) => "tool_error", _ => "server_error", } } + /// Short machine-readable error code for the API error envelope. + #[must_use] + pub fn error_code(&self) -> &'static str { + match self.client_visible_error() { + Self::ConversationLocked { .. } => "conversation_locked", + other => other.error_type(), + } + } + + /// Request parameter associated with the API error, when applicable. + #[must_use] + pub fn error_param(&self) -> Option<&'static str> { + matches!(self.client_visible_error(), Self::ConversationLocked { .. }).then_some("conversation") + } + + /// Client-safe message for the API error envelope. + #[must_use] + pub fn error_message(&self) -> String { + self.client_visible_error().to_string() + } + /// Serialise the error into the HTTP response body bytes. /// /// `LLMRequest` bodies are forwarded verbatim; all other variants are @@ -107,10 +158,16 @@ impl ExecutorError { match self { Self::LLMRequest { body, .. } => body.into_bytes(), other => { + let error_type = other.error_type(); let code = other.error_code(); - serialize_to_vec_or_default(&serde_json::json!({ - "error": { "message": other.to_string(), "type": code, "code": code } - })) + let mut error = serde_json::Map::new(); + error.insert("message".to_owned(), serde_json::json!(other.error_message())); + error.insert("type".to_owned(), serde_json::json!(error_type)); + error.insert("code".to_owned(), serde_json::json!(code)); + if let Some(param) = other.error_param() { + error.insert("param".to_owned(), serde_json::json!(param)); + } + serialize_to_vec_or_default(&serde_json::json!({ "error": error })) } } } @@ -160,4 +217,47 @@ mod tests { assert!(exec_err.source().is_some(), "source should be chained"); assert!(exec_err.to_string().contains("json error")); } + + #[test] + fn conversation_locked_response_preserves_conflict_through_persistence() { + use std::error::Error; + + let error = ExecutorError::Persistence(Box::new(ExecutorError::ConversationLocked { + source: StorageError::ConversationConflict { + conversation_id: "conv_internal".to_owned(), + }, + })); + + let conversation_locked = error.source().expect("persistence source must be retained"); + let conflict = conversation_locked + .source() + .expect("conversation conflict source must be retained"); + assert!(matches!( + conflict.downcast_ref::(), + Some(StorageError::ConversationConflict { conversation_id }) + if conversation_id == "conv_internal" + )); + + assert_eq!(error.http_status(), StatusCode::BAD_REQUEST); + assert_eq!( + serde_json::from_slice::(&error.into_response_body()) + .expect("valid error response JSON"), + serde_json::json!({ + "error": { + "message": "conversation changed while the response was being generated; retry the request", + "type": "invalid_request_error", + "code": "conversation_locked", + "param": "conversation" + } + }) + ); + } + + #[test] + fn non_conflict_response_omits_param() { + let body = ExecutorError::InvalidRequest("invalid input".to_owned()).into_response_body(); + let value: serde_json::Value = serde_json::from_slice(&body).expect("valid error response JSON"); + + assert!(!value["error"].as_object().expect("error object").contains_key("param")); + } } diff --git a/crates/agentic-server-core/src/executor/gateway_accumulator.rs b/crates/agentic-server-core/src/executor/gateway_accumulator.rs index 15ff30b5..9512fbc9 100644 --- a/crates/agentic-server-core/src/executor/gateway_accumulator.rs +++ b/crates/agentic-server-core/src/executor/gateway_accumulator.rs @@ -114,18 +114,20 @@ fn terminal_response_frame(payload: &ResponsePayload) -> ExecutorResult EventFrame { + let error_type = error.error_type(); let code = error.error_code(); let mut wire = WireEvent::new("error"); wire.rest .insert("status".to_owned(), serde_json::json!(error.http_status().as_u16())); - wire.rest.insert( - "error".to_owned(), - serde_json::json!({ - "message": error.to_string(), - "type": code, - "code": code, - }), - ); + let mut error_details = serde_json::Map::new(); + error_details.insert("message".to_owned(), serde_json::json!(error.error_message())); + error_details.insert("type".to_owned(), serde_json::json!(error_type)); + error_details.insert("code".to_owned(), serde_json::json!(code)); + if let Some(param) = error.error_param() { + error_details.insert("param".to_owned(), serde_json::json!(param)); + } + wire.rest + .insert("error".to_owned(), serde_json::Value::Object(error_details)); EventFrame { event_type: SSEEventType::Other, payload: EventPayload::None, @@ -187,6 +189,7 @@ fn serialize_sse_frame(frame: &EventFrame) -> ExecutorResult { #[cfg(test)] mod tests { use super::*; + use crate::StorageError; #[test] fn process_sse_line_numbers_and_rebases_output_index() { @@ -218,6 +221,34 @@ mod tests { assert_eq!(event["error"]["message"], "task failed: \"unexpected\"\nretry"); } + #[test] + fn executor_conflict_sse_chunk_uses_client_conflict_contract() { + let mut accumulator = GatewayStreamAccumulator::new(); + let error = ExecutorError::Persistence(Box::new(ExecutorError::ConversationLocked { + source: StorageError::ConversationConflict { + conversation_id: "conv_test".to_owned(), + }, + })); + let chunk = accumulator.executor_error_chunk(&error); + let data = chunk + .trim_end_matches('\n') + .strip_prefix("data: ") + .expect("SSE data prefix"); + let event: serde_json::Value = serde_json::from_str(data).expect("valid error event JSON"); + + assert_eq!(event["type"], "error"); + assert_eq!(event["status"], 400); + assert_eq!( + event["error"], + serde_json::json!({ + "message": "conversation changed while the response was being generated; retry the request", + "type": "invalid_request_error", + "code": "conversation_locked", + "param": "conversation" + }) + ); + } + #[test] fn emits_in_progress_terminal_event_after_lifecycle_event() { let mut accumulator = GatewayStreamAccumulator::new(); diff --git a/crates/agentic-server-core/src/executor/mod.rs b/crates/agentic-server-core/src/executor/mod.rs index e45f8633..72fca307 100644 --- a/crates/agentic-server-core/src/executor/mod.rs +++ b/crates/agentic-server-core/src/executor/mod.rs @@ -23,7 +23,7 @@ pub use inference::call_inference; pub use messages_loop::run_messages_loop; pub use messages_stream::run_messages_stream; pub use modes::{ConversationHandler, ResponseHandler}; -pub use persist::persist_response; +pub use persist::{persist_response, persist_turn}; pub use rehydrate::rehydrate_conversation; pub use request::ExecutionContext; pub use request::RequestContext; diff --git a/crates/agentic-server-core/src/executor/modes/conversation.rs b/crates/agentic-server-core/src/executor/modes/conversation.rs index 1e185708..e5479fab 100644 --- a/crates/agentic-server-core/src/executor/modes/conversation.rs +++ b/crates/agentic-server-core/src/executor/modes/conversation.rs @@ -1,6 +1,8 @@ //! Conversation storage handler — owns all conversation store operations. -use crate::storage::{ConversationData, ConversationStore, InOutItem, ResponseMetadata}; +use crate::storage::{ + ConversationData, ConversationSnapshot, ConversationStore, InOutItem, ResponseMetadata, StorageError, +}; use crate::types::io::OutputItem; use crate::executor::error::{ExecutorError, ExecutorResult}; @@ -67,12 +69,26 @@ impl ConversationHandler { /// Returns `ExecutorError` if `conversation_id` is absent, the store is /// disabled, or the database query fails. pub async fn rehydrate(&self, ctx: &RequestContext) -> ExecutorResult> { + Ok(self.rehydrate_snapshot(ctx).await?.items) + } + + /// Loads the conversation's history items and storage version. + /// + /// Reads `conversation_id` from `ctx.original_request`. + /// + /// # Errors + /// Returns `ExecutorError` if `conversation_id` is absent, the store is + /// disabled, or the database query fails. + pub async fn rehydrate_snapshot(&self, ctx: &RequestContext) -> ExecutorResult { let conv_id = ctx .original_request .conversation_id .as_deref() .ok_or_else(|| ExecutorError::InvalidRequest("conversation_id is required for rehydrate".into()))?; - self.store.rehydrate(conv_id).await.map_err(ExecutorError::Storage) + self.store + .rehydrate_snapshot(conv_id) + .await + .map_err(ExecutorError::Storage) } /// Persists one conversation turn — only the new items from this turn. @@ -88,6 +104,9 @@ impl ConversationHandler { let conversation_id = ctx .conversation_id .ok_or_else(|| ExecutorError::InvalidRequest("conversation_id is required for execute_turn".into()))?; + let conversation_version = ctx + .conversation_version + .ok_or_else(|| ExecutorError::InvalidRequest("conversation version is required for execute_turn".into()))?; let metadata = ResponseMetadata { model: ctx.enriched_request.model, @@ -102,21 +121,26 @@ impl ConversationHandler { new_items.extend(output_items.into_iter().map(InOutItem::Output)); self.store - .persist( + .persist_if_version( &conversation_id, + conversation_version, &ctx.response_id, metadata.previous_response_id.as_deref(), new_items, &metadata, ) .await - .map_err(ExecutorError::Storage) + .map_err(|error| match error { + source @ StorageError::ConversationConflict { .. } => ExecutorError::ConversationLocked { source }, + other => ExecutorError::Storage(other), + }) } } #[cfg(test)] mod tests { use super::*; + use crate::storage::{ConversationVersion, create_pool_with_schema}; use crate::types::io::ResponsesInput; use crate::types::request_response::RequestPayload; @@ -151,6 +175,7 @@ mod tests { new_input_items: vec![], response_id: "resp_test".into(), conversation_id: conversation_id.map(str::to_string), + conversation_version: None, } } @@ -189,4 +214,86 @@ mod tests { let result = disabled_handler().execute_turn(make_ctx(None), vec![]).await; assert!(result.is_err()); } + + #[tokio::test] + async fn execute_turn_rejects_missing_conversation_version_without_writing() + -> Result<(), Box> { + let pool = create_pool_with_schema(Some("sqlite://?mode=memory")).await?; + let store = ConversationStore::new(pool); + let conversation = store.create().await?; + let handler = ConversationHandler::new(store.clone()); + let mut ctx = make_ctx(Some(&conversation.conversation_id)); + ctx.new_input_items = Vec::from(&ctx.original_request.input); + + let error = handler + .execute_turn(ctx, vec![]) + .await + .expect_err("missing captured version must reject the turn"); + + assert!(matches!( + error, + ExecutorError::InvalidRequest(message) + if message == "conversation version is required for execute_turn" + )); + assert!(store.rehydrate(&conversation.conversation_id).await?.is_empty()); + Ok(()) + } + + #[tokio::test] + async fn execute_turn_persists_with_captured_conversation_version() -> Result<(), Box> { + let pool = create_pool_with_schema(Some("sqlite://?mode=memory")).await?; + let store = ConversationStore::new(pool); + let conversation = store.create().await?; + let handler = ConversationHandler::new(store.clone()); + let mut ctx = make_ctx(Some(&conversation.conversation_id)); + ctx.new_input_items = Vec::from(&ctx.original_request.input); + ctx.conversation_version = Some(ConversationVersion::Empty); + + handler.execute_turn(ctx, vec![]).await?; + + let snapshot = store.rehydrate_snapshot(&conversation.conversation_id).await?; + assert_eq!(snapshot.items.len(), 1); + assert_eq!(snapshot.version, ConversationVersion::LastSequence(0)); + Ok(()) + } + + #[tokio::test] + async fn execute_turn_rejects_a_stale_captured_conversation_version() -> Result<(), Box> { + use std::error::Error; + + let pool = create_pool_with_schema(Some("sqlite://?mode=memory")).await?; + let store = ConversationStore::new(pool); + let conversation = store.create().await?; + let handler = ConversationHandler::new(store.clone()); + let mut ctx = make_ctx(Some(&conversation.conversation_id)); + ctx.new_input_items = Vec::from(&ctx.original_request.input); + ctx.conversation_version = Some(ConversationVersion::Empty); + let competing_items = Vec::from(&ResponsesInput::Text("competing input".into())) + .into_iter() + .map(InOutItem::Input) + .collect(); + store + .persist( + &conversation.conversation_id, + "resp_competing", + None, + competing_items, + &ResponseMetadata::default(), + ) + .await?; + + let error = handler + .execute_turn(ctx, vec![]) + .await + .expect_err("stale captured version must reject the turn"); + + let source = error.source().expect("conversation conflict source must be retained"); + assert!(matches!( + source.downcast_ref::(), + Some(StorageError::ConversationConflict { conversation_id }) + if conversation_id == &conversation.conversation_id + )); + assert!(matches!(error, ExecutorError::ConversationLocked { .. })); + Ok(()) + } } diff --git a/crates/agentic-server-core/src/executor/modes/response.rs b/crates/agentic-server-core/src/executor/modes/response.rs index eb933e22..842634ed 100644 --- a/crates/agentic-server-core/src/executor/modes/response.rs +++ b/crates/agentic-server-core/src/executor/modes/response.rs @@ -82,8 +82,9 @@ impl ResponseHandler { new_items.extend(output_items.into_iter().map(InOutItem::Output)); self.store - .persist( + .persist_with_conversation_id( &ctx.response_id, + ctx.conversation_id.as_deref(), metadata.previous_response_id.as_deref(), new_items, &metadata, @@ -130,6 +131,7 @@ mod tests { new_input_items: vec![], response_id: "resp_test".into(), conversation_id: None, + conversation_version: None, } } diff --git a/crates/agentic-server-core/src/executor/persist.rs b/crates/agentic-server-core/src/executor/persist.rs index 0cc9d44f..d67ab5e8 100644 --- a/crates/agentic-server-core/src/executor/persist.rs +++ b/crates/agentic-server-core/src/executor/persist.rs @@ -7,6 +7,7 @@ use crate::executor::error::{ExecutorError, ExecutorResult}; use crate::executor::modes::{ConversationHandler, ResponseHandler}; use crate::executor::request::RequestContext; use crate::types::event::ResponseStatus; +use crate::types::io::OutputItem; use crate::types::request_response::ResponsePayload; use tracing::error; @@ -38,8 +39,8 @@ pub(crate) async fn persist_if_needed( /// Step 3 — Persist the completed response to storage. /// /// Skipped if [`ResponseStatus`] is not `Completed`/`Incomplete` or `payload.id` is empty. -/// Routes to [`ConversationHandler`] when `ctx.conversation_id` is set, -/// otherwise [`ResponseHandler`]. +/// Routes explicit `conversation_id` requests to [`ConversationHandler`] and +/// all other requests, including `previous_response_id` continuations, to [`ResponseHandler`]. /// /// # Errors /// Returns [`ExecutorError`] if the storage operation fails. @@ -58,10 +59,20 @@ pub async fn persist_response( return Ok(()); } - // Move output items from payload; handlers build ResponseMetadata from ctx internally. - let output_items = payload.output; + persist_turn(ctx, payload.output, &conv_handler, &resp_handler).await +} - if ctx.conversation_id.is_some() { +/// Persists one completed turn with the handler selected by its explicit conversation discriminator. +/// +/// # Errors +/// Returns [`ExecutorError`] if the selected storage operation fails. +pub async fn persist_turn( + ctx: RequestContext, + output_items: Vec, + conv_handler: &ConversationHandler, + resp_handler: &ResponseHandler, +) -> ExecutorResult<()> { + if ctx.original_request.conversation_id.is_some() { conv_handler.execute_turn(ctx, output_items).await } else { resp_handler.execute_turn(ctx, output_items).await diff --git a/crates/agentic-server-core/src/executor/rehydrate.rs b/crates/agentic-server-core/src/executor/rehydrate.rs index 933db1ee..6b9fdf08 100644 --- a/crates/agentic-server-core/src/executor/rehydrate.rs +++ b/crates/agentic-server-core/src/executor/rehydrate.rs @@ -38,6 +38,7 @@ pub async fn rehydrate_conversation( new_input_items, response_id, conversation_id: None, + conversation_version: None, }; if ctx.original_request.conversation_id.is_some() && ctx.original_request.previous_response_id.is_some() { @@ -94,7 +95,7 @@ async fn from_response(ctx: &mut RequestContext, exec_ctx: &ExecutionContext) -> /// Gets or creates the conversation (depending on `store`) and rehydrates its /// history in parallel, then prepends the history items to the enriched request input. async fn from_conversation(ctx: &mut RequestContext, exec_ctx: &ExecutionContext) -> ExecutorResult<()> { - let (conv_data, history) = tokio::try_join!( + let (conv_data, snapshot) = tokio::try_join!( async { if ctx.original_request.store { exec_ctx.conv_handler.get_or_create(ctx).await @@ -102,14 +103,123 @@ async fn from_conversation(ctx: &mut RequestContext, exec_ctx: &ExecutionContext exec_ctx.conv_handler.get(ctx).await } }, - exec_ctx.conv_handler.rehydrate(ctx), + exec_ctx.conv_handler.rehydrate_snapshot(ctx), )?; - let mut items = InOutItem::into_input_items(history); + let mut items = InOutItem::into_input_items(snapshot.items); items.reserve(ctx.new_input_items.len()); items.extend(ctx.new_input_items.iter().cloned()); ctx.enriched_request.input = ResponsesInput::Items(items); ctx.conversation_id = Some(conv_data.conversation_id); + ctx.conversation_version = Some(snapshot.version); Ok(()) } + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use super::*; + use crate::executor::modes::{ConversationHandler, ResponseHandler}; + use crate::storage::{ + ConversationStore, ConversationVersion, InOutItem, ResponseMetadata, ResponseStore, create_pool_with_schema, + }; + use crate::types::request_response::RequestPayload; + + fn request(conversation_id: Option<&str>, previous_response_id: Option<&str>) -> RequestPayload { + RequestPayload { + model: "test".into(), + input: ResponsesInput::Text("new input".into()), + instructions: None, + previous_response_id: previous_response_id.map(str::to_owned), + conversation_id: conversation_id.map(str::to_owned), + tools: None, + tool_choice: None, + stream: false, + store: true, + include: None, + temperature: None, + top_p: None, + max_output_tokens: None, + truncation: None, + metadata: None, + parallel_tool_calls: None, + cache_salt: None, + context_management: None, + } + } + + fn execution_context(conversation_store: ConversationStore, response_store: ResponseStore) -> ExecutionContext { + ExecutionContext::new( + ConversationHandler::new(conversation_store), + ResponseHandler::new(response_store), + Arc::new(reqwest::Client::new()), + "http://localhost:8000".to_owned(), + ) + } + + #[tokio::test] + async fn new_conversation_rehydration_captures_empty_version() -> Result<(), Box> { + let pool = create_pool_with_schema(Some("sqlite://?mode=memory")).await?; + let conversation_store = ConversationStore::new(pool); + let conversation = conversation_store.create().await?; + let exec_ctx = execution_context(conversation_store, ResponseStore::disabled()); + + let ctx = rehydrate_conversation(request(Some(&conversation.conversation_id), None), &exec_ctx).await?; + + assert_eq!(ctx.conversation_version, Some(ConversationVersion::Empty)); + Ok(()) + } + + #[tokio::test] + async fn existing_conversation_rehydration_captures_last_sequence() -> Result<(), Box> { + let pool = create_pool_with_schema(Some("sqlite://?mode=memory")).await?; + let conversation_store = ConversationStore::new(pool); + let conversation = conversation_store.create().await?; + let prior_items = Vec::::from(&ResponsesInput::Text("prior input".into())) + .into_iter() + .map(InOutItem::Input) + .collect(); + conversation_store + .persist( + &conversation.conversation_id, + "resp_prior", + None, + prior_items, + &ResponseMetadata::default(), + ) + .await?; + let exec_ctx = execution_context(conversation_store, ResponseStore::disabled()); + + let ctx = rehydrate_conversation(request(Some(&conversation.conversation_id), None), &exec_ctx).await?; + + assert_eq!(ctx.conversation_version, Some(ConversationVersion::LastSequence(0))); + Ok(()) + } + + #[tokio::test] + async fn request_without_continuation_has_no_conversation_version() -> Result<(), Box> { + let exec_ctx = execution_context(ConversationStore::disabled(), ResponseStore::disabled()); + + let ctx = rehydrate_conversation(request(None, None), &exec_ctx).await?; + + assert_eq!(ctx.conversation_version, None); + Ok(()) + } + + #[tokio::test] + async fn previous_response_rehydration_has_no_conversation_version() -> Result<(), Box> { + let pool = create_pool_with_schema(Some("sqlite://?mode=memory")).await?; + let response_store = ResponseStore::new(pool); + response_store + .persist("resp_prior", None, Vec::new(), &ResponseMetadata::default()) + .await?; + let exec_ctx = execution_context(ConversationStore::disabled(), response_store); + + let ctx = rehydrate_conversation(request(None, Some("resp_prior")), &exec_ctx).await?; + + assert_eq!(ctx.conversation_version, None); + Ok(()) + } +} diff --git a/crates/agentic-server-core/src/executor/request.rs b/crates/agentic-server-core/src/executor/request.rs index 2735ff03..e7e1cc9b 100644 --- a/crates/agentic-server-core/src/executor/request.rs +++ b/crates/agentic-server-core/src/executor/request.rs @@ -5,7 +5,9 @@ use crate::config::Config; use crate::error::Error; use crate::executor::modes::{ConversationHandler, ResponseHandler}; use crate::storage::backend::redact_database_urls; -use crate::storage::{ConversationStore, DatabaseBackend, ResponseStore, create_pool_with_schema_and_configs}; +use crate::storage::{ + ConversationStore, ConversationVersion, DatabaseBackend, ResponseStore, create_pool_with_schema_and_configs, +}; use crate::tool::{GatewayExecutor, GatewayExecutors}; use crate::types::io::InputItem; use crate::types::messages::GatewayToolMap; @@ -30,6 +32,9 @@ pub struct RequestContext { pub response_id: String, /// Resolved conversation ID. `None` when `store=false` or non-conversational. pub conversation_id: Option, + /// Conversation version captured with rehydrated history. + /// `None` for non-conversation and `previous_response_id` execution. + pub conversation_version: Option, } impl RequestContext { diff --git a/crates/agentic-server-core/src/storage/conversation.rs b/crates/agentic-server-core/src/storage/conversation.rs index c013d8b3..4792dec2 100644 --- a/crates/agentic-server-core/src/storage/conversation.rs +++ b/crates/agentic-server-core/src/storage/conversation.rs @@ -5,7 +5,9 @@ use std::sync::Arc; use super::models::{conversation, item, response}; use super::pool::DbPool; -use super::types::{ConversationData, InOutItem, ResponseMetadata, StorageError, StoreResult}; +use super::types::{ + ConversationData, ConversationSnapshot, ConversationVersion, InOutItem, ResponseMetadata, StorageError, StoreResult, +}; use crate::utils::common::{serialize_to_string, uuid7_str}; /// Conversation storage operations. @@ -75,12 +77,32 @@ impl ConversationStore { /// /// # Errors /// - /// Returns error if conversation not found or database query fails. + /// Returns an error if a stored item is missing its sequence number or if the database query fails. pub async fn rehydrate(&self, conversation_id: &str) -> StoreResult> { + Ok(self.rehydrate_snapshot(conversation_id).await?.items) + } + + /// Rehydrates a conversation with its items and storage version. + /// + /// # Errors + /// + /// Returns an error if a stored item is missing its sequence number or if the database query fails. + pub async fn rehydrate_snapshot(&self, conversation_id: &str) -> StoreResult { let pool = self.pool()?; let rows = item::get_items_by_conversation(pool, conversation_id).await?; - Ok(rows.into_iter().filter_map(|row| row.as_inout()).collect()) + let mut last_sequence = None; + for row in &rows { + last_sequence = Some(row.seq.ok_or_else(|| StorageError::InvalidConversationSequence { + conversation_id: conversation_id.to_string(), + item_id: row.id.clone(), + })?); + } + + Ok(ConversationSnapshot { + items: rows.into_iter().filter_map(|row| row.as_inout()).collect(), + version: ConversationVersion::from_last_sequence(last_sequence), + }) } /// Persists conversation turn with new items and response metadata. @@ -97,6 +119,51 @@ impl ConversationStore { previous_response_id: Option<&str>, new_items: Vec, metadata: &ResponseMetadata, + ) -> StoreResult<()> { + self.persist_impl( + conversation_id, + None, + response_id, + previous_response_id, + new_items, + metadata, + ) + .await + } + + /// Persists a conversation turn only if its stored version still matches. + /// + /// # Errors + /// + /// Returns [`StorageError`] if the conversation changed, was not found, or a database operation fails. + pub async fn persist_if_version( + &self, + conversation_id: &str, + expected_version: ConversationVersion, + response_id: &str, + previous_response_id: Option<&str>, + new_items: Vec, + metadata: &ResponseMetadata, + ) -> StoreResult<()> { + self.persist_impl( + conversation_id, + Some(expected_version), + response_id, + previous_response_id, + new_items, + metadata, + ) + .await + } + + async fn persist_impl( + &self, + conversation_id: &str, + expected_version: Option, + response_id: &str, + previous_response_id: Option<&str>, + new_items: Vec, + metadata: &ResponseMetadata, ) -> StoreResult<()> { let pool = self.pool()?; @@ -120,6 +187,16 @@ impl ConversationStore { } Err(error) => return Err(error.into()), } + if let Some(expected_version) = expected_version { + let current_version = ConversationVersion::from_last_sequence( + item::last_conversation_sequence_in_tx(&mut tx, conversation_id).await?, + ); + if current_version != expected_version { + return Err(StorageError::ConversationConflict { + conversation_id: conversation_id.to_owned(), + }); + } + } item::create_in_tx(&mut tx, items_, Some(conversation_id)).await?; response::create_in_tx( diff --git a/crates/agentic-server-core/src/storage/mod.rs b/crates/agentic-server-core/src/storage/mod.rs index a34c4661..3a44cf02 100644 --- a/crates/agentic-server-core/src/storage/mod.rs +++ b/crates/agentic-server-core/src/storage/mod.rs @@ -33,4 +33,7 @@ pub use pool::{ }; pub use response::ResponseStore; pub use schema::{PoolWithSchema, SchemaManager}; -pub use types::{ConversationData, InOutItem, ItemKind, ResponseData, ResponseMetadata, StorageError, StoreResult}; +pub use types::{ + ConversationData, ConversationSnapshot, ConversationVersion, InOutItem, ItemKind, ResponseData, ResponseMetadata, + StorageError, StoreResult, +}; diff --git a/crates/agentic-server-core/src/storage/models/item.rs b/crates/agentic-server-core/src/storage/models/item.rs index 01724b9f..e6cc81a8 100644 --- a/crates/agentic-server-core/src/storage/models/item.rs +++ b/crates/agentic-server-core/src/storage/models/item.rs @@ -232,6 +232,20 @@ pub async fn get_items_by_conversation(pool: &DbPool, conversation_id: &str) -> .await } +/// Returns the last stored item sequence for a conversation inside a transaction. +/// +/// # Errors +/// Returns `DbResult::Err` if the database query fails. +pub async fn last_conversation_sequence_in_tx( + tx: &mut DbTransaction<'_>, + conversation_id: &str, +) -> DbResult> { + sqlx::query_scalar("SELECT MAX(seq) FROM items WHERE conversation_id = $1") + .bind(conversation_id) + .fetch_one(&mut **tx) + .await +} + #[cfg(test)] mod tests { use super::*; diff --git a/crates/agentic-server-core/src/storage/response.rs b/crates/agentic-server-core/src/storage/response.rs index 6274a82e..def95991 100644 --- a/crates/agentic-server-core/src/storage/response.rs +++ b/crates/agentic-server-core/src/storage/response.rs @@ -97,6 +97,23 @@ impl ResponseStore { previous_response_id: Option<&str>, new_items: Vec, metadata: &ResponseMetadata, + ) -> StoreResult<()> { + self.persist_with_conversation_id(response_id, None, previous_response_id, new_items, metadata) + .await + } + + /// Persists a response while retaining its inherited conversation ID. + /// + /// # Errors + /// + /// Returns [`StorageError`] if database operation fails or store is disabled. + pub(crate) async fn persist_with_conversation_id( + &self, + response_id: &str, + conversation_id: Option<&str>, + previous_response_id: Option<&str>, + new_items: Vec, + metadata: &ResponseMetadata, ) -> StoreResult<()> { let pool = self.pool()?; @@ -121,7 +138,7 @@ impl ResponseStore { response::create_in_tx( &mut tx, response_id, - None, + conversation_id, previous_response_id, Some(&history_item_ids_json), Some(&metadata_json), diff --git a/crates/agentic-server-core/src/storage/types/conversation.rs b/crates/agentic-server-core/src/storage/types/conversation.rs index 55dc9093..71c0f590 100644 --- a/crates/agentic-server-core/src/storage/types/conversation.rs +++ b/crates/agentic-server-core/src/storage/types/conversation.rs @@ -1,6 +1,34 @@ //! Domain type for conversation storage. use super::super::models::Conversation as StorageDbConversation; +use super::item::InOutItem; + +/// Version of a conversation's stored item history. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum ConversationVersion { + /// The conversation has no stored items. + Empty, + /// The sequence number of the last stored item. + LastSequence(i64), +} + +impl ConversationVersion { + pub(crate) const fn from_last_sequence(last_sequence: Option) -> Self { + match last_sequence { + Some(sequence) => Self::LastSequence(sequence), + None => Self::Empty, + } + } +} + +/// Rehydrated conversation items together with their storage version. +#[derive(Debug, Clone, PartialEq)] +pub struct ConversationSnapshot { + /// Rehydrated input and output items in stored order. + pub items: Vec, + /// Version derived from the last stored item sequence. + pub version: ConversationVersion, +} /// Domain entity for a stored conversation. /// diff --git a/crates/agentic-server-core/src/storage/types/errors.rs b/crates/agentic-server-core/src/storage/types/errors.rs index 4bf5fb06..e064497e 100644 --- a/crates/agentic-server-core/src/storage/types/errors.rs +++ b/crates/agentic-server-core/src/storage/types/errors.rs @@ -15,6 +15,14 @@ pub enum StorageError { #[error("not found: {resource_type} with id '{id}'")] NotFound { resource_type: String, id: String }, + /// A conversation item did not have its required sequence number. + #[error("invalid conversation sequence for conversation '{conversation_id}' item '{item_id}'")] + InvalidConversationSequence { conversation_id: String, item_id: String }, + + /// A conversation changed after its version was read. + #[error("conversation changed while the response was being generated")] + ConversationConflict { conversation_id: String }, + /// Database operation failed. /// /// Wraps `sqlx::Error` and automatically converts from it via `#[from]`. @@ -55,6 +63,12 @@ impl StorageError { matches!(self, Self::NotConfigured) } + /// Returns `true` if this error is a conversation version conflict. + #[must_use] + pub fn is_conversation_conflict(&self) -> bool { + matches!(self, Self::ConversationConflict { .. }) + } + /// Extracts the resource type and ID if this is a "not found" error. #[must_use] pub fn not_found_details(&self) -> Option<(String, String)> { diff --git a/crates/agentic-server-core/src/storage/types/item.rs b/crates/agentic-server-core/src/storage/types/item.rs index df654c49..5159d4ea 100644 --- a/crates/agentic-server-core/src/storage/types/item.rs +++ b/crates/agentic-server-core/src/storage/types/item.rs @@ -7,6 +7,7 @@ use serde_json::Value; use crate::storage::StorageError; use crate::types::io::{InputItem, OutputItem}; +use crate::utils::common::serialize_to_value; pub(crate) const STORED_ITEM_KIND_KEY: &str = "_agentic_item_kind"; @@ -42,6 +43,26 @@ pub enum InOutItem { Output(OutputItem), } +fn serialized_values_equal(left: &T, right: &T) -> bool { + let Ok(left) = serialize_to_value(left) else { + return false; + }; + let Ok(right) = serialize_to_value(right) else { + return false; + }; + left == right +} + +impl PartialEq for InOutItem { + fn eq(&self, other: &Self) -> bool { + match (self, other) { + (Self::Input(left), Self::Input(right)) => serialized_values_equal(left, right), + (Self::Output(left), Self::Output(right)) => serialized_values_equal(left, right), + _ => false, + } + } +} + impl From for InOutItem { fn from(item: InputItem) -> Self { Self::Input(item) diff --git a/crates/agentic-server-core/src/storage/types/mod.rs b/crates/agentic-server-core/src/storage/types/mod.rs index 71f038a8..41dfa538 100644 --- a/crates/agentic-server-core/src/storage/types/mod.rs +++ b/crates/agentic-server-core/src/storage/types/mod.rs @@ -5,7 +5,7 @@ pub mod errors; pub mod item; pub mod response; -pub use conversation::ConversationData; +pub use conversation::{ConversationData, ConversationSnapshot, ConversationVersion}; pub use errors::{StorageError, StoreResult}; pub use item::{InOutItem, ItemKind}; pub use response::{ResponseData, ResponseMetadata}; diff --git a/crates/agentic-server-core/tests/postgres_storage_integration.rs b/crates/agentic-server-core/tests/postgres_storage_integration.rs index e39ec55a..75f22ff0 100644 --- a/crates/agentic-server-core/tests/postgres_storage_integration.rs +++ b/crates/agentic-server-core/tests/postgres_storage_integration.rs @@ -3,7 +3,7 @@ use std::time::Duration; use agentic_core::config::{PostgresConfig, SqliteConfig}; use agentic_core::storage::{ - ConversationStore, InOutItem, ResponseMetadata, ResponseStore, create_pool_with_configs, + ConversationStore, InOutItem, ResponseMetadata, ResponseStore, StorageError, create_pool_with_configs, create_pool_with_schema_and_configs, }; use agentic_core::types::io::{InputItem, InputMessage, InputMessageContent}; @@ -367,6 +367,133 @@ async fn postgres_concurrent_conversation_writes_have_contiguous_sequences() { second_pool.close().await; } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[ignore = "requires TEST_POSTGRES_URL pointing to an isolated PostgreSQL database"] +#[allow( + clippy::too_many_lines, + reason = "keeps the complete two-pool race and persistence assertions in one integration test" +)] +async fn postgres_optimistic_conversation_conflict() { + let database_url = std::env::var("TEST_POSTGRES_URL").expect("TEST_POSTGRES_URL must be set"); + let postgres_config = PostgresConfig { + max_connections: 2, + acquire_timeout: Duration::from_secs(5), + lock_timeout: Duration::from_secs(1), + migration_timeout: Duration::from_secs(5), + statement_timeout: Duration::from_secs(5), + idle_timeout: Some(Duration::from_secs(30)), + max_lifetime: Some(Duration::from_secs(60)), + }; + let first_pool = create_pool_with_schema_and_configs(Some(&database_url), SqliteConfig::default(), postgres_config) + .await + .expect("initialize PostgreSQL database"); + let second_pool = create_pool_with_configs(Some(&database_url), SqliteConfig::default(), postgres_config) + .await + .expect("create independent PostgreSQL pool"); + let store = ConversationStore::new(first_pool.clone()); + let conversation = store.create().await.expect("create conversation"); + let version = store + .rehydrate_snapshot(&conversation.conversation_id) + .await + .expect("capture conversation version") + .version; + let barrier = Arc::new(Barrier::new(2)); + + let writer_one = { + let store = ConversationStore::new(first_pool.clone()); + let conversation_id = conversation.conversation_id.clone(); + let barrier = Arc::clone(&barrier); + tokio::spawn(async move { + let response_id = format!("resp_postgres_{}", uuid::Uuid::now_v7()); + let items = vec![input_item("writer one input"), input_item("writer one follow-up")]; + barrier.wait().await; + let result = store + .persist_if_version( + &conversation_id, + version, + &response_id, + None, + items.clone(), + &ResponseMetadata::default(), + ) + .await; + (result, response_id, items) + }) + }; + let writer_two = { + let store = ConversationStore::new(second_pool.clone()); + let conversation_id = conversation.conversation_id.clone(); + let barrier = Arc::clone(&barrier); + tokio::spawn(async move { + let response_id = format!("resp_postgres_{}", uuid::Uuid::now_v7()); + let items = vec![input_item("writer two input"), input_item("writer two follow-up")]; + barrier.wait().await; + let result = store + .persist_if_version( + &conversation_id, + version, + &response_id, + None, + items.clone(), + &ResponseMetadata::default(), + ) + .await; + (result, response_id, items) + }) + }; + + let (writer_one_result, writer_one_response_id, writer_one_items) = + writer_one.await.expect("join first checked write"); + let (writer_two_result, writer_two_response_id, writer_two_items) = + writer_two.await.expect("join second checked write"); + + assert_eq!( + usize::from(writer_one_result.is_ok()) + usize::from(writer_two_result.is_ok()), + 1 + ); + assert_eq!( + usize::from(matches!( + &writer_one_result, + Err(StorageError::ConversationConflict { conversation_id }) + if conversation_id == conversation.conversation_id.as_str() + )) + usize::from(matches!( + &writer_two_result, + Err(StorageError::ConversationConflict { conversation_id }) + if conversation_id == conversation.conversation_id.as_str() + )), + 1 + ); + let (winner_items, losing_response_id) = if writer_one_result.is_ok() { + (writer_one_items, writer_two_response_id) + } else { + (writer_two_items, writer_one_response_id) + }; + let rows = + agentic_core::storage::models::item::get_items_by_conversation(&first_pool, &conversation.conversation_id) + .await + .expect("load winning conversation items"); + let sequences = rows + .iter() + .map(|row| row.seq.expect("conversation item sequence")) + .collect::>(); + assert_eq!(sequences, vec![0, 1]); + assert_eq!( + store + .rehydrate(&conversation.conversation_id) + .await + .expect("rehydrate winning conversation items"), + winner_items + ); + let response_error = ResponseStore::new(first_pool.clone()) + .get(&losing_response_id) + .await + .expect_err("the losing response must not be stored"); + assert!(response_error.is_not_found()); + + first_pool.close().await; + second_pool.close().await; +} + #[tokio::test] #[ignore = "requires TEST_POSTGRES_URL pointing to an isolated PostgreSQL database"] async fn postgres_lock_wait_is_bounded_without_blocking_other_conversations() { diff --git a/crates/agentic-server-core/tests/stateful_conversation_integration.rs b/crates/agentic-server-core/tests/stateful_conversation_integration.rs index 5c715633..a6174eb5 100644 --- a/crates/agentic-server-core/tests/stateful_conversation_integration.rs +++ b/crates/agentic-server-core/tests/stateful_conversation_integration.rs @@ -259,7 +259,7 @@ async fn test_multi_branch() { // Turn 2 (main branch) let p2 = unwrap_blocking( execute( - make_request(&t2.request.body.input, true, false, None, Some(conv_id)), + make_request(&t2.request.body.input, true, false, None, Some(conv_id.clone())), Arc::clone(ctx), ) .await @@ -290,6 +290,7 @@ async fn test_multi_branch() { ); assert_eq!(p4.status, "completed"); assert_eq!(output_text(&p4), expected_text(t4)); + assert_eq!(p4.conversation_id.as_deref(), Some(conv_id.as_str())); // Branch 2 — off turn 2 let p5 = unwrap_blocking( diff --git a/crates/agentic-server-core/tests/stateful_responses_integration.rs b/crates/agentic-server-core/tests/stateful_responses_integration.rs index 560716c7..f38966eb 100644 --- a/crates/agentic-server-core/tests/stateful_responses_integration.rs +++ b/crates/agentic-server-core/tests/stateful_responses_integration.rs @@ -624,6 +624,7 @@ async fn test_previous_response_id_persists_inherited_tools_and_choice() { new_input_items: vec![], response_id: "resp_lookup".into(), conversation_id: None, + conversation_version: None, }; let stored = fixture diff --git a/crates/agentic-server-core/tests/storage_integration.rs b/crates/agentic-server-core/tests/storage_integration.rs index 60044d7a..62ee2388 100644 --- a/crates/agentic-server-core/tests/storage_integration.rs +++ b/crates/agentic-server-core/tests/storage_integration.rs @@ -1,11 +1,11 @@ mod support; use agentic_core::config::SqliteConfig; -use agentic_core::storage::InOutItem; use agentic_core::storage::ResponseMetadata; use agentic_core::storage::{ ConversationStore, ResponseStore, create_pool_with_schema, create_pool_with_schema_and_sqlite_config, }; +use agentic_core::storage::{ConversationVersion, InOutItem, StorageError}; use agentic_core::types::event::MessageStatus; use agentic_core::types::io::{InputItem, InputMessage, InputMessageContent, OutputItem, OutputMessage}; use std::sync::Arc; @@ -60,6 +60,297 @@ async fn test_conversation_store_persist_and_rehydrate() { assert_eq!(rehydrated.len(), 2); } +#[tokio::test] +async fn conversation_snapshot_reports_empty_and_last_sequence() -> Result<(), Box> { + let pool = setup_pool().await; + let store = ConversationStore::new(pool); + let conversation = store.create().await?; + + let snapshot = store.rehydrate_snapshot(&conversation.conversation_id).await?; + assert!(snapshot.items.is_empty()); + assert_eq!(snapshot.version, ConversationVersion::Empty); + + store + .persist( + &conversation.conversation_id, + "resp_1", + None, + vec![create_input_item("hello"), create_output_item("msg_1")], + &ResponseMetadata::default(), + ) + .await?; + + let snapshot = store.rehydrate_snapshot(&conversation.conversation_id).await?; + assert_eq!(snapshot.items.len(), 2); + assert_eq!(snapshot.version, ConversationVersion::LastSequence(1)); + assert_eq!(store.rehydrate(&conversation.conversation_id).await?, snapshot.items); + + Ok(()) +} + +#[tokio::test] +async fn conversation_snapshot_version_includes_an_undecodable_final_row() -> Result<(), Box> { + let pool = setup_pool().await; + let store = ConversationStore::new(Arc::clone(&pool)); + let conversation = store.create().await?; + let stored_item = create_input_item("hello"); + + store + .persist( + &conversation.conversation_id, + "resp_1", + None, + vec![stored_item.clone()], + &ResponseMetadata::default(), + ) + .await?; + sqlx::query("INSERT INTO items (id, data, created_at, conversation_id, seq) VALUES ($1, $2, $3, $4, $5)") + .bind("item_undecodable") + .bind("not valid JSON") + .bind(0_i64) + .bind(&conversation.conversation_id) + .bind(1_i64) + .execute(pool.as_ref()) + .await?; + + let snapshot = store.rehydrate_snapshot(&conversation.conversation_id).await?; + + assert_eq!(snapshot.items, vec![stored_item]); + assert_eq!(snapshot.version, ConversationVersion::LastSequence(1)); + + Ok(()) +} + +#[tokio::test] +async fn conversation_snapshot_rejects_items_without_a_sequence() -> Result<(), Box> { + let pool = setup_pool().await; + let store = ConversationStore::new(Arc::clone(&pool)); + let conversation = store.create().await?; + + store + .persist( + &conversation.conversation_id, + "resp_1", + None, + vec![create_input_item("hello")], + &ResponseMetadata::default(), + ) + .await?; + + let item_id: String = sqlx::query_scalar("SELECT id FROM items WHERE conversation_id = $1") + .bind(&conversation.conversation_id) + .fetch_one(pool.as_ref()) + .await?; + sqlx::query("UPDATE items SET seq = NULL WHERE id = $1") + .bind(&item_id) + .execute(pool.as_ref()) + .await?; + + let error = store + .rehydrate_snapshot(&conversation.conversation_id) + .await + .expect_err("snapshot must reject an item without a sequence"); + assert!(matches!( + error, + StorageError::InvalidConversationSequence { + conversation_id, + item_id: invalid_item_id, + } if conversation_id == conversation.conversation_id && invalid_item_id == item_id + )); + + Ok(()) +} + +#[tokio::test] +async fn conversation_version_empty_checked_persist_succeeds() -> Result<(), Box> { + let pool = setup_pool().await; + let store = ConversationStore::new(pool); + let conversation = store.create().await?; + let items = vec![create_input_item("first input"), create_output_item("msg_first")]; + + store + .persist_if_version( + &conversation.conversation_id, + ConversationVersion::Empty, + "resp_first", + None, + items.clone(), + &ResponseMetadata::default(), + ) + .await?; + + let snapshot = store.rehydrate_snapshot(&conversation.conversation_id).await?; + assert_eq!(snapshot.items, items); + assert_eq!(snapshot.version, ConversationVersion::LastSequence(1)); + + Ok(()) +} + +#[tokio::test] +async fn conversation_version_stale_checked_persist_rolls_back_items_and_response() +-> Result<(), Box> { + let pool = setup_pool().await; + let store = ConversationStore::new(Arc::clone(&pool)); + let response_store = ResponseStore::new(Arc::clone(&pool)); + let conversation = store.create().await?; + let snapshot = store.rehydrate_snapshot(&conversation.conversation_id).await?; + let competing_items = vec![ + create_input_item("competing input"), + create_output_item("msg_competing"), + ]; + store + .persist( + &conversation.conversation_id, + "resp_competing", + None, + competing_items.clone(), + &ResponseMetadata::default(), + ) + .await?; + let rejected_items = vec![create_input_item("stale input"), create_output_item("msg_stale")]; + + let error = store + .persist_if_version( + &conversation.conversation_id, + snapshot.version, + "resp_stale", + None, + rejected_items, + &ResponseMetadata::default(), + ) + .await + .expect_err("a stale conversation version must be rejected"); + + assert!(error.is_conversation_conflict()); + assert!(matches!( + error, + StorageError::ConversationConflict { conversation_id } + if conversation_id == conversation.conversation_id + )); + assert_eq!(store.rehydrate(&conversation.conversation_id).await?, competing_items); + let response_error = response_store + .get("resp_stale") + .await + .expect_err("the rejected response must not be stored"); + assert!(response_error.is_not_found()); + + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn conversation_version_racing_checked_persists_allow_exactly_one_winner() +-> Result<(), Box> { + let pool = setup_pool().await; + let store = ConversationStore::new(Arc::clone(&pool)); + let conversation = store.create().await?; + let version = store.rehydrate_snapshot(&conversation.conversation_id).await?.version; + let barrier = Arc::new(tokio::sync::Barrier::new(2)); + + let writer_one = { + let store = ConversationStore::new(Arc::clone(&pool)); + let conversation_id = conversation.conversation_id.clone(); + let barrier = Arc::clone(&barrier); + tokio::spawn(async move { + let items = vec![create_input_item("writer one"), create_output_item("msg_writer_one")]; + barrier.wait().await; + let result = store + .persist_if_version( + &conversation_id, + version, + "resp_writer_one", + None, + items.clone(), + &ResponseMetadata::default(), + ) + .await; + (result, items) + }) + }; + let writer_two = { + let store = ConversationStore::new(pool); + let conversation_id = conversation.conversation_id.clone(); + tokio::spawn(async move { + let items = vec![create_input_item("writer two"), create_output_item("msg_writer_two")]; + barrier.wait().await; + let result = store + .persist_if_version( + &conversation_id, + version, + "resp_writer_two", + None, + items.clone(), + &ResponseMetadata::default(), + ) + .await; + (result, items) + }) + }; + + let (writer_one_result, writer_one_items) = writer_one.await?; + let (writer_two_result, writer_two_items) = writer_two.await?; + + assert_eq!( + usize::from(writer_one_result.is_ok()) + usize::from(writer_two_result.is_ok()), + 1 + ); + assert_eq!( + usize::from( + writer_one_result + .as_ref() + .is_err_and(StorageError::is_conversation_conflict) + ) + usize::from( + writer_two_result + .as_ref() + .is_err_and(StorageError::is_conversation_conflict) + ), + 1 + ); + let winner_items = if writer_one_result.is_ok() { + writer_one_items + } else { + writer_two_items + }; + assert_eq!(store.rehydrate(&conversation.conversation_id).await?, winner_items); + + Ok(()) +} + +#[tokio::test] +async fn conversation_version_is_scoped_per_conversation() -> Result<(), Box> { + let pool = setup_pool().await; + let store = ConversationStore::new(pool); + let first = store.create().await?; + let second = store.create().await?; + let first_items = vec![create_input_item("first conversation")]; + store + .persist( + &first.conversation_id, + "resp_first_conversation", + None, + first_items.clone(), + &ResponseMetadata::default(), + ) + .await?; + let first_snapshot = store.rehydrate_snapshot(&first.conversation_id).await?; + + store + .persist_if_version( + &second.conversation_id, + ConversationVersion::Empty, + "resp_second_conversation", + None, + vec![create_input_item("second conversation")], + &ResponseMetadata::default(), + ) + .await?; + + let first_after = store.rehydrate_snapshot(&first.conversation_id).await?; + assert_eq!(first_after.items, first_items); + assert_eq!(first_after.version, first_snapshot.version); + + Ok(()) +} + #[tokio::test] async fn test_conversation_store_multiple_turns() { let pool = setup_pool().await; diff --git a/crates/agentic-server-core/tests/tool_normalization_test.rs b/crates/agentic-server-core/tests/tool_normalization_test.rs index d108692f..f9de2aea 100644 --- a/crates/agentic-server-core/tests/tool_normalization_test.rs +++ b/crates/agentic-server-core/tests/tool_normalization_test.rs @@ -102,6 +102,7 @@ fn upstream_request_value(payload: RequestPayload, stream: bool) -> Value { new_input_items: Vec::new(), response_id: "resp_test".to_string(), conversation_id: None, + conversation_version: None, }; let upstream_request = ctx .enriched_request diff --git a/crates/agentic-server/Cargo.toml b/crates/agentic-server/Cargo.toml index 2d0f3564..af52b3e5 100644 --- a/crates/agentic-server/Cargo.toml +++ b/crates/agentic-server/Cargo.toml @@ -30,6 +30,7 @@ criterion.workspace = true futures.workspace = true reqwest = { workspace = true, features = ["json"] } serde_json.workspace = true +sqlx = { version = "0.8", features = ["runtime-tokio-rustls", "any", "sqlite"] } tokio = { workspace = true, features = ["test-util"] } tokio-tungstenite.workspace = true uuid = { version = "1", features = ["v7"] } diff --git a/crates/agentic-server/src/app.rs b/crates/agentic-server/src/app.rs index 1c62a10f..48825006 100644 --- a/crates/agentic-server/src/app.rs +++ b/crates/agentic-server/src/app.rs @@ -5,6 +5,8 @@ use axum::Router; use axum::routing::{get, post}; use http::HeaderValue; use tokio::sync::Notify; +#[cfg(debug_assertions)] +use tokio::sync::oneshot; use tokio_util::sync::CancellationToken; use tower_http::cors::{AllowOrigin, Any, CorsLayer}; @@ -24,6 +26,14 @@ pub struct WebSocketTracker { struct WebSocketTrackerInner { active: AtomicUsize, idle: Notify, + #[cfg(debug_assertions)] + local_completion_barrier: std::sync::Mutex>, +} + +#[cfg(debug_assertions)] +struct LocalCompletionBarrier { + rehydrated: oneshot::Sender<()>, + release: oneshot::Receiver<()>, } pub(crate) struct WebSocketGuard { @@ -50,6 +60,40 @@ impl WebSocketTracker { idle.await; } } + + /// Installs a one-shot test barrier after local WebSocket rehydration. + #[cfg(debug_assertions)] + #[doc(hidden)] + #[must_use] + pub fn install_local_completion_test_barrier(&self) -> (oneshot::Receiver<()>, oneshot::Sender<()>) { + let (rehydrated_tx, rehydrated_rx) = oneshot::channel(); + let (release_tx, release_rx) = oneshot::channel(); + let barrier = LocalCompletionBarrier { + rehydrated: rehydrated_tx, + release: release_rx, + }; + self.inner + .local_completion_barrier + .lock() + .expect("local completion test barrier mutex poisoned") + .replace(barrier); + (rehydrated_rx, release_tx) + } + + #[cfg(debug_assertions)] + pub(crate) async fn pause_local_completion_after_rehydration(&self) { + let barrier = self + .inner + .local_completion_barrier + .lock() + .expect("local completion test barrier mutex poisoned") + .take(); + if let Some(barrier) = barrier { + if barrier.rehydrated.send(()).is_ok() { + let _ = barrier.release.await; + } + } + } } impl Drop for WebSocketGuard { diff --git a/crates/agentic-server/src/handler/websocket/error.rs b/crates/agentic-server/src/handler/websocket/error.rs index 43b013d6..3b2936c5 100644 --- a/crates/agentic-server/src/handler/websocket/error.rs +++ b/crates/agentic-server/src/handler/websocket/error.rs @@ -51,6 +51,29 @@ impl WsError { } } + fn error_type(&self) -> &'static str { + match self { + Self::Executor(err) => err.error_type(), + Self::InvalidJson(_) => "invalid_json", + Self::UnexpectedType | Self::BinaryFrame => "invalid_request_error", + Self::SerializeJson(_) | Self::SendFailed | Self::ClientDisconnected | Self::Receive(_) => "server_error", + } + } + + fn param(&self) -> Option<&'static str> { + match self { + Self::Executor(err) => err.error_param(), + _ => None, + } + } + + fn message(&self) -> String { + match self { + Self::Executor(err) => err.error_message(), + _ => self.to_string(), + } + } + pub(super) fn to_ws_frame(&self) -> Option { if matches!( self, @@ -59,15 +82,57 @@ impl WsError { return None; } - let code = self.code(); + let mut error = serde_json::Map::new(); + error.insert("message".to_owned(), Value::String(self.message())); + error.insert("type".to_owned(), Value::String(self.error_type().to_owned())); + error.insert("code".to_owned(), Value::String(self.code().to_owned())); + if let Some(param) = self.param() { + error.insert("param".to_owned(), Value::String(param.to_owned())); + } Some(json!({ "type": "error", "status": self.status().as_u16(), - "error": { - "message": self.to_string(), - "type": code, - "code": code - } + "error": error })) } } + +#[cfg(test)] +mod tests { + use super::*; + use agentic_core::StorageError; + + #[test] + fn executor_conflict_ws_frame_uses_client_conflict_contract() { + let error = WsError::Executor(ExecutorError::Persistence(Box::new( + ExecutorError::ConversationLocked { + source: StorageError::ConversationConflict { + conversation_id: "conv_test".to_owned(), + }, + }, + ))); + + assert_eq!( + error.to_ws_frame().expect("client-visible websocket error"), + json!({ + "type": "error", + "status": 400, + "error": { + "message": "conversation changed while the response was being generated; retry the request", + "type": "invalid_request_error", + "code": "conversation_locked", + "param": "conversation" + } + }) + ); + } + + #[test] + fn non_conflict_ws_frame_omits_param() { + let frame = WsError::UnexpectedType + .to_ws_frame() + .expect("client-visible websocket error"); + + assert!(!frame["error"].as_object().expect("error object").contains_key("param")); + } +} diff --git a/crates/agentic-server/src/handler/websocket/responses.rs b/crates/agentic-server/src/handler/websocket/responses.rs index b5264a1b..3f9f4264 100644 --- a/crates/agentic-server/src/handler/websocket/responses.rs +++ b/crates/agentic-server/src/handler/websocket/responses.rs @@ -13,7 +13,9 @@ use tokio_util::sync::CancellationToken; use tracing::{debug, warn}; use agentic_core::ResponseUsage; -use agentic_core::executor::{BoxStream, ExecuteRequest, ExecutorError, RequestContext, rehydrate_conversation}; +use agentic_core::executor::{ + BoxStream, ExecuteRequest, ExecutorError, RequestContext, persist_turn, rehydrate_conversation, +}; use agentic_core::types::request_response::RequestPayload; use agentic_core::utils::common::utcnow_str; @@ -226,7 +228,15 @@ async fn complete_without_inference( Some(ResponseUsage::default()), ); - state.exec_ctx.resp_handler.execute_turn(ctx, Vec::new()).await?; + #[cfg(debug_assertions)] + state.websocket_tracker.pause_local_completion_after_rehydration().await; + persist_turn( + ctx, + Vec::new(), + &state.exec_ctx.conv_handler, + &state.exec_ctx.resp_handler, + ) + .await?; send_ws_json(sender, created_event).await?; send_ws_json(sender, completed_event).await diff --git a/crates/agentic-server/tests/responses_test.rs b/crates/agentic-server/tests/responses_test.rs index 3307144a..183c4762 100644 --- a/crates/agentic-server/tests/responses_test.rs +++ b/crates/agentic-server/tests/responses_test.rs @@ -2,15 +2,332 @@ mod common; use axum::Router; use axum::body::Bytes; +use axum::http::header; use axum::response::IntoResponse; use axum::routing::post; use http::StatusCode; +use std::convert::Infallible; +use std::future::Future; +use std::path::PathBuf; +use std::pin::Pin; use std::sync::Arc; +use std::task::{Context, Poll}; use tokio::net::TcpListener; -use tokio::sync::Mutex; +use tokio::sync::{Mutex, oneshot}; +use tokio_util::sync::CancellationToken; + +use agentic_core::executor::{ConversationHandler, ExecutionContext, ResponseHandler}; +use agentic_core::proxy::ProxyState; +use agentic_core::storage::{ + ConversationStore, DbPool, InOutItem, ResponseMetadata, ResponseStore, create_pool_with_schema, +}; +use agentic_core::types::io::{InputItem, ResponsesInput}; +use agentic_server::app::{AppState, WebSocketTracker}; use common::{spawn_gateway, spawn_mock_llm, test_config, test_state}; +const COMPETING_RESPONSE_ID: &str = "resp_competing"; +const CONFLICT_MESSAGE: &str = "conversation changed while the response was being generated; retry the request"; + +enum MockResponse { + GatedJson { + body: String, + arrived: oneshot::Sender<()>, + release: oneshot::Receiver<()>, + }, + GatedSse { + first_chunk: String, + terminal_chunk: String, + arrived: oneshot::Sender<()>, + release: oneshot::Receiver<()>, + }, +} + +struct MockResponsesServer { + url: String, + handle: tokio::task::JoinHandle<()>, +} + +struct GatedSse { + first_chunk: Option, + terminal_chunk: Option, + release: oneshot::Receiver<()>, +} + +impl futures::Stream for GatedSse { + type Item = Result; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + if let Some(first_chunk) = self.first_chunk.take() { + return Poll::Ready(Some(Ok(first_chunk))); + } + if self.terminal_chunk.is_none() { + return Poll::Ready(None); + } + match Pin::new(&mut self.release).poll(cx) { + Poll::Ready(_) => Poll::Ready(self.terminal_chunk.take().map(Ok)), + Poll::Pending => Poll::Pending, + } + } +} + +impl MockResponsesServer { + async fn start_gated_json(body: String) -> (Self, oneshot::Receiver<()>, oneshot::Sender<()>) { + let (arrived_tx, arrived_rx) = oneshot::channel(); + let (release_tx, release_rx) = oneshot::channel(); + let server = Self::start(MockResponse::GatedJson { + body, + arrived: arrived_tx, + release: release_rx, + }) + .await; + (server, arrived_rx, release_tx) + } + + async fn start_gated_sse( + first_chunk: String, + terminal_chunk: String, + ) -> (Self, oneshot::Receiver<()>, oneshot::Sender<()>) { + let (arrived_tx, arrived_rx) = oneshot::channel(); + let (release_tx, release_rx) = oneshot::channel(); + let server = Self::start(MockResponse::GatedSse { + first_chunk, + terminal_chunk, + arrived: arrived_tx, + release: release_rx, + }) + .await; + (server, arrived_rx, release_tx) + } + + async fn start(response: MockResponse) -> Self { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let response = Arc::new(Mutex::new(Some(response))); + let route_response = Arc::clone(&response); + let app = Router::new().route( + "/v1/responses", + post(move || { + let response = Arc::clone(&route_response); + async move { + let response = response.lock().await.take().expect("mock response already consumed"); + match response { + MockResponse::GatedJson { body, arrived, release } => { + let _ = arrived.send(()); + let _ = release.await; + axum::response::Response::builder() + .status(StatusCode::OK) + .header(header::CONTENT_TYPE, "application/json") + .body(axum::body::Body::from(body)) + .unwrap() + .into_response() + } + MockResponse::GatedSse { + first_chunk, + terminal_chunk, + arrived, + release, + } => { + let _ = arrived.send(()); + axum::response::Response::builder() + .status(StatusCode::OK) + .header(header::CONTENT_TYPE, "text/event-stream; charset=utf-8") + .body(axum::body::Body::from_stream(GatedSse { + first_chunk: Some(Bytes::from(first_chunk)), + terminal_chunk: Some(Bytes::from(terminal_chunk)), + release, + })) + .unwrap() + .into_response() + } + } + } + }), + ); + let handle = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + Self { + url: format!("http://{addr}"), + handle, + } + } +} + +impl Drop for MockResponsesServer { + fn drop(&mut self) { + self.handle.abort(); + } +} + +struct TestDb { + path: PathBuf, +} + +impl TestDb { + fn new() -> Self { + Self { + path: std::env::temp_dir().join(format!("agentic_http_test_{}.db", uuid::Uuid::now_v7())), + } + } + + fn url(&self) -> String { + format!("sqlite://{}", self.path.display()) + } +} + +impl Drop for TestDb { + fn drop(&mut self) { + let _ = std::fs::remove_file(&self.path); + let _ = std::fs::remove_file(self.path.with_extension("db-shm")); + let _ = std::fs::remove_file(self.path.with_extension("db-wal")); + } +} + +struct StorageBackedState { + state: AppState, + pool: Arc, + _db: TestDb, +} + +async fn storage_backed_state(llm_url: &str) -> StorageBackedState { + let db = TestDb::new(); + let pool = create_pool_with_schema(Some(&db.url())).await.unwrap(); + let config = test_config(llm_url); + let client = Arc::new(reqwest::Client::new()); + let exec_ctx = Arc::new(ExecutionContext::new( + ConversationHandler::new(ConversationStore::new(Arc::clone(&pool))), + ResponseHandler::new(ResponseStore::new(Arc::clone(&pool))), + client, + config.llm_api_base.clone(), + )); + let proxy_state = ProxyState::new(config.clone()).expect("proxy state"); + let state = AppState { + proxy_state, + exec_ctx, + shutdown_token: CancellationToken::new(), + websocket_tracker: WebSocketTracker::default(), + llm_api_base: config.llm_api_base, + openai_api_key: config.openai_api_key, + }; + StorageBackedState { state, pool, _db: db } +} + +async fn create_conversation(client: &reqwest::Client, gateway_url: &str) -> String { + let response = client + .post(format!("{gateway_url}/v1/conversations")) + .json(&serde_json::json!({"store": true})) + .send() + .await + .expect("conversation request"); + assert_eq!(response.status(), StatusCode::OK); + response + .json::() + .await + .expect("conversation response JSON")["id"] + .as_str() + .expect("conversation ID") + .to_owned() +} + +fn competing_turn_items() -> Vec { + Vec::::from(&ResponsesInput::Text("competing turn".to_owned())) + .into_iter() + .map(InOutItem::Input) + .collect() +} + +async fn persist_competing_turn(pool: &Arc, conversation_id: &str) { + ConversationStore::new(Arc::clone(pool)) + .persist( + conversation_id, + COMPETING_RESPONSE_ID, + None, + competing_turn_items(), + &ResponseMetadata { + model: "competing-model".to_owned(), + ..ResponseMetadata::default() + }, + ) + .await + .expect("competing turn should persist"); +} + +fn conflict_error() -> serde_json::Value { + serde_json::json!({ + "message": CONFLICT_MESSAGE, + "type": "invalid_request_error", + "code": "conversation_locked", + "param": "conversation" + }) +} + +fn sse_events(body: &str) -> Vec { + body.split("\n\n") + .filter_map(|frame| frame.strip_prefix("data: ")) + .filter(|data| *data != "[DONE]") + .map(|data| serde_json::from_str(data).expect("SSE data should be JSON")) + .collect() +} + +fn gated_sse_chunks() -> (String, String) { + let created = serde_json::json!({ + "type": "response.created", + "sequence_number": 0, + "response": {"id": "resp_upstream_stale_sse", "status": "in_progress"} + }); + let added = serde_json::json!({ + "type": "response.output_item.added", + "sequence_number": 1, + "output_index": 0, + "item": {"id": "msg_upstream_stale_sse", "type": "message"} + }); + let delta = serde_json::json!({ + "type": "response.output_text.delta", + "sequence_number": 2, + "item_id": "msg_upstream_stale_sse", + "output_index": 0, + "content_index": 0, + "delta": "partial" + }); + let completed = serde_json::json!({ + "type": "response.completed", + "sequence_number": 3, + "response": {"id": "resp_upstream_stale_sse", "status": "completed", "usage": null} + }); + ( + format!("data: {created}\n\ndata: {added}\n\ndata: {delta}\n\n"), + format!("data: {completed}\n\ndata: [DONE]\n\n"), + ) +} + +async fn assert_only_competing_turn_persisted(pool: &Arc, conversation_id: &str) { + let conversation_store = ConversationStore::new(Arc::clone(pool)); + assert_eq!( + conversation_store + .rehydrate(conversation_id) + .await + .expect("conversation history"), + competing_turn_items() + ); + let response = ResponseStore::new(Arc::clone(pool)) + .get(COMPETING_RESPONSE_ID) + .await + .expect("competing response should remain"); + assert_eq!(response.conversation_id.as_deref(), Some(conversation_id)); + let response_count = sqlx::query_scalar::<_, i64>("SELECT COUNT(*) FROM responses") + .fetch_one(pool.as_ref()) + .await + .expect("response count"); + assert_eq!(response_count, 1, "only the competing response should be stored"); +} + +async fn assert_response_not_persisted(pool: &Arc, response_id: &str) { + let error = ResponseStore::new(Arc::clone(pool)) + .get(response_id) + .await + .expect_err("rejected response must not be persisted"); + assert!(error.is_not_found(), "expected missing response, got {error}"); +} + /// Spawn a mock vLLM that returns a minimal valid JSON response. async fn spawn_mock_vllm_json() -> (String, tokio::task::JoinHandle<()>) { let app = Router::new().route( @@ -264,6 +581,131 @@ async fn test_streaming_store_true_hides_persistence_details_without_sequence_ga assert!(!body.contains("\"type\":\"response.completed\""), "{body}"); } +#[tokio::test] +async fn http_json_conversation_conflict_rejects_stale_turn_without_persisting_it() { + // Arrange + let upstream_body = serde_json::json!({ + "id": "resp_upstream_stale_json", + "object": "response", + "status": "completed", + "model": "test-model", + "output": [], + "created_at": 0 + }) + .to_string(); + let (mock, arrived, release) = MockResponsesServer::start_gated_json(upstream_body).await; + let fixture = storage_backed_state(&mock.url).await; + let (gateway_url, _gateway) = spawn_gateway(fixture.state.clone()).await; + let client = reqwest::Client::new(); + let conversation_id = create_conversation(&client, &gateway_url).await; + + // Act + let response_task = { + let client = client.clone(); + let gateway_url = gateway_url.clone(); + let conversation_id = conversation_id.clone(); + tokio::spawn(async move { + client + .post(format!("{gateway_url}/v1/responses")) + .json(&serde_json::json!({ + "model": "test-model", + "input": [{"type": "message", "role": "user", "content": "stale turn"}], + "conversation_id": conversation_id, + "store": true, + "stream": false + })) + .send() + .await + .expect("response request") + }) + }; + arrived.await.expect("upstream request should arrive after rehydration"); + persist_competing_turn(&fixture.pool, &conversation_id).await; + release.send(()).expect("release gated JSON response"); + let response = response_task.await.expect("response task"); + + // Assert + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + assert_eq!( + response.json::().await.expect("error response JSON"), + serde_json::json!({"error": conflict_error()}) + ); + assert_only_competing_turn_persisted(&fixture.pool, &conversation_id).await; +} + +#[tokio::test] +async fn http_sse_conversation_conflict_terminates_after_observable_delta_without_persisting_stale_turn() { + // Arrange + let (first_chunk, terminal_chunk) = gated_sse_chunks(); + let (mock, arrived, release) = MockResponsesServer::start_gated_sse(first_chunk, terminal_chunk).await; + let fixture = storage_backed_state(&mock.url).await; + let (gateway_url, _gateway) = spawn_gateway(fixture.state.clone()).await; + let client = reqwest::Client::new(); + let conversation_id = create_conversation(&client, &gateway_url).await; + + // Act + let mut response = client + .post(format!("{gateway_url}/v1/responses")) + .json(&serde_json::json!({ + "model": "test-model", + "input": [{"type": "message", "role": "user", "content": "stale turn"}], + "conversation_id": conversation_id, + "store": true, + "stream": true + })) + .send() + .await + .expect("streaming response request"); + assert_eq!(response.status(), StatusCode::OK); + arrived.await.expect("upstream request should arrive after rehydration"); + + let mut body = String::new(); + while !body.contains("\"type\":\"response.output_text.delta\"") { + let chunk = response + .chunk() + .await + .expect("stream chunk") + .expect("stream should contain a delta before completion"); + body.push_str(std::str::from_utf8(&chunk).expect("SSE should be UTF-8")); + } + assert!(body.contains("\"delta\":\"partial\""), "{body}"); + + persist_competing_turn(&fixture.pool, &conversation_id).await; + release.send(()).expect("release gated SSE response"); + while let Some(chunk) = response.chunk().await.expect("stream chunk") { + body.push_str(std::str::from_utf8(&chunk).expect("SSE should be UTF-8")); + } + + // Assert + let events = sse_events(&body); + let stale_response_id = events + .iter() + .find(|event| event["type"] == "response.created") + .and_then(|event| event["response"]["id"].as_str()) + .expect("gateway response ID from response.created"); + assert!( + events + .iter() + .any(|event| event["type"] == "response.output_text.delta" && event["delta"] == "partial"), + "{body}" + ); + let errors = events + .iter() + .filter(|event| event["type"] == "error") + .collect::>(); + assert_eq!(errors.len(), 1, "{body}"); + assert_eq!(errors[0]["status"], StatusCode::BAD_REQUEST.as_u16()); + assert_eq!(errors[0]["error"], conflict_error()); + assert_eq!(events.last().expect("terminal SSE event")["type"], "error"); + assert!(body.contains("data: [DONE]"), "{body}"); + assert!( + events.iter().all(|event| event["type"] != "response.completed"), + "{body}" + ); + assert_only_competing_turn_persisted(&fixture.pool, &conversation_id).await; + assert_response_not_persisted(&fixture.pool, stale_response_id).await; +} + #[tokio::test] async fn test_oversized_body_returns_413() { // Arrange — LLM is never reached (gateway rejects the body first) diff --git a/crates/agentic-server/tests/responses_websocket_test.rs b/crates/agentic-server/tests/responses_websocket_test.rs index 02c8e57e..3c0268be 100644 --- a/crates/agentic-server/tests/responses_websocket_test.rs +++ b/crates/agentic-server/tests/responses_websocket_test.rs @@ -24,9 +24,12 @@ use tokio_util::sync::CancellationToken; use agentic_core::executor::{ConversationHandler, ExecutionContext, RequestContext, ResponseHandler}; use agentic_core::proxy::ProxyState; -use agentic_core::storage::{ConversationStore, ResponseStore, create_pool_with_schema}; +use agentic_core::storage::{ + ConversationStore, DbPool, InOutItem, ResponseMetadata, ResponseStore, create_pool_with_schema, +}; use agentic_core::tool::{WebSearchHandler, model_visible_namespace_member_name}; use agentic_core::types::RequestPayload; +use agentic_core::types::io::{InputItem, ResponsesInput}; use agentic_core::types::tools::ResponsesTool; use agentic_server::app::{AppState, WebSocketTracker}; @@ -97,6 +100,7 @@ enum MockResponse { Static(String), Gated { response: String, + arrived: oneshot::Sender<()>, release: oneshot::Receiver<()>, }, Hanging { @@ -150,14 +154,16 @@ impl MockResponsesServer { (server, drop_rx) } - async fn start_gated(response: String) -> (Self, oneshot::Sender<()>) { + async fn start_gated(response: String) -> (Self, oneshot::Receiver<()>, oneshot::Sender<()>) { + let (arrived_tx, arrived_rx) = oneshot::channel(); let (release_tx, release_rx) = oneshot::channel(); let server = Self::start_with_responses(vec![MockResponse::Gated { response, + arrived: arrived_tx, release: release_rx, }]) .await; - (server, release_tx) + (server, arrived_rx, release_tx) } async fn start_with_responses(responses: Vec) -> Self { @@ -179,7 +185,12 @@ impl MockResponsesServer { let response = queue.lock().await.pop_front().expect("mock response queue exhausted"); let body = match response { MockResponse::Static(response) => axum::body::Body::from(response), - MockResponse::Gated { response, release } => { + MockResponse::Gated { + response, + arrived, + release, + } => { + let _ = arrived.send(()); let _ = release.await; axum::body::Body::from(response) } @@ -241,6 +252,7 @@ impl Drop for TestDb { struct StorageBackedState { state: AppState, + pool: Arc, _db: TestDb, } @@ -276,7 +288,7 @@ async fn storage_backed_state_with_web_search(llm_url: &str, web_search_base_url let client = Arc::new(reqwest::Client::new()); let mut exec_ctx = ExecutionContext::new( ConversationHandler::new(ConversationStore::new(Arc::clone(&pool))), - ResponseHandler::new(ResponseStore::new(pool)), + ResponseHandler::new(ResponseStore::new(Arc::clone(&pool))), Arc::clone(&client), config.llm_api_base.clone(), ); @@ -298,7 +310,78 @@ async fn storage_backed_state_with_web_search(llm_url: &str, web_search_base_url llm_api_base: config.llm_api_base, openai_api_key: config.openai_api_key, }; - StorageBackedState { state, _db: db } + StorageBackedState { state, pool, _db: db } +} + +const COMPETING_RESPONSE_ID: &str = "resp_competing"; +const CONFLICT_MESSAGE: &str = "conversation changed while the response was being generated; retry the request"; + +async fn create_conversation(gateway_url: &str) -> String { + let response = reqwest::Client::new() + .post(format!("{gateway_url}/v1/conversations")) + .json(&json!({"store": true})) + .send() + .await + .expect("conversation request"); + assert_eq!(response.status(), StatusCode::OK); + response.json::().await.expect("conversation response JSON")["id"] + .as_str() + .expect("conversation ID") + .to_owned() +} + +fn competing_turn_items() -> Vec { + Vec::::from(&ResponsesInput::Text("competing turn".to_owned())) + .into_iter() + .map(InOutItem::Input) + .collect() +} + +async fn persist_competing_turn(pool: &Arc, conversation_id: &str) { + ConversationStore::new(Arc::clone(pool)) + .persist( + conversation_id, + COMPETING_RESPONSE_ID, + None, + competing_turn_items(), + &ResponseMetadata { + model: "competing-model".to_owned(), + ..ResponseMetadata::default() + }, + ) + .await + .expect("competing turn should persist"); +} + +async fn assert_conflicting_websocket_turn_not_persisted( + pool: &Arc, + conversation_id: &str, + stale_response_id: &str, +) { + let conversation_store = ConversationStore::new(Arc::clone(pool)); + assert_eq!( + conversation_store + .rehydrate(conversation_id) + .await + .expect("conversation history"), + competing_turn_items() + ); + let response_store = ResponseStore::new(Arc::clone(pool)); + let competing = response_store + .get(COMPETING_RESPONSE_ID) + .await + .expect("competing response should remain"); + assert_eq!(competing.conversation_id.as_deref(), Some(conversation_id)); + let error = response_store + .get(stale_response_id) + .await + .expect_err("rejected response must not be persisted"); + assert!(error.is_not_found(), "expected missing response, got {error}"); + let item_count = sqlx::query_scalar::<_, i64>("SELECT COUNT(*) FROM items") + .fetch_one(pool.as_ref()) + .await + .expect("stored item count"); + assert_eq!(item_count, 1, "rejected turn must not leave orphaned items"); } fn ws_url(gateway_url: &str) -> String { @@ -705,6 +788,7 @@ async fn test_websocket_generate_false_prewarm_redacts_mcp_runtime_credentials() new_input_items: vec![], response_id: "resp_lookup".to_owned(), conversation_id: None, + conversation_version: None, }; let stored = fixture .state @@ -768,6 +852,179 @@ async fn test_websocket_first_turn_forwards_incremental_events_and_final_payload assert!(requests[0].get("type").is_none()); } +#[tokio::test] +async fn websocket_conversation_conflict_ends_request_without_persisting_stale_turn() { + // Arrange + let (mock, arrived, release) = + MockResponsesServer::start_gated(sse_response("resp_upstream_stale_ws", "msg_upstream_stale_ws", "STALE")) + .await; + let fixture = storage_backed_state(&mock.url).await; + let (gateway_url, _gateway) = spawn_gateway(fixture.state.clone()).await; + let conversation_id = create_conversation(&gateway_url).await; + let mut ws = connect_responses_ws(&gateway_url).await; + + // Act + send_json( + &mut ws, + json!({ + "type": "response.create", + "model": "test-model", + "input": [{"type": "message", "role": "user", "content": "stale turn"}], + "conversation_id": conversation_id, + "store": true, + "stream": true + }), + ) + .await; + arrived.await.expect("upstream request should arrive after rehydration"); + persist_competing_turn(&fixture.pool, &conversation_id).await; + release.send(()).expect("release gated WebSocket response"); + let events = recv_until_completed(&mut ws).await; + + // Assert + let stale_response_id = events + .iter() + .find(|event| event["type"] == "response.created") + .and_then(|event| event["response"]["id"].as_str()) + .expect("gateway response ID from response.created") + .to_owned(); + let error = events.last().expect("terminal conflict error"); + assert_eq!(error["type"], "error"); + assert_eq!(error["status"], StatusCode::BAD_REQUEST.as_u16()); + assert_eq!( + error["error"], + json!({ + "message": CONFLICT_MESSAGE, + "type": "invalid_request_error", + "code": "conversation_locked", + "param": "conversation" + }) + ); + assert!(events.iter().all(|event| event["type"] != "response.completed")); + + // A local response queued after the error is an ordering barrier: it can only + // begin after the conflicted executor stream has ended. + send_json( + &mut ws, + json!({ + "type": "response.create", + "model": "test-model", + "input": [], + "generate": false, + "store": false, + "stream": true + }), + ) + .await; + let barrier_created = recv_json(&mut ws).await; + assert_eq!(barrier_created["type"], "response.created"); + let barrier_response_id = barrier_created["response"]["id"] + .as_str() + .expect("barrier response ID") + .to_owned(); + let barrier_completed = recv_json(&mut ws).await; + assert_eq!(barrier_completed["type"], "response.completed"); + assert_eq!(barrier_completed["response"]["id"], barrier_response_id); + assert_ne!(barrier_response_id, stale_response_id); + + assert_conflicting_websocket_turn_not_persisted(&fixture.pool, &conversation_id, &stale_response_id).await; +} + +#[tokio::test] +async fn websocket_generate_false_conversation_conflict_rejects_stale_local_completion() { + // Arrange + let mock = MockResponsesServer::start(vec![]).await; + let fixture = storage_backed_state(&mock.url).await; + let (rehydrated, release) = fixture.state.websocket_tracker.install_local_completion_test_barrier(); + let (gateway_url, _gateway) = spawn_gateway(fixture.state.clone()).await; + let conversation_id = create_conversation(&gateway_url).await; + let mut ws = connect_responses_ws(&gateway_url).await; + + // Act: the one-shot barrier deterministically pauses local completion after + // rehydration and before persistence, without relying on timing sleeps. + send_json( + &mut ws, + json!({ + "type": "response.create", + "model": "test-model", + "input": [{"type": "message", "role": "user", "content": "stale local turn"}], + "conversation_id": conversation_id, + "generate": false, + "store": true, + "stream": true + }), + ) + .await; + rehydrated + .await + .expect("local completion should pause after rehydration"); + persist_competing_turn(&fixture.pool, &conversation_id).await; + release + .send(()) + .expect("local completion should remain paused before persistence"); + let error = recv_json(&mut ws).await; + + // Assert + assert_eq!( + error, + json!({ + "type": "error", + "status": StatusCode::BAD_REQUEST.as_u16(), + "error": { + "message": CONFLICT_MESSAGE, + "type": "invalid_request_error", + "code": "conversation_locked", + "param": "conversation" + } + }) + ); + + // A local response queued after the error is an ordering barrier. Its first + // event proves the stale turn emitted no response.completed after the error. + send_json( + &mut ws, + json!({ + "type": "response.create", + "model": "test-model", + "input": [], + "generate": false, + "store": false, + "stream": true + }), + ) + .await; + let barrier_created = recv_json(&mut ws).await; + assert_eq!(barrier_created["type"], "response.created"); + let barrier_response_id = barrier_created["response"]["id"] + .as_str() + .expect("barrier response ID") + .to_owned(); + let barrier_completed = recv_json(&mut ws).await; + assert_eq!(barrier_completed["type"], "response.completed"); + assert_eq!(barrier_completed["response"]["id"], barrier_response_id); + + assert!(mock.request_bodies().await.is_empty()); + let conversation_store = ConversationStore::new(Arc::clone(&fixture.pool)); + assert_eq!( + conversation_store + .rehydrate(&conversation_id) + .await + .expect("conversation history"), + competing_turn_items() + ); + let response_ids = sqlx::query_scalar::<_, String>("SELECT id FROM responses WHERE id != $1 ORDER BY id") + .bind(&barrier_response_id) + .fetch_all(fixture.pool.as_ref()) + .await + .expect("stored response IDs"); + assert_eq!(response_ids, vec![COMPETING_RESPONSE_ID]); + let item_count = sqlx::query_scalar::<_, i64>("SELECT COUNT(*) FROM items") + .fetch_one(fixture.pool.as_ref()) + .await + .expect("stored item count"); + assert_eq!(item_count, 1, "rejected turn must not leave orphaned items"); +} + #[tokio::test] async fn test_websocket_streaming_persistence_error_uses_standard_envelope() { let mock = MockResponsesServer::start(vec![sse_response("resp_upstream_1", "msg_upstream_1", "HELLO")]).await; @@ -1399,7 +1656,7 @@ async fn test_websocket_shutdown_token_closes_idle_connection() { #[tokio::test] async fn test_websocket_shutdown_drains_active_response_before_closing() { - let (mock, release) = + let (mock, arrived, release) = MockResponsesServer::start_gated(sse_response("resp_upstream_shutdown", "msg_upstream_shutdown", "DONE")).await; let fixture = storage_backed_state(&mock.url).await; let shutdown_token = fixture.state.shutdown_token.clone(); @@ -1418,7 +1675,7 @@ async fn test_websocket_shutdown_drains_active_response_before_closing() { }), ) .await; - wait_for_request_count(&mock, 1).await; + arrived.await.expect("upstream request should arrive"); shutdown_token.cancel(); send_json( diff --git a/docs/deploying/container.md b/docs/deploying/container.md index 73866d72..061bf586 100644 --- a/docs/deploying/container.md +++ b/docs/deploying/container.md @@ -93,7 +93,7 @@ For an existing large database, schedule the first upgraded replica during a mai Drain replicas running an older release before enabling writes through this release. Older replicas do not take the per-conversation row lock and can allocate duplicate sequence numbers if they write alongside upgraded replicas. -Stored requests now fail if their response or conversation state cannot be persisted. For streaming requests, the gateway sends an error event instead of `response.completed`. Client responses use the generic message `failed to persist response`; the underlying database error is written only to gateway logs. This prevents clients from receiving a response ID that cannot be continued after a lock timeout or other database failure without exposing database schema or constraint details. +Stored requests now fail if their response or conversation state cannot be persisted. For streaming requests, the gateway sends an error event instead of `response.completed`. Most client responses use the generic message `failed to persist response`; the underlying database error is written only to gateway logs. The exception is an optimistic conversation conflict, which returns status `400`, type `invalid_request_error`, code `conversation_locked`, and param `conversation`. No part of the stale turn is persisted, so the client can retry the request against the conversation's latest state. This prevents clients from receiving a response ID that cannot be continued after a lock timeout or other database failure without exposing database schema or constraint details. `AGENTIC_API_SCHEMA_READY` keeps schema changes under supervisor control. Startup performs a read-only compatibility check and fails if required persistence columns, types, nullability, primary/foreign-key constraints, or the conversation