1use std::future::Future;
17use std::pin::Pin;
18
19use futures::StreamExt as _;
20use futures::stream::FuturesUnordered;
21
22use zeph_common::memory::{AsyncMemoryRouter, CompressionLevel, GraphRecallParams, TokenCounting};
23use zeph_llm::provider::{Message, MessageMetadata, MessagePart, Role};
24
25use crate::error::AssemblerError;
26use crate::input::ContextAssemblyInput;
27use crate::slot::ContextSlot;
28
29pub(crate) fn levels_to_flags(levels: &[CompressionLevel]) -> (bool, bool, bool) {
37 if levels.is_empty() {
38 return (true, true, true);
39 }
40 let episodic = levels.contains(&CompressionLevel::Episodic);
41 let procedural = levels.contains(&CompressionLevel::Procedural);
42 let declarative = levels.contains(&CompressionLevel::Declarative);
43 (episodic, procedural, declarative)
44}
45
46pub const SUMMARY_PREFIX: &str = "[conversation summaries]\n";
48pub const CROSS_SESSION_PREFIX: &str = "[cross-session context]\n";
50pub const RECALL_PREFIX: &str = "[semantic recall]\n";
52pub const CORRECTIONS_PREFIX: &str = "[past corrections]\n";
54pub const DOCUMENT_RAG_PREFIX: &str = "## Relevant documents\n";
56pub const GRAPH_FACTS_PREFIX: &str = "[known facts]\n";
58
59const MEMORY_FETCH_TIMEOUT_MS: u64 = 1000;
67
68#[derive(Default)]
73pub struct PreparedContext {
74 pub graph_facts: Option<Message>,
76 pub doc_rag: Option<Message>,
78 pub corrections: Option<Message>,
80 pub recall: Option<Message>,
82 pub recall_confidence: Option<f32>,
84 pub cross_session: Option<Message>,
86 pub summaries: Option<Message>,
88 pub code_context: Option<String>,
90 pub persona_facts: Option<Message>,
92 pub trajectory_hints: Option<Message>,
94 pub tree_memory: Option<Message>,
96 pub reasoning_hints: Option<Message>,
98 pub memory_first: bool,
100 pub recent_history_budget: usize,
102 pub background_tasks: Vec<tokio::task::JoinHandle<()>>,
107}
108
109pub struct ContextAssembler;
113
114type CtxFuture<'a> = Pin<Box<dyn Future<Output = Result<ContextSlot, AssemblerError>> + Send + 'a>>;
115
116fn empty_prepared_context() -> PreparedContext {
117 PreparedContext::default()
118}
119
120fn resolve_effective_strategy(
121 memory: &crate::input::ContextMemoryView,
122 sidequest_turn_counter: u64,
123) -> zeph_config::ContextStrategy {
124 match memory.context_strategy {
125 zeph_config::ContextStrategy::MemoryFirst => zeph_config::ContextStrategy::MemoryFirst,
126 zeph_config::ContextStrategy::Adaptive => {
127 if sidequest_turn_counter >= u64::from(memory.crossover_turn_threshold) {
128 zeph_config::ContextStrategy::MemoryFirst
129 } else {
130 zeph_config::ContextStrategy::FullHistory
131 }
132 }
133 _ => zeph_config::ContextStrategy::FullHistory,
134 }
135}
136
137fn correction_params(cfg: Option<&crate::input::CorrectionConfig>) -> (usize, f32) {
138 cfg.filter(|c| c.correction_detection)
139 .map_or((3, 0.75), |c| {
140 (
141 c.correction_recall_limit as usize,
142 c.correction_min_similarity,
143 )
144 })
145}
146
147#[allow(clippy::too_many_arguments)]
154fn schedule_context_fetchers<'r>(
155 memory: &'r crate::input::ContextMemoryView,
156 tc: &'r dyn TokenCounting,
157 query: &'r str,
158 scrub: fn(&str) -> std::borrow::Cow<'_, str>,
159 index: Option<&'r dyn crate::input::IndexAccess>,
160 router_ref: &'r dyn AsyncMemoryRouter,
161 summaries_budget: usize,
162 cross_session_budget: usize,
163 semantic_recall_budget: usize,
164 code_context_budget: usize,
165 graph_facts_budget: usize,
166 recall_limit: usize,
167 min_sim: f32,
168 active_levels: &[CompressionLevel],
169) -> FuturesUnordered<CtxFuture<'r>> {
170 let (episodic_active, procedural_active, declarative_active) = levels_to_flags(active_levels);
174
175 let fetchers: FuturesUnordered<CtxFuture<'r>> = FuturesUnordered::new();
176
177 if episodic_active && summaries_budget > 0 {
178 fetchers.push(Box::pin(async move {
179 fetch_summaries(memory, summaries_budget, tc)
180 .await
181 .map(ContextSlot::Summaries)
182 }));
183 }
184 if episodic_active && cross_session_budget > 0 {
185 fetchers.push(Box::pin(async move {
186 fetch_cross_session(memory, query, cross_session_budget, tc)
187 .await
188 .map(ContextSlot::CrossSession)
189 }));
190 }
191 if episodic_active && semantic_recall_budget > 0 {
192 fetchers.push(Box::pin(async move {
193 fetch_semantic_recall(memory, query, semantic_recall_budget, tc, Some(router_ref))
194 .await
195 .map(|(msg, score)| ContextSlot::SemanticRecall(msg, score))
196 }));
197 fetchers.push(Box::pin(async move {
198 fetch_document_rag(memory, query, semantic_recall_budget, tc)
199 .await
200 .map(ContextSlot::DocumentRag)
201 }));
202 }
203 fetchers.push(Box::pin(async move {
205 fetch_corrections(memory, query, recall_limit, min_sim, scrub)
206 .await
207 .map(ContextSlot::Corrections)
208 }));
209 if code_context_budget > 0
211 && let Some(idx) = index
212 {
213 fetchers.push(Box::pin(async move {
214 let result: Result<Option<String>, AssemblerError> = if let Ok(r) =
215 tokio::time::timeout(
216 std::time::Duration::from_millis(MEMORY_FETCH_TIMEOUT_MS),
217 idx.fetch_code_rag(query, code_context_budget),
218 )
219 .await
220 {
221 r
222 } else {
223 tracing::warn!("code RAG fetch timed out ({MEMORY_FETCH_TIMEOUT_MS}ms)");
224 Ok(None)
225 };
226 result.map(ContextSlot::CodeContext)
227 }));
228 }
229 if declarative_active && graph_facts_budget > 0 {
230 fetchers.push(Box::pin(async move {
231 fetch_graph_facts(memory, query, graph_facts_budget, tc)
232 .await
233 .map(ContextSlot::GraphFacts)
234 }));
235 }
236 if declarative_active && memory.persona_config.context_budget_tokens > 0 {
237 fetchers.push(Box::pin(async move {
238 let persona_budget = memory.persona_config.context_budget_tokens;
239 fetch_persona_facts(memory, persona_budget, tc)
240 .await
241 .map(ContextSlot::PersonaFacts)
242 }));
243 }
244 if procedural_active && memory.trajectory_config.context_budget_tokens > 0 {
245 fetchers.push(Box::pin(async move {
246 let tbudget = memory.trajectory_config.context_budget_tokens;
247 fetch_trajectory_hints(memory, tbudget, tc)
248 .await
249 .map(ContextSlot::TrajectoryHints)
250 }));
251 }
252 if declarative_active && memory.tree_config.context_budget_tokens > 0 {
253 fetchers.push(Box::pin(async move {
254 let tbudget = memory.tree_config.context_budget_tokens;
255 fetch_tree_memory(memory, tbudget, tc)
256 .await
257 .map(ContextSlot::TreeMemory)
258 }));
259 }
260 if procedural_active
261 && memory.reasoning_config.enabled
262 && memory.reasoning_config.context_budget_tokens > 0
263 {
264 fetchers.push(Box::pin(async move {
265 let rbudget = memory.reasoning_config.context_budget_tokens;
266 let top_k = memory.reasoning_config.top_k;
267 fetch_reasoning_strategies(memory, query, rbudget, top_k, tc)
268 .await
269 .map(|(msg, handle)| ContextSlot::ReasoningStrategies(msg, handle))
270 }));
271 }
272
273 fetchers
274}
275
276async fn drive_fetchers(
277 mut fetchers: FuturesUnordered<CtxFuture<'_>>,
278 prepared: &mut PreparedContext,
279) -> Result<(), AssemblerError> {
280 while let Some(result) = fetchers.next().await {
281 match result {
282 Ok(slot) => match slot {
283 ContextSlot::Summaries(msg) => prepared.summaries = msg,
284 ContextSlot::CrossSession(msg) => prepared.cross_session = msg,
285 ContextSlot::SemanticRecall(msg, score) => {
286 prepared.recall = msg;
287 prepared.recall_confidence = score;
288 }
289 ContextSlot::DocumentRag(msg) => prepared.doc_rag = msg,
290 ContextSlot::Corrections(msg) => prepared.corrections = msg,
291 ContextSlot::CodeContext(text) => prepared.code_context = text,
292 ContextSlot::GraphFacts(msg) => prepared.graph_facts = msg,
293 ContextSlot::PersonaFacts(msg) => prepared.persona_facts = msg,
294 ContextSlot::TrajectoryHints(msg) => prepared.trajectory_hints = msg,
295 ContextSlot::TreeMemory(msg) => prepared.tree_memory = msg,
296 ContextSlot::ReasoningStrategies(msg, handle) => {
297 prepared.reasoning_hints = msg;
298 if let Some(h) = handle {
299 prepared.background_tasks.push(h);
300 }
301 }
302 },
303 Err(e) => return Err(e),
304 }
305 }
306 Ok(())
307}
308
309impl ContextAssembler {
310 #[tracing::instrument(name = "context.assembler.gather", skip_all)]
318 pub async fn gather(
319 input: &ContextAssemblyInput<'_>,
320 ) -> Result<PreparedContext, AssemblerError> {
321 let Some(ref budget) = input.context_manager.budget else {
322 return Ok(empty_prepared_context());
323 };
324
325 let memory = input.memory;
326 let tc = input.token_counter;
327
328 let effective_strategy = resolve_effective_strategy(memory, input.sidequest_turn_counter);
329 let memory_first = effective_strategy == zeph_config::ContextStrategy::MemoryFirst;
330
331 let system_prompt = input
332 .messages
333 .first()
334 .filter(|m| m.role == Role::System)
335 .map_or("", |m| m.content.as_str());
336
337 let digest_tokens = memory
338 .cached_session_digest
339 .as_ref()
340 .map_or(0, |(_, tokens)| *tokens);
341
342 let alloc = budget.allocate_with_opts(
343 system_prompt,
344 input.skills_prompt,
345 tc,
346 memory.graph_config.enabled,
347 digest_tokens,
348 memory_first,
349 );
350
351 let (recall_limit, min_sim) = correction_params(input.correction_config.as_ref());
352
353 let router_ref: &dyn AsyncMemoryRouter = input.router.as_ref();
354
355 tracing::debug!(
356 active_sources = alloc.active_sources(),
357 active_levels = ?input.active_levels,
358 "context budget allocated"
359 );
360
361 let fetchers = schedule_context_fetchers(
362 memory,
363 tc,
364 input.query,
365 input.scrub,
366 input.index,
367 router_ref,
368 alloc.summaries,
369 alloc.cross_session,
370 alloc.semantic_recall,
371 alloc.code_context,
372 alloc.graph_facts,
373 recall_limit,
374 min_sim,
375 input.active_levels,
376 );
377
378 let mut prepared = empty_prepared_context();
379 prepared.memory_first = memory_first;
380 prepared.recent_history_budget = alloc.recent_history;
381
382 drive_fetchers(fetchers, &mut prepared).await?;
383 Ok(prepared)
384 }
385}
386
387pub fn effective_recall_timeout_ms(configured: u64) -> u64 {
392 if configured == 0 {
393 tracing::warn!(
394 "recall_timeout_ms is 0, which would disable spreading activation recall; \
395 clamping to 100ms"
396 );
397 100
398 } else {
399 configured
400 }
401}
402
403use crate::input::ContextMemoryView;
404
405#[tracing::instrument(name = "context.graph_facts", skip_all)]
406#[allow(clippy::too_many_lines)] pub(crate) async fn fetch_graph_facts(
408 memory: &ContextMemoryView,
409 query: &str,
410 budget_tokens: usize,
411 tc: &dyn TokenCounting,
412) -> Result<Option<Message>, AssemblerError> {
413 use zeph_common::memory::{RecallView, SpreadingActivationParams, classify_graph_subgraph};
414
415 if budget_tokens == 0 || !memory.graph_config.enabled {
416 return Ok(None);
417 }
418 let Some(ref mem) = memory.memory else {
419 return Ok(None);
420 };
421 let recall_limit = memory.graph_config.recall_limit;
422 let temporal_decay_rate = memory.graph_config.temporal_decay_rate;
423 let sa_config = &memory.graph_config.spreading_activation;
424
425 let fused_query;
427 let effective_query = if let Some(ref state) = memory.memcot_state {
428 let max_state_chars = 2 * query.len();
429 let state_slice = if state.len() > max_state_chars {
430 let boundary = state.floor_char_boundary(max_state_chars);
431 &state[..boundary]
432 } else {
433 state.as_str()
434 };
435 fused_query = format!("[state] {state_slice}\n{query}");
436 &fused_query as &str
437 } else {
438 query
439 };
440
441 let edge_types = classify_graph_subgraph(effective_query);
442
443 let view = match memory.memcot_config.recall_view {
444 zeph_config::RecallViewConfig::ZoomIn => RecallView::ZoomIn,
445 zeph_config::RecallViewConfig::ZoomOut => RecallView::ZoomOut,
446 _ => RecallView::Head,
447 };
448
449 let sa_params = if sa_config.enabled {
450 Some(SpreadingActivationParams {
451 decay_lambda: sa_config.decay_lambda,
452 max_hops: sa_config.max_hops,
453 activation_threshold: sa_config.activation_threshold,
454 inhibition_threshold: sa_config.inhibition_threshold,
455 max_activated_nodes: sa_config.max_activated_nodes,
456 temporal_decay_rate,
457 seed_structural_weight: sa_config.seed_structural_weight,
458 seed_community_cap: sa_config.seed_community_cap,
459 alpha: sa_config.alpha,
460 })
461 } else {
462 None
463 };
464
465 let timeout_ms = effective_recall_timeout_ms(sa_config.recall_timeout_ms);
466 let recall_fut = mem.recall_graph_facts(
467 effective_query,
468 GraphRecallParams {
469 limit: recall_limit,
470 view,
471 zoom_out_neighbor_cap: memory.memcot_config.zoom_out_neighbor_cap,
472 max_hops: memory.graph_config.max_hops,
473 temporal_decay_rate,
474 edge_types: &edge_types,
475 spreading_activation: sa_params,
476 },
477 );
478 let recalled = match tokio::time::timeout(
479 std::time::Duration::from_millis(timeout_ms),
480 recall_fut,
481 )
482 .await
483 {
484 Ok(Ok(facts)) => facts,
485 Ok(Err(e)) => {
486 tracing::warn!("graph recall failed: {e:#}");
487 Vec::new()
488 }
489 Err(_) => {
490 tracing::warn!("graph recall timed out ({timeout_ms}ms)");
491 Vec::new()
492 }
493 };
494
495 if recalled.is_empty() {
496 return Ok(None);
497 }
498
499 let mut body = String::from(GRAPH_FACTS_PREFIX);
500 let mut tokens_so_far = tc.count_tokens(&body);
501
502 for rf in &recalled {
503 let fact_text = rf.fact.replace(['\n', '\r', '<', '>'], " ");
504 let line = if let Some(score) = rf.activation_score {
505 format!(
506 "- {} (confidence: {:.2}, activation: {:.2})\n",
507 fact_text, rf.confidence, score
508 )
509 } else {
510 format!("- {} (confidence: {:.2})\n", fact_text, rf.confidence)
511 };
512 let line_tokens = tc.count_tokens(&line);
513 if tokens_so_far + line_tokens > budget_tokens {
514 break;
515 }
516 body.push_str(&line);
517 tokens_so_far += line_tokens;
518
519 for nb in &rf.neighbors {
521 let nb_text = nb.fact.replace(['\n', '\r', '<', '>'], " ");
522 let nb_line = format!(" ~ {} (confidence: {:.2})\n", nb_text, nb.confidence);
523 let nb_tokens = tc.count_tokens(&nb_line);
524 if tokens_so_far + nb_tokens > budget_tokens {
525 break;
526 }
527 body.push_str(&nb_line);
528 tokens_so_far += nb_tokens;
529 }
530
531 if let Some(ref snippet) = rf.provenance_snippet {
533 let snip_line = format!(
534 " [source: {}]\n",
535 snippet.replace(['\n', '\r', '<', '>'], " ")
536 );
537 let snip_tokens = tc.count_tokens(&snip_line);
538 if tokens_so_far + snip_tokens <= budget_tokens {
539 body.push_str(&snip_line);
540 tokens_so_far += snip_tokens;
541 }
542 }
543 }
544
545 if body == GRAPH_FACTS_PREFIX {
546 return Ok(None);
547 }
548
549 Ok(Some(Message::from_legacy(Role::System, body)))
550}
551
552fn append_budgeted_lines(
558 prefix: &str,
559 lines: impl Iterator<Item = String>,
560 budget_tokens: usize,
561 tc: &dyn TokenCounting,
562) -> Option<String> {
563 let mut body = String::from(prefix);
564 let mut tokens_so_far = tc.count_tokens(&body);
565
566 for line in lines {
567 let line_tokens = tc.count_tokens(&line);
568 if tokens_so_far + line_tokens > budget_tokens {
569 break;
570 }
571 body.push_str(&line);
572 tokens_so_far += line_tokens;
573 }
574
575 if body == prefix { None } else { Some(body) }
576}
577
578#[tracing::instrument(name = "context.persona_facts", skip_all)]
579pub(crate) async fn fetch_persona_facts(
580 memory: &ContextMemoryView,
581 budget_tokens: usize,
582 tc: &dyn TokenCounting,
583) -> Result<Option<Message>, AssemblerError> {
584 if budget_tokens == 0 || !memory.persona_config.enabled {
585 return Ok(None);
586 }
587 let Some(ref mem) = memory.memory else {
588 return Ok(None);
589 };
590
591 let min_confidence = memory.persona_config.min_confidence;
592 let facts = if let Ok(result) = tokio::time::timeout(
593 std::time::Duration::from_millis(MEMORY_FETCH_TIMEOUT_MS),
594 mem.load_persona_facts(min_confidence),
595 )
596 .await
597 {
598 result.map_err(AssemblerError::Memory)?
599 } else {
600 tracing::warn!("persona facts load timed out ({MEMORY_FETCH_TIMEOUT_MS}ms)");
601 Vec::new()
602 };
603
604 if facts.is_empty() {
605 return Ok(None);
606 }
607
608 let lines = facts
609 .iter()
610 .map(|fact| format!("[{}] {}\n", fact.category, fact.content));
611 Ok(
612 append_budgeted_lines(crate::slot::PERSONA_PREFIX, lines, budget_tokens, tc)
613 .map(|body| Message::from_legacy(Role::System, body)),
614 )
615}
616
617#[tracing::instrument(name = "context.trajectory_hints", skip_all)]
618pub(crate) async fn fetch_trajectory_hints(
619 memory: &ContextMemoryView,
620 budget_tokens: usize,
621 tc: &dyn TokenCounting,
622) -> Result<Option<Message>, AssemblerError> {
623 if budget_tokens == 0 || !memory.trajectory_config.enabled {
624 return Ok(None);
625 }
626 let Some(ref mem) = memory.memory else {
627 return Ok(None);
628 };
629
630 let top_k = memory.trajectory_config.recall_top_k;
631 let min_conf = memory.trajectory_config.min_confidence;
632 let entries = if let Ok(result) = tokio::time::timeout(
636 std::time::Duration::from_millis(MEMORY_FETCH_TIMEOUT_MS),
637 mem.load_trajectory_entries(Some("procedural"), top_k),
638 )
639 .await
640 {
641 result.map_err(AssemblerError::Memory)?
642 } else {
643 tracing::warn!("trajectory entries load timed out ({MEMORY_FETCH_TIMEOUT_MS}ms)");
644 Vec::new()
645 };
646
647 if entries.is_empty() {
648 return Ok(None);
649 }
650
651 let lines = entries
652 .iter()
653 .filter(|e| e.confidence >= min_conf)
654 .take(top_k)
655 .map(|entry| format!("- {}: {}\n", entry.intent, entry.outcome));
656 Ok(
657 append_budgeted_lines(crate::slot::TRAJECTORY_PREFIX, lines, budget_tokens, tc)
658 .map(|body| Message::from_legacy(Role::System, body)),
659 )
660}
661
662#[tracing::instrument(name = "context.tree_memory", skip_all)]
663pub(crate) async fn fetch_tree_memory(
664 memory: &ContextMemoryView,
665 budget_tokens: usize,
666 tc: &dyn TokenCounting,
667) -> Result<Option<Message>, AssemblerError> {
668 if budget_tokens == 0 || !memory.tree_config.enabled {
669 return Ok(None);
670 }
671 let Some(ref mem) = memory.memory else {
672 return Ok(None);
673 };
674
675 let top_k = memory.tree_config.recall_top_k;
676 let nodes = if let Ok(result) = tokio::time::timeout(
677 std::time::Duration::from_millis(MEMORY_FETCH_TIMEOUT_MS),
678 mem.load_tree_nodes(1, top_k),
679 )
680 .await
681 {
682 result.map_err(AssemblerError::Memory)?
683 } else {
684 tracing::warn!("tree nodes load timed out ({MEMORY_FETCH_TIMEOUT_MS}ms)");
685 Vec::new()
686 };
687
688 if nodes.is_empty() {
689 return Ok(None);
690 }
691
692 let lines = nodes
693 .iter()
694 .take(top_k)
695 .map(|node| format!("- {}\n", node.content));
696 Ok(
697 append_budgeted_lines(crate::slot::TREE_MEMORY_PREFIX, lines, budget_tokens, tc)
698 .map(|body| Message::from_legacy(Role::System, body)),
699 )
700}
701
702#[tracing::instrument(name = "context.reasoning_strategies", skip_all)]
703pub(crate) async fn fetch_reasoning_strategies(
704 memory: &ContextMemoryView,
705 query: &str,
706 budget_tokens: usize,
707 top_k: usize,
708 tc: &dyn TokenCounting,
709) -> Result<(Option<Message>, Option<tokio::task::JoinHandle<()>>), AssemblerError> {
710 let budget_tokens = budget_tokens.min(500);
712 if budget_tokens == 0 {
713 return Ok((None, None));
714 }
715 let Some(ref mem) = memory.memory else {
716 return Ok((None, None));
717 };
718
719 let strategies = if let Ok(result) = tokio::time::timeout(
720 std::time::Duration::from_millis(MEMORY_FETCH_TIMEOUT_MS),
721 mem.retrieve_reasoning_strategies(query, top_k),
722 )
723 .await
724 {
725 result.map_err(AssemblerError::Memory)?
726 } else {
727 tracing::warn!("reasoning strategies retrieval timed out ({MEMORY_FETCH_TIMEOUT_MS}ms)");
728 Vec::new()
729 };
730
731 if strategies.is_empty() {
732 return Ok((None, None));
733 }
734
735 let mut body = String::from(crate::slot::REASONING_PREFIX);
736 let mut tokens_so_far = tc.count_tokens(&body);
737 let mut injected_ids: Vec<String> = Vec::new();
738
739 for s in strategies.iter().take(top_k) {
740 let safe_summary = s.summary.replace(['\n', '\r', '<', '>'], " ");
743 let line = format!("- [{}] {}\n", s.outcome, safe_summary);
744 let line_tokens = tc.count_tokens(&line);
745 if tokens_so_far + line_tokens > budget_tokens {
746 break;
747 }
748 body.push_str(&line);
749 tokens_so_far += line_tokens;
750 injected_ids.push(s.id.clone());
751 }
752
753 if body == crate::slot::REASONING_PREFIX {
754 return Ok((None, None));
755 }
756
757 let handle = if injected_ids.is_empty() {
761 None
762 } else {
763 let mem_clone = mem.clone();
764 let mark_used = async move {
765 if let Err(e) = mem_clone.mark_reasoning_used(&injected_ids).await {
766 tracing::warn!(error = %e, "reasoning: mark_used failed");
767 }
768 };
769 Some(tokio::spawn(mark_used)) };
771
772 Ok((Some(Message::from_legacy(Role::System, body)), handle))
773}
774
775#[tracing::instrument(name = "context.corrections", skip_all)]
776pub(crate) async fn fetch_corrections(
777 memory: &ContextMemoryView,
778 query: &str,
779 limit: usize,
780 min_score: f32,
781 scrub: fn(&str) -> std::borrow::Cow<'_, str>,
782) -> Result<Option<Message>, AssemblerError> {
783 let Some(ref mem) = memory.memory else {
784 return Ok(None);
785 };
786 let corrections = if let Ok(result) = tokio::time::timeout(
787 std::time::Duration::from_millis(MEMORY_FETCH_TIMEOUT_MS),
788 mem.retrieve_corrections(query, limit, min_score),
789 )
790 .await
791 {
792 result.map_err(AssemblerError::Memory)?
793 } else {
794 tracing::warn!("corrections retrieval timed out ({MEMORY_FETCH_TIMEOUT_MS}ms)");
795 Vec::new()
796 };
797 if corrections.is_empty() {
798 return Ok(None);
799 }
800 let mut text = String::from(CORRECTIONS_PREFIX);
801 for c in &corrections {
802 text.push_str("- Past user correction: \"");
803 text.push_str(&scrub(&c.correction_text));
804 text.push_str("\"\n");
805 }
806 Ok(Some(Message::from_legacy(Role::System, text)))
807}
808
809#[tracing::instrument(name = "context.semantic_recall", skip_all)]
810pub(crate) async fn fetch_semantic_recall(
811 memory: &ContextMemoryView,
812 query: &str,
813 token_budget: usize,
814 tc: &dyn TokenCounting,
815 router: Option<&dyn AsyncMemoryRouter>,
816) -> Result<(Option<Message>, Option<f32>), AssemblerError> {
817 let Some(ref mem) = memory.memory else {
818 return Ok((None, None));
819 };
820 if memory.recall_limit == 0 || token_budget == 0 {
821 return Ok((None, None));
822 }
823
824 let recalled = if let Ok(result) = tokio::time::timeout(
825 std::time::Duration::from_millis(MEMORY_FETCH_TIMEOUT_MS),
826 mem.recall(query, memory.recall_limit, router),
827 )
828 .await
829 {
830 result.map_err(AssemblerError::Memory)?
831 } else {
832 tracing::warn!("semantic recall timed out ({MEMORY_FETCH_TIMEOUT_MS}ms)");
833 Vec::new()
834 };
835 if recalled.is_empty() {
836 return Ok((None, None));
837 }
838
839 let top_score = recalled.first().map(|r| r.score);
840
841 let mut recall_text = String::with_capacity(token_budget * 3);
842 recall_text.push_str(RECALL_PREFIX);
843 let mut tokens_used = tc.count_tokens(&recall_text);
844
845 for item in &recalled {
846 if item.content.starts_with("[skipped]") || item.content.starts_with("[stopped]") {
847 continue;
848 }
849 let entry = format!("- [{}] {}\n", item.role, item.content);
850 let entry_tokens = tc.count_tokens(&entry);
851 if tokens_used + entry_tokens > token_budget {
852 break;
853 }
854 recall_text.push_str(&entry);
855 tokens_used += entry_tokens;
856 }
857
858 if tokens_used > tc.count_tokens(RECALL_PREFIX) {
859 Ok((
860 Some(Message::from_parts(
861 Role::System,
862 vec![MessagePart::Recall { text: recall_text }],
863 )),
864 top_score,
865 ))
866 } else {
867 Ok((None, None))
868 }
869}
870
871#[tracing::instrument(name = "context.document_rag", skip_all)]
872pub(crate) async fn fetch_document_rag(
873 memory: &ContextMemoryView,
874 query: &str,
875 token_budget: usize,
876 tc: &dyn TokenCounting,
877) -> Result<Option<Message>, AssemblerError> {
878 if !memory.document_config.rag_enabled || token_budget == 0 {
879 return Ok(None);
880 }
881 let Some(ref mem) = memory.memory else {
882 return Ok(None);
883 };
884
885 let collection = &memory.document_config.collection;
886 let top_k = memory.document_config.top_k;
887 let chunks = if let Ok(result) = tokio::time::timeout(
888 std::time::Duration::from_millis(MEMORY_FETCH_TIMEOUT_MS),
889 mem.search_document_collection(collection, query, top_k),
890 )
891 .await
892 {
893 result.map_err(AssemblerError::Memory)?
894 } else {
895 tracing::warn!("document RAG search timed out ({MEMORY_FETCH_TIMEOUT_MS}ms)");
896 Vec::new()
897 };
898 if chunks.is_empty() {
899 return Ok(None);
900 }
901
902 let mut text = String::from(DOCUMENT_RAG_PREFIX);
903 let mut tokens_used = tc.count_tokens(&text);
904
905 for chunk in &chunks {
906 if chunk.text.is_empty() {
907 continue;
908 }
909 let entry = format!("{}\n", chunk.text);
910 let cost = tc.count_tokens(&entry);
911 if tokens_used + cost > token_budget {
912 break;
913 }
914 text.push_str(&entry);
915 tokens_used += cost;
916 }
917
918 if tokens_used > tc.count_tokens(DOCUMENT_RAG_PREFIX) {
919 Ok(Some(Message {
920 role: Role::System,
921 content: text,
922 parts: vec![],
923 metadata: MessageMetadata::default(),
924 }))
925 } else {
926 Ok(None)
927 }
928}
929
930#[tracing::instrument(name = "context.summaries", skip_all)]
931pub(crate) async fn fetch_summaries(
932 memory: &ContextMemoryView,
933 token_budget: usize,
934 tc: &dyn TokenCounting,
935) -> Result<Option<Message>, AssemblerError> {
936 let (Some(mem), Some(cid)) = (&memory.memory, memory.conversation_id) else {
937 return Ok(None);
938 };
939 if token_budget == 0 {
940 return Ok(None);
941 }
942
943 let summaries = if let Ok(result) = tokio::time::timeout(
944 std::time::Duration::from_millis(MEMORY_FETCH_TIMEOUT_MS),
945 mem.load_summaries(cid),
946 )
947 .await
948 {
949 result.map_err(AssemblerError::Memory)?
950 } else {
951 tracing::warn!("summaries load timed out ({MEMORY_FETCH_TIMEOUT_MS}ms)");
952 Vec::new()
953 };
954 if summaries.is_empty() {
955 return Ok(None);
956 }
957
958 let mut summary_text = String::from(SUMMARY_PREFIX);
959 let mut tokens_used = tc.count_tokens(&summary_text);
960
961 for summary in summaries.iter().rev() {
962 let first = summary.first_message_id.unwrap_or(0);
963 let last = summary.last_message_id.unwrap_or(0);
964 let entry = format!("- Messages {first}-{last}: {}\n", summary.content);
965 let cost = tc.count_tokens(&entry);
966 if tokens_used + cost > token_budget {
967 break;
968 }
969 summary_text.push_str(&entry);
970 tokens_used += cost;
971 }
972
973 if tokens_used > tc.count_tokens(SUMMARY_PREFIX) {
974 Ok(Some(Message::from_parts(
975 Role::System,
976 vec![MessagePart::Summary { text: summary_text }],
977 )))
978 } else {
979 Ok(None)
980 }
981}
982
983#[tracing::instrument(name = "context.cross_session", skip_all)]
984pub(crate) async fn fetch_cross_session(
985 memory: &ContextMemoryView,
986 query: &str,
987 token_budget: usize,
988 tc: &dyn TokenCounting,
989) -> Result<Option<Message>, AssemblerError> {
990 let (Some(mem), Some(cid)) = (&memory.memory, memory.conversation_id) else {
991 return Ok(None);
992 };
993 if token_budget == 0 {
994 return Ok(None);
995 }
996
997 let threshold = memory.cross_session_score_threshold;
998 let summaries = if let Ok(result) = tokio::time::timeout(
999 std::time::Duration::from_millis(MEMORY_FETCH_TIMEOUT_MS),
1000 mem.search_session_summaries(query, 5, Some(cid)),
1001 )
1002 .await
1003 {
1004 result.map_err(AssemblerError::Memory)?
1005 } else {
1006 tracing::warn!("cross-session search timed out ({MEMORY_FETCH_TIMEOUT_MS}ms)");
1007 Vec::new()
1008 };
1009 let results: Vec<_> = summaries
1010 .into_iter()
1011 .filter(|r| r.score >= threshold)
1012 .collect();
1013 if results.is_empty() {
1014 return Ok(None);
1015 }
1016
1017 let mut text = String::from(CROSS_SESSION_PREFIX);
1018 let mut tokens_used = tc.count_tokens(&text);
1019
1020 for item in &results {
1021 let entry = format!("- {}\n", item.summary_text);
1022 let cost = tc.count_tokens(&entry);
1023 if tokens_used + cost > token_budget {
1024 break;
1025 }
1026 text.push_str(&entry);
1027 tokens_used += cost;
1028 }
1029
1030 if tokens_used > tc.count_tokens(CROSS_SESSION_PREFIX) {
1031 Ok(Some(Message::from_parts(
1032 Role::System,
1033 vec![MessagePart::CrossSession { text }],
1034 )))
1035 } else {
1036 Ok(None)
1037 }
1038}
1039
1040pub const MAX_KEEP_TAIL_SCAN: usize = 50;
1043
1044#[must_use]
1052pub fn memory_first_keep_tail(messages: &[Message], history_start: usize) -> usize {
1053 use zeph_llm::provider::MessagePart;
1054
1055 let mut keep_tail = 2usize;
1056 let len = messages.len();
1057 let max = len.saturating_sub(history_start);
1058
1059 while keep_tail < max {
1060 let first_retained = &messages[len - keep_tail];
1061 let is_tool_result = first_retained.role == Role::User
1062 && first_retained
1063 .parts
1064 .iter()
1065 .any(|p| matches!(p, MessagePart::ToolResult { .. }));
1066
1067 if is_tool_result {
1068 keep_tail += 1;
1069 } else {
1070 break;
1071 }
1072
1073 if keep_tail >= MAX_KEEP_TAIL_SCAN {
1074 let preceding_idx = len.saturating_sub(keep_tail + 1);
1075 if preceding_idx >= history_start {
1076 let preceding = &messages[preceding_idx];
1077 let is_tool_use = preceding.role == Role::Assistant
1078 && preceding
1079 .parts
1080 .iter()
1081 .any(|p| matches!(p, MessagePart::ToolUse { .. }));
1082 if is_tool_use {
1083 keep_tail += 1;
1084 }
1085 }
1086 break;
1087 }
1088 }
1089
1090 keep_tail
1091}
1092
1093#[cfg(test)]
1094mod tests {
1095 use super::*;
1096 use crate::input::ContextMemoryView;
1097 use zeph_common::memory::CompressionLevel;
1098 use zeph_config::{
1099 ContextStrategy, DocumentConfig, GraphConfig, PersonaConfig, ReasoningConfig,
1100 TrajectoryConfig, TreeConfig,
1101 };
1102
1103 struct NaiveTokenCounter;
1104 impl zeph_common::memory::TokenCounting for NaiveTokenCounter {
1105 fn count_tokens(&self, text: &str) -> usize {
1106 text.split_whitespace().count()
1107 }
1108 fn count_tool_schema_tokens(&self, schema: &serde_json::Value) -> usize {
1109 schema.to_string().split_whitespace().count()
1110 }
1111 }
1112
1113 fn empty_view() -> ContextMemoryView {
1114 ContextMemoryView {
1115 memory: None,
1116 conversation_id: None,
1117 recall_limit: 10,
1118 cross_session_score_threshold: 0.5,
1119 context_strategy: ContextStrategy::default(),
1120 crossover_turn_threshold: 5,
1121 cached_session_digest: None,
1122 graph_config: GraphConfig::default(),
1123 document_config: DocumentConfig::default(),
1124 persona_config: PersonaConfig::default(),
1125 trajectory_config: TrajectoryConfig::default(),
1126 reasoning_config: ReasoningConfig::default(),
1127 memcot_config: zeph_config::MemCotConfig::default(),
1128 memcot_state: None,
1129 tree_config: TreeConfig::default(),
1130 }
1131 }
1132
1133 #[tokio::test]
1136 async fn fetch_graph_facts_returns_none_when_memory_is_none() {
1137 let view = empty_view();
1138 let tc = NaiveTokenCounter;
1139 let result = fetch_graph_facts(&view, "test", 1000, &tc).await.unwrap();
1140 assert!(result.is_none());
1141 }
1142
1143 #[tokio::test]
1144 async fn fetch_graph_facts_returns_none_when_budget_zero() {
1145 let mut view = empty_view();
1146 view.graph_config.enabled = true;
1147 let tc = NaiveTokenCounter;
1148 let result = fetch_graph_facts(&view, "test", 0, &tc).await.unwrap();
1149 assert!(result.is_none());
1150 }
1151
1152 #[tokio::test]
1153 async fn fetch_graph_facts_returns_none_when_graph_disabled() {
1154 let mut view = empty_view();
1155 view.graph_config.enabled = false;
1156 let tc = NaiveTokenCounter;
1157 let result = fetch_graph_facts(&view, "test", 1000, &tc).await.unwrap();
1158 assert!(result.is_none());
1159 }
1160
1161 #[tokio::test]
1164 async fn fetch_persona_facts_returns_none_when_memory_is_none() {
1165 let view = empty_view();
1166 let tc = NaiveTokenCounter;
1167 let result = fetch_persona_facts(&view, 1000, &tc).await.unwrap();
1168 assert!(result.is_none());
1169 }
1170
1171 #[tokio::test]
1172 async fn fetch_persona_facts_returns_none_when_budget_zero() {
1173 let mut view = empty_view();
1174 view.persona_config.enabled = true;
1175 let tc = NaiveTokenCounter;
1176 let result = fetch_persona_facts(&view, 0, &tc).await.unwrap();
1177 assert!(result.is_none());
1178 }
1179
1180 #[tokio::test]
1183 async fn fetch_trajectory_hints_returns_none_when_memory_is_none() {
1184 let view = empty_view();
1185 let tc = NaiveTokenCounter;
1186 let result = fetch_trajectory_hints(&view, 1000, &tc).await.unwrap();
1187 assert!(result.is_none());
1188 }
1189
1190 #[tokio::test]
1191 async fn fetch_trajectory_hints_returns_none_when_budget_zero() {
1192 let mut view = empty_view();
1193 view.trajectory_config.enabled = true;
1194 let tc = NaiveTokenCounter;
1195 let result = fetch_trajectory_hints(&view, 0, &tc).await.unwrap();
1196 assert!(result.is_none());
1197 }
1198
1199 #[tokio::test]
1202 async fn fetch_tree_memory_returns_none_when_memory_is_none() {
1203 let view = empty_view();
1204 let tc = NaiveTokenCounter;
1205 let result = fetch_tree_memory(&view, 1000, &tc).await.unwrap();
1206 assert!(result.is_none());
1207 }
1208
1209 #[tokio::test]
1210 async fn fetch_tree_memory_returns_none_when_budget_zero() {
1211 let mut view = empty_view();
1212 view.tree_config.enabled = true;
1213 let tc = NaiveTokenCounter;
1214 let result = fetch_tree_memory(&view, 0, &tc).await.unwrap();
1215 assert!(result.is_none());
1216 }
1217
1218 #[tokio::test]
1221 async fn fetch_corrections_returns_none_when_memory_is_none() {
1222 let view = empty_view();
1223 let result = fetch_corrections(&view, "test", 10, 0.5, |s| s.into())
1224 .await
1225 .unwrap();
1226 assert!(result.is_none());
1227 }
1228
1229 #[tokio::test]
1232 async fn fetch_semantic_recall_returns_none_when_memory_is_none() {
1233 let view = empty_view();
1234 let tc = NaiveTokenCounter;
1235 let result = fetch_semantic_recall(&view, "test", 1000, &tc, None)
1236 .await
1237 .unwrap();
1238 assert!(result.0.is_none() && result.1.is_none());
1239 }
1240
1241 #[tokio::test]
1242 async fn fetch_semantic_recall_returns_none_when_budget_zero() {
1243 let view = empty_view();
1244 let tc = NaiveTokenCounter;
1245 let result = fetch_semantic_recall(&view, "test", 0, &tc, None)
1246 .await
1247 .unwrap();
1248 assert!(result.0.is_none() && result.1.is_none());
1249 }
1250
1251 #[tokio::test]
1254 async fn fetch_document_rag_returns_none_when_memory_is_none() {
1255 let mut view = empty_view();
1256 view.document_config.rag_enabled = true;
1257 let tc = NaiveTokenCounter;
1258 let result = fetch_document_rag(&view, "test", 1000, &tc).await.unwrap();
1259 assert!(result.is_none());
1260 }
1261
1262 #[tokio::test]
1263 async fn fetch_document_rag_returns_none_when_rag_disabled() {
1264 let view = empty_view();
1265 let tc = NaiveTokenCounter;
1266 let result = fetch_document_rag(&view, "test", 1000, &tc).await.unwrap();
1267 assert!(result.is_none());
1268 }
1269
1270 #[tokio::test]
1273 async fn fetch_summaries_returns_none_when_memory_is_none() {
1274 let view = empty_view();
1275 let tc = NaiveTokenCounter;
1276 let result = fetch_summaries(&view, 1000, &tc).await.unwrap();
1277 assert!(result.is_none());
1278 }
1279
1280 #[tokio::test]
1283 async fn fetch_cross_session_returns_none_when_memory_is_none() {
1284 let view = empty_view();
1285 let tc = NaiveTokenCounter;
1286 let result = fetch_cross_session(&view, "test", 1000, &tc).await.unwrap();
1287 assert!(result.is_none());
1288 }
1289
1290 #[test]
1293 fn levels_to_flags_empty_slice_enables_all_tiers() {
1294 let (e, p, d) = levels_to_flags(&[]);
1295 assert!(e, "episodic should be active for empty slice");
1296 assert!(p, "procedural should be active for empty slice");
1297 assert!(d, "declarative should be active for empty slice");
1298 }
1299
1300 #[test]
1301 fn levels_to_flags_full_set_enables_all_tiers() {
1302 let all = &[
1303 CompressionLevel::Episodic,
1304 CompressionLevel::Procedural,
1305 CompressionLevel::Declarative,
1306 ];
1307 let (e, p, d) = levels_to_flags(all);
1308 assert!(e);
1309 assert!(p);
1310 assert!(d);
1311 }
1312
1313 #[test]
1314 fn levels_to_flags_episodic_only() {
1315 let (e, p, d) = levels_to_flags(&[CompressionLevel::Episodic]);
1316 assert!(e);
1317 assert!(!p, "procedural should be inactive");
1318 assert!(!d, "declarative should be inactive");
1319 }
1320
1321 #[test]
1322 fn levels_to_flags_episodic_and_procedural() {
1323 let (e, p, d) =
1324 levels_to_flags(&[CompressionLevel::Episodic, CompressionLevel::Procedural]);
1325 assert!(e);
1326 assert!(p);
1327 assert!(!d, "declarative should be inactive");
1328 }
1329
1330 #[test]
1331 fn levels_to_flags_declarative_only() {
1332 let (e, p, d) = levels_to_flags(&[CompressionLevel::Declarative]);
1333 assert!(!e, "episodic should be inactive");
1334 assert!(!p, "procedural should be inactive");
1335 assert!(d);
1336 }
1337
1338 #[tokio::test]
1341 async fn fetch_reasoning_strategies_returns_none_when_memory_is_none() {
1342 let mut view = empty_view();
1343 view.reasoning_config.enabled = true;
1344 let tc = NaiveTokenCounter;
1345 let (result, handle) = fetch_reasoning_strategies(&view, "query", 1000, 3, &tc)
1346 .await
1347 .unwrap();
1348 assert!(result.is_none());
1349 assert!(handle.is_none());
1350 }
1351
1352 #[tokio::test]
1353 async fn fetch_reasoning_strategies_returns_none_when_budget_zero() {
1354 let mut view = empty_view();
1355 view.reasoning_config.enabled = true;
1356 let tc = NaiveTokenCounter;
1357 let (result, handle) = fetch_reasoning_strategies(&view, "query", 0, 3, &tc)
1358 .await
1359 .unwrap();
1360 assert!(result.is_none());
1361 assert!(handle.is_none());
1362 }
1363
1364 use std::sync::{Arc, Mutex};
1367 use zeph_common::memory::{
1368 ContextMemoryBackend, GraphRecallParams, MemCorrection, MemDocumentChunk, MemGraphFact,
1369 MemPersonaFact, MemReasoningStrategy, MemRecalledMessage, MemSessionSummary, MemSummary,
1370 MemTrajectoryEntry, MemTreeNode,
1371 };
1372
1373 const KNOWN_FAIL_ON: &[&str] = &[
1375 "load_persona_facts",
1376 "load_trajectory_entries",
1377 "load_tree_nodes",
1378 "load_summaries",
1379 "retrieve_reasoning_strategies",
1380 "mark_reasoning_used",
1381 "retrieve_corrections",
1382 "recall",
1383 "recall_graph_facts",
1384 "search_session_summaries",
1385 "search_document_collection",
1386 ];
1387
1388 #[derive(Default)]
1389 struct MockMemoryBackend {
1390 persona_facts: Vec<MemPersonaFact>,
1391 trajectory_entries: Vec<MemTrajectoryEntry>,
1392 tree_nodes: Vec<MemTreeNode>,
1393 summaries: Vec<MemSummary>,
1394 reasoning_strategies: Vec<MemReasoningStrategy>,
1395 corrections: Vec<MemCorrection>,
1396 recalled: Vec<MemRecalledMessage>,
1397 graph_facts: Vec<MemGraphFact>,
1398 session_summaries: Vec<MemSessionSummary>,
1399 document_chunks: Vec<MemDocumentChunk>,
1400 fail_on: Option<&'static str>,
1402 delay: Option<std::time::Duration>,
1405 marked_ids: Mutex<Vec<String>>,
1407 }
1408
1409 impl MockMemoryBackend {
1410 fn with_fail_on(method: &'static str) -> Self {
1411 debug_assert!(
1412 KNOWN_FAIL_ON.contains(&method),
1413 "unknown fail_on method name: {method}"
1414 );
1415 Self {
1416 fail_on: Some(method),
1417 ..Default::default()
1418 }
1419 }
1420
1421 fn fail_err(method: &str) -> Box<dyn std::error::Error + Send + Sync> {
1422 format!("mock error in {method}").into()
1423 }
1424 }
1425
1426 impl ContextMemoryBackend for MockMemoryBackend {
1427 fn load_persona_facts<'a>(
1428 &'a self,
1429 _min_confidence: f64,
1430 ) -> std::pin::Pin<
1431 Box<
1432 dyn std::future::Future<
1433 Output = Result<
1434 Vec<MemPersonaFact>,
1435 Box<dyn std::error::Error + Send + Sync>,
1436 >,
1437 > + Send
1438 + 'a,
1439 >,
1440 > {
1441 let result = if self.fail_on == Some("load_persona_facts") {
1442 Err(Self::fail_err("load_persona_facts"))
1443 } else {
1444 Ok(self.persona_facts.clone())
1445 };
1446 let delay = self.delay;
1447 Box::pin(async move {
1448 if let Some(d) = delay {
1449 tokio::time::sleep(d).await;
1450 }
1451 result
1452 })
1453 }
1454
1455 fn load_trajectory_entries<'a>(
1456 &'a self,
1457 _tier: Option<&'a str>,
1458 _top_k: usize,
1459 ) -> std::pin::Pin<
1460 Box<
1461 dyn std::future::Future<
1462 Output = Result<
1463 Vec<MemTrajectoryEntry>,
1464 Box<dyn std::error::Error + Send + Sync>,
1465 >,
1466 > + Send
1467 + 'a,
1468 >,
1469 > {
1470 let result = if self.fail_on == Some("load_trajectory_entries") {
1471 Err(Self::fail_err("load_trajectory_entries"))
1472 } else {
1473 Ok(self.trajectory_entries.clone())
1474 };
1475 Box::pin(async move { result })
1476 }
1477
1478 fn load_tree_nodes<'a>(
1479 &'a self,
1480 _level: u32,
1481 _top_k: usize,
1482 ) -> std::pin::Pin<
1483 Box<
1484 dyn std::future::Future<
1485 Output = Result<Vec<MemTreeNode>, Box<dyn std::error::Error + Send + Sync>>,
1486 > + Send
1487 + 'a,
1488 >,
1489 > {
1490 let result = if self.fail_on == Some("load_tree_nodes") {
1491 Err(Self::fail_err("load_tree_nodes"))
1492 } else {
1493 Ok(self.tree_nodes.clone())
1494 };
1495 Box::pin(async move { result })
1496 }
1497
1498 fn load_summaries<'a>(
1499 &'a self,
1500 _conversation_id: i64,
1501 ) -> std::pin::Pin<
1502 Box<
1503 dyn std::future::Future<
1504 Output = Result<Vec<MemSummary>, Box<dyn std::error::Error + Send + Sync>>,
1505 > + Send
1506 + 'a,
1507 >,
1508 > {
1509 let result = if self.fail_on == Some("load_summaries") {
1510 Err(Self::fail_err("load_summaries"))
1511 } else {
1512 Ok(self.summaries.clone())
1513 };
1514 Box::pin(async move { result })
1515 }
1516
1517 fn retrieve_reasoning_strategies<'a>(
1518 &'a self,
1519 _query: &'a str,
1520 _top_k: usize,
1521 ) -> std::pin::Pin<
1522 Box<
1523 dyn std::future::Future<
1524 Output = Result<
1525 Vec<MemReasoningStrategy>,
1526 Box<dyn std::error::Error + Send + Sync>,
1527 >,
1528 > + Send
1529 + 'a,
1530 >,
1531 > {
1532 let result = if self.fail_on == Some("retrieve_reasoning_strategies") {
1533 Err(Self::fail_err("retrieve_reasoning_strategies"))
1534 } else {
1535 Ok(self.reasoning_strategies.clone())
1536 };
1537 Box::pin(async move { result })
1538 }
1539
1540 fn mark_reasoning_used<'a>(
1541 &'a self,
1542 ids: &'a [String],
1543 ) -> std::pin::Pin<
1544 Box<
1545 dyn std::future::Future<
1546 Output = Result<(), Box<dyn std::error::Error + Send + Sync>>,
1547 > + Send
1548 + 'a,
1549 >,
1550 > {
1551 if self.fail_on == Some("mark_reasoning_used") {
1552 return Box::pin(async move { Err(Self::fail_err("mark_reasoning_used")) });
1553 }
1554 let mut guard = self.marked_ids.lock().expect("marked_ids poisoned");
1555 guard.extend_from_slice(ids);
1556 Box::pin(async move { Ok(()) })
1557 }
1558
1559 fn retrieve_corrections<'a>(
1560 &'a self,
1561 _query: &'a str,
1562 _limit: usize,
1563 _min_score: f32,
1564 ) -> std::pin::Pin<
1565 Box<
1566 dyn std::future::Future<
1567 Output = Result<
1568 Vec<MemCorrection>,
1569 Box<dyn std::error::Error + Send + Sync>,
1570 >,
1571 > + Send
1572 + 'a,
1573 >,
1574 > {
1575 let result = if self.fail_on == Some("retrieve_corrections") {
1576 Err(Self::fail_err("retrieve_corrections"))
1577 } else {
1578 Ok(self.corrections.clone())
1579 };
1580 Box::pin(async move { result })
1581 }
1582
1583 fn recall<'a>(
1584 &'a self,
1585 _query: &'a str,
1586 _limit: usize,
1587 _router: Option<&'a dyn zeph_common::memory::AsyncMemoryRouter>,
1588 ) -> std::pin::Pin<
1589 Box<
1590 dyn std::future::Future<
1591 Output = Result<
1592 Vec<MemRecalledMessage>,
1593 Box<dyn std::error::Error + Send + Sync>,
1594 >,
1595 > + Send
1596 + 'a,
1597 >,
1598 > {
1599 let result = if self.fail_on == Some("recall") {
1600 Err(Self::fail_err("recall"))
1601 } else {
1602 Ok(self.recalled.clone())
1603 };
1604 let delay = self.delay;
1605 Box::pin(async move {
1606 if let Some(d) = delay {
1607 tokio::time::sleep(d).await;
1608 }
1609 result
1610 })
1611 }
1612
1613 fn recall_graph_facts<'a>(
1614 &'a self,
1615 _query: &'a str,
1616 _params: GraphRecallParams<'a>,
1617 ) -> std::pin::Pin<
1618 Box<
1619 dyn std::future::Future<
1620 Output = Result<
1621 Vec<MemGraphFact>,
1622 Box<dyn std::error::Error + Send + Sync>,
1623 >,
1624 > + Send
1625 + 'a,
1626 >,
1627 > {
1628 let result = if self.fail_on == Some("recall_graph_facts") {
1629 Err(Self::fail_err("recall_graph_facts"))
1630 } else {
1631 Ok(self.graph_facts.clone())
1632 };
1633 Box::pin(async move { result })
1634 }
1635
1636 fn search_session_summaries<'a>(
1637 &'a self,
1638 _query: &'a str,
1639 _limit: usize,
1640 _current_conversation_id: Option<i64>,
1641 ) -> std::pin::Pin<
1642 Box<
1643 dyn std::future::Future<
1644 Output = Result<
1645 Vec<MemSessionSummary>,
1646 Box<dyn std::error::Error + Send + Sync>,
1647 >,
1648 > + Send
1649 + 'a,
1650 >,
1651 > {
1652 let result = if self.fail_on == Some("search_session_summaries") {
1653 Err(Self::fail_err("search_session_summaries"))
1654 } else {
1655 Ok(self.session_summaries.clone())
1656 };
1657 Box::pin(async move { result })
1658 }
1659
1660 fn search_document_collection<'a>(
1661 &'a self,
1662 _collection: &'a str,
1663 _query: &'a str,
1664 _top_k: usize,
1665 ) -> std::pin::Pin<
1666 Box<
1667 dyn std::future::Future<
1668 Output = Result<
1669 Vec<MemDocumentChunk>,
1670 Box<dyn std::error::Error + Send + Sync>,
1671 >,
1672 > + Send
1673 + 'a,
1674 >,
1675 > {
1676 let result = if self.fail_on == Some("search_document_collection") {
1677 Err(Self::fail_err("search_document_collection"))
1678 } else {
1679 Ok(self.document_chunks.clone())
1680 };
1681 Box::pin(async move { result })
1682 }
1683 }
1684
1685 fn mock_view(mock: MockMemoryBackend) -> ContextMemoryView {
1686 let mut v = empty_view();
1687 v.memory = Some(Arc::new(mock));
1688 v
1689 }
1690
1691 #[tokio::test]
1694 async fn fetch_graph_facts_returns_message_when_memory_present() {
1695 let mock = MockMemoryBackend {
1696 graph_facts: vec![zeph_common::memory::MemGraphFact {
1697 fact: "Rust is fast".to_string(),
1698 confidence: 0.9,
1699 activation_score: None,
1700 neighbors: vec![],
1701 provenance_snippet: None,
1702 }],
1703 ..Default::default()
1704 };
1705 let mut view = mock_view(mock);
1706 view.graph_config.enabled = true;
1707 view.graph_config.spreading_activation.recall_timeout_ms = 5000;
1709 let tc = NaiveTokenCounter;
1710 let result = fetch_graph_facts(&view, "test", 1000, &tc).await.unwrap();
1711 assert!(result.is_some(), "expected Some message");
1712 let msg = result.unwrap();
1713 assert!(
1714 msg.content.contains("Rust is fast"),
1715 "expected fact text in output, got: {}",
1716 msg.content
1717 );
1718 assert!(
1719 msg.content.starts_with(GRAPH_FACTS_PREFIX),
1720 "expected GRAPH_FACTS_PREFIX"
1721 );
1722 }
1723
1724 #[tokio::test]
1725 async fn fetch_graph_facts_swallows_error_and_returns_none() {
1726 let mock = MockMemoryBackend::with_fail_on("recall_graph_facts");
1727 let mut view = mock_view(mock);
1728 view.graph_config.enabled = true;
1729 view.graph_config.spreading_activation.recall_timeout_ms = 5000;
1730 let tc = NaiveTokenCounter;
1731 let result = fetch_graph_facts(&view, "test", 1000, &tc).await.unwrap();
1733 assert!(
1734 result.is_none(),
1735 "expected None when recall_graph_facts errors"
1736 );
1737 }
1738
1739 #[tokio::test]
1740 async fn fetch_graph_facts_returns_none_when_facts_empty() {
1741 let mock = MockMemoryBackend::default(); let mut view = mock_view(mock);
1743 view.graph_config.enabled = true;
1744 view.graph_config.spreading_activation.recall_timeout_ms = 5000;
1745 let tc = NaiveTokenCounter;
1746 let result = fetch_graph_facts(&view, "test", 1000, &tc).await.unwrap();
1747 assert!(result.is_none());
1748 }
1749
1750 #[tokio::test]
1753 async fn fetch_persona_facts_returns_message_when_persona_enabled() {
1754 let mock = MockMemoryBackend {
1755 persona_facts: vec![MemPersonaFact {
1756 category: "preference".to_string(),
1757 content: "prefers concise answers".to_string(),
1758 }],
1759 ..Default::default()
1760 };
1761 let mut view = mock_view(mock);
1762 view.persona_config.enabled = true;
1763 view.persona_config.context_budget_tokens = 1000;
1764 let tc = NaiveTokenCounter;
1765 let result = fetch_persona_facts(&view, 1000, &tc).await.unwrap();
1766 assert!(result.is_some());
1767 let msg = result.unwrap();
1768 assert!(msg.content.contains("preference"));
1769 assert!(msg.content.contains("prefers concise answers"));
1770 assert!(msg.content.starts_with(crate::slot::PERSONA_PREFIX));
1771 }
1772
1773 #[tokio::test]
1774 async fn fetch_persona_facts_propagates_error() {
1775 let mock = MockMemoryBackend::with_fail_on("load_persona_facts");
1776 let mut view = mock_view(mock);
1777 view.persona_config.enabled = true;
1778 let tc = NaiveTokenCounter;
1779 let result = fetch_persona_facts(&view, 1000, &tc).await;
1780 assert!(
1781 result.is_err(),
1782 "expected Err from load_persona_facts failure"
1783 );
1784 }
1785
1786 #[tokio::test]
1789 async fn fetch_trajectory_hints_returns_message_when_trajectory_enabled() {
1790 let mock = MockMemoryBackend {
1791 trajectory_entries: vec![MemTrajectoryEntry {
1792 intent: "summarize code".to_string(),
1793 outcome: "produced concise summary".to_string(),
1794 confidence: 0.9,
1795 }],
1796 ..Default::default()
1797 };
1798 let mut view = mock_view(mock);
1799 view.trajectory_config.enabled = true;
1800 view.trajectory_config.context_budget_tokens = 1000;
1801 view.trajectory_config.min_confidence = 0.5;
1802 let tc = NaiveTokenCounter;
1803 let result = fetch_trajectory_hints(&view, 1000, &tc).await.unwrap();
1804 assert!(result.is_some());
1805 let msg = result.unwrap();
1806 assert!(msg.content.contains("summarize code"));
1807 assert!(msg.content.starts_with(crate::slot::TRAJECTORY_PREFIX));
1808 }
1809
1810 #[tokio::test]
1811 async fn fetch_trajectory_hints_passes_tier_filter() {
1812 let mock = MockMemoryBackend {
1815 trajectory_entries: vec![
1816 MemTrajectoryEntry {
1817 intent: "debug async code".to_string(),
1818 outcome: "fixed deadlock".to_string(),
1819 confidence: 0.85,
1820 },
1821 MemTrajectoryEntry {
1822 intent: "low confidence task".to_string(),
1823 outcome: "irrelevant".to_string(),
1824 confidence: 0.3,
1825 },
1826 ],
1827 ..Default::default()
1828 };
1829 let mut view = mock_view(mock);
1830 view.trajectory_config.enabled = true;
1831 view.trajectory_config.context_budget_tokens = 1000;
1832 view.trajectory_config.min_confidence = 0.5;
1833 let tc = NaiveTokenCounter;
1834 let result = fetch_trajectory_hints(&view, 1000, &tc).await.unwrap();
1835 assert!(result.is_some(), "expected Some message");
1836 let msg = result.unwrap();
1837 assert!(
1838 msg.content.contains("debug async code"),
1839 "high-confidence entry must be included"
1840 );
1841 assert!(
1842 !msg.content.contains("low confidence task"),
1843 "entry below min_confidence must be filtered out"
1844 );
1845 }
1846
1847 #[tokio::test]
1848 async fn fetch_trajectory_hints_propagates_error() {
1849 let mock = MockMemoryBackend::with_fail_on("load_trajectory_entries");
1850 let mut view = mock_view(mock);
1851 view.trajectory_config.enabled = true;
1852 let tc = NaiveTokenCounter;
1853 let result = fetch_trajectory_hints(&view, 1000, &tc).await;
1854 assert!(result.is_err());
1855 }
1856
1857 #[tokio::test]
1860 async fn fetch_tree_memory_returns_message_when_tree_enabled() {
1861 let mock = MockMemoryBackend {
1862 tree_nodes: vec![MemTreeNode {
1863 content: "Topic: async Rust patterns".to_string(),
1864 }],
1865 ..Default::default()
1866 };
1867 let mut view = mock_view(mock);
1868 view.tree_config.enabled = true;
1869 view.tree_config.context_budget_tokens = 1000;
1870 let tc = NaiveTokenCounter;
1871 let result = fetch_tree_memory(&view, 1000, &tc).await.unwrap();
1872 assert!(result.is_some());
1873 let msg = result.unwrap();
1874 assert!(msg.content.contains("async Rust patterns"));
1875 assert!(msg.content.starts_with(crate::slot::TREE_MEMORY_PREFIX));
1876 }
1877
1878 #[tokio::test]
1879 async fn fetch_tree_memory_propagates_error() {
1880 let mock = MockMemoryBackend::with_fail_on("load_tree_nodes");
1881 let mut view = mock_view(mock);
1882 view.tree_config.enabled = true;
1883 let tc = NaiveTokenCounter;
1884 let result = fetch_tree_memory(&view, 1000, &tc).await;
1885 assert!(result.is_err());
1886 }
1887
1888 #[tokio::test]
1891 async fn fetch_corrections_returns_message_when_corrections_present() {
1892 let mock = MockMemoryBackend {
1893 corrections: vec![MemCorrection {
1894 correction_text: "use snake_case not camelCase".to_string(),
1895 }],
1896 ..Default::default()
1897 };
1898 let view = mock_view(mock);
1899 let result = fetch_corrections(&view, "query", 10, 0.5, |s| s.into())
1900 .await
1901 .unwrap();
1902 assert!(result.is_some());
1903 let msg = result.unwrap();
1904 assert!(msg.content.contains("snake_case"));
1905 assert!(msg.content.starts_with(CORRECTIONS_PREFIX));
1906 }
1907
1908 #[tokio::test]
1909 async fn fetch_corrections_propagates_error() {
1910 let mock = MockMemoryBackend::with_fail_on("retrieve_corrections");
1913 let view = mock_view(mock);
1914 let result = fetch_corrections(&view, "query", 10, 0.5, |s| s.into()).await;
1915 assert!(result.is_err(), "expected Err, got {result:?}");
1916 }
1917
1918 #[tokio::test]
1921 async fn fetch_semantic_recall_returns_message_with_content() {
1922 let mock = MockMemoryBackend {
1923 recalled: vec![
1924 MemRecalledMessage {
1925 role: "user".to_string(),
1926 content: "how does tokio work".to_string(),
1927 score: 0.95,
1928 },
1929 MemRecalledMessage {
1930 role: "assistant".to_string(),
1931 content: "tokio is an async runtime".to_string(),
1932 score: 0.88,
1933 },
1934 ],
1935 ..Default::default()
1936 };
1937 let mut view = mock_view(mock);
1938 view.recall_limit = 10;
1939 let tc = NaiveTokenCounter;
1940 let (msg, score) = fetch_semantic_recall(&view, "tokio", 1000, &tc, None)
1941 .await
1942 .unwrap();
1943 assert!(msg.is_some(), "expected Some message");
1944 assert!(score.is_some_and(|s| (s - 0.95_f32).abs() < f32::EPSILON));
1946 let msg = msg.unwrap();
1947 let has_recall_part = msg.parts.iter().any(|p| {
1949 if let zeph_llm::provider::MessagePart::Recall { text } = p {
1950 text.contains("how does tokio work")
1951 } else {
1952 false
1953 }
1954 });
1955 assert!(has_recall_part, "expected recalled content in Recall part");
1956 }
1957
1958 #[tokio::test]
1959 async fn fetch_semantic_recall_returns_none_when_recalled_empty() {
1960 let mock = MockMemoryBackend::default();
1961 let mut view = mock_view(mock);
1962 view.recall_limit = 10;
1963 let tc = NaiveTokenCounter;
1964 let (msg, score) = fetch_semantic_recall(&view, "query", 1000, &tc, None)
1965 .await
1966 .unwrap();
1967 assert!(msg.is_none());
1968 assert!(score.is_none());
1969 }
1970
1971 #[tokio::test]
1972 async fn fetch_semantic_recall_propagates_error() {
1973 let mock = MockMemoryBackend::with_fail_on("recall");
1974 let mut view = mock_view(mock);
1975 view.recall_limit = 10;
1976 let tc = NaiveTokenCounter;
1977 let result = fetch_semantic_recall(&view, "query", 1000, &tc, None).await;
1978 assert!(result.is_err());
1979 }
1980
1981 #[tokio::test]
1984 async fn fetch_document_rag_returns_message_when_rag_enabled() {
1985 let mock = MockMemoryBackend {
1986 document_chunks: vec![MemDocumentChunk {
1987 text: "Rust ownership rules prevent data races".to_string(),
1988 }],
1989 ..Default::default()
1990 };
1991 let mut view = mock_view(mock);
1992 view.document_config.rag_enabled = true;
1993 let tc = NaiveTokenCounter;
1994 let result = fetch_document_rag(&view, "ownership", 1000, &tc)
1995 .await
1996 .unwrap();
1997 assert!(result.is_some());
1998 let msg = result.unwrap();
1999 assert!(msg.content.contains("ownership rules"));
2000 assert!(msg.content.starts_with(DOCUMENT_RAG_PREFIX));
2001 }
2002
2003 #[tokio::test]
2004 async fn fetch_document_rag_propagates_error() {
2005 let mock = MockMemoryBackend::with_fail_on("search_document_collection");
2006 let mut view = mock_view(mock);
2007 view.document_config.rag_enabled = true;
2008 let tc = NaiveTokenCounter;
2009 let result = fetch_document_rag(&view, "query", 1000, &tc).await;
2010 assert!(result.is_err());
2011 }
2012
2013 #[tokio::test]
2016 async fn fetch_summaries_returns_message_when_summaries_present() {
2017 let mock = MockMemoryBackend {
2018 summaries: vec![MemSummary {
2019 first_message_id: Some(1),
2020 last_message_id: Some(5),
2021 content: "User asked about async Rust".to_string(),
2022 }],
2023 ..Default::default()
2024 };
2025 let mut view = mock_view(mock);
2026 view.conversation_id = Some(42);
2027 let tc = NaiveTokenCounter;
2028 let result = fetch_summaries(&view, 1000, &tc).await.unwrap();
2029 assert!(result.is_some());
2030 let msg = result.unwrap();
2031 let has_summary_part = msg.parts.iter().any(|p| {
2032 if let zeph_llm::provider::MessagePart::Summary { text } = p {
2033 text.contains("Messages 1-5") && text.contains("async Rust")
2034 } else {
2035 false
2036 }
2037 });
2038 assert!(
2039 has_summary_part,
2040 "expected Summary part with messages range"
2041 );
2042 }
2043
2044 #[tokio::test]
2045 async fn fetch_summaries_returns_none_without_conversation_id() {
2046 let mock = MockMemoryBackend {
2047 summaries: vec![MemSummary {
2048 first_message_id: Some(1),
2049 last_message_id: Some(5),
2050 content: "some content".to_string(),
2051 }],
2052 ..Default::default()
2053 };
2054 let mut view = mock_view(mock);
2055 view.conversation_id = None; let tc = NaiveTokenCounter;
2057 let result = fetch_summaries(&view, 1000, &tc).await.unwrap();
2058 assert!(result.is_none());
2059 }
2060
2061 #[tokio::test]
2062 async fn fetch_summaries_propagates_error() {
2063 let mock = MockMemoryBackend::with_fail_on("load_summaries");
2064 let mut view = mock_view(mock);
2065 view.conversation_id = Some(42);
2066 let tc = NaiveTokenCounter;
2067 let result = fetch_summaries(&view, 1000, &tc).await;
2068 assert!(result.is_err());
2069 }
2070
2071 #[tokio::test]
2074 async fn fetch_cross_session_returns_message_when_results_present() {
2075 let mock = MockMemoryBackend {
2076 session_summaries: vec![MemSessionSummary {
2077 summary_text: "Previous session: debugging tokio deadlock".to_string(),
2078 score: 0.9,
2079 }],
2080 ..Default::default()
2081 };
2082 let mut view = mock_view(mock);
2083 view.conversation_id = Some(1);
2084 view.cross_session_score_threshold = 0.5;
2085 let tc = NaiveTokenCounter;
2086 let result = fetch_cross_session(&view, "async", 1000, &tc)
2087 .await
2088 .unwrap();
2089 assert!(result.is_some());
2090 let msg = result.unwrap();
2091 let has_cross_session_part = msg.parts.iter().any(|p| {
2092 if let zeph_llm::provider::MessagePart::CrossSession { text } = p {
2093 text.contains("tokio deadlock")
2094 } else {
2095 false
2096 }
2097 });
2098 assert!(has_cross_session_part);
2099 }
2100
2101 #[tokio::test]
2102 async fn fetch_cross_session_propagates_error() {
2103 let mock = MockMemoryBackend::with_fail_on("search_session_summaries");
2104 let mut view = mock_view(mock);
2105 view.conversation_id = Some(1);
2106 let tc = NaiveTokenCounter;
2107 let result = fetch_cross_session(&view, "query", 1000, &tc).await;
2108 assert!(result.is_err());
2109 }
2110
2111 #[tokio::test]
2114 async fn fetch_reasoning_strategies_returns_message_and_marks_used() {
2115 let mock = Arc::new(MockMemoryBackend {
2116 reasoning_strategies: vec![
2117 MemReasoningStrategy {
2118 id: "strat-1".to_string(),
2119 outcome: "success".to_string(),
2120 summary: "break the problem into small steps".to_string(),
2121 },
2122 MemReasoningStrategy {
2123 id: "strat-2".to_string(),
2124 outcome: "success".to_string(),
2125 summary: "use tracing spans for debugging".to_string(),
2126 },
2127 ],
2128 ..Default::default()
2129 });
2130 let marked_ids = Arc::clone(&mock);
2131 let mut view = empty_view();
2132 view.memory = Some(mock);
2133 view.reasoning_config.enabled = true;
2134 view.reasoning_config.context_budget_tokens = 1000;
2135 let tc = NaiveTokenCounter;
2136 let (result, handle) = fetch_reasoning_strategies(&view, "debug", 1000, 5, &tc)
2137 .await
2138 .unwrap();
2139 assert!(result.is_some());
2140 let msg = result.unwrap();
2141 assert!(msg.content.starts_with(crate::slot::REASONING_PREFIX));
2142 assert!(msg.content.contains("break the problem"));
2143
2144 if let Some(h) = handle {
2146 h.await.expect("mark_reasoning_used task panicked");
2147 }
2148
2149 let ids = marked_ids.marked_ids.lock().expect("marked_ids poisoned");
2150 assert!(
2151 ids.contains(&"strat-1".to_string()),
2152 "expected strat-1 marked"
2153 );
2154 assert!(
2155 ids.contains(&"strat-2".to_string()),
2156 "expected strat-2 marked"
2157 );
2158 }
2159
2160 #[tokio::test]
2161 async fn fetch_reasoning_strategies_propagates_error() {
2162 let mock = MockMemoryBackend::with_fail_on("retrieve_reasoning_strategies");
2163 let mut view = mock_view(mock);
2164 view.reasoning_config.enabled = true;
2165 let tc = NaiveTokenCounter;
2166 let result = fetch_reasoning_strategies(&view, "query", 1000, 3, &tc).await;
2167 assert!(result.is_err());
2168 }
2169
2170 #[tokio::test]
2173 async fn fetch_semantic_recall_skips_skipped_and_stopped_messages() {
2174 let mock = MockMemoryBackend {
2175 recalled: vec![
2176 MemRecalledMessage {
2177 role: "user".to_string(),
2178 content: "[skipped] some content".to_string(),
2179 score: 0.95,
2180 },
2181 MemRecalledMessage {
2182 role: "user".to_string(),
2183 content: "[stopped] other content".to_string(),
2184 score: 0.90,
2185 },
2186 MemRecalledMessage {
2187 role: "user".to_string(),
2188 content: "valid content to recall".to_string(),
2189 score: 0.85,
2190 },
2191 ],
2192 ..Default::default()
2193 };
2194 let mut view = mock_view(mock);
2195 view.recall_limit = 10;
2196 let tc = NaiveTokenCounter;
2197 let (msg, _) = fetch_semantic_recall(&view, "query", 1000, &tc, None)
2198 .await
2199 .unwrap();
2200 assert!(msg.is_some());
2201 let msg = msg.unwrap();
2202 let full_text = msg.parts.iter().find_map(|p| {
2203 if let zeph_llm::provider::MessagePart::Recall { text } = p {
2204 Some(text.clone())
2205 } else {
2206 None
2207 }
2208 });
2209 let text = full_text.unwrap_or_default();
2210 assert!(
2211 !text.contains("[skipped]"),
2212 "skipped messages must be excluded"
2213 );
2214 assert!(
2215 !text.contains("[stopped]"),
2216 "stopped messages must be excluded"
2217 );
2218 assert!(
2219 text.contains("valid content to recall"),
2220 "valid messages must be included"
2221 );
2222 }
2223
2224 #[tokio::test]
2225 async fn fetch_cross_session_filters_below_threshold() {
2226 let mock = MockMemoryBackend {
2227 session_summaries: vec![
2228 MemSessionSummary {
2229 summary_text: "high relevance session".to_string(),
2230 score: 0.9,
2231 },
2232 MemSessionSummary {
2233 summary_text: "low relevance session".to_string(),
2234 score: 0.2,
2235 },
2236 ],
2237 ..Default::default()
2238 };
2239 let mut view = mock_view(mock);
2240 view.conversation_id = Some(1);
2241 view.cross_session_score_threshold = 0.5;
2242 let tc = NaiveTokenCounter;
2243 let result = fetch_cross_session(&view, "query", 1000, &tc)
2244 .await
2245 .unwrap();
2246 assert!(result.is_some());
2247 let msg = result.unwrap();
2248 let text = msg
2249 .parts
2250 .iter()
2251 .find_map(|p| {
2252 if let zeph_llm::provider::MessagePart::CrossSession { text } = p {
2253 Some(text.clone())
2254 } else {
2255 None
2256 }
2257 })
2258 .unwrap_or_default();
2259 assert!(
2260 text.contains("high relevance"),
2261 "high score must be included"
2262 );
2263 assert!(
2264 !text.contains("low relevance"),
2265 "low score must be filtered out"
2266 );
2267 }
2268
2269 #[tokio::test]
2270 async fn fetch_document_rag_skips_empty_chunks() {
2271 let mock = MockMemoryBackend {
2272 document_chunks: vec![
2273 MemDocumentChunk {
2274 text: String::new(),
2275 }, MemDocumentChunk {
2277 text: "real content here".to_string(),
2278 },
2279 ],
2280 ..Default::default()
2281 };
2282 let mut view = mock_view(mock);
2283 view.document_config.rag_enabled = true;
2284 let tc = NaiveTokenCounter;
2285 let result = fetch_document_rag(&view, "query", 1000, &tc).await.unwrap();
2286 assert!(result.is_some());
2287 let msg = result.unwrap();
2288 assert!(msg.content.contains("real content here"));
2289 assert!(!msg.content.contains("\n\n\n"));
2291 }
2292
2293 #[tokio::test]
2294 async fn fetch_graph_facts_sanitizes_injection_payloads() {
2295 let mock = MockMemoryBackend {
2297 graph_facts: vec![zeph_common::memory::MemGraphFact {
2298 fact: "fact with <script>alert(1)</script> and\nnewline".to_string(),
2299 confidence: 0.8,
2300 activation_score: None,
2301 neighbors: vec![],
2302 provenance_snippet: None,
2303 }],
2304 ..Default::default()
2305 };
2306 let mut view = mock_view(mock);
2307 view.graph_config.enabled = true;
2308 view.graph_config.spreading_activation.recall_timeout_ms = 5000;
2309 let tc = NaiveTokenCounter;
2310 let result = fetch_graph_facts(&view, "test", 1000, &tc).await.unwrap();
2311 assert!(result.is_some());
2312 let msg = result.unwrap();
2313 assert!(
2314 !msg.content.contains('<'),
2315 "angle brackets must be sanitized"
2316 );
2317 assert!(
2320 !msg.content.contains("\n\n"),
2321 "embedded newlines must be sanitized, no double-newline sequences expected"
2322 );
2323 }
2324
2325 #[tokio::test]
2326 async fn fetch_reasoning_strategies_sanitizes_injection_payloads() {
2327 let mock = MockMemoryBackend {
2329 reasoning_strategies: vec![MemReasoningStrategy {
2330 id: "s1".to_string(),
2331 outcome: "success".to_string(),
2332 summary: "strategy with <b>bold</b> and\nnewline".to_string(),
2333 }],
2334 ..Default::default()
2335 };
2336 let mut view = mock_view(mock);
2337 view.reasoning_config.enabled = true;
2338 let tc = NaiveTokenCounter;
2339 let (result, _handle) = fetch_reasoning_strategies(&view, "query", 1000, 3, &tc)
2340 .await
2341 .unwrap();
2342 assert!(result.is_some());
2343 let msg = result.unwrap();
2344 assert!(
2345 !msg.content.contains('<'),
2346 "angle brackets must be sanitized in strategy summaries"
2347 );
2348 }
2349
2350 #[tokio::test]
2353 async fn fetch_persona_facts_truncates_at_budget() {
2354 let tc = NaiveTokenCounter;
2355 let first_line = "[pref] brief\n";
2357 let budget = tc.count_tokens(crate::slot::PERSONA_PREFIX) + tc.count_tokens(first_line);
2358 let mock = MockMemoryBackend {
2359 persona_facts: vec![
2360 MemPersonaFact {
2361 category: "pref".to_string(),
2362 content: "brief".to_string(),
2363 },
2364 MemPersonaFact {
2365 category: "lang".to_string(),
2366 content: "english".to_string(),
2367 },
2368 ],
2369 ..Default::default()
2370 };
2371 let mut view = mock_view(mock);
2372 view.persona_config.enabled = true;
2373 let result = fetch_persona_facts(&view, budget, &tc).await.unwrap();
2374 let msg = result.unwrap();
2375 assert!(msg.content.contains("brief"), "first fact must be included");
2376 assert!(
2377 !msg.content.contains("english"),
2378 "second fact must be truncated by budget"
2379 );
2380 }
2381
2382 #[tokio::test]
2383 async fn fetch_semantic_recall_truncates_at_budget() {
2384 let tc = NaiveTokenCounter;
2385 let first_entry = "- [user] first message\n";
2387 let budget = tc.count_tokens(RECALL_PREFIX) + tc.count_tokens(first_entry);
2388 let mock = MockMemoryBackend {
2389 recalled: vec![
2390 MemRecalledMessage {
2391 role: "user".to_string(),
2392 content: "first message".to_string(),
2393 score: 0.95,
2394 },
2395 MemRecalledMessage {
2396 role: "user".to_string(),
2397 content: "second message that should be truncated".to_string(),
2398 score: 0.80,
2399 },
2400 ],
2401 ..Default::default()
2402 };
2403 let mut view = mock_view(mock);
2404 view.recall_limit = 10;
2405 let (msg, _) = fetch_semantic_recall(&view, "query", budget, &tc, None)
2406 .await
2407 .unwrap();
2408 assert!(msg.is_some());
2409 let text = msg
2410 .unwrap()
2411 .parts
2412 .iter()
2413 .find_map(|p| {
2414 if let zeph_llm::provider::MessagePart::Recall { text } = p {
2415 Some(text.clone())
2416 } else {
2417 None
2418 }
2419 })
2420 .unwrap_or_default();
2421 assert!(
2422 text.contains("first message"),
2423 "first entry must be included"
2424 );
2425 assert!(
2426 !text.contains("second message"),
2427 "second entry must be truncated by budget"
2428 );
2429 }
2430
2431 #[tokio::test]
2434 async fn fetch_graph_facts_sanitizes_provenance_snippet() {
2435 use zeph_common::memory::MemGraphNeighbor;
2436 let mock = MockMemoryBackend {
2437 graph_facts: vec![zeph_common::memory::MemGraphFact {
2438 fact: "safe fact".to_string(),
2439 confidence: 0.9,
2440 activation_score: None,
2441 neighbors: vec![MemGraphNeighbor {
2442 fact: "neighbor".to_string(),
2443 confidence: 0.7,
2444 }],
2445 provenance_snippet: Some("source with <injection>\nand newline".to_string()),
2446 }],
2447 ..Default::default()
2448 };
2449 let mut view = mock_view(mock);
2450 view.graph_config.enabled = true;
2451 view.graph_config.spreading_activation.recall_timeout_ms = 5000;
2452 let tc = NaiveTokenCounter;
2453 let result = fetch_graph_facts(&view, "test", 1000, &tc).await.unwrap();
2454 assert!(result.is_some());
2455 let msg = result.unwrap();
2456 assert!(
2457 !msg.content.contains('<'),
2458 "angle brackets in provenance_snippet must be sanitized"
2459 );
2460 assert!(
2461 !msg.content.contains("\n\n"),
2462 "newlines in provenance_snippet must be sanitized"
2463 );
2464 assert!(
2465 msg.content.contains("[source:"),
2466 "provenance snippet must be rendered"
2467 );
2468 }
2469
2470 #[tokio::test(start_paused = true)]
2478 async fn fetch_persona_facts_degrades_to_empty_on_timeout() {
2479 let mock = MockMemoryBackend {
2480 persona_facts: vec![MemPersonaFact {
2481 category: "pref".to_string(),
2482 content: "would have been returned".to_string(),
2483 }],
2484 delay: Some(std::time::Duration::from_millis(
2485 MEMORY_FETCH_TIMEOUT_MS + 1000,
2486 )),
2487 ..Default::default()
2488 };
2489 let mut view = mock_view(mock);
2490 view.persona_config.enabled = true;
2491 let tc = NaiveTokenCounter;
2492 let result = fetch_persona_facts(&view, 1000, &tc).await;
2493 assert!(
2494 result.is_ok(),
2495 "timeout must degrade gracefully, not propagate as an error: {result:?}"
2496 );
2497 assert!(
2498 result.unwrap().is_none(),
2499 "timed-out fetch must yield no message, not the stale backend data"
2500 );
2501 }
2502
2503 #[tokio::test(start_paused = true)]
2504 async fn fetch_semantic_recall_degrades_to_empty_on_timeout() {
2505 let mock = MockMemoryBackend {
2506 recalled: vec![MemRecalledMessage {
2507 role: "user".to_string(),
2508 content: "would have been returned".to_string(),
2509 score: 0.95,
2510 }],
2511 delay: Some(std::time::Duration::from_millis(
2512 MEMORY_FETCH_TIMEOUT_MS + 1000,
2513 )),
2514 ..Default::default()
2515 };
2516 let mut view = mock_view(mock);
2517 view.recall_limit = 10;
2518 let tc = NaiveTokenCounter;
2519 let result = fetch_semantic_recall(&view, "query", 1000, &tc, None).await;
2520 assert!(
2521 result.is_ok(),
2522 "timeout must degrade gracefully, not propagate as an error: {result:?}"
2523 );
2524 let (msg, score) = result.unwrap();
2525 assert!(msg.is_none(), "timed-out recall must yield no message");
2526 assert!(score.is_none(), "timed-out recall must yield no score");
2527 }
2528
2529 #[test]
2532 fn append_budgeted_lines_empty_input_returns_none() {
2533 let tc = NaiveTokenCounter;
2534 let result = append_budgeted_lines("prefix\n", std::iter::empty(), 1000, &tc);
2535 assert!(result.is_none());
2536 }
2537
2538 #[test]
2539 fn append_budgeted_lines_all_items_fit() {
2540 let tc = NaiveTokenCounter;
2541 let lines = vec![
2542 "one\n".to_string(),
2543 "two\n".to_string(),
2544 "three\n".to_string(),
2545 ];
2546 let result = append_budgeted_lines("prefix\n", lines.into_iter(), 1000, &tc).unwrap();
2547 assert!(result.starts_with("prefix\n"));
2548 assert!(result.contains("one"));
2549 assert!(result.contains("two"));
2550 assert!(result.contains("three"));
2551 }
2552
2553 #[test]
2554 fn append_budgeted_lines_truncates_at_budget() {
2555 let tc = NaiveTokenCounter;
2556 let prefix = "prefix\n";
2557 let first = "one\n";
2558 let budget = tc.count_tokens(prefix) + tc.count_tokens(first);
2560 let lines = vec![first.to_string(), "two extra words here\n".to_string()];
2561 let result = append_budgeted_lines(prefix, lines.into_iter(), budget, &tc).unwrap();
2562 assert!(result.contains("one"), "first line must fit in budget");
2563 assert!(
2564 !result.contains("two extra words"),
2565 "second line must be truncated by budget"
2566 );
2567 }
2568
2569 #[test]
2570 fn append_budgeted_lines_zero_budget_returns_none() {
2571 let tc = NaiveTokenCounter;
2572 let lines = vec!["one\n".to_string()];
2573 let result = append_budgeted_lines("prefix\n", lines.into_iter(), 0, &tc);
2574 assert!(result.is_none(), "no line can fit within a zero budget");
2575 }
2576}