use std::collections::HashMap;
use yo_common::{Code, Error, Result, parse_i64};
use yo_doc::{Builder, Doc};
use yo_graph::{Dir, Graph};
use yo_kv::{Foreign, Keyspace, value::Kind};
use super::args::{self, Args};
use super::table::Spec;
use crate::reply::Out;
const NOT_A_GRAPH: &str = "WRONGTYPE Operation against a key holding the wrong kind of value";
const BAD_DEPTH: &str = "DEPTH must be a positive integer";
const BAD_COUNT: &str = "COUNT must be a positive integer";
const COUNT: usize = 10;
const DEPTH: usize = 2;
const MAXDEPTH: usize = 6;
#[derive(Debug, Default)]
pub(super) struct GraphBody {
g: Graph,
ids: HashMap<Box<[u8]>, u64>,
of: Vec<Box<[u8]>>,
labels: HashMap<Box<[u8]>, u32>,
names: Vec<Box<[u8]>>,
build: Builder,
frontier: Vec<u64>,
next: Vec<u64>,
}
impl Foreign for GraphBody {
fn type_name(&self) -> &'static str {
"graph"
}
fn encoding(&self) -> &'static str {
"adjacency"
}
fn memory_bytes(&self) -> usize {
let ids: usize = self
.of
.iter()
.map(|id| id.len() + std::mem::size_of::<Box<[u8]>>())
.sum();
let names: usize = self
.names
.iter()
.map(|n| n.len() + std::mem::size_of::<Box<[u8]>>())
.sum();
self.g.memory_bytes() + ids * 2 + names * 2
}
fn is_empty(&self) -> bool {
self.g.nodes() == 0
}
}
impl GraphBody {
fn dense(&self, id: &[u8]) -> Option<u64> {
self.ids.get(id).copied()
}
fn dense_or_add(&mut self, id: &[u8]) -> Result<u64> {
if let Some(at) = self.ids.get(id) {
return Ok(*at);
}
let at = self.of.len() as u64;
self.g.add_node(at)?;
self.of.push(id.into());
self.ids.insert(id.into(), at);
Ok(at)
}
fn label(&self, name: &[u8]) -> Option<u32> {
self.labels.get(name).copied()
}
fn label_or_add(&mut self, name: &[u8]) -> u32 {
if let Some(n) = self.labels.get(name) {
return *n;
}
let n = self.names.len() as u32;
self.names.push(name.into());
self.labels.insert(name.into(), n);
n
}
fn client(&self, at: u64) -> &[u8] {
self.of
.get(at as usize)
.expect("the plane only holds nodes these tables made")
}
fn document(&mut self, args: Args<'_>, from: usize) -> Result<&[u8]> {
self.build.clear();
self.build.begin_object()?;
let mut i = from;
while i + 1 < args.len() {
self.build.key(args.get(i))?;
self.build.text_bytes(args.get(i + 1))?;
i += 2;
}
self.build.end_object()?;
self.build.finish()
}
fn slot(&self, src: u64, dst: u64, label: u32) -> Option<u32> {
self.g
.hop(src, label, Dir::Out)
.find(|(to, _)| *to == dst)
.map(|(_, slot)| slot)
}
}
pub(super) fn execute(db: &mut Keyspace, spec: &Spec, args: Args<'_>, out: &mut Out) -> Result<()> {
match spec.name {
"g.nadd" => nadd(db, args, out),
"g.nget" => nget(db, args, out),
"g.ndel" => ndel(db, args, out),
"g.eadd" => eadd(db, args, out),
"g.edel" => edel(db, args, out),
"g.out" => step(db, args, Dir::Out, out),
"g.in" => step(db, args, Dir::In, out),
"g.deg" => deg(db, args, out),
"g.neigh" => neigh(db, args, out),
"g.path" => path(db, args, out),
other => unreachable!("{other} is not a graph command"),
}
}
fn nadd(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
pairs(args, 3)?;
let key = args.get(1);
let body = open(db, key)?;
let fresh = body.dense(args.get(2)).is_none();
let at = body.dense_or_add(args.get(2))?;
let doc = body.document(args, 3)?;
let doc = doc.to_vec();
body.g.put_node(at, &doc)?;
out.int(i64::from(fresh));
Ok(())
}
fn nget(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let Some(body) = read(db, args.get(1))? else {
out.nil();
return Ok(());
};
let Some(at) = body.dense(args.get(2)) else {
out.nil();
return Ok(());
};
match body.g.node(at) {
Some(doc) => fields(&doc, out),
None => out.nil(),
}
Ok(())
}
fn ndel(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let Some(body) = write(db, args.get(1))? else {
out.int(0);
return Ok(());
};
let Some(at) = body.dense(args.get(2)) else {
out.int(0);
return Ok(());
};
let gone = body.g.remove_node(at)?;
if gone {
body.ids.remove(args.get(2));
body.of[at as usize] = Box::default();
}
out.int(i64::from(gone));
db.reap_foreign(args.get(1));
Ok(())
}
fn eadd(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
pairs(args, 5)?;
let body = open(db, args.get(1))?;
let src = body.dense_or_add(args.get(2))?;
let dst = body.dense_or_add(args.get(3))?;
let label = body.label_or_add(args.get(4));
let doc = body.document(args, 5)?.to_vec();
match body.slot(src, dst, label) {
Some(slot) => {
body.g.put_edge(slot, &doc)?;
out.int(0);
}
None => {
body.g.link(src, dst, label, &doc)?;
out.int(1);
}
}
Ok(())
}
fn edel(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let Some(body) = write(db, args.get(1))? else {
out.int(0);
return Ok(());
};
let (Some(src), Some(dst), Some(label)) = (
body.dense(args.get(2)),
body.dense(args.get(3)),
body.label(args.get(4)),
) else {
out.int(0);
return Ok(());
};
out.int(i64::from(body.g.unlink(src, dst, label).is_some()));
Ok(())
}
fn step(db: &mut Keyspace, args: Args<'_>, dir: Dir, out: &mut Out) -> Result<()> {
let (count, cursor) = page(args, 4)?;
let Some(body) = read(db, args.get(1))? else {
return empty_page(out);
};
let (Some(at), Some(label)) = (body.dense(args.get(2)), body.label(args.get(3))) else {
return empty_page(out);
};
let run = body.g.neighbours(at, label, dir);
let from = cursor.min(run.len());
let to = from.saturating_add(count).min(run.len());
let next = if to < run.len() { to } else { 0 };
out.array(2);
out.bulk_u64(next as u64);
out.array(to - from);
for hop in &run[from..to] {
out.bulk(body.client(*hop));
}
Ok(())
}
fn deg(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let dir = match args.opt(4) {
None => Dir::Out,
Some(w) if args::is(w, b"out") => Dir::Out,
Some(w) if args::is(w, b"in") => Dir::In,
Some(w) if args::is(w, b"both") => Dir::Out,
Some(_) => return Err(args::syntax()),
};
let both = args.opt(4).is_some_and(|w| args::is(w, b"both"));
let Some(body) = read(db, args.get(1))? else {
out.int(0);
return Ok(());
};
let (Some(at), Some(label)) = (body.dense(args.get(2)), body.label(args.get(3))) else {
out.int(0);
return Ok(());
};
let mut n = body.g.degree(at, label, dir);
if both {
n += body.g.degree(at, label, Dir::In);
}
out.int(i64::try_from(n).unwrap_or(i64::MAX));
Ok(())
}
fn neigh(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let mut depth = DEPTH;
let mut count = usize::MAX;
let mut i = 4;
while i < args.len() {
let rest = args.len() - i;
if args::is(args.get(i), b"depth") && rest >= 2 {
depth = positive(args.get(i + 1), BAD_DEPTH)?;
} else if args::is(args.get(i), b"count") && rest >= 2 {
count = positive(args.get(i + 1), BAD_COUNT)?;
} else {
return Err(args::syntax());
}
i += 2;
}
let Some(body) = write(db, args.get(1))? else {
out.array(0);
return Ok(());
};
let (Some(at), Some(label)) = (body.dense(args.get(2)), body.label(args.get(3))) else {
out.array(0);
return Ok(());
};
let mut seen = vec![at];
let mut frontier = std::mem::take(&mut body.frontier);
let mut next = std::mem::take(&mut body.next);
frontier.clear();
frontier.push(at);
let start = out.len();
let mut written = 0;
for _ in 0..depth {
next.clear();
for node in &frontier {
body.g.prefetch(*node, label, Dir::Out);
}
for node in &frontier {
next.extend_from_slice(body.g.neighbours(*node, label, Dir::Out));
}
next.sort_unstable();
next.dedup();
frontier.clear();
for node in &next {
if seen.binary_search(node).is_ok() {
continue;
}
frontier.push(*node);
}
if frontier.is_empty() {
break;
}
for node in &frontier {
if written == count {
break;
}
out.bulk(body.client(*node));
written += 1;
}
seen.extend_from_slice(&frontier);
seen.sort_unstable();
if written == count {
break;
}
}
body.frontier = frontier;
body.next = next;
out.close_array(start, written);
Ok(())
}
fn path(db: &mut Keyspace, args: Args<'_>, out: &mut Out) -> Result<()> {
let mut max = MAXDEPTH;
if let Some(word) = args.opt(4) {
if args.len() != 6 || !args::is(word, b"maxdepth") {
return Err(args::syntax());
}
max = positive(args.get(5), BAD_DEPTH)?;
}
let Some(body) = read(db, args.get(1))? else {
out.array(0);
return Ok(());
};
let (Some(src), Some(dst)) = (body.dense(args.get(2)), body.dense(args.get(3))) else {
out.array(0);
return Ok(());
};
match search(body, src, dst, max) {
Some(nodes) => {
out.array(nodes.len());
for node in nodes {
out.bulk(body.client(node));
}
}
None => out.array(0),
}
Ok(())
}
fn search(body: &GraphBody, src: u64, dst: u64, max: usize) -> Option<Vec<u64>> {
if src == dst {
return Some(vec![src]);
}
let mut from_src: HashMap<u64, u64> = HashMap::from([(src, src)]);
let mut from_dst: HashMap<u64, u64> = HashMap::from([(dst, dst)]);
let mut a = vec![src];
let mut b = vec![dst];
for _ in 0..max {
let (near, far, seen, other, dir) = if a.len() <= b.len() {
(&mut a, &mut b, &mut from_src, &from_dst, Dir::Out)
} else {
(&mut b, &mut a, &mut from_dst, &from_src, Dir::In)
};
let mut grown = Vec::new();
for node in near.iter() {
for label in body.g.labels() {
for to in body.g.neighbours(*node, *label, dir) {
if seen.contains_key(to) {
continue;
}
seen.insert(*to, *node);
if other.contains_key(to) {
return Some(join(&from_src, &from_dst, *to, src, dst));
}
grown.push(*to);
}
}
}
if grown.is_empty() {
return None;
}
*near = grown;
let _ = far;
}
None
}
fn join(
from_src: &HashMap<u64, u64>,
from_dst: &HashMap<u64, u64>,
meet: u64,
src: u64,
dst: u64,
) -> Vec<u64> {
let mut head = vec![meet];
let mut at = meet;
while at != src {
at = from_src[&at];
head.push(at);
}
head.reverse();
let mut at = meet;
while at != dst {
at = from_dst[&at];
head.push(at);
}
head
}
fn open<'d>(db: &'d mut Keyspace, key: &[u8]) -> Result<&'d mut GraphBody> {
if db.kind_of(key).is_none() {
db.put_foreign(key, Box::new(GraphBody::default()));
}
match write(db, key)? {
Some(body) => Ok(body),
None => unreachable!("the graph was just created"),
}
}
fn write<'d>(db: &'d mut Keyspace, key: &[u8]) -> Result<Option<&'d mut GraphBody>> {
match db.foreign_mut(key)? {
Some(body) => match body.downcast_mut::<GraphBody>() {
Some(body) => Ok(Some(body)),
None => Err(Error::new(Code::WrongType, NOT_A_GRAPH)),
},
None => Ok(None),
}
}
fn read<'d>(db: &'d mut Keyspace, key: &[u8]) -> Result<Option<&'d GraphBody>> {
match db.foreign(key)? {
Some(body) => match body.downcast_ref::<GraphBody>() {
Some(body) => Ok(Some(body)),
None => Err(Error::new(Code::WrongType, NOT_A_GRAPH)),
},
None => Ok(None),
}
}
fn fields(doc: &Doc<'_>, out: &mut Out) {
out.map(doc.len());
for (name, value) in doc.members() {
out.bulk(name);
match value.text_bytes() {
Some(text) => out.bulk(text),
None => match value.as_int() {
Some(n) => out.bulk_int(n),
None => out.bulk(b""),
},
}
}
}
fn page(args: Args<'_>, from: usize) -> Result<(usize, usize)> {
let mut count = COUNT;
let mut cursor = 0;
let mut i = from;
while i < args.len() {
let rest = args.len() - i;
if args::is(args.get(i), b"count") && rest >= 2 {
count = positive(args.get(i + 1), BAD_COUNT)?;
} else if args::is(args.get(i), b"cursor") && rest >= 2 {
cursor = match parse_i64(args.get(i + 1)) {
Some(n) if n >= 0 => usize::try_from(n).unwrap_or(usize::MAX),
_ => return Err(args::syntax()),
};
} else {
return Err(args::syntax());
}
i += 2;
}
Ok((count, cursor))
}
fn empty_page(out: &mut Out) -> Result<()> {
out.array(2);
out.bulk_u64(0);
out.array(0);
Ok(())
}
fn positive(arg: &[u8], msg: &'static str) -> Result<usize> {
match parse_i64(arg) {
Some(n) if n > 0 => Ok(usize::try_from(n).unwrap_or(usize::MAX)),
_ => Err(Error::new(Code::Invalid, msg)),
}
}
fn pairs(args: Args<'_>, from: usize) -> Result<()> {
if (args.len() - from).is_multiple_of(2) {
Ok(())
} else {
Err(args::syntax())
}
}
pub(super) fn is_graph(db: &mut Keyspace, key: &[u8]) -> bool {
db.kind_of(key) == Some(Kind::Foreign)
}