use std::collections::HashMap;
use std::fmt;
use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, Mutex, MutexGuard, PoisonError};
use std::task::{Context, Poll};
use async_trait::async_trait;
use everruns_contracts::error::{AgentLoopError, Result};
use everruns_contracts::typed_id::{AgentId, HarnessId, MessageId, SessionId, TurnId};
use tokio_util::sync::CancellationToken;
use uuid::Uuid;
use super::runtime::{AcceptedTurnInput, InProcessRuntime, TurnResult, TurnSteering};
use crate::events::ToolCompletedData;
#[async_trait]
pub trait TurnBackend: Send + Sync {
async fn start_turn(&self, request: TurnRequest) -> Result<TurnTicket>;
async fn cancel(&self, session_id: SessionId) -> Result<bool>;
async fn is_running(&self, session_id: SessionId) -> bool;
async fn active_count(&self) -> usize;
}
#[derive(Debug)]
#[non_exhaustive]
pub enum TurnInput {
Message(Box<AcceptedTurnInput>),
StoredMessage {
message_id: MessageId,
},
ResumeInterrupted,
ToolResults(Vec<ToolCompletedData>),
RecordedToolResults {
resolution_id: Uuid,
},
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub struct TurnScope {
pub org_id: i64,
pub harness_id: HarnessId,
pub agent_id: Option<AgentId>,
}
impl TurnScope {
pub fn new(org_id: i64, harness_id: HarnessId, agent_id: Option<AgentId>) -> Self {
Self {
org_id,
harness_id,
agent_id,
}
}
}
#[derive(Debug)]
#[non_exhaustive]
pub struct TurnRequest {
pub session_id: SessionId,
pub turn_id: TurnId,
pub input: TurnInput,
pub steering: TurnSteering,
pub scope: Option<TurnScope>,
pub request_id: Option<String>,
}
impl TurnRequest {
pub fn new(session_id: SessionId, turn_id: TurnId, input: TurnInput) -> Self {
Self {
session_id,
turn_id,
input,
steering: TurnSteering::new(),
scope: None,
request_id: None,
}
}
pub fn with_scope(mut self, scope: TurnScope) -> Self {
self.scope = Some(scope);
self
}
pub fn with_request_id(mut self, request_id: Option<String>) -> Self {
self.request_id = request_id;
self
}
pub fn with_steering(mut self, steering: TurnSteering) -> Self {
self.steering = steering;
self
}
}
type TurnCompletion = Pin<Box<dyn Future<Output = Result<TurnResult>> + Send>>;
pub struct TurnTicket {
session_id: SessionId,
turn_id: TurnId,
completion: TurnCompletion,
}
impl TurnTicket {
pub fn new(
session_id: SessionId,
turn_id: TurnId,
completion: impl Future<Output = Result<TurnResult>> + Send + 'static,
) -> Self {
Self {
session_id,
turn_id,
completion: Box::pin(completion),
}
}
pub fn session_id(&self) -> SessionId {
self.session_id
}
pub fn turn_id(&self) -> TurnId {
self.turn_id
}
}
impl Future for TurnTicket {
type Output = Result<TurnResult>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
self.completion.as_mut().poll(cx)
}
}
impl fmt::Debug for TurnTicket {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("TurnTicket")
.field("session_id", &self.session_id)
.field("turn_id", &self.turn_id)
.finish_non_exhaustive()
}
}
#[derive(Clone)]
pub struct InProcessBackend {
runtime: InProcessRuntime,
turns: Arc<Mutex<Turns>>,
}
#[derive(Default)]
struct Turns {
next_generation: u64,
running: HashMap<SessionId, RunningTurn>,
}
struct RunningTurn {
generation: u64,
cancel: CancellationToken,
}
impl InProcessBackend {
pub fn new(runtime: InProcessRuntime) -> Self {
Self {
runtime,
turns: Arc::default(),
}
}
fn turns(&self) -> MutexGuard<'_, Turns> {
lock_turns(&self.turns)
}
}
fn lock_turns(turns: &Mutex<Turns>) -> MutexGuard<'_, Turns> {
turns.lock().unwrap_or_else(PoisonError::into_inner)
}
impl fmt::Debug for InProcessBackend {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("InProcessBackend")
.field("active", &self.turns().running.len())
.finish_non_exhaustive()
}
}
struct Registration {
turns: Arc<Mutex<Turns>>,
session_id: SessionId,
generation: u64,
}
impl Drop for Registration {
fn drop(&mut self) {
let mut turns = lock_turns(&self.turns);
if turns
.running
.get(&self.session_id)
.is_some_and(|turn| turn.generation == self.generation)
{
turns.running.remove(&self.session_id);
}
}
}
#[async_trait]
impl TurnBackend for InProcessBackend {
async fn start_turn(&self, request: TurnRequest) -> Result<TurnTicket> {
let TurnRequest {
session_id,
turn_id,
input,
steering,
..
} = request;
let cancel = CancellationToken::new();
let registration = {
let mut turns = self.turns();
if turns.running.contains_key(&session_id) {
return Err(AgentLoopError::store(format!(
"session {session_id} already runs a turn"
)));
}
let generation = turns.next_generation;
turns.next_generation += 1;
turns.running.insert(
session_id,
RunningTurn {
generation,
cancel: cancel.clone(),
},
);
Registration {
turns: self.turns.clone(),
session_id,
generation,
}
};
let runtime = self.runtime.clone();
let turn = async move {
match input {
TurnInput::Message(input) => {
runtime
.run_steerable_turn(session_id, *input, turn_id, steering)
.await
}
TurnInput::ResumeInterrupted => {
runtime.resume_interrupted_turn(session_id, steering).await
}
TurnInput::ToolResults(results) => {
runtime
.resume_steerable_turn(session_id, results, steering)
.await
}
TurnInput::StoredMessage { message_id } => {
runtime
.run_stored_turn(session_id, message_id, turn_id, steering)
.await
}
TurnInput::RecordedToolResults { resolution_id: _ } => {
runtime
.resume_steerable_turn(session_id, Vec::new(), steering)
.await
}
}
};
let completion = async move {
let _registration = registration;
tokio::select! {
biased;
() = cancel.cancelled() => Err(AgentLoopError::Cancelled),
result = turn => result,
}
};
Ok(TurnTicket::new(session_id, turn_id, completion))
}
async fn cancel(&self, session_id: SessionId) -> Result<bool> {
let turns = self.turns();
Ok(match turns.running.get(&session_id) {
Some(turn) => {
turn.cancel.cancel();
true
}
None => false,
})
}
async fn is_running(&self, session_id: SessionId) -> bool {
self.turns()
.running
.get(&session_id)
.is_some_and(|turn| !turn.cancel.is_cancelled())
}
async fn active_count(&self) -> usize {
self.turns()
.running
.values()
.filter(|turn| !turn.cancel.is_cancelled())
.count()
}
}