Skip to main content

llm_agent_runtime/
providers.rs

1//! # Module: Providers
2//!
3//! ## Responsibility
4//! Provides built-in LLM inference integrations behind the `LlmProvider` trait.
5//! Optional feature flags gate each provider:
6//! - `anthropic` — Anthropic Messages API
7//! - `openai`    — OpenAI Chat Completions API (and compatible endpoints)
8//!
9//! ## Guarantees
10//! - `LlmProvider` is an async, object-safe trait
11//! - Both providers are `Send + Sync` and usable behind `Arc<dyn LlmProvider>`
12//! - Non-panicking: all operations return `Result`
13//!
14//! ## Feature Gate
15//! This module is only compiled when the `providers` feature (or a sub-feature) is enabled.
16
17use crate::error::AgentRuntimeError;
18use async_trait::async_trait;
19
20// ── CompletionOptions ─────────────────────────────────────────────────────────
21
22/// Options for a completion request, passed to [`LlmProvider::complete_with_options`].
23///
24/// Allows callers to supply per-request parameters (max tokens, temperature,
25/// timeout) in addition to the model name, without changing the base
26/// [`LlmProvider::complete`] signature.
27#[derive(Debug)]
28pub struct CompletionOptions<'a> {
29    /// Model identifier (e.g. `"claude-sonnet-4-6"`, `"gpt-4o"`).
30    pub model: &'a str,
31    /// Maximum number of output tokens. Overrides provider defaults when set.
32    pub max_tokens: Option<usize>,
33    /// Sampling temperature in `[0.0, 2.0]`. Higher = more random.
34    pub temperature: Option<f32>,
35    /// Per-request wall-clock timeout.
36    pub timeout: Option<std::time::Duration>,
37    /// Stop sequences: the model will stop generating when it produces any of
38    /// these strings.  An empty slice means no stop sequences.
39    pub stop_sequences: Vec<String>,
40}
41
42impl<'a> CompletionOptions<'a> {
43    /// Create options with just a model name and all other fields unset.
44    pub fn new(model: &'a str) -> Self {
45        Self {
46            model,
47            max_tokens: None,
48            temperature: None,
49            timeout: None,
50            stop_sequences: vec![],
51        }
52    }
53
54    /// Set the maximum output tokens.
55    pub fn with_max_tokens(mut self, n: usize) -> Self {
56        self.max_tokens = Some(n);
57        self
58    }
59
60    /// Set the sampling temperature.
61    pub fn with_temperature(mut self, t: f32) -> Self {
62        self.temperature = Some(t);
63        self
64    }
65
66    /// Set the per-request timeout.
67    pub fn with_timeout(mut self, d: std::time::Duration) -> Self {
68        self.timeout = Some(d);
69        self
70    }
71
72    /// Set stop sequences for this request.
73    pub fn with_stop_sequences(mut self, sequences: Vec<String>) -> Self {
74        self.stop_sequences = sequences;
75        self
76    }
77
78    /// Set the per-request timeout from a number of seconds.
79    pub fn with_timeout_secs(self, secs: u64) -> Self {
80        self.with_timeout(std::time::Duration::from_secs(secs))
81    }
82
83    /// Set the per-request timeout from a number of milliseconds.
84    pub fn with_timeout_ms(self, ms: u64) -> Self {
85        self.with_timeout(std::time::Duration::from_millis(ms))
86    }
87
88    /// Return `true` if at least one stop sequence has been configured.
89    pub fn has_stop_sequences(&self) -> bool {
90        !self.stop_sequences.is_empty()
91    }
92
93    /// Return the number of stop sequences configured.
94    pub fn stop_sequence_count(&self) -> usize {
95        self.stop_sequences.len()
96    }
97}
98
99// ── LlmProvider ───────────────────────────────────────────────────────────────
100
101/// Abstraction over an LLM inference endpoint.
102///
103/// Implement this trait to integrate any model API with `AgentRuntime`.
104/// Built-in implementations are provided for Anthropic and OpenAI
105/// when the corresponding feature flags are enabled.
106#[async_trait]
107pub trait LlmProvider: Send + Sync {
108    /// Send a prompt to the model and return the completion text.
109    ///
110    /// # Arguments
111    /// * `prompt` — the full prompt / context string
112    /// * `model`  — model identifier (e.g. `"claude-sonnet-4-6"`, `"gpt-4o"`)
113    async fn complete(&self, prompt: &str, model: &str) -> Result<String, AgentRuntimeError>;
114
115    /// Send a prompt with additional per-request options.
116    ///
117    /// The default implementation ignores `options.max_tokens`,
118    /// `options.temperature`, and `options.timeout` and delegates to
119    /// `complete(prompt, options.model)`.  Override this method to honour those
120    /// fields in your provider implementation.
121    async fn complete_with_options(
122        &self,
123        prompt: &str,
124        options: CompletionOptions<'_>,
125    ) -> Result<String, AgentRuntimeError> {
126        self.complete(prompt, options.model).await
127    }
128
129    /// Stream the completion token-by-token.
130    ///
131    /// Returns a `Receiver` that yields string chunks as they arrive.
132    /// The channel closes when the stream is complete or an error occurs.
133    ///
134    /// # Default implementation
135    ///
136    /// The default wraps `complete` into a single-chunk stream using a channel
137    /// with a buffer of 64 slots.  Custom providers that support true token
138    /// streaming should override this method; the 64-slot buffer is sized so
139    /// that fast producers do not block waiting for a slow consumer to drain
140    /// the first chunk.
141    ///
142    /// # Note for implementors
143    ///
144    /// If you override this method, choose a channel capacity that balances
145    /// memory use against throughput for your expected token rate.  A capacity
146    /// of 1 will cause the producer to block after each token; a capacity of
147    /// 0 is unbounded and may exhaust memory on a slow consumer.
148    async fn stream_complete(
149        &self,
150        prompt: &str,
151        model: &str,
152    ) -> Result<tokio::sync::mpsc::Receiver<Result<String, AgentRuntimeError>>, AgentRuntimeError>
153    {
154        let result = self.complete(prompt, model).await;
155        let (tx, rx) = tokio::sync::mpsc::channel(64);
156        // Ignore send error — receiver may already be dropped
157        let _ = tx.send(result).await;
158        Ok(rx)
159    }
160}
161
162// ── AnthropicProvider ─────────────────────────────────────────────────────────
163
164#[cfg(feature = "anthropic")]
165/// Built-in provider for the Anthropic Messages API.
166///
167/// Requires the `anthropic` feature flag.
168///
169/// # Example
170/// ```no_run
171/// use llm_agent_runtime::providers::AnthropicProvider;
172/// let provider = AnthropicProvider::new("sk-ant-...");
173/// ```
174pub struct AnthropicProvider {
175    api_key: String,
176    /// Messages API endpoint. Overridable via [`with_base_url`](AnthropicProvider::with_base_url)
177    /// to point at a mock server or proxy during testing.
178    api_url: String,
179    client: reqwest::Client,
180    /// Semaphore bounding concurrent SSE background tasks.
181    stream_semaphore: std::sync::Arc<tokio::sync::Semaphore>,
182    /// Maximum output tokens for `stream_complete`. Falls back to `MAX_TOKENS`
183    /// when `None`.
184    stream_max_tokens: Option<u32>,
185}
186
187#[cfg(feature = "anthropic")]
188impl AnthropicProvider {
189    const DEFAULT_API_URL: &'static str = "https://api.anthropic.com/v1/messages";
190    const API_VERSION: &'static str = "2023-06-01";
191    const MAX_TOKENS: u32 = 1024;
192    /// Default maximum number of concurrent streaming tasks.
193    const DEFAULT_STREAM_CONCURRENCY: usize = 32;
194
195    /// Create a new Anthropic provider with the given API key.
196    pub fn new(api_key: impl Into<String>) -> Self {
197        Self {
198            api_key: api_key.into(),
199            api_url: Self::DEFAULT_API_URL.to_owned(),
200            client: reqwest::Client::new(),
201            stream_semaphore: std::sync::Arc::new(tokio::sync::Semaphore::new(
202                Self::DEFAULT_STREAM_CONCURRENCY,
203            )),
204            stream_max_tokens: None,
205        }
206    }
207
208    /// Create a provider pointing to a custom API endpoint.
209    ///
210    /// Useful for testing against a mock server or routing through a proxy.
211    pub fn with_base_url(api_key: impl Into<String>, api_url: impl Into<String>) -> Self {
212        Self {
213            api_key: api_key.into(),
214            api_url: api_url.into(),
215            client: reqwest::Client::new(),
216            stream_semaphore: std::sync::Arc::new(tokio::sync::Semaphore::new(
217                Self::DEFAULT_STREAM_CONCURRENCY,
218            )),
219            stream_max_tokens: None,
220        }
221    }
222
223    /// Create a provider with a custom limit on concurrent streaming tasks.
224    ///
225    /// Useful when many agents share a single provider and you want to cap
226    /// the number of simultaneous SSE parse tasks.
227    pub fn with_max_concurrent_streams(api_key: impl Into<String>, max: usize) -> Self {
228        Self {
229            api_key: api_key.into(),
230            api_url: Self::DEFAULT_API_URL.to_owned(),
231            client: reqwest::Client::new(),
232            stream_semaphore: std::sync::Arc::new(tokio::sync::Semaphore::new(max)),
233            stream_max_tokens: None,
234        }
235    }
236
237    /// Set the maximum output tokens used by [`stream_complete`].
238    ///
239    /// By default `stream_complete` uses the hard-coded `MAX_TOKENS` constant
240    /// (1024).  Call this method to override that value, e.g. when a task
241    /// requires longer streamed responses.
242    ///
243    /// [`stream_complete`]: AnthropicProvider::stream_complete
244    pub fn with_stream_max_tokens(mut self, max_tokens: u32) -> Self {
245        self.stream_max_tokens = Some(max_tokens);
246        self
247    }
248}
249
250#[cfg(feature = "anthropic")]
251#[async_trait]
252impl LlmProvider for AnthropicProvider {
253    async fn complete(&self, prompt: &str, model: &str) -> Result<String, AgentRuntimeError> {
254        self.complete_with_options(prompt, CompletionOptions::new(model))
255            .await
256    }
257
258    /// Complete with per-request options (max_tokens, temperature).
259    ///
260    /// Falls back to `Self::MAX_TOKENS` when `options.max_tokens` is not set.
261    #[tracing::instrument(skip(self, prompt, options), fields(model = options.model, provider = "anthropic"))]
262    async fn complete_with_options(
263        &self,
264        prompt: &str,
265        options: CompletionOptions<'_>,
266    ) -> Result<String, AgentRuntimeError> {
267        let max_tokens = options
268            .max_tokens
269            .unwrap_or(Self::MAX_TOKENS as usize) as u32;
270
271        let mut body = serde_json::json!({
272            "model": options.model,
273            "max_tokens": max_tokens,
274            "messages": [{ "role": "user", "content": prompt }]
275        });
276        if let Some(t) = options.temperature {
277            body["temperature"] = serde_json::json!(t);
278        }
279
280        let mut req = self
281            .client
282            .post(&self.api_url)
283            .header("x-api-key", &self.api_key)
284            .header("anthropic-version", Self::API_VERSION)
285            .header("content-type", "application/json")
286            .json(&body);
287        if let Some(timeout) = options.timeout {
288            req = req.timeout(timeout);
289        }
290        let response = req
291            .send()
292            .await
293            .map_err(|e| AgentRuntimeError::Provider(format!("Anthropic request failed: {e}")))?;
294
295        if !response.status().is_success() {
296            let status = response.status();
297            let text = response.text().await.unwrap_or_default();
298            return Err(AgentRuntimeError::Provider(format!(
299                "Anthropic API error {status}: {text}"
300            )));
301        }
302
303        let json: serde_json::Value = response
304            .json()
305            .await
306            .map_err(|e| AgentRuntimeError::Provider(format!("Anthropic parse failed: {e}")))?;
307
308        let text = json["content"]
309            .as_array()
310            .and_then(|arr| arr.first())
311            .and_then(|block| block["text"].as_str())
312            .ok_or_else(|| {
313                AgentRuntimeError::Provider("Anthropic response missing content[0].text".into())
314            })?;
315
316        Ok(text.to_owned())
317    }
318
319    /// Stream the completion token-by-token using `chunk()` for true streaming.
320    ///
321    /// Uses `response.chunk()` to consume the SSE body incrementally so tokens
322    /// are emitted as they arrive rather than after the entire response has been
323    /// buffered, fixing the original `bytes().await` implementation.
324    async fn stream_complete(
325        &self,
326        prompt: &str,
327        model: &str,
328    ) -> Result<tokio::sync::mpsc::Receiver<Result<String, AgentRuntimeError>>, AgentRuntimeError>
329    {
330        let max_tokens = self.stream_max_tokens.unwrap_or(Self::MAX_TOKENS);
331        let body = serde_json::json!({
332            "model": model,
333            "max_tokens": max_tokens,
334            "stream": true,
335            "messages": [{ "role": "user", "content": prompt }]
336        });
337
338        let response = self
339            .client
340            .post(&self.api_url)
341            .header("x-api-key", &self.api_key)
342            .header("anthropic-version", Self::API_VERSION)
343            .header("content-type", "application/json")
344            .json(&body)
345            .send()
346            .await
347            .map_err(|e| {
348                AgentRuntimeError::Provider(format!("Anthropic stream request failed: {e}"))
349            })?;
350
351        if !response.status().is_success() {
352            let status = response.status();
353            let text = response.text().await.unwrap_or_default();
354            return Err(AgentRuntimeError::Provider(format!(
355                "Anthropic stream API error {status}: {text}"
356            )));
357        }
358
359        let (tx, rx) = tokio::sync::mpsc::channel::<Result<String, AgentRuntimeError>>(32);
360
361        let permit = std::sync::Arc::clone(&self.stream_semaphore)
362            .acquire_owned()
363            .await
364            .map_err(|_| {
365                AgentRuntimeError::Provider("Anthropic stream semaphore closed".into())
366            })?;
367
368        // Parse SSE incrementally using chunk() so tokens are emitted as they
369        // arrive rather than after the full response body has been buffered.
370        tokio::spawn(async move {
371            let _permit = permit;
372            let mut response = response;
373            let mut buffer = String::new();
374            loop {
375                match response.chunk().await {
376                    Ok(Some(chunk)) => {
377                        match String::from_utf8(chunk.to_vec()) {
378                            Ok(s) => buffer.push_str(&s),
379                            Err(e) => {
380                                let _ = tx
381                                    .send(Err(AgentRuntimeError::Provider(format!(
382                                        "Anthropic stream: invalid UTF-8 in chunk: {e}"
383                                    ))))
384                                    .await;
385                                return;
386                            }
387                        }
388                        // Drain complete SSE lines from the buffer.
389                        while let Some(newline) = buffer.find('\n') {
390                            let line = buffer[..newline].trim().to_owned();
391                            buffer = buffer[newline + 1..].to_owned();
392                            if let Some(data) = line.strip_prefix("data: ") {
393                                if data == "[DONE]" {
394                                    return;
395                                }
396                                if let Ok(json) =
397                                    serde_json::from_str::<serde_json::Value>(data)
398                                {
399                                    if let Some(delta) = json["delta"]["text"].as_str() {
400                                        if tx.send(Ok(delta.to_owned())).await.is_err() {
401                                            return;
402                                        }
403                                    }
404                                }
405                            }
406                        }
407                    }
408                    Ok(None) => break,
409                    Err(e) => {
410                        let _ = tx
411                            .send(Err(AgentRuntimeError::Provider(format!(
412                                "Anthropic stream chunk error: {e}"
413                            ))))
414                            .await;
415                        return;
416                    }
417                }
418            }
419        });
420
421        Ok(rx)
422    }
423}
424
425// ── OpenAiProvider ────────────────────────────────────────────────────────────
426
427#[cfg(feature = "openai")]
428/// Built-in provider for the OpenAI Chat Completions API.
429///
430/// Also compatible with Azure OpenAI and any OpenAI-compatible endpoint.
431/// Requires the `openai` feature flag.
432///
433/// # Example
434/// ```no_run
435/// use llm_agent_runtime::providers::OpenAiProvider;
436/// let provider = OpenAiProvider::new("sk-...");
437/// // For Azure or custom endpoints:
438/// let custom = OpenAiProvider::with_base_url("sk-...", "https://my-endpoint/v1");
439/// ```
440pub struct OpenAiProvider {
441    api_key: String,
442    base_url: String,
443    client: reqwest::Client,
444    /// Semaphore bounding concurrent SSE background tasks (item 10).
445    stream_semaphore: std::sync::Arc<tokio::sync::Semaphore>,
446}
447
448#[cfg(feature = "openai")]
449impl OpenAiProvider {
450    const DEFAULT_BASE_URL: &'static str = "https://api.openai.com/v1";
451    /// Default maximum number of concurrent streaming tasks.
452    const DEFAULT_STREAM_CONCURRENCY: usize = 32;
453
454    /// Create a new OpenAI provider with the default base URL.
455    pub fn new(api_key: impl Into<String>) -> Self {
456        Self {
457            api_key: api_key.into(),
458            base_url: Self::DEFAULT_BASE_URL.to_owned(),
459            client: reqwest::Client::new(),
460            stream_semaphore: std::sync::Arc::new(tokio::sync::Semaphore::new(
461                Self::DEFAULT_STREAM_CONCURRENCY,
462            )),
463        }
464    }
465
466    /// Create a provider pointing to a custom base URL (e.g. Azure, local models).
467    pub fn with_base_url(api_key: impl Into<String>, base_url: impl Into<String>) -> Self {
468        Self {
469            api_key: api_key.into(),
470            base_url: base_url.into(),
471            client: reqwest::Client::new(),
472            stream_semaphore: std::sync::Arc::new(tokio::sync::Semaphore::new(
473                Self::DEFAULT_STREAM_CONCURRENCY,
474            )),
475        }
476    }
477
478    /// Create a provider with a custom limit on concurrent streaming tasks.
479    pub fn with_max_concurrent_streams(
480        api_key: impl Into<String>,
481        base_url: impl Into<String>,
482        max: usize,
483    ) -> Self {
484        Self {
485            api_key: api_key.into(),
486            base_url: base_url.into(),
487            client: reqwest::Client::new(),
488            stream_semaphore: std::sync::Arc::new(tokio::sync::Semaphore::new(max)),
489        }
490    }
491}
492
493#[cfg(feature = "openai")]
494#[async_trait]
495impl LlmProvider for OpenAiProvider {
496    #[tracing::instrument(skip(self, prompt), fields(model, provider = "openai"))]
497    async fn complete(&self, prompt: &str, model: &str) -> Result<String, AgentRuntimeError> {
498        self.complete_with_options(prompt, CompletionOptions::new(model))
499            .await
500    }
501
502    /// Complete with per-request options (max_tokens, temperature, timeout).
503    ///
504    /// Honours `options.max_tokens` and `options.temperature` when set, falling
505    /// back to the API default when unset.
506    #[tracing::instrument(skip(self, prompt, options), fields(model = options.model, provider = "openai"))]
507    async fn complete_with_options(
508        &self,
509        prompt: &str,
510        options: CompletionOptions<'_>,
511    ) -> Result<String, AgentRuntimeError> {
512        let url = format!("{}/chat/completions", self.base_url);
513        let mut body = serde_json::json!({
514            "model": options.model,
515            "messages": [{ "role": "user", "content": prompt }]
516        });
517        if let Some(max_tokens) = options.max_tokens {
518            body["max_tokens"] = serde_json::json!(max_tokens);
519        }
520        if let Some(temp) = options.temperature {
521            body["temperature"] = serde_json::json!(temp);
522        }
523
524        let mut req = self
525            .client
526            .post(&url)
527            .bearer_auth(&self.api_key)
528            .header("content-type", "application/json")
529            .json(&body);
530        if let Some(timeout) = options.timeout {
531            req = req.timeout(timeout);
532        }
533        let response = req
534            .send()
535            .await
536            .map_err(|e| AgentRuntimeError::Provider(format!("OpenAI request failed: {e}")))?;
537
538        if !response.status().is_success() {
539            let status = response.status();
540            let text = response.text().await.unwrap_or_default();
541            return Err(AgentRuntimeError::Provider(format!(
542                "OpenAI API error {status}: {text}"
543            )));
544        }
545
546        let json: serde_json::Value = response
547            .json()
548            .await
549            .map_err(|e| AgentRuntimeError::Provider(format!("OpenAI parse failed: {e}")))?;
550
551        let text = json["choices"]
552            .as_array()
553            .and_then(|arr| arr.first())
554            .and_then(|choice| choice["message"]["content"].as_str())
555            .ok_or_else(|| {
556                AgentRuntimeError::Provider(
557                    "OpenAI response missing choices[0].message.content".into(),
558                )
559            })?;
560
561        Ok(text.to_owned())
562    }
563
564    async fn stream_complete(
565        &self,
566        prompt: &str,
567        model: &str,
568    ) -> Result<tokio::sync::mpsc::Receiver<Result<String, AgentRuntimeError>>, AgentRuntimeError>
569    {
570        let url = format!("{}/chat/completions", self.base_url);
571        let body = serde_json::json!({
572            "model": model,
573            "stream": true,
574            "messages": [{ "role": "user", "content": prompt }]
575        });
576
577        let mut response = self
578            .client
579            .post(&url)
580            .bearer_auth(&self.api_key)
581            .header("content-type", "application/json")
582            .json(&body)
583            .send()
584            .await
585            .map_err(|e| {
586                AgentRuntimeError::Provider(format!("OpenAI stream request failed: {e}"))
587            })?;
588
589        if !response.status().is_success() {
590            let status = response.status();
591            let text = response.text().await.unwrap_or_default();
592            return Err(AgentRuntimeError::Provider(format!(
593                "OpenAI stream API error {status}: {text}"
594            )));
595        }
596
597        let (tx, rx) = tokio::sync::mpsc::channel::<Result<String, AgentRuntimeError>>(32);
598
599        // Item 10 — bound concurrent tasks via semaphore.
600        let permit = std::sync::Arc::clone(&self.stream_semaphore)
601            .acquire_owned()
602            .await
603            .map_err(|_| {
604                AgentRuntimeError::Provider("OpenAI stream semaphore closed".into())
605            })?;
606
607        tokio::spawn(async move {
608            let _permit = permit;
609            // Incremental chunk-by-chunk reading — avoids buffering the full
610            // response body before emitting any tokens.
611            let mut buffer = String::new();
612            loop {
613                match response.chunk().await {
614                    Ok(Some(chunk)) => {
615                        let text = match std::str::from_utf8(&chunk) {
616                            Ok(t) => t,
617                            Err(e) => {
618                                let _ = tx
619                                    .send(Err(AgentRuntimeError::Provider(format!(
620                                        "OpenAI stream chunk is not valid UTF-8: {e}"
621                                    ))))
622                                    .await;
623                                return;
624                            }
625                        };
626                        buffer.push_str(text);
627                        // Drain complete SSE lines from the buffer.
628                        while let Some(newline) = buffer.find('\n') {
629                            let line: String = buffer.drain(..=newline).collect();
630                            let line = line.trim_end_matches(['\r', '\n']);
631                            if let Some(data) = line.strip_prefix("data: ") {
632                                if data == "[DONE]" {
633                                    return;
634                                }
635                                if let Ok(json) =
636                                    serde_json::from_str::<serde_json::Value>(data)
637                                {
638                                    if let Some(content) = json["choices"]
639                                        .as_array()
640                                        .and_then(|c| c.first())
641                                        .and_then(|c| c["delta"]["content"].as_str())
642                                    {
643                                        if tx.send(Ok(content.to_owned())).await.is_err() {
644                                            return;
645                                        }
646                                    }
647                                }
648                            }
649                        }
650                    }
651                    Ok(None) => break,
652                    Err(e) => {
653                        let _ = tx
654                            .send(Err(AgentRuntimeError::Provider(format!(
655                                "OpenAI stream read failed: {e}"
656                            ))))
657                            .await;
658                        return;
659                    }
660                }
661            }
662        });
663
664        Ok(rx)
665    }
666}
667
668// ── Tests ─────────────────────────────────────────────────────────────────────
669
670#[cfg(test)]
671mod tests {
672    use super::*;
673    use std::sync::Arc;
674
675    /// A stub provider for testing.
676    struct StubProvider {
677        response: String,
678    }
679
680    #[async_trait]
681    impl LlmProvider for StubProvider {
682        async fn complete(&self, _prompt: &str, _model: &str) -> Result<String, AgentRuntimeError> {
683            Ok(self.response.clone())
684        }
685        // Uses default stream_complete implementation which wraps complete().
686    }
687
688    #[tokio::test]
689    async fn test_stub_provider_returns_configured_response() {
690        let p = StubProvider {
691            response: "hello".into(),
692        };
693        let result = p.complete("prompt", "stub-model").await.unwrap();
694        assert_eq!(result, "hello");
695    }
696
697    #[tokio::test]
698    async fn test_llm_provider_is_object_safe() {
699        let p: Arc<dyn LlmProvider> = Arc::new(StubProvider {
700            response: "ok".into(),
701        });
702        let result = p.complete("test", "model").await.unwrap();
703        assert_eq!(result, "ok");
704    }
705
706    #[tokio::test]
707    async fn test_stub_provider_ignores_model_parameter() {
708        let p = StubProvider {
709            response: "42".into(),
710        };
711        let r1 = p.complete("q", "model-a").await.unwrap();
712        let r2 = p.complete("q", "model-b").await.unwrap();
713        assert_eq!(r1, r2);
714    }
715
716    #[tokio::test]
717    async fn test_stub_provider_stream_returns_single_chunk() {
718        let p = StubProvider {
719            response: "hello world".into(),
720        };
721        let mut rx = p.stream_complete("prompt", "model").await.unwrap();
722        let mut collected = String::new();
723        while let Some(chunk) = rx.recv().await {
724            collected.push_str(&chunk.unwrap());
725        }
726        assert_eq!(collected, "hello world");
727    }
728
729    #[tokio::test]
730    async fn test_stream_receiver_closes_after_completion() {
731        let p = StubProvider {
732            response: "done".into(),
733        };
734        let mut rx = p.stream_complete("prompt", "model").await.unwrap();
735        // Drain all chunks
736        while let Some(_chunk) = rx.recv().await {}
737        // Channel should now be closed — next recv returns None
738        assert!(rx.recv().await.is_none());
739    }
740}