1use serde::{Deserialize, Serialize, ser::SerializeStruct};
4use std::fmt;
5
6use crate::sanitizer::sanitize_provider_diagnostic;
7
8#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
9pub enum BackendKind {
10 Gemini,
11 OpenAI,
12 Anthropic,
13 DeepSeek,
14 Meta,
15 Mistral,
16 OpenRouter,
17 Ollama,
18 LlamaCpp,
19 ZAI,
20 Moonshot,
21 HuggingFace,
22 Minimax,
23 MiMo,
24 OpenCodeZen,
25 OpenCodeGo,
26 Qwen,
27 StepFun,
28 Evolink,
29 Poolside,
30 Xai,
31 Nvidia,
32 MergeGateway,
33 Vercel,
34}
35
36#[derive(Debug, Clone, Default, Serialize, Deserialize, PartialEq, Eq)]
37pub struct Usage {
38 pub prompt_tokens: u32,
39 pub completion_tokens: u32,
40 pub total_tokens: u32,
41 pub cached_prompt_tokens: Option<u32>,
42 pub cache_creation_tokens: Option<u32>,
43 pub cache_read_tokens: Option<u32>,
44 #[serde(default, skip_serializing_if = "Option::is_none")]
47 pub iterations: Option<Vec<serde_json::Value>>,
48}
49
50#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
56pub struct UsageTotals {
57 pub prompt_tokens: u32,
58 pub completion_tokens: u32,
59 pub total_tokens: u32,
60 pub cache_read_tokens: u32,
61 pub cache_creation_tokens: u32,
62}
63
64impl Usage {
65 #[inline]
66 pub fn billable_totals(&self) -> UsageTotals {
67 let mut totals = UsageTotals::default();
68 let mut recognized_iteration = false;
69 let mut saw_iteration_cache_read = false;
70 let mut saw_iteration_cache_creation = false;
71
72 if let Some(iterations) = &self.iterations {
73 for iteration in iterations {
74 let Some(object) = iteration.as_object() else {
75 continue;
76 };
77 let Some(iteration_type) = object.get("type").and_then(serde_json::Value::as_str) else {
78 continue;
79 };
80 if !matches!(iteration_type, "message" | "fallback_message" | "advisor_message" | "compaction") {
81 continue;
82 }
83
84 recognized_iteration = true;
85 totals.prompt_tokens = totals.prompt_tokens.saturating_add(json_u32(object.get("input_tokens")));
86 totals.completion_tokens =
87 totals.completion_tokens.saturating_add(json_u32(object.get("output_tokens")));
88
89 if let Some(cache_read) = object.get("cache_read_input_tokens").and_then(serde_json::Value::as_u64) {
90 saw_iteration_cache_read = true;
91 totals.cache_read_tokens = totals.cache_read_tokens.saturating_add(saturating_u32(cache_read));
92 }
93 if let Some(cache_creation) =
94 object.get("cache_creation_input_tokens").and_then(serde_json::Value::as_u64)
95 {
96 saw_iteration_cache_creation = true;
97 totals.cache_creation_tokens =
98 totals.cache_creation_tokens.saturating_add(saturating_u32(cache_creation));
99 }
100 }
101 }
102
103 if !recognized_iteration {
104 totals.prompt_tokens = self.prompt_tokens;
105 totals.completion_tokens = self.completion_tokens;
106 totals.cache_read_tokens = self.cache_read_tokens_or_fallback();
107 totals.cache_creation_tokens = self.cache_creation_tokens_or_zero();
108 } else {
109 if !saw_iteration_cache_read {
114 totals.cache_read_tokens = self.cache_read_tokens_or_fallback();
115 }
116 if !saw_iteration_cache_creation {
117 totals.cache_creation_tokens = self.cache_creation_tokens_or_zero();
118 }
119 }
120
121 totals.total_tokens = totals.prompt_tokens.saturating_add(totals.completion_tokens);
122 totals
123 }
124
125 #[inline]
126 fn has_cache_read_metric(&self) -> bool {
127 self.cache_read_tokens.is_some() || self.cached_prompt_tokens.is_some()
128 }
129
130 #[inline]
131 fn has_any_cache_metrics(&self) -> bool {
132 self.has_cache_read_metric() || self.cache_creation_tokens.is_some()
133 }
134
135 #[inline]
136 pub fn cache_read_tokens_or_fallback(&self) -> u32 {
137 self.cache_read_tokens.or(self.cached_prompt_tokens).unwrap_or(0)
138 }
139
140 #[inline]
141 pub fn cache_creation_tokens_or_zero(&self) -> u32 {
142 self.cache_creation_tokens.unwrap_or(0)
143 }
144
145 #[inline]
146 pub fn cache_hit_rate(&self) -> Option<f64> {
147 if !self.has_any_cache_metrics() {
148 return None;
149 }
150 let read = self.cache_read_tokens_or_fallback() as f64;
151 let creation = self.cache_creation_tokens_or_zero() as f64;
152 let total = read + creation;
153 if total > 0.0 {
154 Some((read / total) * 100.0)
155 } else {
156 None
157 }
158 }
159
160 #[inline]
161 fn is_cache_hit(&self) -> Option<bool> {
162 self.has_any_cache_metrics().then(|| self.cache_read_tokens_or_fallback() > 0)
163 }
164
165 #[inline]
166 fn is_cache_miss(&self) -> Option<bool> {
167 self.has_any_cache_metrics()
168 .then(|| self.cache_creation_tokens_or_zero() > 0 && self.cache_read_tokens_or_fallback() == 0)
169 }
170
171 #[inline]
172 fn total_cache_tokens(&self) -> u32 {
173 let read = self.cache_read_tokens_or_fallback();
174 let creation = self.cache_creation_tokens_or_zero();
175 read + creation
176 }
177
178 #[inline]
179 fn cache_savings_ratio(&self) -> Option<f64> {
180 if !self.has_cache_read_metric() {
181 return None;
182 }
183 let read = self.cache_read_tokens_or_fallback() as f64;
184 let prompt = self.prompt_tokens as f64;
185 if prompt > 0.0 { Some(read / prompt) } else { None }
186 }
187}
188
189#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
191pub struct BalanceInfo {
192 pub display: String,
194 pub is_available: bool,
196}
197
198#[derive(Debug, Clone, Serialize, Deserialize)]
200pub struct DeepSeekBalanceResponse {
201 is_available: bool,
202 balance_infos: Vec<DeepSeekCurrencyBalance>,
203}
204
205#[derive(Debug, Clone, Serialize, Deserialize)]
206pub struct DeepSeekCurrencyBalance {
207 currency: String,
208 total_balance: String,
209 #[serde(default)]
210 granted_balance: String,
211 #[serde(default)]
212 topped_up_balance: String,
213}
214
215impl From<DeepSeekBalanceResponse> for BalanceInfo {
216 fn from(resp: DeepSeekBalanceResponse) -> Self {
217 let display = resp
218 .balance_infos
219 .first()
220 .map(|b| {
221 let symbol = match b.currency.as_str() {
222 "CNY" => "¥",
223 "USD" => "$",
224 _ => &b.currency,
225 };
226 format!("{}{}", b.total_balance, symbol)
227 })
228 .unwrap_or_else(|| "N/A".to_string());
229 BalanceInfo { display, is_available: resp.is_available }
230 }
231}
232
233#[inline]
234fn saturating_u32(value: u64) -> u32 {
235 u32::try_from(value.min(u64::from(u32::MAX))).unwrap_or(u32::MAX)
236}
237
238#[inline]
239fn json_u32(value: Option<&serde_json::Value>) -> u32 {
240 value.and_then(serde_json::Value::as_u64).map_or(0, saturating_u32)
241}
242
243#[cfg(test)]
244mod usage_tests {
245 use super::Usage;
246 use serde_json::json;
247
248 #[test]
249 fn cache_helpers_fall_back_to_cached_prompt_tokens() {
250 let usage = Usage {
251 prompt_tokens: 1_000,
252 completion_tokens: 200,
253 total_tokens: 1_200,
254 cached_prompt_tokens: Some(600),
255 cache_creation_tokens: Some(150),
256 cache_read_tokens: None,
257 iterations: None,
258 };
259
260 assert_eq!(usage.cache_read_tokens_or_fallback(), 600);
261 assert_eq!(usage.cache_creation_tokens_or_zero(), 150);
262 assert_eq!(usage.total_cache_tokens(), 750);
263 assert_eq!(usage.is_cache_hit(), Some(true));
264 assert_eq!(usage.is_cache_miss(), Some(false));
265 assert_eq!(usage.cache_savings_ratio(), Some(0.6));
266 assert_eq!(usage.cache_hit_rate(), Some(80.0));
267 }
268
269 #[test]
270 fn cache_helpers_preserve_unknown_without_metrics() {
271 let usage = Usage {
272 prompt_tokens: 1_000,
273 completion_tokens: 200,
274 total_tokens: 1_200,
275 cached_prompt_tokens: None,
276 cache_creation_tokens: None,
277 cache_read_tokens: None,
278 iterations: None,
279 };
280
281 assert_eq!(usage.total_cache_tokens(), 0);
282 assert_eq!(usage.is_cache_hit(), None);
283 assert_eq!(usage.is_cache_miss(), None);
284 assert_eq!(usage.cache_savings_ratio(), None);
285 assert_eq!(usage.cache_hit_rate(), None);
286 }
287
288 #[test]
289 fn billable_totals_aggregate_recognized_iterations() {
290 let usage = Usage {
291 prompt_tokens: 10,
292 completion_tokens: 1,
293 total_tokens: 11,
294 cached_prompt_tokens: None,
295 cache_creation_tokens: Some(2),
296 cache_read_tokens: Some(3),
297 iterations: Some(vec![
298 json!({
299 "type": "compaction",
300 "input_tokens": 50,
301 "output_tokens": 5,
302 "cache_creation_input_tokens": 4,
303 "cache_read_input_tokens": 6,
304 }),
305 json!({
306 "type": "message",
307 "input_tokens": 10,
308 "output_tokens": 2,
309 "cache_creation_input_tokens": 1,
310 "cache_read_input_tokens": 2,
311 }),
312 ]),
313 };
314
315 assert_eq!(
316 usage.billable_totals(),
317 super::UsageTotals {
318 prompt_tokens: 60,
319 completion_tokens: 7,
320 total_tokens: 67,
321 cache_read_tokens: 8,
322 cache_creation_tokens: 5,
323 }
324 );
325 }
326
327 #[test]
328 fn billable_totals_fall_back_to_top_level_without_known_iterations() {
329 let usage = Usage {
330 prompt_tokens: 10,
331 completion_tokens: 1,
332 total_tokens: 11,
333 cached_prompt_tokens: Some(4),
334 cache_creation_tokens: Some(2),
335 cache_read_tokens: None,
336 iterations: Some(vec![json!({"type": "future_iteration"})]),
337 };
338
339 assert_eq!(
340 usage.billable_totals(),
341 super::UsageTotals {
342 prompt_tokens: 10,
343 completion_tokens: 1,
344 total_tokens: 11,
345 cache_read_tokens: 4,
346 cache_creation_tokens: 2,
347 }
348 );
349 }
350}
351
352#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, Default)]
353pub enum FinishReason {
354 #[default]
355 Stop,
356 Length,
357 ToolCalls,
358 ContentFilter,
359 Pause,
360 Refusal,
361 Error(String),
362}
363
364#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
366pub struct ToolCall {
367 pub id: String,
369
370 #[serde(rename = "type")]
372 pub call_type: String,
373
374 #[serde(skip_serializing_if = "Option::is_none")]
376 pub function: Option<FunctionCall>,
377
378 #[serde(skip_serializing_if = "Option::is_none")]
380 pub text: Option<String>,
381
382 #[serde(skip_serializing_if = "Option::is_none")]
384 pub thought_signature: Option<String>,
385}
386
387#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
389pub struct FunctionCall {
390 #[serde(default, skip_serializing_if = "Option::is_none")]
392 pub namespace: Option<String>,
393
394 pub name: String,
396
397 pub arguments: String,
399}
400
401impl ToolCall {
402 pub fn function(id: String, name: String, arguments: String) -> Self {
404 Self::function_with_namespace(id, None, name, arguments)
405 }
406
407 pub fn function_with_namespace(id: String, namespace: Option<String>, name: String, arguments: String) -> Self {
409 Self {
410 id,
411 call_type: "function".to_owned(),
412 function: Some(FunctionCall { namespace, name, arguments }),
413 text: None,
414 thought_signature: None,
415 }
416 }
417
418 pub fn custom(id: String, name: String, text: String) -> Self {
420 Self {
421 id,
422 call_type: "custom".to_owned(),
423 function: Some(FunctionCall { namespace: None, name, arguments: text.clone() }),
424 text: Some(text),
425 thought_signature: None,
426 }
427 }
428
429 pub fn is_custom(&self) -> bool {
431 self.call_type == "custom"
432 }
433
434 pub fn tool_name(&self) -> Option<&str> {
436 self.function.as_ref().map(|function| function.name.as_str())
437 }
438
439 pub fn raw_input(&self) -> Option<&str> {
441 self.text
442 .as_deref()
443 .or_else(|| self.function.as_ref().map(|function| function.arguments.as_str()))
444 }
445
446 pub fn parsed_arguments(&self) -> Result<serde_json::Value, serde_json::Error> {
448 if let Some(ref func) = self.function {
449 parse_tool_arguments(&func.arguments)
450 } else {
451 serde_json::from_str("")
453 }
454 }
455
456 pub fn execution_arguments(&self) -> Result<serde_json::Value, serde_json::Error> {
462 if self.is_custom() {
463 return Ok(serde_json::Value::String(self.raw_input().unwrap_or_default().to_string()));
464 }
465
466 self.parsed_arguments()
467 }
468
469 pub fn validate(&self) -> Result<(), String> {
471 if self.id.is_empty() {
472 return Err("Tool call ID cannot be empty".to_owned());
473 }
474
475 match self.call_type.as_str() {
476 "function" => {
477 if let Some(func) = &self.function {
478 if func.name.is_empty() {
479 return Err("Function name cannot be empty".to_owned());
480 }
481 if let Err(e) = self.parsed_arguments() {
483 return Err(format!("Invalid JSON in function arguments: {e}"));
484 }
485 } else {
486 return Err("Function tool call missing function details".to_owned());
487 }
488 }
489 "custom" => {
490 if let Some(func) = &self.function {
492 if func.name.is_empty() {
493 return Err("Custom tool name cannot be empty".to_owned());
494 }
495 } else {
496 return Err("Custom tool call missing function details".to_owned());
497 }
498 }
499 _ => return Err(format!("Unsupported tool call type: {}", self.call_type)),
500 }
501
502 Ok(())
503 }
504}
505
506fn parse_tool_arguments(raw_arguments: &str) -> Result<serde_json::Value, serde_json::Error> {
507 let trimmed = raw_arguments.trim();
508 match serde_json::from_str(trimmed) {
509 Ok(parsed) => Ok(parsed),
510 Err(primary_error) => {
511 if let Some(candidate) = extract_balanced_json(trimmed)
512 && let Ok(parsed) = serde_json::from_str(candidate)
513 {
514 return Ok(parsed);
515 }
516 if let Some(candidate) = repair_tag_polluted_json(trimmed)
517 && let Ok(parsed) = serde_json::from_str(&candidate)
518 {
519 return Ok(parsed);
520 }
521 if let Some(repaired) = close_incomplete_json_prefix(trimmed)
522 && let Ok(parsed) = serde_json::from_str(&repaired)
523 {
524 return Ok(parsed);
525 }
526 Err(primary_error)
527 }
528 }
529}
530
531fn extract_balanced_json(input: &str) -> Option<&str> {
532 let start = input.find(['{', '['])?;
533 let opening = input.as_bytes().get(start).copied()?;
534 let closing = match opening {
535 b'{' => b'}',
536 b'[' => b']',
537 _ => return None,
538 };
539
540 let mut depth = 0usize;
541 let mut in_string = false;
542 let mut escaped = false;
543
544 for (offset, ch) in input.get(start..)?.char_indices() {
545 if in_string {
546 if escaped {
547 escaped = false;
548 continue;
549 }
550 if ch == '\\' {
551 escaped = true;
552 continue;
553 }
554 if ch == '"' {
555 in_string = false;
556 }
557 continue;
558 }
559
560 match ch {
561 '"' => in_string = true,
562 _ if ch as u32 == opening as u32 => depth += 1,
563 _ if ch as u32 == closing as u32 => {
564 depth = depth.saturating_sub(1);
565 if depth == 0 {
566 let end = start + offset + ch.len_utf8();
567 return input.get(start..end);
568 }
569 }
570 _ => {}
571 }
572 }
573
574 None
575}
576
577fn repair_tag_polluted_json(input: &str) -> Option<String> {
578 let start = input.find(['{', '['])?;
579 let candidate = input.get(start..)?;
580 let boundary = find_provider_markup_boundary(candidate)?;
581 if boundary == 0 {
582 return None;
583 }
584
585 close_incomplete_json_prefix(candidate.get(..boundary)?.trim_end())
586}
587
588fn find_provider_markup_boundary(input: &str) -> Option<usize> {
589 const PROVIDER_MARKERS: &[&str] = &[
590 "<</",
591 "</parameter>",
592 "</invoke>",
593 "</minimax:tool_call>",
594 "<minimax:tool_call>",
595 "<parameter name=\"",
596 "<invoke name=\"",
597 "<tool_call>",
598 "</tool_call>",
599 ];
600
601 input.char_indices().find_map(|(offset, _)| {
602 let rest = input.get(offset..)?;
603 PROVIDER_MARKERS.iter().any(|marker| rest.starts_with(marker)).then_some(offset)
604 })
605}
606
607fn close_incomplete_json_prefix(prefix: &str) -> Option<String> {
608 if prefix.is_empty() {
609 return None;
610 }
611
612 let mut repaired = String::with_capacity(prefix.len() + 8);
613 let mut expected_closers = Vec::new();
614 let mut in_string = false;
615 let mut escaped = false;
616
617 for ch in prefix.chars() {
618 repaired.push(ch);
619
620 if in_string {
621 if escaped {
622 escaped = false;
623 continue;
624 }
625
626 match ch {
627 '\\' => escaped = true,
628 '"' => in_string = false,
629 _ => {}
630 }
631 continue;
632 }
633
634 match ch {
635 '"' => in_string = true,
636 '{' => expected_closers.push('}'),
637 '[' => expected_closers.push(']'),
638 '}' | ']' if expected_closers.pop() != Some(ch) => return None,
639 '}' | ']' => {}
640 _ => {}
641 }
642 }
643
644 if in_string {
645 repaired.push('"');
646 }
647 for closer in expected_closers.drain(..) {
648 repaired.push(closer);
649 }
650
651 Some(repaired)
652}
653
654#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Default)]
656pub struct LLMResponse {
657 pub content: Option<String>,
659
660 pub tool_calls: Option<Vec<ToolCall>>,
662
663 pub model: String,
665
666 pub usage: Option<Usage>,
668
669 pub finish_reason: FinishReason,
671
672 pub reasoning: Option<String>,
674
675 pub reasoning_details: Option<Vec<String>>,
677
678 pub tool_references: Vec<String>,
680
681 pub request_id: Option<String>,
683
684 pub organization_id: Option<String>,
686
687 pub compaction: Option<String>,
692}
693
694impl LLMResponse {
695 pub fn new(model: impl Into<String>, content: impl Into<String>) -> Self {
697 Self {
698 content: Some(content.into()),
699 tool_calls: None,
700 model: model.into(),
701 usage: None,
702 finish_reason: FinishReason::Stop,
703 reasoning: None,
704 reasoning_details: None,
705 tool_references: Vec::new(),
706 request_id: None,
707 organization_id: None,
708 compaction: None,
709 }
710 }
711
712 pub fn content_text(&self) -> &str {
714 self.content.as_deref().unwrap_or("")
715 }
716
717 pub fn content_string(&self) -> String {
719 self.content.clone().unwrap_or_default()
720 }
721}
722
723#[derive(Clone, Deserialize, PartialEq, Eq)]
724pub struct LLMErrorMetadata {
725 provider: Option<String>,
726 pub status: Option<u16>,
727 pub code: Option<String>,
728 request_id: Option<String>,
729 organization_id: Option<String>,
730 pub retry_after: Option<String>,
731 pub message: Option<String>,
732}
733
734impl fmt::Debug for LLMErrorMetadata {
735 #[inline(never)]
739 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
740 formatter
741 .debug_struct("LLMErrorMetadata")
742 .field("provider", &self.provider)
743 .field("status", &self.status)
744 .field("code", &self.code)
745 .field("request_id", &self.request_id)
746 .field("organization_id", &self.organization_id)
747 .field("retry_after", &self.retry_after)
748 .field(
749 "message",
750 &self
751 .message
752 .as_deref()
753 .map(|message| sanitize_provider_diagnostic(message.as_bytes())),
754 )
755 .finish()
756 }
757}
758
759impl Serialize for LLMErrorMetadata {
760 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
761 where
762 S: serde::Serializer,
763 {
764 let mut state = serializer.serialize_struct("LLMErrorMetadata", 7)?;
765 state.serialize_field("provider", &self.provider)?;
766 state.serialize_field("status", &self.status)?;
767 state.serialize_field("code", &self.code)?;
768 state.serialize_field("request_id", &self.request_id)?;
769 state.serialize_field("organization_id", &self.organization_id)?;
770 state.serialize_field("retry_after", &self.retry_after)?;
771 let message = self
772 .message
773 .as_deref()
774 .map(|message| sanitize_provider_diagnostic(message.as_bytes()));
775 state.serialize_field("message", &message)?;
776 state.end()
777 }
778}
779
780impl LLMErrorMetadata {
781 #[must_use]
784 pub fn new(
785 provider: impl Into<String>,
786 status: Option<u16>,
787 code: Option<String>,
788 request_id: Option<String>,
789 organization_id: Option<String>,
790 retry_after: Option<String>,
791 message: Option<String>,
792 ) -> Box<Self> {
793 Box::new(Self {
794 provider: Some(provider.into()),
795 status,
796 code,
797 request_id,
798 organization_id,
799 retry_after,
800 message: message.map(|message| sanitize_provider_diagnostic(message.as_bytes())),
801 })
802 }
803}
804
805#[derive(Deserialize, Clone)]
807#[serde(tag = "type", rename_all = "snake_case")]
808pub enum LLMError {
809 Authentication {
810 message: String,
811 metadata: Option<Box<LLMErrorMetadata>>,
812 },
813 RateLimit {
814 metadata: Option<Box<LLMErrorMetadata>>,
815 },
816 InvalidRequest {
817 message: String,
818 metadata: Option<Box<LLMErrorMetadata>>,
819 },
820 Network {
821 message: String,
822 metadata: Option<Box<LLMErrorMetadata>>,
823 },
824 Provider {
825 message: String,
826 metadata: Option<Box<LLMErrorMetadata>>,
827 },
828}
829
830impl fmt::Debug for LLMError {
831 #[inline(never)]
834 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
835 match self {
836 Self::Authentication { message, metadata } => formatter
837 .debug_struct("Authentication")
838 .field("message", &sanitize_provider_diagnostic(message.as_bytes()))
839 .field("metadata", metadata)
840 .finish(),
841 Self::RateLimit { metadata } => formatter.debug_struct("RateLimit").field("metadata", metadata).finish(),
842 Self::InvalidRequest { message, metadata } => formatter
843 .debug_struct("InvalidRequest")
844 .field("message", &sanitize_provider_diagnostic(message.as_bytes()))
845 .field("metadata", metadata)
846 .finish(),
847 Self::Network { message, metadata } => formatter
848 .debug_struct("Network")
849 .field("message", &sanitize_provider_diagnostic(message.as_bytes()))
850 .field("metadata", metadata)
851 .finish(),
852 Self::Provider { message, metadata } => formatter
853 .debug_struct("Provider")
854 .field("message", &sanitize_provider_diagnostic(message.as_bytes()))
855 .field("metadata", metadata)
856 .finish(),
857 }
858 }
859}
860
861impl fmt::Display for LLMError {
862 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
863 match self {
864 Self::Authentication { message, .. } => {
865 write!(formatter, "Authentication failed: {}", sanitize_provider_diagnostic(message.as_bytes()))
866 }
867 Self::RateLimit { .. } => formatter.write_str("Rate limit exceeded"),
868 Self::InvalidRequest { message, .. } => {
869 write!(formatter, "Invalid request: {}", sanitize_provider_diagnostic(message.as_bytes()))
870 }
871 Self::Network { message, .. } => {
872 write!(formatter, "Network error: {}", sanitize_provider_diagnostic(message.as_bytes()))
873 }
874 Self::Provider { message, .. } => {
875 write!(formatter, "Provider error: {}", sanitize_provider_diagnostic(message.as_bytes()))
876 }
877 }
878 }
879}
880
881impl std::error::Error for LLMError {}
882
883impl Serialize for LLMError {
884 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
885 where
886 S: serde::Serializer,
887 {
888 match self {
889 Self::Authentication { message, metadata } => {
890 let mut state = serializer.serialize_struct("LLMError", 3)?;
891 state.serialize_field("type", "authentication")?;
892 state.serialize_field("message", &sanitize_provider_diagnostic(message.as_bytes()))?;
893 state.serialize_field("metadata", metadata)?;
894 state.end()
895 }
896 Self::RateLimit { metadata } => {
897 let mut state = serializer.serialize_struct("LLMError", 2)?;
898 state.serialize_field("type", "rate_limit")?;
899 state.serialize_field("metadata", metadata)?;
900 state.end()
901 }
902 Self::InvalidRequest { message, metadata } => {
903 let mut state = serializer.serialize_struct("LLMError", 3)?;
904 state.serialize_field("type", "invalid_request")?;
905 state.serialize_field("message", &sanitize_provider_diagnostic(message.as_bytes()))?;
906 state.serialize_field("metadata", metadata)?;
907 state.end()
908 }
909 Self::Network { message, metadata } => {
910 let mut state = serializer.serialize_struct("LLMError", 3)?;
911 state.serialize_field("type", "network")?;
912 state.serialize_field("message", &sanitize_provider_diagnostic(message.as_bytes()))?;
913 state.serialize_field("metadata", metadata)?;
914 state.end()
915 }
916 Self::Provider { message, metadata } => {
917 let mut state = serializer.serialize_struct("LLMError", 3)?;
918 state.serialize_field("type", "provider")?;
919 state.serialize_field("message", &sanitize_provider_diagnostic(message.as_bytes()))?;
920 state.serialize_field("metadata", metadata)?;
921 state.end()
922 }
923 }
924 }
925}
926
927#[cfg(test)]
928mod tests {
929 use super::{LLMError, LLMErrorMetadata, ToolCall};
930 use serde_json::json;
931
932 #[test]
933 fn parsed_arguments_accepts_trailing_characters() {
934 let call = ToolCall::function(
935 "call_read".to_string(),
936 "exec_command".to_string(),
937 r#"{"path":"src/main.rs"} trailing text"#.to_string(),
938 );
939
940 let parsed = call.parsed_arguments().expect("arguments with trailing text should recover");
941 assert_eq!(parsed, json!({"path":"src/main.rs"}));
942 }
943
944 #[test]
945 fn parsed_arguments_accepts_code_fenced_json() {
946 let call = ToolCall::function(
947 "call_read".to_string(),
948 "exec_command".to_string(),
949 "```json\n{\"path\":\"src/lib.rs\",\"limit\":25}\n```".to_string(),
950 );
951
952 let parsed = call.parsed_arguments().expect("code-fenced arguments should recover");
953 assert_eq!(parsed, json!({"path":"src/lib.rs","limit":25}));
954 }
955
956 #[test]
957 fn parsed_arguments_recovers_truncated_json_missing_closing_brace() {
958 let call = ToolCall::function(
959 "call_search".to_string(),
960 "code_search".to_string(),
961 r#"{"query":"context","path":".","file_types":["rust"],"result_types":["definition"],"max_results":20"#
962 .to_string(),
963 );
964
965 let parsed = call
966 .parsed_arguments()
967 .expect("truncated JSON missing closing brace should recover");
968 assert_eq!(
969 parsed,
970 json!({
971 "query": "context",
972 "path": ".",
973 "file_types": ["rust"],
974 "result_types": ["definition"],
975 "max_results": 20
976 })
977 );
978 }
979
980 #[test]
981 fn parsed_arguments_rejects_incomplete_json() {
982 let call = ToolCall::function(
983 "call_read".to_string(),
984 "exec_command".to_string(),
985 r#"{"path":"src/main.rs","limit""#.to_string(),
986 );
987
988 assert!(call.parsed_arguments().is_err());
989 }
990
991 #[test]
992 fn llm_error_debug_and_json_redact_provider_secrets() {
993 let secret = concat!("sk-", "test1234567890abcdefghij");
994 let error = LLMError::Provider {
995 message: format!("response body api_key={secret} bearer Bearer abcdefghijklmnop"),
996 metadata: Some(LLMErrorMetadata::new(
997 "OpenAI",
998 Some(401),
999 Some("invalid_api_key".to_owned()),
1000 Some("req-123".to_owned()),
1001 None,
1002 None,
1003 Some(format!("{}={}", "AWS_SECRET_ACCESS_KEY", "cloud-secret-value")),
1004 )),
1005 };
1006
1007 let debug = format!("{error:?}");
1008 let json = serde_json::to_string(&error).expect("LLM errors should serialize");
1009
1010 assert!(!debug.contains(secret));
1011 assert!(!debug.contains("cloud-secret-value"));
1012 assert!(!json.contains(secret));
1013 assert!(!json.contains("cloud-secret-value"));
1014 assert!(json.contains("req-123"));
1015 assert!(json.contains("401"));
1016 }
1017
1018 #[test]
1019 fn parsed_arguments_recovers_truncated_minimax_markup() {
1020 let call = ToolCall::function(
1021 "call_search".to_string(),
1022 "code_search".to_string(),
1023 "{\"query\":\"persistent_memory\",\"file_types\":[\"rust\"],\"result_types\":[\"text\"],\"max_results\":20,\"path\":\"crates/codegen/vtcode-core/src</parameter>\n<</invoke>\n</minimax:tool_call>".to_string(),
1024 );
1025
1026 let parsed = call.parsed_arguments().expect("minimax markup spillover should recover");
1027 assert_eq!(
1028 parsed,
1029 json!({
1030 "query": "persistent_memory",
1031 "path": "crates/codegen/vtcode-core/src",
1032 "file_types": ["rust"],
1033 "result_types": ["text"],
1034 "max_results": 20
1035 })
1036 );
1037 }
1038
1039 #[test]
1040 fn function_call_serializes_optional_namespace() {
1041 let call = ToolCall::function_with_namespace(
1042 "call_read".to_string(),
1043 Some("workspace".to_string()),
1044 "exec_command".to_string(),
1045 r#"{"path":"src/main.rs"}"#.to_string(),
1046 );
1047
1048 let json = serde_json::to_value(&call).expect("tool call should serialize");
1049 assert_eq!(json["function"]["namespace"], "workspace");
1050 assert_eq!(json["function"]["name"], "exec_command");
1051 }
1052
1053 #[test]
1054 fn custom_tool_call_exposes_raw_execution_arguments() {
1055 let patch = "*** Begin Patch\n*** End Patch\n".to_string();
1056 let call = ToolCall::custom("call_patch".to_string(), "apply_patch".to_string(), patch.clone());
1057
1058 assert!(call.is_custom());
1059 assert_eq!(call.tool_name(), Some("apply_patch"));
1060 assert_eq!(call.raw_input(), Some(patch.as_str()));
1061 assert_eq!(call.execution_arguments().expect("custom arguments"), json!(patch));
1062 assert!(call.parsed_arguments().is_err(), "custom tool payload should stay freeform rather than JSON");
1063 }
1064}