1use 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#[async_trait]
19pub trait LlmProvider: Send + Sync {
20 async fn stream(&self, request: ChatRequest) -> Result<ChatStream, LlmError>;
24
25 async fn chat(&self, request: ChatRequest) -> Result<ChatResponse, LlmError>;
29
30 fn capabilities(&self) -> Capabilities;
32
33 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 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}