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