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