use anyhow::Result;
use rusqlite::{types::ToSql, Connection};
use crate::retrieval::temporal::types::{TemporalConstraint, TemporalField};
pub fn search_by_time(
conn: &Connection,
constraint: &TemporalConstraint,
project: Option<&str>,
limit: i64,
) -> Result<Vec<i64>> {
search_by_time_filtered(conn, constraint, project, None, None, limit, false)
}
pub fn search_by_time_filtered(
conn: &Connection,
constraint: &TemporalConstraint,
project: Option<&str>,
memory_type: Option<&str>,
branch: Option<&str>,
limit: i64,
include_inactive: bool,
) -> Result<Vec<i64>> {
let mut ids = Vec::new();
let has_memory_facts = table_exists(conn, "memory_facts")?;
let has_memory_fact_invalidations =
has_memory_facts && crate::memory::facts::invalidated_at_epoch_available(conn)?;
let event_time_expr = if column_exists(conn, "memories", "reference_time_epoch")? {
"COALESCE(reference_time_epoch, created_at_epoch)"
} else {
"created_at_epoch"
};
let (temporal_condition, order_epoch) = temporal_sql(
constraint.field,
has_memory_facts,
has_memory_fact_invalidations,
event_time_expr,
);
let mut conditions = vec![temporal_condition];
let mut params_vec: Vec<Box<dyn ToSql>> = vec![
Box::new(constraint.start_epoch),
Box::new(constraint.end_epoch),
];
let mut idx = 3;
conditions.push(crate::memory::memory_current_filter_sql(
"status",
"expires_at_epoch",
include_inactive,
));
if let Some(project) = project {
conditions.push(crate::retrieval::memory_search::project_or_global_clause(
"project", idx,
));
params_vec.push(Box::new(project.to_string()));
idx += 1;
}
if let Some(memory_type) = memory_type {
conditions.push(format!("memory_type = ?{idx}"));
params_vec.push(Box::new(memory_type.to_string()));
idx += 1;
}
if let Some(branch) = branch {
conditions.push(format!("(branch = ?{idx} OR branch IS NULL)"));
params_vec.push(Box::new(branch.to_string()));
idx += 1;
}
let sql = format!(
"SELECT id FROM memories
WHERE {}
ORDER BY {order_epoch} DESC, id DESC LIMIT ?{}",
conditions.join(" AND "),
idx
);
let mut stmt = conn.prepare(&sql)?;
params_vec.push(Box::new(limit));
let refs = crate::db::to_sql_refs(¶ms_vec);
let rows = stmt.query_map(refs.as_slice(), |row| row.get::<_, i64>(0))?;
for row in rows {
ids.push(row?);
}
Ok(ids)
}
fn temporal_sql(
field: TemporalField,
has_memory_facts: bool,
has_memory_fact_invalidations: bool,
event_time_expr: &str,
) -> (String, String) {
match field {
TemporalField::UpdatedAt => (
"updated_at_epoch BETWEEN ?1 AND ?2".to_string(),
"updated_at_epoch".to_string(),
),
TemporalField::EventTime if has_memory_facts => {
let current_fact_filter =
crate::memory::facts::current_fact_filter_sql("f", has_memory_fact_invalidations);
let fact_event_overlap = format!(
"f.source_memory_id = memories.id \
AND {current_fact_filter} \
AND f.valid_from_epoch IS NOT NULL \
AND f.valid_from_epoch <= ?2 \
AND (f.valid_to_epoch IS NULL OR f.valid_to_epoch > ?1)"
);
let any_fact_event = format!(
"f.source_memory_id = memories.id \
AND {current_fact_filter} \
AND f.valid_from_epoch IS NOT NULL"
);
(
format!(
"(EXISTS (
SELECT 1 FROM memory_facts f
WHERE {fact_event_overlap}
)
OR (
NOT EXISTS (
SELECT 1 FROM memory_facts f
WHERE {any_fact_event}
)
AND {event_time_expr} BETWEEN ?1 AND ?2
))"
),
format!(
"COALESCE((
SELECT MAX(f.valid_from_epoch)
FROM memory_facts f
WHERE {fact_event_overlap}
), {event_time_expr})"
),
)
}
TemporalField::EventTime => (
format!("{event_time_expr} BETWEEN ?1 AND ?2"),
event_time_expr.to_string(),
),
}
}
fn table_exists(conn: &Connection, table_name: &str) -> Result<bool> {
let count: i64 = conn.query_row(
"SELECT COUNT(*)
FROM sqlite_master
WHERE type = 'table' AND name = ?1",
[table_name],
|row| row.get(0),
)?;
Ok(count > 0)
}
fn column_exists(conn: &Connection, table_name: &str, column_name: &str) -> Result<bool> {
let sql = format!("PRAGMA table_info({})", quote_identifier(table_name));
let mut stmt = conn.prepare(&sql)?;
let mut rows = stmt.query([])?;
while let Some(row) = rows.next()? {
let name: String = row.get(1)?;
if name == column_name {
return Ok(true);
}
}
Ok(false)
}
fn quote_identifier(identifier: &str) -> String {
format!("\"{}\"", identifier.replace('"', "\"\""))
}