use std::sync::Arc;
use super::pipeline::{FieldState, build_field_state, materialise_fields_with_permissions};
use crate::catalog::providers::TableProvider;
use crate::catalog::{DatabaseId, NamespaceId, table_select_permission};
use crate::exec::{ControlFlowExt, EvalContext, ExecutionContext, PhysicalExpr};
use crate::expr::{ControlFlow, FlowResult};
use crate::kvs::{CachePolicy, DatastoreError, Transaction};
use crate::val::{RecordId, RecordIdKey, TableName, Value};
pub(crate) fn ensure_below_memory_threshold() -> FlowResult<()> {
if crate::mem::ALLOC.is_beyond_threshold() {
return Err(ControlFlow::Err(anyhow::Error::new(
DatastoreError::QueryBeyondMemoryThreshold,
)));
}
Ok(())
}
pub(crate) struct FetchedBatch {
pub(crate) values: Vec<Value>,
pub(crate) fetched_bytes: usize,
pub(crate) fetched_rows: usize,
}
const PROBE_ROW_BYTES: usize = 64 << 10;
pub(crate) fn probe_batch_len(budget: usize, max_rows: usize) -> usize {
(budget / PROBE_ROW_BYTES).clamp(1, max_rows)
}
pub(crate) fn next_batch_len(
fetched_bytes: usize,
fetched_rows: usize,
current: usize,
budget: usize,
max_rows: usize,
) -> usize {
if fetched_rows == 0 {
return current;
}
let per_row = (fetched_bytes / fetched_rows).max(1);
(budget / per_row).min(current.saturating_mul(2)).clamp(1, max_rows)
}
pub(crate) fn approx_value_size(value: &Value) -> usize {
let inline = std::mem::size_of::<Value>();
match value {
Value::String(s) => inline + s.len(),
Value::Bytes(b) => inline + b.0.len(),
Value::Array(a) => inline + a.iter().map(approx_value_size).sum::<usize>(),
Value::Set(s) => inline + s.0.iter().map(approx_value_size).sum::<usize>(),
Value::Object(o) => {
let entry = std::mem::size_of::<surrealdb_strand::Strand>();
inline + o.iter().map(|(k, v)| entry + k.len() + approx_value_size(v)).sum::<usize>()
}
Value::Geometry(g) => inline + geometry_heap_size(g),
Value::RecordId(rid) => inline + rid.table.as_str().len() + approx_key_size(&rid.key),
Value::Range(r) => {
inline
+ std::mem::size_of::<crate::val::Range>()
+ approx_bound_size(&r.start, approx_value_size)
+ approx_bound_size(&r.end, approx_value_size)
}
_ => inline,
}
}
fn approx_key_size(key: &RecordIdKey) -> usize {
match key {
RecordIdKey::Number(_) | RecordIdKey::Uuid(_) => 0,
RecordIdKey::String(s) => s.len(),
RecordIdKey::Array(a) => a.iter().map(approx_value_size).sum(),
RecordIdKey::Object(o) => o
.iter()
.map(|(k, v)| {
std::mem::size_of::<surrealdb_strand::Strand>() + k.len() + approx_value_size(v)
})
.sum(),
RecordIdKey::Range(r) => {
std::mem::size_of::<crate::val::RecordIdKeyRange>()
+ approx_bound_size(&r.start, approx_key_size)
+ approx_bound_size(&r.end, approx_key_size)
}
}
}
fn approx_bound_size<T>(bound: &std::ops::Bound<T>, size: fn(&T) -> usize) -> usize {
match bound {
std::ops::Bound::Included(v) | std::ops::Bound::Excluded(v) => size(v),
std::ops::Bound::Unbounded => 0,
}
}
fn geometry_heap_size(geometry: &crate::val::Geometry) -> usize {
use std::mem::size_of;
use crate::val::Geometry;
fn line(line: &geo::LineString<f64>) -> usize {
size_of::<geo::Coord<f64>>() * line.0.len()
}
fn polygon(polygon: &geo::Polygon<f64>) -> usize {
line(polygon.exterior())
+ polygon
.interiors()
.iter()
.map(|ring| size_of::<geo::LineString<f64>>() + line(ring))
.sum::<usize>()
}
match geometry {
Geometry::Point(_) => 0,
Geometry::Line(l) => line(l),
Geometry::Polygon(p) => polygon(p),
Geometry::MultiPoint(m) => size_of::<geo::Point<f64>>() * m.0.len(),
Geometry::MultiLine(m) => {
m.0.iter().map(|l| size_of::<geo::LineString<f64>>() + line(l)).sum()
}
Geometry::MultiPolygon(m) => {
m.0.iter().map(|p| size_of::<geo::Polygon<f64>>() + polygon(p)).sum()
}
Geometry::Collection(c) => {
c.iter().map(|g| size_of::<Geometry>() + geometry_heap_size(g)).sum()
}
}
}
pub(crate) fn value_to_record_id_key(val: Value) -> RecordIdKey {
match val {
Value::Number(n) => RecordIdKey::Number(n.as_int()),
Value::String(s) => RecordIdKey::String(s),
Value::Uuid(u) => RecordIdKey::Uuid(u),
Value::Array(a) => RecordIdKey::Array(a),
Value::Object(o) => RecordIdKey::Object(o),
other => RecordIdKey::String(other.to_raw_string().into()),
}
}
pub(crate) fn extract_record_ids_into(val: Value, rids: &mut Vec<RecordId>) {
match val {
Value::RecordId(rid) => rids.push(rid),
Value::Object(mut obj) => {
if let Some(id_val) = obj.remove("id") {
extract_record_ids_into(id_val, rids);
}
}
Value::Array(arr) => {
for v in arr {
extract_record_ids_into(v, rids);
}
}
_ => {}
}
}
pub(crate) async fn evaluate_bound_key(
expr: &Arc<dyn PhysicalExpr>,
ctx: &ExecutionContext,
) -> Result<RecordIdKey, ControlFlow> {
let eval_ctx = EvalContext::from_exec_ctx(ctx);
let val = expr.evaluate(eval_ctx).await?;
Ok(value_to_record_id_key(val))
}
pub(crate) async fn resolve_version_stamp(
ctx: &ExecutionContext,
version_expr: Option<&Arc<dyn PhysicalExpr>>,
) -> Result<Option<u64>, ControlFlow> {
if let Some(stamp) = ctx.version_stamp() {
return Ok(Some(stamp));
}
let Some(expr) = version_expr else {
return Ok(None);
};
let eval_ctx = EvalContext::from_exec_ctx(ctx);
let v = expr.evaluate(eval_ctx).await?;
let stamp = v
.cast_to::<crate::val::Datetime>()
.map_err(|e| anyhow::anyhow!("{e}"))?
.to_version_stamp(ctx.txn().timestamp_impl().as_ref())?;
Ok(Some(stamp))
}
#[allow(clippy::too_many_arguments)]
pub(crate) async fn resolve_record_batch(
ctx: &ExecutionContext,
txn: &Transaction,
ns_id: NamespaceId,
db_id: DatabaseId,
rids: &[RecordId],
fetch_full: bool,
check_perms: bool,
version: Option<u64>,
cache_policy: CachePolicy,
perm_cache: &mut std::collections::HashMap<
surrealdb_strand::TableName,
crate::exec::permission::PhysicalPermission,
>,
) -> Result<Vec<Value>, ControlFlow> {
use crate::exec::permission::{PhysicalPermission, check_permission_for_value};
if !check_perms && !fetch_full {
return Ok(rids.iter().map(|rid| Value::RecordId(rid.clone())).collect());
}
if check_perms {
let db_ctx = ctx.database().context("permission resolution requires database context")?;
for rid in rids {
if perm_cache.contains_key(&rid.table) {
continue;
}
let table_def = db_ctx
.get_table_def(&rid.table, version)
.await
.context("Failed to get table definition")?;
let catalog_perm = table_select_permission(table_def.as_deref());
let perm =
crate::exec::permission::convert_permission_to_physical_runtime(catalog_perm, ctx)
.await
.context("Failed to convert permission")?;
perm_cache.insert(rid.table.clone(), perm);
}
}
let records = txn
.get_records(ns_id, db_id, rids, version, cache_policy)
.await
.context("Failed to fetch records")?;
ensure_below_memory_threshold()?;
let mut field_state_cache: std::collections::HashMap<TableName, FieldState> =
std::collections::HashMap::new();
let skip_fetch_perms = ctx.root().skip_fetch_perms;
let mut values = Vec::with_capacity(rids.len());
for (rid, record) in rids.iter().zip(records) {
if record.data.is_none() {
continue;
}
if check_perms {
let perm = perm_cache.get(&rid.table).map_or(&PhysicalPermission::Deny, |p| p);
let allowed = check_permission_for_value(perm, &record.data, None, ctx)
.await
.context("Failed to check permission")?;
if !allowed {
continue;
}
}
if fetch_full {
let mut value = match Arc::try_unwrap(record) {
Ok(rec) => rec.data,
Err(arc) => arc.data.clone(),
};
if !field_state_cache.contains_key(&rid.table) {
let fs = build_field_state(ctx, &rid.table, check_perms, None).await?;
field_state_cache.insert(rid.table.clone(), fs);
}
let field_state = &field_state_cache[&rid.table];
materialise_fields_with_permissions(
ctx,
field_state,
&mut value,
skip_fetch_perms,
check_perms,
)
.await?;
values.push(value);
} else {
values.push(Value::RecordId(rid.clone()));
}
}
if fetch_full {
ensure_below_memory_threshold()?;
}
Ok(values)
}
#[allow(clippy::too_many_arguments)]
pub(crate) async fn fetch_and_filter_records_batch(
ctx: &ExecutionContext,
txn: &Transaction,
ns_id: NamespaceId,
db_id: DatabaseId,
rids: &[RecordId],
select_permission: &crate::exec::permission::PhysicalPermission,
check_perms: bool,
version: Option<u64>,
cache_policy: CachePolicy,
) -> Result<Vec<Value>, ControlFlow> {
let fetched = fetch_and_filter(
ctx,
txn,
ns_id,
db_id,
rids,
select_permission,
check_perms,
version,
cache_policy,
false,
)
.await?;
Ok(fetched.values)
}
#[allow(clippy::too_many_arguments)]
pub(crate) async fn fetch_and_filter_records_measured(
ctx: &ExecutionContext,
txn: &Transaction,
ns_id: NamespaceId,
db_id: DatabaseId,
rids: &[RecordId],
select_permission: &crate::exec::permission::PhysicalPermission,
check_perms: bool,
version: Option<u64>,
cache_policy: CachePolicy,
) -> Result<FetchedBatch, ControlFlow> {
fetch_and_filter(
ctx,
txn,
ns_id,
db_id,
rids,
select_permission,
check_perms,
version,
cache_policy,
true,
)
.await
}
#[allow(clippy::too_many_arguments)]
async fn fetch_and_filter(
ctx: &ExecutionContext,
txn: &Transaction,
ns_id: NamespaceId,
db_id: DatabaseId,
rids: &[RecordId],
select_permission: &crate::exec::permission::PhysicalPermission,
check_perms: bool,
version: Option<u64>,
cache_policy: CachePolicy,
measure: bool,
) -> Result<FetchedBatch, ControlFlow> {
let records = txn
.get_records(ns_id, db_id, rids, version, cache_policy)
.await
.context("Failed to fetch records")?;
ensure_below_memory_threshold()?;
let mut values = Vec::with_capacity(rids.len());
let mut fetched_bytes = 0;
let mut fetched_rows = 0;
let mut cloned = false;
for record in records {
if record.data.is_none() {
continue;
}
fetched_rows += 1;
if measure {
fetched_bytes += approx_value_size(&record.data);
}
if check_perms {
let allowed = crate::exec::permission::check_permission_for_value(
select_permission,
&record.data,
None,
ctx,
)
.await
.context("Failed to check permission")?;
if !allowed {
continue;
}
}
let value = match Arc::try_unwrap(record) {
Ok(rec) => rec.data,
Err(arc) => {
cloned = true;
arc.data.clone()
}
};
values.push(value);
}
if cloned {
ensure_below_memory_threshold()?;
}
Ok(FetchedBatch {
values,
fetched_bytes,
fetched_rows,
})
}
#[cfg(test)]
mod tests {
use super::*;
const MAX_ROWS: usize = 1000;
fn bytes(rows: &[Value]) -> usize {
rows.iter().map(approx_value_size).sum()
}
#[test]
fn a_record_id_counts_its_key() {
let key = RecordIdKey::Array(vec![Value::from("x".repeat(4096))].into());
let rid = Value::RecordId(RecordId {
table: "t".into(),
key,
});
assert!(approx_value_size(&rid) > 4096, "a wide compound key is sized by its contents");
}
fn row_with_floats(n: usize) -> Value {
Value::from(vec![Value::from(0.5f64); n])
}
#[test]
fn an_empty_previous_batch_keeps_the_row_cap() {
assert_eq!(next_batch_len(0, 0, MAX_ROWS, 1 << 20, MAX_ROWS), MAX_ROWS);
}
#[test]
fn narrow_rows_keep_the_row_cap() {
let rows = vec![Value::from("narrow"); 10];
assert_eq!(next_batch_len(bytes(&rows), rows.len(), 600, 8 << 20, MAX_ROWS), MAX_ROWS);
}
#[test]
fn the_probe_follows_the_budget() {
assert_eq!(probe_batch_len(2 << 20, 32), 32);
assert_eq!(probe_batch_len(256 << 10, 32), 4);
assert_eq!(probe_batch_len(8 << 20, 32), 32);
assert_eq!(probe_batch_len(1, 32), 1);
}
#[test]
fn a_batch_at_most_doubles() {
let rows = vec![Value::from("narrow"); 10];
assert_eq!(next_batch_len(bytes(&rows), rows.len(), 32, 8 << 20, MAX_ROWS), 64);
}
#[test]
fn wide_rows_shrink_the_batch_to_the_budget() {
let rows = vec![row_with_floats(1024); 4];
let per_row = approx_value_size(&rows[0]);
assert_eq!(per_row, std::mem::size_of::<Value>() * 1025);
assert_eq!(
next_batch_len(bytes(&rows), rows.len(), MAX_ROWS, 8 << 20, MAX_ROWS),
(8 << 20) / per_row
);
}
#[test]
fn a_row_wider_than_the_budget_still_takes_one_row() {
let rows = vec![row_with_floats(1024)];
assert_eq!(next_batch_len(bytes(&rows), rows.len(), 32, 1, MAX_ROWS), 1);
}
#[test]
fn an_empty_batch_keeps_the_current_length() {
assert_eq!(next_batch_len(0, 0, 31, 2 << 20, MAX_ROWS), 31);
}
#[test]
fn a_geometry_counts_its_containers() {
let empty = || geo::Polygon::new(geo::LineString::<f64>::new(vec![]), vec![]);
let members: Vec<crate::val::Geometry> =
(0..5000).map(|_| crate::val::Geometry::Polygon(empty())).collect();
let collection = Value::Geometry(crate::val::Geometry::Collection(members));
assert!(
approx_value_size(&collection) >= 5000 * std::mem::size_of::<crate::val::Geometry>(),
"a collection of 5000 empty polygons is sized by its members"
);
let rings: Vec<geo::LineString<f64>> =
(0..5000).map(|_| geo::LineString::new(vec![])).collect();
let polygon = Value::Geometry(crate::val::Geometry::Polygon(geo::Polygon::new(
geo::LineString::new(vec![]),
rings,
)));
assert!(
approx_value_size(&polygon) >= 5000 * std::mem::size_of::<geo::LineString<f64>>(),
"a polygon of 5000 empty interior rings is sized by its rings"
);
}
#[test]
fn a_geometry_counts_its_coordinates() {
let ring: Vec<(f64, f64)> = (0..5000).map(|i| (i as f64, 0.0)).collect();
let polygon =
Value::Geometry(crate::val::Geometry::Polygon(geo::Polygon::new(ring.into(), vec![])));
assert!(
approx_value_size(&polygon) >= 5000 * std::mem::size_of::<geo::Coord<f64>>(),
"a 5000-point polygon is sized by its coordinates"
);
}
}