Skip to main content

theway_daemon/
agent_session.rs

1//! AgentSession — auto-retry wrapper around [`theway_core::AgentHarness`]. 1:1 port of
2//! `packages/coding-agent/src/core/agent-session.ts` retry logic (`_isRetryableError` +
3//! `_prepareRetry`).
4//!
5//! On a retryable LLM error the wrapper:
6//! 1. Drops the failed assistant message from agent state and rewinds the active session leaf
7//!    while keeping the append-only session log intact.
8//! 2. Waits exponentially (`base_delay_ms * 2^(attempt-1)`, capped).
9//! 3. Calls `harness.continue_()` to re-run.
10//! 4. Up to `max_retries` times. After that, the error surfaces to the caller.
11//!
12//! Retryable patterns: `overloaded`, `rate.?limit`, `429`, `5xx`, network/connection errors,
13//! `ended without`, `stream ended before message_stop`, `terminated`, `retry delay`. Mirrors
14//! the TS regex.
15
16use std::sync::Arc;
17
18use once_cell::sync::Lazy;
19use regex::Regex;
20use theway_core::{AgentHarness, AgentMessage, AgentRunError, LoopListener, SessionTreeEntry};
21use theway_llm_provider::{AssistantMessage as PiAssistantMessage, Message as PiMessage};
22
23#[derive(Clone, Debug)]
24pub struct RetrySettings {
25    pub enabled: bool,
26    pub max_retries: u32,
27    pub base_delay_ms: u64,
28    pub max_delay_ms: u64,
29    /// When set and the primary model's retries exhaust on a retryable error, swap to this
30    /// `(provider, model_id)` and retry once. Returns whatever the fallback says (success or
31    /// definitive error). Each model is tried at most once per fallback chain — no looping.
32    pub fallback_model: Option<(String, String)>,
33}
34
35impl Default for RetrySettings {
36    fn default() -> Self {
37        Self {
38            enabled: true,
39            max_retries: 5,
40            base_delay_ms: 1_000,
41            max_delay_ms: 60_000,
42            fallback_model: None,
43        }
44    }
45}
46
47/// Regex matching error messages we should retry. Compiled once.
48static RETRYABLE_RE: Lazy<Regex> = Lazy::new(|| {
49    Regex::new(
50        r"(?i)overloaded|provider.?returned.?error|rate.?limit|too many requests|429|500|502|503|504|service.?unavailable|server.?error|internal.?error|network.?error|connection.?error|connection.?refused|connection.?lost|websocket.?closed|websocket.?error|other side closed|fetch failed|upstream.?connect|reset before headers|socket hang up|ended without|stream ended before message_stop|http2 request did not get a response|timed? out|timeout|terminated|retry delay",
51    )
52    .expect("retry regex")
53});
54
55pub fn is_retryable_error(err_message: &str) -> bool {
56    RETRYABLE_RE.is_match(err_message)
57}
58
59/// Lightweight wrapper. Holds the harness + retry settings; not a deep clone of TS
60/// `AgentSession` (no extension orchestration, no event-bus fan-out — just retry).
61pub struct AgentSession {
62    harness: Arc<AgentHarness>,
63    settings: RetrySettings,
64}
65
66impl AgentSession {
67    pub fn new(harness: Arc<AgentHarness>, settings: RetrySettings) -> Self {
68        Self { harness, settings }
69    }
70
71    #[allow(dead_code)] // public API for embedders; not used by the binary itself.
72    pub fn harness(&self) -> &AgentHarness {
73        &self.harness
74    }
75
76    /// Subscribe to underlying lifecycle events. Useful for UI listeners.
77    #[allow(dead_code)] // public API for embedders; not used by the binary itself.
78    pub fn subscribe(&self, listener: LoopListener) -> impl FnOnce() {
79        self.harness.subscribe(listener)
80    }
81
82    /// Prompt with retry. Same signature as `AgentHarness::prompt` but loops on retryable LLM
83    /// errors with exponential backoff. If a `fallback_model` is configured and the primary
84    /// run exhausts retries, the model is swapped exactly once and the loop restarts.
85    pub async fn prompt(&self, text: impl Into<String>) -> Result<(), AgentRunError> {
86        let text = text.into();
87        let mut attempt: u32 = 0;
88        let mut fallback_used = false;
89        loop {
90            let r = if attempt == 0 {
91                self.harness.prompt(text.clone()).await
92            } else {
93                self.harness.continue_().await
94            };
95            let err = match r {
96                Ok(()) => match self.assistant_error_message(&self.last_assistant()) {
97                    Some(error_message) => AgentRunError::Other(error_message),
98                    None => return Ok(()),
99                },
100                Err(e) => e,
101            };
102            // Successful prompt() can still leave a synthesized error assistant message
103            // (provider stream encoded the error). Re-evaluate via retry policy.
104
105            if !self.settings.enabled {
106                return Err(err);
107            }
108
109            if !is_retryable_error(&err.to_string()) {
110                return Err(err);
111            }
112
113            if attempt >= self.settings.max_retries {
114                // Exhausted retries on the current model. If a fallback is configured and we
115                // haven't already used it, swap and restart from attempt=0.
116                if let Some((provider, model_id)) = &self.settings.fallback_model {
117                    if !fallback_used {
118                        fallback_used = true;
119                        if let Some(m) = theway_llm_provider::get_model(
120                            &theway_llm_provider::Provider::from(provider.as_str()),
121                            model_id,
122                        ) {
123                            self.rewind_failed_assistant().await?;
124                            if let Err(e) = self.harness.set_model(m).await {
125                                return Err(AgentRunError::Other(format!(
126                                    "fallback set_model failed: {e}"
127                                )));
128                            }
129                            attempt = 0;
130                            continue;
131                        }
132                    }
133                }
134                return Err(err);
135            }
136
137            attempt += 1;
138            let delay_ms = backoff_ms(
139                attempt,
140                self.settings.base_delay_ms,
141                self.settings.max_delay_ms,
142            );
143            tokio::time::sleep(std::time::Duration::from_millis(delay_ms)).await;
144
145            // Drop the failed assistant message from agent state so continue_() doesn't replay
146            // a context that ends in an error.
147            self.rewind_failed_assistant().await?;
148        }
149    }
150
151    fn last_assistant(&self) -> Option<PiAssistantMessage> {
152        let s = self.harness.agent().state();
153        for m in s.messages.iter().rev() {
154            if let AgentMessage::Llm(PiMessage::Assistant(a)) = m {
155                return Some(a.clone());
156            }
157        }
158        None
159    }
160
161    fn assistant_error_message(&self, a: &Option<PiAssistantMessage>) -> Option<String> {
162        let Some(a) = a else { return None };
163        if a.stop_reason != theway_llm_provider::StopReason::Error {
164            return None;
165        }
166        a.error_message
167            .clone()
168            .or_else(|| Some("assistant stopped with an error".to_string()))
169    }
170
171    async fn rewind_failed_assistant(&self) -> Result<(), AgentRunError> {
172        let mut s = self.harness.agent().state();
173        while let Some(last) = s.messages.last() {
174            if matches!(last, AgentMessage::Llm(PiMessage::Assistant(a)) if a.stop_reason == theway_llm_provider::StopReason::Error)
175            {
176                s.messages.pop();
177            } else {
178                break;
179            }
180        }
181        drop(s);
182
183        let session = self.harness.session();
184        let Some(leaf_id) = session
185            .leaf_id()
186            .await
187            .map_err(|e| AgentRunError::Other(format!("session retry leaf lookup: {e}")))?
188        else {
189            return Ok(());
190        };
191        let Some(SessionTreeEntry::Message {
192            parent_id,
193            message: AgentMessage::Llm(PiMessage::Assistant(a)),
194            ..
195        }) = session
196            .get_entry(&leaf_id)
197            .await
198            .map_err(|e| AgentRunError::Other(format!("session retry leaf entry lookup: {e}")))?
199        else {
200            return Ok(());
201        };
202        if a.stop_reason == theway_llm_provider::StopReason::Error {
203            session
204                .move_to(parent_id.as_deref(), None)
205                .await
206                .map_err(|e| AgentRunError::Other(format!("session retry rewind: {e}")))?;
207        }
208        Ok(())
209    }
210}
211
212fn backoff_ms(attempt: u32, base: u64, max: u64) -> u64 {
213    let exponent = attempt.saturating_sub(1).min(10);
214    let n = (base as u128) << exponent;
215    n.min(max as u128) as u64
216}
217
218#[cfg(test)]
219mod tests {
220    use super::*;
221    use theway_core::{AgentHarnessOptions, MemorySessionStorage, Session, SessionStorage};
222    use theway_llm_provider::{
223        AssistantMessage, AssistantMessageEvent, AssistantMessageEventStream, AssistantRole,
224        ContentBlock, DoneReason, ModelCost, StopReason, Usage,
225    };
226
227    fn faux_model() -> theway_llm_provider::Model {
228        theway_llm_provider::Model {
229            id: "faux".into(),
230            name: "Faux".into(),
231            api: theway_llm_provider::Api::from("faux"),
232            provider: theway_llm_provider::Provider::from("faux"),
233            base_url: String::new(),
234            reasoning: false,
235            thinking_level_map: None,
236            input: vec![],
237            cost: ModelCost::default(),
238            context_window: 0,
239            max_tokens: 0,
240            headers: None,
241            compat: None,
242        }
243    }
244
245    fn assistant(
246        text: &str,
247        stop_reason: StopReason,
248        error_message: Option<&str>,
249    ) -> AssistantMessage {
250        AssistantMessage {
251            role: AssistantRole::Assistant,
252            content: vec![ContentBlock::text(text)],
253            api: theway_llm_provider::Api::from("faux"),
254            provider: theway_llm_provider::Provider::from("faux"),
255            model: "faux".into(),
256            response_model: None,
257            response_id: None,
258            diagnostics: None,
259            usage: Usage::default(),
260            stop_reason,
261            error_message: error_message.map(str::to_string),
262            timestamp: 0,
263        }
264    }
265
266    fn stream_fn_with(
267        responses: Arc<tokio::sync::Mutex<Vec<AssistantMessage>>>,
268    ) -> theway_core::StreamFn {
269        Arc::new(move |_, _, _| {
270            let (stream, mut sender) = AssistantMessageEventStream::new();
271            let responses = responses.clone();
272            tokio::spawn(async move {
273                let msg = responses.lock().await.remove(0);
274                sender.push(AssistantMessageEvent::Start {
275                    partial: msg.clone(),
276                });
277                let reason = match msg.stop_reason {
278                    StopReason::ToolUse => DoneReason::ToolUse,
279                    StopReason::Length => DoneReason::Length,
280                    _ => DoneReason::Stop,
281                };
282                sender.push(AssistantMessageEvent::Done {
283                    reason,
284                    message: msg,
285                });
286            });
287            stream
288        })
289    }
290
291    #[test]
292    fn retryable_patterns_match_ts_regex() {
293        assert!(is_retryable_error("overloaded_error"));
294        assert!(is_retryable_error(
295            "Provider returned error: 429 Too Many Requests"
296        ));
297        assert!(is_retryable_error("rate limit exceeded"));
298        assert!(is_retryable_error("HTTP 503 Service Unavailable"));
299        assert!(is_retryable_error("websocket closed"));
300        assert!(is_retryable_error("stream ended before message_stop"));
301        assert!(is_retryable_error("socket hang up"));
302        assert!(is_retryable_error("reset before headers"));
303        assert!(!is_retryable_error("bad request: missing field"));
304        assert!(!is_retryable_error("Unauthorized"));
305        assert!(!is_retryable_error("model not found"));
306    }
307
308    #[test]
309    fn backoff_grows_and_caps() {
310        assert_eq!(backoff_ms(1, 1000, 60_000), 1000);
311        assert_eq!(backoff_ms(2, 1000, 60_000), 2000);
312        assert_eq!(backoff_ms(9, 1000, 60_000), 60_000);
313    }
314
315    #[tokio::test]
316    async fn retry_rewinds_failed_assistant_out_of_active_session_branch() {
317        let storage = Arc::new(MemorySessionStorage::new());
318        let session = Session::new(storage as Arc<dyn SessionStorage>);
319        let responses = Arc::new(tokio::sync::Mutex::new(vec![
320            assistant("temporary failure", StopReason::Error, Some("HTTP 503")),
321            assistant("ok", StopReason::Stop, None),
322        ]));
323
324        let mut opts = AgentHarnessOptions::new(Some(faux_model()), session.clone());
325        opts.stream_fn = Some(stream_fn_with(responses));
326        let harness = Arc::new(AgentHarness::new(opts));
327        let runner = AgentSession::new(
328            harness,
329            RetrySettings {
330                base_delay_ms: 0,
331                max_delay_ms: 0,
332                max_retries: 1,
333                ..RetrySettings::default()
334            },
335        );
336
337        runner.prompt("hi").await.unwrap();
338
339        let entries = session.entries().await.unwrap();
340        assert!(entries.iter().any(|e| matches!(
341            e,
342            SessionTreeEntry::Message {
343                message: AgentMessage::Llm(PiMessage::Assistant(a)),
344                ..
345            } if a.stop_reason == StopReason::Error
346        )));
347
348        let active = session.build_context().await.unwrap();
349        assert!(!active.messages.iter().any(|m| matches!(
350            m,
351            AgentMessage::Llm(PiMessage::Assistant(a)) if a.stop_reason == StopReason::Error
352        )));
353        assert!(active.messages.iter().any(|m| matches!(
354            m,
355            AgentMessage::Llm(PiMessage::Assistant(a))
356                if a.stop_reason == StopReason::Stop
357        )));
358    }
359}
360
361#[cfg(test)]
362mod agent_session_mirrored_tests {
363    //! Mirrored tests live in `tests/agent_session/`; wrapped because the
364    //! top-level `mod tests` slot is already used by the inline tests.
365    tests_bridge_macro::tests_bridge!("agent_session");
366}