Skip to main content

ai_agents_llm/
multi.rs

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}