Skip to main content

mnemo_grpc/
lib.rs

1//! gRPC API server for Mnemo.
2//!
3//! This crate exposes Mnemo's core memory operations (remember, recall, forget)
4//! over gRPC using [`tonic`]. The protobuf service definition lives in
5//! `proto/mnemo.proto` and code is generated at build time via `tonic-build`.
6//!
7//! # Usage
8//!
9//! ```rust,ignore
10//! use std::sync::Arc;
11//! use mnemo_grpc::router;
12//!
13//! let engine: Arc<mnemo_core::query::MnemoEngine> = /* ... */;
14//! let grpc_router = router(engine);
15//! ```
16
17use std::sync::Arc;
18
19use tonic::{Request, Response, Status};
20use uuid::Uuid;
21
22use mnemo_core::model::acl::Permission;
23use mnemo_core::model::delegation::{Delegation, DelegationScope};
24use mnemo_core::model::memory::{MemoryType, Scope, SourceType};
25use mnemo_core::query::MnemoEngine;
26use mnemo_core::query::branch::BranchRequest as CoreBranchRequest;
27use mnemo_core::query::checkpoint::CheckpointRequest as CoreCheckpointRequest;
28use mnemo_core::query::forget::{
29    ForgetRequest as CoreForgetRequest, ForgetStrategy,
30    ForgetSubjectRequest as CoreForgetSubjectRequest,
31};
32use mnemo_core::query::merge::{MergeRequest as CoreMergeRequest, MergeStrategy};
33use mnemo_core::query::recall::RecallRequest as CoreRecallRequest;
34use mnemo_core::query::remember::RememberRequest as CoreRememberRequest;
35use mnemo_core::query::replay::ReplayRequest as CoreReplayRequest;
36use mnemo_core::query::share::ShareRequest as CoreShareRequest;
37
38// ---------------------------------------------------------------------------
39// Generated protobuf code
40// ---------------------------------------------------------------------------
41
42pub mod proto {
43    tonic::include_proto!("mnemo.v1");
44}
45
46use proto::mnemo_service_server::{MnemoService, MnemoServiceServer};
47use proto::{
48    BranchRequest as ProtoBranchRequest, BranchResponse as ProtoBranchResponse,
49    CheckpointRequest as ProtoCheckpointRequest, CheckpointResponse as ProtoCheckpointResponse,
50    DelegateRequest as ProtoDelegateRequest, DelegateResponse as ProtoDelegateResponse,
51    ForgetError as ProtoForgetError, ForgetRequest as ProtoForgetRequest,
52    ForgetResponse as ProtoForgetResponse, ForgetSubjectRequest as ProtoForgetSubjectRequest,
53    ForgetSubjectResponse as ProtoForgetSubjectResponse, HealthRequest, HealthResponse,
54    MergeRequest as ProtoMergeRequest, MergeResponse as ProtoMergeResponse,
55    RecallRequest as ProtoRecallRequest, RecallResponse as ProtoRecallResponse,
56    RememberRequest as ProtoRememberRequest, RememberResponse as ProtoRememberResponse,
57    ReplayMemory as ProtoReplayMemory, ReplayRequest as ProtoReplayRequest,
58    ReplayResponse as ProtoReplayResponse, ScoredMemory as ProtoScoredMemory,
59    ShareRequest as ProtoShareRequest, ShareResponse as ProtoShareResponse,
60    VerifyRequest as ProtoVerifyRequest, VerifyResponse as ProtoVerifyResponse,
61};
62
63// ---------------------------------------------------------------------------
64// Server implementation
65// ---------------------------------------------------------------------------
66
67/// gRPC server backed by a shared [`MnemoEngine`].
68#[derive(Clone)]
69pub struct MnemoGrpcServer {
70    engine: Arc<MnemoEngine>,
71}
72
73impl MnemoGrpcServer {
74    /// Create a new server wrapping the given engine.
75    pub fn new(engine: Arc<MnemoEngine>) -> Self {
76        Self { engine }
77    }
78}
79
80#[tonic::async_trait]
81impl MnemoService for MnemoGrpcServer {
82    // -- Remember ----------------------------------------------------------
83
84    async fn remember(
85        &self,
86        request: Request<ProtoRememberRequest>,
87    ) -> Result<Response<ProtoRememberResponse>, Status> {
88        let req = request.into_inner();
89
90        let memory_type = match req.memory_type {
91            Some(ref s) => match s.parse::<MemoryType>() {
92                Ok(mt) => Some(mt),
93                Err(_) => {
94                    return Err(Status::invalid_argument(format!(
95                        "invalid memory_type '{}': expected one of: episodic, semantic, procedural, working",
96                        s
97                    )));
98                }
99            },
100            None => None,
101        };
102
103        let scope = match req.scope {
104            Some(ref s) => match s.parse::<Scope>() {
105                Ok(sc) => Some(sc),
106                Err(_) => {
107                    return Err(Status::invalid_argument(format!(
108                        "invalid scope '{}': expected one of: private, shared, public, global",
109                        s
110                    )));
111                }
112            },
113            None => None,
114        };
115
116        let source_type = match req.source_type {
117            Some(ref s) => match s.parse::<SourceType>() {
118                Ok(st) => Some(st),
119                Err(_) => {
120                    return Err(Status::invalid_argument(format!(
121                        "invalid source_type '{}': expected one of: agent, human, system, user_input, tool_output, model_response, retrieval, consolidation, import",
122                        s
123                    )));
124                }
125            },
126            None => None,
127        };
128
129        let metadata: Option<serde_json::Value> = match req.metadata {
130            Some(ref s) => match serde_json::from_str(s) {
131                Ok(v) => Some(v),
132                Err(e) => {
133                    return Err(Status::invalid_argument(format!(
134                        "invalid metadata JSON: {}",
135                        e
136                    )));
137                }
138            },
139            None => None,
140        };
141
142        let tags = if req.tags.is_empty() {
143            None
144        } else {
145            Some(req.tags)
146        };
147
148        let related_to = if req.related_to.is_empty() {
149            None
150        } else {
151            Some(req.related_to)
152        };
153
154        let core_req = CoreRememberRequest {
155            content: req.content,
156            agent_id: req.agent_id,
157            memory_type,
158            scope,
159            importance: req.importance,
160            tags,
161            metadata,
162            source_type,
163            source_id: req.source_id,
164            org_id: req.org_id,
165            thread_id: req.thread_id,
166            ttl_seconds: req.ttl_seconds,
167            related_to,
168            decay_rate: req.decay_rate,
169            created_by: req.created_by,
170        };
171
172        let result = self
173            .engine
174            .remember(core_req)
175            .await
176            .map_err(core_error_to_status)?;
177
178        Ok(Response::new(ProtoRememberResponse {
179            id: result.id.to_string(),
180            content_hash: result.content_hash,
181        }))
182    }
183
184    // -- Recall ------------------------------------------------------------
185
186    async fn recall(
187        &self,
188        request: Request<ProtoRecallRequest>,
189    ) -> Result<Response<ProtoRecallResponse>, Status> {
190        let req = request.into_inner();
191
192        let memory_type = match req.memory_type {
193            Some(ref s) => match s.parse::<MemoryType>() {
194                Ok(mt) => Some(mt),
195                Err(_) => {
196                    return Err(Status::invalid_argument(format!(
197                        "invalid memory_type '{}': expected one of: episodic, semantic, procedural, working",
198                        s
199                    )));
200                }
201            },
202            None => None,
203        };
204
205        let scope = match req.scope {
206            Some(ref s) => match s.parse::<Scope>() {
207                Ok(sc) => Some(sc),
208                Err(_) => {
209                    return Err(Status::invalid_argument(format!(
210                        "invalid scope '{}': expected one of: private, shared, public, global",
211                        s
212                    )));
213                }
214            },
215            None => None,
216        };
217
218        let tags = if req.tags.is_empty() {
219            None
220        } else {
221            Some(req.tags)
222        };
223
224        let hybrid_weights = if req.hybrid_weights.is_empty() {
225            None
226        } else {
227            Some(req.hybrid_weights)
228        };
229
230        let core_req = CoreRecallRequest {
231            query: req.query,
232            agent_id: req.agent_id,
233            limit: req.limit.map(|l| l as usize),
234            memory_type,
235            memory_types: None,
236            scope,
237            min_importance: req.min_importance,
238            tags,
239            org_id: req.org_id,
240            strategy: req.strategy,
241            temporal_range: None,
242            recency_half_life_hours: None,
243            hybrid_weights,
244            rrf_k: req.rrf_k,
245            as_of: req.as_of,
246            explain: req.explain,
247            with_provenance: None,
248            mode: None,
249        };
250
251        let result = self
252            .engine
253            .recall(core_req)
254            .await
255            .map_err(core_error_to_status)?;
256
257        let memories: Vec<ProtoScoredMemory> = result
258            .memories
259            .into_iter()
260            .map(|m| ProtoScoredMemory {
261                id: m.id.to_string(),
262                content: m.content,
263                memory_type: format!("{:?}", m.memory_type),
264                importance: m.importance,
265                score: m.score,
266                created_at: m.created_at,
267                agent_id: m.agent_id,
268                scope: format!("{:?}", m.scope),
269                tags: m.tags,
270                metadata: m.metadata.to_string(),
271                access_count: m.access_count,
272                updated_at: m.updated_at,
273                score_breakdown: m.score_breakdown.map(|b| proto::ScoreBreakdown {
274                    vector: b.vector,
275                    bm25: b.bm25,
276                    graph: b.graph,
277                    recency: b.recency,
278                    rrf_rank: b.rrf_rank,
279                }),
280            })
281            .collect();
282
283        let total = result.total as u32;
284
285        Ok(Response::new(ProtoRecallResponse { memories, total }))
286    }
287
288    // -- Forget ------------------------------------------------------------
289
290    async fn forget(
291        &self,
292        request: Request<ProtoForgetRequest>,
293    ) -> Result<Response<ProtoForgetResponse>, Status> {
294        let req = request.into_inner();
295
296        let memory_ids: Vec<Uuid> = req
297            .memory_ids
298            .iter()
299            .map(|s| {
300                Uuid::parse_str(s)
301                    .map_err(|e| Status::invalid_argument(format!("invalid UUID '{s}': {e}")))
302            })
303            .collect::<Result<Vec<_>, _>>()?;
304
305        let strategy = match req.strategy {
306            Some(ref s) => {
307                let st = match s.as_str() {
308                    "soft_delete" => ForgetStrategy::SoftDelete,
309                    "hard_delete" => ForgetStrategy::HardDelete,
310                    "decay" => ForgetStrategy::Decay,
311                    "consolidate" => ForgetStrategy::Consolidate,
312                    "archive" => ForgetStrategy::Archive,
313                    "redact" => ForgetStrategy::Redact,
314                    _ => {
315                        return Err(Status::invalid_argument(format!(
316                            "invalid forget strategy '{}': expected one of: soft_delete, hard_delete, decay, consolidate, archive, redact",
317                            s
318                        )));
319                    }
320                };
321                Some(st)
322            }
323            None => None,
324        };
325
326        let core_req = CoreForgetRequest {
327            memory_ids,
328            agent_id: req.agent_id,
329            strategy,
330            criteria: None,
331        };
332
333        let result = self
334            .engine
335            .forget(core_req)
336            .await
337            .map_err(core_error_to_status)?;
338
339        let forgotten: Vec<String> = result.forgotten.iter().map(|id| id.to_string()).collect();
340
341        let errors: Vec<ProtoForgetError> = result
342            .errors
343            .into_iter()
344            .map(|e| ProtoForgetError {
345                id: e.id.to_string(),
346                error: e.error,
347            })
348            .collect();
349
350        Ok(Response::new(ProtoForgetResponse { forgotten, errors }))
351    }
352
353    // -- Health ------------------------------------------------------------
354
355    async fn health(
356        &self,
357        _request: Request<HealthRequest>,
358    ) -> Result<Response<HealthResponse>, Status> {
359        Ok(Response::new(HealthResponse {
360            status: "ok".to_string(),
361            version: env!("CARGO_PKG_VERSION").to_string(),
362        }))
363    }
364
365    // -- Share -------------------------------------------------------------
366
367    async fn share(
368        &self,
369        request: Request<ProtoShareRequest>,
370    ) -> Result<Response<ProtoShareResponse>, Status> {
371        let req = request.into_inner();
372        let memory_id = Uuid::parse_str(&req.memory_id)
373            .map_err(|e| Status::invalid_argument(format!("invalid UUID: {e}")))?;
374        let permission = match req.permission {
375            Some(ref s) => match s.parse::<Permission>() {
376                Ok(p) => Some(p),
377                Err(_) => {
378                    return Err(Status::invalid_argument(format!(
379                        "invalid permission '{}': expected one of: read, write, delete, share, delegate, admin",
380                        s
381                    )));
382                }
383            },
384            None => None,
385        };
386        let target_agent_ids = if req.target_agent_ids.is_empty() {
387            None
388        } else {
389            Some(req.target_agent_ids)
390        };
391
392        let core_req = CoreShareRequest {
393            memory_id,
394            agent_id: req.agent_id,
395            target_agent_id: req.target_agent_id,
396            target_agent_ids,
397            permission,
398            expires_in_hours: req.expires_in_hours,
399        };
400        let result = self
401            .engine
402            .share(core_req)
403            .await
404            .map_err(core_error_to_status)?;
405
406        Ok(Response::new(ProtoShareResponse {
407            acl_id: result.acl_id.to_string(),
408            acl_ids: result.acl_ids.iter().map(|id| id.to_string()).collect(),
409            memory_id: result.memory_id.to_string(),
410            shared_with: result.shared_with,
411            shared_with_all: result.shared_with_all,
412            permission: result.permission.to_string(),
413        }))
414    }
415
416    // -- Checkpoint --------------------------------------------------------
417
418    async fn checkpoint(
419        &self,
420        request: Request<ProtoCheckpointRequest>,
421    ) -> Result<Response<ProtoCheckpointResponse>, Status> {
422        let req = request.into_inner();
423        let state_snapshot: serde_json::Value = serde_json::from_str(&req.state_snapshot)
424            .map_err(|e| Status::invalid_argument(format!("invalid JSON state_snapshot: {e}")))?;
425        let metadata: Option<serde_json::Value> = match req.metadata {
426            Some(ref s) => match serde_json::from_str(s) {
427                Ok(v) => Some(v),
428                Err(e) => {
429                    return Err(Status::invalid_argument(format!(
430                        "invalid metadata JSON: {}",
431                        e
432                    )));
433                }
434            },
435            None => None,
436        };
437
438        let core_req = CoreCheckpointRequest {
439            thread_id: req.thread_id,
440            agent_id: req.agent_id,
441            branch_name: req.branch_name,
442            state_snapshot,
443            label: req.label,
444            metadata,
445        };
446        let result = self
447            .engine
448            .checkpoint(core_req)
449            .await
450            .map_err(core_error_to_status)?;
451
452        Ok(Response::new(ProtoCheckpointResponse {
453            checkpoint_id: result.id.to_string(),
454            parent_id: result.parent_id.map(|id| id.to_string()),
455            branch_name: result.branch_name,
456        }))
457    }
458
459    // -- Branch ------------------------------------------------------------
460
461    async fn branch(
462        &self,
463        request: Request<ProtoBranchRequest>,
464    ) -> Result<Response<ProtoBranchResponse>, Status> {
465        let req = request.into_inner();
466        let source_checkpoint_id = match req.source_checkpoint_id {
467            Some(ref s) => match Uuid::parse_str(s) {
468                Ok(id) => Some(id),
469                Err(e) => {
470                    return Err(Status::invalid_argument(format!(
471                        "invalid source_checkpoint_id '{}': {}",
472                        s, e
473                    )));
474                }
475            },
476            None => None,
477        };
478
479        let core_req = CoreBranchRequest {
480            thread_id: req.thread_id,
481            agent_id: req.agent_id,
482            new_branch_name: req.new_branch_name,
483            source_checkpoint_id,
484            source_branch: req.source_branch,
485        };
486        let result = self
487            .engine
488            .branch(core_req)
489            .await
490            .map_err(core_error_to_status)?;
491
492        Ok(Response::new(ProtoBranchResponse {
493            checkpoint_id: result.checkpoint_id.to_string(),
494            branch_name: result.branch_name,
495            source_checkpoint_id: result.source_checkpoint_id.to_string(),
496        }))
497    }
498
499    // -- Merge -------------------------------------------------------------
500
501    async fn merge(
502        &self,
503        request: Request<ProtoMergeRequest>,
504    ) -> Result<Response<ProtoMergeResponse>, Status> {
505        let req = request.into_inner();
506        let strategy = match req.strategy {
507            Some(ref s) => {
508                let st = match s.as_str() {
509                    "full_merge" => MergeStrategy::FullMerge,
510                    "cherry_pick" => MergeStrategy::CherryPick,
511                    "squash" => MergeStrategy::Squash,
512                    _ => {
513                        return Err(Status::invalid_argument(format!(
514                            "invalid merge strategy '{}': expected one of: full_merge, cherry_pick, squash",
515                            s
516                        )));
517                    }
518                };
519                Some(st)
520            }
521            None => None,
522        };
523        let cherry_pick_ids = if req.cherry_pick_ids.is_empty() {
524            None
525        } else {
526            let ids: Result<Vec<Uuid>, _> = req
527                .cherry_pick_ids
528                .iter()
529                .map(|s| Uuid::parse_str(s))
530                .collect();
531            Some(ids.map_err(|e| Status::invalid_argument(format!("invalid UUID: {e}")))?)
532        };
533
534        let core_req = CoreMergeRequest {
535            thread_id: req.thread_id,
536            agent_id: req.agent_id,
537            source_branch: req.source_branch,
538            target_branch: req.target_branch,
539            strategy,
540            cherry_pick_ids,
541        };
542        let result = self
543            .engine
544            .merge(core_req)
545            .await
546            .map_err(core_error_to_status)?;
547
548        Ok(Response::new(ProtoMergeResponse {
549            checkpoint_id: result.checkpoint_id.to_string(),
550            target_branch: result.target_branch,
551            merged_memory_count: result.merged_memory_count as u32,
552        }))
553    }
554
555    // -- Replay ------------------------------------------------------------
556
557    async fn replay(
558        &self,
559        request: Request<ProtoReplayRequest>,
560    ) -> Result<Response<ProtoReplayResponse>, Status> {
561        let req = request.into_inner();
562        let checkpoint_id = match req.checkpoint_id {
563            Some(ref s) => match Uuid::parse_str(s) {
564                Ok(id) => Some(id),
565                Err(e) => {
566                    return Err(Status::invalid_argument(format!(
567                        "invalid checkpoint_id '{}': {}",
568                        s, e
569                    )));
570                }
571            },
572            None => None,
573        };
574
575        let core_req = CoreReplayRequest {
576            thread_id: req.thread_id,
577            agent_id: req.agent_id,
578            checkpoint_id,
579            branch_name: req.branch_name,
580            as_of: req.as_of,
581        };
582        let result = self
583            .engine
584            .replay(core_req)
585            .await
586            .map_err(core_error_to_status)?;
587
588        let checkpoint_json =
589            serde_json::to_string(&result.checkpoint).unwrap_or_else(|_| "{}".to_string());
590        let memories: Vec<ProtoReplayMemory> = result
591            .memories
592            .iter()
593            .map(|m| ProtoReplayMemory {
594                id: m.id.to_string(),
595                content: m.content.clone(),
596                memory_type: format!("{:?}", m.memory_type),
597                created_at: m.created_at.clone(),
598            })
599            .collect();
600
601        let (chain_valid, chain_total, chain_verified) =
602            if let Some(ref cv) = result.chain_verification {
603                (
604                    Some(cv.valid),
605                    Some(cv.total_records as u32),
606                    Some(cv.verified_records as u32),
607                )
608            } else {
609                (None, None, None)
610            };
611
612        Ok(Response::new(ProtoReplayResponse {
613            checkpoint_json,
614            memories,
615            event_count: result.events.len() as u32,
616            chain_valid,
617            chain_total,
618            chain_verified,
619        }))
620    }
621
622    // -- Delegate ----------------------------------------------------------
623
624    async fn delegate(
625        &self,
626        request: Request<ProtoDelegateRequest>,
627    ) -> Result<Response<ProtoDelegateResponse>, Status> {
628        let req = request.into_inner();
629        let permission: Permission = req
630            .permission
631            .parse()
632            .map_err(|e: mnemo_core::error::Error| Status::invalid_argument(e.to_string()))?;
633
634        let scope = if !req.memory_ids.is_empty() {
635            let ids: Vec<Uuid> = req
636                .memory_ids
637                .iter()
638                .map(|s| {
639                    Uuid::parse_str(s)
640                        .map_err(|e| Status::invalid_argument(format!("invalid UUID: {e}")))
641                })
642                .collect::<Result<Vec<_>, _>>()?;
643            DelegationScope::ByMemoryId(ids)
644        } else if !req.tags.is_empty() {
645            DelegationScope::ByTag(req.tags)
646        } else {
647            DelegationScope::AllMemories
648        };
649
650        let now = chrono::Utc::now();
651        let expires_at = req
652            .expires_in_hours
653            .map(|h| (now + chrono::Duration::seconds((h * 3600.0) as i64)).to_rfc3339());
654
655        let delegation = Delegation {
656            id: Uuid::now_v7(),
657            delegator_id: req.delegator_id,
658            delegate_id: req.delegate_id,
659            permission,
660            scope,
661            max_depth: req.max_depth.unwrap_or(0),
662            current_depth: 0,
663            parent_delegation_id: None,
664            created_at: now.to_rfc3339(),
665            expires_at,
666            revoked_at: None,
667        };
668
669        self.engine
670            .storage
671            .insert_delegation(&delegation)
672            .await
673            .map_err(core_error_to_status)?;
674
675        Ok(Response::new(ProtoDelegateResponse {
676            delegation_id: delegation.id.to_string(),
677        }))
678    }
679
680    // -- Verify ------------------------------------------------------------
681
682    async fn verify(
683        &self,
684        request: Request<ProtoVerifyRequest>,
685    ) -> Result<Response<ProtoVerifyResponse>, Status> {
686        let req = request.into_inner();
687        let result = self
688            .engine
689            .verify_integrity(req.agent_id, req.thread_id.as_deref())
690            .await
691            .map_err(core_error_to_status)?;
692
693        Ok(Response::new(ProtoVerifyResponse {
694            valid: result.valid,
695            total_records: result.total_records as u32,
696            verified_records: result.verified_records as u32,
697            first_broken_at: result.first_broken_at.map(|id| id.to_string()),
698            error_message: result.error_message,
699        }))
700    }
701
702    // -- ForgetSubject -----------------------------------------------------
703
704    async fn forget_subject(
705        &self,
706        request: Request<ProtoForgetSubjectRequest>,
707    ) -> Result<Response<ProtoForgetSubjectResponse>, Status> {
708        let req = request.into_inner();
709
710        let strategy = match req.strategy.as_deref().unwrap_or("redact") {
711            "redact" => ForgetStrategy::Redact,
712            "hard_delete" => ForgetStrategy::HardDelete,
713            "soft_delete" => ForgetStrategy::SoftDelete,
714            other => {
715                return Err(Status::invalid_argument(format!(
716                    "invalid forget_subject strategy '{}': expected one of: redact, hard_delete, soft_delete",
717                    other
718                )));
719            }
720        };
721
722        let core_req = CoreForgetSubjectRequest {
723            subject_id: req.subject_id,
724            agent_id: req.agent_id,
725            strategy,
726        };
727
728        let result = self
729            .engine
730            .forget_subject(core_req)
731            .await
732            .map_err(core_error_to_status)?;
733
734        let errors: Vec<ProtoForgetError> = result
735            .errors
736            .into_iter()
737            .map(|e| ProtoForgetError {
738                id: e.id.to_string(),
739                error: e.error,
740            })
741            .collect();
742
743        let strategy_str = match result.strategy {
744            ForgetStrategy::SoftDelete => "soft_delete",
745            ForgetStrategy::HardDelete => "hard_delete",
746            ForgetStrategy::Decay => "decay",
747            ForgetStrategy::Consolidate => "consolidate",
748            ForgetStrategy::Archive => "archive",
749            ForgetStrategy::Redact => "redact",
750        }
751        .to_string();
752
753        Ok(Response::new(ProtoForgetSubjectResponse {
754            subject_id: result.subject_id,
755            strategy: strategy_str,
756            matched: result.matched as u32,
757            forgotten: result.forgotten.iter().map(|id| id.to_string()).collect(),
758            cascaded_events: result.cascaded_events as u32,
759            errors,
760        }))
761    }
762}
763
764// ---------------------------------------------------------------------------
765// Router constructor
766// ---------------------------------------------------------------------------
767
768/// Build a [`tonic::transport::server::Router`] serving the Mnemo gRPC API.
769///
770/// The returned router can be composed with other tonic services or served
771/// directly via `tonic::transport::Server`.
772///
773/// # Example
774///
775/// ```rust,ignore
776/// use std::sync::Arc;
777/// use mnemo_grpc::router;
778///
779/// let engine: Arc<mnemo_core::query::MnemoEngine> = /* ... */;
780/// let grpc_router = router(engine);
781/// tonic::transport::Server::builder()
782///     .add_routes(grpc_router.into_service())
783///     .serve("[::1]:50051".parse().unwrap())
784///     .await
785///     .unwrap();
786/// ```
787pub fn router(engine: Arc<MnemoEngine>) -> tonic::transport::server::Router {
788    let svc = MnemoGrpcServer::new(engine);
789    tonic::transport::Server::builder().add_service(MnemoServiceServer::new(svc))
790}
791
792// ---------------------------------------------------------------------------
793// Helpers
794// ---------------------------------------------------------------------------
795
796/// Map a `mnemo_core::error::Error` to a tonic `Status`.
797fn core_error_to_status(err: mnemo_core::error::Error) -> Status {
798    use mnemo_core::error::Error;
799
800    match err {
801        Error::Validation(msg) => Status::invalid_argument(msg),
802        Error::PermissionDenied(msg) => Status::permission_denied(msg),
803        Error::NotFound(msg) => Status::not_found(msg),
804        other => Status::internal(other.to_string()),
805    }
806}
807
808// ---------------------------------------------------------------------------
809// Tests
810// ---------------------------------------------------------------------------
811
812#[cfg(test)]
813mod tests {
814    use super::*;
815
816    #[test]
817    fn core_error_maps_correctly() {
818        let validation =
819            core_error_to_status(mnemo_core::error::Error::Validation("bad input".into()));
820        assert_eq!(validation.code(), tonic::Code::InvalidArgument);
821
822        let perm = core_error_to_status(mnemo_core::error::Error::PermissionDenied(
823            "forbidden".into(),
824        ));
825        assert_eq!(perm.code(), tonic::Code::PermissionDenied);
826
827        let not_found = core_error_to_status(mnemo_core::error::Error::NotFound("missing".into()));
828        assert_eq!(not_found.code(), tonic::Code::NotFound);
829    }
830}