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 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 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 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}