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