switchyard_libsy/algorithms/util/
prompts.rs1use std::collections::BTreeMap;
19
20use async_trait::async_trait;
21use switchyard_protocol::{ContentBlock, InstructionBlock, Message, Request, Role};
22
23use crate::Result;
24use crate::core::processor::{Event, Processor};
25
26pub fn append_note(request: &mut Request, note: &str) {
35 match request.llm_request.messages.last_mut() {
36 Some(last) if last.role == Role::User => last.content.push(ContentBlock::Text {
37 text: note.to_string(),
38 }),
39 _ => request
40 .llm_request
41 .messages
42 .push(Message::text(Role::User, note)),
43 }
44}
45
46#[derive(Clone, Debug, Default)]
49pub struct TargetPrompts {
50 by_target: BTreeMap<String, String>,
51}
52
53impl TargetPrompts {
54 pub fn with(mut self, target: impl Into<String>, prompt: impl Into<String>) -> Self {
56 self.by_target.insert(target.into(), prompt.into());
57 self
58 }
59
60 pub fn get(&self, target: &str) -> Option<&str> {
62 self.by_target.get(target).map(String::as_str)
63 }
64
65 pub fn is_empty(&self) -> bool {
68 self.by_target.is_empty()
69 }
70}
71
72pub struct SystemPromptProcessor {
74 prompts: TargetPrompts,
75}
76
77impl SystemPromptProcessor {
78 pub fn new(prompts: TargetPrompts) -> Self {
80 Self { prompts }
81 }
82}
83
84#[async_trait]
85impl<S: Send> Processor<S> for SystemPromptProcessor {
86 async fn process(&self, _state: &mut S, event: Event<'_>) -> Result<()> {
87 let Event::Decision { request, decision } = event else {
91 return Ok(());
92 };
93 let Some(prompt) = self.prompts.get(decision.selected_model()) else {
94 return Ok(());
95 };
96 request.llm_request.instructions.insert(
99 0,
100 InstructionBlock {
101 role: Role::System,
102 content: vec![ContentBlock::Text {
103 text: prompt.to_string(),
104 }],
105 },
106 );
107 Ok(())
108 }
109}
110
111#[cfg(test)]
112mod tests {
113 use super::*;
114 use switchyard_protocol::{LlmRequest, ToolResult, text_request};
115
116 const NOTE: &str = "recovering from an error";
117 const STRONG_PROMPT: &str = "diagnose before you edit";
118 const WEAK_PROMPT: &str = "follow the settled plan";
119
120 fn request_with(messages: Vec<Message>) -> Request {
121 Request {
122 llm_request: LlmRequest {
123 messages,
124 ..LlmRequest::default()
125 },
126 raw_request: None,
127 metadata: None,
128 }
129 }
130
131 #[test]
132 fn a_note_joins_a_trailing_user_turn_after_its_tool_result() {
133 let tool_result = ContentBlock::ToolResult(ToolResult {
136 tool_call_id: "call_1".to_string(),
137 content: vec![ContentBlock::Text {
138 text: "exit 1".to_string(),
139 }],
140 is_error: Some(true),
141 });
142 let mut request = request_with(vec![Message {
143 role: Role::User,
144 content: vec![tool_result.clone()],
145 }]);
146
147 append_note(&mut request, NOTE);
148
149 let messages = &request.llm_request.messages;
150 assert_eq!(messages.len(), 1, "no second consecutive user turn");
151 assert_eq!(
152 messages[0].content,
153 vec![
154 tool_result,
155 ContentBlock::Text {
156 text: NOTE.to_string()
157 }
158 ]
159 );
160 }
161
162 #[test]
163 fn a_note_opens_a_user_turn_after_an_assistant_turn() {
164 let mut request = request_with(vec![Message::text(Role::Assistant, "done")]);
165
166 append_note(&mut request, NOTE);
167
168 let messages = &request.llm_request.messages;
169 assert_eq!(messages.len(), 2);
170 assert_eq!(messages[1].role, Role::User);
171 assert_eq!(messages[1].text_content(""), Some(NOTE.to_string()));
172 }
173
174 #[test]
175 fn a_note_leaves_the_rest_of_the_conversation_untouched() {
176 let mut request = Request {
177 llm_request: text_request(Some("auto".to_string()), "fix the build"),
178 raw_request: None,
179 metadata: None,
180 };
181
182 append_note(&mut request, NOTE);
183
184 let trail: Vec<String> = request
185 .llm_request
186 .messages
187 .iter()
188 .filter_map(|message| message.text_content("|"))
189 .collect();
190 assert_eq!(trail, vec![format!("fix the build|{NOTE}")]);
191 }
192
193 struct RoutedTo(&'static str);
195 impl switchyard_protocol::Decision for RoutedTo {
196 fn selected_model(&self) -> &str {
197 self.0
198 }
199 fn reasoning(&self) -> Option<&str> {
200 None
201 }
202 fn as_any(&self) -> &dyn std::any::Any {
203 self
204 }
205 }
206
207 fn instructions(request: &Request) -> Vec<String> {
209 request
210 .llm_request
211 .instructions
212 .iter()
213 .filter_map(|block| {
214 block.content.iter().find_map(|content| match content {
215 ContentBlock::Text { text } => Some(text.clone()),
216 _ => None,
217 })
218 })
219 .collect()
220 }
221
222 async fn run(processor: &SystemPromptProcessor, target: &'static str) -> Result<Request> {
224 let mut request = Request::default();
225 processor
226 .process(
227 &mut (),
228 Event::Decision {
229 request: &mut request,
230 decision: &RoutedTo(target),
231 },
232 )
233 .await?;
234 Ok(request)
235 }
236
237 fn request_with_preserved_body() -> Request {
239 let mut request = Request::default();
240 request.llm_request.preservation.requests.insert(
241 "openai_chat".into(),
242 serde_json::json!({
243 "model": "weak-model",
244 "messages": [{"role": "user", "content": "hi"}],
245 }),
246 );
247 request.llm_request.seal_preservation();
249 request
250 }
251
252 fn prompts() -> TargetPrompts {
253 TargetPrompts::default()
254 .with("strong", STRONG_PROMPT)
255 .with("weak", WEAK_PROMPT)
256 }
257
258 #[tokio::test]
259 async fn each_target_gets_its_own_prompt() -> Result<()> {
260 let processor = SystemPromptProcessor::new(prompts());
261 for (target, expected) in [("strong", STRONG_PROMPT), ("weak", WEAK_PROMPT)] {
262 assert_eq!(
263 instructions(&run(&processor, target).await?),
264 vec![expected]
265 );
266 }
267 Ok(())
268 }
269
270 #[tokio::test]
271 async fn an_unconfigured_target_is_left_untouched() -> Result<()> {
272 let processor =
274 SystemPromptProcessor::new(TargetPrompts::default().with("strong", STRONG_PROMPT));
275 assert_eq!(
276 instructions(&run(&processor, "strong").await?),
277 vec![STRONG_PROMPT]
278 );
279 assert!(instructions(&run(&processor, "weak").await?).is_empty());
280 Ok(())
281 }
282
283 #[tokio::test]
284 async fn the_prompt_leads_the_client_instructions() -> Result<()> {
285 let processor = SystemPromptProcessor::new(prompts());
286 let mut request = Request::default();
287 request.llm_request.instructions.push(InstructionBlock {
288 role: Role::System,
289 content: vec![ContentBlock::Text {
290 text: "you are a coding agent".to_string(),
291 }],
292 });
293
294 processor
295 .process(
296 &mut (),
297 Event::Decision {
298 request: &mut request,
299 decision: &RoutedTo("strong"),
300 },
301 )
302 .await?;
303
304 assert_eq!(
305 instructions(&request),
306 vec![STRONG_PROMPT, "you are a coding agent"]
307 );
308 Ok(())
309 }
310
311 #[tokio::test]
312 async fn the_inbound_request_is_left_alone() -> Result<()> {
313 let processor = SystemPromptProcessor::new(prompts());
315 let mut request = Request::default();
316 processor
317 .process(&mut (), Event::Request(&mut request))
318 .await?;
319 assert!(instructions(&request).is_empty());
320 Ok(())
321 }
322
323 #[tokio::test]
328 async fn any_tier_prompt_invalidates_exact_replay() -> Result<()> {
329 let processor = SystemPromptProcessor::new(prompts());
330 for (target, expected) in [("strong", STRONG_PROMPT), ("weak", WEAK_PROMPT)] {
331 let mut request = request_with_preserved_body();
332 processor
333 .process(
334 &mut (),
335 Event::Decision {
336 request: &mut request,
337 decision: &RoutedTo(target),
338 },
339 )
340 .await?;
341
342 assert_eq!(instructions(&request), vec![expected]);
343 assert!(
344 !request.llm_request.preserved_request_is_current(),
345 "{target}: preserved inbound body would replay without the tier prompt"
346 );
347 }
348 Ok(())
349 }
350
351 #[test]
354 fn note_drops_preserved_body() {
355 let mut request = request_with_preserved_body();
356 append_note(&mut request, NOTE);
357 assert!(
358 !request.llm_request.preserved_request_is_current(),
359 "preserved inbound body would replay without the note"
360 );
361 }
362
363 #[tokio::test]
365 async fn unprompted_target_keeps_preserved_body() -> Result<()> {
366 let processor = SystemPromptProcessor::new(prompts());
367 let mut request = request_with_preserved_body();
368 processor
369 .process(
370 &mut (),
371 Event::Decision {
372 request: &mut request,
373 decision: &RoutedTo("unconfigured"),
374 },
375 )
376 .await?;
377 assert!(
378 request.llm_request.preserved_request_is_current(),
379 "an untouched request keeps exact replay"
380 );
381 Ok(())
382 }
383}