use std::sync::Arc;
use tonic::{Request, Response, Status};
use uuid::Uuid;
use mnemo_core::model::acl::Permission;
use mnemo_core::model::delegation::{Delegation, DelegationScope};
use mnemo_core::model::memory::{MemoryType, Scope, SourceType};
use mnemo_core::query::MnemoEngine;
use mnemo_core::query::branch::BranchRequest as CoreBranchRequest;
use mnemo_core::query::checkpoint::CheckpointRequest as CoreCheckpointRequest;
use mnemo_core::query::consolidate::ConsolidateRequest as CoreConsolidateRequest;
use mnemo_core::query::forget::{
ForgetRequest as CoreForgetRequest, ForgetStrategy,
ForgetSubjectRequest as CoreForgetSubjectRequest,
};
use mnemo_core::query::merge::{MergeRequest as CoreMergeRequest, MergeStrategy};
use mnemo_core::query::recall::RecallRequest as CoreRecallRequest;
use mnemo_core::query::remember::RememberRequest as CoreRememberRequest;
use mnemo_core::query::replay::ReplayRequest as CoreReplayRequest;
use mnemo_core::query::share::ShareRequest as CoreShareRequest;
pub mod proto {
tonic::include_proto!("mnemo.v1");
}
use proto::mnemo_service_server::{MnemoService, MnemoServiceServer};
use proto::{
BranchRequest as ProtoBranchRequest, BranchResponse as ProtoBranchResponse,
CheckpointRequest as ProtoCheckpointRequest, CheckpointResponse as ProtoCheckpointResponse,
ConsolidateRequest as ProtoConsolidateRequest, ConsolidateResponse as ProtoConsolidateResponse,
DelegateRequest as ProtoDelegateRequest, DelegateResponse as ProtoDelegateResponse,
ForgetError as ProtoForgetError, ForgetRequest as ProtoForgetRequest,
ForgetResponse as ProtoForgetResponse, ForgetSubjectRequest as ProtoForgetSubjectRequest,
ForgetSubjectResponse as ProtoForgetSubjectResponse, HealthRequest, HealthResponse,
MergeRequest as ProtoMergeRequest, MergeResponse as ProtoMergeResponse,
RecallRequest as ProtoRecallRequest, RecallResponse as ProtoRecallResponse,
RememberRequest as ProtoRememberRequest, RememberResponse as ProtoRememberResponse,
ReplayMemory as ProtoReplayMemory, ReplayRequest as ProtoReplayRequest,
ReplayResponse as ProtoReplayResponse, ScoredMemory as ProtoScoredMemory,
ShareRequest as ProtoShareRequest, ShareResponse as ProtoShareResponse,
TrajectoryAuditRequest as ProtoTrajectoryAuditRequest,
TrajectoryAuditResponse as ProtoTrajectoryAuditResponse,
TrajectoryFinding as ProtoTrajectoryFinding, VerifyRequest as ProtoVerifyRequest,
VerifyResponse as ProtoVerifyResponse,
};
#[derive(Clone)]
pub struct MnemoGrpcServer {
engine: Arc<MnemoEngine>,
}
impl MnemoGrpcServer {
pub fn new(engine: Arc<MnemoEngine>) -> Self {
Self { engine }
}
}
#[tonic::async_trait]
impl MnemoService for MnemoGrpcServer {
async fn remember(
&self,
request: Request<ProtoRememberRequest>,
) -> Result<Response<ProtoRememberResponse>, Status> {
let req = request.into_inner();
let memory_type = match req.memory_type {
Some(ref s) => match s.parse::<MemoryType>() {
Ok(mt) => Some(mt),
Err(_) => {
return Err(Status::invalid_argument(format!(
"invalid memory_type '{}': expected one of: episodic, semantic, procedural, working",
s
)));
}
},
None => None,
};
let scope = match req.scope {
Some(ref s) => match s.parse::<Scope>() {
Ok(sc) => Some(sc),
Err(_) => {
return Err(Status::invalid_argument(format!(
"invalid scope '{}': expected one of: private, shared, public, global",
s
)));
}
},
None => None,
};
let source_type = match req.source_type {
Some(ref s) => match s.parse::<SourceType>() {
Ok(st) => Some(st),
Err(_) => {
return Err(Status::invalid_argument(format!(
"invalid source_type '{}': expected one of: agent, human, system, user_input, tool_output, model_response, retrieval, consolidation, import",
s
)));
}
},
None => None,
};
let metadata: Option<serde_json::Value> = match req.metadata {
Some(ref s) => match serde_json::from_str(s) {
Ok(v) => Some(v),
Err(e) => {
return Err(Status::invalid_argument(format!(
"invalid metadata JSON: {}",
e
)));
}
},
None => None,
};
let tags = if req.tags.is_empty() {
None
} else {
Some(req.tags)
};
let related_to = if req.related_to.is_empty() {
None
} else {
Some(req.related_to)
};
let core_req = CoreRememberRequest {
content: req.content,
agent_id: req.agent_id,
memory_type,
scope,
importance: req.importance,
tags,
metadata,
source_type,
source_id: req.source_id,
org_id: req.org_id,
thread_id: req.thread_id,
ttl_seconds: req.ttl_seconds,
related_to,
decay_rate: req.decay_rate,
created_by: req.created_by,
};
let result = self
.engine
.remember(core_req)
.await
.map_err(core_error_to_status)?;
Ok(Response::new(ProtoRememberResponse {
id: result.id.to_string(),
content_hash: result.content_hash,
}))
}
async fn recall(
&self,
request: Request<ProtoRecallRequest>,
) -> Result<Response<ProtoRecallResponse>, Status> {
let req = request.into_inner();
let memory_type = match req.memory_type {
Some(ref s) => match s.parse::<MemoryType>() {
Ok(mt) => Some(mt),
Err(_) => {
return Err(Status::invalid_argument(format!(
"invalid memory_type '{}': expected one of: episodic, semantic, procedural, working",
s
)));
}
},
None => None,
};
let scope = match req.scope {
Some(ref s) => match s.parse::<Scope>() {
Ok(sc) => Some(sc),
Err(_) => {
return Err(Status::invalid_argument(format!(
"invalid scope '{}': expected one of: private, shared, public, global",
s
)));
}
},
None => None,
};
let tags = if req.tags.is_empty() {
None
} else {
Some(req.tags)
};
let hybrid_weights = if req.hybrid_weights.is_empty() {
None
} else {
Some(req.hybrid_weights)
};
let orientation_cache_cfg = req.orientation_cache.map(|o| {
mnemo_core::query::orientation_cache::OrientationCacheConfig {
namespace: o.namespace,
token_budget: o.token_budget,
include_in_response: o.include_in_response.unwrap_or(true),
distill: o.distill.unwrap_or(true),
}
});
let core_req = CoreRecallRequest {
query: req.query,
agent_id: req.agent_id,
limit: req.limit.map(|l| l as usize),
memory_type,
memory_types: None,
scope,
min_importance: req.min_importance,
tags,
org_id: req.org_id,
strategy: req.strategy,
temporal_range: None,
recency_half_life_hours: None,
hybrid_weights,
rrf_k: req.rrf_k,
as_of: req.as_of,
explain: req.explain,
with_provenance: None,
mode: None,
current_fact_resolver: None,
orientation_cache: orientation_cache_cfg,
evidence_budget: None,
retained_token_budget: None,
domain_scope: None,
reasoning_trust: None,
};
let result = self
.engine
.recall(core_req)
.await
.map_err(core_error_to_status)?;
let memories: Vec<ProtoScoredMemory> = result
.memories
.into_iter()
.map(|m| ProtoScoredMemory {
id: m.id.to_string(),
content: m.content,
memory_type: format!("{:?}", m.memory_type),
importance: m.importance,
score: m.score,
created_at: m.created_at,
agent_id: m.agent_id,
scope: format!("{:?}", m.scope),
tags: m.tags,
metadata: m.metadata.to_string(),
access_count: m.access_count,
updated_at: m.updated_at,
score_breakdown: m.score_breakdown.map(|b| proto::ScoreBreakdown {
vector: b.vector,
bm25: b.bm25,
graph: b.graph,
recency: b.recency,
rrf_rank: b.rrf_rank,
}),
})
.collect();
let total = result.total as u32;
let orientation_cache = result
.orientation_cache
.map(|r| proto::OrientationCacheResponse {
namespace: r.namespace,
entities: r
.entities
.into_iter()
.map(|e| proto::OrientationEntry {
key: e.key,
value: e.value,
freq: e.freq,
token_estimate: e.token_estimate,
})
.collect(),
constants: r
.constants
.into_iter()
.map(|e| proto::OrientationEntry {
key: e.key,
value: e.value,
freq: e.freq,
token_estimate: e.token_estimate,
})
.collect(),
schemas: r
.schemas
.into_iter()
.map(|e| proto::OrientationEntry {
key: e.key,
value: e.value,
freq: e.freq,
token_estimate: e.token_estimate,
})
.collect(),
token_estimate: r.token_estimate,
budget: r.budget,
hit_count: r.hit_count,
});
let reconstruction = result.reconstruction.map(|b| proto::Reconstruction {
cue: b.cue,
summary: b.summary,
source_ids: b.source_ids.iter().map(|id| id.to_string()).collect(),
linked_context_ids: b
.linked_context_ids
.iter()
.map(|id| id.to_string())
.collect(),
confidence: b.confidence,
});
Ok(Response::new(ProtoRecallResponse {
memories,
total,
orientation_cache,
reconstruction,
}))
}
async fn forget(
&self,
request: Request<ProtoForgetRequest>,
) -> Result<Response<ProtoForgetResponse>, Status> {
let req = request.into_inner();
let memory_ids: Vec<Uuid> = req
.memory_ids
.iter()
.map(|s| {
Uuid::parse_str(s)
.map_err(|e| Status::invalid_argument(format!("invalid UUID '{s}': {e}")))
})
.collect::<Result<Vec<_>, _>>()?;
let strategy = match req.strategy {
Some(ref s) => {
let st = match s.as_str() {
"soft_delete" => ForgetStrategy::SoftDelete,
"hard_delete" => ForgetStrategy::HardDelete,
"decay" => ForgetStrategy::Decay,
"consolidate" => ForgetStrategy::Consolidate,
"archive" => ForgetStrategy::Archive,
"redact" => ForgetStrategy::Redact,
_ => {
return Err(Status::invalid_argument(format!(
"invalid forget strategy '{}': expected one of: soft_delete, hard_delete, decay, consolidate, archive, redact",
s
)));
}
};
Some(st)
}
None => None,
};
let core_req = CoreForgetRequest {
memory_ids,
agent_id: req.agent_id,
strategy,
criteria: None,
};
let result = self
.engine
.forget(core_req)
.await
.map_err(core_error_to_status)?;
let forgotten: Vec<String> = result.forgotten.iter().map(|id| id.to_string()).collect();
let errors: Vec<ProtoForgetError> = result
.errors
.into_iter()
.map(|e| ProtoForgetError {
id: e.id.to_string(),
error: e.error,
})
.collect();
Ok(Response::new(ProtoForgetResponse { forgotten, errors }))
}
async fn health(
&self,
_request: Request<HealthRequest>,
) -> Result<Response<HealthResponse>, Status> {
Ok(Response::new(HealthResponse {
status: "ok".to_string(),
version: env!("CARGO_PKG_VERSION").to_string(),
}))
}
async fn share(
&self,
request: Request<ProtoShareRequest>,
) -> Result<Response<ProtoShareResponse>, Status> {
let req = request.into_inner();
let memory_id = Uuid::parse_str(&req.memory_id)
.map_err(|e| Status::invalid_argument(format!("invalid UUID: {e}")))?;
let permission = match req.permission {
Some(ref s) => match s.parse::<Permission>() {
Ok(p) => Some(p),
Err(_) => {
return Err(Status::invalid_argument(format!(
"invalid permission '{}': expected one of: read, write, delete, share, delegate, admin",
s
)));
}
},
None => None,
};
let target_agent_ids = if req.target_agent_ids.is_empty() {
None
} else {
Some(req.target_agent_ids)
};
let core_req = CoreShareRequest {
memory_id,
agent_id: req.agent_id,
target_agent_id: req.target_agent_id,
target_agent_ids,
permission,
expires_in_hours: req.expires_in_hours,
};
let result = self
.engine
.share(core_req)
.await
.map_err(core_error_to_status)?;
Ok(Response::new(ProtoShareResponse {
acl_id: result.acl_id.to_string(),
acl_ids: result.acl_ids.iter().map(|id| id.to_string()).collect(),
memory_id: result.memory_id.to_string(),
shared_with: result.shared_with,
shared_with_all: result.shared_with_all,
permission: result.permission.to_string(),
}))
}
async fn checkpoint(
&self,
request: Request<ProtoCheckpointRequest>,
) -> Result<Response<ProtoCheckpointResponse>, Status> {
let req = request.into_inner();
let state_snapshot: serde_json::Value = serde_json::from_str(&req.state_snapshot)
.map_err(|e| Status::invalid_argument(format!("invalid JSON state_snapshot: {e}")))?;
let metadata: Option<serde_json::Value> = match req.metadata {
Some(ref s) => match serde_json::from_str(s) {
Ok(v) => Some(v),
Err(e) => {
return Err(Status::invalid_argument(format!(
"invalid metadata JSON: {}",
e
)));
}
},
None => None,
};
let core_req = CoreCheckpointRequest {
thread_id: req.thread_id,
agent_id: req.agent_id,
branch_name: req.branch_name,
state_snapshot,
label: req.label,
metadata,
};
let result = self
.engine
.checkpoint(core_req)
.await
.map_err(core_error_to_status)?;
Ok(Response::new(ProtoCheckpointResponse {
checkpoint_id: result.id.to_string(),
parent_id: result.parent_id.map(|id| id.to_string()),
branch_name: result.branch_name,
}))
}
async fn consolidate(
&self,
request: Request<ProtoConsolidateRequest>,
) -> Result<Response<ProtoConsolidateResponse>, Status> {
let req = request.into_inner();
let mut memory_ids = Vec::with_capacity(req.memory_ids.len());
for s in &req.memory_ids {
memory_ids.push(
Uuid::parse_str(s).map_err(|e| {
Status::invalid_argument(format!("invalid memory id '{s}': {e}"))
})?,
);
}
let supersede = match req.supersede {
Some(ref s) => Some(Uuid::parse_str(s).map_err(|e| {
Status::invalid_argument(format!("invalid supersede id '{s}': {e}"))
})?),
None => None,
};
let metadata: Option<serde_json::Value> = match req.metadata {
Some(ref s) => Some(
serde_json::from_str(s)
.map_err(|e| Status::invalid_argument(format!("invalid metadata JSON: {e}")))?,
),
None => None,
};
let mut core_req = CoreConsolidateRequest::new(memory_ids, req.topic_name);
core_req.agent_id = req.agent_id;
core_req.summary = req.summary;
core_req.supersede = supersede;
core_req.thread_id = req.thread_id;
core_req.metadata = metadata;
let result = self
.engine
.consolidate(core_req)
.await
.map_err(core_error_to_status)?;
Ok(Response::new(ProtoConsolidateResponse {
topic_document_id: result.topic_document_id.to_string(),
topic_name: result.topic_name,
source_count: result.source_count as u64,
version: result.version,
superseded_id: result.superseded_id.map(|id| id.to_string()),
member_ids: result.member_ids.iter().map(|id| id.to_string()).collect(),
content_hash: result.content_hash,
consolidation_event_id: result.consolidation_event_id.to_string(),
revision_event_id: result.revision_event_id.map(|id| id.to_string()),
}))
}
async fn branch(
&self,
request: Request<ProtoBranchRequest>,
) -> Result<Response<ProtoBranchResponse>, Status> {
let req = request.into_inner();
let source_checkpoint_id = match req.source_checkpoint_id {
Some(ref s) => match Uuid::parse_str(s) {
Ok(id) => Some(id),
Err(e) => {
return Err(Status::invalid_argument(format!(
"invalid source_checkpoint_id '{}': {}",
s, e
)));
}
},
None => None,
};
let core_req = CoreBranchRequest {
thread_id: req.thread_id,
agent_id: req.agent_id,
new_branch_name: req.new_branch_name,
source_checkpoint_id,
source_branch: req.source_branch,
};
let result = self
.engine
.branch(core_req)
.await
.map_err(core_error_to_status)?;
Ok(Response::new(ProtoBranchResponse {
checkpoint_id: result.checkpoint_id.to_string(),
branch_name: result.branch_name,
source_checkpoint_id: result.source_checkpoint_id.to_string(),
}))
}
async fn merge(
&self,
request: Request<ProtoMergeRequest>,
) -> Result<Response<ProtoMergeResponse>, Status> {
let req = request.into_inner();
let strategy = match req.strategy {
Some(ref s) => {
let st = match s.as_str() {
"full_merge" => MergeStrategy::FullMerge,
"cherry_pick" => MergeStrategy::CherryPick,
"squash" => MergeStrategy::Squash,
_ => {
return Err(Status::invalid_argument(format!(
"invalid merge strategy '{}': expected one of: full_merge, cherry_pick, squash",
s
)));
}
};
Some(st)
}
None => None,
};
let cherry_pick_ids = if req.cherry_pick_ids.is_empty() {
None
} else {
let ids: Result<Vec<Uuid>, _> = req
.cherry_pick_ids
.iter()
.map(|s| Uuid::parse_str(s))
.collect();
Some(ids.map_err(|e| Status::invalid_argument(format!("invalid UUID: {e}")))?)
};
let core_req = CoreMergeRequest {
thread_id: req.thread_id,
agent_id: req.agent_id,
source_branch: req.source_branch,
target_branch: req.target_branch,
strategy,
cherry_pick_ids,
};
let result = self
.engine
.merge(core_req)
.await
.map_err(core_error_to_status)?;
Ok(Response::new(ProtoMergeResponse {
checkpoint_id: result.checkpoint_id.to_string(),
target_branch: result.target_branch,
merged_memory_count: result.merged_memory_count as u32,
}))
}
async fn replay(
&self,
request: Request<ProtoReplayRequest>,
) -> Result<Response<ProtoReplayResponse>, Status> {
let req = request.into_inner();
let checkpoint_id = match req.checkpoint_id {
Some(ref s) => match Uuid::parse_str(s) {
Ok(id) => Some(id),
Err(e) => {
return Err(Status::invalid_argument(format!(
"invalid checkpoint_id '{}': {}",
s, e
)));
}
},
None => None,
};
let core_req = CoreReplayRequest {
thread_id: req.thread_id,
agent_id: req.agent_id,
checkpoint_id,
branch_name: req.branch_name,
as_of: req.as_of,
};
let result = self
.engine
.replay(core_req)
.await
.map_err(core_error_to_status)?;
let checkpoint_json =
serde_json::to_string(&result.checkpoint).unwrap_or_else(|_| "{}".to_string());
let memories: Vec<ProtoReplayMemory> = result
.memories
.iter()
.map(|m| ProtoReplayMemory {
id: m.id.to_string(),
content: m.content.clone(),
memory_type: format!("{:?}", m.memory_type),
created_at: m.created_at.clone(),
})
.collect();
let (chain_valid, chain_total, chain_verified) =
if let Some(ref cv) = result.chain_verification {
(
Some(cv.valid),
Some(cv.total_records as u32),
Some(cv.verified_records as u32),
)
} else {
(None, None, None)
};
Ok(Response::new(ProtoReplayResponse {
checkpoint_json,
memories,
event_count: result.events.len() as u32,
chain_valid,
chain_total,
chain_verified,
}))
}
async fn delegate(
&self,
request: Request<ProtoDelegateRequest>,
) -> Result<Response<ProtoDelegateResponse>, Status> {
let req = request.into_inner();
let permission: Permission = req
.permission
.parse()
.map_err(|e: mnemo_core::error::Error| Status::invalid_argument(e.to_string()))?;
let scope = if !req.memory_ids.is_empty() {
let ids: Vec<Uuid> = req
.memory_ids
.iter()
.map(|s| {
Uuid::parse_str(s)
.map_err(|e| Status::invalid_argument(format!("invalid UUID: {e}")))
})
.collect::<Result<Vec<_>, _>>()?;
DelegationScope::ByMemoryId(ids)
} else if !req.tags.is_empty() {
DelegationScope::ByTag(req.tags)
} else {
DelegationScope::AllMemories
};
let now = chrono::Utc::now();
let expires_at = req
.expires_in_hours
.map(|h| (now + chrono::Duration::seconds((h * 3600.0) as i64)).to_rfc3339());
let delegation = Delegation {
id: Uuid::now_v7(),
delegator_id: req.delegator_id,
delegate_id: req.delegate_id,
permission,
scope,
max_depth: req.max_depth.unwrap_or(0),
current_depth: 0,
parent_delegation_id: None,
created_at: now.to_rfc3339(),
expires_at,
revoked_at: None,
};
self.engine
.storage
.insert_delegation(&delegation)
.await
.map_err(core_error_to_status)?;
Ok(Response::new(ProtoDelegateResponse {
delegation_id: delegation.id.to_string(),
}))
}
async fn verify(
&self,
request: Request<ProtoVerifyRequest>,
) -> Result<Response<ProtoVerifyResponse>, Status> {
let req = request.into_inner();
let result = self
.engine
.verify_integrity(req.agent_id, req.thread_id.as_deref())
.await
.map_err(core_error_to_status)?;
Ok(Response::new(ProtoVerifyResponse {
valid: result.valid,
total_records: result.total_records as u32,
verified_records: result.verified_records as u32,
first_broken_at: result.first_broken_at.map(|id| id.to_string()),
error_message: result.error_message,
}))
}
async fn trajectory_audit(
&self,
request: Request<ProtoTrajectoryAuditRequest>,
) -> Result<Response<ProtoTrajectoryAuditResponse>, Status> {
let req = request.into_inner();
let agent_id = req
.agent_id
.clone()
.unwrap_or_else(|| self.engine.default_agent_id.clone());
let mut events = self
.engine
.storage
.list_events(&agent_id, mnemo_core::query::MAX_BATCH_QUERY_LIMIT, 0)
.await
.map_err(core_error_to_status)?;
events.reverse();
let mut audit_req = mnemo_compliance::trajectory::TrajectoryAuditRequest {
agent_id: Some(agent_id),
thread_id: req.thread_id.clone(),
..Default::default()
};
if let Some(c) = req.active_bank_ceiling {
audit_req.active_bank_ceiling = c as usize;
}
if let Some(k) = req.fact_key {
audit_req.fact_key = k;
}
if !req.named_forget_strategies.is_empty() {
audit_req.named_forget_strategies = req.named_forget_strategies;
}
let report = mnemo_compliance::trajectory::trajectory_audit(&events, &audit_req)
.map_err(|e| Status::invalid_argument(e.to_string()))?;
let report_json = serde_json::to_string(&report)
.map_err(|e| Status::internal(format!("report serialisation: {e}")))?;
Ok(Response::new(ProtoTrajectoryAuditResponse {
scope_label: report.scope_label.clone(),
event_count: report.event_count as u32,
all_ok: report.all_ok(),
unregulated_growth: Some(ProtoTrajectoryFinding {
severity: severity_to_str(report.unregulated_growth.severity).to_string(),
count: report.unregulated_growth.breach_count as u32,
}),
missing_semantic_revision: Some(ProtoTrajectoryFinding {
severity: severity_to_str(report.missing_semantic_revision.severity).to_string(),
count: report.missing_semantic_revision.stale_facts.len() as u32,
}),
capacity_driven_forgetting: Some(ProtoTrajectoryFinding {
severity: severity_to_str(report.capacity_driven_forgetting.severity).to_string(),
count: report
.capacity_driven_forgetting
.unlabelled_forget_event_ids
.len() as u32,
}),
read_only_retrieval: Some(ProtoTrajectoryFinding {
severity: severity_to_str(report.read_only_retrieval.severity).to_string(),
count: report.read_only_retrieval.read_only_scopes.len() as u32,
}),
report_json,
}))
}
async fn forget_subject(
&self,
request: Request<ProtoForgetSubjectRequest>,
) -> Result<Response<ProtoForgetSubjectResponse>, Status> {
let req = request.into_inner();
let strategy = match req.strategy.as_deref().unwrap_or("redact") {
"redact" => ForgetStrategy::Redact,
"hard_delete" => ForgetStrategy::HardDelete,
"soft_delete" => ForgetStrategy::SoftDelete,
other => {
return Err(Status::invalid_argument(format!(
"invalid forget_subject strategy '{}': expected one of: redact, hard_delete, soft_delete",
other
)));
}
};
let core_req = CoreForgetSubjectRequest {
subject_id: req.subject_id,
agent_id: req.agent_id,
strategy,
};
let result = self
.engine
.forget_subject(core_req)
.await
.map_err(core_error_to_status)?;
let errors: Vec<ProtoForgetError> = result
.errors
.into_iter()
.map(|e| ProtoForgetError {
id: e.id.to_string(),
error: e.error,
})
.collect();
let strategy_str = match result.strategy {
ForgetStrategy::SoftDelete => "soft_delete",
ForgetStrategy::HardDelete => "hard_delete",
ForgetStrategy::Decay => "decay",
ForgetStrategy::Consolidate => "consolidate",
ForgetStrategy::Archive => "archive",
ForgetStrategy::Redact => "redact",
}
.to_string();
Ok(Response::new(ProtoForgetSubjectResponse {
subject_id: result.subject_id,
strategy: strategy_str,
matched: result.matched as u32,
forgotten: result.forgotten.iter().map(|id| id.to_string()).collect(),
cascaded_events: result.cascaded_events as u32,
errors,
}))
}
}
pub fn router(engine: Arc<MnemoEngine>) -> tonic::transport::server::Router {
let token = std::env::var("MNEMO_AUTH_TOKEN")
.ok()
.filter(|s| !s.is_empty());
router_with_auth(engine, token)
}
pub fn router_with_auth(
engine: Arc<MnemoEngine>,
auth_token: Option<String>,
) -> tonic::transport::server::Router {
let svc = MnemoGrpcServer::new(engine);
match auth_token {
Some(token) if !token.is_empty() => {
tracing::info!(
"gRPC bearer-token auth ENABLED (authorization metadata = MNEMO_AUTH_TOKEN)"
);
let expected = Arc::new(token);
let interceptor = move |req: Request<()>| -> Result<Request<()>, Status> {
let provided = req
.metadata()
.get("authorization")
.and_then(|v| v.to_str().ok());
if mnemo_core::auth::bearer_token_matches(provided, &expected) {
Ok(req)
} else {
Err(Status::unauthenticated(
"missing or invalid bearer token (set `authorization` metadata)",
))
}
};
tonic::transport::Server::builder()
.add_service(MnemoServiceServer::with_interceptor(svc, interceptor))
}
_ => {
tracing::warn!(
"gRPC API running WITHOUT authentication — set MNEMO_AUTH_TOKEN to require a \
bearer token. Do not expose an unauthenticated memory server."
);
tonic::transport::Server::builder().add_service(MnemoServiceServer::new(svc))
}
}
}
fn core_error_to_status(err: mnemo_core::error::Error) -> Status {
use mnemo_core::error::Error;
match err {
Error::Validation(msg) => Status::invalid_argument(msg),
Error::PermissionDenied(msg) => Status::permission_denied(msg),
Error::NotFound(msg) => Status::not_found(msg),
other => Status::internal(other.to_string()),
}
}
fn severity_to_str(s: mnemo_compliance::Severity) -> &'static str {
match s {
mnemo_compliance::Severity::Ok => "ok",
mnemo_compliance::Severity::Warn => "warn",
mnemo_compliance::Severity::Fail => "fail",
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn core_error_maps_correctly() {
let validation =
core_error_to_status(mnemo_core::error::Error::Validation("bad input".into()));
assert_eq!(validation.code(), tonic::Code::InvalidArgument);
let perm = core_error_to_status(mnemo_core::error::Error::PermissionDenied(
"forbidden".into(),
));
assert_eq!(perm.code(), tonic::Code::PermissionDenied);
let not_found = core_error_to_status(mnemo_core::error::Error::NotFound("missing".into()));
assert_eq!(not_found.code(), tonic::Code::NotFound);
}
}