use std::collections::BTreeMap;
use async_trait::async_trait;
use switchyard_protocol::{ContentBlock, InstructionBlock, Message, Request, Role};
use crate::Result;
use crate::core::processor::{Event, Processor};
pub fn append_note(request: &mut Request, note: &str) {
match request.llm_request.messages.last_mut() {
Some(last) if last.role == Role::User => last.content.push(ContentBlock::Text {
text: note.to_string(),
}),
_ => request
.llm_request
.messages
.push(Message::text(Role::User, note)),
}
drop_exact_replay(request);
}
fn drop_exact_replay(request: &mut Request) {
request.llm_request.preservation.requests.clear();
}
#[derive(Clone, Debug, Default)]
pub struct TargetPrompts {
by_target: BTreeMap<String, String>,
}
impl TargetPrompts {
pub fn with(mut self, target: impl Into<String>, prompt: impl Into<String>) -> Self {
self.by_target.insert(target.into(), prompt.into());
self
}
pub fn get(&self, target: &str) -> Option<&str> {
self.by_target.get(target).map(String::as_str)
}
pub fn is_empty(&self) -> bool {
self.by_target.is_empty()
}
}
pub struct SystemPromptProcessor {
prompts: TargetPrompts,
}
impl SystemPromptProcessor {
pub fn new(prompts: TargetPrompts) -> Self {
Self { prompts }
}
}
#[async_trait]
impl<S: Send> Processor<S> for SystemPromptProcessor {
async fn process(&self, _state: &mut S, event: Event<'_>) -> Result<()> {
let Event::Decision { request, decision } = event else {
return Ok(());
};
let Some(prompt) = self.prompts.get(decision.selected_model()) else {
return Ok(());
};
request.llm_request.instructions.insert(
0,
InstructionBlock {
role: Role::System,
content: vec![ContentBlock::Text {
text: prompt.to_string(),
}],
},
);
drop_exact_replay(request);
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
use switchyard_protocol::{LlmRequest, ToolResult, text_request};
const NOTE: &str = "recovering from an error";
const STRONG_PROMPT: &str = "diagnose before you edit";
const WEAK_PROMPT: &str = "follow the settled plan";
fn request_with(messages: Vec<Message>) -> Request {
Request {
llm_request: LlmRequest {
messages,
preservation: preserved_body(),
..LlmRequest::default()
},
raw_request: None,
metadata: None,
}
}
fn preserved_body() -> switchyard_protocol::PreservationMetadata {
let mut preservation = switchyard_protocol::PreservationMetadata::default();
preservation.requests.insert(
"openai_chat".into(),
serde_json::json!({"model": "weak", "messages": [{"role": "user", "content": "hi"}]}),
);
preservation
}
fn replays_exactly(request: &Request) -> bool {
!request.llm_request.preservation.requests.is_empty()
}
#[test]
fn a_note_joins_a_trailing_user_turn_after_its_tool_result() {
let tool_result = ContentBlock::ToolResult(ToolResult {
tool_call_id: "call_1".to_string(),
content: vec![ContentBlock::Text {
text: "exit 1".to_string(),
}],
is_error: Some(true),
});
let mut request = request_with(vec![Message {
role: Role::User,
content: vec![tool_result.clone()],
}]);
append_note(&mut request, NOTE);
let messages = &request.llm_request.messages;
assert_eq!(messages.len(), 1, "no second consecutive user turn");
assert_eq!(
messages[0].content,
vec![
tool_result,
ContentBlock::Text {
text: NOTE.to_string()
}
]
);
}
#[test]
fn a_note_opens_a_user_turn_after_an_assistant_turn() {
let mut request = request_with(vec![Message::text(Role::Assistant, "done")]);
append_note(&mut request, NOTE);
let messages = &request.llm_request.messages;
assert_eq!(messages.len(), 2);
assert_eq!(messages[1].role, Role::User);
assert_eq!(messages[1].text_content(""), Some(NOTE.to_string()));
}
#[test]
fn a_note_leaves_the_rest_of_the_conversation_untouched() {
let mut request = Request {
llm_request: LlmRequest {
preservation: preserved_body(),
..text_request(Some("auto".to_string()), "fix the build")
},
raw_request: None,
metadata: None,
};
append_note(&mut request, NOTE);
let trail: Vec<String> = request
.llm_request
.messages
.iter()
.filter_map(|message| message.text_content("|"))
.collect();
assert_eq!(trail, vec![format!("fix the build|{NOTE}")]);
assert!(
!replays_exactly(&request),
"a same-format hop would replay the body captured before the note"
);
}
struct RoutedTo(&'static str);
impl switchyard_protocol::Decision for RoutedTo {
fn selected_model(&self) -> &str {
self.0
}
fn reasoning(&self) -> Option<&str> {
None
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
}
fn instructions(request: &Request) -> Vec<String> {
request
.llm_request
.instructions
.iter()
.filter_map(|block| {
block.content.iter().find_map(|content| match content {
ContentBlock::Text { text } => Some(text.clone()),
_ => None,
})
})
.collect()
}
async fn run(processor: &SystemPromptProcessor, target: &'static str) -> Result<Request> {
let mut request = Request {
llm_request: LlmRequest {
preservation: preserved_body(),
..LlmRequest::default()
},
..Request::default()
};
processor
.process(
&mut (),
Event::Decision {
request: &mut request,
decision: &RoutedTo(target),
},
)
.await?;
Ok(request)
}
fn prompts() -> TargetPrompts {
TargetPrompts::default()
.with("strong", STRONG_PROMPT)
.with("weak", WEAK_PROMPT)
}
#[tokio::test]
async fn each_target_gets_its_own_prompt() -> Result<()> {
let processor = SystemPromptProcessor::new(prompts());
for (target, expected) in [("strong", STRONG_PROMPT), ("weak", WEAK_PROMPT)] {
let request = run(&processor, target).await?;
assert_eq!(instructions(&request), vec![expected]);
assert!(
!replays_exactly(&request),
"{target}: a same-format hop would replay the body captured before the prompt"
);
}
Ok(())
}
#[tokio::test]
async fn an_unconfigured_target_is_left_untouched() -> Result<()> {
let processor =
SystemPromptProcessor::new(TargetPrompts::default().with("strong", STRONG_PROMPT));
assert_eq!(
instructions(&run(&processor, "strong").await?),
vec![STRONG_PROMPT]
);
let untouched = run(&processor, "weak").await?;
assert!(instructions(&untouched).is_empty());
assert!(
replays_exactly(&untouched),
"an untouched request must keep its lossless same-format replay"
);
Ok(())
}
#[tokio::test]
async fn the_prompt_leads_the_client_instructions() -> Result<()> {
let processor = SystemPromptProcessor::new(prompts());
let mut request = Request::default();
request.llm_request.instructions.push(InstructionBlock {
role: Role::System,
content: vec![ContentBlock::Text {
text: "you are a coding agent".to_string(),
}],
});
processor
.process(
&mut (),
Event::Decision {
request: &mut request,
decision: &RoutedTo("strong"),
},
)
.await?;
assert_eq!(
instructions(&request),
vec![STRONG_PROMPT, "you are a coding agent"]
);
Ok(())
}
#[tokio::test]
async fn the_inbound_request_is_left_alone() -> Result<()> {
let processor = SystemPromptProcessor::new(prompts());
let mut request = Request::default();
processor
.process(&mut (), Event::Request(&mut request))
.await?;
assert!(instructions(&request).is_empty());
Ok(())
}
}