use std::convert::Infallible;
use axum::Json;
use axum::body::Bytes;
use axum::extract::{Path, Query, State};
use axum::http::header::ACCEPT;
use axum::http::{HeaderMap, StatusCode};
use axum::response::sse::{Event as SseEvent, KeepAlive, Sse};
use axum::response::{IntoResponse, Response};
use salvor_core::{Effect, Event, EventEnvelope, LogValidator, RunId, SequenceNumber, TokenUsage};
use salvor_llm::{ContentDelta, MessageAccumulator, StreamEvent};
use salvor_runtime::{RuntimeError, hash_value, response_value, usage_of, validate_labels};
use salvor_tools::{ToolCtx, ToolOutcome};
use serde::Deserialize;
use serde_json::{Value, json};
use time::format_description::well_known::Rfc3339;
use tokio::sync::mpsc;
use tokio_stream::wrappers::ReceiverStream;
use uuid::Uuid;
use crate::error::ApiError;
use crate::executor::{ModelExecutor, ModelStream};
use crate::state::{AppState, ClientRunLease};
use std::sync::Arc;
const DRIVE_TOKEN_HEADER: &str = "x-drive-token";
const MAX_EVENTS_BODY: usize = 8 * 1024 * 1024;
const MAX_EVENTS_PER_BATCH: usize = 1024;
#[derive(Debug, Deserialize)]
struct OpenRequest {
#[serde(default)]
agent: Option<String>,
#[serde(default)]
input: Value,
#[serde(default)]
run_id: Option<String>,
#[serde(default)]
record_prompts: bool,
}
#[derive(Debug, Deserialize)]
struct AppendRequest {
events: Vec<EventEnvelope>,
}
#[derive(Debug, Default, Deserialize)]
pub struct LogQuery {
#[serde(default)]
from_seq: Option<u64>,
}
#[derive(Debug, Deserialize)]
struct ModelStepRequest {
seq: u64,
request: Value,
}
#[derive(Debug, Default, Deserialize)]
pub struct ModelStepQuery {
#[serde(default)]
stream: Option<String>,
}
#[derive(Debug, Deserialize)]
struct ToolStepRequest {
seq: u64,
tool: String,
input: Value,
#[serde(default)]
idempotency_key: Option<String>,
#[serde(default)]
#[allow(dead_code)]
effect: Option<Effect>,
}
#[derive(Debug, Deserialize)]
struct ResolveRequest {
output: Value,
}
pub async fn open(
State(state): State<AppState>,
body: Bytes,
) -> Result<impl IntoResponse, ApiError> {
let request: OpenRequest = parse_body(&body)?;
let _ = (&request.agent, &request.input);
let run_id = match &request.run_id {
Some(text) => parse_run_id(text)?,
None => RunId::new(),
};
if state.is_client_run(run_id) {
let log = state.store().read_log(run_id).await.map_err(store_error)?;
let drive_token = state.lease_client_run(run_id, request.record_prompts);
return Ok((StatusCode::OK, Json(open_body(run_id, &drive_token, &log))));
}
let log = state.store().read_log(run_id).await.map_err(store_error)?;
if !log.is_empty() {
return Err(ApiError::RunExists(format!(
"run {} already has recorded history and is not a client-driven run on this server; \
it cannot be opened for client-driven runs",
run_id.as_uuid()
)));
}
let drive_token = state.lease_client_run(run_id, request.record_prompts);
Ok((
StatusCode::CREATED,
Json(open_body(run_id, &drive_token, &[])),
))
}
pub async fn get_log(
State(state): State<AppState>,
Path(run_id_text): Path<String>,
Query(query): Query<LogQuery>,
) -> Result<impl IntoResponse, ApiError> {
let run_id = parse_run_id(&run_id_text)?;
if !state.is_client_run(run_id) {
return Err(unknown_client_run(run_id));
}
let mut log = state.store().read_log(run_id).await.map_err(store_error)?;
if let Some(from) = query.from_seq {
log.retain(|env| env.seq.get() >= from);
}
Ok(Json(json!({ "log": log })))
}
pub async fn append(
State(state): State<AppState>,
Path(run_id_text): Path<String>,
headers: HeaderMap,
body: Bytes,
) -> Result<impl IntoResponse, ApiError> {
let run_id = parse_run_id(&run_id_text)?;
authorize_drive(&state, run_id, &headers)?;
if body.len() > MAX_EVENTS_BODY {
return Err(ApiError::PayloadTooLarge(format!(
"append body is {} bytes, over the {MAX_EVENTS_BODY}-byte cap",
body.len()
)));
}
let request: AppendRequest = parse_body(&body)?;
if request.events.len() > MAX_EVENTS_PER_BATCH {
return Err(ApiError::PayloadTooLarge(format!(
"append batch carries {} events, over the {MAX_EVENTS_PER_BATCH} cap",
request.events.len()
)));
}
let stored = state.store().read_log(run_id).await.map_err(store_error)?;
let mut validator = LogValidator::new(stored);
let mut appended: Vec<u64> = Vec::with_capacity(request.events.len());
let mut to_append: Vec<EventEnvelope> = Vec::new();
for mut candidate in request.events {
if candidate.run_id != run_id {
return Err(ApiError::Divergence(format!(
"event names run {} but the path is run {}",
candidate.run_id.as_uuid(),
run_id.as_uuid()
)));
}
reject_side_effecting_kind(&candidate)?;
if let Event::RunStarted {
labels: Some(labels),
..
} = &candidate.event
{
validate_labels(labels).map_err(ApiError::BadRequest)?;
}
let next_seq = validator.next_seq();
if candidate.seq < next_seq {
let index = candidate.seq.get() as usize;
let recorded = &validator.log()[index];
candidate.recorded_at = recorded.recorded_at;
if *recorded == candidate {
appended.push(candidate.seq.get());
continue;
}
return Err(ApiError::Divergence(format!(
"different bytes submitted at the already-recorded seq {}",
candidate.seq.get()
)));
}
candidate.recorded_at = state.now();
validator
.push(candidate.clone())
.map_err(|error| ApiError::Divergence(error.to_string()))?;
appended.push(candidate.seq.get());
to_append.push(candidate);
}
for envelope in &to_append {
state.store().append(envelope).await.map_err(append_error)?;
}
Ok((StatusCode::OK, Json(json!({ "appended": appended }))))
}
pub async fn model_step(
State(state): State<AppState>,
Path(run_id_text): Path<String>,
Query(query): Query<ModelStepQuery>,
headers: HeaderMap,
body: Bytes,
) -> Result<Response, ApiError> {
let run_id = parse_run_id(&run_id_text)?;
let lease = authorize_drive(&state, run_id, &headers)?;
if body.len() > MAX_EVENTS_BODY {
return Err(ApiError::PayloadTooLarge(format!(
"model-step body is {} bytes, over the {MAX_EVENTS_BODY}-byte cap",
body.len()
)));
}
let ModelStepRequest { seq, request } = parse_body(&body)?;
let request_hash = hash_value(&request);
let log = state.store().read_log(run_id).await.map_err(store_error)?;
let plan = plan_model_step(&log, seq, &request_hash)?;
let streaming = wants_stream(&headers, &query);
match plan {
ModelStepPlan::Replay { response, usage } => {
if streaming {
Ok(single_complete_stream(&response, usage))
} else {
Ok(completion_body(&response, usage).into_response())
}
}
ModelStepPlan::Perform { append_intent } => {
let executor = state.model_executor().ok_or_else(|| {
ApiError::ModelExecutorUnavailable(
"this server has no model executor wired, so it cannot perform a model step"
.to_owned(),
)
})?;
if append_intent {
let request_body = lease.record_prompts.then(|| request.clone());
let intent = EventEnvelope::new(
run_id,
SequenceNumber::new(seq),
state.now(),
Event::ModelCallRequested {
seq: SequenceNumber::new(seq),
request_hash: request_hash.clone(),
request_body,
},
);
let mut validator = LogValidator::new(log);
validator
.push(intent.clone())
.map_err(|error| ApiError::Divergence(error.to_string()))?;
state.store().append(&intent).await.map_err(append_error)?;
}
if streaming {
perform_streaming(state, run_id, seq, request, executor).await
} else {
perform_unary(&state, run_id, seq, request, executor.as_ref()).await
}
}
}
}
enum ModelStepPlan {
Replay {
response: Value,
usage: TokenUsage,
},
Perform {
append_intent: bool,
},
}
fn plan_model_step(
log: &[EventEnvelope],
seq: u64,
request_hash: &str,
) -> Result<ModelStepPlan, ApiError> {
let next = log.len() as u64;
if seq == next {
return Ok(ModelStepPlan::Perform {
append_intent: true,
});
}
if seq > next {
return Err(ApiError::Divergence(format!(
"model-step seq {seq} is beyond the log end {next}"
)));
}
let recorded = &log[seq as usize];
let Event::ModelCallRequested {
request_hash: recorded_hash,
..
} = &recorded.event
else {
return Err(ApiError::Divergence(format!(
"seq {seq} already holds a non-model event; it is not a model-step position"
)));
};
if recorded_hash != request_hash {
return Err(ApiError::Divergence(format!(
"model-step at seq {seq} carries a request hash that differs from the recorded intent"
)));
}
match log.get(seq as usize + 1) {
Some(next_env) => match &next_env.event {
Event::ModelCallCompleted {
seq: corr,
response,
usage,
} if corr.get() == seq => Ok(ModelStepPlan::Replay {
response: response.clone(),
usage: *usage,
}),
_ => Err(ApiError::Divergence(format!(
"the event after the intent at seq {seq} is not its completion"
))),
},
None => Ok(ModelStepPlan::Perform {
append_intent: false,
}),
}
}
async fn perform_unary(
state: &AppState,
run_id: RunId,
seq: u64,
request: Value,
executor: &dyn ModelExecutor,
) -> Result<Response, ApiError> {
let response = executor
.execute(request)
.await
.map_err(ApiError::ModelExecution)?;
let usage = usage_of(&response);
let response_value = response_value(&response);
append_completion(state, run_id, seq, &response_value, usage).await?;
Ok(completion_body(&response_value, usage).into_response())
}
async fn perform_streaming(
state: AppState,
run_id: RunId,
seq: u64,
request: Value,
executor: Arc<dyn ModelExecutor>,
) -> Result<Response, ApiError> {
let stream = executor
.open_stream(request)
.await
.map_err(ApiError::ModelExecution)?;
let (tx, rx) = mpsc::channel::<Result<SseEvent, Infallible>>(64);
tokio::spawn(drive_model_stream(state, run_id, seq, stream, tx));
Ok(Sse::new(ReceiverStream::new(rx))
.keep_alive(KeepAlive::default())
.into_response())
}
async fn drive_model_stream(
state: AppState,
run_id: RunId,
seq: u64,
mut stream: Box<dyn ModelStream>,
tx: mpsc::Sender<Result<SseEvent, Infallible>>,
) {
let mut accumulator = MessageAccumulator::new();
loop {
match stream.next_event().await {
Some(Ok(event)) => {
if let Err(error) = accumulator.apply(&event) {
let _ = tx.send(Ok(error_frame(&error.to_string()))).await;
return;
}
if let Some(frame) = ticker_frame(&event)
&& tx
.send(Ok(SseEvent::default()
.event("delta")
.data(frame.to_string())))
.await
.is_err()
{
return;
}
}
Some(Err(message)) => {
let _ = tx.send(Ok(error_frame(&message))).await;
return;
}
None => break,
}
}
let response = match accumulator.into_message() {
Ok(response) => response,
Err(error) => {
let _ = tx.send(Ok(error_frame(&error.to_string()))).await;
return;
}
};
let usage = usage_of(&response);
let response_value = response_value(&response);
if append_completion(&state, run_id, seq, &response_value, usage)
.await
.is_err()
{
let _ = tx
.send(Ok(error_frame("recording the model completion failed")))
.await;
return;
}
let complete = completion_json(&response_value, usage);
let _ = tx
.send(Ok(SseEvent::default()
.event("complete")
.data(complete.to_string())))
.await;
}
async fn append_completion(
state: &AppState,
run_id: RunId,
seq: u64,
response: &Value,
usage: TokenUsage,
) -> Result<(), ApiError> {
let completion = EventEnvelope::new(
run_id,
SequenceNumber::new(seq + 1),
state.now(),
Event::ModelCallCompleted {
seq: SequenceNumber::new(seq),
response: response.clone(),
usage,
},
);
let log = state.store().read_log(run_id).await.map_err(store_error)?;
let mut validator = LogValidator::new(log);
validator
.push(completion.clone())
.map_err(|error| ApiError::Divergence(error.to_string()))?;
state
.store()
.append(&completion)
.await
.map_err(append_error)
}
fn wants_stream(headers: &HeaderMap, query: &ModelStepQuery) -> bool {
if let Some(flag) = &query.stream
&& (flag == "1" || flag == "true")
{
return true;
}
headers
.get(ACCEPT)
.and_then(|value| value.to_str().ok())
.is_some_and(|accept| accept.contains("text/event-stream"))
}
fn ticker_frame(event: &StreamEvent) -> Option<Value> {
match event {
StreamEvent::ContentBlockDelta { index, delta } => match delta {
ContentDelta::Text { text } => {
Some(json!({ "type": "text_delta", "index": index, "text": text }))
}
ContentDelta::Thinking { thinking } => {
Some(json!({ "type": "thinking_delta", "index": index, "thinking": thinking }))
}
_ => None,
},
StreamEvent::MessageDelta { usage, .. } => {
Some(json!({ "type": "usage", "output_tokens": usage.output_tokens }))
}
_ => None,
}
}
fn single_complete_stream(response: &Value, usage: TokenUsage) -> Response {
let frame = SseEvent::default()
.event("complete")
.data(completion_json(response, usage).to_string());
Sse::new(tokio_stream::once(Ok::<_, Infallible>(frame)))
.keep_alive(KeepAlive::default())
.into_response()
}
fn completion_json(response: &Value, usage: TokenUsage) -> Value {
json!({ "response": response, "usage": usage })
}
fn completion_body(response: &Value, usage: TokenUsage) -> Json<Value> {
Json(completion_json(response, usage))
}
fn error_frame(message: &str) -> SseEvent {
SseEvent::default()
.event("error")
.data(json!({ "message": message }).to_string())
}
fn open_body(run_id: RunId, drive_token: &str, log: &[EventEnvelope]) -> Value {
json!({
"run": run_id.as_uuid().to_string(),
"drive_token": drive_token,
"log": log,
})
}
pub async fn tool_step(
State(state): State<AppState>,
Path(run_id_text): Path<String>,
headers: HeaderMap,
body: Bytes,
) -> Result<Json<Value>, ApiError> {
let run_id = parse_run_id(&run_id_text)?;
authorize_drive(&state, run_id, &headers)?;
if body.len() > MAX_EVENTS_BODY {
return Err(ApiError::PayloadTooLarge(format!(
"tool-step body is {} bytes, over the {MAX_EVENTS_BODY}-byte cap",
body.len()
)));
}
let request: ToolStepRequest = parse_body(&body)?;
let registry = state.tool_registry().ok_or_else(|| {
ApiError::ToolRegistryUnavailable(
"this server has no tool registry wired, so it cannot perform a tool step".to_owned(),
)
})?;
let tool = registry.get(&request.tool).ok_or_else(|| {
ApiError::UnknownTool(format!(
"no tool named `{}` is registered on this server",
request.tool
))
})?;
let effect = tool.effect();
let ToolStepRequest {
seq,
tool: tool_name,
input,
idempotency_key,
effect: _,
} = request;
let log = state.store().read_log(run_id).await.map_err(store_error)?;
let plan = plan_tool_step(
&log,
seq,
&tool_name,
&input,
effect,
idempotency_key.as_deref(),
)?;
match plan {
ToolStepPlan::Replay { output } => Ok(tool_output_body(&output)),
ToolStepPlan::Reconcile { intent } => Err(ApiError::NeedsReconciliation {
message: format!(
"run {} needs reconciliation: a write was recorded but never completed, so it \
may or may not have taken effect. Verify externally, then resolve it",
run_id.as_uuid()
),
intent,
}),
ToolStepPlan::Perform {
append_intent,
exec_key,
} => {
if append_intent {
let intent = EventEnvelope::new(
run_id,
SequenceNumber::new(seq),
state.now(),
Event::ToolCallRequested {
seq: SequenceNumber::new(seq),
tool: tool_name.clone(),
input: input.clone(),
effect,
idempotency_key: exec_key.clone(),
},
);
let mut validator = LogValidator::new(log);
validator
.push(intent.clone())
.map_err(|error| ApiError::Divergence(error.to_string()))?;
state.store().append(&intent).await.map_err(append_error)?;
}
let ctx = ToolCtx::new(exec_key);
let outcome = tool
.call_json(&ctx, input)
.await
.map_err(|error| ApiError::ToolExecution(error.to_string()))?;
let output = match outcome {
ToolOutcome::Output(value) => value,
ToolOutcome::Suspend(_) => {
return Err(ApiError::ToolExecution(format!(
"tool `{tool_name}` suspended, which a server-performed tool step does \
not support; no completion recorded"
)));
}
};
append_tool_completion(&state, run_id, seq, &output).await?;
Ok(tool_output_body(&output))
}
}
}
enum ToolStepPlan {
Replay {
output: Value,
},
Reconcile {
intent: Value,
},
Perform {
append_intent: bool,
exec_key: Option<String>,
},
}
fn plan_tool_step(
log: &[EventEnvelope],
seq: u64,
tool: &str,
input: &Value,
effect: Effect,
idempotency_key: Option<&str>,
) -> Result<ToolStepPlan, ApiError> {
let next = log.len() as u64;
if seq == next {
return Ok(ToolStepPlan::Perform {
append_intent: true,
exec_key: idempotency_key.map(ToOwned::to_owned),
});
}
if seq > next {
return Err(ApiError::Divergence(format!(
"tool-step seq {seq} is beyond the log end {next}"
)));
}
let recorded = &log[seq as usize];
let Event::ToolCallRequested {
tool: recorded_tool,
input: recorded_input,
effect: recorded_effect,
idempotency_key: recorded_key,
..
} = &recorded.event
else {
return Err(ApiError::Divergence(format!(
"seq {seq} already holds a non-tool event; it is not a tool-step position"
)));
};
if recorded_tool != tool
|| recorded_input != input
|| *recorded_effect != effect
|| recorded_key.as_deref() != idempotency_key
{
return Err(ApiError::Divergence(format!(
"tool-step at seq {seq} diverges from the recorded intent (tool, input, effect, or key)"
)));
}
match log.get(seq as usize + 1) {
Some(next_env) => match &next_env.event {
Event::ToolCallCompleted { seq: corr, output } if corr.get() == seq => {
Ok(ToolStepPlan::Replay {
output: output.clone(),
})
}
_ => Err(ApiError::Divergence(format!(
"the event after the intent at seq {seq} is not its completion"
))),
},
None => match effect {
Effect::Write => Ok(ToolStepPlan::Reconcile {
intent: intent_evidence(recorded),
}),
Effect::Read | Effect::Idempotent => Ok(ToolStepPlan::Perform {
append_intent: false,
exec_key: recorded_key.clone(),
}),
},
}
}
fn intent_evidence(envelope: &EventEnvelope) -> Value {
let Event::ToolCallRequested {
seq,
tool,
input,
effect,
idempotency_key,
} = &envelope.event
else {
return Value::Null;
};
json!({
"kind": "tool",
"seq": seq.get(),
"tool": tool,
"input": input,
"effect": effect,
"idempotency_key": idempotency_key,
"recorded_at": envelope.recorded_at.format(&Rfc3339).unwrap_or_default(),
})
}
async fn append_tool_completion(
state: &AppState,
run_id: RunId,
seq: u64,
output: &Value,
) -> Result<(), ApiError> {
let completion = EventEnvelope::new(
run_id,
SequenceNumber::new(seq + 1),
state.now(),
Event::ToolCallCompleted {
seq: SequenceNumber::new(seq),
output: output.clone(),
},
);
let log = state.store().read_log(run_id).await.map_err(store_error)?;
let mut validator = LogValidator::new(log);
validator
.push(completion.clone())
.map_err(|error| ApiError::Divergence(error.to_string()))?;
state
.store()
.append(&completion)
.await
.map_err(append_error)
}
fn tool_output_body(output: &Value) -> Json<Value> {
Json(json!({ "output": output }))
}
pub async fn resolve(
State(state): State<AppState>,
Path(run_id_text): Path<String>,
headers: HeaderMap,
body: Bytes,
) -> Result<Json<Value>, ApiError> {
let run_id = parse_run_id(&run_id_text)?;
authorize_drive(&state, run_id, &headers)?;
let request: ResolveRequest = parse_body(&body)?;
match state.runtime().resolve(run_id, request.output).await {
Ok(_) => Ok(Json(json!({
"run": run_id.as_uuid().to_string(),
"resolved": true,
}))),
Err(RuntimeError::NotReconcilable { status, .. }) => Err(ApiError::WrongState(format!(
"run {} does not need reconciliation (status: {status}); there is no dangling write \
to resolve",
run_id.as_uuid()
))),
Err(error) => Err(ApiError::Internal(error.to_string())),
}
}
fn reject_side_effecting_kind(candidate: &EventEnvelope) -> Result<(), ApiError> {
use salvor_core::Event;
let kind = match &candidate.event {
Event::ModelCallRequested { .. } => "ModelCallRequested",
Event::ModelCallCompleted { .. } => "ModelCallCompleted",
Event::ToolCallRequested { .. } => "ToolCallRequested",
Event::ToolCallCompleted { .. } => "ToolCallCompleted",
_ => return Ok(()),
};
Err(ApiError::UnsupportedEventKind(format!(
"the generic append accepts control and context events only; `{kind}` is recorded through \
the model-step or tool-step endpoint"
)))
}
fn authorize_drive(
state: &AppState,
run_id: RunId,
headers: &HeaderMap,
) -> Result<ClientRunLease, ApiError> {
let lease = state
.client_run(run_id)
.ok_or_else(|| unknown_client_run(run_id))?;
let presented = headers
.get(DRIVE_TOKEN_HEADER)
.and_then(|value| value.to_str().ok());
match presented {
None => Err(ApiError::MissingDriveToken(format!(
"run {} requires a drive token in the `{DRIVE_TOKEN_HEADER}` header",
run_id.as_uuid()
))),
Some(token) if token != lease.drive_token => Err(ApiError::InvalidDriveToken(format!(
"the presented drive token is not the current lease for run {}",
run_id.as_uuid()
))),
Some(_) => {
state.touch_client_run(run_id);
Ok(lease)
}
}
}
fn parse_body<T: for<'de> Deserialize<'de>>(body: &Bytes) -> Result<T, ApiError> {
serde_json::from_slice(body)
.map_err(|error| ApiError::BadRequest(format!("request body is not valid JSON: {error}")))
}
fn parse_run_id(text: &str) -> Result<RunId, ApiError> {
Uuid::parse_str(text).map(RunId::from_uuid).map_err(|_| {
ApiError::BadRequest(format!("`{text}` is not a valid run id (expected a UUID)"))
})
}
fn unknown_client_run(run_id: RunId) -> ApiError {
ApiError::UnknownRun(format!(
"no client-driven run {} on this server; open it first",
run_id.as_uuid()
))
}
fn store_error(error: salvor_store::StoreError) -> ApiError {
ApiError::Internal(format!("store: {error}"))
}
fn append_error(error: salvor_store::StoreError) -> ApiError {
match error {
salvor_store::StoreError::Conflict { seq, .. } => ApiError::Divergence(format!(
"seq {} was taken by another writer before the append landed",
SequenceNumber::get(seq)
)),
other => ApiError::Internal(format!("store: {other}")),
}
}