Skip to main content

rskit_llm/
provider.rs

1//! Provider trait — the canonical abstraction over LLM backends.
2//!
3//! This is the single full LLM provider trait. Implementors supply
4//! [`Provider::complete`], [`rskit_provider::Provider::name`], and
5//! [`rskit_provider::RequestResponse::execute`] (which typically delegates to
6//! `complete`).
7//!
8//! The trait extends
9//! `rskit_provider::RequestResponse<CompletionRequest, CompletionResponse>` so
10//! any LLM provider is natively usable in `dag`, `pipeline`, `chain`, `worker`,
11//! and `process` consumers without adapter shims.
12
13use 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/// A fully-featured LLM provider with streaming and capability introspection.
28///
29/// An adapter MUST implement [`Provider::complete`],
30/// [`rskit_provider::Provider::name`] (`&'static str`), and
31/// [`rskit_provider::RequestResponse::execute`] (typically delegates to
32/// `complete`). The
33/// default [`Provider::stream`] synthesizes a
34/// four-event sequence (`message.start` → `text.delta` → `usage.delta` →
35/// `message.stop`) by awaiting `complete`. Adapters whose backend supports
36/// native streaming SHOULD override `stream` to emit incremental events.
37///
38/// # Native provider shape
39///
40/// This trait requires
41/// `rskit_provider::RequestResponse<CompletionRequest, CompletionResponse>` as
42/// supertrait, so every `llm::Provider` carries the canonical
43/// identity/availability + request/response contract natively. The optional
44/// [`LlmStream`] wrapper remains available for consumers that specifically need
45/// the provider `Stream` shape.
46#[async_trait]
47pub trait Provider: rskit_provider::RequestResponse<CompletionRequest, CompletionResponse> {
48    /// Send a chat completion request and return the full response.
49    async fn complete(&self, request: CompletionRequest) -> Result<CompletionResponse, AppError>;
50
51    /// Stream a chat completion as a series of stream event objects.
52    ///
53    /// Default impl synthesizes events from [`Provider::complete`] for
54    /// adapters whose backend has no native streaming endpoint.
55    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    /// Describe what this provider / model supports. Default returns an
79    /// empty [`Capabilities`]; adapters SHOULD override to advertise tool use,
80    /// streaming, vision, etc.
81    fn capabilities(&self) -> Capabilities {
82        Capabilities::default()
83    }
84
85    /// Estimate the number of tokens consumed by the given messages. Default
86    /// uses the shared whitespace-based approximation from `rskit_ai::chat`.
87    fn count_tokens(&self, messages: &[Message]) -> usize {
88        count_tokens_approx(messages)
89    }
90}
91
92/// Adapter wrapping an `llm::Provider` as `provider::RequestResponse<CompletionRequest, CompletionResponse>`.
93///
94/// Use this to plug an LLM provider directly into pipeline/dag/chain consumers.
95pub 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
113/// Type alias for the provider-shaped boxed stream (mirrors `rskit_provider::traits::BoxStream`).
114type ProviderBoxStream<O> = Pin<Box<dyn FutStream<Item = AppResult<O>> + Send + 'static>>;
115
116/// Adapter wrapping an `llm::Provider` as `provider::Stream<CompletionRequest, StreamEventRef>`.
117///
118/// Use this to plug an LLM provider's streaming into pipeline/dag consumers.
119pub 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    /// `MockProvider` only implements `complete` to verify default impls
172    /// (stream, capabilities, `count_tokens`) compose correctly.
173    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}