schwab_cli/agent/
resilience.rs1#[derive(Debug, Clone, Copy, PartialEq, Eq)]
5pub enum AgentErrorClass {
6 Recoverable,
8 AuthFatal,
10 Unexpected,
12}
13
14const MIN_BACKOFF_SECS: u64 = 5;
15const MAX_BACKOFF_SECS: u64 = 300;
16const AUTH_FATAL_BACKOFF_SECS: u64 = 60;
17
18pub fn classify_agent_error(err: &anyhow::Error) -> AgentErrorClass {
20 classify_error_message(&format!("{err:#}"))
21}
22
23pub fn classify_error_message(msg: &str) -> AgentErrorClass {
24 let lower = msg.to_ascii_lowercase();
25
26 if lower.contains("invalid_grant")
27 || lower.contains("refresh token is invalid")
28 || lower.contains("refresh token") && (lower.contains("expired") || lower.contains("revoked"))
29 || lower.contains("not authenticated")
30 || lower.contains("no refresh token")
31 || lower.contains("run: schwab auth login")
32 || lower.contains("schwab auth login")
33 {
34 return AgentErrorClass::AuthFatal;
35 }
36
37 if lower.contains("timeout")
39 || lower.contains("timed out")
40 || lower.contains("connection reset")
41 || lower.contains("connection refused")
42 || lower.contains("temporarily unavailable")
43 || lower.contains("dns")
44 || lower.contains("error sending request")
45 || lower.contains("hyper::error")
46 || lower.contains("api error 429")
47 || lower.contains("api error 500")
48 || lower.contains("api error 502")
49 || lower.contains("api error 503")
50 || lower.contains("api error 504")
51 || lower.contains("status: 429")
52 || lower.contains("status: 500")
53 || lower.contains("status: 502")
54 || lower.contains("status: 503")
55 || lower.contains("status: 504")
56 {
57 return AgentErrorClass::Recoverable;
58 }
59
60 if lower.contains("api error 401") || lower.contains("status: 401") || lower.contains("unauthorized")
63 {
64 return AgentErrorClass::Recoverable;
65 }
66
67 AgentErrorClass::Unexpected
68}
69
70pub fn backoff_seconds(class: AgentErrorClass, consecutive_failures: u32) -> u64 {
72 match class {
73 AgentErrorClass::AuthFatal => AUTH_FATAL_BACKOFF_SECS,
74 AgentErrorClass::Recoverable | AgentErrorClass::Unexpected => {
75 let n = consecutive_failures.max(1).saturating_sub(1);
76 let exp = MIN_BACKOFF_SECS.saturating_mul(2u64.saturating_pow(n.min(6)));
77 exp.min(MAX_BACKOFF_SECS)
78 }
79 }
80}
81
82pub fn class_label(class: AgentErrorClass) -> &'static str {
83 match class {
84 AgentErrorClass::Recoverable => "recoverable",
85 AgentErrorClass::AuthFatal => "auth_fatal",
86 AgentErrorClass::Unexpected => "unexpected",
87 }
88}
89
90#[cfg(test)]
91mod tests {
92 use super::*;
93
94 #[test]
95 fn classifies_invalid_grant_as_auth_fatal() {
96 let msg = r#"OAuth error: HTTP 400: 400 Bad Request: "{"error_description":"Refresh token is invalid, expired or revoked","error":"invalid_grant"}""#;
97 assert_eq!(classify_error_message(msg), AgentErrorClass::AuthFatal);
98 }
99
100 #[test]
101 fn classifies_not_authenticated_as_auth_fatal() {
102 assert_eq!(
103 classify_error_message("Not authenticated: No refresh token on disk"),
104 AgentErrorClass::AuthFatal
105 );
106 }
107
108 #[test]
109 fn classifies_401_as_recoverable() {
110 assert_eq!(
111 classify_error_message("API error 401: "),
112 AgentErrorClass::Recoverable
113 );
114 }
115
116 #[test]
117 fn classifies_5xx_and_timeout_as_recoverable() {
118 assert_eq!(
119 classify_error_message("API error 503: service unavailable"),
120 AgentErrorClass::Recoverable
121 );
122 assert_eq!(
123 classify_error_message("HTTP request failed: timeout"),
124 AgentErrorClass::Recoverable
125 );
126 }
127
128 #[test]
129 fn classifies_unknown_as_unexpected() {
130 assert_eq!(
131 classify_error_message("something weird broke"),
132 AgentErrorClass::Unexpected
133 );
134 }
135
136 #[test]
137 fn backoff_grows_and_caps() {
138 assert_eq!(backoff_seconds(AgentErrorClass::Recoverable, 1), 5);
139 assert_eq!(backoff_seconds(AgentErrorClass::Recoverable, 2), 10);
140 assert_eq!(backoff_seconds(AgentErrorClass::Recoverable, 3), 20);
141 assert_eq!(backoff_seconds(AgentErrorClass::Recoverable, 10), 300);
142 assert_eq!(backoff_seconds(AgentErrorClass::AuthFatal, 1), 60);
143 assert_eq!(backoff_seconds(AgentErrorClass::AuthFatal, 99), 60);
144 }
145}