mod merge;
pub use merge::{
clear_memory_graph_bindings, count_relationships_by_relation, create_or_fetch_relationship,
delete_entities_by_ids, delete_relationship_by_id, delete_relationships_by_relation,
find_entity_id, find_orphan_entity_ids, find_relationship, increment_degree,
link_memory_entity, link_memory_relationship, list_entity_names_by_relation,
recalculate_degree, unlink_memory_entity, RelationshipRow,
};
use crate::embedder::f32_to_bytes;
use crate::entity_type::EntityType;
use crate::errors::AppError;
use crate::parsers::normalize_entity_name;
use crate::storage::utils::with_busy_retry;
use rusqlite::{params, Connection};
use serde::{Deserialize, Serialize};
#[derive(Debug, Serialize, Deserialize, Clone)]
#[serde(deny_unknown_fields)]
pub struct NewEntity {
pub name: String,
#[serde(alias = "type")]
pub entity_type: EntityType,
pub description: Option<String>,
}
#[derive(Debug, Serialize, Deserialize, Clone)]
#[serde(deny_unknown_fields)]
pub struct NewRelationship {
#[serde(alias = "from")]
pub source: String,
#[serde(alias = "to")]
pub target: String,
#[serde(alias = "type")]
pub relation: String,
#[serde(alias = "weight", default = "default_relationship_strength")]
pub strength: f64,
pub description: Option<String>,
}
fn default_relationship_strength() -> f64 {
crate::constants::DEFAULT_RELATION_WEIGHT
}
pub fn validate_entity_name(name: &str) -> Result<(), AppError> {
if name.len() < 2 {
return Err(AppError::Validation(crate::i18n::validation::entity_name_too_short(name)));
}
if name.contains('\n') || name.contains('\r') {
return Err(AppError::Validation(
"entity name must not contain newline characters".to_string(),
));
}
if name.chars().all(|c| c.is_ascii_digit()) {
return Err(AppError::Validation(crate::i18n::validation::entity_name_purely_numeric(name)));
}
if name.len() <= 4
&& name
.chars()
.all(|c| c.is_ascii_uppercase() || c == '_' || c == '-')
{
return Err(AppError::Validation(crate::i18n::validation::entity_name_all_caps_noise(name)));
}
Ok(())
}
#[derive(Debug, Clone)]
pub struct FuzzyEntityMatch {
pub id: i64,
pub name: String,
pub score: f64,
}
pub fn entity_name_similarity(query: &str, name: &str) -> f64 {
let q = query.trim().to_ascii_lowercase();
let n = name.trim().to_ascii_lowercase();
if q.is_empty() || n.is_empty() {
return 0.0;
}
if q == n {
return 1.0;
}
if n.starts_with(&q) {
let rest = &n[q.len()..];
if rest.is_empty()
|| rest.starts_with('-')
|| rest.starts_with('_')
|| rest.starts_with(' ')
{
return 0.95;
}
return 0.88;
}
if q.starts_with(&n) && n.len() >= 3 {
return 0.80;
}
let first_token = n
.split(|c: char| c == '-' || c == '_' || c.is_whitespace())
.next()
.unwrap_or(n.as_str());
if first_token == q {
return 0.92;
}
if n.contains(&q) && q.len() >= 3 {
return 0.82;
}
rapidfuzz::distance::jaro_winkler::normalized_similarity(q.chars(), n.chars())
}
pub fn suggest_entity_names(
conn: &Connection,
namespace: &str,
query: &str,
limit: usize,
min_score: f64,
) -> Result<Vec<FuzzyEntityMatch>, AppError> {
let entities = list_entities(conn, Some(namespace))?;
let mut scored: Vec<FuzzyEntityMatch> = entities
.into_iter()
.filter_map(|e| {
let score = entity_name_similarity(query, &e.name);
if score >= min_score {
Some(FuzzyEntityMatch {
id: e.id,
name: e.name,
score,
})
} else {
None
}
})
.collect();
scored.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
.then_with(|| a.name.cmp(&b.name))
});
scored.truncate(limit.max(1));
Ok(scored)
}
pub fn resolve_entity_fuzzy(
conn: &Connection,
namespace: &str,
name: &str,
auto_fuzzy: bool,
) -> Result<Option<(i64, String, bool)>, AppError> {
if let Some(id) = find_entity_id(conn, namespace, name)? {
return Ok(Some((id, name.to_string(), false)));
}
let normalized = crate::parsers::normalize_entity_name(name);
if normalized != name {
if let Some(id) = find_entity_id(conn, namespace, &normalized)? {
return Ok(Some((id, normalized, false)));
}
}
if !auto_fuzzy {
return Ok(None);
}
let suggestions = suggest_entity_names(conn, namespace, name, 5, 0.75)?;
if suggestions.is_empty() {
return Ok(None);
}
let top = &suggestions[0];
let clear_winner =
top.score >= 0.90 && (suggestions.len() == 1 || top.score - suggestions[1].score >= 0.05);
let single_strong = suggestions.len() == 1 && top.score >= 0.85;
if clear_winner || single_strong {
tracing::warn!(
target: "entities",
query = %name,
resolved = %top.name,
score = top.score,
"fuzzy entity resolution: exact match failed; using best candidate"
);
return Ok(Some((top.id, top.name.clone(), true)));
}
Ok(None)
}
pub fn entity_not_found_with_suggestions(
conn: &Connection,
namespace: &str,
name: &str,
) -> AppError {
let suggestions = suggest_entity_names(conn, namespace, name, 5, 0.70).unwrap_or_default();
if suggestions.is_empty() {
return AppError::NotFound(format!(
"entity '{name}' not found in namespace '{namespace}'"
));
}
let list: Vec<String> = suggestions
.iter()
.map(|s| format!("{} (score={:.2})", s.name, s.score))
.collect();
AppError::NotFound(format!(
"entity '{name}' not found in namespace '{namespace}'. Did you mean: {}? \
Re-run with --fuzzy to auto-resolve a clear match, or pass the canonical name.",
list.join(", ")
))
}
pub fn upsert_entity(conn: &Connection, namespace: &str, e: &NewEntity) -> Result<i64, AppError> {
validate_entity_name(&e.name)?;
let normalized_name = normalize_entity_name(&e.name);
if normalized_name.chars().count() < 2 {
return Err(AppError::Validation(crate::i18n::validation::entity_name_normalizes_too_short(&e.name, &normalized_name)));
}
conn.execute(
"INSERT INTO entities (namespace, name, type, description)
VALUES (?1, ?2, ?3, ?4)
ON CONFLICT(namespace, name) DO UPDATE SET
type = excluded.type,
description = COALESCE(excluded.description, entities.description),
updated_at = unixepoch()",
params![namespace, normalized_name, e.entity_type, e.description],
)?;
let id: i64 = conn.query_row(
"SELECT id FROM entities WHERE namespace = ?1 AND name = ?2",
params![namespace, normalized_name],
|r| r.get(0),
)?;
Ok(id)
}
pub fn upsert_entity_vec(
conn: &Connection,
entity_id: i64,
namespace: &str,
_entity_type: EntityType,
embedding: &[f32],
_name: &str,
) -> Result<(), AppError> {
if embedding.is_empty() {
tracing::debug!(
entity_id,
"empty entity embedding: skipping entity_embeddings row (backfill via enrich re-embed --target entities)"
);
return Ok(());
}
let embedding_bytes = f32_to_bytes(embedding);
with_busy_retry(|| {
conn.execute(
"DELETE FROM entity_embeddings WHERE entity_id = ?1",
params![entity_id],
)?;
conn.execute(
"INSERT INTO entity_embeddings(entity_id, namespace, embedding, source, model, dim)
VALUES (?1, ?2, ?3, ?4, ?5, ?6)",
params![
entity_id,
namespace,
&embedding_bytes,
"llm-headless",
crate::constants::SQLITE_GRAPHRAG_VERSION,
crate::constants::embedding_dim() as i64,
],
)?;
Ok(())
})
}
pub fn upsert_relationship(
conn: &Connection,
namespace: &str,
source_id: i64,
target_id: i64,
rel: &NewRelationship,
) -> Result<i64, AppError> {
conn.execute(
"INSERT INTO relationships (namespace, source_id, target_id, relation, weight, description)
VALUES (?1, ?2, ?3, ?4, ?5, ?6)
ON CONFLICT(source_id, target_id, relation) DO UPDATE SET
weight = excluded.weight,
description = COALESCE(excluded.description, relationships.description)",
params![
namespace,
source_id,
target_id,
rel.relation,
rel.strength,
rel.description
],
)?;
let id: i64 = conn.query_row(
"SELECT id FROM relationships WHERE source_id=?1 AND target_id=?2 AND relation=?3",
params![source_id, target_id, rel.relation],
|r| r.get(0),
)?;
Ok(id)
}
#[derive(Debug, Serialize, Clone)]
pub struct EntityNode {
pub id: i64,
pub name: String,
pub namespace: String,
pub kind: String,
}
pub fn list_entities(
conn: &Connection,
namespace: Option<&str>,
) -> Result<Vec<EntityNode>, AppError> {
if let Some(ns) = namespace {
let mut stmt = conn.prepare_cached(
"SELECT id, name, namespace, type FROM entities WHERE namespace = ?1 ORDER BY id",
)?;
let rows = stmt
.query_map(params![ns], |r| {
Ok(EntityNode {
id: r.get(0)?,
name: r.get(1)?,
namespace: r.get(2)?,
kind: r.get(3)?,
})
})?
.collect::<Result<Vec<_>, _>>()?;
Ok(rows)
} else {
let mut stmt = conn.prepare_cached(
"SELECT id, name, namespace, type FROM entities ORDER BY namespace, id",
)?;
let rows = stmt
.query_map([], |r| {
Ok(EntityNode {
id: r.get(0)?,
name: r.get(1)?,
namespace: r.get(2)?,
kind: r.get(3)?,
})
})?
.collect::<Result<Vec<_>, _>>()?;
Ok(rows)
}
}
pub fn list_relationships_by_namespace(
conn: &Connection,
namespace: Option<&str>,
) -> Result<Vec<RelationshipRow>, AppError> {
if let Some(ns) = namespace {
let mut stmt = conn.prepare_cached(
"SELECT r.id, r.namespace, r.source_id, r.target_id, r.relation, r.weight, r.description
FROM relationships r
JOIN entities se ON se.id = r.source_id AND se.namespace = ?1
JOIN entities te ON te.id = r.target_id AND te.namespace = ?1
ORDER BY r.id",
)?;
let rows = stmt
.query_map(params![ns], |r| {
Ok(RelationshipRow {
id: r.get(0)?,
namespace: r.get(1)?,
source_id: r.get(2)?,
target_id: r.get(3)?,
relation: r.get(4)?,
weight: r.get(5)?,
description: r.get(6)?,
})
})?
.collect::<Result<Vec<_>, _>>()?;
Ok(rows)
} else {
let mut stmt = conn.prepare_cached(
"SELECT id, namespace, source_id, target_id, relation, weight, description
FROM relationships ORDER BY id",
)?;
let rows = stmt
.query_map([], |r| {
Ok(RelationshipRow {
id: r.get(0)?,
namespace: r.get(1)?,
source_id: r.get(2)?,
target_id: r.get(3)?,
relation: r.get(4)?,
weight: r.get(5)?,
description: r.get(6)?,
})
})?
.collect::<Result<Vec<_>, _>>()?;
Ok(rows)
}
}
pub fn knn_search(
conn: &Connection,
embedding: &[f32],
namespace: &str,
k: usize,
) -> Result<Vec<(i64, f32)>, AppError> {
if embedding.len() != crate::constants::embedding_dim() {
return Err(AppError::Embedding(
crate::i18n::validation::embedding_knn_search_dim_mismatch(
embedding.len(),
crate::constants::embedding_dim(),
),
));
}
let mut stmt = conn.prepare_cached(
"SELECT entity_id, embedding FROM entity_embeddings WHERE namespace = ?1",
)?;
let mut scored: Vec<(i64, f32)> = stmt
.query_map(params![namespace], |r| {
let id: i64 = r.get(0)?;
let bytes: Vec<u8> = r.get(1)?;
Ok((id, bytes))
})?
.filter_map(|row| {
row.ok().and_then(|(id, bytes)| {
let stored = crate::embedder::bytes_to_f32(&bytes);
if stored.len() != embedding.len() {
return None;
}
let score = crate::similarity::cosine_similarity(embedding, &stored);
Some((id, score))
})
})
.collect();
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
scored.truncate(k);
Ok(scored)
}
#[cfg(test)]
mod tests_a;
#[cfg(test)]
mod tests_b;