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