Skip to main content

ironflow_core/providers/
router.rs

1//! Provider router for multi-provider workflows.
2//!
3//! [`ProviderRouter`] implements [`AgentProvider`] and dispatches invocations
4//! to different providers based on the model name or other criteria.
5//!
6//! # Examples
7//!
8//! ```no_run
9//! use std::sync::Arc;
10//! use ironflow_core::providers::router::{ProviderRouter, ProviderMatcher};
11//! use ironflow_core::providers::claude::ClaudeCodeProvider;
12//! use ironflow_core::provider::AgentProvider;
13//!
14//! let claude = Arc::new(ClaudeCodeProvider::new());
15//! // let nvidia = Arc::new(nvidia_provider);
16//!
17//! let router = ProviderRouter::new(claude.clone())
18//!     // .route(ProviderMatcher::ModelPrefix("nvidia/".into()), nvidia)
19//!     ;
20//!
21//! // router implements AgentProvider, pass it to Engine as usual
22//! let provider: Arc<dyn AgentProvider> = Arc::new(router);
23//! ```
24
25use std::sync::Arc;
26
27use tracing::debug;
28
29use crate::provider::{AgentConfig, AgentProvider, InvokeFuture, LogSink};
30
31/// Matching strategy for routing invocations to providers.
32#[derive(Debug, Clone)]
33pub enum ProviderMatcher {
34    /// Match when the model starts with the given prefix.
35    ///
36    /// Example: `ModelPrefix("nvidia/".into())` matches `"nvidia/deepseek-v4-flash"`.
37    ModelPrefix(String),
38    /// Match when the model is exactly the given string.
39    ///
40    /// Example: `ModelExact("sonnet".into())` matches only `"sonnet"`.
41    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
53/// Routes agent invocations to different providers based on model/config matching.
54///
55/// Evaluates matchers in registration order; first match wins. If no matcher
56/// matches, the fallback provider handles the request.
57///
58/// # Examples
59///
60/// ```no_run
61/// use std::sync::Arc;
62/// use ironflow_core::providers::router::{ProviderRouter, ProviderMatcher};
63/// use ironflow_core::providers::claude::ClaudeCodeProvider;
64/// use ironflow_core::provider::{AgentConfig, AgentProvider};
65///
66/// # async fn example() -> Result<(), ironflow_core::error::AgentError> {
67/// let claude = Arc::new(ClaudeCodeProvider::new());
68/// let router = ProviderRouter::new(claude.clone());
69///
70/// // Uses the fallback (claude) since no routes match "sonnet"
71/// let config = AgentConfig::new("hello");
72/// let output = router.invoke(&config).await?;
73/// # Ok(())
74/// # }
75/// ```
76pub struct ProviderRouter {
77    routes: Vec<(ProviderMatcher, Arc<dyn AgentProvider>)>,
78    fallback: Arc<dyn AgentProvider>,
79}
80
81impl ProviderRouter {
82    /// Create a router with a fallback provider for unmatched models.
83    pub fn new(fallback: Arc<dyn AgentProvider>) -> Self {
84        Self {
85            routes: Vec::new(),
86            fallback,
87        }
88    }
89
90    /// Add a routing rule. Routes are evaluated in order; first match wins.
91    pub fn route(mut self, matcher: ProviderMatcher, provider: Arc<dyn AgentProvider>) -> Self {
92        self.routes.push((matcher, provider));
93        self
94    }
95
96    /// Resolve which provider handles a given config.
97    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}