diff --git a/crates/libsy-llm-client/src/run.rs b/crates/libsy-llm-client/src/run.rs index 01e1c36fc..4128dd03c 100644 --- a/crates/libsy-llm-client/src/run.rs +++ b/crates/libsy-llm-client/src/run.rs @@ -518,7 +518,11 @@ impl StateOwners { } } if let Some(id) = conversation_id { - self.by_id.entry(id.to_owned()).or_insert(state); + if materialized { + self.by_id.insert(id.to_owned(), state); + } else { + self.by_id.entry(id.to_owned()).or_insert(state); + } } Ok(()) } @@ -673,12 +677,11 @@ impl ClientRouter { } fn canonical_input(&self, request: &Request) -> CanonicalInput { - let parent = request - .llm_request - .extensions - .fields + let fields = &request.llm_request.extensions.fields; + let parent = fields .get("previous_response_id") .and_then(Value::as_str) + .or_else(|| conversation_id(fields)) .and_then(|id| self.inner.state_owners.lock().owner(id)?.history.clone()); let parent_len = parent.as_ref().map_or(0, |history| history.len); let (parent, messages) = match request.llm_request.messages.get(parent_len..) { @@ -719,8 +722,14 @@ impl ClientRouter { .map_err(|error| LibsyError::client_call(model.clone(), error))?; } } else if let Some(input) = &canonical_input { - self.remember_canonical_response(&agg, &model, store, input) - .map_err(|error| LibsyError::client_call(model.clone(), error))?; + self.remember_canonical_response( + &agg, + &model, + store, + conversation.as_deref(), + input, + ) + .map_err(|error| LibsyError::client_call(model.clone(), error))?; } LlmResponse::Agg(agg) } @@ -756,6 +765,7 @@ impl ClientRouter { &accumulator.finish(), &model, store, + conversation.as_deref(), input, )?; } @@ -841,10 +851,12 @@ impl ClientRouter { response: &AggLlmResponse, model: &ModelId, store: bool, + conversation: Option<&str>, input: &CanonicalInput, ) -> std::result::Result<(), LlmClientError> { let response_id = response.id.as_deref().filter(|_| store); - if response_id.is_none() { + // `store: false` disables response-ID lookup, not conversation history. + if response_id.is_none() && conversation.is_none() { return Ok(()); } let mut segment = input.messages.to_vec(); @@ -856,11 +868,12 @@ impl ClientRouter { input.parent.clone(), Arc::from(segment), )); - let result = - self.inner - .state_owners - .lock() - .remember(response_id, None, model, Some(history)); + let result = self.inner.state_owners.lock().remember( + response_id, + conversation, + model, + Some(history), + ); if let Err(error) = result { if matches!(error, LlmClientError::ResponseStateLimitExceeded { .. }) { tracing::warn!(%error, "cross-format Responses state capacity reached; history was not retained"); @@ -1285,6 +1298,22 @@ mod tests { #[tokio::test] async fn responses_stored_tool_continuation_materializes_for_anthropic() -> Result<()> { + for (field, conversation, store) in [ + ("previous_response_id", Value::Null, true), + ("conversation", json!("conv_tools"), true), + ("conversation", json!({"id": "conv_tools"}), true), + ("conversation", json!("conv_tools"), false), + ] { + check_stored_tool_continuation(field, conversation, store).await?; + } + Ok(()) + } + + async fn check_stored_tool_continuation( + field: &str, + conversation: Value, + store: bool, + ) -> Result<()> { let server = MockServer::start().await; Mock::given(method("POST")) .respond_with(|request: &wiremock::Request| { @@ -1369,11 +1398,12 @@ mod tests { &json!({ "model": "route", "input": "Call get_weather for Paris", + "conversation": conversation, "tools": [{ "type": "function", "name": "get_weather", "parameters": {"type": "object"} }], - "store": true + "store": store }), ) .map_err(|error| LibsyError::external("decoding seed request", error))?, @@ -1401,7 +1431,7 @@ mod tests { WireFormat::OpenAiResponses, &json!({ "model": "route", - "previous_response_id": "msg_seed", + (field): if field == "conversation" { conversation.clone() } else { json!("msg_seed") }, "input": [{ "type": "function_call_output", "call_id": "toolu_weather", @@ -1411,7 +1441,7 @@ mod tests { "type": "function", "name": "get_weather", "parameters": {"type": "object"} }], - "store": true + "store": store }), ) .map_err(|error| LibsyError::external("decoding continuation request", error))?, @@ -1441,9 +1471,9 @@ mod tests { WireFormat::OpenAiResponses, &json!({ "model": "route", - "previous_response_id": "msg_follow", + (field): if field == "conversation" { conversation.clone() } else { json!("msg_follow") }, "input": "thanks", - "store": true + "store": store }), ) .map_err(|error| LibsyError::external("decoding chained request", error))?, @@ -1457,6 +1487,24 @@ mod tests { assert_eq!(history.len, 4); assert_eq!(history.segment.len(), 2); assert!(history.parent.is_some()); + assert_eq!( + clients + .inner + .state_owners + .lock() + .owner("msg_seed") + .is_some(), + store + ); + assert_eq!( + clients + .inner + .state_owners + .lock() + .owner("msg_follow") + .is_some(), + store + ); let (_, response) = run( Arc::new(switchyard_libsy::Passthrough), clients, @@ -1483,6 +1531,10 @@ mod tests { assert_eq!(follow["messages"][2]["content"][0]["type"], "tool_result"); let third: Value = serde_json::from_slice(&requests[2].body) .map_err(|error| LibsyError::external("decoding captured request", error))?; + assert_eq!( + third["messages"][0]["content"], + "Call get_weather for Paris" + ); assert_eq!(third["messages"][1]["content"][0]["type"], "tool_use"); assert_eq!(third["messages"][2]["content"][0]["type"], "tool_result"); assert_eq!(third["messages"][3]["content"], "sunny"); @@ -1490,8 +1542,20 @@ mod tests { Ok(()) } + // Stream state is recorded under the response and conversation ids only after completion. #[tokio::test] async fn streamed_cross_format_state_is_recorded_only_after_completion() -> Result<()> { + for (store, response_id) in [ + (true, Some("msg_stream")), + (false, Some("msg_stream")), + (true, None), + ] { + check_streamed_state(store, response_id).await?; + } + Ok(()) + } + + async fn check_streamed_state(store: bool, response_id: Option<&str>) -> Result<()> { let client = Arc::new(CandidateClient { calls: Mutex::new(Vec::new()), requests: Mutex::new(Vec::new()), @@ -1507,9 +1571,17 @@ mod tests { .preservation .requests .insert(WireFormat::OpenAiResponses.into(), json!({})); + seed.llm_request + .extensions + .fields + .insert("store".to_string(), json!(store)); + seed.llm_request + .extensions + .fields + .insert("conversation".to_string(), json!("conv_stream")); let mut response = stream_response(vec![ LlmResponseChunk::MessageStart { - id: Some("msg_stream".to_string()), + id: response_id.map(str::to_owned), model: Some("weak".to_string()), }, LlmResponseChunk::ToolCallDelta { @@ -1525,45 +1597,62 @@ mod tests { response.set_served_model(&ModelId::from("weak")); let response = clients.remember_state_owner(&seed, response)?; - let mut follow = request(); - follow - .llm_request - .extensions - .fields - .insert("previous_response_id".to_string(), json!("msg_stream")); - follow.llm_request.messages = vec![Message { - role: Role::Tool, - content: vec![ContentBlock::ToolResult(ToolResult { - tool_call_id: "toolu_stream".to_string(), - content: vec![ContentBlock::Text { - text: "ready".to_string(), - }], - is_error: None, - })], - }]; - assert!(clients.stored_state_owner(&follow).is_none()); + let follow = |field: &str, value: Value| { + let mut follow = request(); + follow + .llm_request + .extensions + .fields + .insert(field.to_string(), value); + follow.llm_request.messages = vec![Message { + role: Role::Tool, + content: vec![ContentBlock::ToolResult(ToolResult { + tool_call_id: "toolu_stream".to_string(), + content: vec![ContentBlock::Text { + text: "ready".to_string(), + }], + is_error: None, + })], + }]; + follow + }; + let follows = [ + ( + follow("previous_response_id", json!("msg_stream")), + store && response_id.is_some(), + ), + (follow("conversation", json!("conv_stream")), true), + ]; + for (follow, _) in &follows { + assert!(clients.stored_state_owner(follow).is_none()); + } response .llm_response .into_agg() .await .map_err(|error| LibsyError::client_call("weak", error))?; - let outcome = continue_on( - clients - .stored_state_owner(&follow) - .expect("completed stream state"), - "passthrough", - follow, - ); - assert_eq!(outcome.request.llm_request.messages.len(), 3); - assert!(matches!( - outcome.request.llm_request.messages[1].content.as_slice(), - [ContentBlock::ToolCall(call)] if call.id == "toolu_stream" - )); - assert!(matches!( - outcome.request.llm_request.messages[2].content.as_slice(), - [ContentBlock::ToolResult(result)] if result.tool_call_id == "toolu_stream" - )); + for (follow, retained) in follows { + let Some(owner) = clients.stored_state_owner(&follow) else { + assert!(!retained); + continue; + }; + assert!(retained); + let outcome = continue_on(owner, "passthrough", follow); + assert_eq!(outcome.request.llm_request.messages.len(), 3); + assert!(matches!( + outcome.request.llm_request.messages[0].content.as_slice(), + [ContentBlock::Text { text }] if text == "inspect" + )); + assert!(matches!( + outcome.request.llm_request.messages[1].content.as_slice(), + [ContentBlock::ToolCall(call)] if call.id == "toolu_stream" + )); + assert!(matches!( + outcome.request.llm_request.messages[2].content.as_slice(), + [ContentBlock::ToolResult(result)] if result.tool_call_id == "toolu_stream" + )); + } Ok(()) } @@ -1709,6 +1798,48 @@ mod tests { Ok(()) } + // Materialized conversation records keep the latest history. + #[test] + fn materialized_conversation_state_keeps_latest_history() + -> std::result::Result<(), LlmClientError> { + let mut owners = StateOwners::default(); + let model = ModelId::from("model/a"); + let first = Arc::new(MessageHistory::new( + None, + Arc::from(vec![Message::text(Role::User, "seed question")]), + )); + let second = Arc::new(MessageHistory::new( + None, + Arc::from(vec![Message::text(Role::User, "recall question")]), + )); + owners.remember(Some("resp_1"), Some("conv_1"), &model, Some(first))?; + owners.remember(Some("resp_2"), Some("conv_1"), &model, Some(second))?; + let latest = owners + .owner("conv_1") + .and_then(|state| state.history.clone()) + .expect("materialized conversation state"); + assert!(matches!( + latest.segment.as_ref(), + [Message { content, .. }] if matches!( + content.as_slice(), + [ContentBlock::Text { text }] if text == "recall question" + ) + )); + owners.remember(Some("resp_3"), Some("conv_1"), &model, None)?; + let kept = owners + .owner("conv_1") + .and_then(|state| state.history.clone()) + .expect("conversation state survived a provider-owned record"); + assert!(matches!( + kept.segment.as_ref(), + [Message { content, .. }] if matches!( + content.as_slice(), + [ContentBlock::Text { text }] if text == "recall question" + ) + )); + Ok(()) + } + #[test] fn concurrent_state_owners_do_not_exceed_capacity() { let model = ModelId::from("model/a"); diff --git a/crates/switchyard-translation/src/codecs/responses/buffered.rs b/crates/switchyard-translation/src/codecs/responses/buffered.rs index 998505913..e31763d65 100644 --- a/crates/switchyard-translation/src/codecs/responses/buffered.rs +++ b/crates/switchyard-translation/src/codecs/responses/buffered.rs @@ -91,12 +91,17 @@ impl FormatCodec for OpenAiResponsesCodec { }], }); } - // With `previous_response_id`, the provider holds the earlier turns, so a tool output - // may answer a call that is not in this body. + // Stored continuations may answer a tool call that is not in this body. let stored_state = body .get("previous_response_id") .and_then(Value::as_str) - .is_some_and(|id| !id.is_empty()); + .is_some_and(|id| !id.is_empty()) + || body.get("conversation").is_some_and(|conversation| { + conversation + .as_str() + .or_else(|| conversation.get("id").and_then(Value::as_str)) + .is_some_and(|id| !id.is_empty()) + }); let mut custom_call_outputs = Vec::new(); let (messages, instructions) = decode_responses_input( body.get("input").unwrap_or(&Value::String(String::new())), @@ -623,7 +628,7 @@ fn decode_responses_input( }], }; // An output whose call is not in this body answers a call the provider - // holds behind `previous_response_id`; it stays a tool result so routing + // holds behind a continuation ID; it stays a tool result so routing // sees a tool continuation, not a new user turn. Without stored state // the request is malformed, and the output becomes readable user text. let answers_pending_call = pending_tool_calls