Skip to main content

mempal_runtime/
knowledge_card_retrieval.rs

1#![warn(clippy::all)]
2
3use std::collections::BTreeMap;
4use std::path::PathBuf;
5
6use serde::Serialize;
7use thiserror::Error;
8
9use crate::core::{
10    anchor,
11    db::{Database, DbError},
12    types::{
13        AnchorKind, KnowledgeCard, KnowledgeEvidenceLink, KnowledgeEvidenceRole, KnowledgeStatus,
14        MemoryDomain, MemoryKind, RouteDecision,
15    },
16};
17use crate::embed::{EmbedError, Embedder};
18use crate::search::{SearchError, SearchFilters, SearchOptions, search_with_vector_options};
19
20pub type Result<T> = std::result::Result<T, KnowledgeCardRetrievalError>;
21
22#[derive(Debug, Error)]
23pub enum KnowledgeCardRetrievalError {
24    #[error("failed to derive retrieval anchors")]
25    DeriveAnchor(#[from] anchor::AnchorError),
26    #[error("failed to embed card retrieval query")]
27    EmbedQuery(#[source] EmbedError),
28    #[error("embedder returned no card retrieval query vector")]
29    MissingQueryVector,
30    #[error("failed to search linked evidence")]
31    SearchEvidence(#[source] SearchError),
32    #[error("failed to load card retrieval metadata")]
33    LoadMetadata(#[source] DbError),
34}
35
36#[derive(Debug, Clone)]
37pub struct KnowledgeCardRetrievalRequest {
38    pub query: String,
39    pub domain: MemoryDomain,
40    pub field: String,
41    pub cwd: PathBuf,
42    pub top_k: usize,
43    pub evidence_top_k: usize,
44}
45
46#[derive(Debug, Clone, Serialize)]
47pub struct RetrievedKnowledgeCard {
48    pub card: KnowledgeCard,
49    pub evidence_citations: Vec<RetrievedEvidenceCitation>,
50    pub score: f32,
51}
52
53#[derive(Debug, Clone, Serialize)]
54pub struct RetrievedEvidenceCitation {
55    pub evidence_drawer_id: String,
56    pub role: KnowledgeEvidenceRole,
57    pub source_file: String,
58    pub score: f32,
59}
60
61#[derive(Debug, Clone)]
62struct AnchorCandidate {
63    anchor_kind: AnchorKind,
64    anchor_id: String,
65    domain: MemoryDomain,
66}
67
68pub async fn retrieve_knowledge_cards<E: Embedder + ?Sized>(
69    db: &Database,
70    embedder: &E,
71    request: KnowledgeCardRetrievalRequest,
72) -> Result<Vec<RetrievedKnowledgeCard>> {
73    if request.top_k == 0 {
74        return Ok(Vec::new());
75    }
76    let query_vector = embedder
77        .embed(&[request.query.as_str()])
78        .await
79        .map_err(KnowledgeCardRetrievalError::EmbedQuery)?
80        .into_iter()
81        .next()
82        .ok_or(KnowledgeCardRetrievalError::MissingQueryVector)?;
83    retrieve_knowledge_cards_with_vector(db, request, &query_vector)
84}
85
86pub fn retrieve_knowledge_cards_with_vector(
87    db: &Database,
88    request: KnowledgeCardRetrievalRequest,
89    query_vector: &[f32],
90) -> Result<Vec<RetrievedKnowledgeCard>> {
91    if request.top_k == 0 {
92        return Ok(Vec::new());
93    }
94
95    let mut by_card = BTreeMap::<String, RetrievedKnowledgeCard>::new();
96    let route = RouteDecision {
97        wing: None,
98        room: None,
99        confidence: 0.0,
100        reason: "knowledge card linked-evidence retrieval".to_string(),
101    };
102
103    for anchor in retrieval_anchors(&request)? {
104        let filters = SearchFilters {
105            memory_kind: Some(memory_kind_slug(&MemoryKind::Evidence).to_string()),
106            domain: Some(domain_slug(&anchor.domain).to_string()),
107            field: Some(request.field.clone()),
108            tier: None,
109            status: None,
110            anchor_kind: Some(anchor_kind_slug(&anchor.anchor_kind).to_string()),
111        };
112        let evidence_results = search_with_vector_options(
113            db,
114            &request.query,
115            query_vector,
116            route.clone(),
117            SearchOptions {
118                filters,
119                with_neighbors: false,
120            },
121            request.evidence_top_k.max(request.top_k),
122        )
123        .map_err(KnowledgeCardRetrievalError::SearchEvidence)?;
124
125        for evidence in evidence_results {
126            if evidence.anchor_id != anchor.anchor_id {
127                continue;
128            }
129            let links = db
130                .knowledge_evidence_links_for_drawer(&evidence.drawer_id)
131                .map_err(KnowledgeCardRetrievalError::LoadMetadata)?;
132            for link in links {
133                let Some(card) = db
134                    .get_knowledge_card(&link.card_id)
135                    .map_err(KnowledgeCardRetrievalError::LoadMetadata)?
136                else {
137                    continue;
138                };
139                if !card_is_retrievable(&card, &request, &anchor) {
140                    continue;
141                }
142                let citation =
143                    citation_from_link(&link, &evidence.source_file, evidence.similarity);
144                match by_card.get_mut(&card.id) {
145                    Some(existing) => {
146                        if citation.score > existing.score {
147                            existing.score = citation.score;
148                        }
149                        existing.evidence_citations.push(citation);
150                    }
151                    None => {
152                        by_card.insert(
153                            card.id.clone(),
154                            RetrievedKnowledgeCard {
155                                card,
156                                score: citation.score,
157                                evidence_citations: vec![citation],
158                            },
159                        );
160                    }
161                }
162            }
163        }
164    }
165
166    let mut results = by_card.into_values().collect::<Vec<_>>();
167    results.sort_by(|left, right| {
168        right
169            .score
170            .partial_cmp(&left.score)
171            .unwrap_or(std::cmp::Ordering::Equal)
172            .then_with(|| left.card.id.cmp(&right.card.id))
173    });
174    results.truncate(request.top_k);
175    Ok(results)
176}
177
178fn retrieval_anchors(request: &KnowledgeCardRetrievalRequest) -> Result<Vec<AnchorCandidate>> {
179    let derived = anchor::derive_anchor_from_cwd(Some(&request.cwd))?;
180    let mut anchors = Vec::new();
181    anchors.push(AnchorCandidate {
182        anchor_kind: AnchorKind::Worktree,
183        anchor_id: derived.anchor_id,
184        domain: request.domain.clone(),
185    });
186
187    let repo_anchor_id = derived
188        .parent_anchor_id
189        .unwrap_or_else(|| anchor::LEGACY_REPO_ANCHOR_ID.to_string());
190    anchors.push(AnchorCandidate {
191        anchor_kind: AnchorKind::Repo,
192        anchor_id: repo_anchor_id,
193        domain: request.domain.clone(),
194    });
195    anchors.push(AnchorCandidate {
196        anchor_kind: AnchorKind::Repo,
197        anchor_id: anchor::LEGACY_REPO_ANCHOR_ID.to_string(),
198        domain: request.domain.clone(),
199    });
200    anchors.push(AnchorCandidate {
201        anchor_kind: AnchorKind::Global,
202        anchor_id: "global://default".to_string(),
203        domain: MemoryDomain::Global,
204    });
205
206    let mut seen = BTreeMap::new();
207    Ok(anchors
208        .into_iter()
209        .filter(|anchor| {
210            seen.insert(
211                (
212                    anchor_kind_slug(&anchor.anchor_kind).to_string(),
213                    anchor.anchor_id.clone(),
214                ),
215                (),
216            )
217            .is_none()
218        })
219        .collect())
220}
221
222fn card_is_retrievable(
223    card: &KnowledgeCard,
224    request: &KnowledgeCardRetrievalRequest,
225    anchor: &AnchorCandidate,
226) -> bool {
227    matches!(
228        card.status,
229        KnowledgeStatus::Canonical | KnowledgeStatus::Promoted
230    ) && card.domain == anchor.domain
231        && card.field == request.field
232        && card.anchor_kind == anchor.anchor_kind
233        && card.anchor_id == anchor.anchor_id
234}
235
236fn citation_from_link(
237    link: &KnowledgeEvidenceLink,
238    source_file: &str,
239    score: f32,
240) -> RetrievedEvidenceCitation {
241    RetrievedEvidenceCitation {
242        evidence_drawer_id: link.evidence_drawer_id.clone(),
243        role: link.role.clone(),
244        source_file: source_file.to_string(),
245        score,
246    }
247}
248
249fn memory_kind_slug(value: &MemoryKind) -> &'static str {
250    match value {
251        MemoryKind::Evidence => "evidence",
252        MemoryKind::Knowledge => "knowledge",
253    }
254}
255
256fn domain_slug(value: &MemoryDomain) -> &'static str {
257    match value {
258        MemoryDomain::Project => "project",
259        MemoryDomain::Agent => "agent",
260        MemoryDomain::Skill => "skill",
261        MemoryDomain::Global => "global",
262    }
263}
264
265fn anchor_kind_slug(value: &AnchorKind) -> &'static str {
266    match value {
267        AnchorKind::Global => "global",
268        AnchorKind::Repo => "repo",
269        AnchorKind::Worktree => "worktree",
270    }
271}