use std::collections::BTreeMap;
use graph_storage_sdk::models::NodeFilterField as Field;
use graph_storage_sdk::models::{
AdjacencyEntry, AdjacencySide, EdgeKey, EdgeView, ElementEnvelope, GraphRevision, NodeId,
NodeKey, NodeRow, NodeView, ProjectionRequest, ReadSnapshot, Subject,
};
use graph_storage_sdk::plugin_api::{EmbeddingState, GraphStoreError, StoreCtx};
use sea_orm::sea_query::Expr;
use sea_orm::{ColumnTrait, Condition, EntityTrait, ExprTrait, QuerySelect};
use toolkit_db::odata::{LimitCfg, paginate_odata};
use toolkit_db::secure::{DBRunner, SecureEntityExt};
use toolkit_odata::{Page as OdataPage, SortDir};
use uuid::Uuid;
use crate::infra::projections::{
EdgeHop, NodeIdent, NodeTyped, TypeName, TypeTraits, edge_hop_columns, node_ident_columns,
node_typed_columns, type_name_columns, type_traits_columns,
};
use crate::infra::storage::entity::{edge, graph_meta, gts_type, node};
use crate::infra::storage::odata_mapper::NodeODataMapper;
use crate::infra::store::{PgGraphStore, map_db_error, map_scope_err};
fn map_odata_err(error: toolkit_odata::Error) -> GraphStoreError {
use toolkit_odata::Error as E;
match error {
E::InvalidLimit => GraphStoreError::LimitExceeded {
what: error.to_string(),
},
E::InvalidFilter(_) | E::InvalidOrderByField(_) => GraphStoreError::InvalidQuery {
what: format!("{error}; the projection accepts {}", declared_fields()),
},
other => GraphStoreError::InvalidQuery {
what: other.to_string(),
},
}
}
fn declared_fields() -> String {
use toolkit_odata::filter::FilterField as _;
let columns = Field::FIELDS
.iter()
.map(|f| f.name().to_owned())
.collect::<Vec<_>>()
.join(", ");
format!(
"{columns}, and `payload/<path>` for the paths the selected types declare in their \
`index` trait"
)
}
pub async fn revision(
store: &PgGraphStore,
ctx: &StoreCtx<'_>,
) -> Result<GraphRevision, GraphStoreError> {
let conn = store.db().conn().map_err(|error| map_db_error(&error))?;
read_revision(ctx, &conn).await
}
pub async fn read_revision(
ctx: &StoreCtx<'_>,
runner: &impl DBRunner,
) -> Result<GraphRevision, GraphStoreError> {
let rows = graph_meta::Entity::find()
.secure()
.scope_with(ctx.scope)
.filter(
Condition::all().add(
graph_meta::Column::Key
.is_in([graph_meta::KEY_GRAPH_REVISION, graph_meta::KEY_SOURCE_EPOCH]),
),
)
.all(runner)
.await
.map_err(map_scope_err)?;
let mut revision = GraphRevision {
source_epoch: 1,
revision: 0,
};
for row in rows {
let value = row.value.as_i64().unwrap_or(0);
if row.key == graph_meta::KEY_GRAPH_REVISION {
revision.revision = value;
} else if row.key == graph_meta::KEY_SOURCE_EPOCH {
revision.source_epoch = value;
}
}
Ok(revision)
}
pub async fn begin_read(
store: &PgGraphStore,
ctx: &StoreCtx<'_>,
) -> Result<ReadSnapshot, GraphStoreError> {
Ok(ReadSnapshot {
id: Uuid::now_v7(),
revision: revision(store, ctx).await?,
})
}
pub async fn resolve_node_ids(
store: &PgGraphStore,
ctx: &StoreCtx<'_>,
keys: &[NodeKey],
) -> Result<Vec<(NodeKey, NodeId)>, GraphStoreError> {
if keys.is_empty() {
return Ok(Vec::new());
}
let conn = store.db().conn().map_err(|error| map_db_error(&error))?;
let rows = node::Entity::find()
.secure()
.scope_with(ctx.scope)
.filter(Condition::all().add(node::Column::NodeKey.is_in(keys.to_vec())))
.filter(Condition::all().add(node::Column::DeletedAt.is_null()))
.project_all(&conn, |query| {
node_ident_columns(query).into_model::<NodeIdent>()
})
.await
.map_err(map_scope_err)?;
Ok(rows.into_iter().map(|r| (r.node_key, r.id)).collect())
}
pub async fn embedding_state(
store: &PgGraphStore,
ctx: &StoreCtx<'_>,
keys: &[NodeKey],
) -> Result<Vec<Option<EmbeddingState>>, GraphStoreError> {
if keys.is_empty() {
return Ok(Vec::new());
}
let conn = store.db().conn().map_err(|error| map_db_error(&error))?;
let rows: Vec<StoredVector> = node::Entity::find()
.secure()
.scope_with(ctx.scope)
.filter(Condition::all().add(node::Column::NodeKey.is_in(keys.to_vec())))
.filter(Condition::all().add(node::Column::DeletedAt.is_null()))
.project_all(&conn, |query| {
stored_vector_columns(query).into_model::<StoredVector>()
})
.await
.map_err(map_scope_err)?;
let by_key: BTreeMap<String, EmbeddingState> = rows
.into_iter()
.map(|r| {
(
r.node_key,
EmbeddingState {
input_hash: r.embedding_input_hash,
vector_epoch: if r.has_vector {
r.embedding_epoch
} else {
None
},
},
)
})
.collect();
Ok(keys.iter().map(|key| by_key.get(key).cloned()).collect())
}
fn stored_vector_columns(query: sea_orm::Select<node::Entity>) -> sea_orm::Select<node::Entity> {
query
.select_only()
.column(node::Column::NodeKey)
.column(node::Column::EmbeddingInputHash)
.column(node::Column::EmbeddingEpoch)
.column_as(
Expr::col(node::Column::Embedding).is_not_null(),
"has_vector",
)
}
#[derive(Debug, sea_orm::FromQueryResult)]
struct StoredVector {
node_key: String,
embedding_input_hash: Option<String>,
embedding_epoch: Option<i64>,
has_vector: bool,
}
async fn type_names(
ctx: &StoreCtx<'_>,
runner: &impl DBRunner,
ids: &[i32],
) -> Result<BTreeMap<i32, String>, GraphStoreError> {
if ids.is_empty() {
return Ok(BTreeMap::new());
}
let rows = gts_type::Entity::find()
.secure()
.scope_with(ctx.scope)
.filter(Condition::all().add(gts_type::Column::Id.is_in(ids.to_vec())))
.project_all(runner, |query| {
type_name_columns(query).into_model::<TypeName>()
})
.await
.map_err(map_scope_err)?;
Ok(rows.into_iter().map(|r| (r.id, r.gts_type_id)).collect())
}
fn envelope_of_node(model: &node::Model, revision: GraphRevision) -> ElementEnvelope {
ElementEnvelope {
tenant_id: model.tenant_id,
key: model.node_key.clone(),
created_at: model.created_at,
created_by: Subject {
subject_id: model.created_by_subject_id,
subject_type: model.created_by_subject_type.clone(),
},
updated_at: model.updated_at,
updated_by: Subject {
subject_id: model.updated_by_subject_id,
subject_type: model.updated_by_subject_type.clone(),
},
deleted_at: model.deleted_at,
deleted_by: model.deleted_by_subject_id.map(|subject_id| Subject {
subject_id,
subject_type: model.deleted_by_subject_type.clone(),
}),
graph_revision: revision,
}
}
async fn observed_revision(
ctx: &StoreCtx<'_>,
runner: &impl DBRunner,
) -> Result<GraphRevision, GraphStoreError> {
match ctx.snapshot {
Some(snapshot) => Ok(snapshot.revision),
None => read_revision(ctx, runner).await,
}
}
fn to_view(
model: node::Model,
type_id: String,
adjacency: Vec<AdjacencyEntry>,
truncated: bool,
revision: GraphRevision,
) -> NodeView {
let envelope = envelope_of_node(&model, revision);
NodeView {
node_key: model.node_key,
type_id,
name: (!model.name.is_empty()).then_some(model.name),
payload: Some(model.payload),
has_embedding: model.embedding.is_some(),
labels: Vec::new(),
adjacency,
adjacency_truncated: truncated,
envelope,
}
}
pub async fn get_node(
store: &PgGraphStore,
ctx: &StoreCtx<'_>,
key: &NodeKey,
adjacency_limit: u32,
) -> Result<NodeView, GraphStoreError> {
let conn = store.db().conn().map_err(|error| map_db_error(&error))?;
let model = node::Entity::find()
.secure()
.scope_with(ctx.scope)
.filter(Condition::all().add(node::Column::NodeKey.eq(key.clone())))
.filter(Condition::all().add(node::Column::DeletedAt.is_null()))
.one(&conn)
.await
.map_err(map_scope_err)?
.ok_or(GraphStoreError::NotFound)?;
let probe = u64::from(adjacency_limit) + 1;
let outgoing: Vec<EdgeHop> = edge::Entity::find()
.secure()
.scope_with(ctx.scope)
.filter(Condition::all().add(edge::Column::SrcNodeId.eq(model.id)))
.filter(Condition::all().add(edge::Column::DeletedAt.is_null()))
.order_by(edge::Column::Id, sea_orm::Order::Asc)
.limit(probe)
.project_all(&conn, |query| {
edge_hop_columns(query).into_model::<EdgeHop>()
})
.await
.map_err(map_scope_err)?;
let incoming: Vec<EdgeHop> = edge::Entity::find()
.secure()
.scope_with(ctx.scope)
.filter(Condition::all().add(edge::Column::DstNodeId.eq(model.id)))
.filter(Condition::all().add(edge::Column::DeletedAt.is_null()))
.order_by(edge::Column::Id, sea_orm::Order::Asc)
.limit(probe)
.project_all(&conn, |query| {
edge_hop_columns(query).into_model::<EdgeHop>()
})
.await
.map_err(map_scope_err)?;
let truncated = outgoing.len() as u64 > u64::from(adjacency_limit)
|| incoming.len() as u64 > u64::from(adjacency_limit);
let mut neighbour_ids: Vec<i64> = outgoing
.iter()
.map(|e| e.dst_node_id)
.chain(incoming.iter().map(|e| e.src_node_id))
.collect();
neighbour_ids.sort_unstable();
neighbour_ids.dedup();
let neighbours = node::Entity::find()
.secure()
.scope_with(ctx.scope)
.filter(Condition::all().add(node::Column::Id.is_in(neighbour_ids)))
.filter(Condition::all().add(node::Column::DeletedAt.is_null()))
.project_all(&conn, |query| {
node_typed_columns(query).into_model::<NodeTyped>()
})
.await
.map_err(map_scope_err)?;
let by_id: BTreeMap<i64, NodeTyped> = neighbours.into_iter().map(|n| (n.id, n)).collect();
let mut type_ids: Vec<i32> = outgoing
.iter()
.chain(incoming.iter())
.map(|e| e.gts_edge_type_id)
.chain(by_id.values().map(|n| n.gts_node_type_id))
.chain(std::iter::once(model.gts_node_type_id))
.collect();
type_ids.sort_unstable();
type_ids.dedup();
let names = type_names(ctx, &conn, &type_ids).await?;
let unknown = || String::new();
let mut adjacency = Vec::new();
for (edges, side) in [
(&outgoing, AdjacencySide::Outgoing),
(&incoming, AdjacencySide::Incoming),
] {
for e in edges.iter().take(adjacency_limit as usize) {
let neighbour_id = match side {
AdjacencySide::Outgoing => e.dst_node_id,
AdjacencySide::Incoming => e.src_node_id,
};
let Some(neighbour) = by_id.get(&neighbour_id) else {
continue;
};
adjacency.push(AdjacencyEntry {
edge_key: e.edge_key.clone(),
edge_type_id: names
.get(&e.gts_edge_type_id)
.cloned()
.unwrap_or_else(unknown),
side,
neighbor_key: neighbour.node_key.clone(),
neighbor_type_id: names
.get(&neighbour.gts_node_type_id)
.cloned()
.unwrap_or_else(unknown),
});
}
}
let type_id = names
.get(&model.gts_node_type_id)
.cloned()
.unwrap_or_else(unknown);
let revision = observed_revision(ctx, &conn).await?;
Ok(to_view(model, type_id, adjacency, truncated, revision))
}
fn envelope_of_edge(model: &edge::Model, revision: GraphRevision) -> ElementEnvelope {
ElementEnvelope {
tenant_id: model.tenant_id,
key: model.edge_key.clone(),
created_at: model.created_at,
created_by: Subject {
subject_id: model.created_by_subject_id,
subject_type: model.created_by_subject_type.clone(),
},
updated_at: model.updated_at,
updated_by: Subject {
subject_id: model.updated_by_subject_id,
subject_type: model.updated_by_subject_type.clone(),
},
deleted_at: model.deleted_at,
deleted_by: model.deleted_by_subject_id.map(|subject_id| Subject {
subject_id,
subject_type: model.deleted_by_subject_type.clone(),
}),
graph_revision: revision,
}
}
pub async fn get_edge(
store: &PgGraphStore,
ctx: &StoreCtx<'_>,
key: &EdgeKey,
) -> Result<EdgeView, GraphStoreError> {
let conn = store.db().conn().map_err(|error| map_db_error(&error))?;
let model = edge::Entity::find()
.secure()
.scope_with(ctx.scope)
.filter(Condition::all().add(edge::Column::EdgeKey.eq(key.clone())))
.filter(Condition::all().add(edge::Column::DeletedAt.is_null()))
.one(&conn)
.await
.map_err(map_scope_err)?
.ok_or(GraphStoreError::NotFound)?;
let endpoints = node::Entity::find()
.secure()
.scope_with(ctx.scope)
.filter(
Condition::all().add(node::Column::Id.is_in([model.src_node_id, model.dst_node_id])),
)
.filter(Condition::all().add(node::Column::DeletedAt.is_null()))
.project_all(&conn, |query| {
node_ident_columns(query).into_model::<NodeIdent>()
})
.await
.map_err(map_scope_err)?;
let by_id: BTreeMap<i64, NodeIdent> = endpoints.into_iter().map(|n| (n.id, n)).collect();
let (Some(src), Some(dst)) = (by_id.get(&model.src_node_id), by_id.get(&model.dst_node_id))
else {
return Err(GraphStoreError::NotFound);
};
let names = type_names(ctx, &conn, &[model.gts_edge_type_id]).await?;
let revision = observed_revision(ctx, &conn).await?;
let envelope = envelope_of_edge(&model, revision);
Ok(EdgeView {
edge_key: model.edge_key.clone(),
edge_type_id: names
.get(&model.gts_edge_type_id)
.cloned()
.unwrap_or_default(),
src: src.node_key.clone(),
dst: dst.node_key.clone(),
discriminator: model.discriminator.clone(),
payload: Some(model.payload.clone()),
envelope,
})
}
pub async fn hydrate_nodes(
store: &PgGraphStore,
ctx: &StoreCtx<'_>,
ids: &[NodeId],
) -> Result<Vec<NodeView>, GraphStoreError> {
if ids.is_empty() {
return Ok(Vec::new());
}
let conn = store.db().conn().map_err(|error| map_db_error(&error))?;
let models = node::Entity::find()
.secure()
.scope_with(ctx.scope)
.filter(Condition::all().add(node::Column::Id.is_in(ids.to_vec())))
.filter(Condition::all().add(node::Column::DeletedAt.is_null()))
.all(&conn)
.await
.map_err(map_scope_err)?;
let mut type_ids: Vec<i32> = models.iter().map(|m| m.gts_node_type_id).collect();
type_ids.sort_unstable();
type_ids.dedup();
let names = type_names(ctx, &conn, &type_ids).await?;
let revision = observed_revision(ctx, &conn).await?;
let mut by_id: BTreeMap<i64, node::Model> = models.into_iter().map(|m| (m.id, m)).collect();
let mut views = Vec::new();
for id in ids {
if let Some(model) = by_id.remove(id) {
let type_id = names
.get(&model.gts_node_type_id)
.cloned()
.unwrap_or_default();
views.push(to_view(model, type_id, Vec::new(), false, revision));
}
}
Ok(views)
}
pub async fn node_types(
store: &PgGraphStore,
ctx: &StoreCtx<'_>,
ids: &[NodeId],
) -> Result<Vec<(NodeId, String)>, GraphStoreError> {
if ids.is_empty() {
return Ok(Vec::new());
}
let conn = store.db().conn().map_err(|error| map_db_error(&error))?;
let rows = node::Entity::find()
.secure()
.scope_with(ctx.scope)
.filter(Condition::all().add(node::Column::Id.is_in(ids.to_vec())))
.filter(Condition::all().add(node::Column::DeletedAt.is_null()))
.project_all(&conn, |query| {
node_typed_columns(query).into_model::<NodeTyped>()
})
.await
.map_err(map_scope_err)?;
let mut type_ids: Vec<i32> = rows.iter().map(|row| row.gts_node_type_id).collect();
type_ids.sort_unstable();
type_ids.dedup();
let names = type_names(ctx, &conn, &type_ids).await?;
Ok(rows
.into_iter()
.filter_map(|row| {
names
.get(&row.gts_node_type_id)
.map(|name| (row.id, name.clone()))
})
.collect())
}
pub async fn project_table(
store: &PgGraphStore,
ctx: &StoreCtx<'_>,
req: ProjectionRequest,
) -> Result<OdataPage<NodeRow>, GraphStoreError> {
let conn = store.db().conn().map_err(|error| map_db_error(&error))?;
let mut select = node::Entity::find()
.secure()
.scope_with(ctx.scope)
.filter(Condition::all().add(node::Column::DeletedAt.is_null()));
let mut declared_kinds: Vec<BTreeMap<String, crate::domain::ontology::ScalarKind>> = Vec::new();
if let Some(set) = &req.type_set {
let names: Vec<String> = set.0.iter().cloned().collect();
let types = gts_type::Entity::find()
.secure()
.scope_with(ctx.scope)
.filter(Condition::all().add(gts_type::Column::GtsTypeId.is_in(names)))
.project_all(&conn, |query| {
type_traits_columns(query).into_model::<TypeTraits>()
})
.await
.map_err(map_scope_err)?;
let ids: Vec<i32> = types.iter().map(|t| t.id).collect();
declared_kinds = types
.iter()
.map(|t| super::types::index_kinds_from_json(&t.effective_traits))
.collect();
select = select.filter(Condition::all().add(node::Column::GtsNodeTypeId.is_in(ids)));
}
if crate::domain::projection::mentions_payload(&req.query) {
let admitted = req
.type_set
.as_ref()
.map(|_| crate::domain::projection::admitted_paths(&declared_kinds));
return project_over_payload(store, ctx, &conn, select, &req, admitted.as_ref()).await;
}
let page = paginate_odata::<Field, NodeODataMapper, _, node::Model, _, _>(
select,
&conn,
&req.query,
("node_key", SortDir::Asc),
LimitCfg {
default: u64::from(store.config().projection_max_page),
max: u64::from(store.config().projection_max_page),
},
|model| model,
)
.await
.map_err(map_odata_err)?;
rows_to_page(ctx, &conn, page.items, page.page_info).await
}
async fn project_over_payload(
store: &PgGraphStore,
ctx: &StoreCtx<'_>,
conn: &impl DBRunner,
select: toolkit_db::secure::SecureSelect<node::Entity, toolkit_db::secure::Scoped>,
req: &ProjectionRequest,
admitted: Option<&BTreeMap<String, crate::domain::ontology::ScalarKind>>,
) -> Result<OdataPage<NodeRow>, GraphStoreError> {
use crate::domain::projection as plan;
use crate::infra::store::projection as sql;
let invalid = |what: String| GraphStoreError::InvalidQuery { what };
let plan = plan::plan(&req.query, admitted).map_err(|error| invalid(error.0))?;
let max = u64::from(store.config().projection_max_page);
let limit = req.query.limit.unwrap_or(max).clamp(1, max);
if let Some(cursor) = &req.query.cursor {
if cursor.d == "bwd" {
return Err(invalid(
"backward paging is not available over a payload ordering".to_owned(),
));
}
if let (Some(hash), Some(recorded)) =
(req.query.filter_hash.as_deref(), cursor.f.as_deref())
&& hash != recorded
{
return Err(invalid(
"the cursor was minted under a different $filter".to_owned(),
));
}
if cursor.s != sql::signed_tokens(&plan) {
return Err(invalid(
"the cursor was minted under a different $orderby".to_owned(),
));
}
}
let mut statement = select;
if let Some(predicate) = &plan.filter {
statement = statement.filter(sql::condition(predicate));
}
if let Some(cursor) = &req.query.cursor {
statement = statement.filter(sql::keyset(&plan, cursor).map_err(invalid)?);
}
for (expr, order) in sql::order_terms(&plan) {
statement = statement.order_by(expr, order);
}
statement = statement.limit(limit + 1);
let mut rows = statement.all(conn).await.map_err(map_scope_err)?;
let has_more = rows.len() as u64 > limit;
if has_more {
rows.truncate(usize::try_from(limit).unwrap_or(usize::MAX));
}
let next_cursor = match (has_more, rows.last()) {
(true, Some(last)) => {
let keys = sql::cursor_keys(last, &plan)
.map_err(|error| GraphStoreError::Corrupt { reason: error })?;
let cursor = toolkit_odata::CursorV1 {
k: keys,
o: plan.order.first().map_or(SortDir::Asc, |t| t.dir),
s: sql::signed_tokens(&plan),
f: req.query.filter_hash.clone(),
d: "fwd".to_owned(),
};
Some(cursor.encode().map_err(|error| GraphStoreError::Corrupt {
reason: format!("cannot encode the continuation token: {error}"),
})?)
}
_ => None,
};
rows_to_page(
ctx,
conn,
rows,
toolkit_odata::page::PageInfo {
next_cursor,
prev_cursor: None,
limit,
},
)
.await
}
async fn rows_to_page(
ctx: &StoreCtx<'_>,
conn: &impl DBRunner,
rows: Vec<node::Model>,
page_info: toolkit_odata::page::PageInfo,
) -> Result<OdataPage<NodeRow>, GraphStoreError> {
let revision = observed_revision(ctx, conn).await?;
let mut type_ids: Vec<i32> = rows.iter().map(|m| m.gts_node_type_id).collect();
type_ids.sort_unstable();
type_ids.dedup();
let names = type_names(ctx, conn, &type_ids).await?;
Ok(OdataPage {
items: rows
.into_iter()
.map(|m| NodeRow {
type_id: names.get(&m.gts_node_type_id).cloned().unwrap_or_default(),
envelope: envelope_of_node(&m, revision),
node_key: m.node_key,
name: (!m.name.is_empty()).then_some(m.name),
payload: Some(m.payload),
})
.collect(),
page_info,
})
}
#[cfg(test)]
mod tests {
use super::{node, stored_vector_columns};
#[test]
fn the_embedding_state_statement_reads_four_values_and_no_more() {
use sea_orm::{EntityTrait, QueryTrait};
let sql = stored_vector_columns(node::Entity::find())
.build(sea_orm::DatabaseBackend::Postgres)
.to_string();
assert!(
!sql.contains("payload") && !sql.contains("search_text"),
"embedding state must not read the wide columns: {sql}"
);
assert!(
sql.contains("IS NOT NULL"),
"presence is asked of the server, not decided here: {sql}"
);
for needed in ["node_key", "embedding_input_hash", "embedding_epoch"] {
assert!(
sql.contains(needed),
"embedding state needs `{needed}`: {sql}"
);
}
}
}