1use std::pin::Pin;
14use std::sync::Arc;
15
16use async_trait::async_trait;
17use futures::Stream as FutStream;
18use rskit_ai::chat::{Message, count_tokens_approx};
19use rskit_ai::{
20 Capabilities, FinishReason, MessageStart, MessageStop, Role, StreamEventRef, TextDelta,
21 UsageDelta, text_of,
22};
23use rskit_errors::{AppError, AppResult};
24
25use crate::types::{CompletionRequest, CompletionResponse};
26
27#[async_trait]
47pub trait Provider: rskit_provider::RequestResponse<CompletionRequest, CompletionResponse> {
48 async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse, AppError>;
50
51 async fn stream(
56 &self,
57 request: CompletionRequest,
58 ) -> Result<Pin<Box<dyn FutStream<Item = StreamEventRef> + Send>>, AppError> {
59 let resp = self.complete(request).await?;
60 let text = text_of(&resp.message.content);
61 let model = resp.model.clone();
62 let usage = resp.usage;
63 let finish_reason = resp.stop_reason.unwrap_or(FinishReason::Stop);
64 let mut events: Vec<StreamEventRef> = Vec::with_capacity(4);
65 events.push(Arc::new(MessageStart {
66 role: Role::Assistant,
67 model,
68 request_id: None,
69 }));
70 if !text.is_empty() {
71 events.push(Arc::new(TextDelta { text }));
72 }
73 events.push(Arc::new(UsageDelta { usage }));
74 events.push(Arc::new(MessageStop { finish_reason }));
75 Ok(Box::pin(futures::stream::iter(events)))
76 }
77
78 fn capabilities(&self) -> Capabilities {
82 Capabilities::default()
83 }
84
85 fn count_tokens(&self, messages: &[Message]) -> usize {
88 count_tokens_approx(messages)
89 }
90}
91
92pub struct LlmRequestResponse<P: Provider>(pub Arc<P>);
96
97#[async_trait]
98impl<P: Provider + 'static> rskit_provider::Provider for LlmRequestResponse<P> {
99 fn name(&self) -> &'static str {
100 self.0.name()
101 }
102}
103
104#[async_trait]
105impl<P: Provider + 'static> rskit_provider::RequestResponse<CompletionRequest, CompletionResponse>
106 for LlmRequestResponse<P>
107{
108 async fn execute(&self, input: CompletionRequest) -> AppResult<CompletionResponse> {
109 self.0.complete(input).await
110 }
111}
112
113type ProviderBoxStream<O> = Pin<Box<dyn FutStream<Item = AppResult<O>> + Send + 'static>>;
115
116pub struct LlmStream<P: Provider>(pub Arc<P>);
120
121#[async_trait]
122impl<P: Provider + 'static> rskit_provider::Provider for LlmStream<P> {
123 fn name(&self) -> &'static str {
124 self.0.name()
125 }
126}
127
128impl<P: Provider + 'static> rskit_provider::Stream<CompletionRequest, StreamEventRef>
129 for LlmStream<P>
130{
131 async fn execute(
132 &self,
133 input: CompletionRequest,
134 ) -> AppResult<ProviderBoxStream<StreamEventRef>> {
135 use futures::StreamExt;
136 let raw = Provider::stream(&*self.0, input).await?;
137 Ok(Box::pin(raw.map(Ok)) as ProviderBoxStream<StreamEventRef>)
138 }
139}
140
141#[cfg(test)]
142mod tests {
143 use super::*;
144 use crate::{self as llm, types};
145 use futures::StreamExt;
146 use rskit_provider::RequestResponse;
147
148 #[test]
149 fn test_capabilities_default() {
150 let cap = Capabilities::default();
151 assert!(!cap.tool_use);
152 assert!(!cap.vision);
153 assert!(!cap.reasoning_tokens);
154 assert!(!cap.streaming);
155 assert_eq!(cap.max_input_tokens.unwrap_or_default(), 0);
156 assert!(cap.max_output_tokens.is_none());
157 }
158
159 #[test]
160 fn test_count_tokens_approx_user() {
161 let msgs = vec![types::user("hello world")];
162 assert!(count_tokens_approx(&msgs) > 0);
163 }
164
165 #[test]
166 fn test_count_tokens_approx_empty() {
167 let msgs: Vec<Message> = vec![];
168 assert_eq!(count_tokens_approx(&msgs), 0);
169 }
170
171 struct MockProvider;
174
175 #[async_trait]
176 impl rskit_provider::Provider for MockProvider {
177 fn name(&self) -> &'static str {
178 "mock"
179 }
180 }
181
182 #[async_trait]
183 impl rskit_provider::RequestResponse<CompletionRequest, CompletionResponse> for MockProvider {
184 async fn execute(&self, input: CompletionRequest) -> AppResult<CompletionResponse> {
185 self.complete(input).await
186 }
187 }
188
189 #[async_trait]
190 impl Provider for MockProvider {
191 async fn complete(
192 &self,
193 _request: CompletionRequest,
194 ) -> Result<CompletionResponse, AppError> {
195 Ok(CompletionResponse {
196 message: llm::AssistantMessage {
197 content: llm::text_content("Hi"),
198 tool_calls: vec![],
199 usage: None,
200 },
201 model: "mock".to_string(),
202 usage: rskit_ai::Usage {
203 input_tokens: 1,
204 output_tokens: 1,
205 cached_tokens: 0,
206 reasoning_tokens: 0,
207 },
208 stop_reason: Some(FinishReason::Stop),
209 })
210 }
211 }
212
213 #[tokio::test]
214 async fn test_mock_provider_complete() {
215 let provider = MockProvider;
216 let request = CompletionRequest {
217 model: "mock".to_string(),
218 messages: vec![types::user("hi")],
219 max_tokens: None,
220 temperature: None,
221 stream: false,
222 tools: None,
223 tool_choice: None,
224 };
225 let resp = provider.complete(request).await.unwrap();
226 assert_eq!(resp.model, "mock");
227 }
228
229 #[tokio::test]
230 async fn test_default_stream_synthesizes_from_complete() {
231 let provider = MockProvider;
232 let request = CompletionRequest {
233 model: "mock".to_string(),
234 messages: vec![types::user("hi")],
235 max_tokens: None,
236 temperature: None,
237 stream: true,
238 tools: None,
239 tool_choice: None,
240 };
241 let mut stream = provider.stream(request).await.unwrap();
242 let mut event_types = vec![];
243 while let Some(event) = stream.next().await {
244 event_types.push(event.event_type());
245 }
246 assert_eq!(
247 event_types,
248 vec!["message.start", "text.delta", "usage.delta", "message.stop"]
249 );
250 }
251
252 #[tokio::test]
253 async fn test_default_count_tokens_uses_approx() {
254 let provider = MockProvider;
255 let msgs = vec![types::user("hello world")];
256 assert_eq!(provider.count_tokens(&msgs), count_tokens_approx(&msgs));
257 }
258
259 #[tokio::test]
260 async fn test_llm_request_response_adapter() {
261 let provider = Arc::new(MockProvider);
262 let adapter = LlmRequestResponse(provider);
263 let request = CompletionRequest {
264 model: "mock".to_string(),
265 messages: vec![types::user("hi")],
266 max_tokens: None,
267 temperature: None,
268 stream: false,
269 tools: None,
270 tool_choice: None,
271 };
272 let resp = adapter.execute(request).await.unwrap();
273 assert_eq!(resp.model, "mock");
274 }
275}