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)
}
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}")
}