Skip to main content

schwab_cli/agent/
resilience.rs

1//! Classify agent tick failures so the loop can backoff instead of exiting.
2
3/// How the agent should react to a tick (or startup) error.
4#[derive(Debug, Clone, Copy, PartialEq, Eq)]
5pub enum AgentErrorClass {
6    /// Network blip, 5xx, timeouts — short backoff and retry.
7    Recoverable,
8    /// Refresh token revoked / missing — stay alive, long backoff, wait for re-login.
9    AuthFatal,
10    /// Unexpected errors — backoff and keep trying (unattended bot must not die).
11    Unexpected,
12}
13
14const MIN_BACKOFF_SECS: u64 = 5;
15const MAX_BACKOFF_SECS: u64 = 300;
16const AUTH_FATAL_BACKOFF_SECS: u64 = 60;
17
18/// Classify an error from its Display / Debug chain text.
19pub 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    // Transient HTTP / transport
38    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    // Access-token 401 often clears after refresh; treat as recoverable unless
61    // paired with invalid_grant (already handled above).
62    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
70/// Exponential backoff from consecutive failure count (1-based).
71pub 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}