diff --git a/crates/libsy-llm-client/src/run.rs b/crates/libsy-llm-client/src/run.rs index 01e1c36fc..25b23f447 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, @@ -1490,6 +1538,136 @@ mod tests { Ok(()) } + // Every prior turn must reach the backend on each conversation continuation. + #[tokio::test] + async fn responses_conversation_continuation_materializes_for_anthropic() -> Result<()> { + let server = MockServer::start().await; + Mock::given(method("POST")) + .respond_with(|request: &wiremock::Request| { + let body: Value = serde_json::from_slice(&request.body).expect("request JSON"); + let text = body.to_string(); + let id = if text.contains("thanks question") { + "msg_third" + } else if text.contains("recall question") { + "msg_follow" + } else { + "msg_seed" + }; + ResponseTemplate::new(200).set_body_json(json!({ + "id": id, "type": "message", "role": "assistant", "model": "weak", + "content": [{"type": "text", "text": "ok"}], + "stop_reason": "end_turn", + "usage": {"input_tokens": 1, "output_tokens": 1} + })) + }) + .mount(&server) + .await; + + let client: Arc = Arc::new( + TranslatingLlmClient::new(&[ModelConfig::new( + "weak", + Backend::Anthropic(HttpBackendConfig { + base_url: server.uri(), + api_key: None, + forward_auth: false, + extra_headers: BTreeMap::new(), + extra_body: BTreeMap::new(), + reasoning_effort: None, + max_retries: 0, + timeout: None, + }), + None, + )]) + .map_err(|error| LibsyError::external("building test client", error))?, + ); + let clients = ClientRouter::new(HashMap::from([(ModelId::from("weak"), client)])); + let models = to_category_map(&["weak"]); + + let request = |input: &str| { + let llm_request = switchyard_translation::decode_request( + WireFormat::OpenAiResponses, + &json!({ + "model": "route", + "input": input, + "conversation": "conv_081", + "store": true + }), + ) + .map_err(|error| LibsyError::external("decoding request", error)); + Ok::(Request { + llm_request: llm_request?, + raw_request: None, + metadata: None, + }) + }; + + let (_, response) = run( + Arc::new(switchyard_libsy::Passthrough), + clients.clone(), + request("seed question")?, + models.clone(), + None, + ) + .await?; + assert_eq!( + response + .llm_response + .as_agg() + .and_then(|agg| agg.id.as_deref()), + Some("msg_seed") + ); + + let (_, response) = run( + Arc::new(switchyard_libsy::Passthrough), + clients.clone(), + request("recall question")?, + models.clone(), + None, + ) + .await?; + assert_eq!( + completion_text( + response + .llm_response + .as_agg() + .expect("buffered follow-up response") + ), + "ok" + ); + + let (_, response) = run( + Arc::new(switchyard_libsy::Passthrough), + clients, + request("thanks question")?, + models, + None, + ) + .await?; + assert_eq!( + completion_text( + response + .llm_response + .as_agg() + .expect("buffered chained response") + ), + "ok" + ); + + let requests = server.received_requests().await.expect("request recording"); + assert_eq!(requests.len(), 3); + let seed = String::from_utf8_lossy(&requests[0].body); + assert!(seed.contains("seed question")); + assert!(!seed.contains("recall question")); + let follow = String::from_utf8_lossy(&requests[1].body); + assert!(follow.contains("seed question")); + assert!(follow.contains("recall question")); + let third = String::from_utf8_lossy(&requests[2].body); + assert!(third.contains("seed question")); + assert!(third.contains("recall question")); + assert!(third.contains("thanks question")); + Ok(()) + } + #[tokio::test] async fn streamed_cross_format_state_is_recorded_only_after_completion() -> Result<()> { let client = Arc::new(CandidateClient { @@ -1567,6 +1745,104 @@ mod tests { Ok(()) } + // Conversation history is recorded only once the stream completes. + #[tokio::test] + async fn streamed_cross_format_conversation_state_is_recorded_after_completion() -> Result<()> { + for (store, response_id) in [ + (true, Some("msg_stream")), + (false, Some("msg_stream")), + (true, None), + ] { + check_streamed_conversation_state(store, response_id).await?; + } + Ok(()) + } + + async fn check_streamed_conversation_state( + store: bool, + response_id: Option<&str>, + ) -> Result<()> { + let client = Arc::new(CandidateClient { + calls: Mutex::new(Vec::new()), + requests: Mutex::new(Vec::new()), + first: FirstOutcome::StreamSuccess, + }); + let clients = ClientRouter::new(HashMap::from([( + ModelId::from("weak"), + client as Arc, + )])); + let mut seed = request(); + seed.llm_request.messages = vec![Message::text(Role::User, "seed question")]; + 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: response_id.map(str::to_owned), + model: Some("weak".to_string()), + }, + LlmResponseChunk::TextDelta { + index: 0, + text: "streamed".to_string(), + }, + LlmResponseChunk::MessageStop { + reason: Some("stop".to_string()), + }, + ]); + 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("conversation".to_string(), json!("conv_stream")); + follow.llm_request.messages = vec![Message::text(Role::User, "recall question")]; + 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 conversation stream state"), + "passthrough", + follow, + ); + assert_eq!(outcome.request.llm_request.messages.len(), 3); + assert_eq!( + clients + .inner + .state_owners + .lock() + .owner("msg_stream") + .is_some(), + store && response_id.is_some() + ); + assert!(matches!( + outcome.request.llm_request.messages[0].content.as_slice(), + [ContentBlock::Text { text }] if text == "seed question" + )); + assert!(matches!( + outcome.request.llm_request.messages[1].content.as_slice(), + [ContentBlock::Text { text }] if text == "streamed" + )); + assert!(matches!( + outcome.request.llm_request.messages[2].content.as_slice(), + [ContentBlock::Text { text }] if text == "recall question" + )); + Ok(()) + } + #[test] fn cross_format_response_with_store_false_is_not_retained() -> Result<()> { let client = Arc::new(CandidateClient { @@ -1709,6 +1985,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