use std::sync::Arc;
use surrealdb_rpc::{QUERY_STREAM_BUFFER, QueryResult};
use surrealdb_types::Value as PublicValue;
use crate::catalog::providers::CatalogProvider;
use crate::dbs::{QueryStreamItem, Session};
use crate::kvs::{Datastore, QueryRequest, TransactionType};
#[derive(Default)]
struct Statement {
values: Vec<PublicValue>,
single: Option<PublicValue>,
finished: Option<QueryResult>,
}
impl Statement {
fn finish(self) -> QueryResult {
let mut result = self.finished.expect("every statement is terminated by a Finished item");
if result.result.is_ok() {
result.result = Ok(match self.single {
Some(value) => value,
None => PublicValue::Array(self.values.into()),
});
}
result
}
}
async fn datastore() -> Arc<Datastore> {
let ds = Datastore::new("memory").await.expect("a memory datastore");
let tx = ds.transaction(TransactionType::Write).await.expect("a write transaction");
tx.ensure_ns_db(None, "test", "test").await.expect("the test namespace and database");
tx.commit().await.expect("committing the namespace");
for result in ds.execute(SEED, &session(), None).await.expect("the seed runs") {
result.result.expect("every seed statement succeeds");
}
ds
}
fn session() -> Session {
Session::owner().with_ns("test").with_db("test")
}
async fn streamed(ds: &Arc<Datastore>, sql: &str) -> Result<Vec<QueryResult>, String> {
let (tx, rx) = async_channel::bounded(QUERY_STREAM_BUFFER);
let job =
ds.run_streaming(QueryRequest::new(sql, &session()), tx).map_err(|e| e.to_string())?;
let drain = async {
let mut statements: Vec<Option<Statement>> = Vec::new();
while let Ok(item) = rx.recv().await {
let index = item.index();
if index >= statements.len() {
statements.resize_with(index + 1, || None);
}
let statement = statements[index].get_or_insert_with(Statement::default);
match item {
QueryStreamItem::Rows {
values,
..
} => statement.values.extend(values),
QueryStreamItem::Value {
value,
..
} => statement.single = Some(value),
QueryStreamItem::Finished {
time,
query_type,
error,
..
} => {
statement.finished = Some(QueryResult {
time,
result: match error {
Some(error) => Err(error),
None => Ok(PublicValue::None),
},
query_type,
});
}
}
}
statements
};
let (outcome, statements) = futures::future::join(job.run, drain).await;
outcome.map_err(|e| e.to_string())?;
Ok(statements.into_iter().flatten().map(Statement::finish).collect())
}
fn same(streamed: &[QueryResult], buffered: &[QueryResult], sql: &str) {
assert_eq!(
streamed.len(),
buffered.len(),
"{sql}: streamed produced {} results, buffered {}",
streamed.len(),
buffered.len()
);
for (index, (s, b)) in streamed.iter().zip(buffered).enumerate() {
assert_eq!(s.query_type, b.query_type, "{sql}: statement {index} query type");
match (&s.result, &b.result) {
(Ok(s), Ok(b)) => assert_eq!(s, b, "{sql}: statement {index} value"),
(Err(s), Err(b)) => {
assert_eq!(s.to_string(), b.to_string(), "{sql}: statement {index} error")
}
(s, b) => panic!("{sql}: statement {index} differs: streamed {s:?}, buffered {b:?}"),
}
}
}
async fn assert_parity(sql: &str) {
let ds = datastore().await;
let streamed = streamed(&ds, sql).await;
let ds = datastore().await;
let buffered = ds.execute(sql, &session(), None).await;
match (streamed, buffered) {
(Ok(streamed), Ok(buffered)) => same(&streamed, &buffered, sql),
(Err(streamed), Err(buffered)) => {
assert_eq!(streamed, buffered.to_string(), "{sql}: both paths must fail the same way")
}
(s, b) => panic!("{sql}: one path failed and the other did not: {s:?} vs {b:?}"),
}
}
const SEED: &str = "CREATE thing:1 SET n = 1; CREATE thing:2 SET n = 2; CREATE thing:3 SET n = 3;";
#[tokio::test]
async fn streaming_matches_buffered() {
for sql in [
"SELECT * FROM thing",
"SELECT * FROM thing ORDER BY n DESC",
"SELECT * FROM thing WHERE n > 99",
"SELECT * FROM ONLY thing:1",
"RETURN 1 + 2",
"SELECT n FROM thing; RETURN 'done'; SELECT * FROM thing WHERE n = 2",
"CREATE thing:4 SET n = 4",
"UPDATE thing:1 SET n = 10",
"DELETE thing:1",
"INFO FOR DB",
"LET $x = SELECT * FROM thing; RETURN $x",
"SELECT * FROM thing; THROW 'nope'; SELECT * FROM thing WHERE n = 1",
"SELECT count() FROM thing GROUP ALL",
"BEGIN; CREATE thing:5 SET n = 5; SELECT * FROM thing; COMMIT;",
"BEGIN; CREATE thing:6 SET n = 6; CANCEL;",
"BEGIN; SELECT * FROM thing; THROW 'nope'; SELECT * FROM thing; COMMIT;",
"BEGIN; SELECT * FROM thing; RETURN 'early'; SELECT * FROM thing; COMMIT;",
] {
assert_parity(sql).await;
}
}
#[tokio::test]
async fn a_rolled_back_block_retracts_its_rows() {
let ds = datastore().await;
let (tx, rx) = async_channel::bounded(QUERY_STREAM_BUFFER);
let job = ds
.run_streaming(
QueryRequest::new("BEGIN; SELECT * FROM thing; THROW 'nope'; COMMIT;", &session()),
tx,
)
.expect("the query parses");
let drain = async {
let mut items = Vec::new();
while let Ok(item) = rx.recv().await {
items.push(item);
}
items
};
let (outcome, items) = futures::future::join(job.run, drain).await;
outcome.expect("the execution itself succeeds; the statement is what failed");
let rows: Vec<&QueryStreamItem> =
items.iter().filter(|i| matches!(i, QueryStreamItem::Rows { .. })).collect();
assert!(!rows.is_empty(), "the SELECT's rows go out before the block resolves");
let first_finished = items
.iter()
.position(|i| matches!(i, QueryStreamItem::Finished { .. }))
.expect("the block resolves");
let last_row = items
.iter()
.rposition(|i| matches!(i, QueryStreamItem::Rows { .. }))
.expect("rows were sent");
assert!(last_row < first_finished, "rows are provisional until the block resolves");
let select_finished = items
.iter()
.find(|i| {
matches!(
i,
QueryStreamItem::Finished {
index: 1,
..
}
)
})
.expect("the SELECT is statement 1");
assert!(
matches!(
select_finished,
QueryStreamItem::Finished {
error: Some(_),
..
}
),
"the rolled-back SELECT must be reported as failed, got {select_finished:?}"
);
}
#[tokio::test]
async fn each_statement_is_terminated_exactly_once() {
let ds = datastore().await;
let (tx, rx) = async_channel::bounded(QUERY_STREAM_BUFFER);
let job = ds
.run_streaming(
QueryRequest::new("SELECT * FROM thing; RETURN 1; SELECT * FROM thing", &session()),
tx,
)
.expect("the query parses");
assert_eq!(job.statement_count, 3, "statement_count counts what parsed");
let drain = async {
let mut items = Vec::new();
while let Ok(item) = rx.recv().await {
items.push(item);
}
items
};
let (outcome, items) = futures::future::join(job.run, drain).await;
outcome.expect("the query succeeds");
for index in 0..3 {
let terminals = items
.iter()
.filter(|i| matches!(i, QueryStreamItem::Finished { index: i, .. } if *i == index))
.count();
assert_eq!(terminals, 1, "statement {index} must be terminated exactly once");
let last =
items.iter().rposition(|i| i.index() == index).expect("every statement produced items");
assert!(
matches!(items[last], QueryStreamItem::Finished { .. }),
"statement {index}'s last item must be its terminal one"
);
let payload = items
.iter()
.find(|i| i.index() == index && !matches!(i, QueryStreamItem::Finished { .. }))
.expect("every statement here produces a value");
match index {
1 => assert!(
matches!(payload, QueryStreamItem::Value { .. }),
"RETURN 1 is a single value, not rows"
),
_ => assert!(
matches!(payload, QueryStreamItem::Rows { .. }),
"a SELECT streams its rows, got {payload:?}"
),
}
}
}
#[tokio::test]
async fn dropping_the_consumer_stops_the_execution() {
let ds = datastore().await;
let (tx, rx) = async_channel::bounded(QUERY_STREAM_BUFFER);
let job =
ds.run_streaming(QueryRequest::new("SELECT * FROM thing", &session()), tx).expect("parses");
drop(rx);
let outcome = tokio::time::timeout(std::time::Duration::from_secs(10), job.run).await.expect(
"a dropped consumer must end the execution rather than blocking it on a full channel",
);
outcome.expect("abandoning the results is not itself an error");
}
#[tokio::test]
async fn a_dropped_consumer_stops_later_statements() {
let ds = datastore().await;
for result in ds
.execute(
"DEFINE TABLE marker;
FOR $i IN array::range(0, 5000) { CREATE type::record('big', $i) SET n = $i }",
&session(),
None,
)
.await
.expect("the seed runs")
{
result.result.expect("seeding succeeds");
}
let (tx, rx) = async_channel::bounded(QUERY_STREAM_BUFFER);
let job = ds
.run_streaming(QueryRequest::new("SELECT * FROM big; CREATE marker:1", &session()), tx)
.expect("the query parses");
let running = tokio::spawn(job.run);
let first = rx.recv().await.expect("at least one item");
assert!(matches!(first, QueryStreamItem::Rows { .. }), "the SELECT streams first");
drop(rx);
running
.await
.expect("the execution task completes")
.expect("abandoning the results is not itself an error");
let mut created = ds
.execute("SELECT VALUE id FROM marker", &session(), None)
.await
.expect("the follow-up query runs");
let rows = created.remove(0).result.expect("a result");
assert_eq!(
rows,
PublicValue::Array(Vec::<PublicValue>::new().into()),
"a write after the abandoned statement must not have run"
);
}