arrow_graph/graph/
arrow_graph.rs

1use arrow::record_batch::RecordBatch;
2use arrow::datatypes::{DataType, Field, Schema};
3use std::sync::Arc;
4use std::path::Path;
5use crate::error::{GraphError, Result};
6use crate::graph::GraphIndexes;
7
8#[derive(Debug, Clone)]
9pub struct ArrowGraph {
10    pub nodes: RecordBatch,
11    pub edges: RecordBatch,
12    pub indexes: GraphIndexes,
13}
14
15impl ArrowGraph {
16    /// Create a new graph from nodes and edges RecordBatches
17    pub fn new(nodes: RecordBatch, edges: RecordBatch) -> Result<Self> {
18        let indexes = GraphIndexes::build(&nodes, &edges)?;
19        
20        Ok(ArrowGraph {
21            nodes,
22            edges,
23            indexes,
24        })
25    }
26    
27    /// Create a graph from just edges (nodes will be inferred)
28    pub fn from_edges(edges: RecordBatch) -> Result<Self> {
29        // Create an empty nodes RecordBatch with proper schema
30        let nodes_schema = Arc::new(Schema::new(vec![
31            Field::new("id", DataType::Utf8, false),
32        ]));
33        
34        let empty_nodes = RecordBatch::new_empty(nodes_schema);
35        Self::new(empty_nodes, edges)
36    }
37
38    /// Create an empty graph with no nodes or edges
39    pub fn empty() -> Result<Self> {
40        // Create empty nodes RecordBatch
41        let nodes_schema = Arc::new(Schema::new(vec![
42            Field::new("id", DataType::Utf8, false),
43        ]));
44        let empty_nodes = RecordBatch::new_empty(nodes_schema);
45
46        // Create empty edges RecordBatch
47        let edges_schema = Arc::new(Schema::new(vec![
48            Field::new("source", DataType::Utf8, false),
49            Field::new("target", DataType::Utf8, false),
50            Field::new("weight", DataType::Float64, true),
51        ]));
52        let empty_edges = RecordBatch::new_empty(edges_schema);
53
54        Self::new(empty_nodes, empty_edges)
55    }
56    
57    /// Load graph from Arrow/Parquet files
58    pub async fn from_files<P: AsRef<Path>>(
59        _nodes_path: P,
60        _edges_path: P,
61    ) -> Result<Self> {
62        todo!("Implement loading from files - will read Arrow/Parquet files")
63    }
64    
65    /// Create graph from RecordBatches (alias for new)
66    pub fn from_tables(
67        nodes: RecordBatch,
68        edges: RecordBatch,
69    ) -> Result<Self> {
70        Self::new(nodes, edges)
71    }
72    
73    /// Execute SQL query with graph functions
74    pub async fn sql(&self, _query: &str) -> Result<RecordBatch> {
75        todo!("Implement SQL execution using DataFusion with graph functions")
76    }
77    
78    /// Get number of nodes in the graph
79    pub fn node_count(&self) -> usize {
80        self.indexes.node_count
81    }
82    
83    /// Get number of edges in the graph  
84    pub fn edge_count(&self) -> usize {
85        self.indexes.edge_count
86    }
87    
88    /// Calculate graph density (edges / max_possible_edges)
89    pub fn density(&self) -> f64 {
90        let n = self.node_count() as f64;
91        let m = self.edge_count() as f64;
92        
93        if n <= 1.0 {
94            0.0
95        } else {
96            m / (n * (n - 1.0))
97        }
98    }
99    
100    /// Get neighbors of a node
101    pub fn neighbors(&self, node_id: &str) -> Option<&Vec<String>> {
102        self.indexes.neighbors(node_id)
103    }
104    
105    /// Get predecessors of a node (incoming edges)
106    pub fn predecessors(&self, node_id: &str) -> Option<&Vec<String>> {
107        self.indexes.predecessors(node_id)
108    }
109    
110    /// Check if node exists in graph
111    pub fn has_node(&self, node_id: &str) -> bool {
112        self.indexes.has_node(node_id)
113    }
114    
115    /// Get edge weight between two nodes
116    pub fn edge_weight(&self, source: &str, target: &str) -> Option<f64> {
117        self.indexes.edge_weight(source, target)
118    }
119    
120    /// Get all node IDs
121    pub fn node_ids(&self) -> impl Iterator<Item = &String> {
122        self.indexes.all_nodes()
123    }
124    
125    /// Add a new node to the graph
126    pub fn add_node(&mut self, node_id: String) -> Result<()> {
127        // Check if node already exists
128        if self.has_node(&node_id) {
129            return Err(GraphError::invalid_parameter(
130                &format!("Node '{}' already exists in the graph", node_id)
131            ));
132        }
133        
134        // Create new nodes RecordBatch with the added node
135        let mut node_ids: Vec<String> = self.node_ids().cloned().collect();
136        node_ids.push(node_id);
137        
138        let nodes_schema = Arc::new(Schema::new(vec![
139            Field::new("id", DataType::Utf8, false),
140        ]));
141        
142        let new_nodes = RecordBatch::try_new(
143            nodes_schema,
144            vec![Arc::new(arrow::array::StringArray::from(node_ids))],
145        ).map_err(GraphError::from)?;
146        
147        // Rebuild graph with new nodes
148        self.nodes = new_nodes;
149        self.indexes = GraphIndexes::build(&self.nodes, &self.edges)?;
150        
151        Ok(())
152    }
153    
154    /// Remove a node from the graph (also removes all incident edges)
155    pub fn remove_node(&mut self, node_id: &str) -> Result<()> {
156        // Check if node exists
157        if !self.has_node(node_id) {
158            return Err(GraphError::invalid_parameter(
159                &format!("Node '{}' does not exist in the graph", node_id)
160            ));
161        }
162        
163        // Filter out the node
164        let remaining_nodes: Vec<String> = self.node_ids()
165            .filter(|&id| id != node_id)
166            .cloned()
167            .collect();
168        
169        // Create new nodes RecordBatch
170        let nodes_schema = Arc::new(Schema::new(vec![
171            Field::new("id", DataType::Utf8, false),
172        ]));
173        
174        let new_nodes = RecordBatch::try_new(
175            nodes_schema,
176            vec![Arc::new(arrow::array::StringArray::from(remaining_nodes))],
177        ).map_err(GraphError::from)?;
178        
179        // Filter out edges that involve the removed node
180        let source_array = self.edges.column(0)
181            .as_any()
182            .downcast_ref::<arrow::array::StringArray>()
183            .ok_or_else(|| GraphError::invalid_parameter("Invalid source column type"))?;
184            
185        let target_array = self.edges.column(1)
186            .as_any()
187            .downcast_ref::<arrow::array::StringArray>()
188            .ok_or_else(|| GraphError::invalid_parameter("Invalid target column type"))?;
189        
190        let mut remaining_sources = Vec::new();
191        let mut remaining_targets = Vec::new();
192        let mut remaining_weights = Vec::new();
193        
194        for i in 0..self.edges.num_rows() {
195            let source = source_array.value(i);
196            let target = target_array.value(i);
197            
198            // Keep edge only if neither source nor target is the removed node
199            if source != node_id && target != node_id {
200                remaining_sources.push(source.to_string());
201                remaining_targets.push(target.to_string());
202                
203                // Handle optional weight column
204                if self.edges.num_columns() > 2 {
205                    if let Some(weight_array) = self.edges.column(2)
206                        .as_any()
207                        .downcast_ref::<arrow::array::Float64Array>()
208                    {
209                        remaining_weights.push(weight_array.value(i));
210                    }
211                }
212            }
213        }
214        
215        // Create new edges RecordBatch
216        let mut edge_fields = vec![
217            Field::new("source", DataType::Utf8, false),
218            Field::new("target", DataType::Utf8, false),
219        ];
220        
221        let mut edge_columns: Vec<Arc<dyn arrow::array::Array>> = vec![
222            Arc::new(arrow::array::StringArray::from(remaining_sources)),
223            Arc::new(arrow::array::StringArray::from(remaining_targets)),
224        ];
225        
226        if !remaining_weights.is_empty() {
227            edge_fields.push(Field::new("weight", DataType::Float64, true));
228            edge_columns.push(Arc::new(arrow::array::Float64Array::from(remaining_weights)));
229        }
230        
231        let edges_schema = Arc::new(Schema::new(edge_fields));
232        let new_edges = RecordBatch::try_new(edges_schema, edge_columns)
233            .map_err(GraphError::from)?;
234        
235        // Rebuild graph
236        self.nodes = new_nodes;
237        self.edges = new_edges;
238        self.indexes = GraphIndexes::build(&self.nodes, &self.edges)?;
239        
240        Ok(())
241    }
242    
243    /// Add an edge to the graph
244    pub fn add_edge(&mut self, source: String, target: String, weight: Option<f64>) -> Result<()> {
245        // Ensure both nodes exist (add them if they don't)
246        if !self.has_node(&source) {
247            self.add_node(source.clone())?;
248        }
249        if !self.has_node(&target) {
250            self.add_node(target.clone())?;
251        }
252        
253        // Check if edge already exists
254        if self.edge_weight(&source, &target).is_some() {
255            return Err(GraphError::invalid_parameter(
256                &format!("Edge from '{}' to '{}' already exists", source, target)
257            ));
258        }
259        
260        // Get existing edges
261        let source_array = self.edges.column(0)
262            .as_any()
263            .downcast_ref::<arrow::array::StringArray>()
264            .ok_or_else(|| GraphError::invalid_parameter("Invalid source column type"))?;
265            
266        let target_array = self.edges.column(1)
267            .as_any()
268            .downcast_ref::<arrow::array::StringArray>()
269            .ok_or_else(|| GraphError::invalid_parameter("Invalid target column type"))?;
270        
271        let mut new_sources: Vec<String> = source_array.iter()
272            .map(|s| s.unwrap_or("").to_string())
273            .collect();
274        let mut new_targets: Vec<String> = target_array.iter()
275            .map(|t| t.unwrap_or("").to_string())
276            .collect();
277        
278        // Add new edge
279        new_sources.push(source);
280        new_targets.push(target);
281        
282        let mut edge_fields = vec![
283            Field::new("source", DataType::Utf8, false),
284            Field::new("target", DataType::Utf8, false),
285        ];
286        
287        let mut edge_columns: Vec<Arc<dyn arrow::array::Array>> = vec![
288            Arc::new(arrow::array::StringArray::from(new_sources)),
289            Arc::new(arrow::array::StringArray::from(new_targets)),
290        ];
291        
292        // Handle weights
293        if self.edges.num_columns() > 2 || weight.is_some() {
294            let mut new_weights = Vec::new();
295            
296            // Copy existing weights
297            if self.edges.num_columns() > 2 {
298                if let Some(weight_array) = self.edges.column(2)
299                    .as_any()
300                    .downcast_ref::<arrow::array::Float64Array>()
301                {
302                    for i in 0..weight_array.len() {
303                        new_weights.push(Some(weight_array.value(i)));
304                    }
305                }
306            } else {
307                // Fill with None for existing edges if this is the first weighted edge
308                for _ in 0..self.edges.num_rows() {
309                    new_weights.push(None);
310                }
311            }
312            
313            // Add weight for new edge
314            new_weights.push(weight);
315            
316            edge_fields.push(Field::new("weight", DataType::Float64, true));
317            edge_columns.push(Arc::new(arrow::array::Float64Array::from(new_weights)));
318        }
319        
320        let edges_schema = Arc::new(Schema::new(edge_fields));
321        let new_edges = RecordBatch::try_new(edges_schema, edge_columns)
322            .map_err(GraphError::from)?;
323        
324        // Rebuild graph
325        self.edges = new_edges;
326        self.indexes = GraphIndexes::build(&self.nodes, &self.edges)?;
327        
328        Ok(())
329    }
330    
331    /// Remove an edge from the graph
332    pub fn remove_edge(&mut self, source: &str, target: &str) -> Result<()> {
333        // Check if edge exists
334        if self.edge_weight(source, target).is_none() {
335            return Err(GraphError::invalid_parameter(
336                &format!("Edge from '{}' to '{}' does not exist", source, target)
337            ));
338        }
339        
340        // Filter out the edge
341        let source_array = self.edges.column(0)
342            .as_any()
343            .downcast_ref::<arrow::array::StringArray>()
344            .ok_or_else(|| GraphError::invalid_parameter("Invalid source column type"))?;
345            
346        let target_array = self.edges.column(1)
347            .as_any()
348            .downcast_ref::<arrow::array::StringArray>()
349            .ok_or_else(|| GraphError::invalid_parameter("Invalid target column type"))?;
350        
351        let mut remaining_sources = Vec::new();
352        let mut remaining_targets = Vec::new();
353        let mut remaining_weights = Vec::new();
354        
355        for i in 0..self.edges.num_rows() {
356            let edge_source = source_array.value(i);
357            let edge_target = target_array.value(i);
358            
359            // Keep edge only if it's not the one we want to remove
360            if !(edge_source == source && edge_target == target) {
361                remaining_sources.push(edge_source.to_string());
362                remaining_targets.push(edge_target.to_string());
363                
364                // Handle optional weight column
365                if self.edges.num_columns() > 2 {
366                    if let Some(weight_array) = self.edges.column(2)
367                        .as_any()
368                        .downcast_ref::<arrow::array::Float64Array>()
369                    {
370                        remaining_weights.push(Some(weight_array.value(i)));
371                    }
372                }
373            }
374        }
375        
376        // Create new edges RecordBatch
377        let mut edge_fields = vec![
378            Field::new("source", DataType::Utf8, false),
379            Field::new("target", DataType::Utf8, false),
380        ];
381        
382        let mut edge_columns: Vec<Arc<dyn arrow::array::Array>> = vec![
383            Arc::new(arrow::array::StringArray::from(remaining_sources)),
384            Arc::new(arrow::array::StringArray::from(remaining_targets)),
385        ];
386        
387        if !remaining_weights.is_empty() {
388            edge_fields.push(Field::new("weight", DataType::Float64, true));
389            edge_columns.push(Arc::new(arrow::array::Float64Array::from(remaining_weights)));
390        }
391        
392        let edges_schema = Arc::new(Schema::new(edge_fields));
393        let new_edges = RecordBatch::try_new(edges_schema, edge_columns)
394            .map_err(GraphError::from)?;
395        
396        // Rebuild graph
397        self.edges = new_edges;
398        self.indexes = GraphIndexes::build(&self.nodes, &self.edges)?;
399        
400        Ok(())
401    }
402}