use std::convert::Infallible;
use std::path::{Path, PathBuf};
use std::sync::{Arc, Mutex};
use axum::extract::{Query, State};
use axum::http::{HeaderMap, StatusCode};
use axum::response::sse::{Event, KeepAlive, Sse};
use axum::routing::{get, post};
use axum::{Json, Router};
use futures_util::stream::{self, Stream, StreamExt};
use serde::{Deserialize, Serialize};
use time::OffsetDateTime;
use tokio_stream::wrappers::BroadcastStream;
use crate::chat::turns::ChatTurn;
use crate::wave::channel::tagged_turn_json;
use crate::wave::journal::{Attribution, MessageOp, PendingMessage};
use crate::wave::registry::{process_alive, StoreObserver};
use crate::wave::runtime::{InboxItem, WaveRuntime};
use crate::wave::state::MindState;
use crate::wave::supervisor::SupervisorHandle;
use crate::wave::wire::{
AttachRequest, AttachResponse, ContextResponse, InFlightWorker, InboxFrame, PostDeltasRequest,
PostDeltasResponse, RESIDENT_TOKEN_FILE, RESIDENT_TOKEN_HEADER,
};
pub const ENDPOINT_FILE: &str = ".wave-endpoint";
#[derive(Debug, Clone)]
pub struct ResidentDoor {
token: String,
seat: Arc<Mutex<Option<u32>>>,
}
impl ResidentDoor {
pub fn new(token: impl Into<String>) -> Self {
Self {
token: token.into(),
seat: Arc::new(Mutex::new(None)),
}
}
pub fn token(&self) -> &str {
&self.token
}
pub fn seat_pid(&self) -> Option<u32> {
*self.seat.lock().expect("resident seat lock poisoned")
}
pub fn record_pid(&self, pid: u32) {
*self.seat.lock().expect("resident seat lock poisoned") = Some(pid);
}
pub fn clear_seat(&self) {
*self.seat.lock().expect("resident seat lock poisoned") = None;
}
fn authorize(&self, headers: &HeaderMap) -> Result<(), (StatusCode, String)> {
let presented = headers
.get(RESIDENT_TOKEN_HEADER)
.and_then(|value| value.to_str().ok())
.unwrap_or_default();
if presented == self.token {
return Ok(());
}
Err((
StatusCode::UNAUTHORIZED,
format!("missing or wrong {RESIDENT_TOKEN_HEADER}"),
))
}
}
pub fn generate_resident_token() -> String {
format!(
"{}{}",
uuid::Uuid::new_v4().simple(),
uuid::Uuid::new_v4().simple()
)
}
#[derive(Debug, Serialize)]
struct HealthBody {
status: String,
mind: Option<String>,
wave: String,
turns: usize,
workers: usize,
paused: bool,
uptime_seconds: i64,
}
#[derive(Debug, Serialize)]
struct ConversationBody {
turns: Vec<ChatTurn>,
}
#[derive(Debug, Deserialize)]
struct ConversationQuery {
limit: Option<usize>,
}
#[derive(Debug, Deserialize)]
struct PostMessage {
op: MessageOp,
text: String,
from: Option<Attribution>,
channel: Option<String>,
}
#[derive(Debug, Deserialize)]
struct PostChannel {
name: String,
run_id: String,
}
#[derive(Debug, Serialize)]
struct PostChannelResponse {
turn: Option<ChatTurn>,
}
#[derive(Debug, Deserialize)]
struct EventsQuery {
channel: Option<String>,
prefix: Option<String>,
inbox: Option<bool>,
}
#[derive(Debug, Serialize)]
struct MemoryBody {
content: String,
}
#[derive(Debug, Clone, Copy, Deserialize)]
#[serde(rename_all = "snake_case")]
enum MemoryOp {
Update,
Add,
}
#[derive(Debug, Deserialize)]
struct PostMemory {
op: MemoryOp,
content: String,
summary: Option<String>,
}
#[derive(Debug, Serialize)]
struct PostMemoryResponse {
summary: String,
}
#[derive(Debug, Serialize)]
struct PostMessageResponse {
turn: Option<ChatTurn>,
state: String,
}
#[derive(Clone)]
struct ServerState {
runtime: Arc<WaveRuntime>,
resident: ResidentDoor,
observer: Option<Arc<StoreObserver>>,
supervisor: Option<SupervisorHandle>,
started_at: OffsetDateTime,
}
pub fn router(
runtime: Arc<WaveRuntime>,
resident: ResidentDoor,
observer: Option<Arc<StoreObserver>>,
supervisor: Option<SupervisorHandle>,
) -> Router {
let state = ServerState {
runtime,
resident,
observer,
supervisor,
started_at: OffsetDateTime::now_utc(),
};
Router::new()
.route("/health", get(health_handler))
.route("/conversation", get(conversation_handler))
.route("/events", get(events_handler))
.route("/messages", post(messages_handler))
.route("/channels", post(channels_handler))
.route("/memory", get(memory_handler).post(memory_write_handler))
.route("/resident/attach", post(resident_attach_handler))
.route("/resident/deltas", post(resident_deltas_handler))
.route("/resident/context", get(resident_context_handler))
.with_state(state)
}
async fn health_handler(State(state): State<ServerState>) -> Json<HealthBody> {
let mind = state
.runtime
.resident_expected()
.then(|| state.runtime.mind_state().name().to_string());
Json(HealthBody {
status: "serving".to_string(),
mind,
wave: state.runtime.name().to_string(),
turns: state.runtime.thread_len(),
workers: state.runtime.in_flight_workers().len(),
paused: state.runtime.paused(),
uptime_seconds: (OffsetDateTime::now_utc() - state.started_at).whole_seconds(),
})
}
async fn resident_attach_handler(
State(state): State<ServerState>,
headers: HeaderMap,
Json(body): Json<AttachRequest>,
) -> Result<Json<AttachResponse>, (StatusCode, String)> {
state.resident.authorize(&headers)?;
if let Some(seated) = state.resident.seat_pid() {
if seated != body.pid && process_alive(seated).await {
return Err((
StatusCode::CONFLICT,
format!(
"wave '{}' already has a live resident on the seat (pid {seated}); \
stop it before attaching, or use `lf wave <name> --force` to take over",
state.runtime.name()
),
));
}
}
state.resident.record_pid(body.pid);
state.runtime.set_resident_expected();
if let Some(supervisor) = &state.supervisor {
supervisor.on_attach(body.pid);
}
if matches!(state.runtime.mind_state(), MindState::Failed { .. }) {
state
.runtime
.transition(MindState::Idle, "resident attached");
}
tracing::info!(pid = body.pid, "resident attached");
Ok(Json(AttachResponse {
wave: state.runtime.name().to_string(),
thread_id: state.runtime.last_thread_id(),
}))
}
async fn resident_deltas_handler(
State(state): State<ServerState>,
headers: HeaderMap,
Json(body): Json<PostDeltasRequest>,
) -> Result<Json<PostDeltasResponse>, (StatusCode, String)> {
state.resident.authorize(&headers)?;
let accepted = body.deltas.len() as u64;
for delta in body.deltas {
state.runtime.apply_resident_delta(delta);
}
Ok(Json(PostDeltasResponse { accepted }))
}
async fn resident_context_handler(
State(state): State<ServerState>,
headers: HeaderMap,
) -> Result<Json<ContextResponse>, (StatusCode, String)> {
state.resident.authorize(&headers)?;
if let Some(observer) = &state.observer {
observer.poll_once().await;
}
let in_flight = state
.runtime
.in_flight_workers()
.into_iter()
.map(|worker| InFlightWorker {
run_id: worker.run_id,
flow: worker.flow,
task: worker.task,
})
.collect();
Ok(Json(ContextResponse {
thread_id: state.runtime.last_thread_id(),
in_flight,
}))
}
async fn conversation_handler(
State(state): State<ServerState>,
Query(query): Query<ConversationQuery>,
) -> Json<ConversationBody> {
Json(ConversationBody {
turns: state.runtime.thread_tail(query.limit),
})
}
async fn messages_handler(
State(state): State<ServerState>,
Json(body): Json<PostMessage>,
) -> Result<Json<PostMessageResponse>, (StatusCode, String)> {
if body.from.is_some() && !matches!(body.op, MessageOp::Say) {
return Err((
StatusCode::BAD_REQUEST,
"`from` is only valid for the say op".to_string(),
));
}
if matches!(body.op, MessageOp::Say) && body.from.is_none() {
return Err((
StatusCode::BAD_REQUEST,
"`from` is required for the say op".to_string(),
));
}
if body.text.trim().is_empty() && !matches!(body.op, MessageOp::Interrupt) {
return Err((
StatusCode::BAD_REQUEST,
"text is required for every op but interrupt".to_string(),
));
}
let channel = body
.channel
.unwrap_or_else(|| state.runtime.name().to_string());
if !state.runtime.in_family(&channel) {
return Err((
StatusCode::NOT_FOUND,
format!(
"channel '{channel}' is not in wave '{}''s family",
state.runtime.name()
),
));
}
let turn = state
.runtime
.deliver_to_channel(&channel, body.op, body.text, body.from)
.map_err(|err| (StatusCode::NOT_FOUND, err.to_string()))?;
Ok(Json(PostMessageResponse {
turn,
state: state.runtime.mind_state().name().to_string(),
}))
}
async fn channels_handler(
State(state): State<ServerState>,
Json(body): Json<PostChannel>,
) -> Result<Json<PostChannelResponse>, (StatusCode, String)> {
if state.runtime.is_primary(&body.name) || !state.runtime.in_family(&body.name) {
return Err((
StatusCode::NOT_FOUND,
format!(
"'{}' is not a child channel of wave '{}'",
body.name,
state.runtime.name()
),
));
}
let turn = state
.runtime
.journal_channel_opened(&body.name, &body.run_id);
Ok(Json(PostChannelResponse { turn }))
}
async fn memory_handler(State(state): State<ServerState>) -> Json<MemoryBody> {
Json(MemoryBody {
content: state.runtime.memory().read(),
})
}
async fn memory_write_handler(
State(state): State<ServerState>,
Json(body): Json<PostMemory>,
) -> Result<Json<PostMemoryResponse>, (StatusCode, String)> {
let summary = body
.summary
.filter(|s| !s.trim().is_empty())
.or_else(|| first_line(&body.content))
.unwrap_or_else(|| "memory cleared".to_string());
let result = match body.op {
MemoryOp::Update => state.runtime.update_memory(&body.content, &summary),
MemoryOp::Add => {
let fact = body.content.trim();
if fact.is_empty() {
return Err((
StatusCode::BAD_REQUEST,
"content is required for the add op".to_string(),
));
}
state.runtime.append_memory(fact, &summary)
}
};
match result {
Ok(()) => Ok(Json(PostMemoryResponse { summary })),
Err(err) => Err((
StatusCode::INTERNAL_SERVER_ERROR,
format!("memory write failed: {err}"),
)),
}
}
fn first_line(content: &str) -> Option<String> {
content
.lines()
.map(str::trim)
.find(|line| !line.is_empty())
.map(str::to_string)
}
async fn events_handler(
State(state): State<ServerState>,
Query(query): Query<EventsQuery>,
) -> Result<axum::response::Response, (StatusCode, String)> {
let wave = state.runtime.name().to_string();
let include_inbox = query.inbox == Some(true);
let (scope, primary) = match (query.channel, query.prefix) {
(Some(_), Some(_)) => {
return Err((
StatusCode::BAD_REQUEST,
"pass channel or prefix, not both".to_string(),
));
}
(Some(channel), None) => {
let primary = state.runtime.is_primary(&channel);
(Scope::Channel(channel), primary)
}
(None, Some(prefix)) => {
let primary = state.runtime.is_primary(&prefix);
(Scope::Prefix(prefix), primary)
}
(None, None) => (
Scope::Prefix(state.runtime.channel_name().to_string()),
true,
),
};
let name = scope.name();
if !state.runtime.in_family(name) {
return Err((
StatusCode::NOT_FOUND,
format!("'{name}' is not in wave '{wave}''s family"),
));
}
let (child_snapshots, family_rx) = match &scope {
Scope::Channel(channel) if state.runtime.is_primary(channel) => (Vec::new(), None),
Scope::Channel(channel) => {
let (snapshot, rx) = state.runtime.subscribe_child(channel);
let snapshots = snapshot
.map(|turns| vec![(channel.clone(), turns)])
.unwrap_or_default();
(snapshots, Some(rx))
}
Scope::Prefix(prefix) => {
let (snapshots, rx) = state.runtime.subscribe_children(prefix);
(snapshots, Some(rx))
}
};
let child_replay = stream::iter(child_snapshots.into_iter().flat_map(|(channel, turns)| {
turns
.into_iter()
.map(move |turn| Ok(tagged_turn_event(&channel, &turn)))
.collect::<Vec<_>>()
}));
let live_children = family_rx.map(|rx| {
let scope = scope.clone();
BroadcastStream::new(rx).filter_map(move |res| {
let out = match res {
Ok(frame) if scope.matches(&frame.channel) => {
Some(Ok(Event::default().event("turn").data(frame.json.as_ref())))
}
_ => None,
};
async move { out }
})
});
let mut streams: Vec<BoxedEventStream> = Vec::new();
if primary {
let sub = state.runtime.subscribe_with_snapshot();
let inbox_replay: Vec<Result<Event, Infallible>> = if include_inbox {
sub.pending
.iter()
.map(|message| Ok(inbox_event(&pending_inbox_frame(message))))
.collect()
} else {
Vec::new()
};
let replay = stream::iter(
std::iter::once(Ok(state_event(&sub.state)))
.chain(sub.turns.into_iter().map(|t| Ok(turn_event(&t))))
.chain(inbox_replay),
);
let live_turns = BroadcastStream::new(sub.turn_rx).filter_map(move |res| {
let out = match res {
Ok(frame) => Some(Ok(Event::default().event("turn").data(frame.json.as_str()))),
Err(_) => None,
};
async move { out }
});
let live_states = BroadcastStream::new(sub.state_rx).filter_map(move |res| {
let out = match res {
Ok(mind_state) => Some(Ok(state_event(&mind_state))),
Err(_) => None,
};
async move { out }
});
let live_memory = BroadcastStream::new(sub.memory_rx).filter_map(move |res| {
let out = match res {
Ok(summary) => Some(Ok(memory_event(&summary))),
Err(_) => None,
};
async move { out }
});
let mut live: BoxedEventStream = Box::pin(stream::select(
live_turns,
stream::select(live_states, live_memory),
));
if include_inbox {
let live_inbox = BroadcastStream::new(sub.inbox_rx).filter_map(move |res| {
let out = match res {
Ok(item) => Some(Ok(inbox_event(&inbox_item_frame(&item)))),
Err(_) => None,
};
async move { out }
});
live = Box::pin(stream::select(live, live_inbox));
}
streams.push(Box::pin(replay.chain(live)));
}
match live_children {
Some(live) => streams.push(Box::pin(child_replay.chain(live))),
None => streams.push(Box::pin(child_replay)),
}
let merged: BoxedEventStream = match streams.len() {
1 => streams.pop().expect("one stream"),
_ => {
let children = streams.pop().expect("child stream");
let primary = streams.pop().expect("primary stream");
Box::pin(stream::select(primary, children))
}
};
Ok(axum::response::IntoResponse::into_response(
Sse::new(merged).keep_alive(KeepAlive::default()),
))
}
type BoxedEventStream =
std::pin::Pin<Box<dyn Stream<Item = Result<Event, Infallible>> + Send + 'static>>;
#[derive(Debug, Clone)]
enum Scope {
Channel(String),
Prefix(String),
}
impl Scope {
fn name(&self) -> &str {
match self {
Self::Channel(name) | Self::Prefix(name) => name,
}
}
fn matches(&self, channel: &str) -> bool {
match self {
Self::Channel(name) => channel == name,
Self::Prefix(prefix) => crate::wave::channel::matches_prefix(channel, prefix),
}
}
}
fn turn_event(turn: &ChatTurn) -> Event {
Event::default()
.event("turn")
.data(serde_json::to_string(turn).unwrap_or_default())
}
fn tagged_turn_event(channel: &str, turn: &ChatTurn) -> Event {
Event::default()
.event("turn")
.data(tagged_turn_json(channel, turn))
}
fn state_event(state: &MindState) -> Event {
Event::default().event("state").data(state.name())
}
fn memory_event(summary: &str) -> Event {
Event::default().event("memory").data(summary)
}
fn inbox_event(frame: &InboxFrame) -> Event {
Event::default()
.event("inbox")
.data(serde_json::to_string(frame).unwrap_or_default())
}
fn pending_inbox_frame(message: &PendingMessage) -> InboxFrame {
InboxFrame {
id: Some(message.id.0.clone()),
op: message.op,
text: message.text.clone(),
from: message.from.clone(),
}
}
fn inbox_item_frame(item: &InboxItem) -> InboxFrame {
match item {
InboxItem::Message(message) => pending_inbox_frame(message),
InboxItem::Interrupt => InboxFrame {
id: None,
op: MessageOp::Interrupt,
text: String::new(),
from: None,
},
}
}
pub fn endpoint_path(repo_root: &Path, wave: &str) -> PathBuf {
repo_root.join("wave").join(wave).join(ENDPOINT_FILE)
}
const ENDPOINT_PROBE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(2);
pub async fn live_endpoint(repo_root: &Path, wave: &str) -> Option<String> {
let addr = std::fs::read_to_string(endpoint_path(repo_root, wave)).ok()?;
let addr = addr.trim().to_string();
if addr.is_empty() {
return None;
}
let client = reqwest::Client::builder()
.timeout(ENDPOINT_PROBE_TIMEOUT)
.build()
.ok()?;
let body: serde_json::Value = client
.get(format!("http://{addr}/health"))
.send()
.await
.ok()?
.json()
.await
.ok()?;
(body.get("wave").and_then(serde_json::Value::as_str) == Some(wave)).then_some(addr)
}
pub fn write_endpoint(
repo_root: &Path,
wave: &str,
addr: std::net::SocketAddr,
) -> std::io::Result<()> {
let path = endpoint_path(repo_root, wave);
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
std::fs::write(path, addr.to_string())
}
pub fn remove_endpoint(repo_root: &Path, wave: &str, own_addr: &str) {
let path = endpoint_path(repo_root, wave);
match std::fs::read_to_string(&path) {
Ok(contents) if contents.trim() == own_addr => {
let _ = std::fs::remove_file(path);
}
_ => {}
}
}
pub fn resident_token_path(repo_root: &Path, wave: &str) -> PathBuf {
repo_root.join("wave").join(wave).join(RESIDENT_TOKEN_FILE)
}
pub fn write_resident_token(repo_root: &Path, wave: &str, token: &str) -> std::io::Result<()> {
let path = resident_token_path(repo_root, wave);
if let Some(parent) = path.parent() {
std::fs::create_dir_all(parent)?;
}
std::fs::write(&path, token)?;
#[cfg(unix)]
{
use std::os::unix::fs::PermissionsExt;
std::fs::set_permissions(&path, std::fs::Permissions::from_mode(0o600))?;
}
Ok(())
}
pub fn read_resident_token(repo_root: &Path, wave: &str) -> Option<String> {
let token = std::fs::read_to_string(resident_token_path(repo_root, wave)).ok()?;
let token = token.trim().to_string();
(!token.is_empty()).then_some(token)
}
pub fn remove_resident_token(repo_root: &Path, wave: &str, own_token: &str) {
let path = resident_token_path(repo_root, wave);
match std::fs::read_to_string(&path) {
Ok(contents) if contents.trim() == own_token => {
let _ = std::fs::remove_file(path);
}
_ => {}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn write_and_remove_endpoint_roundtrips() {
let tmp = tempfile::tempdir().expect("tempdir");
let addr: std::net::SocketAddr = "127.0.0.1:54321".parse().unwrap();
write_endpoint(tmp.path(), "ship", addr).expect("write endpoint");
let path = endpoint_path(tmp.path(), "ship");
assert_eq!(std::fs::read_to_string(&path).unwrap(), "127.0.0.1:54321");
remove_endpoint(tmp.path(), "ship", "127.0.0.1:54321");
assert!(!path.exists());
}
#[test]
fn remove_endpoint_leaves_a_foreign_pointer_untouched() {
let tmp = tempfile::tempdir().expect("tempdir");
let addr: std::net::SocketAddr = "127.0.0.1:50000".parse().unwrap();
write_endpoint(tmp.path(), "ship", addr).expect("write endpoint");
remove_endpoint(tmp.path(), "ship", "127.0.0.1:50001");
let path = endpoint_path(tmp.path(), "ship");
assert_eq!(
std::fs::read_to_string(&path).unwrap(),
"127.0.0.1:50000",
"foreign pointer survives our shutdown"
);
}
#[test]
fn resident_token_file_roundtrips_and_respects_ownership() {
let tmp = tempfile::tempdir().expect("tempdir");
assert!(read_resident_token(tmp.path(), "ship").is_none());
write_resident_token(tmp.path(), "ship", "tok-1").expect("write");
assert_eq!(
read_resident_token(tmp.path(), "ship").as_deref(),
Some("tok-1")
);
remove_resident_token(tmp.path(), "ship", "tok-other");
assert_eq!(
read_resident_token(tmp.path(), "ship").as_deref(),
Some("tok-1")
);
remove_resident_token(tmp.path(), "ship", "tok-1");
assert!(read_resident_token(tmp.path(), "ship").is_none());
}
#[tokio::test]
async fn live_endpoint_is_none_for_a_stale_pointer() {
let tmp = tempfile::tempdir().expect("tempdir");
assert!(live_endpoint(tmp.path(), "ship").await.is_none(), "no file");
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let dead = listener.local_addr().unwrap();
drop(listener);
write_endpoint(tmp.path(), "ship", dead).expect("write endpoint");
assert!(
live_endpoint(tmp.path(), "ship").await.is_none(),
"dead address is stale"
);
}
}