-
Notifications
You must be signed in to change notification settings - Fork 319
fix(plan-execute): replay stored Responses history for handoff #897
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -55,20 +55,20 @@ use crate::{metrics, observability}; | |
| /// Once routing completes, non-timeout failures may try the outcome's ordered fallback candidates. | ||
| pub async fn run( | ||
| algorithm: Arc<dyn Algorithm>, | ||
| clients: ClientRouter, | ||
| request: Request, | ||
| mut clients: ClientRouter, | ||
| mut request: Request, | ||
| models: Arc<RuntimeModels>, | ||
| observer: Option<RunObserver>, | ||
| ) -> Result<(ModelId, Response)> { | ||
| let algorithm_name = algorithm.name().to_string(); | ||
| let run_started = Instant::now(); | ||
| let routing_clients = clients.clone(); | ||
| // This says if we have an observer, put Some(..) in routing_observations. | ||
| // No observer means we don't want any routing_observations. | ||
| let routing_observations = observer.as_ref().map(|_| Arc::new(Mutex::new(Vec::new()))); | ||
| let outcome = match clients.stored_state_owner(&request) { | ||
| let outcome = match clients.prepare_stored_continuation(algorithm.as_ref(), &mut request) { | ||
| Some(owner) => Ok(continue_on(owner, &algorithm_name, request)), | ||
| None => { | ||
| let routing_clients = clients.clone(); | ||
| drive(algorithm, request, models, { | ||
| let routing_observations = routing_observations.clone(); | ||
| move |call| serve(routing_clients.clone(), call, routing_observations.clone()) | ||
|
|
@@ -129,14 +129,14 @@ pub async fn run( | |
| /// target; use [`run`] to execute selected and fallback candidates. | ||
| pub async fn decide( | ||
| algorithm: Arc<dyn Algorithm>, | ||
| clients: ClientRouter, | ||
| request: Request, | ||
| mut clients: ClientRouter, | ||
| mut request: Request, | ||
| models: Arc<RuntimeModels>, | ||
| ) -> Result<RoutingOutcome> { | ||
| let routing_clients = clients.clone(); | ||
| let mut outcome = match clients.stored_state_owner(&request) { | ||
| let mut outcome = match clients.prepare_stored_continuation(algorithm.as_ref(), &mut request) { | ||
| Some(owner) => continue_on(owner, algorithm.name(), request), | ||
| None => { | ||
| let routing_clients = clients.clone(); | ||
| drive(algorithm, request, models, move |call| { | ||
| serve(routing_clients.clone(), call, None) | ||
| }) | ||
|
|
@@ -564,6 +564,7 @@ fn conversation_id(fields: &serde_json::Map<String, Value>) -> Option<&str> { | |
| #[derive(Clone)] | ||
| pub struct ClientRouter { | ||
| inner: Arc<ClientRouting>, | ||
| needs_history_replay: bool, | ||
| } | ||
|
|
||
| struct ClientRouting { | ||
|
|
@@ -631,6 +632,7 @@ impl ClientRouter { | |
| .next() | ||
| .is_some_and(|first| clients.any(|client| !Arc::ptr_eq(first, client))); | ||
| Self { | ||
| needs_history_replay: false, | ||
| inner: Arc::new(ClientRouting { | ||
| routing: Routing::ByModel(by_model), | ||
| target_prompts, | ||
|
|
@@ -648,6 +650,7 @@ impl ClientRouter { | |
| /// only duplicate that. | ||
| pub fn single(client: Arc<dyn RoutedLlmClient>) -> Self { | ||
| Self { | ||
| needs_history_replay: false, | ||
| inner: Arc::new(ClientRouting { | ||
| routing: Routing::Single(client), | ||
| target_prompts: HashMap::new(), | ||
|
|
@@ -693,14 +696,58 @@ impl ClientRouter { | |
| }) | ||
| } | ||
|
|
||
| fn prepare_stored_continuation( | ||
| &mut self, | ||
| algorithm: &dyn Algorithm, | ||
| request: &mut Request, | ||
| ) -> Option<StateOwner> { | ||
| let owner = self.stored_state_owner(request); | ||
| let fields = &request.llm_request.extensions.fields; | ||
| self.needs_history_replay = algorithm.needs_history_replay(request) | ||
| && fields.get("conversation").is_none_or(Value::is_null) | ||
| && match fields | ||
| .get("previous_response_id") | ||
| .filter(|value| !value.is_null()) | ||
| { | ||
| Some(_) => owner.as_ref().is_some_and(|owner| owner.history.is_some()), | ||
| None => request | ||
| .llm_request | ||
| .preservation | ||
| .requests | ||
| .contains_key(&WireFormat::OpenAiResponses.into()), | ||
| }; | ||
| let owner = owner?; | ||
| if self.needs_history_replay | ||
| && let Some(history) = &owner.history | ||
| { | ||
| let mut messages = Vec::with_capacity(history.len + request.llm_request.messages.len()); | ||
| history.extend(&mut messages); | ||
| messages.append(&mut request.llm_request.messages); | ||
| request.llm_request.messages = messages; | ||
| request.llm_request.preservation.requests.clear(); | ||
| None | ||
|
Comment on lines
+727
to
+728
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🗄️ Data Integrity & Integration | 🟠 Major | ⚡ Quick win Remove If an opted-in algorithm makes a routing-time 🤖 Prompt for AI Agents |
||
| } else { | ||
| Some(owner) | ||
| } | ||
| } | ||
|
|
||
| fn canonical_input(&self, request: &Request) -> CanonicalInput { | ||
| let parent = request | ||
| .llm_request | ||
| .extensions | ||
| .fields | ||
| .get("previous_response_id") | ||
| .and_then(Value::as_str) | ||
| .and_then(|id| self.inner.state_owners.lock().owner(id)?.history.clone()); | ||
| .and_then(|id| self.inner.state_owners.lock().owner(id)?.history.clone()) | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🚀 Performance & Scalability | 🟠 Major | 🏗️ Heavy lift Retain the parent history link after removing the provider ID. For a replayed continuation, 🤖 Prompt for AI Agents |
||
| .filter(|history| { | ||
| if !self.needs_history_replay { | ||
| return true; | ||
| } | ||
| // Handoff processing can rewrite the saved prefix. | ||
| let mut messages = Vec::with_capacity(history.len); | ||
| history.extend(&mut messages); | ||
| request.llm_request.messages.starts_with(&messages) | ||
| }); | ||
| let parent_len = parent.as_ref().map_or(0, |history| history.len); | ||
| let (parent, messages) = match request.llm_request.messages.get(parent_len..) { | ||
| Some(messages) => (parent, messages.to_vec()), | ||
|
|
@@ -721,11 +768,12 @@ impl ClientRouter { | |
| let store = fields.get("store").and_then(Value::as_bool) != Some(false); | ||
| let conversation = conversation_id(fields).map(str::to_owned); | ||
| let responses_format = WireFormat::OpenAiResponses.into(); | ||
| let responses_request = request | ||
| .llm_request | ||
| .preservation | ||
| .requests | ||
| .contains_key(&responses_format) | ||
| let responses_request = self.needs_history_replay | ||
| || request | ||
| .llm_request | ||
| .preservation | ||
| .requests | ||
| .contains_key(&responses_format) | ||
| || fields.contains_key("previous_response_id") | ||
| || fields.contains_key("conversation"); | ||
| if !responses_request && !self.inner.track_provider_state { | ||
|
|
@@ -739,6 +787,14 @@ impl ClientRouter { | |
| self.remember_response(body, &model, store, conversation.as_deref()) | ||
| .map_err(|error| LibsyError::client_call(model.clone(), error))?; | ||
| } | ||
| if self.needs_history_replay | ||
| && body["status"] == "completed" | ||
| && let Some(input) = &canonical_input | ||
| { | ||
| let store = body.get("store").and_then(Value::as_bool).unwrap_or(store); | ||
| self.remember_canonical_response(&agg, &model, store, input) | ||
| .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))?; | ||
|
|
@@ -757,6 +813,7 @@ impl ClientRouter { | |
| ResponseAccumulator::new(), | ||
| false, | ||
| true, | ||
| store, | ||
| ), | ||
| move |( | ||
| mut stream, | ||
|
|
@@ -767,6 +824,7 @@ impl ClientRouter { | |
| mut accumulator, | ||
| mut native_responses, | ||
| mut valid, | ||
| mut store, | ||
| )| async move { | ||
| let Some(event) = stream.next().await else { | ||
| if let Some(input) = &canonical_input | ||
|
|
@@ -783,19 +841,22 @@ impl ClientRouter { | |
| return Ok(None); | ||
| }; | ||
| let event = event?; | ||
| let mut is_completed = false; | ||
| if let Some(preserved) = event.preservation() | ||
| && preserved.source().as_str() == WireFormat::OpenAiResponses.as_str() | ||
| { | ||
| native_responses = true; | ||
| if router.inner.track_provider_state | ||
| && let Some(body) = preserved.raw().get("response") | ||
| { | ||
| router.remember_response( | ||
| body, | ||
| &model, | ||
| store, | ||
| conversation.as_deref(), | ||
| )?; | ||
| is_completed = preserved.raw()["type"] == "response.completed"; | ||
| if let Some(body) = preserved.raw().get("response") { | ||
| store = body.get("store").and_then(Value::as_bool).unwrap_or(store); | ||
| if router.inner.track_provider_state { | ||
| router.remember_response( | ||
| body, | ||
| &model, | ||
| store, | ||
| conversation.as_deref(), | ||
| )?; | ||
| } | ||
| } | ||
| } | ||
| if responses_request { | ||
|
|
@@ -811,6 +872,18 @@ impl ClientRouter { | |
| } | ||
| } | ||
| } | ||
| if router.needs_history_replay | ||
| && is_completed | ||
| && valid | ||
| && let Some(input) = &canonical_input | ||
| { | ||
| router.remember_canonical_response( | ||
| &std::mem::take(&mut accumulator).finish(), | ||
| &model, | ||
| store, | ||
| input, | ||
| )?; | ||
| } | ||
| Ok(Some(( | ||
| event, | ||
| ( | ||
|
|
@@ -822,6 +895,7 @@ impl ClientRouter { | |
| accumulator, | ||
| native_responses, | ||
| valid, | ||
| store, | ||
| ), | ||
| ))) | ||
| }, | ||
|
|
@@ -895,6 +969,13 @@ impl ClientRouter { | |
|
|
||
| /// Prepare a completion candidate with its configured target prompt. | ||
| fn prepare_completion_request(&self, mut request: Request, target: &ModelId) -> Request { | ||
| if self.needs_history_replay { | ||
| request | ||
| .llm_request | ||
| .extensions | ||
| .fields | ||
| .remove("previous_response_id"); | ||
| } | ||
| let prompt = self.inner.target_prompts.get(target).map(String::as_str); | ||
| prepare_request_for_target(&mut request.llm_request, target, prompt); | ||
| request | ||
|
|
@@ -1306,6 +1387,123 @@ mod tests { | |
| Ok(()) | ||
| } | ||
|
|
||
| #[tokio::test] | ||
| async fn plan_execute_replays_stored_tool_history() -> Result<()> { | ||
| let decode = |body: &Value| { | ||
| switchyard_translation::decode_request(WireFormat::OpenAiResponses, body).unwrap() | ||
| }; | ||
| for (stream, null_fields) in [(false, false), (true, false), (false, true), (true, true)] { | ||
| let server = MockServer::start().await; | ||
| let call = json!({"type": "function_call", "call_id": "call_file", | ||
| "name": "write_file", "arguments": "{}"}); | ||
| let result = json!({"type": "function_call_output", "call_id": "call_file", | ||
| "output": "success"}); | ||
| let output = call.clone(); | ||
| Mock::given(method("POST")) | ||
| .respond_with(move |request: &wiremock::Request| { | ||
| let body: Value = serde_json::from_slice(&request.body).unwrap(); | ||
| let model = body["model"].as_str().unwrap(); | ||
| let response = json!({"id": format!("resp_{model}"), "object": "response", | ||
| "model": model, "status": "completed", "store": true, "output": [output]}); | ||
| if stream { | ||
| let created = json!({"type": "response.created", "response": { | ||
| "id": response["id"], "model": model, "status": "in_progress" | ||
| }}); | ||
| let event = json!({"type": "response.completed", "response": response}); | ||
| ResponseTemplate::new(200) | ||
| .insert_header("content-type", "text/event-stream") | ||
| .set_body_string(format!("data: {created}\n\ndata: {event}\n\n")) | ||
| } else { | ||
| ResponseTemplate::new(200).set_body_json(response) | ||
| } | ||
| }) | ||
| .expect(2) | ||
| .mount(&server) | ||
| .await; | ||
| let backend = Backend::OpenAiResponses(HttpBackendConfig { | ||
| base_url: server.uri(), | ||
| api_key: None, | ||
| forward_auth: false, | ||
| extra_headers: BTreeMap::new(), | ||
| extra_body: BTreeMap::new(), | ||
| omit_body_fields: BTreeSet::new(), | ||
| reasoning_effort: None, | ||
| max_retries: 0, | ||
| failure_cooldown: std::time::Duration::ZERO, | ||
| timeout: None, | ||
| }); | ||
| let clients: ClientRouter = ["strong", "weak"] | ||
| .into_iter() | ||
| .map(|model| { | ||
| let config = ModelConfig::new(model, backend.clone(), None); | ||
| let client = TranslatingLlmClient::new(&[config]).unwrap(); | ||
| (model.into(), Arc::new(client) as Arc<dyn RoutedLlmClient>) | ||
| }) | ||
| .collect(); | ||
| let models = Arc::new(RuntimeModels::new(HashMap::from([ | ||
| (Category::Capable, vec![ModelId::from("strong")]), | ||
| (Category::Efficient, vec![ModelId::from("weak")]), | ||
| ]))); | ||
| let algorithm: Arc<dyn Algorithm> = | ||
| Arc::new(switchyard_libsy::PlanExecute::new(Default::default())?); | ||
| for (mut body, expected) in [ | ||
| ( | ||
| json!({"model": "route", "input": "Write task.py", "stream": stream}), | ||
| "strong", | ||
| ), | ||
| ( | ||
| json!({"model": "route", "stream": stream, | ||
| "previous_response_id": "resp_strong", "input": [result]}), | ||
| "weak", | ||
| ), | ||
| ] { | ||
| if null_fields { | ||
| body["conversation"] = Value::Null; | ||
| if body.get("previous_response_id").is_none() { | ||
| body["previous_response_id"] = Value::Null; | ||
| } | ||
| } | ||
| let request = Request { | ||
| llm_request: decode(&body), | ||
| ..Request::default() | ||
| }; | ||
| let (selected, response) = run( | ||
| algorithm.clone(), | ||
| clients.clone(), | ||
| request, | ||
| models.clone(), | ||
| None, | ||
| ) | ||
| .await?; | ||
| assert_eq!( | ||
| selected.as_str(), | ||
| expected, | ||
| "stream={stream}, null_fields={null_fields}" | ||
| ); | ||
| if let LlmResponse::Stream(mut events) = response.llm_response { | ||
| while let Some(event) = events.next().await { | ||
| if event | ||
| .unwrap() | ||
| .preservation() | ||
| .is_some_and(|event| event.raw()["type"] == "response.completed") | ||
| { | ||
| break; | ||
| } | ||
| } | ||
| } | ||
| } | ||
| let requests = server.received_requests().await.unwrap(); | ||
| let sent: Value = serde_json::from_slice(&requests[1].body).unwrap(); | ||
| assert_eq!(sent["model"], "weak"); | ||
| assert!(sent.get("previous_response_id").is_none()); | ||
| let expected = json!({"input": [ | ||
| {"role": "user", "content": "Write task.py"}, call, result | ||
| ]}); | ||
| assert_eq!(decode(&sent).messages, decode(&expected).messages); | ||
| } | ||
| Ok(()) | ||
| } | ||
|
|
||
| #[tokio::test] | ||
| async fn responses_stored_tool_continuation_materializes_for_anthropic() -> Result<()> { | ||
| let server = MockServer::start().await; | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
🔒 Security & Privacy | 🟠 Major | 🏗️ Heavy lift
Scope stored-history lookup to its authorized caller.
If one
ClientRouterserves multiple callers, a caller with another caller’s response ID can select that stored history here. The new replay path then sends the prior messages, including tool results, to the model selected for the new request.StateOwnerskeys history only by response ID and retains no caller scope to check. Partition the cache by an authenticated caller identity, or enforce that identity before replay. Do not use a caller-supplied session ID as the sole authorization check.🤖 Prompt for AI Agents