Skip to main content

switchyard_libsy/algorithms/
plan_execute.rs

1// SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4//! Plans coding tasks on a capable model, then hands execution to an efficient model.
5
6use std::collections::HashSet;
7use std::sync::Arc;
8
9use parking_lot::Mutex;
10use switchyard_protocol::{Category, ContentBlock, Request};
11
12use super::util::prompts::{append_note, drop_exact_replay, prepend_system_prompt};
13use super::util::tool_signals::{ToolSemantics, ToolSignals};
14use crate::core::algorithm::{Algorithm, Driver, RoutingIdentity};
15use crate::{LibsyError, Result, RoutingOutcome};
16
17/// Default instruction added while the capable model is planning.
18pub const DEFAULT_PLANNING_PROMPT: &str =
19    include_str!("../prompts/plan-execute/planning-system-prompt.md");
20
21const MAX_EXECUTING_SESSIONS: usize = 4_096;
22
23/// Configuration for [`PlanExecute`].
24#[derive(Clone, Debug)]
25pub struct PlanExecuteConfig {
26    /// Additional tool names whose mutations trigger handoff.
27    pub tool_semantics: ToolSemantics,
28    /// System instruction added until the first edit or write tool call.
29    pub planning_prompt: String,
30    /// Optional instruction appended to the handoff request.
31    pub handoff_prompt: Option<String>,
32    /// Replays visible planner reasoning summaries as assistant text at handoff.
33    pub planner_reasoning_as_text: bool,
34}
35
36impl Default for PlanExecuteConfig {
37    fn default() -> Self {
38        Self {
39            tool_semantics: ToolSemantics::default(),
40            planning_prompt: DEFAULT_PLANNING_PROMPT.trim().to_string(),
41            handoff_prompt: None,
42            planner_reasoning_as_text: false,
43        }
44    }
45}
46
47#[derive(Clone, Copy, Debug, Eq, PartialEq)]
48enum Phase {
49    Plan,
50    Handoff,
51    Execute,
52}
53
54/// Routes planning turns to the runtime capable model, then latches execution
55/// to the runtime efficient model after the first recorded mutation.
56pub struct PlanExecute {
57    config: PlanExecuteConfig,
58    executing_sessions: Mutex<HashSet<RoutingIdentity>>,
59}
60
61impl PlanExecute {
62    /// Creates a plan/execute router.
63    ///
64    /// Returns an error when a prompt is blank or tool semantics are invalid.
65    pub fn new(config: PlanExecuteConfig) -> Result<Self> {
66        if config.planning_prompt.trim().is_empty() {
67            return Err(LibsyError::AlgorithmError {
68                message: "planning_prompt must not be empty".to_string(),
69            });
70        }
71        if config
72            .handoff_prompt
73            .as_deref()
74            .is_some_and(|prompt| prompt.trim().is_empty())
75        {
76            return Err(LibsyError::AlgorithmError {
77                message: "handoff_prompt must not be empty".to_string(),
78            });
79        }
80        config.tool_semantics.validate()?;
81        Ok(Self {
82            config,
83            executing_sessions: Mutex::new(HashSet::new()),
84        })
85    }
86
87    fn phase(&self, request: &Request) -> Result<Phase> {
88        let signals =
89            ToolSignals::from_request_with_semantics(request, None, &self.config.tool_semantics);
90        let mutation_seen = signals.edit_count > 0 || signals.write_count > 0;
91        let Some(identity) = RoutingIdentity::from_request(request) else {
92            return Ok(if mutation_seen {
93                Phase::Handoff
94            } else {
95                Phase::Plan
96            });
97        };
98        let is_session_final = request
99            .metadata
100            .as_ref()
101            .and_then(|metadata| metadata.session_final)
102            == Some(true);
103
104        let mut sessions = self.executing_sessions.lock();
105        let phase = if sessions.contains(&identity) {
106            Phase::Execute
107        } else if mutation_seen {
108            if !is_session_final {
109                if sessions.len() >= MAX_EXECUTING_SESSIONS {
110                    return Err(LibsyError::AlgorithmError {
111                        message: format!(
112                            "plan_execute reached its limit of {MAX_EXECUTING_SESSIONS} executing sessions. \
113                             Finish an existing session with session_final before retrying the handoff"
114                        ),
115                    });
116                }
117                sessions.insert(identity.clone());
118            }
119            Phase::Handoff
120        } else {
121            Phase::Plan
122        };
123        if is_session_final {
124            sessions.remove(&identity);
125        }
126        Ok(phase)
127    }
128
129    fn replay_planner_reasoning_as_text(request: &mut Request) -> usize {
130        let mut converted = 0;
131        for message in &mut request.llm_request.messages {
132            message.content = std::mem::take(&mut message.content)
133                .into_iter()
134                .filter_map(|block| match block {
135                    ContentBlock::Reasoning { text, .. } => {
136                        converted += 1;
137                        (!text.is_empty()).then_some(ContentBlock::Text { text })
138                    }
139                    other => Some(other),
140                })
141                .collect();
142        }
143        request
144            .llm_request
145            .messages
146            .retain(|message| !message.content.is_empty());
147        if converted > 0 {
148            drop_exact_replay(request);
149        }
150        converted
151    }
152
153    fn route_to(driver: &Driver, category: Category, request: Request) -> Result<RoutingOutcome> {
154        let models = driver.models_for(&category);
155        let Some((selected, fallbacks)) = models.split_first() else {
156            return Err(LibsyError::AlgorithmError {
157                message: format!("no models available for category {}", category.as_str()),
158            });
159        };
160        Ok(RoutingOutcome::route_to(
161            selected.clone(),
162            fallbacks.to_vec(),
163            request,
164        ))
165    }
166}
167
168#[async_trait::async_trait]
169impl Algorithm for PlanExecute {
170    fn name(&self) -> &str {
171        "plan_execute"
172    }
173
174    async fn route(
175        self: Arc<Self>,
176        driver: Driver,
177        mut request: Request,
178    ) -> Result<RoutingOutcome> {
179        match self.phase(&request)? {
180            Phase::Plan => {
181                prepend_system_prompt(&mut request, &self.config.planning_prompt);
182                tracing::debug!(phase = "plan", "plan-execute selected capable tier");
183                Self::route_to(&driver, Category::Capable, request)
184            }
185            Phase::Handoff => {
186                let reasoning_converted = if self.config.planner_reasoning_as_text {
187                    Self::replay_planner_reasoning_as_text(&mut request)
188                } else {
189                    0
190                };
191                let prompt_applied = if let Some(prompt) = &self.config.handoff_prompt {
192                    append_note(&mut request, prompt);
193                    true
194                } else {
195                    false
196                };
197                tracing::debug!(
198                    phase = "handoff",
199                    handoff_prompt_applied = prompt_applied,
200                    planner_reasoning_converted = reasoning_converted,
201                    "plan-execute selected efficient tier"
202                );
203                Self::route_to(&driver, Category::Efficient, request)
204            }
205            Phase::Execute => {
206                tracing::debug!(phase = "execute", "plan-execute selected efficient tier");
207                Self::route_to(&driver, Category::Efficient, request)
208            }
209        }
210    }
211}
212
213#[cfg(test)]
214mod tests {
215    use std::collections::HashMap;
216    use std::sync::Arc;
217
218    use serde_json::json;
219    use switchyard_protocol::{
220        ContentBlock, LlmRequest, Message, Metadata, ModelId, Request, Role, ToolCall,
221    };
222
223    use super::*;
224    use crate::RuntimeModels;
225    use crate::core::testing::{reply, test_drive_with_models};
226
227    const CAPABLE: &str = "model/capable";
228    const EFFICIENT: &str = "model/efficient";
229
230    fn algorithm(config: PlanExecuteConfig) -> Arc<dyn Algorithm> {
231        Arc::new(PlanExecute::new(config).expect("config should be valid"))
232    }
233
234    fn request(messages: Vec<Message>, session_id: Option<&str>) -> Request {
235        Request {
236            llm_request: LlmRequest {
237                model: Some("switchyard/plan-execute".to_string()),
238                messages,
239                ..LlmRequest::default()
240            },
241            metadata: session_id.map(|session_id| Metadata {
242                session_id: Some(session_id.to_string()),
243                ..Metadata::default()
244            }),
245            ..Request::default()
246        }
247    }
248
249    fn tool_call(name: &str, arguments: serde_json::Value) -> Message {
250        Message {
251            role: Role::Assistant,
252            content: vec![ContentBlock::ToolCall(ToolCall {
253                id: "call-1".to_string(),
254                name: name.to_string(),
255                arguments,
256            })],
257        }
258    }
259
260    fn models() -> RuntimeModels {
261        RuntimeModels::new(HashMap::from([
262            (Category::Capable, vec![ModelId::from(CAPABLE)]),
263            (Category::Efficient, vec![ModelId::from(EFFICIENT)]),
264        ]))
265    }
266
267    async fn route_and_capture(
268        algorithm: Arc<dyn Algorithm>,
269        request: Request,
270    ) -> (ModelId, Request) {
271        let captured = Arc::new(Mutex::new(None));
272        let capture = Arc::clone(&captured);
273        let (selected, _) =
274            test_drive_with_models(algorithm, request, models(), move |_target, request| {
275                let capture = Arc::clone(&capture);
276                async move {
277                    *capture.lock() = Some(request);
278                    Ok(reply("ok"))
279                }
280            })
281            .await
282            .expect("routing should succeed");
283        let request = captured
284            .lock()
285            .take()
286            .expect("answer request should be captured");
287        (selected, request)
288    }
289
290    #[tokio::test]
291    async fn plans_then_hands_off_and_latches_execution() {
292        const HANDOFF: &str = "Continue from the plan and repository evidence.";
293        let algorithm = algorithm(PlanExecuteConfig {
294            handoff_prompt: Some(HANDOFF.to_string()),
295            planner_reasoning_as_text: true,
296            ..PlanExecuteConfig::default()
297        });
298
299        let read_only = request(
300            vec![tool_call(
301                "exec_command",
302                json!({"cmd": "rg parser crates"}),
303            )],
304            Some("task-1"),
305        );
306        let (selected, routed) = route_and_capture(Arc::clone(&algorithm), read_only).await;
307        assert_eq!(selected, CAPABLE);
308        assert_eq!(routed.llm_request.instructions.len(), 1);
309
310        let first_edit = request(
311            vec![Message {
312                role: Role::Assistant,
313                content: vec![
314                    ContentBlock::Reasoning {
315                        text: "The parser needs a boundary check.".to_string(),
316                        signature: Some("planner-signature".to_string()),
317                        details: vec![json!({"type": "reasoning.encrypted", "data": "opaque"})],
318                    },
319                    ContentBlock::ToolCall(ToolCall {
320                        id: "call-1".to_string(),
321                        name: "apply_patch".to_string(),
322                        arguments: json!({"patch": "*** Begin Patch"}),
323                    }),
324                ],
325            }],
326            Some("task-1"),
327        );
328        let (selected, routed) = route_and_capture(Arc::clone(&algorithm), first_edit).await;
329        assert_eq!(selected, EFFICIENT);
330        assert_eq!(
331            routed.llm_request.messages[0].content[0],
332            ContentBlock::Text {
333                text: "The parser needs a boundary check.".to_string()
334            }
335        );
336        assert_eq!(
337            routed.llm_request.messages.last(),
338            Some(&Message::text(Role::User, HANDOFF))
339        );
340
341        let mut final_request = request(
342            vec![Message::text(Role::User, "Continue after compaction")],
343            Some("task-1"),
344        );
345        final_request
346            .metadata
347            .as_mut()
348            .expect("session metadata should exist")
349            .session_final = Some(true);
350        let (selected, routed) = route_and_capture(Arc::clone(&algorithm), final_request).await;
351        assert_eq!(selected, EFFICIENT);
352        assert!(routed.llm_request.instructions.is_empty());
353        assert_eq!(routed.llm_request.messages.len(), 1);
354
355        let reused = request(vec![Message::text(Role::User, "New task")], Some("task-1"));
356        let (selected, _) = route_and_capture(algorithm, reused).await;
357        assert_eq!(selected, CAPABLE);
358    }
359
360    #[tokio::test]
361    async fn capacity_preserves_execution_after_compaction() {
362        let algorithm = algorithm(PlanExecuteConfig::default());
363        let sessions: Vec<_> = (0..MAX_EXECUTING_SESSIONS)
364            .map(|index| format!("task-{index}"))
365            .collect();
366        for session in &sessions {
367            let first_edit = request(
368                vec![tool_call("Write", json!({"file_path": "task.py"}))],
369                Some(session),
370            );
371            let (selected, _) = route_and_capture(Arc::clone(&algorithm), first_edit).await;
372            assert_eq!(selected, EFFICIENT);
373        }
374
375        let overflow = request(
376            vec![tool_call("Write", json!({"file_path": "task.py"}))],
377            Some("overflow"),
378        );
379        let (driver, _) = Driver::new("plan_execute", Arc::new(models()));
380        let result = Arc::clone(&algorithm).route(driver, overflow.clone()).await;
381        assert!(matches!(
382            result,
383            Err(LibsyError::AlgorithmError { message })
384                if message.contains("limit of 4096 executing sessions")
385        ));
386
387        for session in &sessions {
388            let mut compacted = request(
389                vec![Message::text(Role::User, "Continue after compaction")],
390                Some(session),
391            );
392            if Some(session) == sessions.last() {
393                compacted.metadata.as_mut().unwrap().session_final = Some(true);
394            }
395            let (selected, routed) = route_and_capture(Arc::clone(&algorithm), compacted).await;
396            assert_eq!(selected, EFFICIENT, "session {session} lost its latch");
397            assert!(routed.llm_request.instructions.is_empty());
398        }
399
400        let (selected, _) = route_and_capture(Arc::clone(&algorithm), overflow).await;
401        assert_eq!(selected, EFFICIENT);
402        let compacted = request(
403            vec![Message::text(Role::User, "Continue after compaction")],
404            Some("overflow"),
405        );
406        let (selected, routed) = route_and_capture(algorithm, compacted).await;
407        assert_eq!(selected, EFFICIENT);
408        assert!(routed.llm_request.instructions.is_empty());
409    }
410
411    #[tokio::test]
412    async fn final_handoff_does_not_evict_an_executing_session_at_capacity() {
413        let algorithm = Arc::new(PlanExecute::new(PlanExecuteConfig::default()).unwrap());
414        algorithm.executing_sessions.lock().extend(
415            (0..MAX_EXECUTING_SESSIONS)
416                .map(|index| RoutingIdentity::Session(format!("task-{index}"))),
417        );
418        let mut final_request = request(
419            vec![tool_call("Write", json!({"file_path": "task.py"}))],
420            Some("final-handoff"),
421        );
422        final_request.metadata.as_mut().unwrap().session_final = Some(true);
423
424        let (selected, routed) = route_and_capture(algorithm.clone(), final_request).await;
425        assert_eq!(selected, EFFICIENT);
426        assert!(routed.llm_request.instructions.is_empty());
427        let sessions = algorithm.executing_sessions.lock();
428        assert_eq!(sessions.len(), MAX_EXECUTING_SESSIONS);
429        assert!(!sessions.contains(&RoutingIdentity::Session("final-handoff".to_string())));
430    }
431
432    #[tokio::test]
433    async fn mutation_without_a_session_uses_the_efficient_tier() {
434        let messages = vec![tool_call(
435            "exec_command",
436            json!({"cmd": "printf 'done\\n' > task.txt"}),
437        )];
438
439        let (selected, routed) = route_and_capture(
440            algorithm(PlanExecuteConfig::default()),
441            request(messages.clone(), None),
442        )
443        .await;
444
445        assert_eq!(selected, EFFICIENT);
446        assert_eq!(routed.llm_request.messages, messages);
447        assert!(routed.llm_request.instructions.is_empty());
448    }
449
450    #[tokio::test]
451    async fn editor_view_keeps_planning() {
452        let messages = vec![tool_call(
453            "str_replace_based_edit_tool",
454            json!({"command": "view", "path": "/app/main.py"}),
455        )];
456
457        let (selected, _) = route_and_capture(
458            algorithm(PlanExecuteConfig::default()),
459            request(messages, None),
460        )
461        .await;
462
463        assert_eq!(selected, CAPABLE);
464    }
465
466    #[test]
467    fn rejects_blank_prompts() {
468        for config in [
469            PlanExecuteConfig {
470                planning_prompt: "  ".to_string(),
471                ..PlanExecuteConfig::default()
472            },
473            PlanExecuteConfig {
474                handoff_prompt: Some("  ".to_string()),
475                ..PlanExecuteConfig::default()
476            },
477        ] {
478            assert!(matches!(
479                PlanExecute::new(config),
480                Err(LibsyError::AlgorithmError { .. })
481            ));
482        }
483    }
484}