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