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