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