Skip to main content

lucy/
provider.rs

1use std::collections::BTreeMap;
2use std::io::{self, BufRead, BufReader};
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
14pub const PROVIDER_TIMEOUT: Duration = Duration::from_secs(60);
15const MAX_PROVIDER_TOOL_CALLS: usize = 64;
16const MAX_PROVIDER_CONTENT_BYTES: usize = 1024 * 1024;
17const MAX_PROVIDER_REASONING_DETAILS_BYTES: usize = 1024 * 1024;
18const MAX_PROVIDER_TOOL_ARGUMENT_BYTES: usize = 1024 * 1024;
19const MAX_SSE_LINE_BYTES: usize = 64 * 1024;
20const MAX_SSE_EVENT_BYTES: usize = 1024 * 1024;
21const MAX_SSE_STREAM_BYTES: usize = 8 * 1024 * 1024;
22const MAX_SSE_DATA_LINES: usize = 1024;
23const MAX_PROVIDER_TOOL_CALL_ID_BYTES: usize = 16 * 1024;
24const MAX_PROVIDER_TOOL_NAME_BYTES: usize = 16 * 1024;
25const MAX_PROVIDER_ERROR_BYTES: usize = 16 * 1024;
26const CANCELLATION_POLL_INTERVAL: Duration = Duration::from_millis(10);
27
28#[derive(Debug)]
29pub struct ProviderError {
30    message: String,
31    cancelled: bool,
32    partial: Option<ProviderTurn>,
33}
34
35impl ProviderError {
36    fn new(message: impl Into<String>) -> Self {
37        Self {
38            message: message.into(),
39            cancelled: false,
40            partial: None,
41        }
42    }
43
44    fn cancelled(partial: ProviderTurn) -> Self {
45        Self {
46            message: "provider stream canceled".to_owned(),
47            cancelled: true,
48            partial: Some(partial),
49        }
50    }
51
52    pub fn is_cancelled(&self) -> bool {
53        self.cancelled
54    }
55
56    pub fn partial_turn(&self) -> Option<&ProviderTurn> {
57        self.partial.as_ref()
58    }
59}
60
61impl std::fmt::Display for ProviderError {
62    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
63        formatter.write_str(&self.message)
64    }
65}
66
67impl std::error::Error for ProviderError {}
68
69#[derive(Debug, Clone, PartialEq, Eq)]
70pub struct ProviderTurn {
71    pub content: String,
72    pub tool_calls: Vec<ChatToolCall>,
73    pub reasoning_details: Vec<Value>,
74}
75
76pub struct Provider {
77    client: Client,
78    async_client: AsyncClient,
79    endpoint: String,
80    model: String,
81    effort: Option<String>,
82    api_key_env: String,
83    api_key: String,
84}
85
86fn chat_request(model: &str, messages: &[ChatMessage], effort: &Option<String>) -> Value {
87    let mut request = json!({
88        "model": model,
89        "messages": messages
90            .iter()
91            .map(ChatMessage::to_openai_value)
92            .collect::<Vec<_>>(),
93        "stream": true,
94        "tools": [
95            {
96                "type": "function",
97                "function": {
98                    "name": "cmd",
99                    "description": "Execute a finite shell command in the session starting directory.",
100                    "parameters": {
101                        "type": "object",
102                        "properties": {
103                            "command": { "type": "string" }
104                        },
105                        "required": ["command"],
106                        "additionalProperties": false
107                    }
108                }
109            }
110        ]
111    });
112    if let Some(effort) = effort {
113        request["reasoning_effort"] = json!(effort);
114    }
115    request
116}
117
118impl Provider {
119    pub fn new(settings: &LlmSettings) -> Result<Self, ProviderError> {
120        let api_key = match std::env::var(&settings.api_key_env) {
121            Ok(api_key) if !api_key.is_empty() => api_key,
122            Ok(_) | Err(_) => return Err(ProviderError::new("missing provider API key")),
123        };
124        if conflicts_with_protected_literal(&api_key) {
125            return Err(ProviderError::new(redact_secret(
126                "API key conflicts with a required structured output literal",
127                Some(&api_key),
128            )));
129        }
130        if redaction_marker(&api_key).is_none() {
131            return Err(ProviderError::new(redact_secret(
132                "API key cannot be safely redacted",
133                Some(&api_key),
134            )));
135        }
136        if settings.model.trim().is_empty() {
137            return Err(ProviderError::new(redact_secret(
138                "missing llm.model; set a model in config.toml",
139                Some(&api_key),
140            )));
141        }
142        let effort = match &settings.effort {
143            Some(value) => {
144                let trimmed = value.trim();
145                if trimmed.is_empty() {
146                    return Err(ProviderError::new(redact_secret(
147                        "llm.effort must not be empty",
148                        Some(&api_key),
149                    )));
150                }
151                Some(trimmed.to_owned())
152            }
153            None => None,
154        };
155        let endpoint = format!(
156            "{}/chat/completions",
157            settings.base_url.trim_end_matches('/')
158        );
159        let client = Client::builder()
160            .timeout(PROVIDER_TIMEOUT)
161            .build()
162            .map_err(|_| {
163                ProviderError::new(redact_secret(
164                    "unable to initialize HTTP client",
165                    Some(&api_key),
166                ))
167            })?;
168        let async_client = AsyncClient::builder()
169            .timeout(PROVIDER_TIMEOUT)
170            .build()
171            .map_err(|_| {
172                ProviderError::new(redact_secret(
173                    "unable to initialize HTTP client",
174                    Some(&api_key),
175                ))
176            })?;
177        Ok(Self {
178            client,
179            async_client,
180            endpoint,
181            model: settings.model.clone(),
182            effort,
183            api_key_env: settings.api_key_env.clone(),
184            api_key,
185        })
186    }
187
188    pub fn api_key(&self) -> &str {
189        &self.api_key
190    }
191
192    pub fn api_key_env(&self) -> &str {
193        &self.api_key_env
194    }
195
196    pub fn stream_chat(
197        &self,
198        messages: &[ChatMessage],
199        on_text: &mut dyn FnMut(&str) -> io::Result<()>,
200    ) -> Result<ProviderTurn, ProviderError> {
201        let request = chat_request(&self.model, messages, &self.effort);
202
203        let response = self
204            .client
205            .post(&self.endpoint)
206            .bearer_auth(&self.api_key)
207            .header("accept", "text/event-stream")
208            .json(&request)
209            .send()
210            .map_err(|_| ProviderError::new("provider request failed"))?;
211        if !response.status().is_success() {
212            return Err(ProviderError::new(format!(
213                "provider returned HTTP status {}",
214                response.status().as_u16()
215            )));
216        }
217
218        let mut content = String::new();
219        let mut tool_calls = BTreeMap::<usize, PartialToolCall>::new();
220        let mut reasoning_details = Vec::new();
221        let mut reasoning_details_bytes = 0;
222        let mut tool_argument_bytes: usize = 0;
223        let mut finish_reason = None;
224        {
225            let mut on_data = |data: Value| -> Result<(), ProviderError> {
226                if let Some(message) = provider_error_message(&data) {
227                    return Err(ProviderError::new(format!(
228                        "provider stream error: {}",
229                        redact_secret(message, Some(&self.api_key))
230                    )));
231                }
232                let Some(choice) = data
233                    .get("choices")
234                    .and_then(Value::as_array)
235                    .and_then(|choices| choices.first())
236                else {
237                    return Ok(());
238                };
239                if let Some(reason) = validate_finish_reason(choice)? {
240                    finish_reason = Some(reason.to_owned());
241                }
242                let Some(delta) = choice.get("delta") else {
243                    return Ok(());
244                };
245                append_reasoning_details(
246                    &mut reasoning_details,
247                    &mut reasoning_details_bytes,
248                    delta,
249                )?;
250                if let Some(text) = delta.get("content").and_then(Value::as_str) {
251                    if content.len().saturating_add(text.len()) > MAX_PROVIDER_CONTENT_BYTES {
252                        return Err(ProviderError::new(
253                            "provider assistant content exceeded the response limit",
254                        ));
255                    }
256                    content.push_str(text);
257                    if on_text(text).is_err() {
258                        return Err(ProviderError::new("unable to emit assistant delta"));
259                    }
260                }
261                if let Some(calls) = delta.get("tool_calls").and_then(Value::as_array) {
262                    for (position, call) in calls.iter().enumerate() {
263                        let index = call
264                            .get("index")
265                            .and_then(Value::as_u64)
266                            .map_or(position, |index| index as usize);
267                        if !tool_calls.contains_key(&index)
268                            && tool_calls.len() >= MAX_PROVIDER_TOOL_CALLS
269                        {
270                            return Err(ProviderError::new(
271                                "provider response exceeded the tool-call limit",
272                            ));
273                        }
274                        let partial = tool_calls.entry(index).or_default();
275                        if let Some(id) = call.get("id").and_then(Value::as_str) {
276                            append_provider_field(
277                                &mut partial.id,
278                                id,
279                                MAX_PROVIDER_TOOL_CALL_ID_BYTES,
280                                "provider tool-call id exceeded the response limit",
281                            )?;
282                        }
283                        if let Some(function) = call.get("function") {
284                            if let Some(name) = function.get("name").and_then(Value::as_str) {
285                                append_provider_field(
286                                    &mut partial.name,
287                                    name,
288                                    MAX_PROVIDER_TOOL_NAME_BYTES,
289                                    "provider tool-call name exceeded the response limit",
290                                )?;
291                            }
292                            if let Some(arguments) =
293                                function.get("arguments").and_then(Value::as_str)
294                            {
295                                if tool_argument_bytes.saturating_add(arguments.len())
296                                    > MAX_PROVIDER_TOOL_ARGUMENT_BYTES
297                                {
298                                    return Err(ProviderError::new(
299                                        "provider tool arguments exceeded the response limit",
300                                    ));
301                                }
302                                tool_argument_bytes += arguments.len();
303                                partial.arguments.push_str(arguments);
304                            }
305                        }
306                    }
307                }
308                if let Some(function_call) = delta.get("function_call") {
309                    if !tool_calls.contains_key(&0) && tool_calls.len() >= MAX_PROVIDER_TOOL_CALLS {
310                        return Err(ProviderError::new(
311                            "provider response exceeded the tool-call limit",
312                        ));
313                    }
314                    let partial = tool_calls.entry(0).or_default();
315                    if let Some(name) = function_call.get("name").and_then(Value::as_str) {
316                        append_provider_field(
317                            &mut partial.name,
318                            name,
319                            MAX_PROVIDER_TOOL_NAME_BYTES,
320                            "provider tool-call name exceeded the response limit",
321                        )?;
322                    }
323                    if let Some(arguments) = function_call.get("arguments").and_then(Value::as_str)
324                    {
325                        if tool_argument_bytes.saturating_add(arguments.len())
326                            > MAX_PROVIDER_TOOL_ARGUMENT_BYTES
327                        {
328                            return Err(ProviderError::new(
329                                "provider tool arguments exceeded the response limit",
330                            ));
331                        }
332                        tool_argument_bytes += arguments.len();
333                        partial.arguments.push_str(arguments);
334                    }
335                }
336                Ok(())
337            };
338            let mut reader = BufReader::new(response);
339            let parse_result = parse_sse(&mut reader, &mut on_data)?;
340            if !parse_result.received_payload {
341                return Err(ProviderError::new(
342                    "provider stream contained no valid payload",
343                ));
344            }
345            if !parse_result.received_done {
346                return Err(ProviderError::new("provider stream ended before [DONE]"));
347            }
348        }
349
350        let tool_calls = tool_calls
351            .into_iter()
352            .map(|(index, partial)| ChatToolCall {
353                id: if partial.id.is_empty() {
354                    format!("call_{index}")
355                } else {
356                    partial.id
357                },
358                name: partial.name,
359                arguments: partial.arguments,
360            })
361            .collect::<Vec<_>>();
362        if let Some(reason) = finish_reason.as_deref() {
363            if !tool_calls.is_empty() && !matches!(reason, "tool_calls" | "function_call") {
364                return Err(ProviderError::new(
365                    "provider tool calls ended with an incompatible finish reason",
366                ));
367            }
368            if tool_calls.is_empty() && matches!(reason, "tool_calls" | "function_call") {
369                return Err(ProviderError::new(
370                    "provider reported tool completion without a tool call",
371                ));
372            }
373        }
374        if content.is_empty() && tool_calls.is_empty() {
375            return Err(ProviderError::new(
376                "provider stream contained no assistant content or tool calls",
377            ));
378        }
379        Ok(ProviderTurn {
380            content,
381            tool_calls,
382            reasoning_details,
383        })
384    }
385
386    /// Stream through an async response so cancellation can drop the pending
387    /// socket read instead of waiting for the blocking client's timeout.
388    pub(crate) fn stream_chat_cancellable(
389        &self,
390        messages: &[ChatMessage],
391        on_text: &mut dyn FnMut(&str) -> io::Result<()>,
392        cancellation: &CancellationToken,
393    ) -> Result<ProviderTurn, ProviderError> {
394        let runtime = tokio::runtime::Builder::new_current_thread()
395            .enable_all()
396            .build()
397            .map_err(|_| ProviderError::new("unable to initialize provider runtime"))?;
398        runtime.block_on(self.stream_chat_async(messages, on_text, cancellation))
399    }
400
401    async fn stream_chat_async(
402        &self,
403        messages: &[ChatMessage],
404        on_text: &mut dyn FnMut(&str) -> io::Result<()>,
405        cancellation: &CancellationToken,
406    ) -> Result<ProviderTurn, ProviderError> {
407        if cancellation.is_cancelled() {
408            return Err(ProviderError::cancelled(ProviderTurn {
409                content: String::new(),
410                tool_calls: Vec::new(),
411                reasoning_details: Vec::new(),
412            }));
413        }
414        let request = chat_request(&self.model, messages, &self.effort);
415        let request = self
416            .async_client
417            .post(&self.endpoint)
418            .bearer_auth(&self.api_key)
419            .header("accept", "text/event-stream")
420            .json(&request)
421            .send();
422        let mut request = Box::pin(request);
423        let mut response = loop {
424            if cancellation.is_cancelled() {
425                return Err(ProviderError::cancelled(ProviderTurn {
426                    content: String::new(),
427                    tool_calls: Vec::new(),
428                    reasoning_details: Vec::new(),
429                }));
430            }
431            match tokio::time::timeout(CANCELLATION_POLL_INTERVAL, request.as_mut()).await {
432                Ok(response) => {
433                    break response.map_err(|_| ProviderError::new("provider request failed"))?;
434                }
435                Err(_) => continue,
436            }
437        };
438        if !response.status().is_success() {
439            return Err(ProviderError::new(format!(
440                "provider returned HTTP status {}",
441                response.status().as_u16()
442            )));
443        }
444
445        let mut accumulator = ProviderAccumulator::default();
446        let mut decoder = SseDecoder::default();
447        loop {
448            if cancellation.is_cancelled() {
449                return Err(ProviderError::cancelled(accumulator.partial_turn()));
450            }
451            let chunk =
452                match tokio::time::timeout(CANCELLATION_POLL_INTERVAL, response.chunk()).await {
453                    Ok(chunk) => {
454                        chunk.map_err(|_| ProviderError::new("unable to read provider stream"))?
455                    }
456                    Err(_) => continue,
457                };
458            let Some(chunk) = chunk else {
459                break;
460            };
461            let done = decoder.feed(&chunk, &mut |data| {
462                accumulator.on_data(data, &self.api_key, on_text)
463            })?;
464            if done {
465                break;
466            }
467        }
468        if cancellation.is_cancelled() {
469            return Err(ProviderError::cancelled(accumulator.partial_turn()));
470        }
471        let parse_result =
472            decoder.finish(&mut |data| accumulator.on_data(data, &self.api_key, on_text))?;
473        if !parse_result.received_payload {
474            return Err(ProviderError::new(
475                "provider stream contained no valid payload",
476            ));
477        }
478        if !parse_result.received_done {
479            return Err(ProviderError::new("provider stream ended before [DONE]"));
480        }
481        accumulator.finish()
482    }
483}
484
485#[derive(Debug, Clone, Default)]
486struct PartialToolCall {
487    id: String,
488    name: String,
489    arguments: String,
490}
491
492#[derive(Debug, Default)]
493struct ProviderAccumulator {
494    content: String,
495    tool_calls: BTreeMap<usize, PartialToolCall>,
496    reasoning_details: Vec<Value>,
497    reasoning_details_bytes: usize,
498    tool_argument_bytes: usize,
499    finish_reason: Option<String>,
500}
501
502impl ProviderAccumulator {
503    fn on_data(
504        &mut self,
505        data: Value,
506        api_key: &str,
507        on_text: &mut dyn FnMut(&str) -> io::Result<()>,
508    ) -> Result<(), ProviderError> {
509        if let Some(message) = provider_error_message(&data) {
510            return Err(ProviderError::new(format!(
511                "provider stream error: {}",
512                redact_secret(message, Some(api_key))
513            )));
514        }
515        let Some(choice) = data
516            .get("choices")
517            .and_then(Value::as_array)
518            .and_then(|choices| choices.first())
519        else {
520            return Ok(());
521        };
522        if let Some(reason) = validate_finish_reason(choice)? {
523            self.finish_reason = Some(reason.to_owned());
524        }
525        let Some(delta) = choice.get("delta") else {
526            return Ok(());
527        };
528        append_reasoning_details(
529            &mut self.reasoning_details,
530            &mut self.reasoning_details_bytes,
531            delta,
532        )?;
533        if let Some(text) = delta.get("content").and_then(Value::as_str) {
534            if self.content.len().saturating_add(text.len()) > MAX_PROVIDER_CONTENT_BYTES {
535                return Err(ProviderError::new(
536                    "provider assistant content exceeded the response limit",
537                ));
538            }
539            self.content.push_str(text);
540            on_text(text).map_err(|_| ProviderError::new("unable to emit assistant delta"))?;
541        }
542        if let Some(calls) = delta.get("tool_calls").and_then(Value::as_array) {
543            for (position, call) in calls.iter().enumerate() {
544                let index = call
545                    .get("index")
546                    .and_then(Value::as_u64)
547                    .map_or(position, |index| index as usize);
548                if !self.tool_calls.contains_key(&index)
549                    && self.tool_calls.len() >= MAX_PROVIDER_TOOL_CALLS
550                {
551                    return Err(ProviderError::new(
552                        "provider response exceeded the tool-call limit",
553                    ));
554                }
555                let partial = self.tool_calls.entry(index).or_default();
556                if let Some(id) = call.get("id").and_then(Value::as_str) {
557                    append_provider_field(
558                        &mut partial.id,
559                        id,
560                        MAX_PROVIDER_TOOL_CALL_ID_BYTES,
561                        "provider tool-call id exceeded the response limit",
562                    )?;
563                }
564                if let Some(function) = call.get("function") {
565                    if let Some(name) = function.get("name").and_then(Value::as_str) {
566                        append_provider_field(
567                            &mut partial.name,
568                            name,
569                            MAX_PROVIDER_TOOL_NAME_BYTES,
570                            "provider tool-call name exceeded the response limit",
571                        )?;
572                    }
573                    if let Some(arguments) = function.get("arguments").and_then(Value::as_str) {
574                        if self.tool_argument_bytes.saturating_add(arguments.len())
575                            > MAX_PROVIDER_TOOL_ARGUMENT_BYTES
576                        {
577                            return Err(ProviderError::new(
578                                "provider tool arguments exceeded the response limit",
579                            ));
580                        }
581                        self.tool_argument_bytes += arguments.len();
582                        partial.arguments.push_str(arguments);
583                    }
584                }
585            }
586        }
587        if let Some(function_call) = delta.get("function_call") {
588            if !self.tool_calls.contains_key(&0) && self.tool_calls.len() >= MAX_PROVIDER_TOOL_CALLS
589            {
590                return Err(ProviderError::new(
591                    "provider response exceeded the tool-call limit",
592                ));
593            }
594            let partial = self.tool_calls.entry(0).or_default();
595            if let Some(name) = function_call.get("name").and_then(Value::as_str) {
596                append_provider_field(
597                    &mut partial.name,
598                    name,
599                    MAX_PROVIDER_TOOL_NAME_BYTES,
600                    "provider tool-call name exceeded the response limit",
601                )?;
602            }
603            if let Some(arguments) = function_call.get("arguments").and_then(Value::as_str) {
604                if self.tool_argument_bytes.saturating_add(arguments.len())
605                    > MAX_PROVIDER_TOOL_ARGUMENT_BYTES
606                {
607                    return Err(ProviderError::new(
608                        "provider tool arguments exceeded the response limit",
609                    ));
610                }
611                self.tool_argument_bytes += arguments.len();
612                partial.arguments.push_str(arguments);
613            }
614        }
615        Ok(())
616    }
617
618    fn partial_turn(&self) -> ProviderTurn {
619        ProviderTurn {
620            content: self.content.clone(),
621            tool_calls: self
622                .tool_calls
623                .iter()
624                .map(|(index, partial)| ChatToolCall {
625                    id: if partial.id.is_empty() {
626                        format!("call_{index}")
627                    } else {
628                        partial.id.clone()
629                    },
630                    name: partial.name.clone(),
631                    arguments: partial.arguments.clone(),
632                })
633                .collect(),
634            reasoning_details: self.reasoning_details.clone(),
635        }
636    }
637
638    fn finish(self) -> Result<ProviderTurn, ProviderError> {
639        let tool_calls = self
640            .tool_calls
641            .into_iter()
642            .map(|(index, partial)| ChatToolCall {
643                id: if partial.id.is_empty() {
644                    format!("call_{index}")
645                } else {
646                    partial.id
647                },
648                name: partial.name,
649                arguments: partial.arguments,
650            })
651            .collect::<Vec<_>>();
652        if let Some(reason) = self.finish_reason.as_deref() {
653            if !tool_calls.is_empty() && !matches!(reason, "tool_calls" | "function_call") {
654                return Err(ProviderError::new(
655                    "provider tool calls ended with an incompatible finish reason",
656                ));
657            }
658            if tool_calls.is_empty() && matches!(reason, "tool_calls" | "function_call") {
659                return Err(ProviderError::new(
660                    "provider reported tool completion without a tool call",
661                ));
662            }
663        }
664        if self.content.is_empty() && tool_calls.is_empty() {
665            return Err(ProviderError::new(
666                "provider stream contained no assistant content or tool calls",
667            ));
668        }
669        Ok(ProviderTurn {
670            content: self.content,
671            tool_calls,
672            reasoning_details: self.reasoning_details,
673        })
674    }
675}
676
677fn append_reasoning_details(
678    target: &mut Vec<Value>,
679    serialized_bytes: &mut usize,
680    delta: &Value,
681) -> Result<(), ProviderError> {
682    let Some(details) = delta.get("reasoning_details").and_then(Value::as_array) else {
683        return Ok(());
684    };
685    if details.is_empty() {
686        return Ok(());
687    }
688    let serialized_delta = serde_json::to_vec(details)
689        .map_err(|_| ProviderError::new("provider reasoning details could not be serialized"))?;
690    let combined_bytes = if target.is_empty() {
691        serialized_delta.len()
692    } else {
693        serialized_bytes
694            .saturating_add(serialized_delta.len())
695            .saturating_sub(1)
696    };
697    if combined_bytes > MAX_PROVIDER_REASONING_DETAILS_BYTES {
698        return Err(ProviderError::new(
699            "provider reasoning details exceeded the response limit",
700        ));
701    }
702    target.extend(details.iter().cloned());
703    *serialized_bytes = combined_bytes;
704    Ok(())
705}
706
707fn append_provider_field(
708    target: &mut String,
709    fragment: &str,
710    limit: usize,
711    error_message: &str,
712) -> Result<(), ProviderError> {
713    if target.len().saturating_add(fragment.len()) > limit {
714        return Err(ProviderError::new(error_message));
715    }
716    target.push_str(fragment);
717    Ok(())
718}
719
720fn provider_error_message(data: &Value) -> Option<&str> {
721    let error = data.get("error")?;
722    let message = if let Some(message) = error.get("message").and_then(Value::as_str) {
723        message
724    } else if let Some(message) = error.as_str() {
725        message
726    } else {
727        return Some("provider returned an error payload");
728    };
729    if message.len() > MAX_PROVIDER_ERROR_BYTES {
730        Some("provider error text exceeded the response limit")
731    } else {
732        Some(message)
733    }
734}
735
736#[derive(Debug, Default)]
737struct SseDecoder {
738    line: Vec<u8>,
739    data_lines: Vec<String>,
740    data_event_bytes: usize,
741    stream_bytes: usize,
742    result: SseParseResult,
743    done: bool,
744}
745
746impl SseDecoder {
747    fn feed<F>(&mut self, bytes: &[u8], on_data: &mut F) -> Result<bool, ProviderError>
748    where
749        F: FnMut(Value) -> Result<(), ProviderError>,
750    {
751        if self.stream_bytes.saturating_add(bytes.len()) > MAX_SSE_STREAM_BYTES {
752            return Err(ProviderError::new(
753                "provider SSE stream exceeded the response limit",
754            ));
755        }
756        self.stream_bytes += bytes.len();
757        for byte in bytes {
758            if self.done {
759                break;
760            }
761            if *byte == b'\n' {
762                let line = std::mem::take(&mut self.line);
763                if self.process_line(&line, on_data)? {
764                    return Ok(true);
765                }
766            } else {
767                self.line.push(*byte);
768                if self.line.len() > MAX_SSE_LINE_BYTES {
769                    return Err(ProviderError::new(
770                        "provider SSE line exceeded the response limit",
771                    ));
772                }
773            }
774        }
775        Ok(self.done)
776    }
777
778    fn finish<F>(&mut self, on_data: &mut F) -> Result<SseParseResult, ProviderError>
779    where
780        F: FnMut(Value) -> Result<(), ProviderError>,
781    {
782        if !self.line.is_empty() && !self.done {
783            let line = std::mem::take(&mut self.line);
784            self.process_line(&line, on_data)?;
785        }
786        if !self.done {
787            self.dispatch_data(on_data)?;
788        }
789        Ok(self.result)
790    }
791
792    fn process_line<F>(&mut self, raw_line: &[u8], on_data: &mut F) -> Result<bool, ProviderError>
793    where
794        F: FnMut(Value) -> Result<(), ProviderError>,
795    {
796        let line = std::str::from_utf8(raw_line)
797            .map_err(|_| ProviderError::new("unable to read provider stream"))?
798            .trim_end_matches('\r');
799        if line.is_empty() {
800            return self.dispatch_data(on_data);
801        }
802        if line.starts_with(':') {
803            return Ok(false);
804        }
805        let (field, value) = line
806            .split_once(':')
807            .map_or((line, ""), |(field, value)| (field, value));
808        if field == "data" {
809            let value = value.strip_prefix(' ').unwrap_or(value);
810            let separator_bytes = (!self.data_lines.is_empty()) as usize;
811            let added_bytes = separator_bytes.saturating_add(value.len());
812            if self.data_event_bytes.saturating_add(added_bytes) > MAX_SSE_EVENT_BYTES {
813                return Err(ProviderError::new(
814                    "provider SSE data event exceeded the response limit",
815                ));
816            }
817            if self.data_lines.len() >= MAX_SSE_DATA_LINES {
818                return Err(ProviderError::new(
819                    "provider SSE data line count exceeded the response limit",
820                ));
821            }
822            self.data_event_bytes += added_bytes;
823            self.data_lines.push(value.to_owned());
824        }
825        Ok(false)
826    }
827
828    fn dispatch_data<F>(&mut self, on_data: &mut F) -> Result<bool, ProviderError>
829    where
830        F: FnMut(Value) -> Result<(), ProviderError>,
831    {
832        if self.data_lines.is_empty() {
833            self.data_event_bytes = 0;
834            return Ok(false);
835        }
836        let data = self.data_lines.join("\n");
837        self.data_lines.clear();
838        self.data_event_bytes = 0;
839        if data.trim().is_empty() {
840            return Ok(false);
841        }
842        if data == "[DONE]" {
843            self.result.received_done = true;
844            self.done = true;
845            return Ok(true);
846        }
847        let value: Value = serde_json::from_str(&data)
848            .map_err(|_| ProviderError::new("provider sent malformed SSE data"))?;
849        self.result.received_payload = true;
850        on_data(value)?;
851        Ok(false)
852    }
853}
854
855#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
856pub struct SseParseResult {
857    pub received_payload: bool,
858    pub received_done: bool,
859}
860
861pub fn parse_sse<R, F>(reader: &mut R, mut on_data: F) -> Result<SseParseResult, ProviderError>
862where
863    R: BufRead,
864    F: FnMut(Value) -> Result<(), ProviderError>,
865{
866    let mut data_lines = Vec::new();
867    let mut data_event_bytes = 0;
868    let mut stream_bytes: usize = 0;
869    let mut result = SseParseResult::default();
870    let mut line = Vec::with_capacity(MAX_SSE_LINE_BYTES);
871    loop {
872        let (has_line, line_bytes) = read_sse_line(reader, &mut line)?;
873        if stream_bytes.saturating_add(line_bytes) > MAX_SSE_STREAM_BYTES {
874            return Err(ProviderError::new(
875                "provider SSE stream exceeded the response limit",
876            ));
877        }
878        stream_bytes += line_bytes;
879        if !has_line {
880            if !data_lines.is_empty() {
881                dispatch_data(
882                    &mut data_lines,
883                    &mut data_event_bytes,
884                    &mut on_data,
885                    &mut result,
886                )?;
887            }
888            return Ok(result);
889        }
890
891        let line = std::str::from_utf8(&line)
892            .map_err(|_| ProviderError::new("unable to read provider stream"))?
893            .trim_end_matches('\r');
894        if line.is_empty() {
895            if dispatch_data(
896                &mut data_lines,
897                &mut data_event_bytes,
898                &mut on_data,
899                &mut result,
900            )? {
901                return Ok(result);
902            }
903            continue;
904        }
905        if line.starts_with(':') {
906            continue;
907        }
908        let (field, value) = line
909            .split_once(':')
910            .map_or((line, ""), |(field, value)| (field, value));
911        if field == "data" {
912            let value = value.strip_prefix(' ').unwrap_or(value);
913            let separator_bytes = (!data_lines.is_empty()) as usize;
914            let added_bytes = separator_bytes.saturating_add(value.len());
915            if data_event_bytes.saturating_add(added_bytes) > MAX_SSE_EVENT_BYTES {
916                return Err(ProviderError::new(
917                    "provider SSE data event exceeded the response limit",
918                ));
919            }
920            if data_lines.len() >= MAX_SSE_DATA_LINES {
921                return Err(ProviderError::new(
922                    "provider SSE data line count exceeded the response limit",
923                ));
924            }
925            data_event_bytes += added_bytes;
926            data_lines.push(value.to_owned());
927        }
928    }
929}
930
931fn read_sse_line<R: BufRead>(
932    reader: &mut R,
933    line: &mut Vec<u8>,
934) -> Result<(bool, usize), ProviderError> {
935    line.clear();
936    let mut consumed_bytes = 0;
937    loop {
938        let buffer = reader
939            .fill_buf()
940            .map_err(|_| ProviderError::new("unable to read provider stream"))?;
941        if buffer.is_empty() {
942            return Ok((!line.is_empty(), consumed_bytes));
943        }
944
945        let newline = buffer.iter().position(|byte| *byte == b'\n');
946        let chunk_length = newline.unwrap_or(buffer.len());
947        if line.len().saturating_add(chunk_length) > MAX_SSE_LINE_BYTES {
948            return Err(ProviderError::new(
949                "provider SSE line exceeded the response limit",
950            ));
951        }
952        line.extend_from_slice(&buffer[..chunk_length]);
953        let consumed = newline.map_or(chunk_length, |index| index + 1);
954        reader.consume(consumed);
955        consumed_bytes += consumed;
956        if newline.is_some() {
957            return Ok((true, consumed_bytes));
958        }
959    }
960}
961
962fn dispatch_data<F>(
963    data_lines: &mut Vec<String>,
964    data_event_bytes: &mut usize,
965    on_data: &mut F,
966    result: &mut SseParseResult,
967) -> Result<bool, ProviderError>
968where
969    F: FnMut(Value) -> Result<(), ProviderError>,
970{
971    if data_lines.is_empty() {
972        *data_event_bytes = 0;
973        return Ok(false);
974    }
975    let data = data_lines.join("\n");
976    data_lines.clear();
977    *data_event_bytes = 0;
978    if data.trim().is_empty() {
979        return Ok(false);
980    }
981    if data == "[DONE]" {
982        result.received_done = true;
983        return Ok(true);
984    }
985    let value: Value = serde_json::from_str(&data)
986        .map_err(|_| ProviderError::new("provider sent malformed SSE data"))?;
987    result.received_payload = true;
988    on_data(value)?;
989    Ok(false)
990}
991
992fn validate_finish_reason(choice: &Value) -> Result<Option<&str>, ProviderError> {
993    let Some(reason) = choice.get("finish_reason") else {
994        return Ok(None);
995    };
996    if reason.is_null() {
997        return Ok(None);
998    }
999    match reason.as_str() {
1000        Some("stop") | Some("tool_calls") | Some("function_call") => Ok(reason.as_str()),
1001        Some("length") | Some("content_filter") => Err(ProviderError::new(
1002            "provider response ended before completion",
1003        )),
1004        Some(_) | None => Err(ProviderError::new(
1005            "provider response has an unsupported finish reason",
1006        )),
1007    }
1008}
1009
1010#[cfg(test)]
1011mod tests {
1012    use super::*;
1013    use std::io::{Cursor, Read, Write};
1014    use std::net::TcpListener;
1015    use std::sync::mpsc;
1016    use std::thread;
1017    use std::time::Instant;
1018
1019    #[test]
1020    fn parses_sse_comments_multiline_data_and_done() {
1021        let stream = b": keep-alive\n\ndata: \n\ndata: {\"choices\":[]\ndata: }\n\ndata: [DONE]\n";
1022        let mut values = Vec::new();
1023        let result = parse_sse(&mut Cursor::new(stream), |value| {
1024            values.push(value);
1025            Ok(())
1026        })
1027        .expect("SSE");
1028        assert!(result.received_payload);
1029        assert!(result.received_done);
1030        assert_eq!(values.len(), 1);
1031        assert!(values[0]["choices"].is_array());
1032    }
1033
1034    #[test]
1035    fn parses_text_and_fragmented_tool_calls() {
1036        let first = serde_json::json!({
1037            "choices": [{"delta": {"content": "hi"}}]
1038        });
1039        let second = serde_json::json!({
1040            "choices": [{
1041                "delta": {
1042                    "tool_calls": [{
1043                        "index": 0,
1044                        "id": "c1",
1045                        "function": {"name": "cmd", "arguments": "{command:"}
1046                    }]
1047                }
1048            }]
1049        });
1050        let third = serde_json::json!({
1051            "choices": [{
1052                "delta": {
1053                    "tool_calls": [{
1054                        "index": 0,
1055                        "function": {"arguments": "pwd}"}
1056                    }]
1057                }
1058            }]
1059        });
1060        let stream =
1061            format!("data: {first}\n\n data: {second}\n\ndata: {third}\n\ndata: [DONE]\n\n")
1062                .replace(" data:", "data:");
1063        let mut content = String::new();
1064        let mut calls = BTreeMap::<usize, PartialToolCall>::new();
1065        let result = parse_sse(&mut Cursor::new(stream.as_bytes()), |value| {
1066            let choice = &value["choices"][0];
1067            let delta = &choice["delta"];
1068            if let Some(text) = delta["content"].as_str() {
1069                content.push_str(text);
1070            }
1071            if let Some(tool_calls) = delta["tool_calls"].as_array() {
1072                for call in tool_calls {
1073                    let index = call["index"].as_u64().expect("index") as usize;
1074                    let partial = calls.entry(index).or_default();
1075                    partial.id.push_str(call["id"].as_str().unwrap_or(""));
1076                    partial
1077                        .name
1078                        .push_str(call["function"]["name"].as_str().unwrap_or(""));
1079                    partial
1080                        .arguments
1081                        .push_str(call["function"]["arguments"].as_str().unwrap_or(""));
1082                }
1083            }
1084            Ok(())
1085        })
1086        .expect("SSE");
1087        assert!(result.received_payload);
1088        assert!(result.received_done);
1089        assert_eq!(content, "hi");
1090        assert_eq!(calls[&0].id, "c1");
1091        assert_eq!(calls[&0].name, "cmd");
1092        assert_eq!(calls[&0].arguments, "{command:pwd}");
1093    }
1094
1095    #[test]
1096    fn accumulates_reasoning_details_with_fragmented_tool_calls() {
1097        let mut accumulator = ProviderAccumulator::default();
1098        accumulator
1099            .on_data(
1100                serde_json::json!({
1101                    "choices": [{
1102                        "delta": {
1103                            "reasoning_details": [{
1104                                "type": "reasoning.text",
1105                                "text": "part one"
1106                            }]
1107                        }
1108                    }]
1109                }),
1110                "provider-secret",
1111                &mut |_| Ok(()),
1112            )
1113            .expect("first provider chunk");
1114        accumulator
1115            .on_data(
1116                serde_json::json!({
1117                    "choices": [{
1118                        "delta": {
1119                            "reasoning_details": [{
1120                                "type": "reasoning.text",
1121                                "text": "part two"
1122                            }],
1123                            "tool_calls": [{
1124                                "index": 0,
1125                                "id": "call-1",
1126                                "function": {
1127                                    "name": "cmd",
1128                                    "arguments": "{\"command\":\"true\"}"
1129                                }
1130                            }]
1131                        },
1132                        "finish_reason": "tool_calls"
1133                    }]
1134                }),
1135                "provider-secret",
1136                &mut |_| Ok(()),
1137            )
1138            .expect("second provider chunk");
1139
1140        let partial = accumulator.partial_turn();
1141        assert_eq!(partial.reasoning_details.len(), 2);
1142        let turn = accumulator.finish().expect("provider turn");
1143        assert_eq!(
1144            turn.reasoning_details,
1145            vec![
1146                json!({"type": "reasoning.text", "text": "part one"}),
1147                json!({"type": "reasoning.text", "text": "part two"}),
1148            ]
1149        );
1150        assert_eq!(turn.tool_calls.len(), 1);
1151        assert_eq!(turn.tool_calls[0].name, "cmd");
1152    }
1153
1154    #[test]
1155    fn accumulates_many_small_reasoning_details_and_rejects_overflow_atomically() {
1156        const FRAGMENT_COUNT: usize = 4096;
1157        let mut details = Vec::new();
1158        let mut serialized_bytes = 0;
1159        let delta = serde_json::json!({
1160            "reasoning_details": [{
1161                "type": "reasoning.text",
1162                "text": "x".repeat(64)
1163            }]
1164        });
1165        for _ in 0..FRAGMENT_COUNT {
1166            append_reasoning_details(&mut details, &mut serialized_bytes, &delta)
1167                .expect("small reasoning detail");
1168        }
1169        assert_eq!(details.len(), FRAGMENT_COUNT);
1170
1171        let first_chunk_delta = serde_json::json!({
1172            "reasoning_details": [{
1173                "type": "reasoning.text",
1174                "text": "x".repeat(500 * 1024)
1175            }]
1176        });
1177        let first_chunk_bytes = serde_json::to_vec(&first_chunk_delta["reasoning_details"])
1178            .expect("first reasoning detail chunk")
1179            .len();
1180        assert!(first_chunk_bytes < MAX_PROVIDER_REASONING_DETAILS_BYTES);
1181        append_reasoning_details(&mut details, &mut serialized_bytes, &first_chunk_delta)
1182            .expect("first individually bounded reasoning detail chunk");
1183        assert!(serialized_bytes <= MAX_PROVIDER_REASONING_DETAILS_BYTES);
1184
1185        let second_chunk_delta = serde_json::json!({
1186            "reasoning_details": [{
1187                "type": "reasoning.text",
1188                "text": "x".repeat(200 * 1024)
1189            }]
1190        });
1191        let second_chunk_bytes = serde_json::to_vec(&second_chunk_delta["reasoning_details"])
1192            .expect("second reasoning detail chunk")
1193            .len();
1194        assert!(second_chunk_bytes < MAX_PROVIDER_REASONING_DETAILS_BYTES);
1195        assert!(
1196            serialized_bytes
1197                .saturating_add(second_chunk_bytes)
1198                .saturating_sub(1)
1199                > MAX_PROVIDER_REASONING_DETAILS_BYTES
1200        );
1201
1202        let prior_details = details.clone();
1203        let prior_bytes = serialized_bytes;
1204        let error =
1205            append_reasoning_details(&mut details, &mut serialized_bytes, &second_chunk_delta)
1206                .expect_err("reasoning details limit");
1207
1208        assert_eq!(
1209            error.to_string(),
1210            "provider reasoning details exceeded the response limit"
1211        );
1212        assert_eq!(details, prior_details);
1213        assert_eq!(serialized_bytes, prior_bytes);
1214        assert_eq!(
1215            serde_json::to_vec(&details)
1216                .expect("accumulated reasoning details")
1217                .len(),
1218            serialized_bytes
1219        );
1220        assert!(serialized_bytes <= MAX_PROVIDER_REASONING_DETAILS_BYTES);
1221    }
1222
1223    #[test]
1224    fn rejects_reasoning_details_that_exceed_the_serialized_response_limit_before_retaining_them() {
1225        let mut accumulator = ProviderAccumulator::default();
1226        let retained = serde_json::json!({
1227            "choices": [{
1228                "delta": {
1229                    "reasoning_details": [{
1230                        "type": "reasoning.text",
1231                        "text": "retained"
1232                    }]
1233                }
1234            }]
1235        });
1236        accumulator
1237            .on_data(retained, "provider-secret", &mut |_| Ok(()))
1238            .expect("details within limit");
1239
1240        let oversized = "x".repeat(MAX_PROVIDER_REASONING_DETAILS_BYTES);
1241        let error = accumulator
1242            .on_data(
1243                serde_json::json!({
1244                    "choices": [{
1245                        "delta": {
1246                            "reasoning_details": [{"text": oversized}]
1247                        }
1248                    }]
1249                }),
1250                "provider-secret",
1251                &mut |_| Ok(()),
1252            )
1253            .expect_err("reasoning details limit");
1254        assert_eq!(
1255            error.to_string(),
1256            "provider reasoning details exceeded the response limit"
1257        );
1258        assert_eq!(
1259            accumulator.reasoning_details,
1260            vec![serde_json::json!({
1261                "type": "reasoning.text",
1262                "text": "retained"
1263            })]
1264        );
1265    }
1266
1267    #[test]
1268    fn accepts_compatible_finish_reasons_and_rejects_incomplete_ones() {
1269        for reason in [
1270            None,
1271            Some(Value::Null),
1272            Some(Value::String("stop".to_owned())),
1273        ] {
1274            let mut choice = serde_json::json!({"delta": {}});
1275            if let Some(reason) = reason {
1276                choice["finish_reason"] = reason;
1277            }
1278            validate_finish_reason(&choice).expect("compatible finish reason");
1279        }
1280        for reason in ["tool_calls", "function_call"] {
1281            validate_finish_reason(&serde_json::json!({
1282                "delta": {},
1283                "finish_reason": reason
1284            }))
1285            .expect("tool finish reason");
1286        }
1287        for reason in ["length", "content_filter", "error"] {
1288            assert!(validate_finish_reason(&serde_json::json!({
1289                "delta": {},
1290                "finish_reason": reason
1291            }))
1292            .is_err());
1293        }
1294    }
1295
1296    #[test]
1297    fn rejects_api_keys_that_conflict_with_fixed_literals() {
1298        for (index, secret) in [
1299            "session",
1300            "tool",
1301            "cmd",
1302            "command",
1303            "finite",
1304            "0",
1305            ":",
1306            "[REDACTED]",
1307        ]
1308        .into_iter()
1309        .enumerate()
1310        {
1311            let environment = format!("LUCY_PROVIDER_CONFLICT_{}_{}", std::process::id(), index);
1312            std::env::set_var(&environment, secret);
1313            let settings = LlmSettings {
1314                base_url: "http://localhost".to_owned(),
1315                model: "model".to_owned(),
1316                api_key_env: environment.clone(),
1317                effort: None,
1318            };
1319            let error = match Provider::new(&settings) {
1320                Ok(_) => panic!("fixed literal conflict should be rejected: {secret}"),
1321                Err(error) => error,
1322            };
1323            assert!(error.to_string().contains("structured output"));
1324            assert!(!error.to_string().contains(secret));
1325            std::env::remove_var(environment);
1326        }
1327    }
1328
1329    #[test]
1330    fn accepts_a_normal_long_provider_key() {
1331        let environment = format!("LUCY_PROVIDER_NORMAL_{}", std::process::id());
1332        std::env::set_var(&environment, "provider-secret");
1333        let settings = LlmSettings {
1334            base_url: "http://localhost".to_owned(),
1335            model: "model".to_owned(),
1336            api_key_env: environment.clone(),
1337            effort: None,
1338        };
1339        assert!(Provider::new(&settings).is_ok());
1340        std::env::remove_var(environment);
1341    }
1342
1343    #[test]
1344    fn accepts_a_configurable_effort() {
1345        let environment = format!("LUCY_PROVIDER_EFFORT_OK_{}", std::process::id());
1346        std::env::set_var(&environment, "provider-secret");
1347        let settings = LlmSettings {
1348            base_url: "http://localhost".to_owned(),
1349            model: "model".to_owned(),
1350            api_key_env: environment.clone(),
1351            effort: Some("high".to_owned()),
1352        };
1353        assert!(Provider::new(&settings).is_ok());
1354        std::env::remove_var(environment);
1355    }
1356
1357    #[test]
1358    fn empty_effort_is_rejected_without_echoing_the_key() {
1359        let environment = format!("LUCY_PROVIDER_EFFORT_EMPTY_{}", std::process::id());
1360        std::env::set_var(&environment, "provider-secret");
1361        for effort in ["", "   ", "\t"] {
1362            let settings = LlmSettings {
1363                base_url: "http://localhost".to_owned(),
1364                model: "model".to_owned(),
1365                api_key_env: environment.clone(),
1366                effort: Some(effort.to_owned()),
1367            };
1368            let error = match Provider::new(&settings) {
1369                Ok(_) => panic!("empty effort should be rejected: {effort:?}"),
1370                Err(error) => error,
1371            };
1372            assert!(error.to_string().contains("llm.effort must not be empty"));
1373            assert!(!error.to_string().contains("provider-secret"));
1374        }
1375        std::env::remove_var(environment);
1376    }
1377
1378    #[test]
1379    fn missing_api_key_error_does_not_echo_the_environment_name() {
1380        let environment = format!("LUCY_MISSING_KEY_{}", std::process::id());
1381        std::env::remove_var(&environment);
1382        let settings = LlmSettings {
1383            base_url: "http://localhost".to_owned(),
1384            model: "model".to_owned(),
1385            api_key_env: environment.clone(),
1386            effort: None,
1387        };
1388        let error = match Provider::new(&settings) {
1389            Ok(_) => panic!("missing key should be rejected"),
1390            Err(error) => error,
1391        };
1392        assert_eq!(error.to_string(), "missing provider API key");
1393        assert!(!error.to_string().contains(&environment));
1394    }
1395
1396    #[test]
1397    fn cancellable_stream_stops_a_stalled_provider_without_waiting_for_timeout() {
1398        let listener = TcpListener::bind(("127.0.0.1", 0)).expect("listener");
1399        let address = listener.local_addr().expect("address");
1400        let (sent, sent_receiver) = mpsc::channel();
1401        let server = thread::spawn(move || {
1402            let (mut stream, _) = listener.accept().expect("request");
1403            let mut request = std::io::BufReader::new(stream.try_clone().expect("clone"));
1404            let mut content_length = 0;
1405            loop {
1406                let mut line = String::new();
1407                request.read_line(&mut line).expect("header");
1408                if line == "\r\n" {
1409                    break;
1410                }
1411                if let Some(value) = line.strip_prefix("Content-Length:") {
1412                    content_length = value.trim().parse::<usize>().expect("length");
1413                }
1414            }
1415            let mut body = vec![0; content_length];
1416            request.read_exact(&mut body).expect("body");
1417
1418            let payload = serde_json::json!({
1419                "choices": [{"delta": {"content": "partial"}, "finish_reason": null}]
1420            });
1421            let event = format!("data: {payload}\n\n");
1422            let response = format!(
1423                "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",
1424                event.len(), event
1425            );
1426            stream.write_all(response.as_bytes()).expect("response");
1427            stream.flush().expect("flush");
1428            sent.send(()).expect("body readiness");
1429            thread::sleep(Duration::from_millis(500));
1430        });
1431
1432        let environment = format!("LUCY_PROVIDER_CANCEL_{}", std::process::id());
1433        std::env::set_var(&environment, "provider-secret");
1434        let provider = Provider::new(&LlmSettings {
1435            base_url: format!("http://{address}/v1"),
1436            model: "model".to_owned(),
1437            api_key_env: environment.clone(),
1438            effort: None,
1439        })
1440        .expect("provider");
1441        let token = CancellationToken::new();
1442        let worker_token = token.clone();
1443        let worker = thread::spawn(move || {
1444            let mut received = String::new();
1445            let result = provider.stream_chat_cancellable(
1446                &[ChatMessage::user("hello".to_owned())],
1447                &mut |text| {
1448                    received.push_str(text);
1449                    Ok(())
1450                },
1451                &worker_token,
1452            );
1453            (result, received)
1454        });
1455        sent_receiver
1456            .recv_timeout(Duration::from_secs(1))
1457            .expect("body was sent");
1458        let started = Instant::now();
1459        assert!(token.cancel());
1460        let (result, received) = worker.join().expect("provider worker");
1461        assert!(started.elapsed() < Duration::from_millis(400));
1462        let error = result.expect_err("cancellation");
1463        assert!(error.is_cancelled());
1464        assert!(received.is_empty() || received == "partial");
1465        server.join().expect("server");
1466        std::env::remove_var(environment);
1467    }
1468
1469    #[test]
1470    fn rejects_an_oversized_sse_line_before_json_parsing() {
1471        let stream = format!("data: {}\n\n", "x".repeat(MAX_SSE_LINE_BYTES));
1472        let error = parse_sse(&mut Cursor::new(stream.as_bytes()), |_| Ok(())).expect_err("limit");
1473        assert_eq!(
1474            error.to_string(),
1475            "provider SSE line exceeded the response limit"
1476        );
1477    }
1478
1479    #[test]
1480    fn rejects_an_oversized_sse_data_event_before_json_parsing() {
1481        let payload = "x".repeat(MAX_SSE_LINE_BYTES - "data: ".len());
1482        let line_count = MAX_SSE_EVENT_BYTES / payload.len() + 2;
1483        let mut stream = String::new();
1484        for _ in 0..line_count {
1485            stream.push_str("data: ");
1486            stream.push_str(&payload);
1487            stream.push('\n');
1488        }
1489        stream.push('\n');
1490
1491        let error = parse_sse(&mut Cursor::new(stream.as_bytes()), |_| Ok(())).expect_err("limit");
1492        assert_eq!(
1493            error.to_string(),
1494            "provider SSE data event exceeded the response limit"
1495        );
1496    }
1497
1498    #[test]
1499    fn rejects_an_oversized_sse_stream_of_ignored_fields() {
1500        let line = format!("ignored: {}\n", "x".repeat(1024));
1501        let mut stream = Vec::new();
1502        while stream.len() <= MAX_SSE_STREAM_BYTES {
1503            stream.extend_from_slice(line.as_bytes());
1504        }
1505
1506        let error = parse_sse(&mut Cursor::new(stream), |_| Ok(())).expect_err("limit");
1507        assert_eq!(
1508            error.to_string(),
1509            "provider SSE stream exceeded the response limit"
1510        );
1511    }
1512
1513    #[test]
1514    fn rejects_too_many_empty_sse_data_lines() {
1515        let stream = format!("{}\n", "data:\n".repeat(MAX_SSE_DATA_LINES + 1));
1516        let error = parse_sse(&mut Cursor::new(stream.as_bytes()), |_| Ok(())).expect_err("limit");
1517        assert_eq!(
1518            error.to_string(),
1519            "provider SSE data line count exceeded the response limit"
1520        );
1521    }
1522
1523    #[test]
1524    fn reports_eof_before_done_as_incomplete() {
1525        let stream = b"data: {\"choices\":[{\"delta\":{\"content\":\"partial\"}}]}\n\n";
1526        let result = parse_sse(&mut Cursor::new(stream), |_| Ok(())).expect("SSE parse");
1527        assert!(result.received_payload);
1528        assert!(!result.received_done);
1529    }
1530
1531    #[test]
1532    fn reports_empty_non_sse_input_without_payload_or_done() {
1533        let result =
1534            parse_sse(&mut Cursor::new(b"not an SSE response\n"), |_| Ok(())).expect("SSE parse");
1535        assert_eq!(result, SseParseResult::default());
1536    }
1537
1538    #[test]
1539    fn caps_accumulated_tool_call_id_and_name_fields() {
1540        let fragment = "x".repeat(MAX_PROVIDER_TOOL_CALL_ID_BYTES);
1541        let mut id = String::new();
1542        append_provider_field(
1543            &mut id,
1544            &fragment,
1545            MAX_PROVIDER_TOOL_CALL_ID_BYTES,
1546            "provider tool-call id exceeded the response limit",
1547        )
1548        .expect("id within limit");
1549        let error = append_provider_field(
1550            &mut id,
1551            "x",
1552            MAX_PROVIDER_TOOL_CALL_ID_BYTES,
1553            "provider tool-call id exceeded the response limit",
1554        )
1555        .expect_err("id limit");
1556        assert_eq!(
1557            error.to_string(),
1558            "provider tool-call id exceeded the response limit"
1559        );
1560
1561        let fragment = "x".repeat(MAX_PROVIDER_TOOL_NAME_BYTES);
1562        let mut name = String::new();
1563        append_provider_field(
1564            &mut name,
1565            &fragment,
1566            MAX_PROVIDER_TOOL_NAME_BYTES,
1567            "provider tool-call name exceeded the response limit",
1568        )
1569        .expect("name within limit");
1570        let error = append_provider_field(
1571            &mut name,
1572            "x",
1573            MAX_PROVIDER_TOOL_NAME_BYTES,
1574            "provider tool-call name exceeded the response limit",
1575        )
1576        .expect_err("name limit");
1577        assert_eq!(
1578            error.to_string(),
1579            "provider tool-call name exceeded the response limit"
1580        );
1581    }
1582
1583    #[test]
1584    fn caps_provider_error_text_without_copying_the_full_message() {
1585        let message = "x".repeat(MAX_PROVIDER_ERROR_BYTES + 1);
1586        let value = serde_json::json!({"error": {"message": message}});
1587        assert_eq!(
1588            provider_error_message(&value),
1589            Some("provider error text exceeded the response limit")
1590        );
1591    }
1592
1593    #[test]
1594    fn reports_midstream_error_without_echoing_provider_body() {
1595        let stream = b"data: {\"error\":{\"message\":\"bad request\"}}\n\n";
1596        let error = parse_sse(&mut Cursor::new(stream), |value| {
1597            if let Some(message) = provider_error_message(&value) {
1598                return Err(ProviderError::new(format!(
1599                    "provider stream error: {message}"
1600                )));
1601            }
1602            Ok(())
1603        })
1604        .expect_err("error");
1605        assert_eq!(error.to_string(), "provider stream error: bad request");
1606    }
1607}