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