use crate::{AppState, AuthIdentity};
use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
use axum::extract::{Extension, State};
use axum::http::StatusCode;
use axum::response::{IntoResponse, Response};
use axum::Json;
use core_api::{DbEvent, Subscription};
use serde::Deserialize;
use serde_json::json;
use std::time::Duration;
use tokio::task;
const BRIDGE_IDLE_TIMEOUT: Duration = Duration::from_millis(100);
#[derive(Debug, Deserialize, Default)]
struct SubscribeMsg {
#[serde(default)]
rules: Vec<String>,
#[serde(default)]
writes: bool,
cypher: Option<String>,
}
pub async fn subscribe(
ws: WebSocketUpgrade,
State(state): State<AppState>,
Extension(identity): Extension<AuthIdentity>,
) -> Response {
if let AuthIdentity::Role(_) = identity {
return (
StatusCode::FORBIDDEN,
Json(json!({"error": "role-bound token: /subscribe is not permitted"})),
)
.into_response();
}
ws.on_upgrade(move |socket| run(socket, state))
.into_response()
}
async fn run(mut socket: WebSocket, state: AppState) {
let msg = loop {
match socket.recv().await {
Some(Ok(Message::Text(t))) => match serde_json::from_str::<SubscribeMsg>(&t) {
Ok(m) => break m,
Err(e) => {
let _ = socket
.send(Message::Text(
serde_json::json!({"error": format!("bad subscribe message: {e}")})
.to_string()
.into(),
))
.await;
return;
}
},
Some(Ok(Message::Close(_))) | None => return,
Some(Ok(_)) => continue,
Some(Err(_)) => return,
}
};
let sub_result: Result<Vec<Subscription>, String> = {
let mut db = state.db.write();
let mut subs: Vec<Subscription> = Vec::new();
let mut err: Option<String> = None;
for rule in &msg.rules {
match db.subscribe_rule(rule) {
Ok(sub) => subs.push(sub),
Err(e) => {
err = Some(e.to_string());
break;
}
}
}
if err.is_none() && msg.writes {
match db.subscribe_writes() {
Ok(sub) => subs.push(sub),
Err(e) => err = Some(e.to_string()),
}
}
if err.is_none() {
if let Some(cypher) = &msg.cypher {
match db.subscribe_query(cypher) {
Ok(sub) => subs.push(sub),
Err(e) => err = Some(e.to_string()),
}
}
}
drop(db);
match err {
Some(e) => Err(e),
None => Ok(subs),
}
};
let subs = match sub_result {
Ok(s) => s,
Err(e) => {
let _ = socket
.send(Message::Text(
serde_json::json!({"error": e}).to_string().into(),
))
.await;
return;
}
};
if socket
.send(Message::Text(r#"{"subscribed":true}"#.into()))
.await
.is_err()
{
return;
}
let (event_tx, mut event_rx) = tokio::sync::mpsc::channel::<DbEvent>(256);
let _bridge = task::spawn_blocking(move || bridge_loop(subs, event_tx));
loop {
tokio::select! {
ev = event_rx.recv() => {
let Some(ev) = ev else { break; }; let text = serde_json::to_string(&ev).expect("DbEvent is always serializable");
if socket.send(Message::Text(text.into())).await.is_err() {
break;
}
}
msg = socket.recv() => {
match msg {
Some(Ok(Message::Close(_))) | None | Some(Err(_)) => break,
_ => {} }
}
}
}
}
fn bridge_loop(subs: Vec<Subscription>, tx: tokio::sync::mpsc::Sender<DbEvent>) {
if subs.is_empty() {
return;
}
loop {
if tx.is_closed() {
return;
}
let mut sent_any = false;
for sub in &subs {
while let Some(ev) = sub.try_recv() {
if tx.blocking_send(ev).is_err() {
return; }
sent_any = true;
}
}
if !sent_any {
if let Some(ev) = subs[0].recv_timeout(BRIDGE_IDLE_TIMEOUT) {
if tx.blocking_send(ev).is_err() {
return;
}
for sub in &subs {
while let Some(ev) = sub.try_recv() {
if tx.blocking_send(ev).is_err() {
return;
}
}
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn lagged_event_serializes_to_type_tagged_json() {
let ev = DbEvent::Lagged { missed: 5 };
let text = serde_json::to_string(&ev).unwrap();
assert_eq!(text, r#"{"type":"lagged","missed":5}"#);
}
#[test]
fn edge_fired_serializes_correctly() {
let ev = DbEvent::EdgeFired {
rule: "rel".into(),
src_key: "n1".into(),
dst_key: "n2".into(),
edge_type: "REL".into(),
weight: Some(0.9),
commit_seq: 7,
};
let j: serde_json::Value =
serde_json::from_str(&serde_json::to_string(&ev).unwrap()).unwrap();
assert_eq!(j["type"], "edge_fired");
assert_eq!(j["rule"], "rel");
assert_eq!(j["commit_seq"], 7);
assert_eq!(j["weight"], 0.9);
}
#[test]
fn subscribe_msg_defaults_to_empty() {
let m: SubscribeMsg = serde_json::from_str("{}").unwrap();
assert!(m.rules.is_empty());
assert!(!m.writes);
}
#[test]
fn subscribe_msg_parses_rules_and_writes() {
let m: SubscribeMsg = serde_json::from_str(r#"{"rules":["rel"],"writes":true}"#).unwrap();
assert_eq!(m.rules, ["rel"]);
assert!(m.writes);
}
}