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