Skip to main content

llm_trait/
provider.rs

1//! The unified LLM provider trait.
2
3use async_trait::async_trait;
4
5use super::capabilities::{Capabilities, ProviderInfo};
6use super::error::LlmError;
7use super::request::ChatRequest;
8use super::response::{ChatResponse, ChatStream};
9
10/// LLM Provider unified interface.
11///
12/// The core trait exposed by the framework. All providers implement this.
13/// Supports both streaming and non-streaming call modes.
14///
15/// # Object Safety
16///
17/// This trait is object-safe, so `Arc<dyn LlmProvider>` works.
18#[async_trait]
19pub trait LlmProvider: Send + Sync {
20    /// Streaming call (returns chunks in real-time).
21    ///
22    /// Use for: chat interfaces, real-time output, long text generation.
23    async fn stream(&self, request: ChatRequest) -> Result<ChatStream, LlmError>;
24
25    /// Non-streaming call (returns complete result at once).
26    ///
27    /// Use for: API services, batch processing, testing.
28    async fn chat(&self, request: ChatRequest) -> Result<ChatResponse, LlmError>;
29
30    /// Get provider capabilities.
31    fn capabilities(&self) -> Capabilities;
32
33    /// Get provider info.
34    fn info(&self) -> ProviderInfo;
35}
36
37#[cfg(test)]
38mod tests {
39    use super::*;
40    use crate::message::ChatMessage;
41    use crate::response::{FinishReason, StreamChunk};
42    use crate::types::UsageInfo;
43    use std::sync::Arc;
44
45    /// Mock provider for testing object safety
46    struct MockProvider;
47
48    #[async_trait]
49    impl LlmProvider for MockProvider {
50        async fn stream(&self, _request: ChatRequest) -> Result<ChatStream, LlmError> {
51            let chunks = vec![Ok(StreamChunk::Text("mock response".into()))];
52            Ok(ChatStream::new(Box::pin(futures_util::stream::iter(
53                chunks,
54            ))))
55        }
56
57        async fn chat(&self, _request: ChatRequest) -> Result<ChatResponse, LlmError> {
58            Ok(ChatResponse {
59                content: "mock response".to_string(),
60                reasoning_content: None,
61                thinking_signature: None,
62                tool_calls: vec![],
63                usage: UsageInfo::default(),
64                finish_reason: FinishReason::Stop,
65                raw: None,
66            })
67        }
68
69        fn capabilities(&self) -> Capabilities {
70            Capabilities {
71                supports_streaming: true,
72                supports_tools: true,
73                ..Default::default()
74            }
75        }
76
77        fn info(&self) -> ProviderInfo {
78            ProviderInfo {
79                name: "mock".to_string(),
80                model: "mock-model".to_string(),
81                version: None,
82            }
83        }
84    }
85
86    #[test]
87    fn trait_is_object_safe() {
88        let _provider: Arc<dyn LlmProvider> = Arc::new(MockProvider);
89    }
90
91    #[tokio::test]
92    async fn mock_provider_chat() {
93        let provider = MockProvider;
94        let request = ChatRequest::new(vec![ChatMessage::user("hello")]);
95        let response = provider.chat(request).await.unwrap();
96        assert_eq!(response.content, "mock response");
97        assert_eq!(response.finish_reason, FinishReason::Stop);
98    }
99
100    #[tokio::test]
101    async fn mock_provider_stream() {
102        let provider = MockProvider;
103        let request = ChatRequest::new(vec![ChatMessage::user("hello")]);
104        let stream = provider.stream(request).await.unwrap();
105        let text = stream.collect_text().await.unwrap();
106        assert_eq!(text, "mock response");
107    }
108
109    #[test]
110    fn mock_provider_capabilities() {
111        let provider = MockProvider;
112        let caps = provider.capabilities();
113        assert!(caps.supports_streaming);
114        assert!(caps.supports_tools);
115    }
116
117    #[test]
118    fn mock_provider_info() {
119        let provider = MockProvider;
120        let info = provider.info();
121        assert_eq!(info.name, "mock");
122        assert_eq!(info.model, "mock-model");
123    }
124}