switchyard_libsy/core/
processor.rs1use crate::Result;
5use crate::core::algorithm::Driver;
6use async_trait::async_trait;
7use switchyard_protocol::{AggLlmResponse, Category, ModelId, Request};
8
9pub enum Event<'a> {
15 Request {
17 request: &'a mut Request,
19 driver: &'a Driver,
21 },
22 Decision {
27 request: &'a mut Request,
29 selected_model_id: &'a ModelId,
31 category: Option<Category>,
34 driver: &'a Driver,
36 },
37 ModelResponse(&'a AggLlmResponse),
39}
40
41#[async_trait]
43pub trait Processor<S = ()>: Send + Sync {
44 async fn process(&self, state: &mut S, event: Event<'_>) -> Result<()>;
47}
48
49#[cfg(test)]
50mod tests {
51 use super::*;
52 use crate::core::testing::empty_driver;
53 use std::collections::HashMap;
54 use switchyard_protocol::{text_request, text_response};
55
56 type TestState = HashMap<&'static str, u32>;
57
58 fn event_key(event: &Event<'_>) -> &'static str {
60 match event {
61 Event::Request { .. } => "requests",
62 Event::Decision { .. } => "decisions",
63 Event::ModelResponse(_) => "model_responses",
64 }
65 }
66
67 fn count(state: &TestState, key: &'static str) -> u32 {
69 state.get(key).copied().unwrap_or_default()
70 }
71
72 struct CountingProcessor;
74
75 #[async_trait]
76 impl Processor<TestState> for CountingProcessor {
77 async fn process(&self, state: &mut TestState, event: Event<'_>) -> Result<()> {
78 *state.entry(event_key(&event)).or_default() += 1;
79 Ok(())
80 }
81 }
82
83 fn request() -> Request {
84 Request {
85 llm_request: text_request(Some("auto".to_string()), "hi"),
86 raw_request: None,
87 metadata: None,
88 }
89 }
90
91 #[tokio::test]
92 async fn processor_tallies_each_event_variant_into_state() -> Result<()> {
93 let processor = CountingProcessor;
94 let mut state = TestState::default();
95 let mut req = request();
96 let response = text_response(None, "ok");
97 let selected_model_id = ModelId::from("test/model");
98 processor
100 .process(
101 &mut state,
102 Event::Request {
103 request: &mut req,
104 driver: &empty_driver(),
105 },
106 )
107 .await?;
108 processor
109 .process(&mut state, Event::ModelResponse(&response))
110 .await?;
111 processor
112 .process(
113 &mut state,
114 Event::Decision {
115 request: &mut req,
116 selected_model_id: &selected_model_id,
117 category: None,
118 driver: &empty_driver(),
119 },
120 )
121 .await?;
122 assert_eq!(count(&state, "requests"), 1);
123 assert_eq!(count(&state, "decisions"), 1);
124 assert_eq!(count(&state, "model_responses"), 1);
125 Ok(())
126 }
127
128 #[tokio::test]
129 async fn process_accumulates_state_across_repeated_events() -> Result<()> {
130 let processor = CountingProcessor;
131 let mut state = TestState::default();
132 let mut req = request();
133
134 for _ in 0..3 {
135 processor
136 .process(
137 &mut state,
138 Event::Request {
139 request: &mut req,
140 driver: &empty_driver(),
141 },
142 )
143 .await?;
144 }
145
146 assert_eq!(count(&state, "requests"), 3);
147 Ok(())
148 }
149
150 struct RewritingProcessor;
152
153 #[async_trait]
154 impl Processor for RewritingProcessor {
155 async fn process(&self, _state: &mut (), event: Event<'_>) -> Result<()> {
156 match event {
157 Event::Request { request, .. } | Event::Decision { request, .. } => {
158 request.llm_request.model = Some("rewritten".to_string());
159 }
160 _ => {}
161 }
162 Ok(())
163 }
164 }
165
166 #[tokio::test]
167 async fn processor_rewrites_the_request_in_place() -> Result<()> {
168 let mut state = ();
169 let mut req = request();
170 assert_eq!(req.model_id(), Some("auto".into()));
171
172 RewritingProcessor
173 .process(
174 &mut state,
175 Event::Request {
176 request: &mut req,
177 driver: &empty_driver(),
178 },
179 )
180 .await?;
181
182 assert_eq!(req.model_id(), Some("rewritten".into()));
184 Ok(())
185 }
186}