ironflow_core/providers/
router.rs1use std::sync::Arc;
26
27use tracing::debug;
28
29use crate::provider::{AgentConfig, AgentProvider, InvokeFuture, LogSink};
30
31#[derive(Debug, Clone)]
33pub enum ProviderMatcher {
34 ModelPrefix(String),
38 ModelExact(String),
42}
43
44impl ProviderMatcher {
45 fn matches(&self, config: &AgentConfig) -> bool {
46 match self {
47 Self::ModelPrefix(prefix) => config.model.starts_with(prefix.as_str()),
48 Self::ModelExact(exact) => config.model == *exact,
49 }
50 }
51}
52
53pub struct ProviderRouter {
77 routes: Vec<(ProviderMatcher, Arc<dyn AgentProvider>)>,
78 fallback: Arc<dyn AgentProvider>,
79}
80
81impl ProviderRouter {
82 pub fn new(fallback: Arc<dyn AgentProvider>) -> Self {
84 Self {
85 routes: Vec::new(),
86 fallback,
87 }
88 }
89
90 pub fn route(mut self, matcher: ProviderMatcher, provider: Arc<dyn AgentProvider>) -> Self {
92 self.routes.push((matcher, provider));
93 self
94 }
95
96 fn resolve(&self, config: &AgentConfig) -> &Arc<dyn AgentProvider> {
98 for (matcher, provider) in &self.routes {
99 if matcher.matches(config) {
100 debug!(
101 model = %config.model,
102 matcher = ?matcher,
103 "routed to matched provider"
104 );
105 return provider;
106 }
107 }
108 debug!(model = %config.model, "using fallback provider");
109 &self.fallback
110 }
111}
112
113impl AgentProvider for ProviderRouter {
114 fn invoke<'a>(&'a self, config: &'a AgentConfig) -> InvokeFuture<'a> {
115 let provider = self.resolve(config);
116 provider.invoke(config)
117 }
118
119 fn invoke_with_logs<'a>(
120 &'a self,
121 config: &'a AgentConfig,
122 log_sink: Arc<dyn LogSink>,
123 ) -> InvokeFuture<'a> {
124 let provider = self.resolve(config);
125 provider.invoke_with_logs(config, log_sink)
126 }
127}
128
129#[cfg(test)]
130mod tests {
131 use std::sync::atomic::{AtomicUsize, Ordering};
132
133 use serde_json::json;
134
135 use super::*;
136 use crate::provider::AgentOutput;
137
138 struct CountingProvider {
139 name: &'static str,
140 count: AtomicUsize,
141 }
142
143 impl CountingProvider {
144 fn new(name: &'static str) -> Arc<Self> {
145 Arc::new(Self {
146 name,
147 count: AtomicUsize::new(0),
148 })
149 }
150
151 fn call_count(&self) -> usize {
152 self.count.load(Ordering::Relaxed)
153 }
154 }
155
156 impl AgentProvider for CountingProvider {
157 fn invoke<'a>(&'a self, _config: &'a AgentConfig) -> InvokeFuture<'a> {
158 self.count.fetch_add(1, Ordering::Relaxed);
159 let name = self.name;
160 Box::pin(async move { Ok(AgentOutput::new(json!(name))) })
161 }
162 }
163
164 #[tokio::test]
165 async fn router_fallback_when_no_routes() {
166 let fallback = CountingProvider::new("fallback");
167 let router = ProviderRouter::new(fallback.clone());
168
169 let config = AgentConfig::new("hello");
170 let output = router.invoke(&config).await.expect("should succeed");
171 assert_eq!(output.value, json!("fallback"));
172 assert_eq!(fallback.call_count(), 1);
173 }
174
175 #[tokio::test]
176 async fn router_matches_model_prefix() {
177 let fallback = CountingProvider::new("fallback");
178 let nvidia = CountingProvider::new("nvidia");
179
180 let router = ProviderRouter::new(fallback.clone()).route(
181 ProviderMatcher::ModelPrefix("nvidia/".into()),
182 nvidia.clone(),
183 );
184
185 let config = AgentConfig::new("hello").model("nvidia/deepseek-v4-flash");
186 let output = router.invoke(&config).await.expect("should succeed");
187 assert_eq!(output.value, json!("nvidia"));
188 assert_eq!(nvidia.call_count(), 1);
189 assert_eq!(fallback.call_count(), 0);
190 }
191
192 #[tokio::test]
193 async fn router_matches_model_exact() {
194 let fallback = CountingProvider::new("fallback");
195 let special = CountingProvider::new("special");
196
197 let router = ProviderRouter::new(fallback.clone()).route(
198 ProviderMatcher::ModelExact("my-model".into()),
199 special.clone(),
200 );
201
202 let config = AgentConfig::new("hello").model("my-model");
203 let output = router.invoke(&config).await.expect("should succeed");
204 assert_eq!(output.value, json!("special"));
205 assert_eq!(special.call_count(), 1);
206 }
207
208 #[tokio::test]
209 async fn router_exact_does_not_match_prefix() {
210 let fallback = CountingProvider::new("fallback");
211 let special = CountingProvider::new("special");
212
213 let router = ProviderRouter::new(fallback.clone()).route(
214 ProviderMatcher::ModelExact("nvidia".into()),
215 special.clone(),
216 );
217
218 let config = AgentConfig::new("hello").model("nvidia/something");
219 let output = router.invoke(&config).await.expect("should succeed");
220 assert_eq!(output.value, json!("fallback"));
221 assert_eq!(special.call_count(), 0);
222 assert_eq!(fallback.call_count(), 1);
223 }
224
225 #[tokio::test]
226 async fn router_first_match_wins() {
227 let fallback = CountingProvider::new("fallback");
228 let first = CountingProvider::new("first");
229 let second = CountingProvider::new("second");
230
231 let router = ProviderRouter::new(fallback.clone())
232 .route(
233 ProviderMatcher::ModelPrefix("nvidia/".into()),
234 first.clone(),
235 )
236 .route(
237 ProviderMatcher::ModelPrefix("nvidia/".into()),
238 second.clone(),
239 );
240
241 let config = AgentConfig::new("hello").model("nvidia/test");
242 let output = router.invoke(&config).await.expect("should succeed");
243 assert_eq!(output.value, json!("first"));
244 assert_eq!(first.call_count(), 1);
245 assert_eq!(second.call_count(), 0);
246 }
247
248 #[tokio::test]
249 async fn router_multiple_routes() {
250 let fallback = CountingProvider::new("claude");
251 let nvidia = CountingProvider::new("nvidia");
252 let openai = CountingProvider::new("openai");
253
254 let router = ProviderRouter::new(fallback.clone())
255 .route(
256 ProviderMatcher::ModelPrefix("nvidia/".into()),
257 nvidia.clone(),
258 )
259 .route(ProviderMatcher::ModelPrefix("gpt-".into()), openai.clone());
260
261 let config1 = AgentConfig::new("hello").model("nvidia/nemotron");
262 let config2 = AgentConfig::new("hello").model("gpt-5.5");
263 let config3 = AgentConfig::new("hello").model("sonnet");
264
265 let out1 = router.invoke(&config1).await.expect("should succeed");
266 let out2 = router.invoke(&config2).await.expect("should succeed");
267 let out3 = router.invoke(&config3).await.expect("should succeed");
268
269 assert_eq!(out1.value, json!("nvidia"));
270 assert_eq!(out2.value, json!("openai"));
271 assert_eq!(out3.value, json!("claude"));
272 }
273}