Skip to main content

astraea_server/
grpc.rs

1//! gRPC transport layer for AstraeaDB.
2//!
3//! This module provides a thin adapter that converts gRPC requests into the
4//! existing [`Request`] enum, delegates to [`RequestHandler::handle`], and
5//! converts the [`Response`] back into gRPC response types.
6//!
7//! It does **not** duplicate any business logic -- all processing goes through
8//! the same handler that the TCP server uses.
9
10use std::sync::Arc;
11
12use tonic::{self, Status};
13use tracing::info;
14
15use crate::handler::RequestHandler;
16use crate::protocol::{Request, Response};
17
18// Pull in the generated protobuf / gRPC types.
19pub mod proto {
20    tonic::include_proto!("astraea");
21}
22
23use proto::astraea_service_server::{AstraeaService, AstraeaServiceServer};
24use proto::*;
25
26// ---------------------------------------------------------------------------
27// Service implementation
28// ---------------------------------------------------------------------------
29
30/// The gRPC service implementation. Holds a shared reference to the same
31/// [`RequestHandler`] that the TCP server uses.
32pub struct AstraeaGrpcService {
33    handler: Arc<RequestHandler>,
34}
35
36impl AstraeaGrpcService {
37    pub fn new(handler: Arc<RequestHandler>) -> Self {
38        Self { handler }
39    }
40
41    /// Build a `tonic` service that can be added to a [`tonic::transport::Server`].
42    pub fn into_service(self) -> AstraeaServiceServer<Self> {
43        AstraeaServiceServer::new(self)
44    }
45}
46
47// ---------------------------------------------------------------------------
48// Helper: extract data / error from Response
49// ---------------------------------------------------------------------------
50
51/// Convert our internal [`Response`] to `(success, result_json, error)`.
52fn response_to_parts(resp: Response) -> (bool, String, String) {
53    match resp {
54        Response::Ok { data } => {
55            let json = serde_json::to_string(&data).unwrap_or_default();
56            (true, json, String::new())
57        }
58        Response::Error { message } => (false, String::new(), message),
59    }
60}
61
62// ---------------------------------------------------------------------------
63// Trait implementation
64// ---------------------------------------------------------------------------
65
66#[tonic::async_trait]
67impl AstraeaService for AstraeaGrpcService {
68    // -- Node CRUD ---------------------------------------------------------
69
70    async fn create_node(
71        &self,
72        request: tonic::Request<CreateNodeRequest>,
73    ) -> Result<tonic::Response<MutationResponse>, Status> {
74        let req = request.into_inner();
75
76        let properties: serde_json::Value = if req.properties_json.is_empty() {
77            serde_json::json!({})
78        } else {
79            serde_json::from_str(&req.properties_json)
80                .map_err(|e| Status::invalid_argument(format!("invalid properties JSON: {e}")))?
81        };
82
83        let embedding = if req.embedding.is_empty() {
84            None
85        } else {
86            Some(req.embedding)
87        };
88
89        let internal = Request::CreateNode {
90            labels: req.labels,
91            properties,
92            embedding,
93        };
94
95        let (success, result_json, error) = response_to_parts(self.handler.handle(internal));
96        Ok(tonic::Response::new(MutationResponse {
97            success,
98            result_json,
99            error,
100        }))
101    }
102
103    async fn get_node(
104        &self,
105        request: tonic::Request<GetNodeRequest>,
106    ) -> Result<tonic::Response<GetNodeResponse>, Status> {
107        let req = request.into_inner();
108        let internal = Request::GetNode { id: req.id };
109        let resp = self.handler.handle(internal);
110
111        match resp {
112            Response::Ok { data } => {
113                let id = data.get("id").and_then(|v| v.as_u64()).unwrap_or(0);
114                let labels: Vec<String> = data
115                    .get("labels")
116                    .and_then(|v| serde_json::from_value(v.clone()).ok())
117                    .unwrap_or_default();
118                let properties_json = data
119                    .get("properties")
120                    .map(|v| v.to_string())
121                    .unwrap_or_else(|| "{}".into());
122                let has_embedding = data
123                    .get("has_embedding")
124                    .and_then(|v| v.as_bool())
125                    .unwrap_or(false);
126
127                Ok(tonic::Response::new(GetNodeResponse {
128                    found: true,
129                    id,
130                    labels,
131                    properties_json,
132                    has_embedding,
133                    error: String::new(),
134                }))
135            }
136            Response::Error { message } => Ok(tonic::Response::new(GetNodeResponse {
137                found: false,
138                id: 0,
139                labels: vec![],
140                properties_json: String::new(),
141                has_embedding: false,
142                error: message,
143            })),
144        }
145    }
146
147    async fn update_node(
148        &self,
149        request: tonic::Request<UpdateNodeRequest>,
150    ) -> Result<tonic::Response<MutationResponse>, Status> {
151        let req = request.into_inner();
152
153        let properties: serde_json::Value = if req.properties_json.is_empty() {
154            serde_json::json!({})
155        } else {
156            serde_json::from_str(&req.properties_json)
157                .map_err(|e| Status::invalid_argument(format!("invalid properties JSON: {e}")))?
158        };
159
160        let internal = Request::UpdateNode {
161            id: req.id,
162            properties,
163        };
164        let (success, result_json, error) = response_to_parts(self.handler.handle(internal));
165        Ok(tonic::Response::new(MutationResponse {
166            success,
167            result_json,
168            error,
169        }))
170    }
171
172    async fn delete_node(
173        &self,
174        request: tonic::Request<DeleteNodeRequest>,
175    ) -> Result<tonic::Response<MutationResponse>, Status> {
176        let req = request.into_inner();
177        let internal = Request::DeleteNode { id: req.id };
178        let (success, result_json, error) = response_to_parts(self.handler.handle(internal));
179        Ok(tonic::Response::new(MutationResponse {
180            success,
181            result_json,
182            error,
183        }))
184    }
185
186    // -- Edge CRUD ---------------------------------------------------------
187
188    async fn create_edge(
189        &self,
190        request: tonic::Request<CreateEdgeRequest>,
191    ) -> Result<tonic::Response<MutationResponse>, Status> {
192        let req = request.into_inner();
193
194        let properties: serde_json::Value = if req.properties_json.is_empty() {
195            serde_json::json!({})
196        } else {
197            serde_json::from_str(&req.properties_json)
198                .map_err(|e| Status::invalid_argument(format!("invalid properties JSON: {e}")))?
199        };
200
201        let weight = if req.weight == 0.0 { 1.0 } else { req.weight };
202
203        let internal = Request::CreateEdge {
204            source: req.source,
205            target: req.target,
206            edge_type: req.edge_type,
207            properties,
208            weight,
209            valid_from: req.valid_from,
210            valid_to: req.valid_to,
211        };
212
213        let (success, result_json, error) = response_to_parts(self.handler.handle(internal));
214        Ok(tonic::Response::new(MutationResponse {
215            success,
216            result_json,
217            error,
218        }))
219    }
220
221    async fn get_edge(
222        &self,
223        request: tonic::Request<GetEdgeRequest>,
224    ) -> Result<tonic::Response<GetEdgeResponse>, Status> {
225        let req = request.into_inner();
226        let internal = Request::GetEdge { id: req.id };
227        let resp = self.handler.handle(internal);
228
229        match resp {
230            Response::Ok { data } => {
231                let id = data.get("id").and_then(|v| v.as_u64()).unwrap_or(0);
232                let source = data.get("source").and_then(|v| v.as_u64()).unwrap_or(0);
233                let target = data.get("target").and_then(|v| v.as_u64()).unwrap_or(0);
234                let edge_type = data
235                    .get("edge_type")
236                    .and_then(|v| v.as_str())
237                    .unwrap_or("")
238                    .to_string();
239                let properties_json = data
240                    .get("properties")
241                    .map(|v| v.to_string())
242                    .unwrap_or_else(|| "{}".into());
243                let weight = data.get("weight").and_then(|v| v.as_f64()).unwrap_or(0.0);
244                let valid_from = data.get("valid_from").and_then(|v| v.as_i64());
245                let valid_to = data.get("valid_to").and_then(|v| v.as_i64());
246
247                Ok(tonic::Response::new(GetEdgeResponse {
248                    found: true,
249                    id,
250                    source,
251                    target,
252                    edge_type,
253                    properties_json,
254                    weight,
255                    valid_from,
256                    valid_to,
257                    error: String::new(),
258                }))
259            }
260            Response::Error { message } => Ok(tonic::Response::new(GetEdgeResponse {
261                found: false,
262                id: 0,
263                source: 0,
264                target: 0,
265                edge_type: String::new(),
266                properties_json: String::new(),
267                weight: 0.0,
268                valid_from: None,
269                valid_to: None,
270                error: message,
271            })),
272        }
273    }
274
275    async fn update_edge(
276        &self,
277        request: tonic::Request<UpdateEdgeRequest>,
278    ) -> Result<tonic::Response<MutationResponse>, Status> {
279        let req = request.into_inner();
280
281        let properties: serde_json::Value = if req.properties_json.is_empty() {
282            serde_json::json!({})
283        } else {
284            serde_json::from_str(&req.properties_json)
285                .map_err(|e| Status::invalid_argument(format!("invalid properties JSON: {e}")))?
286        };
287
288        let internal = Request::UpdateEdge {
289            id: req.id,
290            properties,
291        };
292        let (success, result_json, error) = response_to_parts(self.handler.handle(internal));
293        Ok(tonic::Response::new(MutationResponse {
294            success,
295            result_json,
296            error,
297        }))
298    }
299
300    async fn delete_edge(
301        &self,
302        request: tonic::Request<DeleteEdgeRequest>,
303    ) -> Result<tonic::Response<MutationResponse>, Status> {
304        let req = request.into_inner();
305        let internal = Request::DeleteEdge { id: req.id };
306        let (success, result_json, error) = response_to_parts(self.handler.handle(internal));
307        Ok(tonic::Response::new(MutationResponse {
308            success,
309            result_json,
310            error,
311        }))
312    }
313
314    // -- Traversal ---------------------------------------------------------
315
316    async fn neighbors(
317        &self,
318        request: tonic::Request<NeighborsRequest>,
319    ) -> Result<tonic::Response<NeighborsResponse>, Status> {
320        let req = request.into_inner();
321
322        let edge_type = if req.edge_type.is_empty() {
323            None
324        } else {
325            Some(req.edge_type)
326        };
327
328        let internal = Request::Neighbors {
329            id: req.id,
330            direction: if req.direction.is_empty() {
331                "outgoing".into()
332            } else {
333                req.direction
334            },
335            edge_type,
336        };
337
338        let resp = self.handler.handle(internal);
339
340        match resp {
341            Response::Ok { data } => {
342                let neighbors = data
343                    .get("neighbors")
344                    .and_then(|v| v.as_array())
345                    .map(|arr| {
346                        arr.iter()
347                            .map(|entry| NeighborEntry {
348                                edge_id: entry.get("edge_id").and_then(|v| v.as_u64()).unwrap_or(0),
349                                node_id: entry.get("node_id").and_then(|v| v.as_u64()).unwrap_or(0),
350                            })
351                            .collect()
352                    })
353                    .unwrap_or_default();
354
355                Ok(tonic::Response::new(NeighborsResponse {
356                    neighbors,
357                    error: String::new(),
358                }))
359            }
360            Response::Error { message } => Ok(tonic::Response::new(NeighborsResponse {
361                neighbors: vec![],
362                error: message,
363            })),
364        }
365    }
366
367    async fn bfs(
368        &self,
369        request: tonic::Request<BfsRequest>,
370    ) -> Result<tonic::Response<BfsResponse>, Status> {
371        let req = request.into_inner();
372
373        let max_depth = if req.max_depth == 0 {
374            3
375        } else {
376            req.max_depth as usize
377        };
378
379        let internal = Request::Bfs {
380            start: req.start,
381            max_depth,
382        };
383
384        let resp = self.handler.handle(internal);
385
386        match resp {
387            Response::Ok { data } => {
388                let nodes = data
389                    .get("nodes")
390                    .and_then(|v| v.as_array())
391                    .map(|arr| {
392                        arr.iter()
393                            .map(|entry| BfsEntry {
394                                node_id: entry.get("node_id").and_then(|v| v.as_u64()).unwrap_or(0),
395                                depth: entry.get("depth").and_then(|v| v.as_u64()).unwrap_or(0)
396                                    as u32,
397                            })
398                            .collect()
399                    })
400                    .unwrap_or_default();
401
402                Ok(tonic::Response::new(BfsResponse {
403                    nodes,
404                    error: String::new(),
405                }))
406            }
407            Response::Error { message } => Ok(tonic::Response::new(BfsResponse {
408                nodes: vec![],
409                error: message,
410            })),
411        }
412    }
413
414    async fn shortest_path(
415        &self,
416        request: tonic::Request<ShortestPathRequest>,
417    ) -> Result<tonic::Response<ShortestPathResponse>, Status> {
418        let req = request.into_inner();
419
420        let internal = Request::ShortestPath {
421            from: req.from,
422            to: req.to,
423            weighted: req.weighted,
424        };
425
426        let resp = self.handler.handle(internal);
427
428        match resp {
429            Response::Ok { data } => {
430                let path: Vec<u64> = data
431                    .get("path")
432                    .and_then(|v| {
433                        if v.is_null() {
434                            None
435                        } else {
436                            v.as_array()
437                                .map(|arr| arr.iter().filter_map(|val| val.as_u64()).collect())
438                        }
439                    })
440                    .unwrap_or_default();
441
442                let found = !path.is_empty();
443                let length = data.get("length").and_then(|v| v.as_u64()).unwrap_or(0) as u32;
444                let cost = data.get("cost").and_then(|v| v.as_f64());
445
446                Ok(tonic::Response::new(ShortestPathResponse {
447                    found,
448                    path,
449                    length,
450                    cost,
451                    error: String::new(),
452                }))
453            }
454            Response::Error { message } => Ok(tonic::Response::new(ShortestPathResponse {
455                found: false,
456                path: vec![],
457                length: 0,
458                cost: None,
459                error: message,
460            })),
461        }
462    }
463
464    // -- Vector search -----------------------------------------------------
465
466    async fn vector_search(
467        &self,
468        request: tonic::Request<VectorSearchRequest>,
469    ) -> Result<tonic::Response<VectorSearchResponse>, Status> {
470        let req = request.into_inner();
471
472        let k = if req.k == 0 { 10 } else { req.k as usize };
473
474        let internal = Request::VectorSearch {
475            query: req.query,
476            k,
477        };
478
479        let resp = self.handler.handle(internal);
480
481        match resp {
482            Response::Ok { data } => {
483                let results = data
484                    .get("results")
485                    .and_then(|v| v.as_array())
486                    .map(|arr| {
487                        arr.iter()
488                            .map(|entry| VectorSearchResult {
489                                node_id: entry.get("node_id").and_then(|v| v.as_u64()).unwrap_or(0),
490                                // The proto field is now canonically named
491                                // `distance` (astraeadb-issues.md #6), matching
492                                // the TCP JSON wire format exactly, so no
493                                // fallback lookup is needed here anymore.
494                                distance: entry
495                                    .get("distance")
496                                    .and_then(|v| v.as_f64())
497                                    .unwrap_or(0.0)
498                                    as f32,
499                            })
500                            .collect()
501                    })
502                    .unwrap_or_default();
503
504                Ok(tonic::Response::new(VectorSearchResponse {
505                    results,
506                    error: String::new(),
507                }))
508            }
509            Response::Error { message } => Ok(tonic::Response::new(VectorSearchResponse {
510                results: vec![],
511                error: message,
512            })),
513        }
514    }
515
516    // -- GQL query ---------------------------------------------------------
517
518    async fn query(
519        &self,
520        request: tonic::Request<QueryRequest>,
521    ) -> Result<tonic::Response<QueryResponse>, Status> {
522        let req = request.into_inner();
523        let internal = Request::Query { gql: req.gql };
524        let (success, result_json, error) = response_to_parts(self.handler.handle(internal));
525
526        Ok(tonic::Response::new(QueryResponse {
527            success,
528            result_json,
529            error,
530        }))
531    }
532
533    // -- Temporal traversals -----------------------------------------------
534
535    async fn neighbors_at(
536        &self,
537        request: tonic::Request<NeighborsAtRequest>,
538    ) -> Result<tonic::Response<NeighborsResponse>, Status> {
539        let req = request.into_inner();
540        let edge_type = if req.edge_type.is_empty() {
541            None
542        } else {
543            Some(req.edge_type)
544        };
545        let internal = Request::NeighborsAt {
546            id: req.id,
547            direction: if req.direction.is_empty() {
548                "outgoing".into()
549            } else {
550                req.direction
551            },
552            timestamp: req.timestamp,
553            edge_type,
554        };
555        let resp = self.handler.handle(internal);
556        match resp {
557            Response::Ok { data } => {
558                let neighbors = data
559                    .get("neighbors")
560                    .and_then(|v| v.as_array())
561                    .map(|arr| {
562                        arr.iter()
563                            .map(|entry| NeighborEntry {
564                                edge_id: entry.get("edge_id").and_then(|v| v.as_u64()).unwrap_or(0),
565                                node_id: entry.get("node_id").and_then(|v| v.as_u64()).unwrap_or(0),
566                            })
567                            .collect()
568                    })
569                    .unwrap_or_default();
570                Ok(tonic::Response::new(NeighborsResponse {
571                    neighbors,
572                    error: String::new(),
573                }))
574            }
575            Response::Error { message } => Ok(tonic::Response::new(NeighborsResponse {
576                neighbors: vec![],
577                error: message,
578            })),
579        }
580    }
581
582    async fn bfs_at(
583        &self,
584        request: tonic::Request<BfsAtRequest>,
585    ) -> Result<tonic::Response<BfsResponse>, Status> {
586        let req = request.into_inner();
587        let max_depth = if req.max_depth == 0 {
588            3
589        } else {
590            req.max_depth as usize
591        };
592        let internal = Request::BfsAt {
593            start: req.start,
594            max_depth,
595            timestamp: req.timestamp,
596        };
597        let resp = self.handler.handle(internal);
598        match resp {
599            Response::Ok { data } => {
600                let nodes = data
601                    .get("nodes")
602                    .and_then(|v| v.as_array())
603                    .map(|arr| {
604                        arr.iter()
605                            .map(|entry| BfsEntry {
606                                node_id: entry.get("node_id").and_then(|v| v.as_u64()).unwrap_or(0),
607                                depth: entry.get("depth").and_then(|v| v.as_u64()).unwrap_or(0)
608                                    as u32,
609                            })
610                            .collect()
611                    })
612                    .unwrap_or_default();
613                Ok(tonic::Response::new(BfsResponse {
614                    nodes,
615                    error: String::new(),
616                }))
617            }
618            Response::Error { message } => Ok(tonic::Response::new(BfsResponse {
619                nodes: vec![],
620                error: message,
621            })),
622        }
623    }
624
625    async fn shortest_path_at(
626        &self,
627        request: tonic::Request<ShortestPathAtRequest>,
628    ) -> Result<tonic::Response<ShortestPathResponse>, Status> {
629        let req = request.into_inner();
630        let internal = Request::ShortestPathAt {
631            from: req.from,
632            to: req.to,
633            timestamp: req.timestamp,
634            weighted: req.weighted,
635        };
636        let resp = self.handler.handle(internal);
637        match resp {
638            Response::Ok { data } => {
639                let path: Vec<u64> = data
640                    .get("path")
641                    .and_then(|v| {
642                        if v.is_null() {
643                            None
644                        } else {
645                            v.as_array()
646                                .map(|arr| arr.iter().filter_map(|val| val.as_u64()).collect())
647                        }
648                    })
649                    .unwrap_or_default();
650                let found = !path.is_empty();
651                let length = data.get("length").and_then(|v| v.as_u64()).unwrap_or(0) as u32;
652                let cost = data.get("cost").and_then(|v| v.as_f64());
653                Ok(tonic::Response::new(ShortestPathResponse {
654                    found,
655                    path,
656                    length,
657                    cost,
658                    error: String::new(),
659                }))
660            }
661            Response::Error { message } => Ok(tonic::Response::new(ShortestPathResponse {
662                found: false,
663                path: vec![],
664                length: 0,
665                cost: None,
666                error: message,
667            })),
668        }
669    }
670
671    // -- DFS ---------------------------------------------------------------
672
673    async fn dfs(
674        &self,
675        request: tonic::Request<DfsRequest>,
676    ) -> Result<tonic::Response<DfsResponse>, Status> {
677        let req = request.into_inner();
678        let max_depth = if req.max_depth == 0 {
679            3
680        } else {
681            req.max_depth as usize
682        };
683        let internal = Request::Dfs {
684            start: req.start,
685            max_depth,
686        };
687        let resp = self.handler.handle(internal);
688        match resp {
689            Response::Ok { data } => {
690                let nodes: Vec<u64> = data
691                    .get("nodes")
692                    .and_then(|v| v.as_array())
693                    .map(|arr| arr.iter().filter_map(|v| v.as_u64()).collect())
694                    .unwrap_or_default();
695                Ok(tonic::Response::new(DfsResponse {
696                    nodes,
697                    error: String::new(),
698                }))
699            }
700            Response::Error { message } => Ok(tonic::Response::new(DfsResponse {
701                nodes: vec![],
702                error: message,
703            })),
704        }
705    }
706
707    async fn dfs_at(
708        &self,
709        request: tonic::Request<DfsAtRequest>,
710    ) -> Result<tonic::Response<DfsResponse>, Status> {
711        let req = request.into_inner();
712        let max_depth = if req.max_depth == 0 {
713            3
714        } else {
715            req.max_depth as usize
716        };
717        let internal = Request::DfsAt {
718            start: req.start,
719            max_depth,
720            timestamp: req.timestamp,
721        };
722        let resp = self.handler.handle(internal);
723        match resp {
724            Response::Ok { data } => {
725                let nodes: Vec<u64> = data
726                    .get("nodes")
727                    .and_then(|v| v.as_array())
728                    .map(|arr| arr.iter().filter_map(|v| v.as_u64()).collect())
729                    .unwrap_or_default();
730                Ok(tonic::Response::new(DfsResponse {
731                    nodes,
732                    error: String::new(),
733                }))
734            }
735            Response::Error { message } => Ok(tonic::Response::new(DfsResponse {
736                nodes: vec![],
737                error: message,
738            })),
739        }
740    }
741
742    // -- Label lookup ------------------------------------------------------
743
744    async fn find_by_label(
745        &self,
746        request: tonic::Request<FindByLabelRequest>,
747    ) -> Result<tonic::Response<FindByLabelResponse>, Status> {
748        let req = request.into_inner();
749        let internal = Request::FindByLabel { label: req.label };
750        let resp = self.handler.handle(internal);
751        match resp {
752            Response::Ok { data } => {
753                let node_ids: Vec<u64> = data
754                    .get("node_ids")
755                    .and_then(|v| v.as_array())
756                    .map(|arr| arr.iter().filter_map(|v| v.as_u64()).collect())
757                    .unwrap_or_default();
758                Ok(tonic::Response::new(FindByLabelResponse {
759                    node_ids,
760                    error: String::new(),
761                }))
762            }
763            Response::Error { message } => Ok(tonic::Response::new(FindByLabelResponse {
764                node_ids: vec![],
765                error: message,
766            })),
767        }
768    }
769
770    // -- Hybrid / Semantic search ------------------------------------------
771
772    async fn hybrid_search(
773        &self,
774        request: tonic::Request<HybridSearchRequest>,
775    ) -> Result<tonic::Response<GenericJsonResponse>, Status> {
776        let req = request.into_inner();
777        let internal = Request::HybridSearch {
778            anchor: req.anchor,
779            query: req.query,
780            max_hops: if req.max_hops == 0 {
781                3
782            } else {
783                req.max_hops as usize
784            },
785            k: if req.k == 0 { 10 } else { req.k as usize },
786            alpha: if req.alpha == 0.0 { 0.5 } else { req.alpha },
787        };
788        let (success, result_json, error) = response_to_parts(self.handler.handle(internal));
789        Ok(tonic::Response::new(GenericJsonResponse {
790            success,
791            result_json,
792            error,
793        }))
794    }
795
796    async fn semantic_neighbors(
797        &self,
798        request: tonic::Request<SemanticNeighborsRequest>,
799    ) -> Result<tonic::Response<GenericJsonResponse>, Status> {
800        let req = request.into_inner();
801        let internal = Request::SemanticNeighbors {
802            id: req.id,
803            concept: req.concept,
804            direction: if req.direction.is_empty() {
805                "outgoing".into()
806            } else {
807                req.direction
808            },
809            k: if req.k == 0 { 10 } else { req.k as usize },
810        };
811        let (success, result_json, error) = response_to_parts(self.handler.handle(internal));
812        Ok(tonic::Response::new(GenericJsonResponse {
813            success,
814            result_json,
815            error,
816        }))
817    }
818
819    async fn semantic_walk(
820        &self,
821        request: tonic::Request<SemanticWalkRequest>,
822    ) -> Result<tonic::Response<GenericJsonResponse>, Status> {
823        let req = request.into_inner();
824        let internal = Request::SemanticWalk {
825            start: req.start,
826            concept: req.concept,
827            max_hops: if req.max_hops == 0 {
828                3
829            } else {
830                req.max_hops as usize
831            },
832        };
833        let (success, result_json, error) = response_to_parts(self.handler.handle(internal));
834        Ok(tonic::Response::new(GenericJsonResponse {
835            success,
836            result_json,
837            error,
838        }))
839    }
840
841    // -- Graph algorithms --------------------------------------------------
842
843    async fn run_page_rank(
844        &self,
845        request: tonic::Request<RunPageRankRequest>,
846    ) -> Result<tonic::Response<GenericJsonResponse>, Status> {
847        let req = request.into_inner();
848        let nodes = if req.nodes.is_empty() {
849            None
850        } else {
851            Some(req.nodes)
852        };
853        let internal = Request::RunPageRank {
854            nodes,
855            damping: if req.damping == 0.0 {
856                0.85
857            } else {
858                req.damping
859            },
860            max_iterations: if req.max_iterations == 0 {
861                100
862            } else {
863                req.max_iterations as usize
864            },
865            tolerance: if req.tolerance == 0.0 {
866                1e-6
867            } else {
868                req.tolerance
869            },
870        };
871        let (success, result_json, error) = response_to_parts(self.handler.handle(internal));
872        Ok(tonic::Response::new(GenericJsonResponse {
873            success,
874            result_json,
875            error,
876        }))
877    }
878
879    async fn run_louvain(
880        &self,
881        request: tonic::Request<RunLouvainRequest>,
882    ) -> Result<tonic::Response<GenericJsonResponse>, Status> {
883        let req = request.into_inner();
884        let nodes = if req.nodes.is_empty() {
885            None
886        } else {
887            Some(req.nodes)
888        };
889        let internal = Request::RunLouvain { nodes };
890        let (success, result_json, error) = response_to_parts(self.handler.handle(internal));
891        Ok(tonic::Response::new(GenericJsonResponse {
892            success,
893            result_json,
894            error,
895        }))
896    }
897
898    async fn run_connected_components(
899        &self,
900        request: tonic::Request<RunConnectedComponentsRequest>,
901    ) -> Result<tonic::Response<GenericJsonResponse>, Status> {
902        let req = request.into_inner();
903        let nodes = if req.nodes.is_empty() {
904            None
905        } else {
906            Some(req.nodes)
907        };
908        let internal = Request::RunConnectedComponents {
909            nodes,
910            strong: req.strong,
911        };
912        let (success, result_json, error) = response_to_parts(self.handler.handle(internal));
913        Ok(tonic::Response::new(GenericJsonResponse {
914            success,
915            result_json,
916            error,
917        }))
918    }
919
920    async fn run_degree_centrality(
921        &self,
922        request: tonic::Request<RunDegreeCentralityRequest>,
923    ) -> Result<tonic::Response<GenericJsonResponse>, Status> {
924        let req = request.into_inner();
925        let nodes = if req.nodes.is_empty() {
926            None
927        } else {
928            Some(req.nodes)
929        };
930        let internal = Request::RunDegreeCentrality {
931            nodes,
932            direction: if req.direction.is_empty() {
933                "both".into()
934            } else {
935                req.direction
936            },
937        };
938        let (success, result_json, error) = response_to_parts(self.handler.handle(internal));
939        Ok(tonic::Response::new(GenericJsonResponse {
940            success,
941            result_json,
942            error,
943        }))
944    }
945
946    async fn run_betweenness_centrality(
947        &self,
948        request: tonic::Request<RunBetweennessCentralityRequest>,
949    ) -> Result<tonic::Response<GenericJsonResponse>, Status> {
950        let req = request.into_inner();
951        let nodes = if req.nodes.is_empty() {
952            None
953        } else {
954            Some(req.nodes)
955        };
956        let internal = Request::RunBetweennessCentrality { nodes };
957        let (success, result_json, error) = response_to_parts(self.handler.handle(internal));
958        Ok(tonic::Response::new(GenericJsonResponse {
959            success,
960            result_json,
961            error,
962        }))
963    }
964
965    // -- Graph stats & subgraph --------------------------------------------
966
967    async fn graph_stats(
968        &self,
969        _request: tonic::Request<GraphStatsRequest>,
970    ) -> Result<tonic::Response<GenericJsonResponse>, Status> {
971        let internal = Request::GraphStats;
972        let (success, result_json, error) = response_to_parts(self.handler.handle(internal));
973        Ok(tonic::Response::new(GenericJsonResponse {
974            success,
975            result_json,
976            error,
977        }))
978    }
979
980    async fn get_subgraph(
981        &self,
982        request: tonic::Request<GetSubgraphRequest>,
983    ) -> Result<tonic::Response<GenericJsonResponse>, Status> {
984        let req = request.into_inner();
985        let internal = Request::GetSubgraph {
986            center: req.center,
987            hops: if req.hops == 0 { 3 } else { req.hops as usize },
988            max_nodes: if req.max_nodes == 0 {
989                50
990            } else {
991                req.max_nodes as usize
992            },
993        };
994        let (success, result_json, error) = response_to_parts(self.handler.handle(internal));
995        Ok(tonic::Response::new(GenericJsonResponse {
996            success,
997            result_json,
998            error,
999        }))
1000    }
1001
1002    // -- RAG ---------------------------------------------------------------
1003
1004    async fn extract_subgraph(
1005        &self,
1006        request: tonic::Request<ExtractSubgraphRequest>,
1007    ) -> Result<tonic::Response<GenericJsonResponse>, Status> {
1008        let req = request.into_inner();
1009        let internal = Request::ExtractSubgraph {
1010            center: req.center,
1011            hops: if req.hops == 0 { 3 } else { req.hops as usize },
1012            max_nodes: if req.max_nodes == 0 {
1013                50
1014            } else {
1015                req.max_nodes as usize
1016            },
1017            format: if req.format.is_empty() {
1018                "structured".into()
1019            } else {
1020                req.format
1021            },
1022        };
1023        let (success, result_json, error) = response_to_parts(self.handler.handle(internal));
1024        Ok(tonic::Response::new(GenericJsonResponse {
1025            success,
1026            result_json,
1027            error,
1028        }))
1029    }
1030
1031    async fn graph_rag(
1032        &self,
1033        request: tonic::Request<GraphRagRequest>,
1034    ) -> Result<tonic::Response<GenericJsonResponse>, Status> {
1035        let req = request.into_inner();
1036        let question_embedding = if req.question_embedding.is_empty() {
1037            None
1038        } else {
1039            Some(req.question_embedding)
1040        };
1041        let anchor = if req.anchor == 0 {
1042            None
1043        } else {
1044            Some(req.anchor)
1045        };
1046        let internal = Request::GraphRag {
1047            question: req.question,
1048            question_embedding,
1049            anchor,
1050            hops: if req.hops == 0 { 3 } else { req.hops as usize },
1051            max_nodes: if req.max_nodes == 0 {
1052                50
1053            } else {
1054                req.max_nodes as usize
1055            },
1056            format: if req.format.is_empty() {
1057                "structured".into()
1058            } else {
1059                req.format
1060            },
1061        };
1062        let (success, result_json, error) = response_to_parts(self.handler.handle(internal));
1063        Ok(tonic::Response::new(GenericJsonResponse {
1064            success,
1065            result_json,
1066            error,
1067        }))
1068    }
1069
1070    // -- Health check ------------------------------------------------------
1071
1072    async fn ping(
1073        &self,
1074        _request: tonic::Request<PingRequest>,
1075    ) -> Result<tonic::Response<PingResponse>, Status> {
1076        let internal = Request::Ping;
1077        let resp = self.handler.handle(internal);
1078
1079        match resp {
1080            Response::Ok { data } => {
1081                let pong = data.get("pong").and_then(|v| v.as_bool()).unwrap_or(true);
1082                let version = data
1083                    .get("version")
1084                    .and_then(|v| v.as_str())
1085                    .unwrap_or("unknown")
1086                    .to_string();
1087
1088                Ok(tonic::Response::new(PingResponse { pong, version }))
1089            }
1090            Response::Error { message } => Err(Status::internal(message)),
1091        }
1092    }
1093}
1094
1095// ---------------------------------------------------------------------------
1096// Server startup helper
1097// ---------------------------------------------------------------------------
1098
1099/// Start the gRPC server on the given address.
1100///
1101/// This function runs until the server is shut down (e.g. by dropping the
1102/// tokio runtime or sending a signal).
1103pub async fn run_grpc_server(
1104    bind_addr: impl Into<String>,
1105    port: u16,
1106    handler: Arc<RequestHandler>,
1107) -> Result<(), Box<dyn std::error::Error + Send + Sync>> {
1108    let addr = format!("{}:{}", bind_addr.into(), port).parse()?;
1109
1110    let service = AstraeaGrpcService::new(handler);
1111
1112    info!("AstraeaDB gRPC server listening on {}", addr);
1113
1114    tonic::transport::Server::builder()
1115        .add_service(service.into_service())
1116        .serve(addr)
1117        .await?;
1118
1119    Ok(())
1120}
1121
1122// ---------------------------------------------------------------------------
1123// Tests
1124// ---------------------------------------------------------------------------
1125
1126#[cfg(test)]
1127mod tests {
1128    use super::proto::astraea_service_server::AstraeaService;
1129    use super::*;
1130
1131    /// Build a handler backed by an in-memory graph for testing.
1132    fn test_handler() -> Arc<RequestHandler> {
1133        let storage = astraea_graph::test_utils::InMemoryStorage::new();
1134        let graph = astraea_graph::Graph::new(Box::new(storage));
1135        let graph = std::sync::Arc::new(graph);
1136        Arc::new(RequestHandler::new(graph, None))
1137    }
1138
1139    #[tokio::test]
1140    async fn grpc_ping() {
1141        let svc = AstraeaGrpcService::new(test_handler());
1142        let resp = svc
1143            .ping(tonic::Request::new(PingRequest {}))
1144            .await
1145            .unwrap()
1146            .into_inner();
1147        assert!(resp.pong);
1148        assert!(!resp.version.is_empty());
1149    }
1150
1151    #[tokio::test]
1152    async fn grpc_create_and_get_node() {
1153        let svc = AstraeaGrpcService::new(test_handler());
1154
1155        // Create
1156        let create_resp = svc
1157            .create_node(tonic::Request::new(CreateNodeRequest {
1158                labels: vec!["Person".into()],
1159                properties_json: r#"{"name":"Alice"}"#.into(),
1160                embedding: vec![],
1161            }))
1162            .await
1163            .unwrap()
1164            .into_inner();
1165        assert!(create_resp.success, "create failed: {}", create_resp.error);
1166
1167        // Parse the created node_id from the result JSON.
1168        let result: serde_json::Value = serde_json::from_str(&create_resp.result_json).unwrap();
1169        let node_id = result.get("node_id").and_then(|v| v.as_u64()).unwrap();
1170
1171        // Get
1172        let get_resp = svc
1173            .get_node(tonic::Request::new(GetNodeRequest { id: node_id }))
1174            .await
1175            .unwrap()
1176            .into_inner();
1177        assert!(get_resp.found);
1178        assert_eq!(get_resp.id, node_id);
1179        assert_eq!(get_resp.labels, vec!["Person"]);
1180        assert!(get_resp.properties_json.contains("Alice"));
1181    }
1182
1183    #[tokio::test]
1184    async fn grpc_create_and_get_edge() {
1185        let svc = AstraeaGrpcService::new(test_handler());
1186
1187        // Create two nodes first.
1188        let resp1 = svc
1189            .create_node(tonic::Request::new(CreateNodeRequest {
1190                labels: vec!["A".into()],
1191                properties_json: "{}".into(),
1192                embedding: vec![],
1193            }))
1194            .await
1195            .unwrap()
1196            .into_inner();
1197        let n1: serde_json::Value = serde_json::from_str(&resp1.result_json).unwrap();
1198        let nid1 = n1["node_id"].as_u64().unwrap();
1199
1200        let resp2 = svc
1201            .create_node(tonic::Request::new(CreateNodeRequest {
1202                labels: vec!["B".into()],
1203                properties_json: "{}".into(),
1204                embedding: vec![],
1205            }))
1206            .await
1207            .unwrap()
1208            .into_inner();
1209        let n2: serde_json::Value = serde_json::from_str(&resp2.result_json).unwrap();
1210        let nid2 = n2["node_id"].as_u64().unwrap();
1211
1212        // Create edge.
1213        let edge_resp = svc
1214            .create_edge(tonic::Request::new(CreateEdgeRequest {
1215                source: nid1,
1216                target: nid2,
1217                edge_type: "KNOWS".into(),
1218                properties_json: r#"{"since":2024}"#.into(),
1219                weight: 1.0,
1220                valid_from: None,
1221                valid_to: None,
1222            }))
1223            .await
1224            .unwrap()
1225            .into_inner();
1226        assert!(edge_resp.success, "create edge failed: {}", edge_resp.error);
1227
1228        let result: serde_json::Value = serde_json::from_str(&edge_resp.result_json).unwrap();
1229        let edge_id = result["edge_id"].as_u64().unwrap();
1230
1231        // Get edge.
1232        let get_resp = svc
1233            .get_edge(tonic::Request::new(GetEdgeRequest { id: edge_id }))
1234            .await
1235            .unwrap()
1236            .into_inner();
1237        assert!(get_resp.found);
1238        assert_eq!(get_resp.source, nid1);
1239        assert_eq!(get_resp.target, nid2);
1240        assert_eq!(get_resp.edge_type, "KNOWS");
1241    }
1242
1243    #[tokio::test]
1244    async fn grpc_delete_node() {
1245        let svc = AstraeaGrpcService::new(test_handler());
1246
1247        // Create node.
1248        let resp = svc
1249            .create_node(tonic::Request::new(CreateNodeRequest {
1250                labels: vec!["Temp".into()],
1251                properties_json: "{}".into(),
1252                embedding: vec![],
1253            }))
1254            .await
1255            .unwrap()
1256            .into_inner();
1257        let result: serde_json::Value = serde_json::from_str(&resp.result_json).unwrap();
1258        let nid = result["node_id"].as_u64().unwrap();
1259
1260        // Delete.
1261        let del_resp = svc
1262            .delete_node(tonic::Request::new(DeleteNodeRequest { id: nid }))
1263            .await
1264            .unwrap()
1265            .into_inner();
1266        assert!(del_resp.success);
1267
1268        // Verify it is gone.
1269        let get_resp = svc
1270            .get_node(tonic::Request::new(GetNodeRequest { id: nid }))
1271            .await
1272            .unwrap()
1273            .into_inner();
1274        assert!(!get_resp.found);
1275    }
1276
1277    #[tokio::test]
1278    async fn grpc_neighbors() {
1279        let svc = AstraeaGrpcService::new(test_handler());
1280
1281        // Create two nodes and an edge.
1282        let r1 = svc
1283            .create_node(tonic::Request::new(CreateNodeRequest {
1284                labels: vec![],
1285                properties_json: "{}".into(),
1286                embedding: vec![],
1287            }))
1288            .await
1289            .unwrap()
1290            .into_inner();
1291        let nid1 = serde_json::from_str::<serde_json::Value>(&r1.result_json).unwrap()["node_id"]
1292            .as_u64()
1293            .unwrap();
1294
1295        let r2 = svc
1296            .create_node(tonic::Request::new(CreateNodeRequest {
1297                labels: vec![],
1298                properties_json: "{}".into(),
1299                embedding: vec![],
1300            }))
1301            .await
1302            .unwrap()
1303            .into_inner();
1304        let nid2 = serde_json::from_str::<serde_json::Value>(&r2.result_json).unwrap()["node_id"]
1305            .as_u64()
1306            .unwrap();
1307
1308        svc.create_edge(tonic::Request::new(CreateEdgeRequest {
1309            source: nid1,
1310            target: nid2,
1311            edge_type: "LINK".into(),
1312            properties_json: "{}".into(),
1313            weight: 1.0,
1314            valid_from: None,
1315            valid_to: None,
1316        }))
1317        .await
1318        .unwrap();
1319
1320        // Query neighbors.
1321        let resp = svc
1322            .neighbors(tonic::Request::new(NeighborsRequest {
1323                id: nid1,
1324                direction: "outgoing".into(),
1325                edge_type: String::new(),
1326            }))
1327            .await
1328            .unwrap()
1329            .into_inner();
1330        assert!(resp.error.is_empty());
1331        assert_eq!(resp.neighbors.len(), 1);
1332        assert_eq!(resp.neighbors[0].node_id, nid2);
1333    }
1334
1335    #[tokio::test]
1336    async fn grpc_query() {
1337        let svc = AstraeaGrpcService::new(test_handler());
1338
1339        // Create a node via GQL.
1340        let resp = svc
1341            .query(tonic::Request::new(QueryRequest {
1342                gql: "CREATE (n:Test {name: 'GrpcTest'}) RETURN n".into(),
1343            }))
1344            .await
1345            .unwrap()
1346            .into_inner();
1347        assert!(resp.success, "query failed: {}", resp.error);
1348    }
1349
1350    #[tokio::test]
1351    async fn grpc_get_nonexistent_node() {
1352        let svc = AstraeaGrpcService::new(test_handler());
1353
1354        let resp = svc
1355            .get_node(tonic::Request::new(GetNodeRequest { id: 99999 }))
1356            .await
1357            .unwrap()
1358            .into_inner();
1359        assert!(!resp.found);
1360        assert!(!resp.error.is_empty());
1361    }
1362}