use super::cluster::{collect_sysinfo, generate_cluster_status};
use super::system::AppState;
use crate::{server::handlers::auth::AuthParams, storage::StorageEngine};
use axum::{
body::Body,
extract::{
ws::{Message, WebSocket, WebSocketUpgrade},
Query as AxumQuery, State,
},
http::HeaderMap,
http::StatusCode,
response::{IntoResponse, Response},
};
use futures::{SinkExt, StreamExt};
use serde::Deserialize;
use std::sync::Arc;
const MAX_WS_MESSAGE_SIZE: usize = 1024 * 1024;
fn validate_ws_origin(headers: &HeaderMap) -> Result<(), ()> {
let origin = match headers.get("origin").and_then(|o| o.to_str().ok()) {
Some(o) => o,
None => return Ok(()), };
let allowed_raw = std::env::var("SOLIDB_CORS_ALLOWED_ORIGINS").unwrap_or_default();
if allowed_raw == "*" {
return Ok(());
}
if allowed_raw.is_empty() {
tracing::warn!(
"WebSocket: rejecting Origin '{}' — SOLIDB_CORS_ALLOWED_ORIGINS not set",
origin
);
return Err(());
}
let allowed = allowed_raw
.split(',')
.map(str::trim)
.any(|a| a == origin || a == "*");
if allowed {
Ok(())
} else {
tracing::warn!("WebSocket: rejecting disallowed Origin '{}'", origin);
Err(())
}
}
fn forbidden_response() -> Response {
Response::builder()
.status(StatusCode::FORBIDDEN)
.body(Body::empty())
.expect("Valid status code should not fail")
.into_response()
}
fn unauthorized_response() -> Response {
Response::builder()
.status(StatusCode::UNAUTHORIZED)
.body(Body::empty())
.expect("Valid status code should not fail")
.into_response()
}
fn limit_upgrade(ws: WebSocketUpgrade) -> WebSocketUpgrade {
ws.max_message_size(MAX_WS_MESSAGE_SIZE)
.max_frame_size(MAX_WS_MESSAGE_SIZE)
}
const WS_REVALIDATE_INTERVAL: std::time::Duration = std::time::Duration::from_secs(60);
async fn authenticate_ws_token(
token: &str,
storage: &Arc<StorageEngine>,
allow_livequery: bool,
) -> Option<crate::server::auth::Claims> {
let claims = crate::server::auth::AuthService::validate_token(token).ok()?;
if claims.livequery == Some(true) && !allow_livequery {
tracing::warn!("livequery token presented to a non-changefeed WebSocket");
return None;
}
let storage = storage.clone();
tokio::task::spawn_blocking(move || {
let claims = crate::server::auth::refresh_jwt_roles(claims, &storage)?;
check_ws_credential(&claims, &storage)
.map_err(|reason| {
tracing::warn!(
target: "audit",
user = %claims.sub,
"rejecting WebSocket token: {}",
reason
);
})
.ok()?;
Some(claims)
})
.await
.ok()
.flatten()
}
fn check_ws_credential(
claims: &crate::server::auth::Claims,
storage: &StorageEngine,
) -> Result<(), &'static str> {
if claims.livequery != Some(true) {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs() as usize)
.unwrap_or(usize::MAX);
if claims.exp <= now {
return Err("token expired");
}
}
let system_db = storage
.get_database("_system")
.map_err(|_| "system database unavailable")?;
if let Some(name) = claims.sub.strip_prefix("api-key:") {
let coll = system_db
.system_collection(crate::server::auth::API_KEYS_COLL)
.map_err(|_| "API key revoked")?;
let now = chrono::Utc::now();
let alive = coll
.scan(None)
.into_iter()
.filter_map(|d| {
serde_json::from_value::<crate::server::auth::ApiKey>(d.to_value()).ok()
})
.any(|k| {
k.name == name
&& !k
.expires_at
.as_deref()
.and_then(|e| chrono::DateTime::parse_from_rfc3339(e).ok())
.is_some_and(|e| e < now)
});
return if alive {
Ok(())
} else {
Err("API key revoked or expired")
};
}
let admins = system_db
.system_collection(crate::server::auth::ADMIN_COLL)
.map_err(|_| "user no longer exists")?;
if admins.get(&claims.sub).is_err() {
return Err("user no longer exists");
}
let normalized = |roles: Option<Vec<String>>| {
let mut roles = roles.unwrap_or_default();
roles.sort();
roles.dedup();
roles
};
let current = crate::server::auth::AuthService::get_user_roles(storage, &claims.sub);
if normalized(current) != normalized(claims.roles.clone()) {
return Err("roles changed");
}
Ok(())
}
async fn ws_credential_still_valid(
claims: &crate::server::auth::Claims,
storage: &Arc<StorageEngine>,
) -> bool {
let claims = claims.clone();
let storage = storage.clone();
tokio::task::spawn_blocking(move || match check_ws_credential(&claims, &storage) {
Ok(()) => true,
Err(reason) => {
tracing::warn!(
target: "audit",
user = %claims.sub,
"closing WebSocket: {}",
reason
);
false
}
})
.await
.unwrap_or(false)
}
fn session_ended_message() -> Message {
Message::Text(
serde_json::json!({
"type": "error",
"error": "Session no longer valid; reconnect with a fresh token"
})
.to_string()
.into(),
)
}
pub async fn cluster_status_ws(
ws: WebSocketUpgrade,
AxumQuery(params): AxumQuery<AuthParams>,
State(state): State<AppState>,
headers: HeaderMap,
) -> Response {
let Some(claims) = authenticate_ws_token(¶ms.token, &state.storage, false).await else {
return unauthorized_response();
};
if crate::server::authz_middleware::enforce(
&claims,
&state,
crate::server::authorization::PermissionAction::Admin,
None,
)
.await
.is_err()
{
return forbidden_response();
}
if validate_ws_origin(&headers).is_err() {
return forbidden_response();
}
limit_upgrade(ws).on_upgrade(|socket| handle_cluster_ws(socket, state, claims))
}
async fn handle_cluster_ws(
mut socket: WebSocket,
state: AppState,
claims: crate::server::auth::Claims,
) {
use tokio::time::{interval, Duration};
let mut ticker = interval(Duration::from_secs(1));
let mut validated_at = tokio::time::Instant::now();
loop {
tokio::select! {
_ = ticker.tick() => {
if validated_at.elapsed() >= WS_REVALIDATE_INTERVAL {
if !ws_credential_still_valid(&claims, &state.storage).await {
let _ = socket.send(session_ended_message()).await;
break;
}
validated_at = tokio::time::Instant::now();
}
let sysinfo = {
let mut sys = state.system_monitor.lock().unwrap();
collect_sysinfo(&mut sys)
};
let status = generate_cluster_status(&state, &sysinfo);
let json = match serde_json::to_string(&status) {
Ok(j) => j,
Err(_) => continue,
};
if socket.send(Message::Text(json.into())).await.is_err() {
break; }
}
msg = socket.recv() => {
match msg {
Some(Ok(Message::Close(_))) | None => break,
#[allow(clippy::collapsible_match)]
Some(Ok(Message::Ping(data))) => {
if socket.send(Message::Pong(data)).await.is_err() {
break;
}
}
_ => {} }
}
}
}
}
pub async fn monitor_ws_handler(
ws: WebSocketUpgrade,
AxumQuery(params): AxumQuery<AuthParams>,
State(state): State<AppState>,
headers: HeaderMap,
) -> Response {
let Some(claims) = authenticate_ws_token(¶ms.token, &state.storage, false).await else {
return unauthorized_response();
};
if crate::server::authz_middleware::enforce(
&claims,
&state,
crate::server::authorization::PermissionAction::Admin,
None,
)
.await
.is_err()
{
return forbidden_response();
}
if validate_ws_origin(&headers).is_err() {
return forbidden_response();
}
limit_upgrade(ws).on_upgrade(|socket| handle_monitor_socket(socket, state, claims))
}
async fn handle_monitor_socket(
mut socket: WebSocket,
state: AppState,
claims: crate::server::auth::Claims,
) {
use std::sync::atomic::Ordering;
tracing::info!("Monitor WS: Client connected");
let mut interval = tokio::time::interval(std::time::Duration::from_secs(2));
let mut validated_at = tokio::time::Instant::now();
loop {
interval.tick().await;
if validated_at.elapsed() >= WS_REVALIDATE_INTERVAL {
if !ws_credential_still_valid(&claims, &state.storage).await {
let _ = socket.send(session_ended_message()).await;
break;
}
validated_at = tokio::time::Instant::now();
}
let stats = {
let mut sys = state.system_monitor.lock().unwrap();
sys.refresh_cpu_all();
sys.refresh_memory();
let cpu = sys.global_cpu_usage();
let mem_used = sys.used_memory();
let mem_total = sys.total_memory();
let up = sysinfo::System::uptime();
let name = sysinfo::System::name().unwrap_or_else(|| "Unknown".to_string());
let version =
sysinfo::System::kernel_version().unwrap_or_else(|| "Unknown".to_string());
let host = sysinfo::System::host_name().unwrap_or_else(|| "Unknown".to_string());
let cores = sys.cpus().len();
serde_json::json!({
"cpu_usage": cpu,
"memory_usage": mem_used,
"memory_total": mem_total,
"uptime": up,
"os_name": name,
"os_version": version,
"hostname": host,
"num_cpus": cores,
"pid": std::process::id(),
"active_scripts": state.script_stats.active_scripts.load(Ordering::Relaxed),
"active_ws": state.script_stats.active_ws.load(Ordering::Relaxed)
})
};
let msg = match serde_json::to_string(&stats) {
Ok(s) => s,
Err(_) => continue,
};
if socket.send(Message::Text(msg.into())).await.is_err() {
break;
}
}
}
#[derive(Debug, Deserialize)]
pub struct ChangefeedRequest {
#[serde(rename = "type")]
pub type_: String,
pub collection: Option<String>,
pub database: Option<String>,
pub key: Option<String>,
pub local: Option<bool>,
pub query: Option<String>,
pub id: Option<String>,
}
pub async fn ws_changefeed_handler(
ws: WebSocketUpgrade,
headers: HeaderMap,
AxumQuery(params): AxumQuery<AuthParams>,
State(state): State<AppState>,
) -> impl IntoResponse {
let is_cluster_internal = {
let cluster_secret = state.cluster_secret();
let provided_secret = headers
.get("X-Cluster-Secret")
.and_then(|h| h.to_str().ok())
.unwrap_or("");
!cluster_secret.is_empty()
&& crate::server::auth::constant_time_eq(
cluster_secret.as_bytes(),
provided_secret.as_bytes(),
)
};
let claims = if is_cluster_internal {
crate::server::auth::Claims {
sub: "cluster-internal".to_string(),
exp: usize::MAX,
livequery: None,
roles: Some(vec!["admin".to_string()]),
scoped_databases: None,
}
} else {
match authenticate_ws_token(¶ms.token, &state.storage, true).await {
Some(claims) => claims,
None => return unauthorized_response(),
}
};
if validate_ws_origin(&headers).is_err() {
return forbidden_response();
}
let use_htmx = params.htmx.map(|s| s == "true").unwrap_or(false);
limit_upgrade(ws).on_upgrade(move |socket| {
handle_socket(socket, state, claims, use_htmx, is_cluster_internal)
})
}
async fn handle_socket(
socket: WebSocket,
state: AppState,
claims: crate::server::auth::Claims,
use_htmx: bool,
is_cluster_internal: bool,
) {
let (mut sender, mut receiver) = socket.split();
let (tx, mut rx) = tokio::sync::mpsc::channel::<Message>(1000);
let mut send_task = tokio::spawn(async move {
let heartbeat = std::time::Duration::from_secs(30);
let mut heartbeat_interval =
tokio::time::interval_at(tokio::time::Instant::now() + heartbeat, heartbeat);
loop {
tokio::select! {
_ = heartbeat_interval.tick() => {
if sender.send(Message::Ping(vec![].into())).await.is_err() {
tracing::debug!("[WS] Failed to send ping, closing writer");
break;
}
}
msg = rx.recv() => {
let Some(msg) = msg else { break };
if sender.send(msg).await.is_err() {
tracing::debug!("[WS] Failed to send message, closing writer");
break;
}
}
}
}
});
let mut subscriptions = 0usize;
let mut tasks = tokio::task::JoinSet::new();
let mut revalidate = tokio::time::interval_at(
tokio::time::Instant::now() + WS_REVALIDATE_INTERVAL,
WS_REVALIDATE_INTERVAL,
);
loop {
let msg = tokio::select! {
next = receiver.next() => match next {
Some(Ok(msg)) => msg,
_ => break,
},
_ = tx.closed() => break,
_ = revalidate.tick() => {
if !is_cluster_internal
&& !ws_credential_still_valid(&claims, &state.storage).await
{
let _ = tx.send(session_ended_message()).await;
break;
}
continue;
}
Some(_) = tasks.join_next(), if !tasks.is_empty() => continue,
};
let msg_len = match &msg {
Message::Text(text) => text.len(),
Message::Binary(data) => data.len(),
Message::Ping(data) => data.len(),
Message::Pong(data) => data.len(),
_ => 0,
};
if msg_len > MAX_WS_MESSAGE_SIZE {
tracing::warn!(
"[WS] Message size {} exceeds limit {}, closing connection",
msg_len,
MAX_WS_MESSAGE_SIZE
);
let _ = tx
.send(Message::Text(
serde_json::json!({
"error": "Message too large"
})
.to_string()
.into(),
))
.await;
break;
}
match msg {
Message::Text(text) => {
let req_result = serde_json::from_str::<ChangefeedRequest>(&text);
match req_result {
Ok(req) if req.type_ == "subscribe" => {
if subscriptions >= MAX_SUBSCRIPTIONS_PER_CONNECTION {
let _ = tx.send(subscription_limit_message()).await;
continue;
}
subscriptions += 1;
let tx_clone = tx.clone();
let state_clone = state.clone();
let claims_clone = claims.clone();
tasks.spawn(async move {
handle_subscribe_request(
req,
state_clone,
claims_clone,
tx_clone,
use_htmx,
)
.await;
});
}
Ok(req) if req.type_ == "live_query" => {
if subscriptions >= MAX_SUBSCRIPTIONS_PER_CONNECTION {
let _ = tx.send(subscription_limit_message()).await;
continue;
}
subscriptions += 1;
let tx_clone = tx.clone();
let state_clone = state.clone();
let claims_clone = claims.clone();
tasks.spawn(async move {
handle_live_query_request(req, state_clone, claims_clone, tx_clone)
.await;
});
}
_ => {
let _ = tx
.send(Message::Text(
serde_json::json!({
"error": "Invalid subscription request or unknown type"
})
.to_string()
.into(),
))
.await;
}
}
}
Message::Close(_) => break,
Message::Ping(_) => {
}
Message::Pong(_) => {
}
_ => {}
}
}
tasks.shutdown().await;
drop(tx);
if tokio::time::timeout(std::time::Duration::from_secs(1), &mut send_task)
.await
.is_err()
{
send_task.abort();
}
}
const MAX_SUBSCRIPTIONS_PER_CONNECTION: usize = 64;
fn subscription_limit_message() -> Message {
Message::Text(
serde_json::json!({
"type": "error",
"error": format!(
"Subscription limit reached ({} per connection)",
MAX_SUBSCRIPTIONS_PER_CONNECTION
),
})
.to_string()
.into(),
)
}
async fn forward_changes(
mut rx: tokio::sync::broadcast::Receiver<crate::storage::collection::ChangeEvent>,
out: tokio::sync::mpsc::Sender<crate::storage::collection::ChangeEvent>,
key: Option<String>,
source: String,
lag_notice: Option<tokio::sync::mpsc::Sender<Message>>,
) {
use tokio::sync::broadcast::error::RecvError;
loop {
let received = tokio::select! {
_ = out.closed() => break,
received = rx.recv() => received,
};
match received {
Ok(event) => {
if key.as_ref().is_some_and(|k| &event.key != k) {
continue;
}
if out.send(event).await.is_err() {
break;
}
}
Err(RecvError::Lagged(skipped)) => {
tracing::warn!(
"[WS] changefeed on '{}' lagged: {} events dropped",
source,
skipped
);
if let Some(client) = &lag_notice {
let _ = client.try_send(Message::Text(
serde_json::json!({
"type": "lagged",
"collection": source,
"skipped": skipped,
})
.to_string()
.into(),
));
}
}
Err(RecvError::Closed) => break,
}
}
}
async fn forward_remote_changes(
node_addr: String,
db_name: String,
coll_name: String,
secret: String,
out: tokio::sync::mpsc::Sender<crate::storage::collection::ChangeEvent>,
) {
use crate::cluster::ClusterWebsocketClient;
let connect = ClusterWebsocketClient::connect(&node_addr, &db_name, &coll_name, true, &secret);
let stream = tokio::select! {
_ = out.closed() => return,
stream = connect => match stream {
Ok(stream) => stream,
Err(_) => return,
},
};
tokio::pin!(stream);
loop {
let next = tokio::select! {
_ = out.closed() => break,
next = stream.next() => next,
};
match next {
Some(Ok(event)) => {
if out.send(event).await.is_err() {
break;
}
}
_ => break,
}
}
}
fn row_policy_allows(
storage: &StorageEngine,
db_name: String,
collection: &str,
principal: &crate::sdbql::QueryPrincipal,
policy: &str,
doc: serde_json::Value,
) -> bool {
let Ok(mut parser) = crate::sdbql::parser::Parser::new(policy) else {
return false;
};
let Ok(expr) = parser.parse_expression() else {
return false;
};
let executor = crate::sdbql::executor::QueryExecutor::with_database(storage, db_name)
.with_principal(principal.clone())
.with_timeout(std::time::Duration::from_secs(5));
let mut ctx = std::collections::HashMap::new();
ctx.insert(collection.to_string(), doc.clone());
ctx.insert("doc".to_string(), doc);
ctx.insert(
"CURRENT_USER".to_string(),
serde_json::Value::String(principal.user.clone()),
);
executor
.evaluate_expr_with_context(&expr, &ctx)
.map(|v| crate::sdbql::executor::to_bool(&v))
.unwrap_or(false)
}
async fn changefeed_event_visible(
storage: &Arc<StorageEngine>,
db_name: &str,
collection: &crate::storage::Collection,
principal: &crate::sdbql::QueryPrincipal,
event: &crate::storage::collection::ChangeEvent,
) -> bool {
use crate::storage::collection::ChangeType;
if principal.can_admin {
return true;
}
let Some(policy) = collection.get_row_policy() else {
return true;
};
let row = match event.type_ {
ChangeType::Truncate => return true,
ChangeType::Insert | ChangeType::Update => event.data.clone(),
ChangeType::Delete => event.old_data.clone().or_else(|| event.data.clone()),
};
let Some(row) = row else {
return false;
};
let storage = storage.clone();
let db_name = db_name.to_string();
let coll_name = collection.name.clone();
let principal = principal.clone();
tokio::task::spawn_blocking(move || {
row_policy_allows(&storage, db_name, &coll_name, &principal, &policy, row)
})
.await
.unwrap_or(false)
}
async fn handle_subscribe_request(
req: ChangefeedRequest,
state: AppState,
claims: crate::server::auth::Claims,
tx: tokio::sync::mpsc::Sender<Message>,
use_htmx: bool,
) {
let db_name = req.database.clone().unwrap_or("_system".to_string());
if let Err(e) = crate::server::authz_middleware::enforce(
&claims,
&state,
crate::server::authorization::PermissionAction::Read,
Some(&db_name),
)
.await
{
let mut response = serde_json::json!({ "error": e.to_string() });
if let Some(req_id) = &req.id {
response["id"] = serde_json::Value::String(req_id.clone());
}
let _ = tx.send(Message::Text(response.to_string().into())).await;
return;
}
let coll_name = match req.collection.clone() {
Some(c) => c,
None => {
if let Some(query_str) = &req.query {
if let Ok(query_ast) = crate::sdbql::parser::parse(query_str) {
if let Some(first_for) = query_ast.for_clauses.first() {
first_for.collection.clone()
} else {
query_ast
.body_clauses
.iter()
.find_map(|c| {
if let crate::sdbql::ast::BodyClause::For(f) = c {
Some(f.collection.clone())
} else {
None
}
})
.unwrap_or_default()
}
} else {
"".to_string()
}
} else {
"".to_string()
}
}
};
if coll_name.is_empty() {
let _ = tx
.send(Message::Text(
serde_json::json!({
"error": "Collection required for subscribe mode (could not infer from query)"
})
.to_string()
.into(),
))
.await;
return;
}
let collection_result = state
.storage
.get_database(&db_name)
.and_then(|db| db.get_collection(&coll_name));
match collection_result {
Ok(collection) => {
let msg = if use_htmx {
format!(
r#"<div id="connection-status" hx-swap-oob="innerHTML" class="inline-flex items-center gap-2 px-3 py-1.5 rounded-full text-sm bg-success/10 text-success">
<span class="w-2 h-2 rounded-full bg-success animate-pulse"></span>
<span>Connected: {}</span>
</div>
<div id="no-subscriptions" hx-swap-oob="true" class="hidden"></div>
<div id="subscriptions-list" hx-swap-oob="beforeend">
<div class="px-4 py-3 border-b border-border/20 last:border-0 flex items-center justify-between">
<div class="flex items-center gap-3">
<span class="w-2 h-2 rounded-full bg-success animate-pulse"></span>
<div>
<span class="font-medium text-text">{}</span>
</div>
</div>
</div>
</div>"#,
coll_name, coll_name
)
} else {
serde_json::json!({
"type": "subscribed",
"collection": coll_name
})
.to_string()
};
if tx.send(Message::Text(msg.into())).await.is_err() {
return;
}
let (sub_tx, mut sub_rx) =
tokio::sync::mpsc::channel::<crate::storage::collection::ChangeEvent>(1000);
let req_key = req.key.clone();
let lag_notice = if use_htmx { None } else { Some(tx.clone()) };
let mut forwarders = tokio::task::JoinSet::new();
forwarders.spawn(forward_changes(
collection.change_sender.subscribe(),
sub_tx.clone(),
req_key.clone(),
coll_name.clone(),
lag_notice.clone(),
));
if let Some(shard_config) = collection.get_shard_config() {
if shard_config.num_shards > 0 {
if let Ok(database) = state.storage.get_database(&db_name) {
for shard_id in 0..shard_config.num_shards {
let physical_name = format!("{}_s{}", coll_name, shard_id);
if let Ok(physical_coll) = database.get_collection(&physical_name) {
forwarders.spawn(forward_changes(
physical_coll.change_sender.subscribe(),
sub_tx.clone(),
req_key.clone(),
physical_name,
lag_notice.clone(),
));
}
}
}
}
}
let is_local_only = req.local.unwrap_or(false);
if !is_local_only {
if let Some(shard_config) = collection.get_shard_config() {
if let Some(coordinator) = &state.shard_coordinator {
let my_addr = coordinator.my_address();
let all_nodes = coordinator.get_collection_nodes(&shard_config);
let cluster_secret = state.cluster_secret();
let mut remote_nodes = std::collections::HashSet::new();
for node_addr in all_nodes {
if node_addr != my_addr {
remote_nodes.insert(node_addr);
}
}
for node_addr in remote_nodes {
forwarders.spawn(forward_remote_changes(
node_addr,
db_name.clone(),
coll_name.clone(),
cluster_secret.clone(),
sub_tx.clone(),
));
}
}
}
}
drop(sub_tx);
let principal = crate::server::handlers::query::principal_from_claims(&claims);
loop {
let event = tokio::select! {
_ = tx.closed() => break,
event = sub_rx.recv() => match event {
Some(event) => event,
None => break,
},
};
if let Some(ref target_key) = req.key {
if &event.key != target_key {
continue;
}
}
if !changefeed_event_visible(
&state.storage,
&db_name,
&collection,
&principal,
&event,
)
.await
{
continue;
}
let msg_text = if use_htmx {
use crate::storage::collection::ChangeType;
let op_type = match event.type_ {
ChangeType::Insert => "INSERT",
ChangeType::Update => "UPDATE",
ChangeType::Delete => "DELETE",
ChangeType::Truncate => "TRUNCATE",
};
let status_class = match event.type_ {
ChangeType::Insert => "bg-success/10 text-success",
ChangeType::Update => "bg-warning/10 text-warning",
ChangeType::Delete => "bg-error/10 text-error",
ChangeType::Truncate => "bg-error/10 text-error",
};
let data_str = event
.data
.as_ref()
.map(|v| v.to_string())
.unwrap_or_default();
format!(
r#"<div hx-swap-oob="afterbegin:#events-container">
<div class="px-4 py-2 border-b border-border/10 last:border-0 font-mono text-sm hover:bg-white/5 transition-colors">
<div class="flex items-center gap-2 mb-1">
<span class="px-1.5 py-0.5 rounded text-xs {}">{}</span>
<span class="text-text-dim text-xs">{}</span>
<span class="text-text-dim text-xs ml-auto">{}</span>
</div>
<pre class="text-text-muted text-xs overflow-x-auto">{}</pre>
</div>
</div>"#,
status_class,
op_type,
coll_name,
chrono::Local::now().format("%H:%M:%S"),
data_str
)
} else {
serde_json::json!({
"operation": event.type_,
"collection": coll_name,
"key": event.key,
"data": event.data
})
.to_string()
};
if tx.send(Message::Text(msg_text.into())).await.is_err() {
break;
}
}
}
Err(_) => {
let _ = tx
.send(Message::Text(
serde_json::json!({
"error": format!("Collection '{}' not found", coll_name)
})
.to_string()
.into(),
))
.await;
}
}
}
async fn handle_live_query_request(
req: ChangefeedRequest,
state: AppState,
claims: crate::server::auth::Claims,
tx: tokio::sync::mpsc::Sender<Message>,
) {
if let Some(query_str) = req.query {
let db_name = req.database.clone().unwrap_or("_system".to_string());
if let Err(e) = crate::server::authz_middleware::enforce(
&claims,
&state,
crate::server::authorization::PermissionAction::Read,
Some(&db_name),
)
.await
{
let mut response = serde_json::json!({ "error": e.to_string() });
if let Some(req_id) = &req.id {
response["id"] = serde_json::Value::String(req_id.clone());
}
let _ = tx.send(Message::Text(response.to_string().into())).await;
return;
}
match crate::sdbql::parser::parse(&query_str) {
Ok(query) => {
let mut dependencies = std::collections::HashSet::new();
for clause in &query.for_clauses {
dependencies.insert(clause.collection.clone());
}
if dependencies.is_empty() {
let _ = tx
.send(Message::Text(
serde_json::json!({
"error": "Live query must reference at least one collection"
})
.to_string()
.into(),
))
.await;
return;
}
let mut response = serde_json::json!({
"type": "subscribed",
"mode": "live_query",
"collections": dependencies
});
if let Some(req_id) = &req.id {
response["id"] = serde_json::Value::String(req_id.clone());
}
if tx
.send(Message::Text(response.to_string().into()))
.await
.is_err()
{
return;
}
let (dep_tx, mut dep_rx) =
tokio::sync::mpsc::channel::<crate::storage::collection::ChangeEvent>(1000);
let mut forwarders = tokio::task::JoinSet::new();
for coll_name in &dependencies {
let coll_name = coll_name.clone();
if let Ok(collection) = state
.storage
.get_database(&db_name)
.and_then(|db| db.get_collection(&coll_name))
{
forwarders.spawn(forward_changes(
collection.change_sender.subscribe(),
dep_tx.clone(),
None,
coll_name.clone(),
None,
));
if let Some(shard_config) = collection.get_shard_config() {
if shard_config.num_shards > 0 {
if let Ok(database) = state.storage.get_database(&db_name) {
for shard_id in 0..shard_config.num_shards {
let physical_name = format!("{}_s{}", coll_name, shard_id);
if let Ok(physical_coll) =
database.get_collection(&physical_name)
{
forwarders.spawn(forward_changes(
physical_coll.change_sender.subscribe(),
dep_tx.clone(),
None,
physical_name,
None,
));
}
}
}
}
}
let is_local_only = req.local.unwrap_or(false);
if !is_local_only {
if let Some(shard_config) = collection.get_shard_config() {
if let Some(coordinator) = &state.shard_coordinator {
let my_addr = coordinator.my_address();
let all_nodes = coordinator.get_collection_nodes(&shard_config);
let cluster_secret = state.cluster_secret();
let mut remote_nodes = std::collections::HashSet::new();
for node_addr in all_nodes {
if node_addr != my_addr {
remote_nodes.insert(node_addr);
}
}
for node_addr in remote_nodes {
forwarders.spawn(forward_remote_changes(
node_addr,
db_name.clone(),
coll_name.clone(),
cluster_secret.clone(),
dep_tx.clone(),
));
}
}
}
}
}
}
drop(dep_tx);
let principal = crate::server::handlers::query::principal_from_claims(&claims);
if !execute_live_query_step(
&tx,
state.storage.clone(),
query_str.clone(),
db_name.clone(),
state.shard_coordinator.clone(),
principal.clone(),
req.id.clone(),
)
.await
{
return;
}
const DEBOUNCE: std::time::Duration = std::time::Duration::from_millis(150);
const MAX_DELAY: std::time::Duration = std::time::Duration::from_millis(500);
'reactive: loop {
tokio::select! {
_ = tx.closed() => break 'reactive,
first = dep_rx.recv() => {
if first.is_none() {
break 'reactive; }
}
}
let deadline = tokio::time::Instant::now() + MAX_DELAY;
loop {
tokio::select! {
more = dep_rx.recv() => {
if more.is_none() {
break 'reactive; }
if tokio::time::Instant::now() >= deadline {
break; }
}
_ = tokio::time::sleep(DEBOUNCE) => break,
}
}
if !execute_live_query_step(
&tx,
state.storage.clone(),
query_str.clone(),
db_name.clone(),
state.shard_coordinator.clone(),
principal.clone(),
req.id.clone(),
)
.await
{
break;
}
}
}
Err(e) => {
let _ = tx
.send(Message::Text(
serde_json::json!({
"error": format!("Invalid SDBQL query: {}", e)
})
.to_string()
.into(),
))
.await;
}
}
} else {
let _ = tx
.send(Message::Text(
serde_json::json!({
"error": "Missing 'query' field for live_query"
})
.to_string()
.into(),
))
.await;
}
}
async fn execute_live_query_step(
tx: &tokio::sync::mpsc::Sender<Message>,
storage: Arc<StorageEngine>,
query_str: String,
db_name: String,
shard_coordinator: Option<Arc<crate::sharding::ShardCoordinator>>,
principal: crate::sdbql::QueryPrincipal,
req_id: Option<String>,
) -> bool {
let exec_result = tokio::task::spawn_blocking(move || {
match crate::sdbql::parser::parse(&query_str) {
Ok(parsed) => {
if parsed.has_mutations() {
return Err(crate::error::DbError::ExecutionError(
"Live queries are read-only".to_string(),
));
}
let mut executor =
crate::sdbql::executor::QueryExecutor::with_database(&storage, db_name)
.with_principal(principal)
.with_timeout(std::time::Duration::from_secs(30));
if let Some(coord) = shard_coordinator {
executor = executor.with_shard_coordinator(coord);
}
executor.execute(&parsed)
}
Err(e) => Err(crate::error::DbError::ParseError(e.to_string())),
}
})
.await
.unwrap_or_else(|e| {
Err(crate::error::DbError::InternalError(format!(
"Live query task failed: {}",
e
)))
});
match exec_result {
Ok(results) => {
let mut response = serde_json::json!({
"type": "query_result",
"result": results
});
if let Some(id) = req_id {
response["id"] = serde_json::Value::String(id);
}
tx.send(Message::Text(response.to_string().into()))
.await
.is_ok()
}
Err(e) => {
let mut response = serde_json::json!({
"type": "error",
"error": e.to_string()
});
if let Some(id) = req_id {
response["id"] = serde_json::Value::String(id);
}
tx.send(Message::Text(response.to_string().into()))
.await
.is_ok()
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::server::auth::Claims;
fn engine() -> (StorageEngine, tempfile::TempDir) {
let tmp = tempfile::TempDir::new().expect("tempdir");
let engine = StorageEngine::new(tmp.path().to_str().unwrap()).expect("engine");
(engine, tmp)
}
fn claims(sub: &str, exp: usize, livequery: Option<bool>) -> Claims {
Claims {
sub: sub.to_string(),
exp,
livequery,
roles: None,
scoped_databases: None,
}
}
#[test]
fn row_policy_filters_changefeed_rows_by_principal() {
let (engine, _tmp) = engine();
engine.create_database("app".to_string()).unwrap();
let alice = crate::sdbql::QueryPrincipal::from_roles("alice", vec!["viewer".into()]);
let policy = "doc.owner == CURRENT_USER";
let own = serde_json::json!({"_key": "1", "owner": "alice"});
let other = serde_json::json!({"_key": "2", "owner": "bob"});
assert!(row_policy_allows(
&engine,
"app".into(),
"orders",
&alice,
policy,
own
));
assert!(!row_policy_allows(
&engine,
"app".into(),
"orders",
&alice,
policy,
other.clone()
));
assert!(!row_policy_allows(
&engine,
"app".into(),
"orders",
&alice,
"((",
other
));
}
#[test]
fn ws_credential_follows_user_existence_and_expiry() {
let (engine, _tmp) = engine();
engine.create_database("_system".to_string()).unwrap();
let system = engine.get_database("_system").unwrap();
system
.create_collection(crate::server::auth::ADMIN_COLL.to_string(), None)
.unwrap();
system
.system_collection(crate::server::auth::ADMIN_COLL)
.unwrap()
.insert(serde_json::json!({"_key": "ws_cred_alice", "password_hash": "x"}))
.unwrap();
let far = usize::MAX - 1;
assert!(check_ws_credential(&claims("ws_cred_alice", far, None), &engine).is_ok());
assert!(check_ws_credential(&claims("ws_cred_bob", far, None), &engine).is_err());
assert!(check_ws_credential(&claims("ws_cred_alice", 1, None), &engine).is_err());
assert!(check_ws_credential(&claims("ws_cred_alice", 1, Some(true)), &engine).is_ok());
assert!(check_ws_credential(&claims("ws_cred_bob", 1, Some(true)), &engine).is_err());
let mut stale = claims("ws_cred_alice", far, None);
stale.roles = Some(vec!["admin".to_string()]);
assert!(check_ws_credential(&stale, &engine).is_err());
assert!(check_ws_credential(&claims("api-key:ws_cred_gone", far, None), &engine).is_err());
}
}