Skip to main content

remem/memory/preference/
render.rs

1use anyhow::Result;
2use rusqlite::{params, Connection, OptionalExtension};
3
4use crate::memory::poisoning::{scan_instruction_pattern, InstructionPatternMatch};
5use crate::memory::Memory;
6
7use super::{
8    consolidation::{classify_preference_texts, PreferenceConsolidationKind},
9    query_global_preferences, query_project_preferences,
10};
11
12#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
13pub struct PreferenceRenderSummary {
14    pub rendered: usize,
15    pub project_rendered: usize,
16    pub global_rendered: usize,
17}
18
19#[derive(Debug, Clone, Default)]
20pub(crate) struct PreferenceRenderDetails {
21    pub summary: PreferenceRenderSummary,
22    pub rendered_ids: Vec<i64>,
23}
24
25#[derive(Debug, Clone, Copy, PartialEq, Eq)]
26enum PreferenceSource {
27    Project,
28    Global,
29}
30
31pub fn dedup_with_claude_md(prefs: &[Memory], cwd: &str) -> Vec<usize> {
32    let claude_md_path = std::path::Path::new(cwd).join("CLAUDE.md");
33    let claude_md_content = std::fs::read_to_string(&claude_md_path).unwrap_or_default();
34
35    if claude_md_content.is_empty() {
36        return (0..prefs.len()).collect();
37    }
38
39    let claude_lower = claude_md_content.to_lowercase();
40    (0..prefs.len())
41        .filter(|&i| {
42            let title_lower = prefs[i].title.to_lowercase();
43            let search_term = title_lower
44                .strip_prefix("preference: ")
45                .unwrap_or(&title_lower);
46            !claude_lower.contains(search_term)
47        })
48        .collect()
49}
50
51pub fn render_preferences(
52    output: &mut String,
53    conn: &Connection,
54    project: &str,
55    cwd: &str,
56) -> Result<()> {
57    render_preferences_with_limits(output, conn, project, cwd, 20, 0, 1500).map(|_| ())
58}
59
60pub fn render_preferences_with_limits(
61    output: &mut String,
62    conn: &Connection,
63    project: &str,
64    cwd: &str,
65    project_limit: usize,
66    global_limit: usize,
67    char_limit: usize,
68) -> Result<usize> {
69    render_preferences_with_context_details(
70        output,
71        conn,
72        project,
73        cwd,
74        project_limit,
75        global_limit,
76        char_limit,
77    )
78    .map(|details| details.summary.rendered)
79}
80
81pub fn render_preferences_with_limits_detailed(
82    output: &mut String,
83    conn: &Connection,
84    project: &str,
85    cwd: &str,
86    project_limit: usize,
87    global_limit: usize,
88    char_limit: usize,
89) -> Result<PreferenceRenderSummary> {
90    render_preferences_with_context_details(
91        output,
92        conn,
93        project,
94        cwd,
95        project_limit,
96        global_limit,
97        char_limit,
98    )
99    .map(|details| details.summary)
100}
101
102pub(crate) fn render_preferences_with_context_details(
103    output: &mut String,
104    conn: &Connection,
105    project: &str,
106    cwd: &str,
107    project_limit: usize,
108    global_limit: usize,
109    char_limit: usize,
110) -> Result<PreferenceRenderDetails> {
111    let project_prefs = query_project_preferences(conn, project, project_limit)?;
112    let global_prefs = query_global_preferences(conn, global_limit)?;
113
114    let mut all_prefs: Vec<(Memory, PreferenceSource)> = project_prefs
115        .into_iter()
116        .map(|memory| (memory, PreferenceSource::Project))
117        .collect();
118    let project_topics: std::collections::HashSet<String> = all_prefs
119        .iter()
120        .filter_map(|(memory, _)| memory.topic_key.clone())
121        .collect();
122    for global_pref in global_prefs {
123        if let Some(ref topic_key) = global_pref.topic_key {
124            if !project_topics.contains(topic_key) {
125                all_prefs.push((global_pref, PreferenceSource::Global));
126            }
127        }
128    }
129    all_prefs = filter_unacknowledged_poisoned_preferences(conn, all_prefs)?;
130
131    if all_prefs.is_empty() {
132        return Ok(PreferenceRenderDetails::default());
133    }
134
135    let memories = all_prefs
136        .iter()
137        .map(|(memory, _)| memory.clone())
138        .collect::<Vec<_>>();
139    let keep_indices = dedup_with_claude_md(&memories, cwd);
140    if keep_indices.is_empty() {
141        return Ok(PreferenceRenderDetails::default());
142    }
143    let keep_indices = dedup_with_preference_similarity(&memories, &keep_indices);
144    if keep_indices.is_empty() {
145        return Ok(PreferenceRenderDetails::default());
146    }
147
148    output.push_str("## Your Preferences (always apply these)\n");
149    let mut total_chars = 0usize;
150    let mut summary = PreferenceRenderSummary::default();
151    let mut rendered_ids = Vec::new();
152    for &idx in &keep_indices {
153        let (pref, source) = &all_prefs[idx];
154        let text = normalize_rendered_preference_text(&pref.text);
155        let preview: String = text.chars().take(120).collect();
156        let line = if preview.chars().count() < text.chars().count() {
157            format!("- {}...\n", preview)
158        } else {
159            format!("- {text}\n")
160        };
161        let line_chars = line.chars().count();
162        if total_chars + line_chars > char_limit && total_chars > 0 {
163            break;
164        }
165        output.push_str(&line);
166        total_chars += line_chars;
167        summary.rendered += 1;
168        rendered_ids.push(pref.id);
169        match source {
170            PreferenceSource::Project => summary.project_rendered += 1,
171            PreferenceSource::Global => summary.global_rendered += 1,
172        }
173    }
174    output.push('\n');
175
176    Ok(PreferenceRenderDetails {
177        summary,
178        rendered_ids,
179    })
180}
181
182#[derive(Debug, Default)]
183struct PreferencePoisoningState {
184    acknowledged_pattern_id: Option<String>,
185    acknowledged_pattern_version: Option<i64>,
186    source_trust_class: String,
187    source_project: Option<String>,
188}
189
190fn filter_unacknowledged_poisoned_preferences(
191    conn: &Connection,
192    prefs: Vec<(Memory, PreferenceSource)>,
193) -> Result<Vec<(Memory, PreferenceSource)>> {
194    let mut kept = Vec::with_capacity(prefs.len());
195    for (memory, source) in prefs {
196        let Some(pattern_match) =
197            scan_instruction_pattern(&format!("{}\n{}", memory.title, memory.text))
198        else {
199            kept.push((memory, source));
200            continue;
201        };
202        let state = load_preference_poisoning_state(conn, memory.id)?;
203        if state.acknowledged_pattern_id.as_deref() == Some(pattern_match.pattern_id)
204            && state.acknowledged_pattern_version == Some(pattern_match.pattern_set_version)
205        {
206            kept.push((memory, source));
207            continue;
208        }
209        crate::log::error(
210            "context-poisoning",
211            &format!(
212                "dropping unacknowledged poisoned preference memory id={} pattern={}@v{}",
213                memory.id, pattern_match.pattern_id, pattern_match.pattern_set_version
214            ),
215        );
216        record_preference_injection_drop(conn, &memory, &state, pattern_match)?;
217    }
218    Ok(kept)
219}
220
221fn load_preference_poisoning_state(
222    conn: &Connection,
223    memory_id: i64,
224) -> Result<PreferencePoisoningState> {
225    Ok(conn
226        .query_row(
227            "SELECT acknowledged_pattern_id, acknowledged_pattern_version,
228                    source_trust_class, source_project
229             FROM memories WHERE id = ?1",
230            params![memory_id],
231            |row| {
232                Ok(PreferencePoisoningState {
233                    acknowledged_pattern_id: row.get(0)?,
234                    acknowledged_pattern_version: row.get(1)?,
235                    source_trust_class: row.get(2)?,
236                    source_project: row.get(3)?,
237                })
238            },
239        )
240        .optional()?
241        .unwrap_or_else(|| PreferencePoisoningState {
242            acknowledged_pattern_id: None,
243            acknowledged_pattern_version: None,
244            source_trust_class: "external_content".to_string(),
245            source_project: None,
246        }))
247}
248
249fn record_preference_injection_drop(
250    conn: &Connection,
251    memory: &Memory,
252    state: &PreferencePoisoningState,
253    pattern_match: InstructionPatternMatch,
254) -> Result<()> {
255    conn.execute(
256        "INSERT INTO memory_poisoning_injection_drops
257         (memory_id, pattern_id, pattern_version, source_trust_class, source_project,
258          title, created_at_epoch)
259         VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7)",
260        params![
261            memory.id,
262            pattern_match.pattern_id,
263            pattern_match.pattern_set_version,
264            state.source_trust_class.as_str(),
265            state.source_project.as_deref(),
266            memory.title.as_str(),
267            chrono::Utc::now().timestamp(),
268        ],
269    )?;
270    Ok(())
271}
272
273fn normalize_rendered_preference_text(text: &str) -> String {
274    text.trim()
275        .lines()
276        .filter(|line| !line.trim().is_empty())
277        .collect::<Vec<_>>()
278        .join(" ")
279}
280
281fn dedup_with_preference_similarity(prefs: &[Memory], indices: &[usize]) -> Vec<usize> {
282    let mut kept: Vec<usize> = Vec::new();
283    for &idx in indices {
284        let incoming = &prefs[idx];
285        let already_represented = kept.iter().any(|&kept_idx| {
286            let existing = &prefs[kept_idx];
287            classify_preference_texts(existing.id, &existing.text, &incoming.text).is_some_and(
288                |matched| {
289                    matches!(
290                        matched.kind,
291                        PreferenceConsolidationKind::SamePreference
292                            | PreferenceConsolidationKind::Refinement
293                    )
294                },
295            )
296        });
297        if !already_represented {
298            kept.push(idx);
299        }
300    }
301    kept
302}