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