1use async_trait::async_trait;
2use std::sync::Arc;
3
4use ai_agents_core::{
5 ChatMessage, LLMCapability, LLMChunk, LLMConfig, LLMError, LLMFeature, LLMProvider,
6 LLMResponse, LLMToolRequest, TaskContext, ToolChoice, ToolSelection,
7};
8
9use super::capability::DefaultLLMCapability;
10
11#[derive(Clone)]
12pub struct MultiLLMRouter {
13 primary: Arc<dyn LLMProvider>,
14 tool_selector: Option<Arc<dyn LLMProvider>>,
15 guard_evaluator: Option<Arc<dyn LLMProvider>>,
16 classifier: Option<Arc<dyn LLMProvider>>,
17 enable_fallback: bool,
18}
19
20impl MultiLLMRouter {
21 pub fn new(primary: Arc<dyn LLMProvider>) -> Self {
22 Self {
23 primary,
24 tool_selector: None,
25 guard_evaluator: None,
26 classifier: None,
27 enable_fallback: true,
28 }
29 }
30
31 pub fn with_tool_selector(mut self, provider: Arc<dyn LLMProvider>) -> Self {
32 self.tool_selector = Some(provider);
33 self
34 }
35
36 pub fn with_guard_evaluator(mut self, provider: Arc<dyn LLMProvider>) -> Self {
37 self.guard_evaluator = Some(provider);
38 self
39 }
40
41 pub fn with_classifier(mut self, provider: Arc<dyn LLMProvider>) -> Self {
42 self.classifier = Some(provider);
43 self
44 }
45
46 pub fn with_fallback(mut self, enable: bool) -> Self {
47 self.enable_fallback = enable;
48 self
49 }
50
51 fn get_tool_selector(&self) -> Arc<dyn LLMProvider> {
52 self.tool_selector
53 .as_ref()
54 .cloned()
55 .unwrap_or_else(|| self.primary.clone())
56 }
57
58 fn get_guard_evaluator(&self) -> Arc<dyn LLMProvider> {
59 self.guard_evaluator
60 .as_ref()
61 .cloned()
62 .unwrap_or_else(|| self.primary.clone())
63 }
64
65 fn get_classifier(&self) -> Arc<dyn LLMProvider> {
66 self.classifier
67 .as_ref()
68 .cloned()
69 .unwrap_or_else(|| self.primary.clone())
70 }
71
72 #[allow(dead_code)]
73 async fn execute_with_fallback<F, Fut, T>(
74 &self,
75 primary_fn: F,
76 _fallback_provider: Arc<dyn LLMProvider>,
77 operation: &str,
78 ) -> Result<T, LLMError>
79 where
80 F: FnOnce() -> Fut,
81 Fut: std::future::Future<Output = Result<T, LLMError>>,
82 {
83 match primary_fn().await {
84 Ok(result) => Ok(result),
85 Err(e) if self.enable_fallback => {
86 eprintln!(
87 "Multi-LLM: {} failed with specialized provider, falling back to primary: {}",
88 operation, e
89 );
90 Err(e)
91 }
92 Err(e) => Err(e),
93 }
94 }
95}
96
97#[async_trait]
98impl LLMProvider for MultiLLMRouter {
99 async fn complete(
100 &self,
101 messages: &[ChatMessage],
102 config: Option<&LLMConfig>,
103 ) -> Result<LLMResponse, LLMError> {
104 self.primary.complete(messages, config).await
105 }
106
107 async fn complete_with_tools(
108 &self,
109 messages: &[ChatMessage],
110 config: Option<&LLMConfig>,
111 request: &LLMToolRequest,
112 ) -> Result<LLMResponse, LLMError> {
113 self.primary
114 .complete_with_tools(messages, config, request)
115 .await
116 }
117
118 fn configured_tool_choice(&self) -> Option<ToolChoice> {
119 self.primary.configured_tool_choice()
120 }
121
122 fn supports_tool_choice(&self, choice: &ToolChoice) -> bool {
123 self.primary.supports_tool_choice(choice)
124 }
125
126 async fn complete_stream(
127 &self,
128 messages: &[ChatMessage],
129 config: Option<&LLMConfig>,
130 ) -> Result<Box<dyn futures::Stream<Item = Result<LLMChunk, LLMError>> + Unpin + Send>, LLMError>
131 {
132 self.primary.complete_stream(messages, config).await
133 }
134
135 fn provider_name(&self) -> &str {
136 "multi-llm-router"
137 }
138
139 fn supports(&self, feature: LLMFeature) -> bool {
140 self.primary.supports(feature)
141 }
142}
143
144#[async_trait]
145impl LLMCapability for MultiLLMRouter {
146 async fn select_tool(
147 &self,
148 context: &TaskContext,
149 user_input: &str,
150 ) -> Result<ToolSelection, LLMError> {
151 let provider = self.get_tool_selector();
152 let capability = DefaultLLMCapability::new(provider);
153 capability.select_tool(context, user_input).await
154 }
155
156 async fn generate_tool_args(
157 &self,
158 tool_id: &str,
159 user_input: &str,
160 schema: &serde_json::Value,
161 ) -> Result<serde_json::Value, LLMError> {
162 let provider = self.get_tool_selector();
163 let capability = DefaultLLMCapability::new(provider);
164 capability
165 .generate_tool_args(tool_id, user_input, schema)
166 .await
167 }
168
169 async fn evaluate_yesno(
170 &self,
171 question: &str,
172 context: &TaskContext,
173 ) -> Result<(bool, String), LLMError> {
174 let provider = self.get_guard_evaluator();
175 let capability = DefaultLLMCapability::new(provider);
176 capability.evaluate_yesno(question, context).await
177 }
178
179 async fn classify(
180 &self,
181 input: &str,
182 categories: &[String],
183 ) -> Result<(String, f32), LLMError> {
184 let provider = self.get_classifier();
185 let capability = DefaultLLMCapability::new(provider);
186 capability.classify(input, categories).await
187 }
188
189 async fn process_task(
190 &self,
191 context: &TaskContext,
192 system_prompt: &str,
193 ) -> Result<LLMResponse, LLMError> {
194 let capability = DefaultLLMCapability::new(self.primary.clone());
195 capability.process_task(context, system_prompt).await
196 }
197}
198
199#[cfg(test)]
200mod tests {
201 use super::*;
202 use crate::mock::MockLLMProvider;
203 use ai_agents_core::{FinishReason, LLMToolDefinition, LLMToolRequest, Role, ToolChoice};
204 use std::collections::HashMap;
205
206 #[tokio::test]
207 async fn test_router_with_primary_only() {
208 let mut primary = MockLLMProvider::new("primary");
209 primary.add_response(LLMResponse::new("Hello from primary", FinishReason::Stop));
210
211 let router = MultiLLMRouter::new(Arc::new(primary));
212
213 let messages = vec![ChatMessage {
214 timestamp: None,
215 role: Role::User,
216 content: "Test".to_string(),
217 name: None,
218 }];
219
220 let response = router.complete(&messages, None).await.unwrap();
221 assert_eq!(response.content, "Hello from primary");
222 }
223
224 #[tokio::test]
225 async fn test_router_delegates_native_tool_methods() {
226 let mut primary = MockLLMProvider::new("primary");
227 primary.add_response(LLMResponse::new("No call", FinishReason::Stop));
228 let history = primary.clone();
229 let router = MultiLLMRouter::new(Arc::new(primary));
230 let request = LLMToolRequest {
231 tools: vec![LLMToolDefinition {
232 name: "calculator".to_string(),
233 description: "Calculate an expression".to_string(),
234 input_schema: serde_json::json!({"type": "object"}),
235 }],
236 choice: ToolChoice::Auto,
237 };
238
239 router
240 .complete_with_tools(&[ChatMessage::user("Hello")], None, &request)
241 .await
242 .unwrap();
243
244 assert!(router.supports_tool_choice(&ToolChoice::Required));
245 assert_eq!(history.last_call().unwrap().request, Some(request));
246 }
247
248 #[tokio::test]
249 async fn test_router_with_specialized_providers() {
250 let mut primary = MockLLMProvider::new("primary");
251 primary.add_response(LLMResponse::new("Primary response", FinishReason::Stop));
252
253 let mut tool_selector = MockLLMProvider::new("tool-selector");
254 tool_selector.add_response(LLMResponse::new(
255 r#"{"tool_id": "calculator", "confidence": 0.9}"#,
256 FinishReason::Stop,
257 ));
258
259 let mut guard = MockLLMProvider::new("guard");
260 guard.add_response(LLMResponse::new(
261 r#"{"answer": true, "reasoning": "Approved"}"#,
262 FinishReason::Stop,
263 ));
264
265 let router = MultiLLMRouter::new(Arc::new(primary))
266 .with_tool_selector(Arc::new(tool_selector))
267 .with_guard_evaluator(Arc::new(guard));
268
269 let context = TaskContext {
270 current_state: None,
271 available_tools: vec!["calculator".to_string()],
272 memory_slots: HashMap::new(),
273 recent_messages: vec![],
274 };
275
276 let tool_selection = router.select_tool(&context, "Do math").await.unwrap();
277 assert_eq!(tool_selection.tool_id, "calculator");
278 assert_eq!(tool_selection.confidence, 0.9);
279
280 let (answer, reasoning) = router
281 .evaluate_yesno("Is it safe?", &context)
282 .await
283 .unwrap();
284 assert!(answer);
285 assert_eq!(reasoning, "Approved");
286 }
287
288 #[tokio::test]
289 async fn test_router_fallback_to_primary() {
290 let mut primary = MockLLMProvider::new("primary");
291 primary.add_response(LLMResponse::new("Primary response", FinishReason::Stop));
292
293 let router = MultiLLMRouter::new(Arc::new(primary)).with_fallback(true);
294
295 let messages = vec![ChatMessage {
296 timestamp: None,
297 role: Role::User,
298 content: "Test".to_string(),
299 name: None,
300 }];
301
302 let response = router.complete(&messages, None).await.unwrap();
303 assert_eq!(response.content, "Primary response");
304 }
305
306 #[tokio::test]
307 async fn test_router_provider_name() {
308 let primary = MockLLMProvider::new("primary");
309 let router = MultiLLMRouter::new(Arc::new(primary));
310
311 assert_eq!(router.provider_name(), "multi-llm-router");
312 }
313
314 #[tokio::test]
315 async fn test_router_supports() {
316 let mut primary = MockLLMProvider::new("primary");
317 primary.set_feature_support(LLMFeature::Streaming, true);
318
319 let router = MultiLLMRouter::new(Arc::new(primary));
320
321 assert!(router.supports(LLMFeature::Streaming));
322 }
323
324 #[tokio::test]
325 async fn test_classify_with_specialized_provider() {
326 let primary = MockLLMProvider::new("primary");
327
328 let mut classifier = MockLLMProvider::new("classifier");
329 classifier.add_response(LLMResponse::new(
330 r#"{"category": "greeting", "confidence": 0.95}"#,
331 FinishReason::Stop,
332 ));
333
334 let router = MultiLLMRouter::new(Arc::new(primary)).with_classifier(Arc::new(classifier));
335
336 let categories = vec!["greeting".to_string(), "question".to_string()];
337 let (category, confidence) = router.classify("Hello!", &categories).await.unwrap();
338
339 assert_eq!(category, "greeting");
340 assert_eq!(confidence, 0.95);
341 }
342
343 #[tokio::test]
344 async fn test_process_task_uses_primary() {
345 let mut primary = MockLLMProvider::new("primary");
346 primary.add_response(LLMResponse::new(
347 "Task processed by primary",
348 FinishReason::Stop,
349 ));
350
351 let tool_selector = MockLLMProvider::new("tool-selector");
352
353 let router =
354 MultiLLMRouter::new(Arc::new(primary)).with_tool_selector(Arc::new(tool_selector));
355
356 let context = TaskContext {
357 current_state: None,
358 available_tools: vec![],
359 memory_slots: HashMap::new(),
360 recent_messages: vec![],
361 };
362
363 let response = router
364 .process_task(&context, "System prompt")
365 .await
366 .unwrap();
367
368 assert_eq!(response.content, "Task processed by primary");
369 }
370
371 #[tokio::test]
372 async fn test_generate_tool_args_with_specialized() {
373 let primary = MockLLMProvider::new("primary");
374
375 let mut tool_selector = MockLLMProvider::new("tool-selector");
376 tool_selector.add_response(LLMResponse::new(
377 r#"{"expression": "2 + 2"}"#,
378 FinishReason::Stop,
379 ));
380
381 let router =
382 MultiLLMRouter::new(Arc::new(primary)).with_tool_selector(Arc::new(tool_selector));
383
384 let schema = serde_json::json!({
385 "type": "object",
386 "properties": {
387 "expression": {"type": "string"}
388 }
389 });
390
391 let result = router
392 .generate_tool_args("calculator", "Calculate 2 + 2", &schema)
393 .await
394 .unwrap();
395
396 assert_eq!(result["expression"], "2 + 2");
397 }
398
399 #[test]
400 fn test_builder_pattern() {
401 let primary = MockLLMProvider::new("primary");
402 let tool_selector = MockLLMProvider::new("tool-selector");
403 let guard = MockLLMProvider::new("guard");
404 let classifier = MockLLMProvider::new("classifier");
405
406 let router = MultiLLMRouter::new(Arc::new(primary))
407 .with_tool_selector(Arc::new(tool_selector))
408 .with_guard_evaluator(Arc::new(guard))
409 .with_classifier(Arc::new(classifier))
410 .with_fallback(false);
411
412 assert_eq!(router.provider_name(), "multi-llm-router");
413 assert!(!router.enable_fallback);
414 }
415}