use std::sync::Arc;
use std::time::Duration;
use axum::extract::ws::{CloseFrame, Message, WebSocket, WebSocketUpgrade};
use axum::routing::{get, post};
use axum::{Extension, Json, Router};
use serde::Deserialize;
use serde_json::{json, Value};
use tokio::sync::broadcast;
use tokio::time::{interval_at, Instant as TokioInstant};
use crate::auth::{AuthBackend, Principal, Role};
use crate::QueryState;
use pensieve_ingest_core::{ConsumerAction, ConsumerActivity, ConsumerEvents};
const AUTH_TIMEOUT: Duration = Duration::from_secs(5);
const HEARTBEAT: Duration = Duration::from_secs(15);
const BACKFILL_LIMIT: usize = 50;
#[derive(Clone)]
struct Deps {
state: QueryState,
backend: Arc<dyn AuthBackend>,
events: Option<ConsumerEvents>,
traces_db: String,
}
pub fn consumers_live_router(
state: QueryState,
backend: Arc<dyn AuthBackend>,
events: Option<ConsumerEvents>,
traces_db: String,
) -> Router {
let deps = Deps {
state,
backend,
events,
traces_db,
};
Router::new().route(
"/v1/consumers/live",
get(move |ws: WebSocketUpgrade| {
let deps = deps.clone();
async move { ws.on_upgrade(move |sock| session(sock, deps)) }
}),
)
}
pub fn consumers_emit_router(events: Option<ConsumerEvents>) -> Router {
Router::new().route(
"/v1/consumers/emit",
post(
move |Extension(principal): Extension<Principal>,
Json(activity): Json<ConsumerActivity>| {
let events = events.clone();
async move {
if let Some(bus) = events {
let mut a = activity;
a.tenant = principal.tenant.to_string();
bus.publish(a);
}
axum::http::StatusCode::ACCEPTED
}
},
),
)
}
#[derive(Debug, Deserialize)]
#[serde(tag = "type", rename_all = "lowercase")]
enum ClientMsg {
Auth { token: String },
}
async fn session(mut ws: WebSocket, deps: Deps) {
let principal = match await_auth(&mut ws, &deps.backend).await {
Some(p) => p,
None => {
let _ = ws.send(close_policy("unauthorized")).await;
return;
}
};
if principal.role < Role::Read {
let _ = send_json(
&mut ws,
&json!({"type":"error","code":"forbidden","message":"live consumers requires read access"}),
)
.await;
let _ = ws.send(close_policy("forbidden")).await;
return;
}
let tenant = principal.tenant.to_string();
for act in load_backfill(&deps, &tenant).await {
if send_json(&mut ws, &json!({"type":"backfill","activity": act}))
.await
.is_err()
{
return;
}
}
if send_json(&mut ws, &json!({"type":"live"})).await.is_err() {
return;
}
let mut rx = deps.events.as_ref().map(|e| e.subscribe());
let mut hb = interval_at(TokioInstant::now() + HEARTBEAT, HEARTBEAT);
loop {
tokio::select! {
_ = hb.tick() => {
if send_json(&mut ws, &json!({"type":"heartbeat"})).await.is_err() {
return;
}
}
msg = ws.recv() => {
match msg {
Some(Ok(Message::Close(_))) | None | Some(Err(_)) => return,
_ => {}
}
}
ev = recv_event(rx.as_mut()) => {
match ev {
EventOutcome::Event(act) => {
if act.tenant == tenant
&& send_json(&mut ws, &json!({"type":"events","events":[act]}))
.await
.is_err()
{
return;
}
}
EventOutcome::Lagged => {
if send_json(&mut ws, &json!({"type":"lagged"})).await.is_err() {
return;
}
}
EventOutcome::Closed => rx = None,
}
}
}
}
}
async fn await_auth(ws: &mut WebSocket, backend: &Arc<dyn AuthBackend>) -> Option<Principal> {
let deadline = tokio::time::Instant::now() + AUTH_TIMEOUT;
loop {
match tokio::time::timeout_at(deadline, ws.recv()).await {
Ok(Some(Ok(Message::Text(t)))) => {
let token = match serde_json::from_str::<ClientMsg>(&t) {
Ok(ClientMsg::Auth { token }) => token,
Err(_) => return None,
};
return backend.authenticate(&token).await.ok();
}
Ok(Some(Ok(Message::Ping(_))))
| Ok(Some(Ok(Message::Pong(_))))
| Ok(Some(Ok(Message::Binary(_)))) => continue,
_ => return None,
}
}
}
async fn send_json(ws: &mut WebSocket, v: &Value) -> Result<(), ()> {
ws.send(Message::Text(v.to_string())).await.map_err(|_| ())
}
fn close_policy(reason: &str) -> Message {
Message::Close(Some(CloseFrame {
code: 1008,
reason: reason.to_string().into(),
}))
}
enum EventOutcome {
Event(ConsumerActivity),
Lagged,
Closed,
}
async fn recv_event(rx: Option<&mut broadcast::Receiver<ConsumerActivity>>) -> EventOutcome {
match rx {
None => std::future::pending().await,
Some(rx) => match rx.recv().await {
Ok(ev) => EventOutcome::Event(ev),
Err(broadcast::error::RecvError::Lagged(_)) => EventOutcome::Lagged,
Err(broadcast::error::RecvError::Closed) => EventOutcome::Closed,
},
}
}
async fn load_backfill(deps: &Deps, tenant: &str) -> Vec<ConsumerActivity> {
let shared = crate::agent::SharedToolCtx {
realm_scope: Default::default(),
consumer_sink: None,
federation: None,
catalog: deps.state.catalog.clone(),
format: deps.state.format.clone(),
pool: None,
memory: None,
hitl: None,
memory_settings_path: None,
};
let tenant_esc = tenant.replace('\'', "''");
let sql = format!(
"SELECT name, subject, start_time, attributes_json FROM otel_traces \
WHERE (name = 'memory.recall' OR name = 'memory.remember' OR name = 'memory.import') \
AND tenant = '{tenant_esc}' \
ORDER BY start_time DESC LIMIT {BACKFILL_LIMIT}"
);
let res = crate::agent::execute_sql(&shared, &deps.traces_db, &sql, BACKFILL_LIMIT).await;
let Some(rows) = res.get("rows").and_then(|v| v.as_array()) else {
return Vec::new();
};
let mut out: Vec<ConsumerActivity> = rows
.iter()
.filter_map(|row| backfill_row_to_activity(row, tenant))
.collect();
out.reverse();
out
}
fn backfill_row_to_activity(row: &Value, tenant: &str) -> Option<ConsumerActivity> {
let name = row.get("name")?.as_str()?;
let action = if name == "memory.recall" {
ConsumerAction::Recall
} else {
ConsumerAction::Remember
};
let subject = row
.get("subject")
.and_then(|v| v.as_str())
.filter(|s| !s.is_empty())
.map(|s| s.to_string());
let attrs: Option<Value> = row
.get("attributes_json")
.and_then(|v| v.as_str())
.and_then(|s| serde_json::from_str::<Value>(s).ok());
let query_preview = attrs.as_ref().and_then(|a| {
a.get("memory.query")
.and_then(|q| q.as_str())
.map(|q| q.to_string())
});
let ts = row
.get("start_time")
.and_then(|v| v.as_str())
.and_then(|s| chrono::DateTime::parse_from_rfc3339(s).ok())
.map(|dt| dt.timestamp_millis())
.unwrap_or_else(|| chrono::Utc::now().timestamp_millis());
let kind = attrs
.as_ref()
.and_then(|a| a.get("pensieve.client").and_then(|c| c.as_str()))
.filter(|s| !s.is_empty())
.unwrap_or("pensieve")
.to_string();
Some(ConsumerActivity {
consumer_id: format!("{kind}:{}", subject.as_deref().unwrap_or("anon")),
label: subject.clone().unwrap_or_else(|| kind.clone()),
kind,
subject,
tenant: tenant.to_string(),
action,
node_ids: Vec::new(),
namespaces: Vec::new(),
query_preview,
ts,
host: None,
client_version: None,
transport: None,
ip: None,
pid: None,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn backfill_maps_recall_row() {
let row = json!({
"name": "memory.recall",
"subject": "shaked",
"start_time": "2026-06-16T10:00:00Z",
"attributes_json": "{\"memory.query\":\"auth flow\",\"memory.results\":\"3\"}",
});
let act = backfill_row_to_activity(&row, "default").expect("mapped");
assert_eq!(act.action, ConsumerAction::Recall);
assert_eq!(act.subject.as_deref(), Some("shaked"));
assert_eq!(act.query_preview.as_deref(), Some("auth flow"));
assert_eq!(act.tenant, "default");
let expected = chrono::DateTime::parse_from_rfc3339("2026-06-16T10:00:00Z")
.unwrap()
.timestamp_millis();
assert_eq!(act.ts, expected);
}
#[test]
fn backfill_maps_write_row_and_handles_missing_fields() {
let row = json!({ "name": "memory.import" });
let act = backfill_row_to_activity(&row, "t1").expect("mapped");
assert_eq!(act.action, ConsumerAction::Remember);
assert!(act.subject.is_none());
assert!(act.query_preview.is_none());
assert_eq!(act.consumer_id, "pensieve:anon");
assert_eq!(act.kind, "pensieve");
}
#[test]
fn backfill_reads_stamped_client() {
let row = json!({
"name": "memory.recall",
"attributes_json": "{\"pensieve.client\":\"claude-code\",\"memory.query\":\"q\"}",
});
let act = backfill_row_to_activity(&row, "t1").expect("mapped");
assert_eq!(act.kind, "claude-code");
assert_eq!(act.consumer_id, "claude-code:anon");
}
#[test]
fn backfill_skips_rows_without_name() {
assert!(backfill_row_to_activity(&json!({"subject":"x"}), "t").is_none());
}
}