use crate::Result;
use async_trait::async_trait;
use switchyard_protocol::{AggLlmResponse, Decision, Request, Signals};
pub enum Event<'a> {
Request(&'a mut Request),
Signal(&'a Signals),
Decision {
request: &'a mut Request,
decision: &'a dyn Decision,
},
ModelResponse(&'a AggLlmResponse),
}
#[async_trait]
pub trait Processor<S = ()>: Send + Sync {
async fn process(&self, state: &mut S, event: Event<'_>) -> Result<()>;
}
#[cfg(test)]
mod tests {
use super::*;
use std::collections::HashMap;
use switchyard_protocol::{text_request, text_response};
type TestState = HashMap<&'static str, u32>;
fn event_key(event: &Event<'_>) -> &'static str {
match event {
Event::Request(_) => "requests",
Event::Signal(_) => "signals",
Event::Decision { .. } => "decisions",
Event::ModelResponse(_) => "model_responses",
}
}
fn count(state: &TestState, key: &'static str) -> u32 {
state.get(key).copied().unwrap_or_default()
}
struct CountingProcessor;
#[async_trait]
impl Processor<TestState> for CountingProcessor {
async fn process(&self, state: &mut TestState, event: Event<'_>) -> Result<()> {
*state.entry(event_key(&event)).or_default() += 1;
Ok(())
}
}
struct TestDecision;
impl Decision for TestDecision {
fn selected_model(&self) -> &str {
"test/model"
}
fn reasoning(&self) -> Option<&str> {
None
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
}
fn request() -> Request {
Request {
llm_request: text_request(Some("auto".to_string()), "hi"),
raw_request: None,
metadata: None,
}
}
#[tokio::test]
async fn processor_tallies_each_event_variant_into_state() -> Result<()> {
let processor = CountingProcessor;
let mut state = TestState::default();
let mut req = request();
let response = text_response(None, "ok");
let decision = TestDecision;
let signals = Signals {};
processor
.process(&mut state, Event::Request(&mut req))
.await?;
processor
.process(&mut state, Event::ModelResponse(&response))
.await?;
processor
.process(
&mut state,
Event::Decision {
request: &mut req,
decision: &decision,
},
)
.await?;
processor
.process(&mut state, Event::Signal(&signals))
.await?;
assert_eq!(count(&state, "requests"), 1);
assert_eq!(count(&state, "signals"), 1);
assert_eq!(count(&state, "decisions"), 1);
assert_eq!(count(&state, "model_responses"), 1);
Ok(())
}
#[tokio::test]
async fn process_accumulates_state_across_repeated_events() -> Result<()> {
let processor = CountingProcessor;
let mut state = TestState::default();
let mut req = request();
for _ in 0..3 {
processor
.process(&mut state, Event::Request(&mut req))
.await?;
}
assert_eq!(count(&state, "requests"), 3);
Ok(())
}
struct RewritingProcessor;
#[async_trait]
impl Processor for RewritingProcessor {
async fn process(&self, _state: &mut (), event: Event<'_>) -> Result<()> {
match event {
Event::Request(request) | Event::Decision { request, .. } => {
request.llm_request.model = Some("rewritten".to_string());
}
_ => {}
}
Ok(())
}
}
#[tokio::test]
async fn processor_rewrites_the_request_in_place() -> Result<()> {
let mut state = ();
let mut req = request();
assert_eq!(req.requested_model(), Some("auto"));
RewritingProcessor
.process(&mut state, Event::Request(&mut req))
.await?;
assert_eq!(req.requested_model(), Some("rewritten"));
Ok(())
}
}