1use std::path::PathBuf;
29use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, Ordering};
30use std::sync::{Arc, Weak};
31use std::time::Duration;
32
33use async_trait::async_trait;
34use oxi_ai::Message;
35use parking_lot::Mutex;
36use tokio::sync::oneshot;
37
38use crate::advisor::types::AdvisorNote;
39
40#[async_trait]
43pub trait AdvisorAgent: Send + Sync + 'static {
44 async fn prompt(&self, input: String) -> Result<(), String>;
47 fn abort(&self, reason: &str);
49 fn reset(&self);
51 async fn rollback_to(&self, count: usize);
54 fn message_count(&self) -> usize;
56}
57
58pub trait AdvisorRuntimeHost: Send + Sync + 'static {
60 fn snapshot_messages(&self) -> Vec<Message>;
63 fn enqueue_advice(&self, note: AdvisorNote);
66 fn maintain_context(&self, _incoming_tokens: usize) -> bool {
70 false
71 }
72 fn begin_advisor_update(&self) {}
76 fn notify_failure(&self, _error: &str) {}
79}
80
81struct PendingDelta {
83 text: String,
84 turns: u64,
86}
87
88#[derive(Default)]
92struct DrainState {
93 pending: Vec<PendingDelta>,
94 draining: bool,
95}
96
97struct CatchupWaiter {
99 threshold: u64,
100 tx: Option<oneshot::Sender<()>>,
101}
102
103pub struct AdvisorRuntime {
107 agent: Arc<dyn AdvisorAgent>,
108 host: Arc<dyn AdvisorRuntimeHost>,
109 transcript_path: Mutex<Option<PathBuf>>,
110
111 state: Mutex<DrainState>,
112 epoch: AtomicU64,
116 backlog: AtomicU64,
118 last_count: AtomicU64,
120 latest: Mutex<Option<Vec<Message>>>,
122
123 waiters: Mutex<Vec<CatchupWaiter>>,
124
125 consecutive_failures: AtomicU32,
126 failure_notified: AtomicBool,
127 disposed: AtomicBool,
128 retry_delay: Duration,
129
130 self_ref: Mutex<Option<Weak<AdvisorRuntime>>>,
132}
133
134impl AdvisorRuntime {
135 #[must_use]
138 pub fn new(
139 agent: Arc<dyn AdvisorAgent>,
140 host: Arc<dyn AdvisorRuntimeHost>,
141 retry_delay: Duration,
142 ) -> Self {
143 Self {
144 agent,
145 host,
146 transcript_path: Mutex::new(None),
147 state: Mutex::new(DrainState::default()),
148 epoch: AtomicU64::new(0),
149 backlog: AtomicU64::new(0),
150 last_count: AtomicU64::new(0),
151 latest: Mutex::new(None),
152 waiters: Mutex::new(Vec::new()),
153 consecutive_failures: AtomicU32::new(0),
154 failure_notified: AtomicBool::new(false),
155 disposed: AtomicBool::new(false),
156 retry_delay,
157 self_ref: Mutex::new(None),
158 }
159 }
160
161 pub fn set_transcript_path(&self, path: Option<PathBuf>) {
163 *self.transcript_path.lock() = path;
164 }
165
166 #[must_use]
168 pub fn transcript_path(&self) -> Option<PathBuf> {
169 self.transcript_path.lock().clone()
170 }
171
172 pub fn install_self(&self, weak: Weak<AdvisorRuntime>) {
175 *self.self_ref.lock() = Some(weak);
176 }
177
178 #[must_use]
180 pub fn backlog(&self) -> u64 {
181 self.backlog.load(Ordering::SeqCst)
182 }
183
184 #[must_use]
186 pub fn is_disposed(&self) -> bool {
187 self.disposed.load(Ordering::SeqCst)
188 }
189
190 pub fn on_turn_end(&self, messages: Vec<Message>) {
194 if self.disposed.load(Ordering::SeqCst) {
195 return;
196 }
197 *self.latest.lock() = Some(messages.clone());
198 let Some(render) = self.render_delta(&messages) else {
199 return;
200 };
201 let spawn = {
202 let mut s = self.state.lock();
203 s.pending.push(PendingDelta {
204 text: render,
205 turns: 1,
206 });
207 self.backlog.fetch_add(1, Ordering::SeqCst);
208 !s.draining
209 };
210 self.notify_waiters();
211 let drain_handle = self.self_ref.lock().as_ref().and_then(Weak::upgrade);
212 if spawn && let Some(this) = drain_handle {
213 tokio::spawn(async move {
214 this.drain().await;
215 });
216 }
217 }
218
219 pub async fn wait_for_catchup(&self, max: Duration, threshold: u64) {
224 if self.disposed.load(Ordering::SeqCst) || self.backlog.load(Ordering::SeqCst) < threshold {
225 return;
226 }
227 let (tx, rx) = oneshot::channel();
228 {
229 let mut waiters = self.waiters.lock();
230 if self.backlog.load(Ordering::SeqCst) < threshold {
233 return;
234 }
235 waiters.push(CatchupWaiter {
236 threshold,
237 tx: Some(tx),
238 });
239 }
240 let _ = tokio::time::timeout(max, rx).await;
241 }
242
243 pub fn reset(&self) {
247 self.epoch.fetch_add(1, Ordering::SeqCst);
248 self.reset_advisor_context(true);
249 self.wake_all_waiters();
250 }
251
252 pub fn seed_to(&self, count: u64) {
256 self.epoch.fetch_add(1, Ordering::SeqCst);
257 self.last_count.store(count, Ordering::SeqCst);
258 let mut s = self.state.lock();
259 s.pending.clear();
260 self.backlog.store(0, Ordering::SeqCst);
267 self.consecutive_failures.store(0, Ordering::SeqCst);
268 self.failure_notified.store(false, Ordering::SeqCst);
269 drop(s);
270 self.wake_all_waiters();
271 }
272
273 pub fn dispose(&self) {
276 self.disposed.store(true, Ordering::SeqCst);
277 self.epoch.fetch_add(1, Ordering::SeqCst);
278 let mut s = self.state.lock();
279 s.pending.clear();
280 s.draining = false;
281 self.backlog.store(0, Ordering::SeqCst);
282 drop(s);
283 self.wake_all_waiters();
284 self.agent.abort("advisor disposed");
285 }
286
287 fn reset_advisor_context(&self, clear_backlog: bool) {
288 self.last_count.store(0, Ordering::SeqCst);
289 let mut s = self.state.lock();
290 s.pending.clear();
291 if clear_backlog {
292 self.backlog.store(0, Ordering::SeqCst);
293 }
294 self.consecutive_failures.store(0, Ordering::SeqCst);
295 self.failure_notified.store(false, Ordering::SeqCst);
296 drop(s);
297 self.agent.reset();
298 self.agent.abort("advisor reset");
299 }
300
301 fn render_delta(&self, messages: &[Message]) -> Option<String> {
304 let last = self.last_count.load(Ordering::SeqCst) as usize;
305 if messages.len() < last {
306 self.last_count
307 .store(messages.len() as u64, Ordering::SeqCst);
308 return None;
309 }
310 let delta = &messages[last..];
311 self.last_count
312 .store(messages.len() as u64, Ordering::SeqCst);
313 if delta.is_empty() {
314 return None;
315 }
316 let mut parts: Vec<String> = Vec::new();
317 for msg in delta {
318 if let Some(md) = format_message_md(msg) {
319 parts.push(md);
320 }
321 }
322 if parts.is_empty() {
323 return None;
324 }
325 Some(format!("### Session update\n\n{}", parts.join("\n\n")))
326 }
327
328 fn wake_all_waiters(&self) {
329 let mut waiters = self.waiters.lock();
330 for w in waiters.drain(..) {
331 if let Some(tx) = w.tx {
332 let _ = tx.send(());
333 }
334 }
335 }
336
337 fn notify_waiters(&self) {
338 let mut waiters = self.waiters.lock();
339 let backlog = self.backlog.load(Ordering::SeqCst);
340 for w in waiters.iter_mut() {
341 if backlog < w.threshold
342 && let Some(tx) = w.tx.take()
343 {
344 let _ = tx.send(());
345 }
346 }
347 waiters.retain(|w| w.tx.is_some());
348 }
349
350 fn decrement_backlog(&self, by: u64) {
351 let mut prev = self.backlog.load(Ordering::SeqCst);
352 loop {
353 let next = prev.saturating_sub(by);
354 match self
355 .backlog
356 .compare_exchange(prev, next, Ordering::SeqCst, Ordering::SeqCst)
357 {
358 Ok(_) => break,
359 Err(actual) => prev = actual,
360 }
361 }
362 }
363
364 async fn drain(self: Arc<Self>) {
368 {
369 let mut s = self.state.lock();
370 if s.draining || s.pending.is_empty() {
371 return;
372 }
373 s.draining = true;
374 }
375 loop {
376 let (batch_text, turns_covered) = {
378 let mut s = self.state.lock();
379 if s.pending.is_empty() {
380 s.draining = false;
384 return;
385 }
386 let taken: Vec<PendingDelta> = s.pending.drain(..).collect();
387 let turns: u64 = taken.iter().map(|d| d.turns).sum();
388 let joined = taken
389 .into_iter()
390 .map(|d| d.text)
391 .collect::<Vec<_>>()
392 .join("\n\n");
393 (joined, turns)
394 };
395
396 let epoch_start = self.epoch.load(Ordering::SeqCst);
397
398 let should_reprime = self.host.maintain_context(batch_text.len());
401 if self.epoch.load(Ordering::SeqCst) != epoch_start {
402 continue;
403 }
404
405 let (batch, final_turns) = if should_reprime {
406 self.reset_advisor_context(false);
409 let new_turns = self.state.lock().pending.len() as u64;
410 let rendered = self
411 .latest
412 .lock()
413 .as_ref()
414 .and_then(|m| self.render_delta(m));
415 let final_turns = turns_covered.saturating_add(new_turns);
416 match rendered {
417 Some(b) => (b, final_turns),
418 None => {
419 self.decrement_backlog(final_turns);
420 self.notify_waiters();
421 continue;
422 }
423 }
424 } else {
425 (batch_text, turns_covered)
426 };
427
428 if self.disposed.load(Ordering::SeqCst) {
429 self.decrement_backlog(final_turns);
430 self.notify_waiters();
431 continue;
432 }
433
434 let message_snapshot = self.agent.message_count();
435 self.host.begin_advisor_update();
436 let prompt_result = self.agent.prompt(batch.clone()).await;
437
438 if self.epoch.load(Ordering::SeqCst) != epoch_start {
441 continue;
442 }
443
444 let success;
445 match prompt_result {
446 Ok(()) => {
447 self.consecutive_failures.store(0, Ordering::SeqCst);
448 self.failure_notified.store(false, Ordering::SeqCst);
449 success = true;
450 }
451 Err(err) => {
452 self.agent.rollback_to(message_snapshot).await;
453 let failures = self.consecutive_failures.fetch_add(1, Ordering::SeqCst) + 1;
454 if failures >= 3 {
455 tracing::warn!(
456 failures,
457 "advisor failed consecutively; dropping backlog to prevent stall"
458 );
459 if !self.failure_notified.swap(true, Ordering::SeqCst) {
460 self.host.notify_failure(&err);
461 }
462 self.consecutive_failures.store(0, Ordering::SeqCst);
463 success = true;
464 } else {
465 {
467 let mut s = self.state.lock();
468 s.pending.insert(
469 0,
470 PendingDelta {
471 text: batch,
472 turns: final_turns,
473 },
474 );
475 }
476 tokio::time::sleep(self.retry_delay).await;
477 continue;
478 }
479 }
480 }
481
482 if success {
483 self.decrement_backlog(final_turns);
484 self.notify_waiters();
485 }
486 }
487 }
488}
489
490fn format_message_md(msg: &Message) -> Option<String> {
494 let role = match msg {
495 Message::User(_) => "user",
496 Message::Assistant(_) => "assistant",
497 Message::ToolResult(_) => "tool",
498 };
499 let text = msg.text_content().unwrap_or_default();
500 if text.trim().is_empty() {
501 return None;
502 }
503 Some(format!("**[{role}]**\n{text}"))
504}
505
506#[cfg(test)]
507mod tests {
508 #![allow(clippy::unwrap_used)]
509 use super::*;
510 use std::sync::Mutex as StdMutex;
511 type PromptLog = Arc<StdMutex<Vec<String>>>;
512 type AdviceLog = Arc<StdMutex<Vec<AdvisorNote>>>;
513
514 struct FakeAgent {
516 prompts: PromptLog,
517 fail_first_n: AtomicU32,
518 messages_len: AtomicU64,
519 }
520
521 impl FakeAgent {
522 fn new() -> (Arc<Self>, PromptLog) {
523 let prompts = Arc::new(StdMutex::new(Vec::new()));
524 let a = Arc::new(Self {
525 prompts: Arc::clone(&prompts),
526 fail_first_n: AtomicU32::new(0),
527 messages_len: AtomicU64::new(0),
528 });
529 (a, prompts)
530 }
531 }
532
533 #[async_trait]
534 impl AdvisorAgent for FakeAgent {
535 async fn prompt(&self, input: String) -> Result<(), String> {
536 self.messages_len.fetch_add(4, Ordering::SeqCst);
538 self.prompts.lock().unwrap().push(input);
539 let n = self.fail_first_n.load(Ordering::SeqCst);
543 if n > 0 {
544 self.fail_first_n.fetch_sub(1, Ordering::SeqCst);
545 Err("simulated advisor failure".into())
546 } else {
547 Ok(())
548 }
549 }
550 fn abort(&self, _reason: &str) {}
551 fn reset(&self) {
552 self.messages_len.store(0, Ordering::SeqCst);
553 }
554 async fn rollback_to(&self, count: usize) {
555 self.messages_len.store(count as u64, Ordering::SeqCst);
556 }
557 fn message_count(&self) -> usize {
558 self.messages_len.load(Ordering::SeqCst) as usize
559 }
560 }
561
562 struct FakeHost {
564 advice: AdviceLog,
565 }
566 impl AdvisorRuntimeHost for FakeHost {
567 fn snapshot_messages(&self) -> Vec<Message> {
568 Vec::new()
569 }
570 fn enqueue_advice(&self, note: AdvisorNote) {
571 self.advice.lock().unwrap().push(note);
572 }
573 }
574
575 fn build() -> (Arc<AdvisorRuntime>, PromptLog, AdviceLog) {
576 let (agent, prompts) = FakeAgent::new();
577 let advice = Arc::new(StdMutex::new(Vec::new()));
578 let host: Arc<dyn AdvisorRuntimeHost> = Arc::new(FakeHost {
579 advice: Arc::clone(&advice),
580 });
581 let rt = Arc::new(AdvisorRuntime::new(agent, host, Duration::from_millis(10)));
582 rt.install_self(Arc::downgrade(&rt));
583 (rt, prompts, advice)
584 }
585
586 fn user_msg(s: &str) -> Message {
587 Message::user(s)
588 }
589
590 #[tokio::test]
591 async fn drain_prompts_advisor_with_delta() {
592 let (rt, prompts, _advice) = build();
593 rt.on_turn_end(vec![user_msg("turn 1")]);
594 tokio::time::sleep(Duration::from_millis(50)).await;
596 let p = prompts.lock().unwrap();
597 assert_eq!(p.len(), 1);
598 assert!(p[0].contains("turn 1"));
599 assert!(p[0].starts_with("### Session update"));
600 }
601
602 #[tokio::test]
603 async fn reset_aborts_inflight_and_drops_batch() {
604 let (rt, prompts, _advice) = build();
605 rt.on_turn_end(vec![user_msg("turn 1")]);
606 rt.reset(); tokio::time::sleep(Duration::from_millis(50)).await;
608 assert_eq!(rt.backlog(), 0);
611 let _ = prompts.lock().unwrap().len();
612 }
613
614 #[tokio::test]
615 async fn drain_exit_racing_turn_end_no_lost_wakeup() {
616 let (rt, _prompts, _advice) = build();
619 let rt2 = Arc::clone(&rt);
620 let handles: Vec<_> = (0..20)
621 .map(move |i| {
622 let rt3 = Arc::clone(&rt2);
623 tokio::spawn(async move {
624 rt3.on_turn_end(vec![user_msg(&format!("turn {i}"))]);
625 })
626 })
627 .collect();
628 for h in handles {
629 h.await.unwrap();
630 }
631 tokio::time::sleep(Duration::from_millis(120)).await;
633 assert_eq!(rt.backlog(), 0);
634 let pending = rt.state.lock().pending.len();
636 assert_eq!(pending, 0);
637 }
638
639 #[tokio::test]
640 async fn wait_for_catchup_resolves_below_threshold() {
641 let (rt, _prompts, _advice) = build();
642 rt.on_turn_end(vec![user_msg("turn 1")]);
643 rt.wait_for_catchup(Duration::from_millis(50), 0).await;
645 let _ = tokio::time::timeout(Duration::from_millis(200), async {
647 while rt.backlog() > 0 {
648 tokio::time::sleep(Duration::from_millis(5)).await;
649 }
650 })
651 .await;
652 assert_eq!(rt.backlog(), 0);
653 }
654
655 #[tokio::test]
656 async fn seed_to_skips_history() {
657 let (rt, prompts, _advice) = build();
658 rt.seed_to(5); rt.on_turn_end(vec![user_msg("a"), user_msg("b"), user_msg("c")]);
661 tokio::time::sleep(Duration::from_millis(30)).await;
662 assert!(prompts.lock().unwrap().is_empty());
663 }
664
665 #[tokio::test]
666 async fn reprime_via_maintain_context() {
667 struct ReprimeHost {
670 advice: Arc<StdMutex<Vec<AdvisorNote>>>,
671 }
672 impl AdvisorRuntimeHost for ReprimeHost {
673 fn snapshot_messages(&self) -> Vec<Message> {
674 Vec::new()
675 }
676 fn enqueue_advice(&self, n: AdvisorNote) {
677 self.advice.lock().unwrap().push(n);
678 }
679 fn maintain_context(&self, _t: usize) -> bool {
680 true
681 }
682 }
683 let (agent, prompts) = FakeAgent::new();
684 let advice = Arc::new(StdMutex::new(Vec::new()));
685 let host: Arc<dyn AdvisorRuntimeHost> = Arc::new(ReprimeHost {
686 advice: Arc::clone(&advice),
687 });
688 let rt = Arc::new(AdvisorRuntime::new(agent, host, Duration::from_millis(10)));
689 rt.install_self(Arc::downgrade(&rt));
690 rt.on_turn_end(vec![user_msg("turn 1"), user_msg("turn 2")]);
691 tokio::time::sleep(Duration::from_millis(60)).await;
692 let p = prompts.lock().unwrap();
693 assert!(!p.is_empty());
694 assert!(p[0].contains("turn 1") && p[0].contains("turn 2"));
696 }
697}