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