use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::sync::atomic::Ordering;
use std::time::Duration;
use axum::extract::ws::Message;
use surrealdb_core::channel::{Receiver, bounded};
use surrealdb_core::ctx::CancelHandle;
use surrealdb_core::dbs::AuthPrincipalSnapshot;
use surrealdb_core::rpc::{RpcProtocol, live_query_owner};
use surrealdb_rpc::capabilities::MethodTarget;
use surrealdb_rpc::error::{invalid_params, method_not_allowed, stream_exists, too_many_streams};
use surrealdb_rpc::framing::{live_queries_disowned, stream_stopped};
use surrealdb_rpc::{
DbResponse, DbResult, Method, QUERY_STREAM_BUFFER, QueryResult, QueryStreamFrame,
QueryStreamItem, QueryType, StreamFrames,
};
use surrealdb_types::{Array, Error as TypesError, ToSql, Value};
use tokio::sync::mpsc::Sender;
use tokio::time::Instant;
use uuid::Uuid;
use crate::cnf::{WEBSOCKET_MAX_CONCURRENT_STREAMS, WEBSOCKET_STREAM_SEND_TIMEOUT_SECS};
use crate::rpc::format::WsFormat;
use crate::rpc::websocket::Websocket;
const TERMINAL_FRAME_GRACE: Duration = Duration::from_secs(1);
type QueryStreamRun = Pin<Box<dyn Future<Output = Result<Vec<QueryResult>, TypesError>> + Send>>;
pub(crate) struct StreamHandle {
cancel: CancelHandle,
items: Receiver<QueryStreamItem>,
}
impl StreamHandle {
pub(crate) fn stop(&self) {
self.cancel.trip();
self.items.close();
}
}
struct StreamSlot<'a> {
rpc: &'a Websocket,
}
impl<'a> StreamSlot<'a> {
fn claim(rpc: &'a Websocket) -> Option<Self> {
rpc.stream_slots
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |held| {
(held < *WEBSOCKET_MAX_CONCURRENT_STREAMS).then_some(held + 1)
})
.ok()
.map(|_| Self {
rpc,
})
}
}
impl Drop for StreamSlot<'_> {
fn drop(&mut self) {
self.rpc.stream_slots.fetch_sub(1, Ordering::AcqRel);
}
}
struct StreamRegistration<'a> {
rpc: &'a Websocket,
key: String,
_slot: StreamSlot<'a>,
}
impl Drop for StreamRegistration<'_> {
fn drop(&mut self) {
self.rpc.streams.remove(&self.key);
}
}
pub(crate) async fn process_query_stream(
rpc: &Arc<Websocket>,
id: Option<Value>,
session_id: Uuid,
client_session: Option<Uuid>,
txn: Option<Uuid>,
params: Array,
chn: Sender<Message>,
) {
let fmt = rpc.format;
let Some(id) = id else {
let error = invalid_params("The query_stream method requires a request id");
crate::rpc::response::send(DbResponse::failure(None, client_session, error), fmt, chn)
.await;
return;
};
if let Err(error) = preflight(rpc) {
crate::rpc::response::send(DbResponse::failure(Some(id), client_session, error), fmt, chn)
.await;
return;
}
let key = id.to_sql();
let cancel = CancelHandle::new();
let (items_tx, items_rx) = bounded(QUERY_STREAM_BUFFER);
let handle = StreamHandle {
cancel: cancel.clone(),
items: items_rx.clone(),
};
let reserved = match rpc.streams.entry(key.clone()) {
dashmap::mapref::entry::Entry::Occupied(_) => Err(stream_exists()),
dashmap::mapref::entry::Entry::Vacant(entry) => match StreamSlot::claim(rpc) {
Some(slot) => {
entry.insert(handle);
Ok(slot)
}
None => Err(too_many_streams()),
},
};
let slot = match reserved {
Ok(slot) => slot,
Err(error) => {
crate::rpc::response::send(
DbResponse::failure(Some(id), client_session, error),
fmt,
chn,
)
.await;
return;
}
};
let registration = StreamRegistration {
rpc,
key,
_slot: slot,
};
let deadline = rpc.kvs().query_timeout().map(|timeout| Instant::now() + timeout);
let (job, principal) = match RpcProtocol::query_stream(
rpc.as_ref(),
txn,
session_id,
params,
Some(cancel.clone()),
items_tx,
)
.await
{
Ok(started) => started,
Err(error) => {
drop(registration);
crate::rpc::response::send(
DbResponse::failure(Some(id), client_session, error),
fmt,
chn,
)
.await;
return;
}
};
let mut driver = StreamDriver {
rpc,
id,
client_session,
session_id,
chn,
cancel,
items: items_rx,
principal,
deadline,
muted: false,
failure: None,
answered: false,
};
driver.serve(job.statement_count, job.run).await;
drop(registration);
}
fn preflight(rpc: &Websocket) -> Result<(), TypesError> {
if !rpc.kvs().allows_rpc_method(&MethodTarget {
method: Method::QueryStream,
}) {
warn!("Capabilities denied RPC method call attempt, target: 'query_stream'");
return Err(method_not_allowed(Method::QueryStream.to_string()));
}
Ok(())
}
struct StreamDriver<'a> {
rpc: &'a Arc<Websocket>,
id: Value,
client_session: Option<Uuid>,
session_id: Uuid,
chn: Sender<Message>,
cancel: CancelHandle,
items: Receiver<QueryStreamItem>,
principal: AuthPrincipalSnapshot,
deadline: Option<Instant>,
muted: bool,
failure: Option<TypesError>,
answered: bool,
}
impl StreamDriver<'_> {
async fn serve(&mut self, statement_count: usize, run: QueryStreamRun) {
let started = Instant::now();
let conn = self.rpc.cancel.token();
let mut run = Some(run);
let mut frames = StreamFrames::new();
let mut outcome = None;
self.send(
QueryStreamFrame::Begin {
statements: statement_count,
},
&mut frames,
)
.await;
let mut conn_observed = false;
let outcome = loop {
while let Some(frame) = frames.pop() {
self.send(frame, &mut frames).await;
}
if let Some(outcome) = outcome {
break outcome;
}
let running = run.as_mut().expect("the execution is driven until it completes");
tokio::select! {
biased;
_ = conn.cancelled(), if !conn_observed => {
conn_observed = true;
self.stop_sending();
}
item = self.items.recv() => match item {
Ok(item) => frames.absorb(item),
Err(_) => {
outcome = Some(run.take().expect("still running").await);
}
},
result = running => {
run = None;
while let Ok(item) = self.items.try_recv() {
frames.absorb(item);
}
outcome = Some(result);
}
}
};
let error = match &outcome {
Ok(results) => {
let disowned = self.settle_live_queries(Some(results), &frames).await;
self.failure
.take()
.or_else(|| (!disowned.is_empty()).then(|| live_queries_disowned(&disowned)))
.or_else(|| self.cancel.is_cancelled().then(stream_stopped))
}
Err(error) => {
self.settle_live_queries(None, &frames).await;
Some(error.clone())
}
};
let end = QueryStreamFrame::End {
results: frames.delivered_count(),
time: started.elapsed(),
error,
};
self.send(end, &mut frames).await;
}
async fn send(&mut self, frame: QueryStreamFrame, frames: &mut StreamFrames) {
if self.answered {
return;
}
let is_end = matches!(&frame, QueryStreamFrame::End { .. });
let is_begin = matches!(&frame, QueryStreamFrame::Begin { .. });
if is_end {
let floor = Instant::now() + TERMINAL_FRAME_GRACE;
self.deadline = Some(match self.deadline {
Some(deadline) => deadline.max(floor),
None => floor,
});
} else if self.muted {
return;
}
let statement = match &frame {
QueryStreamFrame::Rows {
index,
..
}
| QueryStreamFrame::Value {
index,
..
} => Some(*index),
_ => None,
};
let terminal = match &frame {
QueryStreamFrame::Finished {
index,
..
} => Some(*index),
_ => None,
};
let response = DbResponse::success(
Some(self.id.clone()),
self.client_session,
DbResult::Other(frame.into_value()),
);
let message = match self.rpc.format.res_ws(response) {
Ok((_len, message)) => message,
Err(error) => {
if let Some(index) = statement
&& frames.retract(index, error.clone())
{
return;
}
self.fail_stream(error.clone());
self.answered = true;
let failure =
DbResponse::failure(Some(self.id.clone()), self.client_session, error);
if let Ok((_len, message)) = self.rpc.format.res_ws(failure) {
let conn = self.rpc.cancel.token();
let send = async {
tokio::select! {
biased;
_ = conn.cancelled() => false,
sent = self.chn.send(message) => sent.is_ok(),
}
};
let deadline = Instant::now() + TERMINAL_FRAME_GRACE;
let _ = tokio::time::timeout_at(deadline, send).await;
}
return;
}
};
let conn = self.rpc.cancel.token();
let stream = self.cancel.token();
let brackets = is_end || is_begin;
let send = async {
tokio::select! {
biased;
_ = conn.cancelled() => false,
_ = stream.cancelled(), if !brackets => false,
sent = self.chn.send(message) => sent.is_ok(),
}
};
let send_timeout = Duration::from_secs(*WEBSOCKET_STREAM_SEND_TIMEOUT_SECS);
let now = Instant::now();
let send_deadline = now.checked_add(send_timeout).unwrap_or(now);
let deadline = match self.deadline {
Some(deadline) if is_end => deadline.max(send_deadline),
Some(deadline) => deadline.min(send_deadline),
None => send_deadline,
};
let sent = tokio::time::timeout_at(deadline, send).await.unwrap_or_default();
if sent {
if let Some(index) = terminal {
frames.mark_delivered(index);
}
} else {
self.stop_sending();
}
}
fn stop_sending(&mut self) {
self.cancel.trip();
self.items.close();
self.muted = true;
}
fn fail_stream(&mut self, error: TypesError) {
self.failure = Some(error);
self.stop_sending();
}
async fn settle_live_queries(
&self,
results: Option<&[QueryResult]>,
frames: &StreamFrames,
) -> Vec<Uuid> {
let live: Vec<(usize, Uuid)> = match results {
Some(results) => results
.iter()
.enumerate()
.filter(|(_, r)| matches!(r.query_type, QueryType::Live))
.filter_map(|(index, r)| match &r.result {
Ok(Value::Uuid(id)) => Some((index, id.into_inner())),
_ => None,
})
.collect(),
None => frames.live_queries().to_vec(),
};
if live.is_empty() {
return Vec::new();
}
let session =
live_query_owner(self.rpc.session_map(), self.session_id, &self.principal).await;
let mut orphans = Vec::new();
let mut disowned = Vec::new();
for (index, id) in &live {
let told = frames.was_delivered(*index);
match &session {
Some(owner) if told => {
self.rpc
.handle_live(id, self.session_id, owner.ns.clone(), owner.db.clone())
.await;
}
_ => {
if told {
disowned.push(*id);
}
orphans.push(*id);
}
}
}
drop(session);
if !orphans.is_empty()
&& let Err(err) = self.rpc.kvs().delete_queries(orphans).await
{
error!("Error cleaning up the live queries of a streaming query: {err}");
}
disowned
}
}