1use std::collections::HashMap;
7use std::pin::Pin;
8use std::sync::Arc;
9
10use mentedb_cognitive::stream::{CognitionStream, StreamAlert};
11use mentedb_core::edge::EdgeType;
12use mentedb_core::memory::{AttributeValue, MemoryType};
13use mentedb_core::{MemoryEdge, MemoryNode};
14use tokio::sync::mpsc;
15use tokio_stream::{Stream, StreamExt, wrappers::ReceiverStream};
16use tonic::{Request, Response, Status, Streaming};
17use tracing::{error, info};
18use uuid::Uuid;
19
20use crate::auth;
21use crate::state::AppState;
22
23pub mod pb {
24 tonic::include_proto!("mentedb");
25}
26
27use mentedb_core::types::{AgentId, MemoryId, SpaceId, UserId};
28use pb::cognition_service_server::CognitionService;
29use pb::memory_service_server::MemoryService;
30
31pub struct CognitionServiceImpl {
37 #[allow(dead_code)]
38 pub state: Arc<AppState>,
39}
40
41type CognitionStream_ =
42 Pin<Box<dyn Stream<Item = Result<pb::CognitionAlert, Status>> + Send + 'static>>;
43
44#[tonic::async_trait]
45impl CognitionService for CognitionServiceImpl {
46 type StreamCognitionStream = CognitionStream_;
47
48 async fn stream_cognition(
49 &self,
50 request: Request<Streaming<pb::CognitionRequest>>,
51 ) -> Result<Response<Self::StreamCognitionStream>, Status> {
52 authenticate_grpc_streaming(&self.state, request.metadata())?;
53 info!("gRPC cognition stream opened");
54
55 let stream_engine = Arc::new(CognitionStream::new(1000));
56 let known_facts: Arc<std::sync::Mutex<Vec<(MemoryId, String)>>> =
57 Arc::new(std::sync::Mutex::new(Vec::new()));
58
59 let (tx, rx) = mpsc::channel(128);
60 let mut inbound = request.into_inner();
61
62 let engine = stream_engine.clone();
63 let facts = known_facts.clone();
64
65 tokio::spawn(async move {
66 while let Some(result) = inbound.next().await {
67 let req = match result {
68 Ok(r) => r,
69 Err(e) => {
70 error!("cognition stream receive error: {e}");
71 break;
72 }
73 };
74
75 let payload = match req.payload {
76 Some(p) => p,
77 None => continue,
78 };
79
80 match payload {
81 pb::cognition_request::Payload::Token(t) => {
82 engine.feed_token(&t.text);
83
84 let current_facts = facts.lock().unwrap().clone();
85 let alerts = engine.check_alerts(¤t_facts);
86 for alert in alerts {
87 let proto_alert = stream_alert_to_proto(alert);
88 if tx.send(Ok(proto_alert)).await.is_err() {
89 return;
90 }
91 }
92 }
93 pb::cognition_request::Payload::KnownFact(kf) => {
94 let mid = kf
95 .memory_id
96 .parse::<MemoryId>()
97 .unwrap_or_else(|_| MemoryId::nil());
98 facts.lock().unwrap().push((mid, kf.content));
99 }
100 pb::cognition_request::Payload::EndOfTurn(_) => {
101 let current_facts = facts.lock().unwrap().clone();
102 let alerts = engine.check_alerts(¤t_facts);
103 for alert in alerts {
104 let proto_alert = stream_alert_to_proto(alert);
105 if tx.send(Ok(proto_alert)).await.is_err() {
106 return;
107 }
108 }
109 }
110 pb::cognition_request::Payload::Flush(_) => {
111 let text = engine.drain_buffer();
112 let alert = pb::CognitionAlert {
113 alert: Some(pb::cognition_alert::Alert::BufferFlushed(
114 pb::BufferFlushedAlert {
115 accumulated_text: text,
116 },
117 )),
118 };
119 if tx.send(Ok(alert)).await.is_err() {
120 return;
121 }
122 }
123 }
124 }
125 info!("gRPC cognition stream closed");
126 });
127
128 let output = ReceiverStream::new(rx);
129 Ok(Response::new(Box::pin(output)))
130 }
131}
132
133fn stream_alert_to_proto(alert: StreamAlert) -> pb::CognitionAlert {
134 let inner = match alert {
135 StreamAlert::Contradiction {
136 memory_id,
137 ai_said,
138 stored,
139 } => pb::cognition_alert::Alert::Contradiction(pb::ContradictionAlert {
140 memory_id: memory_id.to_string(),
141 ai_said,
142 stored,
143 }),
144 StreamAlert::Forgotten { memory_id, summary } => {
145 pb::cognition_alert::Alert::Forgotten(pb::ForgottenAlert {
146 memory_id: memory_id.to_string(),
147 summary,
148 })
149 }
150 StreamAlert::Correction {
151 memory_id,
152 old,
153 new,
154 } => pb::cognition_alert::Alert::Correction(pb::CorrectionAlert {
155 memory_id: memory_id.to_string(),
156 old_content: old,
157 new_content: new,
158 }),
159 StreamAlert::Reinforcement { memory_id } => {
160 pb::cognition_alert::Alert::Reinforcement(pb::ReinforcementAlert {
161 memory_id: memory_id.to_string(),
162 })
163 }
164 };
165 pb::CognitionAlert { alert: Some(inner) }
166}
167
168pub struct MemoryServiceImpl {
174 pub state: Arc<AppState>,
176}
177
178#[tonic::async_trait]
179impl MemoryService for MemoryServiceImpl {
180 async fn store(
181 &self,
182 request: Request<pb::StoreRequest>,
183 ) -> Result<Response<pb::StoreResponse>, Status> {
184 let caller = authenticate_grpc_request(&self.state, &request)?;
185 let req = request.into_inner();
186
187 let agent_id = AgentId(parse_uuid_field(&req.agent_id, "agent_id")?);
188 if let Some(ref ta) = caller {
189 let tu: AgentId = ta
190 .parse()
191 .map_err(|_| Status::internal("bad token agent_id"))?;
192 if tu != agent_id {
193 return Err(Status::permission_denied("agent_id mismatch"));
194 }
195 }
196 let memory_type = parse_memory_type_str(&req.memory_type)?;
197 let space_id = if req.space_id.is_empty() {
198 SpaceId::nil()
199 } else {
200 SpaceId(parse_uuid_field(&req.space_id, "space_id")?)
201 };
202
203 let salience = if req.salience == 0.0 {
204 0.5
205 } else {
206 req.salience
207 };
208 let confidence = if req.confidence == 0.0 {
209 1.0
210 } else {
211 req.confidence
212 };
213
214 let now = std::time::SystemTime::now()
215 .duration_since(std::time::UNIX_EPOCH)
216 .unwrap_or_default()
217 .as_micros() as u64;
218
219 let id = MemoryId::new();
220
221 let attributes = req
222 .attributes
223 .into_iter()
224 .map(|(k, v)| (k, AttributeValue::String(v)))
225 .collect::<HashMap<String, AttributeValue>>();
226
227 let node = MemoryNode {
228 id,
229 agent_id,
230 user_id: UserId::nil(),
231 memory_type,
232 embedding: req.embedding,
233 content: req.content,
234 created_at: now,
235 accessed_at: now,
236 access_count: 0,
237 salience,
238 confidence,
239 space_id,
240 attributes,
241 tags: req.tags,
242 valid_from: req.valid_from,
243 valid_until: req.valid_until,
244 };
245
246 let db = &*self.state.db;
247 db.store(node).map_err(|e| {
248 error!("gRPC store failed: {e}");
249 Status::internal(format!("store failed: {e}"))
250 })?;
251
252 Ok(Response::new(pb::StoreResponse {
253 id: id.to_string(),
254 status: "stored".into(),
255 }))
256 }
257
258 async fn recall(
259 &self,
260 request: Request<pb::RecallRequest>,
261 ) -> Result<Response<pb::RecallResponse>, Status> {
262 authenticate_grpc_request(&self.state, &request)?;
263 let req = request.into_inner();
264
265 let db = &*self.state.db;
266 let window = db.recall(&req.query).map_err(|e| {
267 error!("gRPC recall failed: {e}");
268 Status::internal(format!("recall failed: {e}"))
269 })?;
270
271 let memory_count: usize = window.blocks.iter().map(|b| b.memories.len()).sum();
272
273 Ok(Response::new(pb::RecallResponse {
274 context: window.format,
275 total_tokens: window.total_tokens as u64,
276 memory_count: memory_count as u64,
277 }))
278 }
279
280 async fn search(
281 &self,
282 request: Request<pb::SearchRequest>,
283 ) -> Result<Response<pb::SearchResponse>, Status> {
284 authenticate_grpc_request(&self.state, &request)?;
285 let req = request.into_inner();
286 let k = if req.k == 0 { 10 } else { req.k as usize };
287
288 if req.embedding.is_empty() {
289 return Err(Status::invalid_argument("missing embedding vector"));
290 }
291
292 let db = &*self.state.db;
293 let results = db.recall_similar(&req.embedding, k).map_err(|e| {
294 error!("gRPC search failed: {e}");
295 Status::internal(format!("search failed: {e}"))
296 })?;
297
298 let items: Vec<pb::SearchResult> = results
299 .iter()
300 .map(|(id, score)| pb::SearchResult {
301 id: id.to_string(),
302 score: *score,
303 })
304 .collect();
305
306 Ok(Response::new(pb::SearchResponse { results: items }))
307 }
308
309 async fn forget(
310 &self,
311 request: Request<pb::ForgetRequest>,
312 ) -> Result<Response<pb::ForgetResponse>, Status> {
313 authenticate_grpc_request(&self.state, &request)?;
314 let req = request.into_inner();
315 let id = MemoryId(parse_uuid_field(&req.id, "id")?);
316
317 let db = &*self.state.db;
318 db.forget(id).map_err(|e| {
319 error!("gRPC forget failed: {e}");
320 Status::internal(format!("forget failed: {e}"))
321 })?;
322
323 Ok(Response::new(pb::ForgetResponse {
324 status: "deleted".into(),
325 }))
326 }
327
328 async fn relate(
329 &self,
330 request: Request<pb::RelateRequest>,
331 ) -> Result<Response<pb::RelateResponse>, Status> {
332 authenticate_grpc_request(&self.state, &request)?;
333 let req = request.into_inner();
334 let source = MemoryId(parse_uuid_field(&req.source, "source")?);
335 let target = MemoryId(parse_uuid_field(&req.target, "target")?);
336 let edge_type = parse_edge_type_str(&req.edge_type)?;
337 let weight = if req.weight == 0.0 { 1.0 } else { req.weight };
338
339 let now = std::time::SystemTime::now()
340 .duration_since(std::time::UNIX_EPOCH)
341 .unwrap_or_default()
342 .as_micros() as u64;
343
344 let edge = MemoryEdge {
345 source,
346 target,
347 edge_type,
348 weight,
349 created_at: now,
350 valid_from: req.valid_from,
351 valid_until: req.valid_until,
352 label: None,
353 };
354
355 let db = &*self.state.db;
356 db.relate(edge).map_err(|e| {
357 error!("gRPC relate failed: {e}");
358 Status::internal(format!("relate failed: {e}"))
359 })?;
360
361 Ok(Response::new(pb::RelateResponse {
362 status: "created".into(),
363 }))
364 }
365}
366
367#[allow(clippy::result_large_err)]
368fn authenticate_grpc_request<T>(
369 state: &AppState,
370 request: &Request<T>,
371) -> Result<Option<String>, Status> {
372 let secret = match &state.jwt_secret {
373 Some(s) => s,
374 None => return Ok(None),
375 };
376 let token = request
377 .metadata()
378 .get("authorization")
379 .and_then(|v| v.to_str().ok())
380 .and_then(|v| v.strip_prefix("Bearer "))
381 .ok_or_else(|| Status::unauthenticated("missing or invalid authorization metadata"))?;
382 auth::validate_token(secret, token)
383 .map(|c| Some(c.agent_id))
384 .map_err(|e| Status::unauthenticated(format!("invalid token: {e}")))
385}
386#[allow(clippy::result_large_err)]
387fn authenticate_grpc_streaming(
388 state: &AppState,
389 metadata: &tonic::metadata::MetadataMap,
390) -> Result<Option<String>, Status> {
391 let secret = match &state.jwt_secret {
392 Some(s) => s,
393 None => return Ok(None),
394 };
395 let token = metadata
396 .get("authorization")
397 .and_then(|v| v.to_str().ok())
398 .and_then(|v| v.strip_prefix("Bearer "))
399 .ok_or_else(|| Status::unauthenticated("missing or invalid authorization metadata"))?;
400 auth::validate_token(secret, token)
401 .map(|c| Some(c.agent_id))
402 .map_err(|e| Status::unauthenticated(format!("invalid token: {e}")))
403}
404
405#[allow(clippy::result_large_err)]
410fn parse_uuid_field(s: &str, field: &str) -> Result<Uuid, Status> {
411 Uuid::parse_str(s)
412 .map_err(|_| Status::invalid_argument(format!("invalid UUID for '{field}': {s}")))
413}
414
415#[allow(clippy::result_large_err)]
416fn parse_memory_type_str(s: &str) -> Result<MemoryType, Status> {
417 match s.to_lowercase().as_str() {
418 "episodic" => Ok(MemoryType::Episodic),
419 "semantic" => Ok(MemoryType::Semantic),
420 "procedural" => Ok(MemoryType::Procedural),
421 "antipattern" | "anti_pattern" => Ok(MemoryType::AntiPattern),
422 "reasoning" => Ok(MemoryType::Reasoning),
423 "correction" => Ok(MemoryType::Correction),
424 _ => Err(Status::invalid_argument(format!(
425 "unknown memory_type: {s}"
426 ))),
427 }
428}
429
430#[allow(clippy::result_large_err)]
431fn parse_edge_type_str(s: &str) -> Result<EdgeType, Status> {
432 match s.to_lowercase().as_str() {
433 "caused" => Ok(EdgeType::Caused),
434 "before" => Ok(EdgeType::Before),
435 "related" => Ok(EdgeType::Related),
436 "contradicts" => Ok(EdgeType::Contradicts),
437 "supports" => Ok(EdgeType::Supports),
438 "supersedes" => Ok(EdgeType::Supersedes),
439 "derived" => Ok(EdgeType::Derived),
440 "partof" | "part_of" => Ok(EdgeType::PartOf),
441 _ => Err(Status::invalid_argument(format!("unknown edge_type: {s}"))),
442 }
443}
444
445#[cfg(test)]
450mod tests {
451 use super::*;
452 use mentedb::MenteDb;
453 use std::time::Instant;
454 use tempfile::TempDir;
455
456 fn make_test_state() -> (Arc<AppState>, TempDir) {
457 let tmp = TempDir::new().unwrap();
458 let db = MenteDb::open(tmp.path()).unwrap();
459 let state = Arc::new(AppState {
460 db: Arc::new(db),
461 spaces: Arc::new(tokio::sync::RwLock::new(mentedb_core::SpaceManager::new())),
462 jwt_secret: None,
463 admin_key: None,
464 start_time: Instant::now(),
465 extraction_config: None,
466 auto_extract: false,
467 extraction_tx: None,
468 cluster: None,
469 });
470 (state, tmp)
471 }
472
473 #[tokio::test]
474 async fn test_grpc_memory_store_and_recall() {
475 let (state, _tmp) = make_test_state();
476 let svc = MemoryServiceImpl {
477 state: state.clone(),
478 };
479
480 let agent_id = AgentId::new().to_string();
481 let store_req = Request::new(pb::StoreRequest {
482 agent_id: agent_id.clone(),
483 memory_type: "episodic".into(),
484 content: "The user prefers dark mode".into(),
485 embedding: vec![],
486 tags: vec!["preference".into()],
487 attributes: HashMap::new(),
488 space_id: String::new(),
489 salience: 0.8,
490 confidence: 1.0,
491 valid_from: None,
492 valid_until: None,
493 });
494
495 let resp = svc.store(store_req).await.unwrap();
496 let inner = resp.into_inner();
497 assert_eq!(inner.status, "stored");
498 assert!(!inner.id.is_empty());
499
500 let recall_req = Request::new(pb::RecallRequest {
501 query: "RECALL memories LIMIT 100".into(),
502 });
503 let resp = svc.recall(recall_req).await.unwrap();
504 let inner = resp.into_inner();
505 let _ = inner.memory_count;
508 }
509
510 #[tokio::test]
511 async fn test_grpc_memory_forget() {
512 let (state, _tmp) = make_test_state();
513 let svc = MemoryServiceImpl {
514 state: state.clone(),
515 };
516
517 let agent_id = AgentId::new().to_string();
518 let store_resp = svc
519 .store(Request::new(pb::StoreRequest {
520 agent_id,
521 memory_type: "semantic".into(),
522 content: "Temporary memory".into(),
523 embedding: vec![],
524 tags: vec![],
525 attributes: HashMap::new(),
526 space_id: String::new(),
527 salience: 0.5,
528 confidence: 1.0,
529 valid_from: None,
530 valid_until: None,
531 }))
532 .await
533 .unwrap();
534 let stored_id = store_resp.into_inner().id;
535
536 let forget_resp = svc
537 .forget(Request::new(pb::ForgetRequest {
538 id: stored_id.clone(),
539 }))
540 .await
541 .unwrap();
542 assert_eq!(forget_resp.into_inner().status, "deleted");
543 }
544
545 #[tokio::test]
546 async fn test_grpc_memory_relate() {
547 let (state, _tmp) = make_test_state();
548 let svc = MemoryServiceImpl {
549 state: state.clone(),
550 };
551
552 let agent_id = AgentId::new().to_string();
553 let id1 = svc
554 .store(Request::new(pb::StoreRequest {
555 agent_id: agent_id.clone(),
556 memory_type: "episodic".into(),
557 content: "Event A".into(),
558 embedding: vec![],
559 tags: vec![],
560 attributes: HashMap::new(),
561 space_id: String::new(),
562 salience: 0.5,
563 confidence: 1.0,
564 valid_from: None,
565 valid_until: None,
566 }))
567 .await
568 .unwrap()
569 .into_inner()
570 .id;
571
572 let id2 = svc
573 .store(Request::new(pb::StoreRequest {
574 agent_id,
575 memory_type: "episodic".into(),
576 content: "Event B".into(),
577 embedding: vec![],
578 tags: vec![],
579 attributes: HashMap::new(),
580 space_id: String::new(),
581 salience: 0.5,
582 confidence: 1.0,
583 valid_from: None,
584 valid_until: None,
585 }))
586 .await
587 .unwrap()
588 .into_inner()
589 .id;
590
591 let relate_resp = svc
592 .relate(Request::new(pb::RelateRequest {
593 source: id1,
594 target: id2,
595 edge_type: "caused".into(),
596 weight: 0.9,
597 valid_from: None,
598 valid_until: None,
599 }))
600 .await
601 .unwrap();
602 assert_eq!(relate_resp.into_inner().status, "created");
603 }
604
605 #[tokio::test]
606 async fn test_grpc_cognition_stream() {
607 use pb::cognition_service_server::CognitionServiceServer;
608 use tokio::net::TcpListener;
609 use tonic::transport::{Channel, Server};
610
611 let (state, _tmp) = make_test_state();
612
613 let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
614 let addr = listener.local_addr().unwrap();
615
616 let svc = CognitionServiceServer::new(CognitionServiceImpl {
617 state: state.clone(),
618 });
619
620 tokio::spawn(async move {
621 let incoming = tokio_stream::wrappers::TcpListenerStream::new(listener);
622 Server::builder()
623 .add_service(svc)
624 .serve_with_incoming(incoming)
625 .await
626 .unwrap();
627 });
628
629 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
631
632 let channel = Channel::from_shared(format!("http://{addr}"))
633 .unwrap()
634 .connect()
635 .await
636 .unwrap();
637
638 let mut client = pb::cognition_service_client::CognitionServiceClient::new(channel);
639
640 let (tx, rx) = mpsc::channel(32);
641 let outbound = ReceiverStream::new(rx);
642
643 let response = client.stream_cognition(outbound).await.unwrap();
644 let mut inbound = response.into_inner();
645
646 let mid = MemoryId::new();
648 tx.send(pb::CognitionRequest {
649 payload: Some(pb::cognition_request::Payload::KnownFact(pb::KnownFact {
650 memory_id: mid.to_string(),
651 content: "The system uses PostgreSQL for storage".into(),
652 })),
653 })
654 .await
655 .unwrap();
656
657 tx.send(pb::CognitionRequest {
659 payload: Some(pb::cognition_request::Payload::Token(pb::TokenPayload {
660 text: "The system does not use PostgreSQL, actually it uses MySQL".into(),
661 })),
662 })
663 .await
664 .unwrap();
665
666 if let Some(Ok(alert)) = inbound.next().await {
668 match alert.alert {
669 Some(pb::cognition_alert::Alert::Contradiction(c)) => {
670 assert_eq!(c.memory_id, mid.to_string());
671 assert!(!c.ai_said.is_empty());
672 assert!(!c.stored.is_empty());
673 }
674 other => panic!("expected contradiction alert, got: {:?}", other),
675 }
676 } else {
677 panic!("expected an alert from the cognition stream");
678 }
679
680 tx.send(pb::CognitionRequest {
682 payload: Some(pb::cognition_request::Payload::Flush(pb::FlushBuffer {})),
683 })
684 .await
685 .unwrap();
686
687 if let Some(Ok(alert)) = inbound.next().await {
688 match alert.alert {
689 Some(pb::cognition_alert::Alert::BufferFlushed(f)) => {
690 assert!(!f.accumulated_text.is_empty());
691 }
692 other => panic!("expected buffer_flushed alert, got: {:?}", other),
693 }
694 }
695
696 drop(tx);
697 }
698
699 #[tokio::test]
700 async fn test_grpc_invalid_uuid() {
701 let (state, _tmp) = make_test_state();
702 let svc = MemoryServiceImpl { state };
703
704 let result = svc
705 .store(Request::new(pb::StoreRequest {
706 agent_id: "not-a-uuid".into(),
707 memory_type: "episodic".into(),
708 content: "test".into(),
709 embedding: vec![],
710 tags: vec![],
711 attributes: HashMap::new(),
712 space_id: String::new(),
713 salience: 0.5,
714 confidence: 1.0,
715 valid_from: None,
716 valid_until: None,
717 }))
718 .await;
719
720 assert!(result.is_err());
721 let status = result.unwrap_err();
722 assert_eq!(status.code(), tonic::Code::InvalidArgument);
723 }
724
725 #[tokio::test]
726 async fn test_grpc_invalid_memory_type() {
727 let (state, _tmp) = make_test_state();
728 let svc = MemoryServiceImpl { state };
729
730 let result = svc
731 .store(Request::new(pb::StoreRequest {
732 agent_id: AgentId::new().to_string(),
733 memory_type: "nonexistent".into(),
734 content: "test".into(),
735 embedding: vec![],
736 tags: vec![],
737 attributes: HashMap::new(),
738 space_id: String::new(),
739 salience: 0.5,
740 confidence: 1.0,
741 valid_from: None,
742 valid_until: None,
743 }))
744 .await;
745
746 assert!(result.is_err());
747 let status = result.unwrap_err();
748 assert_eq!(status.code(), tonic::Code::InvalidArgument);
749 }
750
751 #[tokio::test]
752 async fn test_grpc_search_missing_embedding() {
753 let (state, _tmp) = make_test_state();
754 let svc = MemoryServiceImpl { state };
755
756 let result = svc
757 .search(Request::new(pb::SearchRequest {
758 embedding: vec![],
759 k: 10,
760 }))
761 .await;
762
763 assert!(result.is_err());
764 assert_eq!(result.unwrap_err().code(), tonic::Code::InvalidArgument);
765 }
766
767 fn make_auth_test_state(secret: &str) -> (Arc<AppState>, TempDir) {
768 let tmp = TempDir::new().unwrap();
769 let db = MenteDb::open(tmp.path()).unwrap();
770 (
771 Arc::new(AppState {
772 db: Arc::new(db),
773 spaces: Arc::new(tokio::sync::RwLock::new(mentedb_core::SpaceManager::new())),
774 jwt_secret: Some(secret.into()),
775 admin_key: Some("ak".into()),
776 start_time: Instant::now(),
777 extraction_config: None,
778 auto_extract: false,
779 extraction_tx: None,
780 cluster: None,
781 }),
782 tmp,
783 )
784 }
785 #[tokio::test]
786 async fn test_grpc_auth_required_when_secret_set() {
787 let (s, _t) = make_auth_test_state("s");
788 let svc = MemoryServiceImpl { state: s };
789 let r = svc
790 .recall(Request::new(pb::RecallRequest {
791 query: "RECALL memories LIMIT 10".into(),
792 }))
793 .await;
794 assert!(r.is_err());
795 assert_eq!(r.unwrap_err().code(), tonic::Code::Unauthenticated);
796 }
797 #[tokio::test]
798 async fn test_grpc_auth_succeeds_with_valid_token() {
799 let (s, _t) = make_auth_test_state("s");
800 let svc = MemoryServiceImpl { state: s };
801 let a = AgentId::new();
802 let tok = crate::auth::create_token("s", &a.to_string(), false, 1);
803 let mut r = Request::new(pb::StoreRequest {
804 agent_id: a.to_string(),
805 memory_type: "episodic".into(),
806 content: "t".into(),
807 embedding: vec![],
808 tags: vec![],
809 attributes: HashMap::new(),
810 space_id: String::new(),
811 salience: 0.5,
812 confidence: 1.0,
813 valid_from: None,
814 valid_until: None,
815 });
816 r.metadata_mut()
817 .insert("authorization", format!("Bearer {tok}").parse().unwrap());
818 assert_eq!(svc.store(r).await.unwrap().into_inner().status, "stored");
819 }
820 #[tokio::test]
821 async fn test_grpc_auth_rejects_wrong_agent_id() {
822 let (s, _t) = make_auth_test_state("s");
823 let svc = MemoryServiceImpl { state: s };
824 let ta = AgentId::new();
825 let ra = AgentId::new();
826 let tok = crate::auth::create_token("s", &ta.to_string(), false, 1);
827 let mut r = Request::new(pb::StoreRequest {
828 agent_id: ra.to_string(),
829 memory_type: "episodic".into(),
830 content: "t".into(),
831 embedding: vec![],
832 tags: vec![],
833 attributes: HashMap::new(),
834 space_id: String::new(),
835 salience: 0.5,
836 confidence: 1.0,
837 valid_from: None,
838 valid_until: None,
839 });
840 r.metadata_mut()
841 .insert("authorization", format!("Bearer {tok}").parse().unwrap());
842 assert_eq!(
843 svc.store(r).await.unwrap_err().code(),
844 tonic::Code::PermissionDenied
845 );
846 }
847 #[tokio::test]
848 async fn test_grpc_auth_invalid_token() {
849 let (s, _t) = make_auth_test_state("s");
850 let svc = MemoryServiceImpl { state: s };
851 let mut r = Request::new(pb::RecallRequest {
852 query: "RECALL memories LIMIT 10".into(),
853 });
854 r.metadata_mut()
855 .insert("authorization", "Bearer bad.tok".parse().unwrap());
856 assert_eq!(
857 svc.recall(r).await.unwrap_err().code(),
858 tonic::Code::Unauthenticated
859 );
860 }
861}