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