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 finite shell command in the session starting directory.",
240 "parameters": {"type": "object", "properties": {"command": {"type": "string"}}, "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 }
1210
1211 #[test]
1212 fn compaction_request_does_not_include_tools() {
1213 let normal = chat_request(
1214 "model",
1215 &[ChatMessage::user("hello".to_owned())],
1216 &None,
1217 true,
1218 );
1219 let compact = chat_request(
1220 "model",
1221 &[ChatMessage::user("hello".to_owned())],
1222 &None,
1223 false,
1224 );
1225
1226 assert!(normal.get("tools").is_some());
1227 assert!(normal.get("max_tokens").is_none());
1228 assert!(compact.get("tools").is_none());
1229 assert_eq!(compact["max_tokens"], COMPACTION_MAX_SUMMARY_TOKENS);
1230 }
1231
1232 #[test]
1233 fn model_catalog_reads_nested_and_compatible_effort_metadata() {
1234 let openrouter = serde_json::json!({
1235 "reasoning": {
1236 "supported_efforts": ["max", "xhigh", "high", "medium", "low", "none"]
1237 }
1238 });
1239 assert_eq!(
1240 model_efforts(&openrouter),
1241 Some(vec![
1242 "max".to_owned(),
1243 "xhigh".to_owned(),
1244 "high".to_owned(),
1245 "medium".to_owned(),
1246 "low".to_owned(),
1247 "none".to_owned(),
1248 ])
1249 );
1250
1251 let compatible = serde_json::json!({
1252 "supported_reasoning_efforts": ["light", "medium", "max", "light", ""]
1253 });
1254 assert_eq!(
1255 model_efforts(&compatible),
1256 Some(vec![
1257 "light".to_owned(),
1258 "medium".to_owned(),
1259 "max".to_owned()
1260 ])
1261 );
1262 assert_eq!(model_efforts(&serde_json::json!({})), None);
1263 }
1264
1265 #[test]
1266 fn model_catalog_context_window_matches_configured_model() {
1267 let payload = serde_json::json!({
1268 "data": [
1269 {"id": "other", "context_length": 8_000},
1270 {"id": "provider/model", "context_length": 128_000}
1271 ]
1272 });
1273
1274 assert_eq!(
1275 context_window_from_models(&payload, "provider/model"),
1276 Some(128_000)
1277 );
1278 assert_eq!(context_window_from_models(&payload, "missing"), None);
1279 }
1280
1281 #[test]
1282 fn model_catalog_context_window_accepts_provider_fallback_fields() {
1283 let payload = serde_json::json!({
1284 "data": [{
1285 "id": "provider/model",
1286 "top_provider": {"context_length": 64_000}
1287 }]
1288 });
1289
1290 assert_eq!(
1291 context_window_from_models(&payload, "provider/model"),
1292 Some(64_000)
1293 );
1294 }
1295
1296 #[test]
1297 fn parses_sse_comments_multiline_data_and_done() {
1298 let stream = b": keep-alive\n\ndata: \n\ndata: {\"choices\":[]\ndata: }\n\ndata: [DONE]\n";
1299 let mut values = Vec::new();
1300 let result = parse_sse(&mut Cursor::new(stream), |value| {
1301 values.push(value);
1302 Ok(())
1303 })
1304 .expect("SSE");
1305 assert!(result.received_payload);
1306 assert!(result.received_done);
1307 assert_eq!(values.len(), 1);
1308 assert!(values[0]["choices"].is_array());
1309 }
1310
1311 #[test]
1312 fn parses_text_and_fragmented_tool_calls() {
1313 let first = serde_json::json!({
1314 "choices": [{"delta": {"content": "hi"}}]
1315 });
1316 let second = serde_json::json!({
1317 "choices": [{
1318 "delta": {
1319 "tool_calls": [{
1320 "index": 0,
1321 "id": "c1",
1322 "function": {"name": "cmd", "arguments": "{command:"}
1323 }]
1324 }
1325 }]
1326 });
1327 let third = serde_json::json!({
1328 "choices": [{
1329 "delta": {
1330 "tool_calls": [{
1331 "index": 0,
1332 "function": {"arguments": "pwd}"}
1333 }]
1334 }
1335 }]
1336 });
1337 let stream =
1338 format!("data: {first}\n\n data: {second}\n\ndata: {third}\n\ndata: [DONE]\n\n")
1339 .replace(" data:", "data:");
1340 let mut content = String::new();
1341 let mut calls = BTreeMap::<usize, PartialToolCall>::new();
1342 let result = parse_sse(&mut Cursor::new(stream.as_bytes()), |value| {
1343 let choice = &value["choices"][0];
1344 let delta = &choice["delta"];
1345 if let Some(text) = delta["content"].as_str() {
1346 content.push_str(text);
1347 }
1348 if let Some(tool_calls) = delta["tool_calls"].as_array() {
1349 for call in tool_calls {
1350 let index = call["index"].as_u64().expect("index") as usize;
1351 let partial = calls.entry(index).or_default();
1352 partial.id.push_str(call["id"].as_str().unwrap_or(""));
1353 partial
1354 .name
1355 .push_str(call["function"]["name"].as_str().unwrap_or(""));
1356 partial
1357 .arguments
1358 .push_str(call["function"]["arguments"].as_str().unwrap_or(""));
1359 }
1360 }
1361 Ok(())
1362 })
1363 .expect("SSE");
1364 assert!(result.received_payload);
1365 assert!(result.received_done);
1366 assert_eq!(content, "hi");
1367 assert_eq!(calls[&0].id, "c1");
1368 assert_eq!(calls[&0].name, "cmd");
1369 assert_eq!(calls[&0].arguments, "{command:pwd}");
1370 }
1371
1372 #[test]
1373 fn cancellable_accumulator_accepts_more_than_sixty_four_tool_calls() {
1374 let mut accumulator = ProviderAccumulator::default();
1375 let tool_calls = (0..65)
1376 .map(|index| {
1377 serde_json::json!({
1378 "index": index,
1379 "id": format!("call-{index}"),
1380 "function": {
1381 "name": "cmd",
1382 "arguments": "{\"command\":\"true\"}"
1383 }
1384 })
1385 })
1386 .collect::<Vec<_>>();
1387 accumulator
1388 .on_data(
1389 serde_json::json!({
1390 "choices": [{
1391 "delta": {"tool_calls": tool_calls},
1392 "finish_reason": "tool_calls"
1393 }]
1394 }),
1395 "provider-secret",
1396 &mut |_| Ok(()),
1397 )
1398 .expect("tool-call chunk");
1399
1400 let turn = accumulator.finish().expect("provider turn");
1401 assert_eq!(turn.tool_calls.len(), 65);
1402 }
1403
1404 #[test]
1405 fn reasoning_stream_event_is_emitted_once_before_assistant_text() {
1406 let mut accumulator = ProviderAccumulator::default();
1407 let mut events = Vec::new();
1408 let mut on_event = |event| {
1409 match event {
1410 ProviderStreamEvent::ReasoningStarted => events.push("started".to_owned()),
1411 ProviderStreamEvent::Text(text) => events.push(text),
1412 }
1413 Ok(())
1414 };
1415
1416 accumulator
1417 .on_data(
1418 serde_json::json!({
1419 "choices": [{
1420 "delta": {
1421 "reasoning_details": [{"type": "reasoning.text", "text": "thinking"}]
1422 }
1423 }]
1424 }),
1425 "provider-secret",
1426 &mut on_event,
1427 )
1428 .expect("reasoning chunk");
1429 accumulator
1430 .on_data(
1431 serde_json::json!({
1432 "choices": [{"delta": {"content": "answer"}}]
1433 }),
1434 "provider-secret",
1435 &mut on_event,
1436 )
1437 .expect("answer chunk");
1438
1439 assert_eq!(events, vec!["started".to_owned(), "answer".to_owned()]);
1440 }
1441
1442 #[test]
1443 fn accumulates_reasoning_details_with_fragmented_tool_calls() {
1444 let mut accumulator = ProviderAccumulator::default();
1445 accumulator
1446 .on_data(
1447 serde_json::json!({
1448 "choices": [{
1449 "delta": {
1450 "reasoning_details": [{
1451 "type": "reasoning.text",
1452 "text": "part one"
1453 }]
1454 }
1455 }]
1456 }),
1457 "provider-secret",
1458 &mut |_| Ok(()),
1459 )
1460 .expect("first provider chunk");
1461 accumulator
1462 .on_data(
1463 serde_json::json!({
1464 "choices": [{
1465 "delta": {
1466 "reasoning_details": [{
1467 "type": "reasoning.text",
1468 "text": "part two"
1469 }],
1470 "tool_calls": [{
1471 "index": 0,
1472 "id": "call-1",
1473 "function": {
1474 "name": "cmd",
1475 "arguments": "{\"command\":\"true\"}"
1476 }
1477 }]
1478 },
1479 "finish_reason": "tool_calls"
1480 }]
1481 }),
1482 "provider-secret",
1483 &mut |_| Ok(()),
1484 )
1485 .expect("second provider chunk");
1486
1487 let partial = accumulator.partial_turn();
1488 assert_eq!(partial.reasoning_details.len(), 2);
1489 let turn = accumulator.finish().expect("provider turn");
1490 assert_eq!(
1491 turn.reasoning_details,
1492 vec![
1493 json!({"type": "reasoning.text", "text": "part one"}),
1494 json!({"type": "reasoning.text", "text": "part two"}),
1495 ]
1496 );
1497 assert_eq!(turn.tool_calls.len(), 1);
1498 assert_eq!(turn.tool_calls[0].name, "cmd");
1499 }
1500
1501 #[test]
1502 fn accumulates_many_small_reasoning_details_and_rejects_overflow_atomically() {
1503 const FRAGMENT_COUNT: usize = 4096;
1504 let mut details = Vec::new();
1505 let mut serialized_bytes = 0;
1506 let delta = serde_json::json!({
1507 "reasoning_details": [{
1508 "type": "reasoning.text",
1509 "text": "x".repeat(64)
1510 }]
1511 });
1512 for _ in 0..FRAGMENT_COUNT {
1513 append_reasoning_details(&mut details, &mut serialized_bytes, &delta)
1514 .expect("small reasoning detail");
1515 }
1516 assert_eq!(details.len(), FRAGMENT_COUNT);
1517
1518 let first_chunk_delta = serde_json::json!({
1519 "reasoning_details": [{
1520 "type": "reasoning.text",
1521 "text": "x".repeat(500 * 1024)
1522 }]
1523 });
1524 let first_chunk_bytes = serde_json::to_vec(&first_chunk_delta["reasoning_details"])
1525 .expect("first reasoning detail chunk")
1526 .len();
1527 assert!(first_chunk_bytes < MAX_PROVIDER_REASONING_DETAILS_BYTES);
1528 append_reasoning_details(&mut details, &mut serialized_bytes, &first_chunk_delta)
1529 .expect("first individually bounded reasoning detail chunk");
1530 assert!(serialized_bytes <= MAX_PROVIDER_REASONING_DETAILS_BYTES);
1531
1532 let second_chunk_delta = serde_json::json!({
1533 "reasoning_details": [{
1534 "type": "reasoning.text",
1535 "text": "x".repeat(200 * 1024)
1536 }]
1537 });
1538 let second_chunk_bytes = serde_json::to_vec(&second_chunk_delta["reasoning_details"])
1539 .expect("second reasoning detail chunk")
1540 .len();
1541 assert!(second_chunk_bytes < MAX_PROVIDER_REASONING_DETAILS_BYTES);
1542 assert!(
1543 serialized_bytes
1544 .saturating_add(second_chunk_bytes)
1545 .saturating_sub(1)
1546 > MAX_PROVIDER_REASONING_DETAILS_BYTES
1547 );
1548
1549 let prior_details = details.clone();
1550 let prior_bytes = serialized_bytes;
1551 let error =
1552 append_reasoning_details(&mut details, &mut serialized_bytes, &second_chunk_delta)
1553 .expect_err("reasoning details limit");
1554
1555 assert_eq!(
1556 error.to_string(),
1557 "provider reasoning details exceeded the response limit"
1558 );
1559 assert_eq!(details, prior_details);
1560 assert_eq!(serialized_bytes, prior_bytes);
1561 assert_eq!(
1562 serde_json::to_vec(&details)
1563 .expect("accumulated reasoning details")
1564 .len(),
1565 serialized_bytes
1566 );
1567 assert!(serialized_bytes <= MAX_PROVIDER_REASONING_DETAILS_BYTES);
1568 }
1569
1570 #[test]
1571 fn rejects_reasoning_details_that_exceed_the_serialized_response_limit_before_retaining_them() {
1572 let mut accumulator = ProviderAccumulator::default();
1573 let retained = serde_json::json!({
1574 "choices": [{
1575 "delta": {
1576 "reasoning_details": [{
1577 "type": "reasoning.text",
1578 "text": "retained"
1579 }]
1580 }
1581 }]
1582 });
1583 accumulator
1584 .on_data(retained, "provider-secret", &mut |_| Ok(()))
1585 .expect("details within limit");
1586
1587 let oversized = "x".repeat(MAX_PROVIDER_REASONING_DETAILS_BYTES);
1588 let error = accumulator
1589 .on_data(
1590 serde_json::json!({
1591 "choices": [{
1592 "delta": {
1593 "reasoning_details": [{"text": oversized}]
1594 }
1595 }]
1596 }),
1597 "provider-secret",
1598 &mut |_| Ok(()),
1599 )
1600 .expect_err("reasoning details limit");
1601 assert_eq!(
1602 error.to_string(),
1603 "provider reasoning details exceeded the response limit"
1604 );
1605 assert_eq!(
1606 accumulator.reasoning_details,
1607 vec![serde_json::json!({
1608 "type": "reasoning.text",
1609 "text": "retained"
1610 })]
1611 );
1612 }
1613
1614 #[test]
1615 fn accepts_compatible_finish_reasons_and_rejects_incomplete_ones() {
1616 for reason in [
1617 None,
1618 Some(Value::Null),
1619 Some(Value::String("stop".to_owned())),
1620 ] {
1621 let mut choice = serde_json::json!({"delta": {}});
1622 if let Some(reason) = reason {
1623 choice["finish_reason"] = reason;
1624 }
1625 validate_finish_reason(&choice).expect("compatible finish reason");
1626 }
1627 for reason in ["tool_calls", "function_call"] {
1628 validate_finish_reason(&serde_json::json!({
1629 "delta": {},
1630 "finish_reason": reason
1631 }))
1632 .expect("tool finish reason");
1633 }
1634 for reason in ["length", "content_filter", "error"] {
1635 assert!(validate_finish_reason(&serde_json::json!({
1636 "delta": {},
1637 "finish_reason": reason
1638 }))
1639 .is_err());
1640 }
1641 }
1642
1643 #[test]
1644 fn rejects_api_keys_that_conflict_with_fixed_literals() {
1645 for (index, secret) in [
1646 "session",
1647 "tool",
1648 "cmd",
1649 "command",
1650 "finite",
1651 "0",
1652 ":",
1653 "[REDACTED]",
1654 ]
1655 .into_iter()
1656 .enumerate()
1657 {
1658 let environment = format!("LUCY_PROVIDER_CONFLICT_{}_{}", std::process::id(), index);
1659 std::env::set_var(&environment, secret);
1660 let settings = LlmSettings {
1661 base_url: "http://localhost".to_owned(),
1662 model: "model".to_owned(),
1663 api_key_env: environment.clone(),
1664 effort: None,
1665 };
1666 let error = match Provider::new(&settings) {
1667 Ok(_) => panic!("fixed literal conflict should be rejected: {secret}"),
1668 Err(error) => error,
1669 };
1670 assert!(error.to_string().contains("structured output"));
1671 assert!(!error.to_string().contains(secret));
1672 std::env::remove_var(environment);
1673 }
1674 }
1675
1676 #[test]
1677 fn accepts_a_normal_long_provider_key() {
1678 let environment = format!("LUCY_PROVIDER_NORMAL_{}", std::process::id());
1679 std::env::set_var(&environment, "provider-secret");
1680 let settings = LlmSettings {
1681 base_url: "http://localhost".to_owned(),
1682 model: "model".to_owned(),
1683 api_key_env: environment.clone(),
1684 effort: None,
1685 };
1686 assert!(Provider::new(&settings).is_ok());
1687 std::env::remove_var(environment);
1688 }
1689
1690 #[test]
1691 fn accepts_a_configurable_effort() {
1692 let environment = format!("LUCY_PROVIDER_EFFORT_OK_{}", std::process::id());
1693 std::env::set_var(&environment, "provider-secret");
1694 let settings = LlmSettings {
1695 base_url: "http://localhost".to_owned(),
1696 model: "model".to_owned(),
1697 api_key_env: environment.clone(),
1698 effort: Some("high".to_owned()),
1699 };
1700 assert!(Provider::new(&settings).is_ok());
1701 std::env::remove_var(environment);
1702 }
1703
1704 #[test]
1705 fn empty_effort_is_rejected_without_echoing_the_key() {
1706 let environment = format!("LUCY_PROVIDER_EFFORT_EMPTY_{}", std::process::id());
1707 std::env::set_var(&environment, "provider-secret");
1708 for effort in ["", " ", "\t"] {
1709 let settings = LlmSettings {
1710 base_url: "http://localhost".to_owned(),
1711 model: "model".to_owned(),
1712 api_key_env: environment.clone(),
1713 effort: Some(effort.to_owned()),
1714 };
1715 let error = match Provider::new(&settings) {
1716 Ok(_) => panic!("empty effort should be rejected: {effort:?}"),
1717 Err(error) => error,
1718 };
1719 assert!(error.to_string().contains("llm.effort must not be empty"));
1720 assert!(!error.to_string().contains("provider-secret"));
1721 }
1722 std::env::remove_var(environment);
1723 }
1724
1725 #[test]
1726 fn missing_api_key_error_does_not_echo_the_environment_name() {
1727 let environment = format!("LUCY_MISSING_KEY_{}", std::process::id());
1728 std::env::remove_var(&environment);
1729 let settings = LlmSettings {
1730 base_url: "http://localhost".to_owned(),
1731 model: "model".to_owned(),
1732 api_key_env: environment.clone(),
1733 effort: None,
1734 };
1735 let error = match Provider::new(&settings) {
1736 Ok(_) => panic!("missing key should be rejected"),
1737 Err(error) => error,
1738 };
1739 assert_eq!(error.to_string(), "missing provider API key");
1740 assert!(!error.to_string().contains(&environment));
1741 }
1742
1743 #[test]
1744 fn cancellable_stream_stops_a_stalled_provider_without_waiting_for_timeout() {
1745 let listener = TcpListener::bind(("127.0.0.1", 0)).expect("listener");
1746 let address = listener.local_addr().expect("address");
1747 let (sent, sent_receiver) = mpsc::channel();
1748 let server = thread::spawn(move || {
1749 let (mut stream, _) = listener.accept().expect("request");
1750 let mut request = std::io::BufReader::new(stream.try_clone().expect("clone"));
1751 let mut content_length = 0;
1752 loop {
1753 let mut line = String::new();
1754 request.read_line(&mut line).expect("header");
1755 if line == "\r\n" {
1756 break;
1757 }
1758 if let Some(value) = line.strip_prefix("Content-Length:") {
1759 content_length = value.trim().parse::<usize>().expect("length");
1760 }
1761 }
1762 let mut body = vec![0; content_length];
1763 request.read_exact(&mut body).expect("body");
1764
1765 let payload = serde_json::json!({
1766 "choices": [{"delta": {"content": "partial"}, "finish_reason": null}]
1767 });
1768 let event = format!("data: {payload}\n\n");
1769 let response = format!(
1770 "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",
1771 event.len(), event
1772 );
1773 stream.write_all(response.as_bytes()).expect("response");
1774 stream.flush().expect("flush");
1775 sent.send(()).expect("body readiness");
1776 thread::sleep(Duration::from_millis(500));
1777 });
1778
1779 let environment = format!("LUCY_PROVIDER_CANCEL_{}", std::process::id());
1780 std::env::set_var(&environment, "provider-secret");
1781 let provider = Provider::new(&LlmSettings {
1782 base_url: format!("http://{address}/v1"),
1783 model: "model".to_owned(),
1784 api_key_env: environment.clone(),
1785 effort: None,
1786 })
1787 .expect("provider");
1788 let token = CancellationToken::new();
1789 let worker_token = token.clone();
1790 let worker = thread::spawn(move || {
1791 let mut received = String::new();
1792 let result = provider.stream_chat_cancellable(
1793 &[ChatMessage::user("hello".to_owned())],
1794 &mut |text| {
1795 received.push_str(text);
1796 Ok(())
1797 },
1798 &worker_token,
1799 );
1800 (result, received)
1801 });
1802 sent_receiver
1803 .recv_timeout(Duration::from_secs(1))
1804 .expect("body was sent");
1805 let started = Instant::now();
1806 assert!(token.cancel());
1807 let (result, received) = worker.join().expect("provider worker");
1808 assert!(started.elapsed() < Duration::from_millis(400));
1809 let error = result.expect_err("cancellation");
1810 assert!(error.is_cancelled());
1811 assert!(received.is_empty() || received == "partial");
1812 server.join().expect("server");
1813 std::env::remove_var(environment);
1814 }
1815
1816 #[test]
1817 fn rejects_an_oversized_sse_line_before_json_parsing() {
1818 let stream = format!("data: {}\n\n", "x".repeat(MAX_SSE_LINE_BYTES));
1819 let error = parse_sse(&mut Cursor::new(stream.as_bytes()), |_| Ok(())).expect_err("limit");
1820 assert_eq!(
1821 error.to_string(),
1822 "provider SSE line exceeded the response limit"
1823 );
1824 }
1825
1826 #[test]
1827 fn rejects_an_oversized_sse_data_event_before_json_parsing() {
1828 let payload = "x".repeat(MAX_SSE_LINE_BYTES - "data: ".len());
1829 let line_count = MAX_SSE_EVENT_BYTES / payload.len() + 2;
1830 let mut stream = String::new();
1831 for _ in 0..line_count {
1832 stream.push_str("data: ");
1833 stream.push_str(&payload);
1834 stream.push('\n');
1835 }
1836 stream.push('\n');
1837
1838 let error = parse_sse(&mut Cursor::new(stream.as_bytes()), |_| Ok(())).expect_err("limit");
1839 assert_eq!(
1840 error.to_string(),
1841 "provider SSE data event exceeded the response limit"
1842 );
1843 }
1844
1845 #[test]
1846 fn rejects_an_oversized_sse_stream_of_ignored_fields() {
1847 let line = format!("ignored: {}\n", "x".repeat(1024));
1848 let mut stream = Vec::new();
1849 while stream.len() <= MAX_SSE_STREAM_BYTES {
1850 stream.extend_from_slice(line.as_bytes());
1851 }
1852
1853 let error = parse_sse(&mut Cursor::new(stream), |_| Ok(())).expect_err("limit");
1854 assert_eq!(
1855 error.to_string(),
1856 "provider SSE stream exceeded the response limit"
1857 );
1858 }
1859
1860 #[test]
1861 fn rejects_too_many_empty_sse_data_lines() {
1862 let stream = format!("{}\n", "data:\n".repeat(MAX_SSE_DATA_LINES + 1));
1863 let error = parse_sse(&mut Cursor::new(stream.as_bytes()), |_| Ok(())).expect_err("limit");
1864 assert_eq!(
1865 error.to_string(),
1866 "provider SSE data line count exceeded the response limit"
1867 );
1868 }
1869
1870 #[test]
1871 fn reports_eof_before_done_as_incomplete() {
1872 let stream = b"data: {\"choices\":[{\"delta\":{\"content\":\"partial\"}}]}\n\n";
1873 let result = parse_sse(&mut Cursor::new(stream), |_| Ok(())).expect("SSE parse");
1874 assert!(result.received_payload);
1875 assert!(!result.received_done);
1876 }
1877
1878 #[test]
1879 fn reports_empty_non_sse_input_without_payload_or_done() {
1880 let result =
1881 parse_sse(&mut Cursor::new(b"not an SSE response\n"), |_| Ok(())).expect("SSE parse");
1882 assert_eq!(result, SseParseResult::default());
1883 }
1884
1885 #[test]
1886 fn caps_accumulated_tool_call_id_and_name_fields() {
1887 let fragment = "x".repeat(MAX_PROVIDER_TOOL_CALL_ID_BYTES);
1888 let mut id = String::new();
1889 append_provider_field(
1890 &mut id,
1891 &fragment,
1892 MAX_PROVIDER_TOOL_CALL_ID_BYTES,
1893 "provider tool-call id exceeded the response limit",
1894 )
1895 .expect("id within limit");
1896 let error = append_provider_field(
1897 &mut id,
1898 "x",
1899 MAX_PROVIDER_TOOL_CALL_ID_BYTES,
1900 "provider tool-call id exceeded the response limit",
1901 )
1902 .expect_err("id limit");
1903 assert_eq!(
1904 error.to_string(),
1905 "provider tool-call id exceeded the response limit"
1906 );
1907
1908 let fragment = "x".repeat(MAX_PROVIDER_TOOL_NAME_BYTES);
1909 let mut name = String::new();
1910 append_provider_field(
1911 &mut name,
1912 &fragment,
1913 MAX_PROVIDER_TOOL_NAME_BYTES,
1914 "provider tool-call name exceeded the response limit",
1915 )
1916 .expect("name within limit");
1917 let error = append_provider_field(
1918 &mut name,
1919 "x",
1920 MAX_PROVIDER_TOOL_NAME_BYTES,
1921 "provider tool-call name exceeded the response limit",
1922 )
1923 .expect_err("name limit");
1924 assert_eq!(
1925 error.to_string(),
1926 "provider tool-call name exceeded the response limit"
1927 );
1928 }
1929
1930 #[test]
1931 fn caps_provider_error_text_without_copying_the_full_message() {
1932 let message = "x".repeat(MAX_PROVIDER_ERROR_BYTES + 1);
1933 let value = serde_json::json!({"error": {"message": message}});
1934 assert_eq!(
1935 provider_error_message(&value),
1936 Some("provider error text exceeded the response limit")
1937 );
1938 }
1939
1940 fn read_request_headers(stream: &TcpStream) {
1941 let mut reader = BufReader::new(stream.try_clone().expect("clone request"));
1942 loop {
1943 let mut line = String::new();
1944 reader.read_line(&mut line).expect("request header");
1945 if line == "\r\n" || line.is_empty() {
1946 return;
1947 }
1948 }
1949 }
1950
1951 fn response_body(text: &str) -> String {
1952 let payload = serde_json::json!({
1953 "choices": [{"delta": {"content": text}, "finish_reason": null}]
1954 });
1955 let finish = serde_json::json!({
1956 "choices": [{"delta": {}, "finish_reason": "stop"}]
1957 });
1958 format!("data: {payload}\n\ndata: {finish}\n\ndata: [DONE]\n\n")
1959 }
1960
1961 fn provider_for(address: std::net::SocketAddr, read_timeout: Duration) -> (Provider, String) {
1962 let environment = format!(
1963 "LUCY_PROVIDER_STREAM_TEST_{}_{}",
1964 std::process::id(),
1965 address.port()
1966 );
1967 std::env::set_var(&environment, "provider-secret");
1968 let settings = LlmSettings {
1969 base_url: format!("http://{address}/v1"),
1970 model: "model".to_owned(),
1971 api_key_env: environment.clone(),
1972 effort: None,
1973 };
1974 let mut provider = Provider::new(&settings).expect("provider");
1975 provider.async_client = AsyncClient::builder()
1976 .connect_timeout(Duration::from_secs(1))
1977 .read_timeout(read_timeout)
1978 .build()
1979 .expect("test async client");
1980 (provider, environment)
1981 }
1982
1983 #[test]
1984 fn worker_stream_can_exceed_idle_interval_without_total_deadline() {
1985 let listener = TcpListener::bind(("127.0.0.1", 0)).expect("listener");
1986 let address = listener.local_addr().expect("address");
1987 let parts = (0..5)
1988 .map(|index| {
1989 let payload = serde_json::json!({
1990 "choices": [{
1991 "delta": {"content": format!("part-{index}")},
1992 "finish_reason": null
1993 }]
1994 });
1995 format!("data: {payload}\n\n")
1996 })
1997 .chain([
1998 format!(
1999 "data: {}\n\n",
2000 serde_json::json!({
2001 "choices": [{"delta": {}, "finish_reason": "stop"}]
2002 })
2003 ),
2004 "data: [DONE]\n\n".to_owned(),
2005 ])
2006 .collect::<Vec<_>>();
2007 let body = parts.concat();
2008 let server = thread::spawn(move || {
2009 let (mut stream, _) = listener.accept().expect("request");
2010 read_request_headers(&stream);
2011 let header = format!(
2012 "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
2013 body.len()
2014 );
2015 stream.write_all(header.as_bytes()).expect("header");
2016 stream.flush().expect("header flush");
2017 for part in parts {
2018 stream.write_all(part.as_bytes()).expect("SSE part");
2019 stream.flush().expect("SSE flush");
2020 thread::sleep(Duration::from_millis(25));
2021 }
2022 });
2023
2024 let (provider, environment) = provider_for(address, Duration::from_millis(80));
2025 let cancellation = CancellationToken::new();
2026 let started = Instant::now();
2027 let mut output = String::new();
2028 let turn = provider
2029 .stream_chat_cancellable_with_options(
2030 &[ChatMessage::user("worker task".to_owned())],
2031 &mut |text| {
2032 output.push_str(text);
2033 Ok(())
2034 },
2035 &cancellation,
2036 true,
2037 )
2038 .expect("long worker stream");
2039
2040 assert!(started.elapsed() >= Duration::from_millis(80));
2041 assert_eq!(turn.content, output);
2042 assert!(output.contains("part-4"));
2043 server.join().expect("server");
2044 std::env::remove_var(environment);
2045 }
2046
2047 #[test]
2048 fn retries_a_pre_payload_stream_failure_once_and_classifies_it() {
2049 let listener = TcpListener::bind(("127.0.0.1", 0)).expect("listener");
2050 listener
2051 .set_nonblocking(true)
2052 .expect("nonblocking listener");
2053 let address = listener.local_addr().expect("address");
2054 let body = response_body("retried");
2055 let server = thread::spawn(move || {
2056 let deadline = Instant::now() + Duration::from_secs(2);
2057 for attempt in 0..2 {
2058 let (mut stream, _) = loop {
2059 match listener.accept() {
2060 Ok(connection) => break connection,
2061 Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
2062 assert!(Instant::now() < deadline, "provider did not retry");
2063 thread::sleep(Duration::from_millis(5));
2064 }
2065 Err(error) => panic!("accept: {error}"),
2066 }
2067 };
2068 read_request_headers(&stream);
2069 if attempt == 0 {
2070 let header = "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nContent-Length: 100\r\nConnection: close\r\n\r\n";
2071 stream.write_all(header.as_bytes()).expect("failed header");
2072 stream.flush().expect("failed flush");
2073 } else {
2074 let header = format!(
2075 "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
2076 body.len()
2077 );
2078 stream.write_all(header.as_bytes()).expect("success header");
2079 stream.write_all(body.as_bytes()).expect("success body");
2080 stream.flush().expect("success flush");
2081 }
2082 }
2083 });
2084
2085 let (provider, environment) = provider_for(address, Duration::from_secs(1));
2086 let cancellation = CancellationToken::new();
2087 let mut output = String::new();
2088 let turn = provider
2089 .stream_chat_cancellable_with_options(
2090 &[ChatMessage::user("retry".to_owned())],
2091 &mut |text| {
2092 output.push_str(text);
2093 Ok(())
2094 },
2095 &cancellation,
2096 true,
2097 )
2098 .expect("retry succeeds");
2099
2100 assert_eq!(turn.content, "retried");
2101 assert_eq!(output, "retried");
2102 server.join().expect("server");
2103 std::env::remove_var(environment);
2104 }
2105
2106 #[test]
2107 fn does_not_retry_after_partial_provider_output() {
2108 let listener = TcpListener::bind(("127.0.0.1", 0)).expect("listener");
2109 listener
2110 .set_nonblocking(true)
2111 .expect("nonblocking listener");
2112 let address = listener.local_addr().expect("address");
2113 let partial = format!(
2114 "data: {}\n\n",
2115 serde_json::json!({
2116 "choices": [{"delta": {"content": "partial"}, "finish_reason": null}]
2117 })
2118 );
2119 let server = thread::spawn(move || {
2120 let deadline = Instant::now() + Duration::from_secs(2);
2121 let (mut stream, _) = loop {
2122 match listener.accept() {
2123 Ok(connection) => break connection,
2124 Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => {
2125 assert!(Instant::now() < deadline, "provider request missing");
2126 thread::sleep(Duration::from_millis(5));
2127 }
2128 Err(error) => panic!("accept: {error}"),
2129 }
2130 };
2131 read_request_headers(&stream);
2132 let header = format!(
2133 "HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
2134 partial.len() + 10
2135 );
2136 stream.write_all(header.as_bytes()).expect("partial header");
2137 stream.write_all(partial.as_bytes()).expect("partial body");
2138 stream.flush().expect("partial flush");
2139 thread::sleep(Duration::from_millis(150));
2140 assert!(matches!(
2141 listener.accept(),
2142 Err(error) if error.kind() == std::io::ErrorKind::WouldBlock
2143 ));
2144 });
2145
2146 let (provider, environment) = provider_for(address, Duration::from_secs(1));
2147 let cancellation = CancellationToken::new();
2148 let mut output = String::new();
2149 let error = provider
2150 .stream_chat_cancellable_with_options(
2151 &[ChatMessage::user("partial".to_owned())],
2152 &mut |text| {
2153 output.push_str(text);
2154 Ok(())
2155 },
2156 &cancellation,
2157 true,
2158 )
2159 .expect_err("partial stream must fail");
2160
2161 let message = error.to_string();
2162 assert!(message.contains("provider stream read failed"));
2163 assert!([
2164 "(timeout)",
2165 "(connection)",
2166 "(body)",
2167 "(decode)",
2168 "(request)",
2169 "(transport)",
2170 ]
2171 .iter()
2172 .any(|kind| message.contains(kind)));
2173 assert_eq!(output, "partial");
2174 server.join().expect("server");
2175 std::env::remove_var(environment);
2176 }
2177
2178 #[test]
2179 fn reports_midstream_error_without_echoing_provider_body() {
2180 let stream = b"data: {\"error\":{\"message\":\"bad request\"}}\n\n";
2181 let error = parse_sse(&mut Cursor::new(stream), |value| {
2182 if let Some(message) = provider_error_message(&value) {
2183 return Err(ProviderError::new(format!(
2184 "provider stream error: {message}"
2185 )));
2186 }
2187 Ok(())
2188 })
2189 .expect_err("error");
2190 assert_eq!(error.to_string(), "provider stream error: bad request");
2191 }
2192}