1use 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 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
47static 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
59pub 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)] pub fn harness(&self) -> &AgentHarness {
73 &self.harness
74 }
75
76 #[allow(dead_code)] pub fn subscribe(&self, listener: LoopListener) -> impl FnOnce() {
79 self.harness.subscribe(listener)
80 }
81
82 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 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 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 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 tests_bridge_macro::tests_bridge!("agent_session");
366}