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