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