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;
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()
}