use crate::canonical::CanonicalUnitEvent;
use crate::config::{InvocationConfig, SessionConfig};
use crate::id::ToolId;
use crate::id::{ChannelId, SessionId, SessionKey, TransactionId};
use crate::input::CanonicalInput;
use crate::safe::SafeDiagnostic;
use crate::tool::ToolLifecycleEvent;
use serde::{Deserialize, Serialize};
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::time::Duration;
use thiserror::Error;
pub type EventDelivery =
Pin<Box<dyn Future<Output = Result<(), EventDeliveryError>> + Send + 'static>>;
pub trait TransactionEventSink: Send + Sync + 'static {
fn deliver(&self, event: TransactionEvent) -> EventDelivery;
}
pub type CompletionDelivery =
Pin<Box<dyn Future<Output = Result<(), CompletionDeliveryError>> + Send + 'static>>;
pub trait CompletionCallback: Send + 'static {
fn call(self: Box<Self>, end: TransactionEnd) -> CompletionDelivery;
}
pub struct FnEventSink<F>(pub F);
impl<F> TransactionEventSink for FnEventSink<F>
where
F: Fn(TransactionEvent) -> EventDelivery + Send + Sync + 'static,
{
fn deliver(&self, event: TransactionEvent) -> EventDelivery {
(self.0)(event)
}
}
pub struct FnCompletionCallback<F>(pub F);
impl<F> CompletionCallback for FnCompletionCallback<F>
where
F: FnOnce(TransactionEnd) -> CompletionDelivery + Send + 'static,
{
fn call(self: Box<Self>, end: TransactionEnd) -> CompletionDelivery {
(self.0)(end)
}
}
pub struct TransactionRequest {
pub channel_id: ChannelId,
pub session_id: Option<SessionId>,
pub input: CanonicalInput,
pub session_config: Option<SessionConfig>,
pub invocation_config: InvocationConfig,
pub tools: Vec<ToolId>,
pub events: Arc<dyn TransactionEventSink>,
pub completion: Box<dyn CompletionCallback>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct AdmissionReceipt {
pub transaction_id: TransactionId,
pub session_id: Option<SessionId>,
}
#[derive(Clone, Debug, PartialEq, Eq, Hash)]
pub enum TransactionSelector {
Transaction(TransactionId),
Session(SessionKey),
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum TerminationMode {
Cancel {
reason: CancellationReason,
},
ForceTerminate {
reason: TerminationReason,
},
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct CancellationReason {
pub code: CancellationReasonCode,
pub detail: Option<SafeDiagnostic>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum CancellationReasonCode {
CallerRequested,
RuntimeShutdown,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct TerminationReason {
pub code: TerminationReasonCode,
pub detail: Option<SafeDiagnostic>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum TerminationReasonCode {
CallerRequested,
CancellationGraceExpired,
RuntimeShutdown,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum TerminationDisposition {
Accepted,
AlreadyRequested,
AlreadyTerminal,
NotFound,
}
pub type Shutdown = Pin<Box<dyn Future<Output = ShutdownDisposition> + Send + 'static>>;
#[derive(Clone, Debug, PartialEq, Eq, Default)]
pub struct ShutdownDisposition {
pub normally_finalized: u64,
pub supervisor_finalized: u64,
pub callback_failed: u64,
pub callback_aborted: u64,
pub invariant_failed: u64,
}
pub trait TransactionRuntime: Send + Sync {
fn submit(&self, request: TransactionRequest) -> Result<AdmissionReceipt, AdmissionError>;
fn terminate(
&self,
selector: TransactionSelector,
mode: TerminationMode,
) -> TerminationDisposition;
fn shutdown(&self, deadline: Duration) -> Shutdown;
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct TransactionEvent {
pub transaction_id: TransactionId,
pub channel_id: ChannelId,
pub session_id: SessionId,
pub sequence: u64,
pub payload: TransactionEventPayload,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub enum TransactionEventPayload {
SessionEstablished {
external_session_id: crate::id::ExternalSessionId,
},
CanonicalUnit(CanonicalUnitEvent),
ToolLifecycle(ToolLifecycleEvent),
Diagnostic(TransactionDiagnostic),
Ended(TransactionEnd),
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct TransactionDiagnostic {
pub diagnostic: SafeDiagnostic,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
pub struct TransactionEnd {
pub transaction_id: TransactionId,
pub session_id: Option<SessionId>,
pub channel_id: ChannelId,
pub kind: TransactionEndKind,
pub prior_terminal_cause: Option<TransactionEndKind>,
pub event_delivery: EventDeliveryOutcome,
pub emitted_events: u64,
pub usage: TransactionUsage,
pub diagnostics: Vec<TransactionDiagnostic>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum TransactionEndKind {
Completed,
ContinuationRequired,
Cancelled,
Terminated,
RuntimeShutdown,
DeadlineExceeded,
ChannelOpenFailed,
EncodingFailed,
ConnectorFailed,
InterpretationFailed,
ToolExchangeFailed,
EventDeliveryFailed,
LimitExceeded,
InvariantFailed,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum EventDeliveryOutcome {
Accepted,
Failed,
}
#[derive(Clone, Debug, PartialEq, Eq, Serialize, Deserialize, Default)]
pub struct TransactionUsage {
pub provider_input_tokens: Option<u64>,
pub provider_output_tokens: Option<u64>,
pub provider_exchanges: u32,
pub tools_started: u32,
pub tools_completed: u32,
}
#[derive(Clone, Debug, Error, PartialEq, Eq)]
pub enum EventDeliveryError {
#[error("event delivery failed")]
Failed,
#[error("event delivery deadline exceeded")]
DeadlineExceeded,
}
#[derive(Clone, Debug, Error, PartialEq, Eq)]
pub enum CompletionDeliveryError {
#[error("completion callback failed")]
Failed,
#[error("completion callback deadline exceeded")]
DeadlineExceeded,
}
#[derive(Clone, Debug, Error, PartialEq, Eq)]
#[error("{kind:?}: {message}")]
pub struct AdmissionError {
pub kind: AdmissionErrorKind,
pub message: String,
}
impl AdmissionError {
pub fn new(kind: AdmissionErrorKind, message: impl Into<String>) -> Self {
Self {
kind,
message: message.into(),
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum AdmissionErrorKind {
RuntimeShuttingDown,
UnknownChannel,
SessionAlreadyActive,
UnknownTool,
DuplicateTool,
InvalidInput,
InvalidConfiguration,
CapabilityMismatch,
CapacityExceeded,
SpawnFailed,
}
#[cfg(test)]
mod tests {
use super::*;
use crate::input::user_text_input;
#[test]
fn end_kind_round_trip() {
let kind = TransactionEndKind::Completed;
let json = serde_json::to_string(&kind).unwrap();
let back: TransactionEndKind = serde_json::from_str(&json).unwrap();
assert_eq!(kind, back);
}
#[tokio::test]
async fn sink_adapters_return_futures() {
let sink = FnEventSink(|_e| Box::pin(async { Ok(()) }) as EventDelivery);
let events: Arc<dyn TransactionEventSink> = Arc::new(sink);
let end = TransactionEnd {
transaction_id: TransactionId::generate(),
session_id: None,
channel_id: ChannelId::try_new("ch").unwrap(),
kind: TransactionEndKind::Completed,
prior_terminal_cause: None,
event_delivery: EventDeliveryOutcome::Accepted,
emitted_events: 1,
usage: TransactionUsage::default(),
diagnostics: vec![],
};
let ev = TransactionEvent {
transaction_id: end.transaction_id,
channel_id: end.channel_id.clone(),
session_id: SessionId::try_new("s").unwrap(),
sequence: 1,
payload: TransactionEventPayload::Ended(end.clone()),
};
events.deliver(ev).await.unwrap();
let cb: Box<dyn CompletionCallback> = Box::new(FnCompletionCallback(|_e| {
Box::pin(async { Ok(()) }) as CompletionDelivery
}));
cb.call(end).await.unwrap();
let _input = user_text_input("hello").unwrap();
}
}