goosedump 0.12.43

Browse, search, compact, and learn from coding-agent sessions
// SPDX-License-Identifier: LGPL-2.1-or-later
// Copyright (C) Jarkko Sakkinen 2026

use std::collections::{BTreeSet, HashMap};
use std::path::PathBuf;

use anyhow::Context as _;
use rusqlite::{Connection, params_from_iter};

use crate::engine::Client;

use super::super::external_id;
use super::super::types::{ClaimRelationship, MemoryError, RelationshipDirection, SourceReference};
use super::util::DISPLAY_ID_CHARS;

// Relationship hydration binds each batch twice; stay below SQLite's 999-variable default.
const HYDRATION_BATCH_SIZE: usize = 400;

#[derive(Debug)]
pub(super) struct ClaimHydration {
    pub(super) display_id: String,
    pub(super) evidence: Vec<SourceReference>,
    pub(super) relationships: Vec<ClaimRelationship>,
}

#[derive(Debug)]
struct RawRelationship {
    claim_id: String,
    direction: RelationshipDirection,
    kind: String,
    related_id: String,
    memory_type: String,
    text: String,
    status: String,
    rationale: String,
}

pub(super) fn hydrate_claims(
    conn: &Connection,
    claim_ids: &[String],
) -> anyhow::Result<HashMap<String, ClaimHydration>> {
    let claim_ids = claim_ids.iter().cloned().collect::<BTreeSet<_>>();
    if claim_ids.is_empty() {
        return Ok(HashMap::new());
    }

    let relationships = load_relationships(conn, &claim_ids)?;
    let display_ids = load_display_ids(
        conn,
        claim_ids.iter().chain(
            relationships
                .iter()
                .map(|relationship| &relationship.related_id),
        ),
    )?;
    let mut hydrated = claim_ids
        .iter()
        .map(|claim_id| {
            let display_id = display_ids
                .get(claim_id)
                .with_context(|| format!("claim '{claim_id}' disappeared during hydration"))?
                .clone();
            Ok((
                claim_id.clone(),
                ClaimHydration {
                    display_id,
                    evidence: Vec::new(),
                    relationships: Vec::new(),
                },
            ))
        })
        .collect::<anyhow::Result<HashMap<_, _>>>()?;
    load_evidence(conn, &claim_ids, &mut hydrated)?;
    for relationship in relationships {
        let display_id = display_ids
            .get(&relationship.related_id)
            .with_context(|| {
                format!(
                    "related claim '{}' disappeared during hydration",
                    relationship.related_id
                )
            })?
            .clone();
        hydrated
            .get_mut(&relationship.claim_id)
            .context("hydrated relationship claim disappeared")?
            .relationships
            .push(ClaimRelationship {
                direction: relationship.direction,
                kind: relationship.kind.parse()?,
                claim_id: external_id(&relationship.related_id),
                display_id,
                memory_type: relationship.memory_type.parse()?,
                text: relationship.text,
                status: relationship.status.parse()?,
                rationale: relationship.rationale,
            });
    }

    Ok(hydrated)
}

fn load_evidence(
    conn: &Connection,
    claim_ids: &BTreeSet<String>,
    hydrated: &mut HashMap<String, ClaimHydration>,
) -> anyhow::Result<()> {
    let claim_ids = claim_ids.iter().collect::<Vec<_>>();
    for batch in claim_ids.chunks(HYDRATION_BATCH_SIZE) {
        let placeholders = placeholders(batch.len());
        let sql = format!(
            r"SELECT claim_evidence.claim_id, evidence.provider, evidence.session_id,
                    evidence.entry_id, evidence.role, evidence.observed_at, projects.path,
                    evidence.source_path, evidence.content_hash,
                    substr(evidence.text, 1, 500), evidence.id
             FROM claim_evidence
             JOIN evidence ON evidence.id = claim_evidence.evidence_id
             JOIN projects ON projects.id = evidence.project_id
             WHERE claim_evidence.claim_id IN ({placeholders})
             ORDER BY claim_evidence.claim_id, evidence.observed_at, evidence.id"
        );
        let mut stmt = conn.prepare(&sql)?;
        let rows = stmt.query_map(params_from_iter(batch.iter().copied()), |row| {
            Ok((
                row.get::<_, String>(0)?,
                row.get::<_, String>(1)?,
                row.get::<_, String>(2)?,
                row.get::<_, String>(3)?,
                row.get::<_, String>(4)?,
                row.get::<_, i64>(5)?,
                row.get::<_, String>(6)?,
                row.get::<_, String>(7)?,
                row.get::<_, String>(8)?,
                row.get::<_, String>(9)?,
                row.get::<_, i64>(10)?,
            ))
        })?;
        for row in rows {
            let (
                claim_id,
                provider,
                session_id,
                entry_id,
                role,
                observed_at,
                project,
                source_path,
                content_hash,
                snippet,
                evidence_id,
            ) = row?;
            hydrated
                .get_mut(&claim_id)
                .context("evidence claim was not requested")?
                .evidence
                .push(SourceReference {
                    citation_id: format!("ev_{evidence_id}"),
                    provider: provider.parse::<Client>()?,
                    session_id,
                    entry_id,
                    role,
                    observed_at,
                    project: PathBuf::from(project),
                    source_path: PathBuf::from(source_path),
                    content_hash,
                    snippet,
                });
        }
    }
    Ok(())
}

fn load_relationships(
    conn: &Connection,
    claim_ids: &BTreeSet<String>,
) -> anyhow::Result<Vec<RawRelationship>> {
    let claim_ids = claim_ids.iter().collect::<Vec<_>>();
    let mut relationships = Vec::new();
    for batch in claim_ids.chunks(HYDRATION_BATCH_SIZE) {
        let placeholders = placeholders(batch.len());
        let sql = format!(
            r"SELECT relations.subject_claim_id, 'outgoing', relations.predicate, claims.id,
                    claims.memory_type, claims.statement, claims.status, relations.rationale
             FROM relations
             JOIN claims ON claims.id = relations.object_claim_id
             WHERE relations.subject_claim_id IN ({placeholders})
             UNION ALL
             SELECT relations.object_claim_id, 'incoming', relations.predicate, claims.id,
                    claims.memory_type, claims.statement, claims.status, relations.rationale
             FROM relations
             JOIN claims ON claims.id = relations.subject_claim_id
             WHERE relations.object_claim_id IN ({placeholders})
             ORDER BY 1, 2, 3, 4"
        );
        let mut stmt = conn.prepare(&sql)?;
        let rows = stmt.query_map(
            params_from_iter(batch.iter().copied().chain(batch.iter().copied())),
            |row| {
                Ok((
                    row.get::<_, String>(0)?,
                    row.get::<_, String>(1)?,
                    row.get::<_, String>(2)?,
                    row.get::<_, String>(3)?,
                    row.get::<_, String>(4)?,
                    row.get::<_, String>(5)?,
                    row.get::<_, String>(6)?,
                    row.get::<_, String>(7)?,
                ))
            },
        )?;
        for row in rows {
            let (claim_id, direction, kind, related_id, memory_type, text, status, rationale) =
                row?;
            relationships.push(RawRelationship {
                claim_id,
                direction: match direction.as_str() {
                    "outgoing" => RelationshipDirection::Outgoing,
                    "incoming" => RelationshipDirection::Incoming,
                    _ => return Err(MemoryError::InvalidRelationshipDirection.into()),
                },
                kind,
                related_id,
                memory_type,
                text,
                status,
                rationale,
            });
        }
    }
    Ok(relationships)
}

fn load_display_ids<'a>(
    conn: &Connection,
    claim_ids: impl IntoIterator<Item = &'a String>,
) -> anyhow::Result<HashMap<String, String>> {
    let claim_ids = claim_ids.into_iter().cloned().collect::<BTreeSet<_>>();
    let claim_ids = claim_ids.iter().collect::<Vec<_>>();
    let mut display_ids = HashMap::new();
    for batch in claim_ids.chunks(HYDRATION_BATCH_SIZE) {
        let placeholders = placeholders(batch.len());
        let sql = format!(
            r"SELECT current.id,
                    (SELECT max(previous.id) FROM claims AS previous
                     WHERE previous.id < current.id),
                    (SELECT min(next.id) FROM claims AS next
                     WHERE next.id > current.id)
             FROM claims AS current
             WHERE current.id IN ({placeholders})"
        );
        let mut stmt = conn.prepare(&sql)?;
        let rows = stmt.query_map(params_from_iter(batch.iter().copied()), |row| {
            Ok((
                row.get::<_, String>(0)?,
                row.get::<_, Option<String>>(1)?,
                row.get::<_, Option<String>>(2)?,
            ))
        })?;
        for row in rows {
            let (id, previous, next) = row?;
            let unique_bytes = previous
                .as_deref()
                .map_or(0, |other| shared_prefix_bytes(&id, other) + 1)
                .max(
                    next.as_deref()
                        .map_or(0, |other| shared_prefix_bytes(&id, other) + 1),
                )
                .max(DISPLAY_ID_CHARS)
                .min(id.len());
            display_ids.insert(id.clone(), format!("mem_{}", &id[..unique_bytes]));
        }
    }
    Ok(display_ids)
}

fn placeholders(count: usize) -> String {
    std::iter::repeat_n("?", count)
        .collect::<Vec<_>>()
        .join(",")
}

fn shared_prefix_bytes(left: &str, right: &str) -> usize {
    left.bytes()
        .zip(right.bytes())
        .take_while(|(left, right)| left == right)
        .count()
}