Skip to main content

lucy/
provider.rs

1use std::collections::BTreeMap;
2use std::io::{self, BufRead};
3use std::time::Duration;
4
5use reqwest::blocking::Client;
6use reqwest::header::{HeaderMap, HeaderValue};
7use reqwest::Client as AsyncClient;
8use serde_json::{json, Value};
9
10use crate::cancellation::CancellationToken;
11use crate::codex_provider::CodexProvider;
12use crate::config::LlmSettings;
13use crate::model::{ChatMessage, ChatToolCall};
14use crate::redaction::{conflicts_with_protected_literal, redact_secret, redaction_marker};
15
16/// Maximum idle interval between provider response reads.
17pub const PROVIDER_TIMEOUT: Duration = Duration::from_secs(60);
18const PROVIDER_CONNECT_TIMEOUT: Duration = Duration::from_secs(10);
19const PROVIDER_RETRY_COUNT: usize = 1;
20const PROVIDER_RETRY_BACKOFF: Duration = Duration::from_millis(250);
21const MAX_PROVIDER_CONTENT_BYTES: usize = 1024 * 1024;
22const MAX_PROVIDER_REASONING_DETAILS_BYTES: usize = 1024 * 1024;
23const MAX_PROVIDER_TOOL_ARGUMENT_BYTES: usize = 1024 * 1024;
24const MAX_SSE_LINE_BYTES: usize = 64 * 1024;
25const MAX_SSE_EVENT_BYTES: usize = 1024 * 1024;
26const MAX_SSE_STREAM_BYTES: usize = 8 * 1024 * 1024;
27const MAX_SSE_DATA_LINES: usize = 1024;
28const MAX_PROVIDER_TOOL_CALL_ID_BYTES: usize = 16 * 1024;
29const MAX_PROVIDER_TOOL_NAME_BYTES: usize = 16 * 1024;
30const MAX_PROVIDER_ERROR_BYTES: usize = 16 * 1024;
31const CANCELLATION_POLL_INTERVAL: Duration = Duration::from_millis(10);
32const MODEL_METADATA_TIMEOUT: Duration = Duration::from_secs(2);
33const MAX_MODEL_METADATA_BYTES: usize = 4 * 1024 * 1024;
34const COMPACTION_MAX_SUMMARY_TOKENS: usize = 4_096;
35const OPENROUTER_HOST: &str = "openrouter.ai";
36const APP_TITLE: &str = "Lucy";
37const APP_URL: &str = "https://lucyna.run";
38
39#[derive(Debug)]
40pub struct ProviderError {
41    message: String,
42    cancelled: bool,
43    partial: Option<ProviderTurn>,
44    retryable: bool,
45}
46
47impl ProviderError {
48    pub(crate) fn new(message: impl Into<String>) -> Self {
49        Self {
50            message: message.into(),
51            cancelled: false,
52            partial: None,
53            retryable: false,
54        }
55    }
56
57    fn retryable(message: impl Into<String>) -> Self {
58        Self {
59            message: message.into(),
60            cancelled: false,
61            partial: None,
62            retryable: true,
63        }
64    }
65
66    pub(crate) fn cancelled(partial: ProviderTurn) -> Self {
67        Self {
68            message: "provider stream canceled".to_owned(),
69            cancelled: true,
70            partial: Some(partial),
71            retryable: false,
72        }
73    }
74
75    pub fn is_cancelled(&self) -> bool {
76        self.cancelled
77    }
78
79    pub fn partial_turn(&self) -> Option<&ProviderTurn> {
80        self.partial.as_ref()
81    }
82
83    fn is_retryable(&self) -> bool {
84        self.retryable
85    }
86}
87
88impl std::fmt::Display for ProviderError {
89    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
90        formatter.write_str(&self.message)
91    }
92}
93
94impl std::error::Error for ProviderError {}
95
96fn transient_http_status(status: u16) -> bool {
97    matches!(status, 408 | 429 | 500 | 502 | 503 | 504)
98}
99
100fn transient_reqwest_error(error: &reqwest::Error) -> bool {
101    error.is_timeout()
102        || error.is_connect()
103        || error.is_request()
104        || error.is_body()
105        || error.is_decode()
106}
107
108fn reqwest_error_kind(error: &reqwest::Error) -> &'static str {
109    if error.is_timeout() {
110        "timeout"
111    } else if error.is_connect() {
112        "connection"
113    } else if error.is_body() {
114        "body"
115    } else if error.is_decode() {
116        "decode"
117    } else if error.is_request() {
118        "request"
119    } else {
120        "transport"
121    }
122}
123
124fn reqwest_failure(
125    prefix: &str,
126    error: reqwest::Error,
127    api_key: &str,
128    retry_before_payload: bool,
129) -> ProviderError {
130    let detail = redact_secret(&error.to_string(), Some(api_key));
131    let message = format!("{prefix} ({}): {detail}", reqwest_error_kind(&error));
132    let mut provider_error = ProviderError::new(message);
133    provider_error.retryable = retry_before_payload && transient_reqwest_error(&error);
134    provider_error
135}
136
137#[derive(Debug, Clone, PartialEq, Eq)]
138pub struct ProviderTurn {
139    pub content: String,
140    pub tool_calls: Vec<ChatToolCall>,
141    pub reasoning_details: Vec<Value>,
142}
143
144fn empty_turn() -> ProviderTurn {
145    ProviderTurn {
146        content: String::new(),
147        tool_calls: Vec::new(),
148        reasoning_details: Vec::new(),
149    }
150}
151
152pub(crate) enum ProviderStreamEvent {
153    ReasoningStarted,
154    Text(String),
155}
156
157#[derive(Debug, Clone, PartialEq, Eq)]
158pub struct ProviderModel {
159    pub id: String,
160    pub efforts: Option<Vec<String>>,
161}
162
163pub struct Provider {
164    client: Client,
165    async_client: AsyncClient,
166    endpoint: String,
167    model: String,
168    effort: Option<String>,
169    api_key_env: String,
170    api_key: String,
171    codex: Option<CodexProvider>,
172    session_id: Option<String>,
173    is_openrouter: bool,
174}
175
176fn model_efforts(entry: &Value) -> Option<Vec<String>> {
177    let values = entry
178        .get("reasoning")
179        .and_then(|reasoning| reasoning.get("supported_efforts"))
180        .or_else(|| {
181            [
182                "supported_reasoning_efforts",
183                "reasoning_efforts",
184                "reasoning_effort",
185                "efforts",
186            ]
187            .into_iter()
188            .find_map(|key| entry.get(key))
189        })
190        .and_then(Value::as_array)?;
191    let efforts = values
192        .iter()
193        .filter_map(Value::as_str)
194        .map(str::trim)
195        .filter(|value| !value.is_empty())
196        .fold(Vec::new(), |mut efforts, value| {
197            if !efforts.iter().any(|effort| effort == value) {
198                efforts.push(value.to_owned());
199            }
200            efforts
201        });
202    (!efforts.is_empty()).then_some(efforts)
203}
204
205fn context_window_from_models(payload: &Value, model: &str) -> Option<usize> {
206    let models = payload.get("data").and_then(Value::as_array)?;
207    let entry = models.iter().find(|entry| {
208        entry.get("id").and_then(Value::as_str) == Some(model)
209            || entry.get("name").and_then(Value::as_str) == Some(model)
210    })?;
211    [
212        entry.get("context_length"),
213        entry.get("context_window"),
214        entry.get("max_context_length"),
215        entry
216            .get("top_provider")
217            .and_then(|provider| provider.get("context_length")),
218    ]
219    .into_iter()
220    .flatten()
221    .find_map(Value::as_u64)
222    .and_then(|value| usize::try_from(value).ok())
223    .filter(|value| *value > 0)
224}
225
226fn chat_request(
227    model: &str,
228    messages: &[ChatMessage],
229    effort: &Option<String>,
230    include_tools: bool,
231    session_id: Option<&str>,
232) -> Value {
233    let mut request = json!({
234        "model": model,
235        "messages": messages
236            .iter()
237            .map(ChatMessage::to_openai_value)
238            .collect::<Vec<_>>(),
239        "stream": true,
240    });
241    if include_tools {
242        let tools = vec![json!({
243            "type": "function",
244            "function": {
245                "name": "cmd",
246                "description": "Execute a shell command in the session starting directory. Set background to return immediately and receive the completed result automatically.",
247                "parameters": {"type": "object", "properties": {"command": {"type": "string"}, "background": {"type": "boolean", "default": false}}, "required": ["command"], "additionalProperties": false}
248            }
249        })];
250        request["tools"] = Value::Array(tools);
251    } else {
252        request["max_tokens"] = json!(COMPACTION_MAX_SUMMARY_TOKENS);
253    }
254    if let Some(effort) = effort {
255        request["reasoning_effort"] = json!(effort);
256    }
257    if let Some(session_id) = session_id {
258        request["session_id"] = json!(session_id);
259    }
260    request
261}
262
263fn is_openrouter_base_url(base_url: &str) -> bool {
264    reqwest::Url::parse(base_url)
265        .ok()
266        .and_then(|url| url.host_str().map(str::to_owned))
267        .is_some_and(|host| host.eq_ignore_ascii_case(OPENROUTER_HOST))
268}
269
270fn provider_headers(is_openrouter: bool) -> HeaderMap {
271    let mut headers = HeaderMap::new();
272    if is_openrouter {
273        headers.insert("x-openrouter-title", HeaderValue::from_static(APP_TITLE));
274        headers.insert("http-referer", HeaderValue::from_static(APP_URL));
275    }
276    headers
277}
278
279fn provider_clients(
280    is_openrouter: bool,
281    api_key: &str,
282) -> Result<(Client, AsyncClient), ProviderError> {
283    let headers = provider_headers(is_openrouter);
284    let client = Client::builder()
285        .connect_timeout(PROVIDER_CONNECT_TIMEOUT)
286        .timeout(PROVIDER_TIMEOUT)
287        .default_headers(headers.clone())
288        .build()
289        .map_err(|_| {
290            ProviderError::new(redact_secret(
291                "unable to initialize HTTP client",
292                Some(api_key),
293            ))
294        })?;
295    let async_client = AsyncClient::builder()
296        .connect_timeout(PROVIDER_CONNECT_TIMEOUT)
297        .read_timeout(PROVIDER_TIMEOUT)
298        .default_headers(headers)
299        .build()
300        .map_err(|_| {
301            ProviderError::new(redact_secret(
302                "unable to initialize HTTP client",
303                Some(api_key),
304            ))
305        })?;
306    Ok((client, async_client))
307}
308
309impl Provider {
310    pub fn new(settings: &LlmSettings) -> Result<Self, ProviderError> {
311        let api_key = match std::env::var(&settings.api_key_env) {
312            Ok(api_key) if !api_key.is_empty() => api_key,
313            Ok(_) | Err(_) => return Err(ProviderError::new("missing provider API key")),
314        };
315        if conflicts_with_protected_literal(&api_key) {
316            return Err(ProviderError::new(redact_secret(
317                "API key conflicts with a required structured output literal",
318                Some(&api_key),
319            )));
320        }
321        if redaction_marker(&api_key).is_none() {
322            return Err(ProviderError::new(redact_secret(
323                "API key cannot be safely redacted",
324                Some(&api_key),
325            )));
326        }
327        if settings.model.trim().is_empty() {
328            return Err(ProviderError::new(redact_secret(
329                "missing llm.model; set a model in config.toml",
330                Some(&api_key),
331            )));
332        }
333        let effort = match &settings.effort {
334            Some(value) => {
335                let trimmed = value.trim();
336                if trimmed.is_empty() {
337                    return Err(ProviderError::new(redact_secret(
338                        "llm.effort must not be empty",
339                        Some(&api_key),
340                    )));
341                }
342                Some(trimmed.to_owned())
343            }
344            None => None,
345        };
346        let endpoint = format!(
347            "{}/chat/completions",
348            settings.base_url.trim_end_matches('/')
349        );
350        let is_openrouter = is_openrouter_base_url(&settings.base_url);
351        let (client, async_client) = provider_clients(is_openrouter, &api_key)?;
352        Ok(Self {
353            client,
354            async_client,
355            endpoint,
356            model: settings.model.clone(),
357            effort,
358            api_key_env: settings.api_key_env.clone(),
359            api_key,
360            codex: None,
361            session_id: None,
362            is_openrouter,
363        })
364    }
365
366    pub fn new_codex(
367        home: &std::path::Path,
368        settings: &LlmSettings,
369    ) -> Result<Self, ProviderError> {
370        let codex = CodexProvider::new(home, settings)?;
371        let api_key = codex.api_key().to_owned();
372        let client = Client::builder()
373            .connect_timeout(PROVIDER_CONNECT_TIMEOUT)
374            .timeout(PROVIDER_TIMEOUT)
375            .build()
376            .map_err(|_| ProviderError::new("unable to initialize HTTP client"))?;
377        let async_client = AsyncClient::builder()
378            .connect_timeout(PROVIDER_CONNECT_TIMEOUT)
379            .read_timeout(PROVIDER_TIMEOUT)
380            .build()
381            .map_err(|_| ProviderError::new("unable to initialize HTTP client"))?;
382        Ok(Self {
383            client,
384            async_client,
385            endpoint: crate::codex_provider::CODEX_ENDPOINT.to_owned(),
386            model: settings.model.clone(),
387            effort: settings.effort.clone(),
388            api_key_env: codex.api_key_env().to_owned(),
389            api_key,
390            codex: Some(codex),
391            session_id: None,
392            is_openrouter: false,
393        })
394    }
395
396    pub(crate) fn with_session_id(mut self, session_id: &str) -> Self {
397        self.session_id = Some(session_id.to_owned());
398        if let Some(codex) = self.codex.take() {
399            self.codex = Some(codex.with_session_id(session_id));
400        }
401        self
402    }
403
404    pub fn api_key(&self) -> String {
405        self.codex
406            .as_ref()
407            .map(CodexProvider::active_access)
408            .unwrap_or_else(|| self.api_key.clone())
409    }
410
411    pub fn api_key_env(&self) -> &str {
412        &self.api_key_env
413    }
414
415    pub(crate) fn models(&self) -> Result<Vec<ProviderModel>, ProviderError> {
416        if let Some(codex) = &self.codex {
417            return Ok(codex.models());
418        }
419        let base_url = self
420            .endpoint
421            .strip_suffix("/chat/completions")
422            .ok_or_else(|| ProviderError::new("invalid provider endpoint"))?;
423        let response = self
424            .client
425            .get(format!("{base_url}/models"))
426            .bearer_auth(&self.api_key)
427            .timeout(MODEL_METADATA_TIMEOUT)
428            .send()
429            .map_err(|_| ProviderError::new("unable to load provider models"))?;
430        if !response.status().is_success() {
431            return Err(ProviderError::new("unable to load provider models"));
432        }
433        let bytes = response
434            .bytes()
435            .map_err(|_| ProviderError::new("unable to load provider models"))?;
436        if bytes.len() > MAX_MODEL_METADATA_BYTES {
437            return Err(ProviderError::new(
438                "provider model catalog exceeded the response limit",
439            ));
440        }
441        let payload: Value = serde_json::from_slice(&bytes)
442            .map_err(|_| ProviderError::new("invalid provider model catalog"))?;
443        let models = payload
444            .get("data")
445            .and_then(Value::as_array)
446            .ok_or_else(|| ProviderError::new("invalid provider model catalog"))?;
447        let mut result = models
448            .iter()
449            .filter_map(|entry| {
450                let id = entry
451                    .get("id")
452                    .or_else(|| entry.get("name"))
453                    .and_then(Value::as_str)?
454                    .trim();
455                if id.is_empty() {
456                    return None;
457                }
458                let efforts = model_efforts(entry);
459                Some(ProviderModel {
460                    id: id.to_owned(),
461                    efforts,
462                })
463            })
464            .collect::<Vec<_>>();
465        result.sort_by(|left, right| left.id.cmp(&right.id));
466        result.dedup_by(|left, right| left.id == right.id);
467        Ok(result)
468    }
469
470    /// Query the OpenAI-compatible model catalog for the configured model's
471    /// context window. Providers that do not expose context metadata simply
472    /// return `None`; this lookup is only used by the interactive statusline.
473    pub(crate) fn context_window(&self) -> Option<usize> {
474        if let Some(codex) = &self.codex {
475            return codex.context_window();
476        }
477        let base_url = self.endpoint.strip_suffix("/chat/completions")?;
478        let response = self
479            .client
480            .get(format!("{base_url}/models"))
481            .bearer_auth(&self.api_key)
482            .timeout(MODEL_METADATA_TIMEOUT)
483            .send()
484            .ok()?;
485        if !response.status().is_success() {
486            return None;
487        }
488        let bytes = response.bytes().ok()?;
489        if bytes.len() > MAX_MODEL_METADATA_BYTES {
490            return None;
491        }
492        let payload: Value = serde_json::from_slice(&bytes).ok()?;
493        context_window_from_models(&payload, &self.model)
494    }
495
496    pub fn stream_chat(
497        &self,
498        messages: &[ChatMessage],
499        on_text: &mut dyn FnMut(&str) -> io::Result<()>,
500    ) -> Result<ProviderTurn, ProviderError> {
501        let cancellation = CancellationToken::new();
502        self.stream_chat_cancellable_with_options(messages, on_text, &cancellation, true)
503    }
504
505    /// Generate an internal compaction summary without exposing `cmd` to the
506    /// summarization request or emitting its text as a normal assistant delta.
507    pub(crate) fn summarize(
508        &self,
509        messages: &[ChatMessage],
510        cancellation: &CancellationToken,
511    ) -> Result<String, ProviderError> {
512        let mut ignored = |_text: &str| Ok(());
513        let turn =
514            self.stream_chat_cancellable_with_options(messages, &mut ignored, cancellation, false)?;
515        if !turn.tool_calls.is_empty() {
516            return Err(ProviderError::new(
517                "compaction summary requested an unsupported tool",
518            ));
519        }
520        if turn.content.trim().is_empty() {
521            return Err(ProviderError::new("compaction summary was empty"));
522        }
523        Ok(turn.content)
524    }
525
526    /// Stream through an async response so cancellation can drop the pending
527    /// socket read instead of waiting for the blocking client's timeout.
528    #[allow(dead_code)]
529    pub(crate) fn stream_chat_cancellable(
530        &self,
531        messages: &[ChatMessage],
532        on_text: &mut dyn FnMut(&str) -> io::Result<()>,
533        cancellation: &CancellationToken,
534    ) -> Result<ProviderTurn, ProviderError> {
535        self.stream_chat_cancellable_with_options(messages, on_text, cancellation, true)
536    }
537
538    pub(crate) fn stream_chat_cancellable_with_options(
539        &self,
540        messages: &[ChatMessage],
541        on_text: &mut dyn FnMut(&str) -> io::Result<()>,
542        cancellation: &CancellationToken,
543        include_tools: bool,
544    ) -> Result<ProviderTurn, ProviderError> {
545        let mut on_event = |event| match event {
546            ProviderStreamEvent::ReasoningStarted => Ok(()),
547            ProviderStreamEvent::Text(text) => on_text(&text),
548        };
549        self.stream_chat_cancellable_with_options_and_events(
550            messages,
551            &mut on_event,
552            cancellation,
553            include_tools,
554        )
555    }
556
557    pub(crate) fn stream_chat_cancellable_with_options_and_events(
558        &self,
559        messages: &[ChatMessage],
560        on_event: &mut dyn FnMut(ProviderStreamEvent) -> io::Result<()>,
561        cancellation: &CancellationToken,
562        include_tools: bool,
563    ) -> Result<ProviderTurn, ProviderError> {
564        if let Some(codex) = &self.codex {
565            let mut on_event = |event| on_event(event);
566            return codex.stream_chat(messages, &mut on_event, cancellation, include_tools);
567        }
568        let runtime = tokio::runtime::Builder::new_current_thread()
569            .enable_all()
570            .build()
571            .map_err(|_| ProviderError::new("unable to initialize provider runtime"))?;
572        runtime.block_on(self.stream_chat_async_with_retries(
573            messages,
574            on_event,
575            cancellation,
576            include_tools,
577        ))
578    }
579
580    async fn stream_chat_async_with_retries(
581        &self,
582        messages: &[ChatMessage],
583        on_event: &mut dyn FnMut(ProviderStreamEvent) -> io::Result<()>,
584        cancellation: &CancellationToken,
585        include_tools: bool,
586    ) -> Result<ProviderTurn, ProviderError> {
587        for attempt in 0..=PROVIDER_RETRY_COUNT {
588            match self
589                .stream_chat_async_once(messages, on_event, cancellation, include_tools)
590                .await
591            {
592                Err(error) if error.is_retryable() && attempt < PROVIDER_RETRY_COUNT => {
593                    if cancellation.is_cancelled() {
594                        return Err(ProviderError::cancelled(empty_turn()));
595                    }
596                    tokio::time::sleep(PROVIDER_RETRY_BACKOFF).await;
597                }
598                result => return result,
599            }
600        }
601        unreachable!("provider retry loop must return");
602    }
603
604    async fn stream_chat_async_once(
605        &self,
606        messages: &[ChatMessage],
607        on_event: &mut dyn FnMut(ProviderStreamEvent) -> io::Result<()>,
608        cancellation: &CancellationToken,
609        include_tools: bool,
610    ) -> Result<ProviderTurn, ProviderError> {
611        if cancellation.is_cancelled() {
612            return Err(ProviderError::cancelled(ProviderTurn {
613                content: String::new(),
614                tool_calls: Vec::new(),
615                reasoning_details: Vec::new(),
616            }));
617        }
618        let request = chat_request(
619            &self.model,
620            messages,
621            &self.effort,
622            include_tools,
623            self.is_openrouter
624                .then_some(self.session_id.as_deref())
625                .flatten(),
626        );
627        let request = self
628            .async_client
629            .post(&self.endpoint)
630            .bearer_auth(&self.api_key)
631            .header("accept", "text/event-stream")
632            .json(&request)
633            .send();
634        let mut request = Box::pin(request);
635        let mut response = loop {
636            if cancellation.is_cancelled() {
637                return Err(ProviderError::cancelled(ProviderTurn {
638                    content: String::new(),
639                    tool_calls: Vec::new(),
640                    reasoning_details: Vec::new(),
641                }));
642            }
643            match tokio::time::timeout(CANCELLATION_POLL_INTERVAL, request.as_mut()).await {
644                Ok(response) => {
645                    break response.map_err(|error| {
646                        reqwest_failure("provider request failed", error, &self.api_key, true)
647                    })?;
648                }
649                Err(_) => continue,
650            }
651        };
652        if !response.status().is_success() {
653            let status = response.status().as_u16();
654            let error = if transient_http_status(status) {
655                ProviderError::retryable(format!("provider returned HTTP status {status}"))
656            } else {
657                ProviderError::new(format!("provider returned HTTP status {status}"))
658            };
659            return Err(error);
660        }
661
662        let mut accumulator = ProviderAccumulator::default();
663        let mut decoder = SseDecoder::default();
664        loop {
665            if cancellation.is_cancelled() {
666                return Err(ProviderError::cancelled(accumulator.partial_turn()));
667            }
668            let chunk =
669                match tokio::time::timeout(CANCELLATION_POLL_INTERVAL, response.chunk()).await {
670                    Ok(chunk) => {
671                        let retry_before_payload = !decoder.result.received_payload;
672                        chunk.map_err(|error| {
673                            reqwest_failure(
674                                "provider stream read failed",
675                                error,
676                                &self.api_key,
677                                retry_before_payload,
678                            )
679                        })?
680                    }
681                    Err(_) => continue,
682                };
683            let Some(chunk) = chunk else {
684                break;
685            };
686            let done = decoder.feed(&chunk, &mut |data| {
687                accumulator.on_data(data, &self.api_key, on_event)
688            })?;
689            if done {
690                break;
691            }
692        }
693        if cancellation.is_cancelled() {
694            return Err(ProviderError::cancelled(accumulator.partial_turn()));
695        }
696        let parse_result =
697            decoder.finish(&mut |data| accumulator.on_data(data, &self.api_key, on_event))?;
698        if !parse_result.received_payload {
699            return Err(ProviderError::new(
700                "provider stream contained no valid payload",
701            ));
702        }
703        if !parse_result.received_done {
704            return Err(ProviderError::new("provider stream ended before [DONE]"));
705        }
706        accumulator.finish()
707    }
708}
709
710#[derive(Debug, Clone, Default)]
711struct PartialToolCall {
712    id: String,
713    name: String,
714    arguments: String,
715}
716
717#[derive(Debug, Default)]
718struct ProviderAccumulator {
719    content: String,
720    tool_calls: BTreeMap<usize, PartialToolCall>,
721    reasoning_details: Vec<Value>,
722    reasoning_details_bytes: usize,
723    tool_argument_bytes: usize,
724    finish_reason: Option<String>,
725    reasoning_started: bool,
726}
727
728impl ProviderAccumulator {
729    fn on_data(
730        &mut self,
731        data: Value,
732        api_key: &str,
733        on_event: &mut dyn FnMut(ProviderStreamEvent) -> io::Result<()>,
734    ) -> Result<(), ProviderError> {
735        if let Some(message) = provider_error_message(&data) {
736            return Err(ProviderError::new(format!(
737                "provider stream error: {}",
738                redact_secret(message, Some(api_key))
739            )));
740        }
741        let Some(choice) = data
742            .get("choices")
743            .and_then(Value::as_array)
744            .and_then(|choices| choices.first())
745        else {
746            return Ok(());
747        };
748        if let Some(reason) = validate_finish_reason(choice)? {
749            self.finish_reason = Some(reason.to_owned());
750        }
751        let Some(delta) = choice.get("delta") else {
752            return Ok(());
753        };
754        let received_reasoning = append_reasoning_details(
755            &mut self.reasoning_details,
756            &mut self.reasoning_details_bytes,
757            delta,
758        )?;
759        if received_reasoning && !self.reasoning_started {
760            self.reasoning_started = true;
761            on_event(ProviderStreamEvent::ReasoningStarted)
762                .map_err(|_| ProviderError::new("unable to emit reasoning state"))?;
763        }
764        if let Some(text) = delta.get("content").and_then(Value::as_str) {
765            if self.content.len().saturating_add(text.len()) > MAX_PROVIDER_CONTENT_BYTES {
766                return Err(ProviderError::new(
767                    "provider assistant content exceeded the response limit",
768                ));
769            }
770            self.content.push_str(text);
771            on_event(ProviderStreamEvent::Text(text.to_owned()))
772                .map_err(|_| ProviderError::new("unable to emit assistant delta"))?;
773        }
774        if let Some(calls) = delta.get("tool_calls").and_then(Value::as_array) {
775            for (position, call) in calls.iter().enumerate() {
776                let index = call
777                    .get("index")
778                    .and_then(Value::as_u64)
779                    .map_or(position, |index| index as usize);
780                let partial = self.tool_calls.entry(index).or_default();
781                if let Some(id) = call.get("id").and_then(Value::as_str) {
782                    append_provider_field(
783                        &mut partial.id,
784                        id,
785                        MAX_PROVIDER_TOOL_CALL_ID_BYTES,
786                        "provider tool-call id exceeded the response limit",
787                    )?;
788                }
789                if let Some(function) = call.get("function") {
790                    if let Some(name) = function.get("name").and_then(Value::as_str) {
791                        append_provider_field(
792                            &mut partial.name,
793                            name,
794                            MAX_PROVIDER_TOOL_NAME_BYTES,
795                            "provider tool-call name exceeded the response limit",
796                        )?;
797                    }
798                    if let Some(arguments) = function.get("arguments").and_then(Value::as_str) {
799                        if self.tool_argument_bytes.saturating_add(arguments.len())
800                            > MAX_PROVIDER_TOOL_ARGUMENT_BYTES
801                        {
802                            return Err(ProviderError::new(
803                                "provider tool arguments exceeded the response limit",
804                            ));
805                        }
806                        self.tool_argument_bytes += arguments.len();
807                        partial.arguments.push_str(arguments);
808                    }
809                }
810            }
811        }
812        if let Some(function_call) = delta.get("function_call") {
813            let partial = self.tool_calls.entry(0).or_default();
814            if let Some(name) = function_call.get("name").and_then(Value::as_str) {
815                append_provider_field(
816                    &mut partial.name,
817                    name,
818                    MAX_PROVIDER_TOOL_NAME_BYTES,
819                    "provider tool-call name exceeded the response limit",
820                )?;
821            }
822            if let Some(arguments) = function_call.get("arguments").and_then(Value::as_str) {
823                if self.tool_argument_bytes.saturating_add(arguments.len())
824                    > MAX_PROVIDER_TOOL_ARGUMENT_BYTES
825                {
826                    return Err(ProviderError::new(
827                        "provider tool arguments exceeded the response limit",
828                    ));
829                }
830                self.tool_argument_bytes += arguments.len();
831                partial.arguments.push_str(arguments);
832            }
833        }
834        Ok(())
835    }
836
837    fn partial_turn(&self) -> ProviderTurn {
838        ProviderTurn {
839            content: self.content.clone(),
840            tool_calls: self
841                .tool_calls
842                .iter()
843                .map(|(index, partial)| ChatToolCall {
844                    id: if partial.id.is_empty() {
845                        format!("call_{index}")
846                    } else {
847                        partial.id.clone()
848                    },
849                    name: partial.name.clone(),
850                    arguments: partial.arguments.clone(),
851                })
852                .collect(),
853            reasoning_details: self.reasoning_details.clone(),
854        }
855    }
856
857    fn finish(self) -> Result<ProviderTurn, ProviderError> {
858        let tool_calls = self
859            .tool_calls
860            .into_iter()
861            .map(|(index, partial)| ChatToolCall {
862                id: if partial.id.is_empty() {
863                    format!("call_{index}")
864                } else {
865                    partial.id
866                },
867                name: partial.name,
868                arguments: partial.arguments,
869            })
870            .collect::<Vec<_>>();
871        if let Some(reason) = self.finish_reason.as_deref() {
872            if !tool_calls.is_empty() && !matches!(reason, "tool_calls" | "function_call") {
873                return Err(ProviderError::new(
874                    "provider tool calls ended with an incompatible finish reason",
875                ));
876            }
877            if tool_calls.is_empty() && matches!(reason, "tool_calls" | "function_call") {
878                return Err(ProviderError::new(
879                    "provider reported tool completion without a tool call",
880                ));
881            }
882        }
883        if self.content.is_empty() && tool_calls.is_empty() {
884            return Err(ProviderError::new(
885                "provider stream contained no assistant content or tool calls",
886            ));
887        }
888        Ok(ProviderTurn {
889            content: self.content,
890            tool_calls,
891            reasoning_details: self.reasoning_details,
892        })
893    }
894}
895
896fn append_reasoning_details(
897    target: &mut Vec<Value>,
898    serialized_bytes: &mut usize,
899    delta: &Value,
900) -> Result<bool, ProviderError> {
901    let Some(details) = delta.get("reasoning_details").and_then(Value::as_array) else {
902        return Ok(false);
903    };
904    if details.is_empty() {
905        return Ok(false);
906    }
907    let serialized_delta = serde_json::to_vec(details)
908        .map_err(|_| ProviderError::new("provider reasoning details could not be serialized"))?;
909    let combined_bytes = if target.is_empty() {
910        serialized_delta.len()
911    } else {
912        serialized_bytes
913            .saturating_add(serialized_delta.len())
914            .saturating_sub(1)
915    };
916    if combined_bytes > MAX_PROVIDER_REASONING_DETAILS_BYTES {
917        return Err(ProviderError::new(
918            "provider reasoning details exceeded the response limit",
919        ));
920    }
921    target.extend(details.iter().cloned());
922    *serialized_bytes = combined_bytes;
923    Ok(true)
924}
925
926fn append_provider_field(
927    target: &mut String,
928    fragment: &str,
929    limit: usize,
930    error_message: &str,
931) -> Result<(), ProviderError> {
932    if target.len().saturating_add(fragment.len()) > limit {
933        return Err(ProviderError::new(error_message));
934    }
935    target.push_str(fragment);
936    Ok(())
937}
938
939fn provider_error_message(data: &Value) -> Option<&str> {
940    let error = data.get("error")?;
941    let message = if let Some(message) = error.get("message").and_then(Value::as_str) {
942        message
943    } else if let Some(message) = error.as_str() {
944        message
945    } else {
946        return Some("provider returned an error payload");
947    };
948    if message.len() > MAX_PROVIDER_ERROR_BYTES {
949        Some("provider error text exceeded the response limit")
950    } else {
951        Some(message)
952    }
953}
954
955#[derive(Debug, Default)]
956struct SseDecoder {
957    line: Vec<u8>,
958    data_lines: Vec<String>,
959    data_event_bytes: usize,
960    stream_bytes: usize,
961    result: SseParseResult,
962    done: bool,
963}
964
965impl SseDecoder {
966    fn feed<F>(&mut self, bytes: &[u8], on_data: &mut F) -> Result<bool, ProviderError>
967    where
968        F: FnMut(Value) -> Result<(), ProviderError>,
969    {
970        if self.stream_bytes.saturating_add(bytes.len()) > MAX_SSE_STREAM_BYTES {
971            return Err(ProviderError::new(
972                "provider SSE stream exceeded the response limit",
973            ));
974        }
975        self.stream_bytes += bytes.len();
976        for byte in bytes {
977            if self.done {
978                break;
979            }
980            if *byte == b'\n' {
981                let line = std::mem::take(&mut self.line);
982                if self.process_line(&line, on_data)? {
983                    return Ok(true);
984                }
985            } else {
986                self.line.push(*byte);
987                if self.line.len() > MAX_SSE_LINE_BYTES {
988                    return Err(ProviderError::new(
989                        "provider SSE line exceeded the response limit",
990                    ));
991                }
992            }
993        }
994        Ok(self.done)
995    }
996
997    fn finish<F>(&mut self, on_data: &mut F) -> Result<SseParseResult, ProviderError>
998    where
999        F: FnMut(Value) -> Result<(), ProviderError>,
1000    {
1001        if !self.line.is_empty() && !self.done {
1002            let line = std::mem::take(&mut self.line);
1003            self.process_line(&line, on_data)?;
1004        }
1005        if !self.done {
1006            self.dispatch_data(on_data)?;
1007        }
1008        Ok(self.result)
1009    }
1010
1011    fn process_line<F>(&mut self, raw_line: &[u8], on_data: &mut F) -> Result<bool, ProviderError>
1012    where
1013        F: FnMut(Value) -> Result<(), ProviderError>,
1014    {
1015        let line = std::str::from_utf8(raw_line)
1016            .map_err(|_| ProviderError::new("provider stream contained invalid UTF-8"))?
1017            .trim_end_matches('\r');
1018        if line.is_empty() {
1019            return self.dispatch_data(on_data);
1020        }
1021        if line.starts_with(':') {
1022            return Ok(false);
1023        }
1024        let (field, value) = line
1025            .split_once(':')
1026            .map_or((line, ""), |(field, value)| (field, value));
1027        if field == "data" {
1028            let value = value.strip_prefix(' ').unwrap_or(value);
1029            let separator_bytes = (!self.data_lines.is_empty()) as usize;
1030            let added_bytes = separator_bytes.saturating_add(value.len());
1031            if self.data_event_bytes.saturating_add(added_bytes) > MAX_SSE_EVENT_BYTES {
1032                return Err(ProviderError::new(
1033                    "provider SSE data event exceeded the response limit",
1034                ));
1035            }
1036            if self.data_lines.len() >= MAX_SSE_DATA_LINES {
1037                return Err(ProviderError::new(
1038                    "provider SSE data line count exceeded the response limit",
1039                ));
1040            }
1041            self.data_event_bytes += added_bytes;
1042            self.data_lines.push(value.to_owned());
1043        }
1044        Ok(false)
1045    }
1046
1047    fn dispatch_data<F>(&mut self, on_data: &mut F) -> Result<bool, ProviderError>
1048    where
1049        F: FnMut(Value) -> Result<(), ProviderError>,
1050    {
1051        if self.data_lines.is_empty() {
1052            self.data_event_bytes = 0;
1053            return Ok(false);
1054        }
1055        let data = self.data_lines.join("\n");
1056        self.data_lines.clear();
1057        self.data_event_bytes = 0;
1058        if data.trim().is_empty() {
1059            return Ok(false);
1060        }
1061        if data == "[DONE]" {
1062            self.result.received_done = true;
1063            self.done = true;
1064            return Ok(true);
1065        }
1066        let value: Value = serde_json::from_str(&data)
1067            .map_err(|_| ProviderError::new("provider sent malformed SSE data"))?;
1068        self.result.received_payload = true;
1069        on_data(value)?;
1070        Ok(false)
1071    }
1072}
1073
1074#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
1075pub struct SseParseResult {
1076    pub received_payload: bool,
1077    pub received_done: bool,
1078}
1079
1080pub fn parse_sse<R, F>(reader: &mut R, mut on_data: F) -> Result<SseParseResult, ProviderError>
1081where
1082    R: BufRead,
1083    F: FnMut(Value) -> Result<(), ProviderError>,
1084{
1085    let mut data_lines = Vec::new();
1086    let mut data_event_bytes = 0;
1087    let mut stream_bytes: usize = 0;
1088    let mut result = SseParseResult::default();
1089    let mut line = Vec::with_capacity(MAX_SSE_LINE_BYTES);
1090    loop {
1091        let (has_line, line_bytes) = match read_sse_line(reader, &mut line) {
1092            Ok(result) => result,
1093            Err(mut error) => {
1094                if result.received_payload {
1095                    error.retryable = false;
1096                }
1097                return Err(error);
1098            }
1099        };
1100        if stream_bytes.saturating_add(line_bytes) > MAX_SSE_STREAM_BYTES {
1101            return Err(ProviderError::new(
1102                "provider SSE stream exceeded the response limit",
1103            ));
1104        }
1105        stream_bytes += line_bytes;
1106        if !has_line {
1107            if !data_lines.is_empty() {
1108                dispatch_data(
1109                    &mut data_lines,
1110                    &mut data_event_bytes,
1111                    &mut on_data,
1112                    &mut result,
1113                )?;
1114            }
1115            return Ok(result);
1116        }
1117
1118        let line = std::str::from_utf8(&line)
1119            .map_err(|_| ProviderError::new("provider stream contained invalid UTF-8"))?
1120            .trim_end_matches('\r');
1121        if line.is_empty() {
1122            if dispatch_data(
1123                &mut data_lines,
1124                &mut data_event_bytes,
1125                &mut on_data,
1126                &mut result,
1127            )? {
1128                return Ok(result);
1129            }
1130            continue;
1131        }
1132        if line.starts_with(':') {
1133            continue;
1134        }
1135        let (field, value) = line
1136            .split_once(':')
1137            .map_or((line, ""), |(field, value)| (field, value));
1138        if field == "data" {
1139            let value = value.strip_prefix(' ').unwrap_or(value);
1140            let separator_bytes = (!data_lines.is_empty()) as usize;
1141            let added_bytes = separator_bytes.saturating_add(value.len());
1142            if data_event_bytes.saturating_add(added_bytes) > MAX_SSE_EVENT_BYTES {
1143                return Err(ProviderError::new(
1144                    "provider SSE data event exceeded the response limit",
1145                ));
1146            }
1147            if data_lines.len() >= MAX_SSE_DATA_LINES {
1148                return Err(ProviderError::new(
1149                    "provider SSE data line count exceeded the response limit",
1150                ));
1151            }
1152            data_event_bytes += added_bytes;
1153            data_lines.push(value.to_owned());
1154        }
1155    }
1156}
1157
1158fn read_sse_line<R: BufRead>(
1159    reader: &mut R,
1160    line: &mut Vec<u8>,
1161) -> Result<(bool, usize), ProviderError> {
1162    line.clear();
1163    let mut consumed_bytes = 0;
1164    loop {
1165        let buffer = reader.fill_buf().map_err(|error| {
1166            ProviderError::retryable(format!("provider stream read failed: {error}"))
1167        })?;
1168        if buffer.is_empty() {
1169            return Ok((!line.is_empty(), consumed_bytes));
1170        }
1171
1172        let newline = buffer.iter().position(|byte| *byte == b'\n');
1173        let chunk_length = newline.unwrap_or(buffer.len());
1174        if line.len().saturating_add(chunk_length) > MAX_SSE_LINE_BYTES {
1175            return Err(ProviderError::new(
1176                "provider SSE line exceeded the response limit",
1177            ));
1178        }
1179        line.extend_from_slice(&buffer[..chunk_length]);
1180        let consumed = newline.map_or(chunk_length, |index| index + 1);
1181        reader.consume(consumed);
1182        consumed_bytes += consumed;
1183        if newline.is_some() {
1184            return Ok((true, consumed_bytes));
1185        }
1186    }
1187}
1188
1189fn dispatch_data<F>(
1190    data_lines: &mut Vec<String>,
1191    data_event_bytes: &mut usize,
1192    on_data: &mut F,
1193    result: &mut SseParseResult,
1194) -> Result<bool, ProviderError>
1195where
1196    F: FnMut(Value) -> Result<(), ProviderError>,
1197{
1198    if data_lines.is_empty() {
1199        *data_event_bytes = 0;
1200        return Ok(false);
1201    }
1202    let data = data_lines.join("\n");
1203    data_lines.clear();
1204    *data_event_bytes = 0;
1205    if data.trim().is_empty() {
1206        return Ok(false);
1207    }
1208    if data == "[DONE]" {
1209        result.received_done = true;
1210        return Ok(true);
1211    }
1212    let value: Value = serde_json::from_str(&data)
1213        .map_err(|_| ProviderError::new("provider sent malformed SSE data"))?;
1214    result.received_payload = true;
1215    on_data(value)?;
1216    Ok(false)
1217}
1218
1219fn validate_finish_reason(choice: &Value) -> Result<Option<&str>, ProviderError> {
1220    let Some(reason) = choice.get("finish_reason") else {
1221        return Ok(None);
1222    };
1223    if reason.is_null() {
1224        return Ok(None);
1225    }
1226    match reason.as_str() {
1227        Some("stop") | Some("tool_calls") | Some("function_call") => Ok(reason.as_str()),
1228        Some("length") | Some("content_filter") => Err(ProviderError::new(
1229            "provider response ended before completion",
1230        )),
1231        Some(_) | None => Err(ProviderError::new(
1232            "provider response has an unsupported finish reason",
1233        )),
1234    }
1235}
1236
1237#[cfg(test)]
1238mod tests {
1239    use super::*;
1240    use std::io::{BufRead, BufReader, Cursor, Read, Write};
1241    use std::net::{TcpListener, TcpStream};
1242    use std::sync::mpsc;
1243    use std::thread;
1244    use std::time::{Duration, Instant};
1245
1246    #[test]
1247    fn openrouter_requests_include_session_and_app_metadata() {
1248        let request = chat_request(
1249            "model",
1250            &[ChatMessage::user("hello".to_owned())],
1251            &None,
1252            true,
1253            Some("lucy-session"),
1254        );
1255        assert_eq!(request["session_id"], "lucy-session");
1256        assert!(is_openrouter_base_url("https://openrouter.ai/api/v1"));
1257        assert!(!is_openrouter_base_url("https://example.com/v1"));
1258
1259        let headers = provider_headers(true);
1260        assert_eq!(headers.get("x-openrouter-title").unwrap(), "Lucy");
1261        assert_eq!(headers.get("http-referer").unwrap(), "https://lucyna.run");
1262
1263        let compact = chat_request(
1264            "model",
1265            &[ChatMessage::user("hello".to_owned())],
1266            &None,
1267            false,
1268            Some("lucy-session"),
1269        );
1270        assert_eq!(compact["session_id"], request["session_id"]);
1271    }
1272
1273    #[test]
1274    fn compatible_requests_omit_provider_specific_session_metadata() {
1275        let request = chat_request(
1276            "model",
1277            &[ChatMessage::user("hello".to_owned())],
1278            &None,
1279            true,
1280            None,
1281        );
1282        assert!(request.get("session_id").is_none());
1283
1284        let headers = provider_headers(false);
1285        assert!(headers.get("x-openrouter-title").is_none());
1286        assert!(headers.get("http-referer").is_none());
1287    }
1288
1289    #[test]
1290    fn retry_policy_only_marks_transient_http_statuses() {
1291        for status in [408, 429, 500, 502, 503, 504] {
1292            assert!(transient_http_status(status));
1293        }
1294        for status in [400, 401, 403, 404, 422] {
1295            assert!(!transient_http_status(status));
1296        }
1297    }
1298
1299    #[test]
1300    fn normal_requests_expose_only_cmd() {
1301        let request = chat_request(
1302            "model",
1303            &[ChatMessage::user("hello".to_owned())],
1304            &None,
1305            true,
1306            None,
1307        );
1308        let tools = request["tools"].as_array().expect("model tools");
1309        assert_eq!(tools.len(), 1);
1310        assert_eq!(tools[0]["function"]["name"], "cmd");
1311        let background = &tools[0]["function"]["parameters"]["properties"]["background"];
1312        assert_eq!(background["type"], "boolean");
1313        assert_eq!(background["default"], false);
1314    }
1315
1316    #[test]
1317    fn compaction_request_does_not_include_tools() {
1318        let normal = chat_request(
1319            "model",
1320            &[ChatMessage::user("hello".to_owned())],
1321            &None,
1322            true,
1323            None,
1324        );
1325        let compact = chat_request(
1326            "model",
1327            &[ChatMessage::user("hello".to_owned())],
1328            &None,
1329            false,
1330            None,
1331        );
1332
1333        assert!(normal.get("tools").is_some());
1334        assert!(normal.get("max_tokens").is_none());
1335        assert!(compact.get("tools").is_none());
1336        assert_eq!(compact["max_tokens"], COMPACTION_MAX_SUMMARY_TOKENS);
1337    }
1338
1339    #[test]
1340    fn model_catalog_reads_nested_and_compatible_effort_metadata() {
1341        let openrouter = serde_json::json!({
1342            "reasoning": {
1343                "supported_efforts": ["max", "xhigh", "high", "medium", "low", "none"]
1344            }
1345        });
1346        assert_eq!(
1347            model_efforts(&openrouter),
1348            Some(vec![
1349                "max".to_owned(),
1350                "xhigh".to_owned(),
1351                "high".to_owned(),
1352                "medium".to_owned(),
1353                "low".to_owned(),
1354                "none".to_owned(),
1355            ])
1356        );
1357
1358        let compatible = serde_json::json!({
1359            "supported_reasoning_efforts": ["light", "medium", "max", "light", ""]
1360        });
1361        assert_eq!(
1362            model_efforts(&compatible),
1363            Some(vec![
1364                "light".to_owned(),
1365                "medium".to_owned(),
1366                "max".to_owned()
1367            ])
1368        );
1369        assert_eq!(model_efforts(&serde_json::json!({})), None);
1370    }
1371
1372    #[test]
1373    fn model_catalog_context_window_matches_configured_model() {
1374        let payload = serde_json::json!({
1375            "data": [
1376                {"id": "other", "context_length": 8_000},
1377                {"id": "provider/model", "context_length": 128_000}
1378            ]
1379        });
1380
1381        assert_eq!(
1382            context_window_from_models(&payload, "provider/model"),
1383            Some(128_000)
1384        );
1385        assert_eq!(context_window_from_models(&payload, "missing"), None);
1386    }
1387
1388    #[test]
1389    fn model_catalog_context_window_accepts_provider_fallback_fields() {
1390        let payload = serde_json::json!({
1391            "data": [{
1392                "id": "provider/model",
1393                "top_provider": {"context_length": 64_000}
1394            }]
1395        });
1396
1397        assert_eq!(
1398            context_window_from_models(&payload, "provider/model"),
1399            Some(64_000)
1400        );
1401    }
1402
1403    #[test]
1404    fn parses_sse_comments_multiline_data_and_done() {
1405        let stream = b": keep-alive\n\ndata: \n\ndata: {\"choices\":[]\ndata: }\n\ndata: [DONE]\n";
1406        let mut values = Vec::new();
1407        let result = parse_sse(&mut Cursor::new(stream), |value| {
1408            values.push(value);
1409            Ok(())
1410        })
1411        .expect("SSE");
1412        assert!(result.received_payload);
1413        assert!(result.received_done);
1414        assert_eq!(values.len(), 1);
1415        assert!(values[0]["choices"].is_array());
1416    }
1417
1418    #[test]
1419    fn parses_text_and_fragmented_tool_calls() {
1420        let first = serde_json::json!({
1421            "choices": [{"delta": {"content": "hi"}}]
1422        });
1423        let second = serde_json::json!({
1424            "choices": [{
1425                "delta": {
1426                    "tool_calls": [{
1427                        "index": 0,
1428                        "id": "c1",
1429                        "function": {"name": "cmd", "arguments": "{command:"}
1430                    }]
1431                }
1432            }]
1433        });
1434        let third = serde_json::json!({
1435            "choices": [{
1436                "delta": {
1437                    "tool_calls": [{
1438                        "index": 0,
1439                        "function": {"arguments": "pwd}"}
1440                    }]
1441                }
1442            }]
1443        });
1444        let stream =
1445            format!("data: {first}\n\n data: {second}\n\ndata: {third}\n\ndata: [DONE]\n\n")
1446                .replace(" data:", "data:");
1447        let mut content = String::new();
1448        let mut calls = BTreeMap::<usize, PartialToolCall>::new();
1449        let result = parse_sse(&mut Cursor::new(stream.as_bytes()), |value| {
1450            let choice = &value["choices"][0];
1451            let delta = &choice["delta"];
1452            if let Some(text) = delta["content"].as_str() {
1453                content.push_str(text);
1454            }
1455            if let Some(tool_calls) = delta["tool_calls"].as_array() {
1456                for call in tool_calls {
1457                    let index = call["index"].as_u64().expect("index") as usize;
1458                    let partial = calls.entry(index).or_default();
1459                    partial.id.push_str(call["id"].as_str().unwrap_or(""));
1460                    partial
1461                        .name
1462                        .push_str(call["function"]["name"].as_str().unwrap_or(""));
1463                    partial
1464                        .arguments
1465                        .push_str(call["function"]["arguments"].as_str().unwrap_or(""));
1466                }
1467            }
1468            Ok(())
1469        })
1470        .expect("SSE");
1471        assert!(result.received_payload);
1472        assert!(result.received_done);
1473        assert_eq!(content, "hi");
1474        assert_eq!(calls[&0].id, "c1");
1475        assert_eq!(calls[&0].name, "cmd");
1476        assert_eq!(calls[&0].arguments, "{command:pwd}");
1477    }
1478
1479    #[test]
1480    fn cancellable_accumulator_accepts_more_than_sixty_four_tool_calls() {
1481        let mut accumulator = ProviderAccumulator::default();
1482        let tool_calls = (0..65)
1483            .map(|index| {
1484                serde_json::json!({
1485                    "index": index,
1486                    "id": format!("call-{index}"),
1487                    "function": {
1488                        "name": "cmd",
1489                        "arguments": "{\"command\":\"true\"}"
1490                    }
1491                })
1492            })
1493            .collect::<Vec<_>>();
1494        accumulator
1495            .on_data(
1496                serde_json::json!({
1497                    "choices": [{
1498                        "delta": {"tool_calls": tool_calls},
1499                        "finish_reason": "tool_calls"
1500                    }]
1501                }),
1502                "provider-secret",
1503                &mut |_| Ok(()),
1504            )
1505            .expect("tool-call chunk");
1506
1507        let turn = accumulator.finish().expect("provider turn");
1508        assert_eq!(turn.tool_calls.len(), 65);
1509    }
1510
1511    #[test]
1512    fn reasoning_stream_event_is_emitted_once_before_assistant_text() {
1513        let mut accumulator = ProviderAccumulator::default();
1514        let mut events = Vec::new();
1515        let mut on_event = |event| {
1516            match event {
1517                ProviderStreamEvent::ReasoningStarted => events.push("started".to_owned()),
1518                ProviderStreamEvent::Text(text) => events.push(text),
1519            }
1520            Ok(())
1521        };
1522
1523        accumulator
1524            .on_data(
1525                serde_json::json!({
1526                    "choices": [{
1527                        "delta": {
1528                            "reasoning_details": [{"type": "reasoning.text", "text": "thinking"}]
1529                        }
1530                    }]
1531                }),
1532                "provider-secret",
1533                &mut on_event,
1534            )
1535            .expect("reasoning chunk");
1536        accumulator
1537            .on_data(
1538                serde_json::json!({
1539                    "choices": [{"delta": {"content": "answer"}}]
1540                }),
1541                "provider-secret",
1542                &mut on_event,
1543            )
1544            .expect("answer chunk");
1545
1546        assert_eq!(events, vec!["started".to_owned(), "answer".to_owned()]);
1547    }
1548
1549    #[test]
1550    fn accumulates_reasoning_details_with_fragmented_tool_calls() {
1551        let mut accumulator = ProviderAccumulator::default();
1552        accumulator
1553            .on_data(
1554                serde_json::json!({
1555                    "choices": [{
1556                        "delta": {
1557                            "reasoning_details": [{
1558                                "type": "reasoning.text",
1559                                "text": "part one"
1560                            }]
1561                        }
1562                    }]
1563                }),
1564                "provider-secret",
1565                &mut |_| Ok(()),
1566            )
1567            .expect("first provider chunk");
1568        accumulator
1569            .on_data(
1570                serde_json::json!({
1571                    "choices": [{
1572                        "delta": {
1573                            "reasoning_details": [{
1574                                "type": "reasoning.text",
1575                                "text": "part two"
1576                            }],
1577                            "tool_calls": [{
1578                                "index": 0,
1579                                "id": "call-1",
1580                                "function": {
1581                                    "name": "cmd",
1582                                    "arguments": "{\"command\":\"true\"}"
1583                                }
1584                            }]
1585                        },
1586                        "finish_reason": "tool_calls"
1587                    }]
1588                }),
1589                "provider-secret",
1590                &mut |_| Ok(()),
1591            )
1592            .expect("second provider chunk");
1593
1594        let partial = accumulator.partial_turn();
1595        assert_eq!(partial.reasoning_details.len(), 2);
1596        let turn = accumulator.finish().expect("provider turn");
1597        assert_eq!(
1598            turn.reasoning_details,
1599            vec![
1600                json!({"type": "reasoning.text", "text": "part one"}),
1601                json!({"type": "reasoning.text", "text": "part two"}),
1602            ]
1603        );
1604        assert_eq!(turn.tool_calls.len(), 1);
1605        assert_eq!(turn.tool_calls[0].name, "cmd");
1606    }
1607
1608    #[test]
1609    fn accumulates_many_small_reasoning_details_and_rejects_overflow_atomically() {
1610        const FRAGMENT_COUNT: usize = 4096;
1611        let mut details = Vec::new();
1612        let mut serialized_bytes = 0;
1613        let delta = serde_json::json!({
1614            "reasoning_details": [{
1615                "type": "reasoning.text",
1616                "text": "x".repeat(64)
1617            }]
1618        });
1619        for _ in 0..FRAGMENT_COUNT {
1620            append_reasoning_details(&mut details, &mut serialized_bytes, &delta)
1621                .expect("small reasoning detail");
1622        }
1623        assert_eq!(details.len(), FRAGMENT_COUNT);
1624
1625        let first_chunk_delta = serde_json::json!({
1626            "reasoning_details": [{
1627                "type": "reasoning.text",
1628                "text": "x".repeat(500 * 1024)
1629            }]
1630        });
1631        let first_chunk_bytes = serde_json::to_vec(&first_chunk_delta["reasoning_details"])
1632            .expect("first reasoning detail chunk")
1633            .len();
1634        assert!(first_chunk_bytes < MAX_PROVIDER_REASONING_DETAILS_BYTES);
1635        append_reasoning_details(&mut details, &mut serialized_bytes, &first_chunk_delta)
1636            .expect("first individually bounded reasoning detail chunk");
1637        assert!(serialized_bytes <= MAX_PROVIDER_REASONING_DETAILS_BYTES);
1638
1639        let second_chunk_delta = serde_json::json!({
1640            "reasoning_details": [{
1641                "type": "reasoning.text",
1642                "text": "x".repeat(200 * 1024)
1643            }]
1644        });
1645        let second_chunk_bytes = serde_json::to_vec(&second_chunk_delta["reasoning_details"])
1646            .expect("second reasoning detail chunk")
1647            .len();
1648        assert!(second_chunk_bytes < MAX_PROVIDER_REASONING_DETAILS_BYTES);
1649        assert!(
1650            serialized_bytes
1651                .saturating_add(second_chunk_bytes)
1652                .saturating_sub(1)
1653                > MAX_PROVIDER_REASONING_DETAILS_BYTES
1654        );
1655
1656        let prior_details = details.clone();
1657        let prior_bytes = serialized_bytes;
1658        let error =
1659            append_reasoning_details(&mut details, &mut serialized_bytes, &second_chunk_delta)
1660                .expect_err("reasoning details limit");
1661
1662        assert_eq!(
1663            error.to_string(),
1664            "provider reasoning details exceeded the response limit"
1665        );
1666        assert_eq!(details, prior_details);
1667        assert_eq!(serialized_bytes, prior_bytes);
1668        assert_eq!(
1669            serde_json::to_vec(&details)
1670                .expect("accumulated reasoning details")
1671                .len(),
1672            serialized_bytes
1673        );
1674        assert!(serialized_bytes <= MAX_PROVIDER_REASONING_DETAILS_BYTES);
1675    }
1676
1677    #[test]
1678    fn rejects_reasoning_details_that_exceed_the_serialized_response_limit_before_retaining_them() {
1679        let mut accumulator = ProviderAccumulator::default();
1680        let retained = serde_json::json!({
1681            "choices": [{
1682                "delta": {
1683                    "reasoning_details": [{
1684                        "type": "reasoning.text",
1685                        "text": "retained"
1686                    }]
1687                }
1688            }]
1689        });
1690        accumulator
1691            .on_data(retained, "provider-secret", &mut |_| Ok(()))
1692            .expect("details within limit");
1693
1694        let oversized = "x".repeat(MAX_PROVIDER_REASONING_DETAILS_BYTES);
1695        let error = accumulator
1696            .on_data(
1697                serde_json::json!({
1698                    "choices": [{
1699                        "delta": {
1700                            "reasoning_details": [{"text": oversized}]
1701                        }
1702                    }]
1703                }),
1704                "provider-secret",
1705                &mut |_| Ok(()),
1706            )
1707            .expect_err("reasoning details limit");
1708        assert_eq!(
1709            error.to_string(),
1710            "provider reasoning details exceeded the response limit"
1711        );
1712        assert_eq!(
1713            accumulator.reasoning_details,
1714            vec![serde_json::json!({
1715                "type": "reasoning.text",
1716                "text": "retained"
1717            })]
1718        );
1719    }
1720
1721    #[test]
1722    fn accepts_compatible_finish_reasons_and_rejects_incomplete_ones() {
1723        for reason in [
1724            None,
1725            Some(Value::Null),
1726            Some(Value::String("stop".to_owned())),
1727        ] {
1728            let mut choice = serde_json::json!({"delta": {}});
1729            if let Some(reason) = reason {
1730                choice["finish_reason"] = reason;
1731            }
1732            validate_finish_reason(&choice).expect("compatible finish reason");
1733        }
1734        for reason in ["tool_calls", "function_call"] {
1735            validate_finish_reason(&serde_json::json!({
1736                "delta": {},
1737                "finish_reason": reason
1738            }))
1739            .expect("tool finish reason");
1740        }
1741        for reason in ["length", "content_filter", "error"] {
1742            assert!(validate_finish_reason(&serde_json::json!({
1743                "delta": {},
1744                "finish_reason": reason
1745            }))
1746            .is_err());
1747        }
1748    }
1749
1750    #[test]
1751    fn rejects_api_keys_that_conflict_with_fixed_literals() {
1752        for (index, secret) in [
1753            "session",
1754            "tool",
1755            "cmd",
1756            "command",
1757            "finite",
1758            "0",
1759            ":",
1760            "[REDACTED]",
1761        ]
1762        .into_iter()
1763        .enumerate()
1764        {
1765            let environment = format!("LUCY_PROVIDER_CONFLICT_{}_{}", std::process::id(), index);
1766            std::env::set_var(&environment, secret);
1767            let settings = LlmSettings {
1768                base_url: "http://localhost".to_owned(),
1769                model: "model".to_owned(),
1770                api_key_env: environment.clone(),
1771                effort: None,
1772            };
1773            let error = match Provider::new(&settings) {
1774                Ok(_) => panic!("fixed literal conflict should be rejected: {secret}"),
1775                Err(error) => error,
1776            };
1777            assert!(error.to_string().contains("structured output"));
1778            assert!(!error.to_string().contains(secret));
1779            std::env::remove_var(environment);
1780        }
1781    }
1782
1783    #[test]
1784    fn accepts_a_normal_long_provider_key() {
1785        let environment = format!("LUCY_PROVIDER_NORMAL_{}", std::process::id());
1786        std::env::set_var(&environment, "provider-secret");
1787        let settings = LlmSettings {
1788            base_url: "http://localhost".to_owned(),
1789            model: "model".to_owned(),
1790            api_key_env: environment.clone(),
1791            effort: None,
1792        };
1793        assert!(Provider::new(&settings).is_ok());
1794        std::env::remove_var(environment);
1795    }
1796
1797    #[test]
1798    fn accepts_a_configurable_effort() {
1799        let environment = format!("LUCY_PROVIDER_EFFORT_OK_{}", std::process::id());
1800        std::env::set_var(&environment, "provider-secret");
1801        let settings = LlmSettings {
1802            base_url: "http://localhost".to_owned(),
1803            model: "model".to_owned(),
1804            api_key_env: environment.clone(),
1805            effort: Some("high".to_owned()),
1806        };
1807        assert!(Provider::new(&settings).is_ok());
1808        std::env::remove_var(environment);
1809    }
1810
1811    #[test]
1812    fn empty_effort_is_rejected_without_echoing_the_key() {
1813        let environment = format!("LUCY_PROVIDER_EFFORT_EMPTY_{}", std::process::id());
1814        std::env::set_var(&environment, "provider-secret");
1815        for effort in ["", "   ", "\t"] {
1816            let settings = LlmSettings {
1817                base_url: "http://localhost".to_owned(),
1818                model: "model".to_owned(),
1819                api_key_env: environment.clone(),
1820                effort: Some(effort.to_owned()),
1821            };
1822            let error = match Provider::new(&settings) {
1823                Ok(_) => panic!("empty effort should be rejected: {effort:?}"),
1824                Err(error) => error,
1825            };
1826            assert!(error.to_string().contains("llm.effort must not be empty"));
1827            assert!(!error.to_string().contains("provider-secret"));
1828        }
1829        std::env::remove_var(environment);
1830    }
1831
1832    #[test]
1833    fn missing_api_key_error_does_not_echo_the_environment_name() {
1834        let environment = format!("LUCY_MISSING_KEY_{}", std::process::id());
1835        std::env::remove_var(&environment);
1836        let settings = LlmSettings {
1837            base_url: "http://localhost".to_owned(),
1838            model: "model".to_owned(),
1839            api_key_env: environment.clone(),
1840            effort: None,
1841        };
1842        let error = match Provider::new(&settings) {
1843            Ok(_) => panic!("missing key should be rejected"),
1844            Err(error) => error,
1845        };
1846        assert_eq!(error.to_string(), "missing provider API key");
1847        assert!(!error.to_string().contains(&environment));
1848    }
1849
1850    #[test]
1851    fn cancellable_stream_stops_a_stalled_provider_without_waiting_for_timeout() {
1852        let listener = TcpListener::bind(("127.0.0.1", 0)).expect("listener");
1853        let address = listener.local_addr().expect("address");
1854        let (sent, sent_receiver) = mpsc::channel();
1855        let server = thread::spawn(move || {
1856            let (mut stream, _) = listener.accept().expect("request");
1857            let mut request = std::io::BufReader::new(stream.try_clone().expect("clone"));
1858            let mut content_length = 0;
1859            loop {
1860                let mut line = String::new();
1861                request.read_line(&mut line).expect("header");
1862                if line == "\r\n" {
1863                    break;
1864                }
1865                if let Some(value) = line.strip_prefix("Content-Length:") {
1866                    content_length = value.trim().parse::<usize>().expect("length");
1867                }
1868            }
1869            let mut body = vec![0; content_length];
1870            request.read_exact(&mut body).expect("body");
1871
1872            let payload = serde_json::json!({
1873                "choices": [{"delta": {"content": "partial"}, "finish_reason": null}]
1874            });
1875            let event = format!("data: {payload}\n\n");
1876            let response = format!(
1877                "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nTransfer-Encoding: chunked\r\nConnection: keep-alive\r\n\r\n{:x}\r\n{}\r\n",
1878                event.len(), event
1879            );
1880            stream.write_all(response.as_bytes()).expect("response");
1881            stream.flush().expect("flush");
1882            sent.send(()).expect("body readiness");
1883            thread::sleep(Duration::from_millis(500));
1884        });
1885
1886        let environment = format!("LUCY_PROVIDER_CANCEL_{}", std::process::id());
1887        std::env::set_var(&environment, "provider-secret");
1888        let provider = Provider::new(&LlmSettings {
1889            base_url: format!("http://{address}/v1"),
1890            model: "model".to_owned(),
1891            api_key_env: environment.clone(),
1892            effort: None,
1893        })
1894        .expect("provider");
1895        let token = CancellationToken::new();
1896        let worker_token = token.clone();
1897        let worker = thread::spawn(move || {
1898            let mut received = String::new();
1899            let result = provider.stream_chat_cancellable(
1900                &[ChatMessage::user("hello".to_owned())],
1901                &mut |text| {
1902                    received.push_str(text);
1903                    Ok(())
1904                },
1905                &worker_token,
1906            );
1907            (result, received)
1908        });
1909        sent_receiver
1910            .recv_timeout(Duration::from_secs(1))
1911            .expect("body was sent");
1912        let started = Instant::now();
1913        assert!(token.cancel());
1914        let (result, received) = worker.join().expect("provider worker");
1915        assert!(started.elapsed() < Duration::from_millis(400));
1916        let error = result.expect_err("cancellation");
1917        assert!(error.is_cancelled());
1918        assert!(received.is_empty() || received == "partial");
1919        server.join().expect("server");
1920        std::env::remove_var(environment);
1921    }
1922
1923    #[test]
1924    fn rejects_an_oversized_sse_line_before_json_parsing() {
1925        let stream = format!("data: {}\n\n", "x".repeat(MAX_SSE_LINE_BYTES));
1926        let error = parse_sse(&mut Cursor::new(stream.as_bytes()), |_| Ok(())).expect_err("limit");
1927        assert_eq!(
1928            error.to_string(),
1929            "provider SSE line exceeded the response limit"
1930        );
1931    }
1932
1933    #[test]
1934    fn rejects_an_oversized_sse_data_event_before_json_parsing() {
1935        let payload = "x".repeat(MAX_SSE_LINE_BYTES - "data: ".len());
1936        let line_count = MAX_SSE_EVENT_BYTES / payload.len() + 2;
1937        let mut stream = String::new();
1938        for _ in 0..line_count {
1939            stream.push_str("data: ");
1940            stream.push_str(&payload);
1941            stream.push('\n');
1942        }
1943        stream.push('\n');
1944
1945        let error = parse_sse(&mut Cursor::new(stream.as_bytes()), |_| Ok(())).expect_err("limit");
1946        assert_eq!(
1947            error.to_string(),
1948            "provider SSE data event exceeded the response limit"
1949        );
1950    }
1951
1952    #[test]
1953    fn rejects_an_oversized_sse_stream_of_ignored_fields() {
1954        let line = format!("ignored: {}\n", "x".repeat(1024));
1955        let mut stream = Vec::new();
1956        while stream.len() <= MAX_SSE_STREAM_BYTES {
1957            stream.extend_from_slice(line.as_bytes());
1958        }
1959
1960        let error = parse_sse(&mut Cursor::new(stream), |_| Ok(())).expect_err("limit");
1961        assert_eq!(
1962            error.to_string(),
1963            "provider SSE stream exceeded the response limit"
1964        );
1965    }
1966
1967    #[test]
1968    fn rejects_too_many_empty_sse_data_lines() {
1969        let stream = format!("{}\n", "data:\n".repeat(MAX_SSE_DATA_LINES + 1));
1970        let error = parse_sse(&mut Cursor::new(stream.as_bytes()), |_| Ok(())).expect_err("limit");
1971        assert_eq!(
1972            error.to_string(),
1973            "provider SSE data line count exceeded the response limit"
1974        );
1975    }
1976
1977    #[test]
1978    fn reports_eof_before_done_as_incomplete() {
1979        let stream = b"data: {\"choices\":[{\"delta\":{\"content\":\"partial\"}}]}\n\n";
1980        let result = parse_sse(&mut Cursor::new(stream), |_| Ok(())).expect("SSE parse");
1981        assert!(result.received_payload);
1982        assert!(!result.received_done);
1983    }
1984
1985    #[test]
1986    fn reports_empty_non_sse_input_without_payload_or_done() {
1987        let result =
1988            parse_sse(&mut Cursor::new(b"not an SSE response\n"), |_| Ok(())).expect("SSE parse");
1989        assert_eq!(result, SseParseResult::default());
1990    }
1991
1992    #[test]
1993    fn caps_accumulated_tool_call_id_and_name_fields() {
1994        let fragment = "x".repeat(MAX_PROVIDER_TOOL_CALL_ID_BYTES);
1995        let mut id = String::new();
1996        append_provider_field(
1997            &mut id,
1998            &fragment,
1999            MAX_PROVIDER_TOOL_CALL_ID_BYTES,
2000            "provider tool-call id exceeded the response limit",
2001        )
2002        .expect("id within limit");
2003        let error = append_provider_field(
2004            &mut id,
2005            "x",
2006            MAX_PROVIDER_TOOL_CALL_ID_BYTES,
2007            "provider tool-call id exceeded the response limit",
2008        )
2009        .expect_err("id limit");
2010        assert_eq!(
2011            error.to_string(),
2012            "provider tool-call id exceeded the response limit"
2013        );
2014
2015        let fragment = "x".repeat(MAX_PROVIDER_TOOL_NAME_BYTES);
2016        let mut name = String::new();
2017        append_provider_field(
2018            &mut name,
2019            &fragment,
2020            MAX_PROVIDER_TOOL_NAME_BYTES,
2021            "provider tool-call name exceeded the response limit",
2022        )
2023        .expect("name within limit");
2024        let error = append_provider_field(
2025            &mut name,
2026            "x",
2027            MAX_PROVIDER_TOOL_NAME_BYTES,
2028            "provider tool-call name exceeded the response limit",
2029        )
2030        .expect_err("name limit");
2031        assert_eq!(
2032            error.to_string(),
2033            "provider tool-call name exceeded the response limit"
2034        );
2035    }
2036
2037    #[test]
2038    fn caps_provider_error_text_without_copying_the_full_message() {
2039        let message = "x".repeat(MAX_PROVIDER_ERROR_BYTES + 1);
2040        let value = serde_json::json!({"error": {"message": message}});
2041        assert_eq!(
2042            provider_error_message(&value),
2043            Some("provider error text exceeded the response limit")
2044        );
2045    }
2046
2047    fn read_request_headers(stream: &TcpStream) {
2048        let mut reader = BufReader::new(stream.try_clone().expect("clone request"));
2049        loop {
2050            let mut line = String::new();
2051            reader.read_line(&mut line).expect("request header");
2052            if line == "\r\n" || line.is_empty() {
2053                return;
2054            }
2055        }
2056    }
2057
2058    fn response_body(text: &str) -> String {
2059        let payload = serde_json::json!({
2060            "choices": [{"delta": {"content": text}, "finish_reason": null}]
2061        });
2062        let finish = serde_json::json!({
2063            "choices": [{"delta": {}, "finish_reason": "stop"}]
2064        });
2065        format!("data: {payload}\n\ndata: {finish}\n\ndata: [DONE]\n\n")
2066    }
2067
2068    fn provider_for(address: std::net::SocketAddr, read_timeout: Duration) -> (Provider, String) {
2069        let environment = format!(
2070            "LUCY_PROVIDER_STREAM_TEST_{}_{}",
2071            std::process::id(),
2072            address.port()
2073        );
2074        std::env::set_var(&environment, "provider-secret");
2075        let settings = LlmSettings {
2076            base_url: format!("http://{address}/v1"),
2077            model: "model".to_owned(),
2078            api_key_env: environment.clone(),
2079            effort: None,
2080        };
2081        let mut provider = Provider::new(&settings).expect("provider");
2082        provider.async_client = AsyncClient::builder()
2083            .connect_timeout(Duration::from_secs(1))
2084            .read_timeout(read_timeout)
2085            .build()
2086            .expect("test async client");
2087        (provider, environment)
2088    }
2089
2090    #[test]
2091    fn worker_stream_can_exceed_idle_interval_without_total_deadline() {
2092        let listener = TcpListener::bind(("127.0.0.1", 0)).expect("listener");
2093        let address = listener.local_addr().expect("address");
2094        let parts = (0..5)
2095            .map(|index| {
2096                let payload = serde_json::json!({
2097                    "choices": [{
2098                        "delta": {"content": format!("part-{index}")},
2099                        "finish_reason": null
2100                    }]
2101                });
2102                format!("data: {payload}\n\n")
2103            })
2104            .chain([
2105                format!(
2106                    "data: {}\n\n",
2107                    serde_json::json!({
2108                        "choices": [{"delta": {}, "finish_reason": "stop"}]
2109                    })
2110                ),
2111                "data: [DONE]\n\n".to_owned(),
2112            ])
2113            .collect::<Vec<_>>();
2114        let body = parts.concat();
2115        let server = thread::spawn(move || {
2116            let (mut stream, _) = listener.accept().expect("request");
2117            read_request_headers(&stream);
2118            let header = format!(
2119                "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
2120                body.len()
2121            );
2122            stream.write_all(header.as_bytes()).expect("header");
2123            stream.flush().expect("header flush");
2124            for part in parts {
2125                stream.write_all(part.as_bytes()).expect("SSE part");
2126                stream.flush().expect("SSE flush");
2127                thread::sleep(Duration::from_millis(25));
2128            }
2129        });
2130
2131        let (provider, environment) = provider_for(address, Duration::from_millis(80));
2132        let cancellation = CancellationToken::new();
2133        let started = Instant::now();
2134        let mut output = String::new();
2135        let turn = provider
2136            .stream_chat_cancellable_with_options(
2137                &[ChatMessage::user("worker task".to_owned())],
2138                &mut |text| {
2139                    output.push_str(text);
2140                    Ok(())
2141                },
2142                &cancellation,
2143                true,
2144            )
2145            .expect("long worker stream");
2146
2147        assert!(started.elapsed() >= Duration::from_millis(80));
2148        assert_eq!(turn.content, output);
2149        assert!(output.contains("part-4"));
2150        server.join().expect("server");
2151        std::env::remove_var(environment);
2152    }
2153
2154    #[test]
2155    fn retries_a_pre_payload_stream_failure_once_and_classifies_it() {
2156        let listener = TcpListener::bind(("127.0.0.1", 0)).expect("listener");
2157        listener
2158            .set_nonblocking(true)
2159            .expect("nonblocking listener");
2160        let address = listener.local_addr().expect("address");
2161        let body = response_body("retried");
2162        let server = thread::spawn(move || {
2163            let deadline = Instant::now() + Duration::from_secs(2);
2164            for attempt in 0..2 {
2165                let (mut stream, _) = loop {
2166                    match listener.accept() {
2167                        Ok(connection) => break connection,
2168                        Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
2169                            assert!(Instant::now() < deadline, "provider did not retry");
2170                            thread::sleep(Duration::from_millis(5));
2171                        }
2172                        Err(error) => panic!("accept: {error}"),
2173                    }
2174                };
2175                read_request_headers(&stream);
2176                if attempt == 0 {
2177                    let header = "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nContent-Length: 100\r\nConnection: close\r\n\r\n";
2178                    stream.write_all(header.as_bytes()).expect("failed header");
2179                    stream.flush().expect("failed flush");
2180                } else {
2181                    let header = format!(
2182                        "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
2183                        body.len()
2184                    );
2185                    stream.write_all(header.as_bytes()).expect("success header");
2186                    stream.write_all(body.as_bytes()).expect("success body");
2187                    stream.flush().expect("success flush");
2188                }
2189            }
2190        });
2191
2192        let (provider, environment) = provider_for(address, Duration::from_secs(1));
2193        let cancellation = CancellationToken::new();
2194        let mut output = String::new();
2195        let turn = provider
2196            .stream_chat_cancellable_with_options(
2197                &[ChatMessage::user("retry".to_owned())],
2198                &mut |text| {
2199                    output.push_str(text);
2200                    Ok(())
2201                },
2202                &cancellation,
2203                true,
2204            )
2205            .expect("retry succeeds");
2206
2207        assert_eq!(turn.content, "retried");
2208        assert_eq!(output, "retried");
2209        server.join().expect("server");
2210        std::env::remove_var(environment);
2211    }
2212
2213    #[test]
2214    fn does_not_retry_after_partial_provider_output() {
2215        let listener = TcpListener::bind(("127.0.0.1", 0)).expect("listener");
2216        listener
2217            .set_nonblocking(true)
2218            .expect("nonblocking listener");
2219        let address = listener.local_addr().expect("address");
2220        let partial = format!(
2221            "data: {}\n\n",
2222            serde_json::json!({
2223                "choices": [{"delta": {"content": "partial"}, "finish_reason": null}]
2224            })
2225        );
2226        let server = thread::spawn(move || {
2227            let deadline = Instant::now() + Duration::from_secs(2);
2228            let (mut stream, _) = loop {
2229                match listener.accept() {
2230                    Ok(connection) => break connection,
2231                    Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
2232                        assert!(Instant::now() < deadline, "provider request missing");
2233                        thread::sleep(Duration::from_millis(5));
2234                    }
2235                    Err(error) => panic!("accept: {error}"),
2236                }
2237            };
2238            read_request_headers(&stream);
2239            let header = format!(
2240                "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
2241                partial.len() + 10
2242            );
2243            stream.write_all(header.as_bytes()).expect("partial header");
2244            stream.write_all(partial.as_bytes()).expect("partial body");
2245            stream.flush().expect("partial flush");
2246            thread::sleep(Duration::from_millis(150));
2247            assert!(matches!(
2248                listener.accept(),
2249                Err(error) if error.kind() == std::io::ErrorKind::WouldBlock
2250            ));
2251        });
2252
2253        let (provider, environment) = provider_for(address, Duration::from_secs(1));
2254        let cancellation = CancellationToken::new();
2255        let mut output = String::new();
2256        let error = provider
2257            .stream_chat_cancellable_with_options(
2258                &[ChatMessage::user("partial".to_owned())],
2259                &mut |text| {
2260                    output.push_str(text);
2261                    Ok(())
2262                },
2263                &cancellation,
2264                true,
2265            )
2266            .expect_err("partial stream must fail");
2267
2268        let message = error.to_string();
2269        assert!(message.contains("provider stream read failed"));
2270        assert!([
2271            "(timeout)",
2272            "(connection)",
2273            "(body)",
2274            "(decode)",
2275            "(request)",
2276            "(transport)",
2277        ]
2278        .iter()
2279        .any(|kind| message.contains(kind)));
2280        assert_eq!(output, "partial");
2281        server.join().expect("server");
2282        std::env::remove_var(environment);
2283    }
2284
2285    #[test]
2286    fn reports_midstream_error_without_echoing_provider_body() {
2287        let stream = b"data: {\"error\":{\"message\":\"bad request\"}}\n\n";
2288        let error = parse_sse(&mut Cursor::new(stream), |value| {
2289            if let Some(message) = provider_error_message(&value) {
2290                return Err(ProviderError::new(format!(
2291                    "provider stream error: {message}"
2292                )));
2293            }
2294            Ok(())
2295        })
2296        .expect_err("error");
2297        assert_eq!(error.to_string(), "provider stream error: bad request");
2298    }
2299}