use std::collections::{BTreeMap, HashSet, VecDeque};
use std::path::PathBuf;
use minijinja::value::Value;
use minijinja::{Environment, Error as JinjaError, ErrorKind, UndefinedBehavior};
use rusqlite::Connection;
use serde::Serialize;
use thiserror::Error;
use crate::sexp::SexpError;
use crate::storage;
use tftio_org::ast::{Block, Document, Inline};
#[derive(Debug, Error)]
pub enum PromptError {
#[error("no template named {0}")]
TemplateNotFound(String),
#[error("failed to read user template {path}: {source}")]
ReadTemplate {
path: PathBuf,
source: std::io::Error,
},
#[error("query layer failed: {0}")]
Sql(#[from] rusqlite::Error),
#[error("corrupt stored document: {0}")]
Decode(#[from] SexpError),
#[error("template error: {0}")]
Render(#[from] JinjaError),
}
pub const SUMMARY_BODY_CHAR_BUDGET: usize = 500;
pub const DEFAULT_RECENT_LIMIT: usize = 50;
pub const DEFAULT_HUBS_LIMIT: usize = 20;
const BUILTIN_TEMPLATES: &[(&str, &str)] = &[(
"cold-start-audit",
include_str!("../templates/cold-start-audit.j2"),
)];
const TEMPLATE_EXT: &str = "j2";
#[derive(Debug, Clone, Serialize, PartialEq, Eq)]
pub struct TemplateSummary {
pub name: String,
pub source: String,
pub path: Option<PathBuf>,
}
#[must_use]
pub fn list_templates() -> Vec<TemplateSummary> {
list_templates_in(user_template_dir().as_deref())
}
#[must_use]
fn list_templates_in(override_dir: Option<&std::path::Path>) -> Vec<TemplateSummary> {
let mut by_name: BTreeMap<String, TemplateSummary> = BTreeMap::new();
for (name, _) in BUILTIN_TEMPLATES {
by_name.insert(
(*name).to_string(),
TemplateSummary {
name: (*name).to_string(),
source: "builtin".to_string(),
path: None,
},
);
}
if let Some(dir) = override_dir
&& let Ok(entries) = std::fs::read_dir(dir)
{
for entry in entries.flatten() {
let path = entry.path();
if path.extension().and_then(|s| s.to_str()) != Some(TEMPLATE_EXT) {
continue;
}
let Some(stem) = path.file_stem().and_then(|s| s.to_str()) else {
continue;
};
by_name.insert(
stem.to_string(),
TemplateSummary {
name: stem.to_string(),
source: "user".to_string(),
path: Some(path),
},
);
}
}
by_name.into_values().collect()
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ResolvedTemplate {
pub name: String,
pub source: String,
pub path: Option<PathBuf>,
pub body: String,
}
pub fn resolve_template_source(name: &str) -> Result<ResolvedTemplate, PromptError> {
resolve_template_source_in(name, user_template_dir().as_deref())
}
fn resolve_template_source_in(
name: &str,
override_dir: Option<&std::path::Path>,
) -> Result<ResolvedTemplate, PromptError> {
if let Some(dir) = override_dir {
let path = dir.join(format!("{name}.{TEMPLATE_EXT}"));
if path.is_file() {
let body =
std::fs::read_to_string(&path).map_err(|source| PromptError::ReadTemplate {
path: path.clone(),
source,
})?;
return Ok(ResolvedTemplate {
name: name.to_string(),
source: "user".to_string(),
path: Some(path),
body,
});
}
}
for (n, body) in BUILTIN_TEMPLATES {
if *n == name {
return Ok(ResolvedTemplate {
name: name.to_string(),
source: "builtin".to_string(),
path: None,
body: (*body).to_string(),
});
}
}
Err(PromptError::TemplateNotFound(name.to_string()))
}
pub fn render_prompt(conn: &Connection, name: &str) -> Result<String, PromptError> {
let resolved = resolve_template_source(name)?;
let context = build_context(conn)?;
let mut env = build_env();
env.add_template_owned(resolved.name.clone(), resolved.body.clone())?;
let tmpl = env.get_template(&resolved.name)?;
Ok(tmpl.render(context)?)
}
#[must_use]
#[allow(
clippy::disallowed_methods,
reason = "XDG_CONFIG_HOME/HOME locate the optional user template override dir; sanctioned bootstrap locators (REPO_INVARIANTS.md #5b)"
)]
pub fn user_template_dir() -> Option<PathBuf> {
if let Ok(xdg) = std::env::var("XDG_CONFIG_HOME")
&& !xdg.is_empty()
{
return Some(PathBuf::from(xdg).join("kb").join("prompts"));
}
if let Ok(home) = std::env::var("HOME")
&& !home.is_empty()
{
return Some(
PathBuf::from(home)
.join(".config")
.join("kb")
.join("prompts"),
);
}
dirs::config_dir().map(|d| d.join("kb").join("prompts"))
}
fn build_env() -> Environment<'static> {
let mut env = Environment::new();
env.set_undefined_behavior(UndefinedBehavior::Strict);
env
}
fn build_context(conn: &Connection) -> Result<Value, PromptError> {
let recent_rows = storage::list_recent(conn, DEFAULT_RECENT_LIMIT)?;
let orphan_rows = storage::list_orphans(conn)?;
let hub_rows = storage::list_hubs(conn, DEFAULT_HUBS_LIMIT)?;
let tag_freq = tag_frequency(conn)?;
let recent: Vec<Value> = recent_rows
.iter()
.map(node_value_from_row)
.collect::<Result<_, _>>()?;
let orphans: Vec<Value> = orphan_rows
.iter()
.map(node_value_from_row)
.collect::<Result<_, _>>()?;
let hubs: Vec<Value> = hub_rows
.iter()
.map(|h| {
Value::from_serialize(serde_json::json!({
"id": h.id,
"title": h.title,
"in_degree": h.in_degree,
}))
})
.collect();
let tag_frequency_val: Vec<Value> = tag_freq
.iter()
.map(|(tag, count)| {
Value::from_serialize(serde_json::json!({
"tag": tag,
"count": count,
}))
})
.collect();
let db_path = db_path_of(conn);
let by_tag = make_by_tag(db_path.clone());
let search_fn = make_search(db_path.clone());
let all_fn = make_all_nodes(db_path.clone());
let get_fn = make_get(db_path.clone());
let links_fn = make_links(db_path.clone());
let link_distance_fn = make_link_distance(db_path);
let mut ctx: BTreeMap<&'static str, Value> = BTreeMap::new();
ctx.insert("recent", Value::from(recent));
ctx.insert("orphans", Value::from(orphans));
ctx.insert("hubs", Value::from(hubs));
ctx.insert("tag_frequency", Value::from(tag_frequency_val));
ctx.insert("by_tag", by_tag);
ctx.insert("search", search_fn);
ctx.insert("all_nodes", all_fn);
ctx.insert("get", get_fn);
ctx.insert("links", links_fn);
ctx.insert("link_distance", link_distance_fn);
Ok(Value::from_serialize(&ctx))
}
fn db_path_of(conn: &Connection) -> String {
conn.path().map_or_else(String::new, ToString::to_string)
}
fn node_value_from_row(row: &storage::NodeRow) -> Result<Value, PromptError> {
let doc = crate::sexp::decode_document(&row.ast_blob)?;
let body_full = first_paragraph_or_all(&doc, usize::MAX);
let body = summary_body(row.title.as_str(), &doc);
let tags = collect_tag_names(&doc);
Ok(Value::from_serialize(serde_json::json!({
"id": row.id,
"title": row.title,
"tags": tags,
"body": body,
"body_full": body_full,
"created_at": row.created_at,
"updated_at": row.updated_at,
})))
}
fn collect_tag_names(doc: &Document) -> Vec<String> {
storage::extract_tags(doc)
.into_iter()
.map(|t| t.0)
.collect()
}
#[must_use]
pub fn summary_body(title: &str, doc: &Document) -> String {
let first_para = first_paragraph_text(&doc.blocks).unwrap_or_default();
let mut out = String::new();
if !title.trim().is_empty() {
out.push_str(title);
}
if !first_para.trim().is_empty() {
if !out.is_empty() {
out.push_str("\n\n");
}
out.push_str(&first_para);
}
truncate_chars(&out, SUMMARY_BODY_CHAR_BUDGET)
}
fn first_paragraph_or_all(doc: &Document, cap: usize) -> String {
let text = storage::extract_body_text(doc);
if cap == usize::MAX {
text
} else {
truncate_chars(&text, cap)
}
}
fn first_paragraph_text(blocks: &[Block]) -> Option<String> {
for block in blocks {
match block {
Block::Paragraph { inlines } => {
let text = inline_plain(inlines);
if !text.trim().is_empty() {
return Some(text);
}
}
Block::Heading { children, .. } | Block::QuoteBlock { children } => {
if let Some(t) = first_paragraph_text(children) {
return Some(t);
}
}
_ => {}
}
}
None
}
fn inline_plain(inlines: &[Inline]) -> String {
let mut out = String::new();
for inl in inlines {
match inl {
Inline::Plain(s) | Inline::InlineCode(s) | Inline::Verbatim(s) => out.push_str(s),
Inline::Bold(xs) | Inline::Italic(xs) | Inline::Strikethrough(xs) => {
out.push_str(&inline_plain(xs));
}
Inline::LineBreak => out.push('\n'),
Inline::Link {
target,
description,
} => {
if let Some(desc) = description {
out.push_str(desc);
} else {
out.push_str(target);
}
}
}
}
out
}
fn truncate_chars(s: &str, cap: usize) -> String {
if s.chars().count() <= cap {
s.to_string()
} else {
let truncated: String = s.chars().take(cap).collect();
format!("{truncated}…")
}
}
pub fn tag_frequency(conn: &Connection) -> Result<Vec<(String, i64)>, rusqlite::Error> {
let mut stmt = conn.prepare(
"SELECT tag, COUNT(*) AS c FROM node_tags GROUP BY tag ORDER BY c DESC, tag ASC",
)?;
let rows = stmt.query_map([], |r| Ok((r.get::<_, String>(0)?, r.get::<_, i64>(1)?)))?;
rows.collect()
}
fn open_conn(path: &str) -> Result<Connection, JinjaError> {
storage::open_db(path).map_err(|e| {
JinjaError::new(
ErrorKind::InvalidOperation,
format!("kb: failed to open db {path}: {e}"),
)
})
}
fn node_value_in_callable(row: &storage::NodeRow) -> Result<Value, JinjaError> {
node_value_from_row(row)
.map_err(|e| JinjaError::new(ErrorKind::InvalidOperation, e.to_string()))
}
fn node_values_in_callable(rows: &[storage::NodeRow]) -> Result<Vec<Value>, JinjaError> {
rows.iter().map(node_value_in_callable).collect()
}
fn make_by_tag(path: String) -> Value {
Value::from_function(move |tag: &str| -> Result<Value, JinjaError> {
let conn = open_conn(&path)?;
let rows = storage::list_by_tag(&conn, tag).map_err(|e| {
JinjaError::new(ErrorKind::InvalidOperation, format!("by_tag failed: {e}"))
})?;
Ok(Value::from(node_values_in_callable(&rows)?))
})
}
fn make_search(path: String) -> Value {
Value::from_function(move |query: &str| -> Result<Value, JinjaError> {
let conn = open_conn(&path)?;
let ids = storage::search_fts(&conn, query).map_err(|e| {
JinjaError::new(ErrorKind::InvalidOperation, format!("search failed: {e}"))
})?;
let mut out = Vec::with_capacity(ids.len());
for id in &ids {
if let Some(row) = storage::get_node_row(&conn, id).map_err(|e| {
JinjaError::new(ErrorKind::InvalidOperation, format!("search fetch: {e}"))
})? {
out.push(node_value_in_callable(&row)?);
}
}
Ok(Value::from(out))
})
}
fn make_all_nodes(path: String) -> Value {
Value::from_function(move || -> Result<Value, JinjaError> {
let conn = open_conn(&path)?;
let pairs = storage::list_all_nodes(&conn, 10_000, 0).map_err(|e| {
JinjaError::new(
ErrorKind::InvalidOperation,
format!("all_nodes failed: {e}"),
)
})?;
let mut out = Vec::with_capacity(pairs.len());
for (id, _title) in &pairs {
if let Some(row) = storage::get_node_row(&conn, &id.0).map_err(|e| {
JinjaError::new(ErrorKind::InvalidOperation, format!("all_nodes fetch: {e}"))
})? {
out.push(node_value_in_callable(&row)?);
}
}
Ok(Value::from(out))
})
}
fn make_get(path: String) -> Value {
Value::from_function(move |id: &str| -> Result<Value, JinjaError> {
let conn = open_conn(&path)?;
storage::get_node_row(&conn, id)
.map_err(|e| JinjaError::new(ErrorKind::InvalidOperation, format!("get failed: {e}")))?
.map_or_else(|| Ok(Value::from(())), |row| node_value_in_callable(&row))
})
}
fn make_links(path: String) -> Value {
Value::from_function(move |id: &str| -> Result<Value, JinjaError> {
let conn = open_conn(&path)?;
let nb = storage::get_links(&conn, id).map_err(|e| {
JinjaError::new(ErrorKind::InvalidOperation, format!("links failed: {e}"))
})?;
let to_row = |r: &storage::LinkRow| {
serde_json::json!({
"source_id": r.source_id,
"link_type": r.link_type,
"target_id": r.target_id,
"target_slug": r.target_slug,
})
};
Ok(Value::from_serialize(serde_json::json!({
"outgoing": nb.outgoing.iter().map(to_row).collect::<Vec<_>>(),
"incoming": nb.incoming.iter().map(to_row).collect::<Vec<_>>(),
})))
})
}
fn make_link_distance(path: String) -> Value {
Value::from_function(
move |id: &str, max_depth: i64| -> Result<Value, JinjaError> {
let conn = open_conn(&path)?;
let max = usize::try_from(max_depth.max(0)).unwrap_or(0);
let ids = bfs_link_distance(&conn, id, max).map_err(|e| {
JinjaError::new(
ErrorKind::InvalidOperation,
format!("link_distance failed: {e}"),
)
})?;
let mut out = Vec::with_capacity(ids.len());
for nid in &ids {
if let Some(row) = storage::get_node_row(&conn, nid).map_err(|e| {
JinjaError::new(
ErrorKind::InvalidOperation,
format!("link_distance fetch: {e}"),
)
})? {
out.push(node_value_in_callable(&row)?);
}
}
Ok(Value::from(out))
},
)
}
fn bfs_link_distance(
conn: &Connection,
start: &str,
max_depth: usize,
) -> Result<Vec<String>, rusqlite::Error> {
let mut visited: HashSet<String> = HashSet::new();
let mut queue: VecDeque<(String, usize)> = VecDeque::new();
let mut order: Vec<String> = Vec::new();
queue.push_back((start.to_string(), 0));
visited.insert(start.to_string());
while let Some((node, depth)) = queue.pop_front() {
order.push(node.clone());
if depth >= max_depth {
continue;
}
let nb = storage::get_links(conn, &node)?;
for row in nb.outgoing.iter().chain(nb.incoming.iter()) {
let candidates = [row.target_id.as_ref(), Some(&row.source_id)];
for c in candidates.into_iter().flatten() {
if c == &node {
continue;
}
if visited.insert(c.clone()) {
queue.push_back((c.clone(), depth + 1));
}
}
}
}
Ok(order)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::parser;
use minijinja::value::ValueKind;
fn setup_db() -> Connection {
let conn = storage::open_db(":memory:").expect("open db");
storage::init_db(&conn).expect("init db");
conn
}
fn insert(conn: &Connection, id: &str, body: &str) {
let doc = parser::parse_document(body).expect("parse");
storage::insert_node(conn, id, &doc).expect("insert");
}
#[test]
fn cold_start_audit_template_ships() {
let resolved =
resolve_template_source("cold-start-audit").expect("cold-start-audit template missing");
assert_eq!(resolved.source, "builtin");
assert!(resolved.path.is_none());
assert!(
resolved.body.contains("Cold-Start Audit"),
"template body should contain its title heading: {}",
resolved.body
);
let summary = list_templates();
assert!(
summary.iter().any(|t| t.name == "cold-start-audit"),
"list_templates must include cold-start-audit: {summary:?}"
);
}
#[test]
fn prompt_renders_to_stdout() {
let conn = setup_db();
insert(
&conn,
"n1",
"#+title: Hello World\n* heading\n\nA short paragraph here.\n",
);
insert(&conn, "n2", "* Another\n\nMore text.\n");
let out = render_prompt(&conn, "cold-start-audit").expect("render");
assert!(out.contains("Cold-Start Audit"), "expected heading: {out}");
assert!(out.contains("Hello World"), "expected n1 title: {out}");
}
#[test]
fn summary_projection_default() {
let long_para: String = "x".repeat(SUMMARY_BODY_CHAR_BUDGET * 3);
let body = format!("#+title: Big\n\n{long_para}\n");
let doc = parser::parse_document(&body).expect("parse");
let summary = summary_body("Big", &doc);
assert!(
summary.chars().count() <= SUMMARY_BODY_CHAR_BUDGET + 1,
"summary should be capped to ~{SUMMARY_BODY_CHAR_BUDGET} chars, got {}",
summary.chars().count()
);
assert!(summary.starts_with("Big"));
let conn = setup_db();
storage::insert_node(&conn, "n1", &doc).expect("insert");
let path = db_path_of(&conn);
let ctx = build_context(&conn).expect("ctx");
let mut env = build_env();
env.add_template(
"t",
"{{ recent[0].body }}|{{ recent[0].body_full | length }}",
)
.expect("compile");
let rendered = env.get_template("t").unwrap().render(ctx).expect("render");
let parts: Vec<&str> = rendered.split('|').collect();
assert_eq!(parts.len(), 2, "two fields rendered: {rendered}");
assert!(
parts[0].chars().count() <= SUMMARY_BODY_CHAR_BUDGET + 1,
"body summary length: {}",
parts[0].chars().count()
);
let full_len: usize = parts[1].parse().expect("body_full length parses");
assert!(
full_len > SUMMARY_BODY_CHAR_BUDGET,
"body_full must exceed summary cap: {full_len}"
);
let _ = path;
}
#[test]
fn template_query_surface() {
let tmp = tempfile::tempdir().expect("tempdir");
let db_path = tmp.path().join("kb.db");
let conn = storage::open_db(&db_path.to_string_lossy()).expect("open db");
storage::init_db(&conn).expect("init db");
insert(
&conn,
"n1",
"#+title: Alpha\n#+filetags: :rust:\n* h\n\nAlpha body.\n",
);
insert(
&conn,
"n2",
"#+title: Beta\n#+filetags: :rust:\n* h\n\nBeta body links to [[id:n1]].\n",
);
let _ = storage::relink_all(&conn).expect("relink");
let mut env = build_env();
let ctx = build_context(&conn).expect("ctx");
let template = r#"
recent={{ recent | length }}
orphans={{ orphans | length }}
hubs={{ hubs | length }}
tags={{ tag_frequency | length }}
bytag={{ by_tag("rust") | length }}
search={{ search("Alpha") | length }}
all={{ all_nodes() | length }}
get={{ get("n1").title }}
links_out={{ links("n2").outgoing | length }}
distance={{ link_distance("n1", 2) | length }}
"#;
env.add_template("t", template).expect("compile");
let out = env.get_template("t").unwrap().render(ctx).expect("render");
assert!(out.contains("recent=2"), "{out}");
assert!(out.contains("tags=1"), "{out}");
assert!(out.contains("bytag=2"), "{out}");
assert!(out.contains("search=1"), "{out}");
assert!(out.contains("all=2"), "{out}");
assert!(out.contains("get=Alpha"), "{out}");
assert!(
out.contains("distance=2") || out.contains("distance=1"),
"{out}"
);
}
#[test]
fn user_override_resolution() {
let tmp = tempfile::tempdir().expect("tempdir");
let prompts_dir = tmp.path().join("kb").join("prompts");
std::fs::create_dir_all(&prompts_dir).expect("mkdir");
let user_path = prompts_dir.join("cold-start-audit.j2");
std::fs::write(&user_path, "USER OVERRIDE BODY").expect("write");
let resolved =
resolve_template_source_in("cold-start-audit", Some(&prompts_dir)).expect("resolve");
assert_eq!(resolved.source, "user", "user override should win");
assert_eq!(resolved.body, "USER OVERRIDE BODY");
assert_eq!(resolved.path.as_deref(), Some(user_path.as_path()));
let summaries = list_templates_in(Some(&prompts_dir));
let cs = summaries
.iter()
.find(|t| t.name == "cold-start-audit")
.expect("present");
assert_eq!(cs.source, "user");
std::fs::write(prompts_dir.join("custom.j2"), "x").expect("write");
let summaries2 = list_templates_in(Some(&prompts_dir));
assert!(
summaries2
.iter()
.any(|t| t.name == "custom" && t.source == "user")
);
let builtins_only = list_templates_in(None);
assert!(
builtins_only.iter().all(|t| t.source == "builtin"),
"no override dir → no user-source rows"
);
}
#[test]
fn node_value_has_expected_keys() {
let conn = setup_db();
insert(&conn, "n1", "#+title: T\n* h\n\nbody.\n");
let row = storage::get_node_row(&conn, "n1").unwrap().unwrap();
let v = node_value_from_row(&row).expect("project row");
assert_eq!(v.kind(), ValueKind::Map);
for key in [
"id",
"title",
"tags",
"body",
"body_full",
"created_at",
"updated_at",
] {
assert!(
v.get_attr(key).is_ok(),
"missing key {key} in node value: {v:?}"
);
}
}
}