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