use std::collections::BTreeMap;
use std::sync::Arc;
use salvor_core::{
Budget, Emitted, EventEnvelope, ModelReply, Outcome, ReplayCursor, RunId, SequenceNumber,
TokenUsage,
};
use salvor_llm::{Client, MessageAccumulator, MessageRequest, MessageResponse, StreamEvent};
use salvor_store::EventStore;
use salvor_tools::{DynTool, RetryPolicy, Suspension, ToolCtx, ToolError, ToolOutcome};
use serde_json::Value;
use time::OffsetDateTime;
use uuid::Uuid;
use crate::error::RuntimeError;
use crate::hash::hash_value;
use crate::labels::validate_labels;
use crate::model::{response_value, usage_of};
use crate::wire::{
ToolFailure, decode_failure, decode_suspension, encode_failure, encode_suspension,
};
pub type ClockFn = Arc<dyn Fn() -> OffsetDateTime + Send + Sync>;
pub type RandomFn = Arc<dyn Fn() -> u64 + Send + Sync>;
pub const MAX_TOOL_ATTEMPTS: u32 = 3;
#[derive(Debug, Clone)]
pub struct ModelTurn {
pub response: MessageResponse,
pub usage: TokenUsage,
}
#[derive(Debug, Clone)]
pub enum ToolCallResult {
Output(Value),
Failed(ToolFailure),
Suspended(Suspension),
}
#[derive(Debug, Clone)]
pub enum Resumption {
Resumed(Value),
Parked,
}
pub struct RunCtx {
cursor: ReplayCursor,
store: Arc<dyn EventStore>,
run_id: RunId,
clock: ClockFn,
random: RandomFn,
resume_input: Option<Value>,
record_prompts: bool,
labels: Option<BTreeMap<String, String>>,
}
impl RunCtx {
pub fn new(
store: Arc<dyn EventStore>,
run_id: RunId,
log: Vec<EventEnvelope>,
) -> Result<Self, RuntimeError> {
Self::with_hooks(
store,
run_id,
log,
Arc::new(OffsetDateTime::now_utc),
Arc::new(os_random),
)
}
pub fn with_hooks(
store: Arc<dyn EventStore>,
run_id: RunId,
log: Vec<EventEnvelope>,
clock: ClockFn,
random: RandomFn,
) -> Result<Self, RuntimeError> {
let cursor = ReplayCursor::new(log)?;
Ok(Self {
cursor,
store,
run_id,
clock,
random,
resume_input: None,
record_prompts: false,
labels: None,
})
}
#[must_use]
pub fn with_record_prompts(mut self, record_prompts: bool) -> Self {
self.record_prompts = record_prompts;
self
}
#[must_use]
pub fn with_labels(mut self, labels: BTreeMap<String, String>) -> Self {
self.labels = Some(labels);
self
}
pub fn set_resume_input(&mut self, input: Value) {
self.resume_input = Some(input);
}
#[must_use]
pub fn run_id(&self) -> RunId {
self.run_id
}
#[must_use]
pub fn is_replaying(&self) -> bool {
self.cursor.is_replaying()
}
#[must_use]
pub fn next_seq(&self) -> SequenceNumber {
self.cursor.next_seq()
}
pub async fn begin(
&mut self,
agent_def_hash: &str,
input: &Value,
) -> Result<Value, RuntimeError> {
match self.cursor.begin(agent_def_hash, self.labels.clone())? {
Outcome::Replayed(recorded) => Ok(recorded),
Outcome::Live(permit) => {
if let Some(labels) = &self.labels {
validate_labels(labels).map_err(RuntimeError::InvalidLabels)?;
}
let emitted = permit.record(input.clone());
persist(self.store.as_ref(), self.run_id, &self.clock, &emitted).await?;
Ok(input.clone())
}
}
}
pub async fn begin_graph(
&mut self,
graph_hash: &str,
input: &Value,
) -> Result<Value, RuntimeError> {
match self
.cursor
.begin_graph(graph_hash, self.labels.clone(), None)?
{
Outcome::Replayed(recorded) => Ok(recorded),
Outcome::Live(permit) => {
if let Some(labels) = &self.labels {
validate_labels(labels).map_err(RuntimeError::InvalidLabels)?;
}
let emitted = permit.record(input.clone());
persist(self.store.as_ref(), self.run_id, &self.clock, &emitted).await?;
Ok(input.clone())
}
}
}
pub async fn node_entered(&mut self, node: &str) -> Result<(), RuntimeError> {
match self.cursor.node_entered(node)? {
Outcome::Replayed(()) => Ok(()),
Outcome::Live(emitted) => {
persist(self.store.as_ref(), self.run_id, &self.clock, &emitted).await
}
}
}
pub async fn node_exited(&mut self, node: &str) -> Result<(), RuntimeError> {
match self.cursor.node_exited(node)? {
Outcome::Replayed(()) => Ok(()),
Outcome::Live(emitted) => {
persist(self.store.as_ref(), self.run_id, &self.clock, &emitted).await
}
}
}
pub async fn node_skipped(&mut self, node: &str, reason: &str) -> Result<(), RuntimeError> {
match self.cursor.node_skipped(node, reason)? {
Outcome::Replayed(()) => Ok(()),
Outcome::Live(emitted) => {
persist(self.store.as_ref(), self.run_id, &self.clock, &emitted).await
}
}
}
pub async fn branch_taken(&mut self, node: &str, case: &str) -> Result<(), RuntimeError> {
match self.cursor.branch_taken(node, case)? {
Outcome::Replayed(()) => Ok(()),
Outcome::Live(emitted) => {
persist(self.store.as_ref(), self.run_id, &self.clock, &emitted).await
}
}
}
pub async fn map_fanned_out(&mut self, node: &str, items: &Value) -> Result<(), RuntimeError> {
match self.cursor.map_fanned_out(node, items)? {
Outcome::Replayed(()) => Ok(()),
Outcome::Live(emitted) => {
persist(self.store.as_ref(), self.run_id, &self.clock, &emitted).await
}
}
}
pub async fn map_iteration_started(
&mut self,
node: &str,
index: u64,
child_run: &str,
) -> Result<(), RuntimeError> {
match self.cursor.map_iteration_started(node, index, child_run)? {
Outcome::Replayed(()) => Ok(()),
Outcome::Live(emitted) => {
persist(self.store.as_ref(), self.run_id, &self.clock, &emitted).await
}
}
}
pub async fn map_iteration_joined(
&mut self,
node: &str,
index: u64,
) -> Result<(), RuntimeError> {
match self.cursor.map_iteration_joined(node, index)? {
Outcome::Replayed(()) => Ok(()),
Outcome::Live(emitted) => {
persist(self.store.as_ref(), self.run_id, &self.clock, &emitted).await
}
}
}
pub async fn now(&mut self) -> Result<OffsetDateTime, RuntimeError> {
match self.cursor.now()? {
Outcome::Replayed(instant) => Ok(instant),
Outcome::Live(permit) => {
let instant = (self.clock)();
let emitted = permit.record(instant);
persist(self.store.as_ref(), self.run_id, &self.clock, &emitted).await?;
Ok(instant)
}
}
}
pub async fn random(&mut self) -> Result<u64, RuntimeError> {
match self.cursor.random()? {
Outcome::Replayed(bits) => Ok(bits),
Outcome::Live(permit) => {
let bits = (self.random)();
let emitted = permit.record(bits);
persist(self.store.as_ref(), self.run_id, &self.clock, &emitted).await?;
Ok(bits)
}
}
}
pub async fn model_call(
&mut self,
client: &Client,
request: &MessageRequest,
) -> Result<ModelTurn, RuntimeError> {
let request_value = serde_json::to_value(request).map_err(RuntimeError::RequestEncode)?;
let request_hash = hash_value(&request_value);
let request_body = if self.record_prompts {
Some(request_value)
} else {
None
};
match self.cursor.model_call(&request_hash, request_body)? {
Outcome::Replayed(ModelReply { response, usage }) => {
let response = serde_json::from_value(response)
.map_err(RuntimeError::RecordedResponseDecode)?;
Ok(ModelTurn { response, usage })
}
Outcome::Live(permit) => {
if let Some(intent) = permit.intent().cloned() {
persist(self.store.as_ref(), self.run_id, &self.clock, &intent).await?;
}
let response = client.send_message(request).await?;
let usage = usage_of(&response);
let completion = permit.record(response_value(&response), usage);
persist(self.store.as_ref(), self.run_id, &self.clock, &completion).await?;
Ok(ModelTurn { response, usage })
}
}
}
pub async fn model_call_streaming(
&mut self,
client: &Client,
request: &MessageRequest,
mut on_event: impl FnMut(&StreamEvent),
) -> Result<ModelTurn, RuntimeError> {
let request_value = serde_json::to_value(request).map_err(RuntimeError::RequestEncode)?;
let request_hash = hash_value(&request_value);
let request_body = if self.record_prompts {
Some(request_value)
} else {
None
};
match self.cursor.model_call(&request_hash, request_body)? {
Outcome::Replayed(ModelReply { response, usage }) => {
let response = serde_json::from_value(response)
.map_err(RuntimeError::RecordedResponseDecode)?;
Ok(ModelTurn { response, usage })
}
Outcome::Live(permit) => {
if let Some(intent) = permit.intent().cloned() {
persist(self.store.as_ref(), self.run_id, &self.clock, &intent).await?;
}
let mut stream = client.stream_message(request).await?;
let mut accumulator = MessageAccumulator::new();
while let Some(event) = stream.next_event().await {
let event = event?;
on_event(&event);
accumulator.apply(&event)?;
}
let response = accumulator.into_message()?;
let usage = usage_of(&response);
let completion = permit.record(response_value(&response), usage);
persist(self.store.as_ref(), self.run_id, &self.clock, &completion).await?;
Ok(ModelTurn { response, usage })
}
}
}
pub async fn tool_call(
&mut self,
tool: &dyn DynTool,
input: &Value,
idempotency_key: Option<&str>,
) -> Result<ToolCallResult, RuntimeError> {
let effect = tool.effect();
match self
.cursor
.tool_call(tool.name(), input, effect, idempotency_key)?
{
Outcome::Replayed(output) => Ok(decode_tool_output(output)),
Outcome::Live(permit) => {
if let Some(intent) = permit.intent().cloned() {
persist(self.store.as_ref(), self.run_id, &self.clock, &intent).await?;
}
let key = permit.idempotency_key().map(ToOwned::to_owned);
let tool_ctx = ToolCtx::new(key);
let policy = RetryPolicy::for_effect(effect);
let mut attempts: u32 = 0;
let outcome = loop {
attempts += 1;
match tool.call_json(&tool_ctx, input.clone()).await {
Ok(outcome) => break Ok(outcome),
Err(error) => {
let may_retry = matches!(error, ToolError::Handler { .. })
&& policy.allows_retry()
&& attempts < MAX_TOOL_ATTEMPTS;
if may_retry {
continue;
}
break Err(error);
}
}
};
let (output, result) = match outcome {
Ok(ToolOutcome::Output(value)) => {
(value.clone(), ToolCallResult::Output(value))
}
Ok(ToolOutcome::Suspend(suspension)) => (
encode_suspension(&suspension),
ToolCallResult::Suspended(suspension),
),
Err(error) => {
let failure = ToolFailure::from_error(&error, attempts);
(encode_failure(&failure), ToolCallResult::Failed(failure))
}
};
let completion = permit.record(output);
persist(self.store.as_ref(), self.run_id, &self.clock, &completion).await?;
Ok(result)
}
}
}
pub async fn suspend(
&mut self,
reason: &str,
input_schema: &Value,
) -> Result<(), RuntimeError> {
match self.cursor.suspend(reason, input_schema)? {
Outcome::Replayed(()) => Ok(()),
Outcome::Live(emitted) => {
persist(self.store.as_ref(), self.run_id, &self.clock, &emitted).await
}
}
}
pub async fn await_resume(&mut self) -> Result<Resumption, RuntimeError> {
match self.cursor.await_resume()? {
Outcome::Replayed(input) => Ok(Resumption::Resumed(input)),
Outcome::Live(parked) => match self.resume_input.take() {
Some(input) => {
let emitted = parked.resume(input.clone());
persist(self.store.as_ref(), self.run_id, &self.clock, &emitted).await?;
Ok(Resumption::Resumed(input))
}
None => Ok(Resumption::Parked),
},
}
}
pub async fn budget_exceeded(
&mut self,
budget: Budget,
observed: f64,
) -> Result<(), RuntimeError> {
match self.cursor.budget_exceeded(budget, observed)? {
Outcome::Replayed(()) => Ok(()),
Outcome::Live(emitted) => {
persist(self.store.as_ref(), self.run_id, &self.clock, &emitted).await
}
}
}
pub async fn complete_run(&mut self, output: &Value) -> Result<(), RuntimeError> {
match self.cursor.complete_run(output)? {
Outcome::Replayed(()) => Ok(()),
Outcome::Live(emitted) => {
persist(self.store.as_ref(), self.run_id, &self.clock, &emitted).await
}
}
}
pub async fn fail_run(&mut self, error: &str) -> Result<(), RuntimeError> {
match self.cursor.fail_run(error)? {
Outcome::Replayed(()) => Ok(()),
Outcome::Live(emitted) => {
persist(self.store.as_ref(), self.run_id, &self.clock, &emitted).await
}
}
}
}
async fn persist(
store: &dyn EventStore,
run_id: RunId,
clock: &ClockFn,
emitted: &Emitted,
) -> Result<(), RuntimeError> {
let envelope = EventEnvelope::new(run_id, emitted.seq, (clock)(), emitted.event.clone());
store.append(&envelope).await?;
crate::progress::emit_step(run_id, envelope.seq, &envelope.event);
Ok(())
}
fn decode_tool_output(output: Value) -> ToolCallResult {
if let Some(suspension) = decode_suspension(&output) {
return ToolCallResult::Suspended(suspension);
}
if let Some(failure) = decode_failure(&output) {
return ToolCallResult::Failed(failure);
}
ToolCallResult::Output(output)
}
pub(crate) fn os_random() -> u64 {
let bits = Uuid::new_v4().as_u128();
(bits as u64) ^ ((bits >> 64) as u64)
}