switchyard_libsy/core/
processor.rs1use crate::Result;
5use async_trait::async_trait;
6use switchyard_protocol::{AggLlmResponse, Decision, Request, Signals};
7
8pub enum Event<'a> {
14 Request(&'a mut Request),
16 Signal(&'a Signals),
18 Decision {
23 request: &'a mut Request,
25 decision: &'a Decision,
27 },
28 ModelResponse(&'a AggLlmResponse),
30}
31
32#[async_trait]
34pub trait Processor<S = ()>: Send + Sync {
35 async fn process(&self, state: &mut S, event: Event<'_>) -> Result<()>;
38}
39
40#[cfg(test)]
41mod tests {
42 use super::*;
43 use std::collections::HashMap;
44 use switchyard_protocol::{text_request, text_response};
45
46 type TestState = HashMap<&'static str, u32>;
47
48 fn event_key(event: &Event<'_>) -> &'static str {
50 match event {
51 Event::Request(_) => "requests",
52 Event::Signal(_) => "signals",
53 Event::Decision { .. } => "decisions",
54 Event::ModelResponse(_) => "model_responses",
55 }
56 }
57
58 fn count(state: &TestState, key: &'static str) -> u32 {
60 state.get(key).copied().unwrap_or_default()
61 }
62
63 struct CountingProcessor;
65
66 #[async_trait]
67 impl Processor<TestState> for CountingProcessor {
68 async fn process(&self, state: &mut TestState, event: Event<'_>) -> Result<()> {
69 *state.entry(event_key(&event)).or_default() += 1;
70 Ok(())
71 }
72 }
73
74 fn request() -> Request {
75 Request {
76 llm_request: text_request(Some("auto".to_string()), "hi"),
77 raw_request: None,
78 metadata: None,
79 }
80 }
81
82 #[tokio::test]
83 async fn processor_tallies_each_event_variant_into_state() -> Result<()> {
84 let processor = CountingProcessor;
85 let mut state = TestState::default();
86 let mut req = request();
87 let response = text_response(None, "ok");
88 let decision = Decision::new("test/model", None, true);
89 let signals = Signals {};
90
91 processor
93 .process(&mut state, Event::Request(&mut req))
94 .await?;
95 processor
96 .process(&mut state, Event::ModelResponse(&response))
97 .await?;
98 processor
99 .process(
100 &mut state,
101 Event::Decision {
102 request: &mut req,
103 decision: &decision,
104 },
105 )
106 .await?;
107 processor
108 .process(&mut state, Event::Signal(&signals))
109 .await?;
110
111 assert_eq!(count(&state, "requests"), 1);
112 assert_eq!(count(&state, "signals"), 1);
113 assert_eq!(count(&state, "decisions"), 1);
114 assert_eq!(count(&state, "model_responses"), 1);
115 Ok(())
116 }
117
118 #[tokio::test]
119 async fn process_accumulates_state_across_repeated_events() -> Result<()> {
120 let processor = CountingProcessor;
121 let mut state = TestState::default();
122 let mut req = request();
123
124 for _ in 0..3 {
125 processor
126 .process(&mut state, Event::Request(&mut req))
127 .await?;
128 }
129
130 assert_eq!(count(&state, "requests"), 3);
131 Ok(())
132 }
133
134 struct RewritingProcessor;
136
137 #[async_trait]
138 impl Processor for RewritingProcessor {
139 async fn process(&self, _state: &mut (), event: Event<'_>) -> Result<()> {
140 match event {
141 Event::Request(request) | Event::Decision { request, .. } => {
142 request.llm_request.model = Some("rewritten".to_string());
143 }
144 _ => {}
145 }
146 Ok(())
147 }
148 }
149
150 #[tokio::test]
151 async fn processor_rewrites_the_request_in_place() -> Result<()> {
152 let mut state = ();
153 let mut req = request();
154 assert_eq!(req.requested_model(), Some("auto"));
155
156 RewritingProcessor
157 .process(&mut state, Event::Request(&mut req))
158 .await?;
159
160 assert_eq!(req.requested_model(), Some("rewritten"));
162 Ok(())
163 }
164}