Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
235 changes: 183 additions & 52 deletions crates/libsy-llm-client/src/run.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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(())
}
Expand Down Expand Up @@ -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..) {
Expand Down Expand Up @@ -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)
}
Expand Down Expand Up @@ -756,6 +765,7 @@ impl ClientRouter {
&accumulator.finish(),
&model,
store,
conversation.as_deref(),
input,
)?;
}
Expand Down Expand Up @@ -841,10 +851,12 @@ impl ClientRouter {
response: &AggLlmResponse,
model: &ModelId,
store: bool,
conversation: Option<&str>,
Comment thread
coderabbitai[bot] marked this conversation as resolved.
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();
Expand All @@ -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");
Expand Down Expand Up @@ -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| {
Expand Down Expand Up @@ -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))?,
Expand Down Expand Up @@ -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",
Expand All @@ -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))?,
Expand Down Expand Up @@ -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))?,
Expand All @@ -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,
Expand All @@ -1483,15 +1531,31 @@ 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");
assert_eq!(third["messages"][4]["content"][0]["text"], "thanks");
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()),
Expand All @@ -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 {
Expand All @@ -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(())
}

Expand Down Expand Up @@ -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");
Expand Down
Loading
Loading