use std::collections::HashMap;
use std::sync::Arc;
use parking_lot::RwLock;
#[derive(Debug, Clone, uniffi::Record)]
pub struct MobileGraphNode {
pub id: u64,
pub label: String,
pub properties_json: Option<String>,
pub vector: Option<Vec<f32>>,
}
#[derive(Debug, Clone, uniffi::Record)]
pub struct MobileGraphEdge {
pub id: u64,
pub source: u64,
pub target: u64,
pub label: String,
pub properties_json: Option<String>,
}
#[derive(Debug, Clone, uniffi::Record)]
pub struct TraversalResult {
pub node_id: u64,
pub depth: u32,
}
#[derive(uniffi::Object)]
pub struct MobileGraphStore {
nodes: RwLock<HashMap<u64, MobileGraphNode>>,
edges: RwLock<HashMap<u64, MobileGraphEdge>>,
outgoing: RwLock<HashMap<u64, Vec<u64>>>,
incoming: RwLock<HashMap<u64, Vec<u64>>>,
}
#[uniffi::export]
impl MobileGraphStore {
#[uniffi::constructor]
pub fn new() -> Arc<Self> {
Arc::new(Self {
nodes: RwLock::new(HashMap::new()),
edges: RwLock::new(HashMap::new()),
outgoing: RwLock::new(HashMap::new()),
incoming: RwLock::new(HashMap::new()),
})
}
pub fn add_node(&self, node: MobileGraphNode) {
let mut nodes = self.nodes.write();
nodes.insert(node.id, node);
}
pub fn add_edge(&self, edge: MobileGraphEdge) -> Result<(), crate::VelesError> {
let mut edges = self.edges.write();
let mut outgoing = self.outgoing.write();
let mut incoming = self.incoming.write();
if edges.contains_key(&edge.id) {
return Err(crate::VelesError::Database {
message: format!("Edge with ID {} already exists", edge.id),
});
}
let source = edge.source;
let target = edge.target;
let id = edge.id;
edges.insert(id, edge);
outgoing.entry(source).or_default().push(id);
incoming.entry(target).or_default().push(id);
Ok(())
}
pub fn get_node(&self, id: u64) -> Option<MobileGraphNode> {
let nodes = self.nodes.read();
nodes.get(&id).cloned()
}
pub fn get_edge(&self, id: u64) -> Option<MobileGraphEdge> {
let edges = self.edges.read();
edges.get(&id).cloned()
}
pub fn node_count(&self) -> u64 {
let nodes = self.nodes.read();
nodes.len() as u64
}
pub fn edge_count(&self) -> u64 {
let edges = self.edges.read();
edges.len() as u64
}
pub fn get_outgoing(&self, node_id: u64) -> Vec<MobileGraphEdge> {
let edges = self.edges.read();
let outgoing = self.outgoing.read();
outgoing
.get(&node_id)
.map(|ids| ids.iter().filter_map(|id| edges.get(id).cloned()).collect())
.unwrap_or_default()
}
pub fn get_incoming(&self, node_id: u64) -> Vec<MobileGraphEdge> {
let edges = self.edges.read();
let incoming = self.incoming.read();
incoming
.get(&node_id)
.map(|ids| ids.iter().filter_map(|id| edges.get(id).cloned()).collect())
.unwrap_or_default()
}
pub fn get_outgoing_by_label(&self, node_id: u64, label: String) -> Vec<MobileGraphEdge> {
self.get_outgoing(node_id)
.into_iter()
.filter(|e| e.label == label)
.collect()
}
pub fn get_neighbors(&self, node_id: u64) -> Vec<u64> {
self.get_outgoing(node_id)
.into_iter()
.map(|e| e.target)
.collect()
}
pub fn bfs_traverse(&self, source_id: u64, max_depth: u32, limit: u32) -> Vec<TraversalResult> {
use std::collections::{HashSet, VecDeque};
let mut results: Vec<TraversalResult> = Vec::new();
let mut visited: HashSet<u64> = HashSet::new();
let mut queue: VecDeque<(u64, u32)> = VecDeque::new();
queue.push_back((source_id, 0));
visited.insert(source_id);
while let Some((node_id, depth)) = queue.pop_front() {
if results.len() >= limit as usize {
break;
}
if depth > 0 {
results.push(TraversalResult { node_id, depth });
}
if depth < max_depth {
for edge in self.get_outgoing(node_id) {
if !visited.contains(&edge.target) {
visited.insert(edge.target);
queue.push_back((edge.target, depth + 1));
}
}
}
}
results
}
pub fn remove_node(&self, node_id: u64) {
let mut edges = self.edges.write();
let mut outgoing = self.outgoing.write();
let mut incoming = self.incoming.write();
let mut nodes = self.nodes.write();
nodes.remove(&node_id);
let outgoing_ids: Vec<u64> = outgoing.remove(&node_id).unwrap_or_default();
for edge_id in outgoing_ids {
if let Some(edge) = edges.remove(&edge_id) {
if let Some(ids) = incoming.get_mut(&edge.target) {
ids.retain(|&id| id != edge_id);
}
}
}
let incoming_ids: Vec<u64> = incoming.remove(&node_id).unwrap_or_default();
for edge_id in incoming_ids {
if let Some(edge) = edges.remove(&edge_id) {
if let Some(ids) = outgoing.get_mut(&edge.source) {
ids.retain(|&id| id != edge_id);
}
}
}
}
pub fn remove_edge(&self, edge_id: u64) {
let mut edges = self.edges.write();
let mut outgoing = self.outgoing.write();
let mut incoming = self.incoming.write();
if let Some(edge) = edges.remove(&edge_id) {
if let Some(ids) = outgoing.get_mut(&edge.source) {
ids.retain(|&id| id != edge_id);
}
if let Some(ids) = incoming.get_mut(&edge.target) {
ids.retain(|&id| id != edge_id);
}
}
}
pub fn clear(&self) {
let mut edges = self.edges.write();
let mut outgoing = self.outgoing.write();
let mut incoming = self.incoming.write();
let mut nodes = self.nodes.write();
edges.clear();
outgoing.clear();
incoming.clear();
nodes.clear();
}
pub fn dfs_traverse(&self, source_id: u64, max_depth: u32, limit: u32) -> Vec<TraversalResult> {
use std::collections::HashSet;
let mut results: Vec<TraversalResult> = Vec::new();
let mut visited: HashSet<u64> = HashSet::new();
let mut stack: Vec<(u64, u32)> = vec![(source_id, 0)];
while let Some((node_id, depth)) = stack.pop() {
if results.len() >= limit as usize {
break;
}
if visited.contains(&node_id) {
continue;
}
visited.insert(node_id);
if depth > 0 {
results.push(TraversalResult { node_id, depth });
}
if depth < max_depth {
let neighbors: Vec<_> = self
.get_outgoing(node_id)
.into_iter()
.filter(|e| !visited.contains(&e.target))
.collect();
for edge in neighbors.into_iter().rev() {
stack.push((edge.target, depth + 1));
}
}
}
results
}
pub fn has_node(&self, id: u64) -> bool {
let nodes = self.nodes.read();
nodes.contains_key(&id)
}
pub fn has_edge(&self, id: u64) -> bool {
let edges = self.edges.read();
edges.contains_key(&id)
}
#[allow(clippy::cast_possible_truncation)]
pub fn out_degree(&self, node_id: u64) -> u32 {
let outgoing = self.outgoing.read();
outgoing.get(&node_id).map_or(0, |v| v.len() as u32)
}
#[allow(clippy::cast_possible_truncation)]
pub fn in_degree(&self, node_id: u64) -> u32 {
let incoming = self.incoming.read();
incoming.get(&node_id).map_or(0, |v| v.len() as u32)
}
pub fn get_nodes_by_label(&self, label: String) -> Vec<MobileGraphNode> {
let nodes = self.nodes.read();
nodes
.values()
.filter(|n| n.label == label)
.cloned()
.collect()
}
pub fn get_edges_by_label(&self, label: String) -> Vec<MobileGraphEdge> {
let edges = self.edges.read();
edges
.values()
.filter(|e| e.label == label)
.cloned()
.collect()
}
}
impl Default for MobileGraphStore {
fn default() -> Self {
Self {
nodes: RwLock::new(HashMap::new()),
edges: RwLock::new(HashMap::new()),
outgoing: RwLock::new(HashMap::new()),
incoming: RwLock::new(HashMap::new()),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_mobile_graph_node_creation() {
let node = MobileGraphNode {
id: 1,
label: "Person".to_string(),
properties_json: Some(r#"{"name": "John"}"#.to_string()),
vector: None,
};
assert_eq!(node.id, 1);
assert_eq!(node.label, "Person");
}
#[test]
fn test_mobile_graph_edge_creation() {
let edge = MobileGraphEdge {
id: 100,
source: 1,
target: 2,
label: "KNOWS".to_string(),
properties_json: None,
};
assert_eq!(edge.id, 100);
assert_eq!(edge.source, 1);
assert_eq!(edge.target, 2);
}
#[test]
fn test_mobile_graph_store_add_nodes() {
let store = MobileGraphStore::new();
store.add_node(MobileGraphNode {
id: 1,
label: "Person".to_string(),
properties_json: None,
vector: None,
});
assert_eq!(store.node_count(), 1);
}
#[test]
fn test_mobile_graph_store_add_edges() {
let store = MobileGraphStore::new();
store.add_node(MobileGraphNode {
id: 1,
label: "Person".to_string(),
properties_json: None,
vector: None,
});
store.add_node(MobileGraphNode {
id: 2,
label: "Person".to_string(),
properties_json: None,
vector: None,
});
let result = store.add_edge(MobileGraphEdge {
id: 100,
source: 1,
target: 2,
label: "KNOWS".to_string(),
properties_json: None,
});
assert!(result.is_ok());
assert_eq!(store.edge_count(), 1);
}
#[test]
fn test_mobile_graph_store_duplicate_edge_error() {
let store = MobileGraphStore::new();
store.add_node(MobileGraphNode {
id: 1,
label: "Person".to_string(),
properties_json: None,
vector: None,
});
store.add_node(MobileGraphNode {
id: 2,
label: "Person".to_string(),
properties_json: None,
vector: None,
});
let _ = store.add_edge(MobileGraphEdge {
id: 100,
source: 1,
target: 2,
label: "KNOWS".to_string(),
properties_json: None,
});
let result = store.add_edge(MobileGraphEdge {
id: 100,
source: 1,
target: 2,
label: "KNOWS".to_string(),
properties_json: None,
});
assert!(result.is_err());
}
#[test]
fn test_mobile_graph_store_get_outgoing() {
let store = MobileGraphStore::new();
store.add_node(MobileGraphNode {
id: 1,
label: "Person".to_string(),
properties_json: None,
vector: None,
});
store.add_node(MobileGraphNode {
id: 2,
label: "Person".to_string(),
properties_json: None,
vector: None,
});
store.add_node(MobileGraphNode {
id: 3,
label: "Person".to_string(),
properties_json: None,
vector: None,
});
let _ = store.add_edge(MobileGraphEdge {
id: 100,
source: 1,
target: 2,
label: "KNOWS".to_string(),
properties_json: None,
});
let _ = store.add_edge(MobileGraphEdge {
id: 101,
source: 1,
target: 3,
label: "KNOWS".to_string(),
properties_json: None,
});
let outgoing = store.get_outgoing(1);
assert_eq!(outgoing.len(), 2);
}
#[test]
fn test_mobile_graph_store_bfs_traverse() {
let store = MobileGraphStore::new();
for i in 1..=4 {
store.add_node(MobileGraphNode {
id: i,
label: "Person".to_string(),
properties_json: None,
vector: None,
});
}
let _ = store.add_edge(MobileGraphEdge {
id: 100,
source: 1,
target: 2,
label: "KNOWS".to_string(),
properties_json: None,
});
let _ = store.add_edge(MobileGraphEdge {
id: 101,
source: 2,
target: 3,
label: "KNOWS".to_string(),
properties_json: None,
});
let _ = store.add_edge(MobileGraphEdge {
id: 102,
source: 3,
target: 4,
label: "KNOWS".to_string(),
properties_json: None,
});
let results = store.bfs_traverse(1, 3, 100);
assert_eq!(results.len(), 3);
assert!(results.iter().any(|r| r.node_id == 2 && r.depth == 1));
assert!(results.iter().any(|r| r.node_id == 3 && r.depth == 2));
assert!(results.iter().any(|r| r.node_id == 4 && r.depth == 3));
}
#[test]
fn test_mobile_graph_store_remove_node() {
let store = MobileGraphStore::new();
store.add_node(MobileGraphNode {
id: 1,
label: "Person".to_string(),
properties_json: None,
vector: None,
});
store.add_node(MobileGraphNode {
id: 2,
label: "Person".to_string(),
properties_json: None,
vector: None,
});
let _ = store.add_edge(MobileGraphEdge {
id: 100,
source: 1,
target: 2,
label: "KNOWS".to_string(),
properties_json: None,
});
assert_eq!(store.node_count(), 2);
assert_eq!(store.edge_count(), 1);
store.remove_node(1);
assert_eq!(store.node_count(), 1);
assert_eq!(store.edge_count(), 0); }
#[test]
fn test_mobile_graph_store_remove_edge() {
let store = MobileGraphStore::new();
store.add_node(MobileGraphNode {
id: 1,
label: "Person".to_string(),
properties_json: None,
vector: None,
});
store.add_node(MobileGraphNode {
id: 2,
label: "Person".to_string(),
properties_json: None,
vector: None,
});
let _ = store.add_edge(MobileGraphEdge {
id: 100,
source: 1,
target: 2,
label: "KNOWS".to_string(),
properties_json: None,
});
assert_eq!(store.edge_count(), 1);
store.remove_edge(100);
assert_eq!(store.edge_count(), 0);
assert!(store.get_outgoing(1).is_empty());
assert!(store.get_incoming(2).is_empty());
}
#[test]
fn test_mobile_graph_store_clear() {
let store = MobileGraphStore::new();
store.add_node(MobileGraphNode {
id: 1,
label: "Person".to_string(),
properties_json: None,
vector: None,
});
store.add_node(MobileGraphNode {
id: 2,
label: "Person".to_string(),
properties_json: None,
vector: None,
});
let _ = store.add_edge(MobileGraphEdge {
id: 100,
source: 1,
target: 2,
label: "KNOWS".to_string(),
properties_json: None,
});
store.clear();
assert_eq!(store.node_count(), 0);
assert_eq!(store.edge_count(), 0);
}
}