use std::sync::{
Arc, Mutex, MutexGuard, TryLockError,
atomic::{AtomicBool, Ordering},
};
const MAX_PENDING_STEERING_MESSAGES: usize = 16;
pub(crate) const STEERING_HEADER: &str = "Steering update from user while current run was active:";
#[derive(Clone, Debug)]
pub(crate) struct AgentSteering {
inner: Arc<Mutex<SteeringQueue>>,
closed: Arc<AtomicBool>,
}
#[derive(Debug, Default)]
struct SteeringQueue {
messages: Vec<String>,
generation: u64,
reserved: Option<SteeringBatch>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct SteeringQueued {
pub(crate) position: usize,
pub(crate) pending_count: usize,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct SteeringBatch {
pub(crate) text: String,
pub(crate) count: usize,
pub(crate) original_prompts: Vec<String>,
generation: u64,
}
pub(crate) struct SteeringReservation {
steering: AgentSteering,
batch: SteeringBatch,
}
impl SteeringReservation {
pub(crate) fn text(&self) -> &str {
&self.batch.text
}
}
impl Drop for SteeringReservation {
fn drop(&mut self) {
let mut queue = self.steering.lock_queue();
if queue.reserved.as_ref() == Some(&self.batch) {
queue.reserved = None;
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) enum SteeringRejected {
Empty,
Full { capacity: usize },
Closed,
Busy,
}
impl AgentSteering {
pub(crate) fn new() -> Self {
Self {
inner: Arc::new(Mutex::new(SteeringQueue::default())),
closed: Arc::new(AtomicBool::new(false)),
}
}
pub(crate) fn try_enqueue(&self, text: String) -> Result<SteeringQueued, SteeringRejected> {
self.enqueue_with_queue(text, self.lock_queue())
}
pub(crate) fn try_enqueue_nonblocking(
&self,
text: String,
) -> Result<SteeringQueued, SteeringRejected> {
let queue = match self.inner.try_lock() {
Ok(queue) => queue,
Err(TryLockError::Poisoned(poisoned)) => poisoned.into_inner(),
Err(TryLockError::WouldBlock) => return Err(SteeringRejected::Busy),
};
self.enqueue_with_queue(text, queue)
}
fn enqueue_with_queue(
&self,
text: String,
mut queue: MutexGuard<'_, SteeringQueue>,
) -> Result<SteeringQueued, SteeringRejected> {
let text = text.trim().to_string();
if text.is_empty() {
return Err(SteeringRejected::Empty);
}
if self.closed.load(Ordering::SeqCst) {
return Err(SteeringRejected::Closed);
}
if queue.messages.len() >= MAX_PENDING_STEERING_MESSAGES {
return Err(SteeringRejected::Full {
capacity: MAX_PENDING_STEERING_MESSAGES,
});
}
let position = queue.messages.len();
queue.messages.push(text);
Ok(SteeringQueued {
position,
pending_count: queue.messages.len()
- queue.reserved.as_ref().map_or(0, |batch| batch.count),
})
}
pub(crate) fn observe_collapsed(&self) -> Option<SteeringBatch> {
let queue = self.lock_queue();
(queue.reserved.is_none() && !queue.messages.is_empty()).then(|| SteeringBatch {
text: collapse_steering_messages(&queue.messages),
count: queue.messages.len(),
original_prompts: queue.messages.clone(),
generation: queue.generation,
})
}
pub(crate) fn reserve_collapsed(&self) -> Option<SteeringReservation> {
self.reserve_or_close(false)
}
pub(crate) fn reserve_collapsed_or_close(&self) -> Option<SteeringReservation> {
self.reserve_or_close(true)
}
pub(crate) fn close(&self) {
self.closed.store(true, Ordering::SeqCst);
}
fn reserve_or_close(&self, close_when_empty: bool) -> Option<SteeringReservation> {
let mut queue = self.lock_queue();
if queue.reserved.is_some() || queue.messages.is_empty() {
if close_when_empty && queue.messages.is_empty() {
self.close();
}
return None;
}
let batch = SteeringBatch {
text: collapse_steering_messages(&queue.messages),
count: queue.messages.len(),
original_prompts: queue.messages.clone(),
generation: queue.generation,
};
queue.reserved = Some(batch.clone());
Some(SteeringReservation {
steering: self.clone(),
batch,
})
}
pub(crate) fn restore_pending_messages(
&self,
restore: impl FnOnce(&str) -> bool,
) -> Result<bool, ()> {
let mut queue = match self.inner.try_lock() {
Ok(queue) => queue,
Err(TryLockError::Poisoned(poisoned)) => poisoned.into_inner(),
Err(TryLockError::WouldBlock) => return Err(()),
};
let reserved_count = queue.reserved.as_ref().map_or(0, |batch| batch.count);
if queue.messages.len() == reserved_count {
return Ok(false);
}
if restore(&queue.messages[reserved_count..].join("\n")) {
queue.messages.truncate(reserved_count);
queue.generation += 1;
}
Ok(true)
}
pub(crate) fn persist_and_acknowledge<E>(
&self,
batch: &SteeringBatch,
persist: impl FnOnce() -> Result<(), E>,
) -> Result<bool, E> {
let mut queue = self.lock_queue();
if queue.generation != batch.generation
|| queue.reserved.is_some()
|| batch.count == 0
|| batch.count > queue.messages.len()
|| collapse_steering_messages(&queue.messages[..batch.count]) != batch.text
{
return Ok(false);
}
persist()?;
queue.messages.drain(..batch.count);
queue.generation += 1;
Ok(true)
}
pub(crate) fn acknowledge_reserved_prompt(&self, text: &str) -> Option<Vec<String>> {
let mut queue = self.lock_queue();
let batch = queue.reserved.as_ref().filter(|batch| batch.text == text)?;
let count = batch.count;
let prompts = queue.messages.drain(..count).collect();
queue.reserved = None;
queue.generation += 1;
Some(prompts)
}
#[cfg(test)]
pub(crate) fn drain_collapsed(&self) -> Option<String> {
let batch = self.observe_collapsed()?;
self.persist_and_acknowledge(&batch, || Ok::<_, std::convert::Infallible>(()))
.expect("infallible persistence")
.then_some(batch.text)
}
pub(crate) fn try_pending_count(&self) -> Option<usize> {
let queue = match self.inner.try_lock() {
Ok(queue) => queue,
Err(TryLockError::Poisoned(poisoned)) => poisoned.into_inner(),
Err(TryLockError::WouldBlock) => return None,
};
Some(queue.messages.len() - queue.reserved.as_ref().map_or(0, |batch| batch.count))
}
pub(crate) fn pending_count(&self) -> usize {
let queue = self.lock_queue();
queue.messages.len() - queue.reserved.as_ref().map_or(0, |batch| batch.count)
}
pub(crate) fn clear(&self) -> usize {
let mut queue = self.lock_queue();
let reserved_count = queue.reserved.as_ref().map_or(0, |batch| batch.count);
let count = queue.messages.len() - reserved_count;
queue.messages.truncate(reserved_count);
if count > 0 {
queue.generation += 1;
}
count
}
fn lock_queue(&self) -> MutexGuard<'_, SteeringQueue> {
self.inner
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
}
impl Default for AgentSteering {
fn default() -> Self {
Self::new()
}
}
fn collapse_steering_messages(messages: &[String]) -> String {
if messages.len() == 1 {
return format!("{STEERING_HEADER}\n\n{}", messages[0]);
}
let numbered = messages
.iter()
.enumerate()
.map(|(index, message)| format!("{}. {message}", index + 1))
.collect::<Vec<_>>()
.join("\n\n");
format!("{STEERING_HEADER}\n\n{numbered}")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn busy_persistence_does_not_block_enqueue_close_or_snapshot() {
let steering = AgentSteering::new();
steering.try_enqueue("accepted".into()).unwrap();
let batch = steering.observe_collapsed().unwrap();
let result = steering.persist_and_acknowledge(&batch, || {
assert_eq!(
steering.try_enqueue_nonblocking("retry".into()),
Err(SteeringRejected::Busy)
);
assert_eq!(steering.try_pending_count(), None);
steering.close();
Err::<(), _>("disk failed")
});
assert_eq!(result, Err("disk failed"));
assert_eq!(
steering.try_enqueue("too late".into()),
Err(SteeringRejected::Closed)
);
assert_eq!(
steering.restore_pending_messages(|text| {
assert_eq!(text, "accepted");
true
}),
Ok(true)
);
}
#[test]
fn reservation_excludes_only_preparing_prefix_from_recall_until_persisted() {
let steering = AgentSteering::new();
steering.try_enqueue("preparing".into()).unwrap();
let reservation = steering.reserve_collapsed().unwrap();
assert_eq!(steering.pending_count(), 0);
assert_eq!(
steering.restore_pending_messages(|_| panic!("reserved input cannot be recalled")),
Ok(false)
);
steering.try_enqueue("later".into()).unwrap();
assert_eq!(
steering.restore_pending_messages(|text| {
assert_eq!(text, "later");
true
}),
Ok(true)
);
steering.try_enqueue("newer".into()).unwrap();
assert!(
steering
.acknowledge_reserved_prompt(reservation.text())
.is_some()
);
drop(reservation);
assert_eq!(steering.pending_count(), 1);
assert_eq!(
steering.restore_pending_messages(|text| {
assert_eq!(text, "newer");
true
}),
Ok(true)
);
}
#[test]
fn abandoned_preparation_releases_unpersisted_input_for_recall() {
let steering = AgentSteering::new();
steering.try_enqueue("retry after failure".into()).unwrap();
let observed = steering.observe_collapsed().unwrap();
let reservation = steering.reserve_collapsed().unwrap();
assert_eq!(
steering.persist_and_acknowledge(&observed, || -> Result<(), ()> {
panic!("another consumer cannot persist reserved input")
}),
Ok(false)
);
assert_eq!(steering.clear(), 0);
assert!(
steering
.acknowledge_reserved_prompt("wrong prompt")
.is_none()
);
drop(reservation);
assert_eq!(
steering.restore_pending_messages(|text| {
assert_eq!(text, "retry after failure");
true
}),
Ok(true)
);
assert!(steering.reserve_collapsed().is_none());
}
#[test]
fn old_reservation_drop_does_not_release_a_new_follow_up() {
let steering = AgentSteering::new();
steering.try_enqueue("same text".into()).unwrap();
let first = steering.reserve_collapsed().unwrap();
assert_eq!(
steering.acknowledge_reserved_prompt(first.text()),
Some(vec!["same text".into()])
);
steering.try_enqueue("same text".into()).unwrap();
let second = steering.reserve_collapsed().unwrap();
drop(first);
assert_eq!(
steering.restore_pending_messages(|_| panic!("new batch remains reserved")),
Ok(false)
);
drop(second);
assert_eq!(
steering.restore_pending_messages(|text| text == "same text"),
Ok(true)
);
}
#[test]
fn recalled_batch_cannot_persist_or_remove_identical_new_messages() {
let steering = AgentSteering::new();
steering.try_enqueue("first".into()).unwrap();
steering.try_enqueue("second\nline".into()).unwrap();
let batch = steering.observe_collapsed().unwrap();
let mut restored = String::new();
assert_eq!(
steering.restore_pending_messages(|text| {
restored = text.to_owned();
true
}),
Ok(true)
);
assert_eq!(restored, "first\nsecond\nline");
assert_eq!(steering.pending_count(), 0);
steering.try_enqueue("first".into()).unwrap();
steering.try_enqueue("second\nline".into()).unwrap();
assert_eq!(
steering.persist_and_acknowledge(&batch, || -> Result<(), ()> {
panic!("recalled input must not be persisted")
}),
Ok(false)
);
assert_eq!(steering.pending_count(), 2);
}
#[test]
fn editor_rejection_preserves_pending_batch_for_injection() {
let steering = AgentSteering::new();
steering.try_enqueue("keep this".into()).unwrap();
let batch = steering.observe_collapsed().unwrap();
assert_eq!(
steering.restore_pending_messages(|text| {
assert_eq!(text, "keep this");
false
}),
Ok(true)
);
assert_eq!(steering.observe_collapsed(), Some(batch.clone()));
assert_eq!(
steering.persist_and_acknowledge(&batch, || Ok::<_, ()>(())),
Ok(true)
);
assert_eq!(steering.pending_count(), 0);
}
#[test]
fn persistence_failure_preserves_batch_and_recall_does_not_wait_for_write() {
let steering = AgentSteering::new();
steering.try_enqueue("retry this".into()).unwrap();
let batch = steering.observe_collapsed().unwrap();
assert_eq!(
steering.persist_and_acknowledge(&batch, || {
assert_eq!(
steering.restore_pending_messages(|_| panic!("queue is busy")),
Err(())
);
Err("write failed")
}),
Err("write failed")
);
assert_eq!(steering.observe_collapsed(), Some(batch.clone()));
assert_eq!(
steering.persist_and_acknowledge(&batch, || Ok::<_, ()>(())),
Ok(true)
);
assert_eq!(steering.pending_count(), 0);
assert_eq!(
steering.restore_pending_messages(|_| panic!("queue is empty")),
Ok(false)
);
}
#[test]
fn changed_prefix_or_removed_batch_is_not_persisted_twice() {
let steering = AgentSteering::new();
steering.try_enqueue("one".into()).unwrap();
let batch = steering.observe_collapsed().unwrap();
let mut mismatched = batch.clone();
mismatched.text = "different input".into();
assert_eq!(
steering.persist_and_acknowledge(&mismatched, || -> Result<(), ()> {
panic!("mismatched input must not be persisted")
}),
Ok(false)
);
assert_eq!(
steering.persist_and_acknowledge(&batch, || Ok::<_, ()>(())),
Ok(true)
);
steering.try_enqueue("one".into()).unwrap();
assert_eq!(
steering.persist_and_acknowledge(&batch, || -> Result<(), ()> {
panic!("already consumed input must not be persisted")
}),
Ok(false)
);
assert_eq!(steering.pending_count(), 1);
}
#[test]
fn enqueue_single_message_drain_returns_header_and_message() {
let steering = AgentSteering::new();
assert_eq!(
steering.try_enqueue("keep going".to_string()),
Ok(SteeringQueued {
position: 0,
pending_count: 1,
})
);
assert_eq!(
steering.drain_collapsed(),
Some(
"Steering update from user while current run was active:\n\nkeep going".to_string()
)
);
}
#[test]
fn enqueue_multiple_messages_drain_returns_numbered_format() {
let steering = AgentSteering::new();
assert_eq!(
steering.try_enqueue("first".to_string()),
Ok(SteeringQueued {
position: 0,
pending_count: 1,
})
);
assert_eq!(
steering.try_enqueue("second".to_string()),
Ok(SteeringQueued {
position: 1,
pending_count: 2,
})
);
assert_eq!(
steering.try_enqueue("third".to_string()),
Ok(SteeringQueued {
position: 2,
pending_count: 3,
})
);
assert_eq!(
steering.drain_collapsed(),
Some(
"Steering update from user while current run was active:\n\n1. first\n\n2. second\n\n3. third"
.to_string()
)
);
}
#[test]
fn enqueue_empty_or_whitespace_only_text_rejects_empty() {
let steering = AgentSteering::new();
assert_eq!(
steering.try_enqueue("".to_string()),
Err(SteeringRejected::Empty)
);
assert_eq!(
steering.try_enqueue(" \n\t ".to_string()),
Err(SteeringRejected::Empty)
);
assert_eq!(steering.pending_count(), 0);
}
#[test]
fn enqueue_over_capacity_rejects_full() {
let steering = AgentSteering::new();
for index in 0..MAX_PENDING_STEERING_MESSAGES {
assert_eq!(
steering.try_enqueue(format!("message {index}")),
Ok(SteeringQueued {
position: index,
pending_count: index + 1,
})
);
}
assert_eq!(
steering.try_enqueue("overflow".to_string()),
Err(SteeringRejected::Full {
capacity: MAX_PENDING_STEERING_MESSAGES,
})
);
assert_eq!(steering.pending_count(), MAX_PENDING_STEERING_MESSAGES);
}
#[test]
fn acknowledge_observed_prefix_preserves_concurrent_append() {
let steering = AgentSteering::new();
steering.try_enqueue("first".to_string()).unwrap();
steering.try_enqueue("second".to_string()).unwrap();
let batch = steering.observe_collapsed().unwrap();
steering.try_enqueue("later".to_string()).unwrap();
assert_eq!(
steering.persist_and_acknowledge(&batch, || Ok::<_, ()>(())),
Ok(true)
);
assert_eq!(steering.pending_count(), 1);
assert_eq!(
steering.drain_collapsed(),
Some("Steering update from user while current run was active:\n\nlater".to_string())
);
}
#[test]
fn drain_collapsed_on_empty_queue_returns_none() {
let steering = AgentSteering::new();
assert_eq!(steering.drain_collapsed(), None);
}
#[test]
fn clear_returns_count_and_empties_queue() {
let steering = AgentSteering::new();
steering.try_enqueue("one".to_string()).unwrap();
steering.try_enqueue("two".to_string()).unwrap();
assert_eq!(steering.clear(), 2);
assert_eq!(steering.pending_count(), 0);
assert_eq!(steering.clear(), 0);
}
#[test]
fn recovers_after_queue_lock_poison() {
let steering = AgentSteering::new();
steering.try_enqueue("before poison".to_string()).unwrap();
let _ = std::panic::catch_unwind(|| {
let _guard = steering.inner.lock().unwrap();
panic!("poison steering queue lock");
});
assert_eq!(steering.pending_count(), 1);
assert_eq!(
steering.try_enqueue("after poison".to_string()),
Ok(SteeringQueued {
position: 1,
pending_count: 2,
})
);
assert_eq!(
steering.drain_collapsed(),
Some(
"Steering update from user while current run was active:\n\n1. before poison\n\n2. after poison"
.to_string()
)
);
assert_eq!(steering.clear(), 0);
}
#[test]
fn pending_count_tracks_after_enqueue_and_drain() {
let steering = AgentSteering::new();
assert_eq!(steering.pending_count(), 0);
steering.try_enqueue("one".to_string()).unwrap();
steering.try_enqueue("two".to_string()).unwrap();
assert_eq!(steering.pending_count(), 2);
steering.drain_collapsed().unwrap();
assert_eq!(steering.pending_count(), 0);
}
#[test]
fn drain_collapsed_drains_only_non_empty_queue() {
let steering = AgentSteering::new();
assert_eq!(steering.drain_collapsed(), None);
steering.try_enqueue("one".to_string()).unwrap();
assert_eq!(
steering.drain_collapsed(),
Some("Steering update from user while current run was active:\n\none".to_string())
);
assert_eq!(steering.drain_collapsed(), None);
steering.try_enqueue("two".to_string()).unwrap();
assert_eq!(
steering.drain_collapsed(),
Some("Steering update from user while current run was active:\n\ntwo".to_string())
);
assert_eq!(steering.drain_collapsed(), None);
}
#[test]
fn enqueue_trims_surrounding_whitespace() {
let steering = AgentSteering::new();
steering
.try_enqueue(" trimmed text\n".to_string())
.unwrap();
assert_eq!(
steering.drain_collapsed(),
Some(
"Steering update from user while current run was active:\n\ntrimmed text"
.to_string()
)
);
}
}