Standing interests get a table of their own plus article_interests, the per-run top-3 matches, and a pure rates() that derives each interest's Beta-smoothed weight from current ratings the way feed affinity does. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01K9PrjtUS16PAQve8D4bHgc
535 lines
16 KiB
Rust
535 lines
16 KiB
Rust
//! Standing-interest storage and rating-derived weights.
|
||
//!
|
||
//! Interest queries stay here so the central database layer remains focused on
|
||
//! the pipeline's shared records.
|
||
|
||
use std::collections::HashMap;
|
||
|
||
use anyhow::{Result, bail};
|
||
use jiff::Timestamp;
|
||
use sqlx::Row as _;
|
||
|
||
use crate::curate::signals::TopInterest;
|
||
use crate::db::{Db, fmt_ts};
|
||
use crate::types::ArticleId;
|
||
|
||
/// One standing interest and its optional prompt category.
|
||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||
pub struct Interest {
|
||
pub id: i64,
|
||
pub name: String,
|
||
pub category: Option<String>,
|
||
pub created_at: String,
|
||
pub categorized_at: Option<String>,
|
||
}
|
||
|
||
/// Result of adding a name whose uniqueness is case-insensitive.
|
||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||
pub enum AddOutcome {
|
||
Added(i64),
|
||
Duplicate,
|
||
}
|
||
|
||
/// One stored article-to-interest match.
|
||
#[derive(Debug, Clone, PartialEq)]
|
||
pub struct MatchRow {
|
||
pub article_id: ArticleId,
|
||
pub interest_id: i64,
|
||
pub name: String,
|
||
pub cos: f64,
|
||
pub z: f64,
|
||
}
|
||
|
||
/// Rating credit accumulated for one interest.
|
||
#[derive(Debug, Clone, Copy, Default, PartialEq)]
|
||
pub struct Rate {
|
||
pub up: f64,
|
||
pub down: f64,
|
||
pub n: usize,
|
||
}
|
||
|
||
impl Rate {
|
||
/// Beta smoothing keeps an unrated interest neutral.
|
||
pub fn weight(&self) -> f64 {
|
||
(self.up + 1.0) / (self.up + self.down + 2.0)
|
||
}
|
||
}
|
||
|
||
/// Interest rates plus the number of ratings that could affect them.
|
||
#[derive(Debug, Clone, Default, PartialEq)]
|
||
pub struct Rates {
|
||
pub by_interest: HashMap<i64, Rate>,
|
||
pub attributable: usize,
|
||
}
|
||
|
||
const INTEREST_COLUMNS: &str = "id, name, category, created_at, categorized_at";
|
||
|
||
fn interest_from(row: &sqlx::sqlite::SqliteRow) -> Interest {
|
||
Interest {
|
||
id: row.get("id"),
|
||
name: row.get("name"),
|
||
category: row.get("category"),
|
||
created_at: row.get("created_at"),
|
||
categorized_at: row.get("categorized_at"),
|
||
}
|
||
}
|
||
|
||
/// All interests, ordered case-insensitively by name.
|
||
pub async fn list(db: &Db) -> Result<Vec<Interest>> {
|
||
let rows = sqlx::query(sqlx::AssertSqlSafe(format!(
|
||
"SELECT {INTEREST_COLUMNS} FROM interests ORDER BY name COLLATE NOCASE, name"
|
||
)))
|
||
.fetch_all(db.pool())
|
||
.await?;
|
||
Ok(rows.iter().map(interest_from).collect())
|
||
}
|
||
|
||
/// Add one trimmed, non-empty name of at most 80 characters.
|
||
pub async fn add(
|
||
db: &Db,
|
||
name: &str,
|
||
category: Option<&str>,
|
||
now: Timestamp,
|
||
) -> Result<AddOutcome> {
|
||
let name = name.trim();
|
||
let len = name.chars().count();
|
||
if !(1..=80).contains(&len) {
|
||
bail!("interest name must be 1–80 characters");
|
||
}
|
||
|
||
let result = sqlx::query(
|
||
"INSERT OR IGNORE INTO interests (name, category, created_at, categorized_at)
|
||
VALUES (?, ?, ?, ?)",
|
||
)
|
||
.bind(name)
|
||
.bind(category)
|
||
.bind(fmt_ts(now))
|
||
.bind(category.map(|_| fmt_ts(now)))
|
||
.execute(db.pool())
|
||
.await?;
|
||
if result.rows_affected() == 0 {
|
||
Ok(AddOutcome::Duplicate)
|
||
} else {
|
||
Ok(AddOutcome::Added(result.last_insert_rowid()))
|
||
}
|
||
}
|
||
|
||
/// Set or clear an interest category and its categorization timestamp together.
|
||
pub async fn set_category(db: &Db, id: i64, category: Option<&str>, now: Timestamp) -> Result<()> {
|
||
sqlx::query("UPDATE interests SET category = ?, categorized_at = ? WHERE id = ?")
|
||
.bind(category)
|
||
.bind(category.map(|_| fmt_ts(now)))
|
||
.bind(id)
|
||
.execute(db.pool())
|
||
.await?;
|
||
Ok(())
|
||
}
|
||
|
||
/// Delete an interest, its match rows, and its name-keyed cached embedding.
|
||
pub async fn delete(db: &Db, id: i64) -> Result<()> {
|
||
let mut tx = db.pool().begin().await?;
|
||
let name: Option<String> = sqlx::query_scalar("SELECT name FROM interests WHERE id = ?")
|
||
.bind(id)
|
||
.fetch_optional(&mut *tx)
|
||
.await?;
|
||
sqlx::query("DELETE FROM interests WHERE id = ?")
|
||
.bind(id)
|
||
.execute(&mut *tx)
|
||
.await?;
|
||
if let Some(name) = name {
|
||
sqlx::query("DELETE FROM interest_embeddings WHERE interest = ?")
|
||
.bind(name)
|
||
.execute(&mut *tx)
|
||
.await?;
|
||
}
|
||
tx.commit().await?;
|
||
Ok(())
|
||
}
|
||
|
||
/// All names in the stable order used for embedding requests.
|
||
pub async fn names(db: &Db) -> Result<Vec<String>> {
|
||
Ok(
|
||
sqlx::query_scalar("SELECT name FROM interests ORDER BY name COLLATE NOCASE, name")
|
||
.fetch_all(db.pool())
|
||
.await?,
|
||
)
|
||
}
|
||
|
||
/// Names grouped for the prompt, with uncategorized interests last.
|
||
pub async fn grouped(db: &Db) -> Result<Vec<(String, Vec<String>)>> {
|
||
let mut by_category: HashMap<String, Vec<String>> = HashMap::new();
|
||
let mut other = Vec::new();
|
||
for interest in list(db).await? {
|
||
if let Some(category) = interest.category {
|
||
by_category.entry(category).or_default().push(interest.name);
|
||
} else {
|
||
other.push(interest.name);
|
||
}
|
||
}
|
||
|
||
let mut groups: Vec<_> = by_category.into_iter().collect();
|
||
groups.sort_by(|left, right| {
|
||
left.0
|
||
.to_lowercase()
|
||
.cmp(&right.0.to_lowercase())
|
||
.then_with(|| left.0.cmp(&right.0))
|
||
});
|
||
for (_, members) in &mut groups {
|
||
members.sort_by(|left, right| {
|
||
left.to_lowercase()
|
||
.cmp(&right.to_lowercase())
|
||
.then_with(|| left.cmp(right))
|
||
});
|
||
}
|
||
if !other.is_empty() {
|
||
groups.push(("Other standing interests".to_string(), other));
|
||
}
|
||
Ok(groups)
|
||
}
|
||
|
||
/// Interests awaiting the categorizer, ordered case-insensitively by name.
|
||
pub async fn uncategorized(db: &Db) -> Result<Vec<Interest>> {
|
||
let rows = sqlx::query(sqlx::AssertSqlSafe(format!(
|
||
"SELECT {INTEREST_COLUMNS} FROM interests WHERE category IS NULL
|
||
ORDER BY name COLLATE NOCASE, name"
|
||
)))
|
||
.fetch_all(db.pool())
|
||
.await?;
|
||
Ok(rows.iter().map(interest_from).collect())
|
||
}
|
||
|
||
/// Upsert the current run's recorded top-interest matches in one transaction.
|
||
pub async fn replace_matches(
|
||
db: &Db,
|
||
run_id: Option<i64>,
|
||
matches: &[(ArticleId, Vec<TopInterest>)],
|
||
ids: &HashMap<String, i64>,
|
||
) -> Result<()> {
|
||
let mut tx = db.pool().begin().await?;
|
||
for (article_id, top_interests) in matches {
|
||
for top in top_interests {
|
||
let Some(interest_id) = ids.get(&top.name) else {
|
||
continue;
|
||
};
|
||
sqlx::query(
|
||
"INSERT INTO article_interests (article_id, interest_id, cos, z, run_id)
|
||
VALUES (?, ?, ?, ?, ?)
|
||
ON CONFLICT(article_id, interest_id) DO UPDATE SET
|
||
cos = excluded.cos, z = excluded.z, run_id = excluded.run_id",
|
||
)
|
||
.bind(article_id)
|
||
.bind(interest_id)
|
||
.bind(top.cos)
|
||
.bind(top.z)
|
||
.bind(run_id)
|
||
.execute(&mut *tx)
|
||
.await?;
|
||
}
|
||
}
|
||
tx.commit().await?;
|
||
Ok(())
|
||
}
|
||
|
||
/// Stored matches for the requested articles.
|
||
pub async fn matches_for_articles(db: &Db, article_ids: &[ArticleId]) -> Result<Vec<MatchRow>> {
|
||
let mut matches = Vec::new();
|
||
for chunk in article_ids.chunks(500) {
|
||
let placeholders = vec!["?"; chunk.len()].join(", ");
|
||
let mut query = sqlx::query(sqlx::AssertSqlSafe(format!(
|
||
"SELECT ai.article_id, ai.interest_id, i.name, ai.cos, ai.z
|
||
FROM article_interests ai
|
||
JOIN interests i ON i.id = ai.interest_id
|
||
WHERE ai.article_id IN ({placeholders})"
|
||
)));
|
||
for article_id in chunk {
|
||
query = query.bind(article_id);
|
||
}
|
||
for row in query.fetch_all(db.pool()).await? {
|
||
matches.push(MatchRow {
|
||
article_id: row.get("article_id"),
|
||
interest_id: row.get("interest_id"),
|
||
name: row.get("name"),
|
||
cos: row.get("cos"),
|
||
z: row.get("z"),
|
||
});
|
||
}
|
||
}
|
||
matches.sort_by(|left, right| {
|
||
left.article_id
|
||
.cmp(&right.article_id)
|
||
.then_with(|| right.z.total_cmp(&left.z))
|
||
.then_with(|| left.interest_id.cmp(&right.interest_id))
|
||
});
|
||
Ok(matches)
|
||
}
|
||
|
||
/// Number of stored article matches for each interest.
|
||
pub async fn match_counts(db: &Db) -> Result<HashMap<i64, i64>> {
|
||
let rows = sqlx::query(
|
||
"SELECT interest_id, COUNT(*) AS matches FROM article_interests GROUP BY interest_id",
|
||
)
|
||
.fetch_all(db.pool())
|
||
.await?;
|
||
Ok(rows
|
||
.iter()
|
||
.map(|row| (row.get("interest_id"), row.get("matches")))
|
||
.collect())
|
||
}
|
||
|
||
/// Derive smoothed interest rates from current ratings and their match rows.
|
||
pub fn rates(ratings: &[(ArticleId, f64, f64)], rows: &[(ArticleId, i64, f64)]) -> Rates {
|
||
let mut rows_by_article: HashMap<ArticleId, Vec<(i64, f64)>> = HashMap::new();
|
||
for &(article_id, interest_id, z) in rows {
|
||
rows_by_article
|
||
.entry(article_id)
|
||
.or_default()
|
||
.push((interest_id, z));
|
||
}
|
||
|
||
let mut result = Rates::default();
|
||
for &(article_id, value, decay) in ratings {
|
||
let mut attributed = false;
|
||
if let Some(matches) = rows_by_article.get(&article_id) {
|
||
for &(interest_id, z) in matches {
|
||
let strength = (z / 3.0).clamp(0.0, 1.0);
|
||
if strength <= 0.0 {
|
||
continue;
|
||
}
|
||
attributed = true;
|
||
let credit = value * decay * strength;
|
||
let rate = result.by_interest.entry(interest_id).or_default();
|
||
rate.up += credit.max(0.0);
|
||
rate.down += (-credit).max(0.0);
|
||
rate.n += 1;
|
||
}
|
||
}
|
||
if attributed {
|
||
result.attributable += 1;
|
||
}
|
||
}
|
||
result
|
||
}
|
||
|
||
#[cfg(test)]
|
||
mod tests {
|
||
use super::*;
|
||
|
||
fn ts(value: &str) -> Timestamp {
|
||
value.parse().unwrap()
|
||
}
|
||
|
||
async fn test_db() -> (tempfile::TempDir, Db) {
|
||
let dir = tempfile::tempdir().unwrap();
|
||
let db = Db::open_and_migrate(&dir.path().join("db.sqlite"))
|
||
.await
|
||
.unwrap();
|
||
(dir, db)
|
||
}
|
||
|
||
async fn seed_article(db: &Db, id: ArticleId) {
|
||
sqlx::query(
|
||
"INSERT INTO articles (id, canonical_url, title, first_seen) VALUES (?, ?, ?, ?)",
|
||
)
|
||
.bind(id)
|
||
.bind(format!("https://example.com/{id}"))
|
||
.bind(format!("Article {id}"))
|
||
.bind("2026-09-12T00:00:00Z")
|
||
.execute(db.pool())
|
||
.await
|
||
.unwrap();
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn add_trims_names_and_uniqueness_is_case_insensitive() {
|
||
let (_dir, db) = test_db().await;
|
||
let now = ts("2026-09-12T12:00:00Z");
|
||
let AddOutcome::Added(id) = add(&db, " Rust ", Some("Software"), now).await.unwrap()
|
||
else {
|
||
panic!("first insert should succeed");
|
||
};
|
||
assert_eq!(
|
||
add(&db, "rust", None, now).await.unwrap(),
|
||
AddOutcome::Duplicate
|
||
);
|
||
assert!(add(&db, " ", None, now).await.is_err());
|
||
assert!(add(&db, &"x".repeat(81), None, now).await.is_err());
|
||
|
||
let interests = list(&db).await.unwrap();
|
||
assert_eq!(interests.len(), 1);
|
||
assert_eq!(interests[0].id, id);
|
||
assert_eq!(interests[0].name, "Rust");
|
||
assert_eq!(interests[0].category.as_deref(), Some("Software"));
|
||
assert_eq!(
|
||
interests[0].categorized_at.as_deref(),
|
||
Some("2026-09-12T12:00:00Z")
|
||
);
|
||
}
|
||
|
||
#[test]
|
||
fn rates_apply_value_decay_strength_and_negative_credit() {
|
||
let ratings = [
|
||
(1, 1.0, 1.0),
|
||
(2, 0.35, 1.0),
|
||
(3, 1.0, 0.5),
|
||
(4, -1.0, 0.5),
|
||
(5, -1.0, 1.0),
|
||
];
|
||
let rows = [
|
||
(1, 10, 3.0),
|
||
(2, 10, 1.5),
|
||
(3, 11, 3.0),
|
||
(4, 10, 0.9),
|
||
(5, 12, 0.0),
|
||
];
|
||
let rates = rates(&ratings, &rows);
|
||
|
||
assert_eq!(rates.attributable, 4);
|
||
let ten = rates.by_interest[&10];
|
||
assert!((ten.up - 1.175).abs() < 1e-12);
|
||
assert!((ten.down - 0.15).abs() < 1e-12);
|
||
assert_eq!(ten.n, 3);
|
||
assert!((ten.weight() - 2.175 / 3.325).abs() < 1e-12);
|
||
assert_eq!(
|
||
rates.by_interest[&11],
|
||
Rate {
|
||
up: 0.5,
|
||
down: 0.0,
|
||
n: 1
|
||
}
|
||
);
|
||
assert!(!rates.by_interest.contains_key(&12));
|
||
assert_eq!(Rate::default().weight(), 0.5);
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn delete_cascades_matches_and_removes_the_embedding() {
|
||
let (_dir, db) = test_db().await;
|
||
seed_article(&db, 1).await;
|
||
let now = ts("2026-09-12T12:00:00Z");
|
||
let AddOutcome::Added(id) = add(&db, "Rust", None, now).await.unwrap() else {
|
||
unreachable!();
|
||
};
|
||
sqlx::query(
|
||
"INSERT INTO interest_embeddings
|
||
(interest, model, dimension, embedding, created_at) VALUES (?, ?, ?, ?, ?)",
|
||
)
|
||
.bind("Rust")
|
||
.bind("test")
|
||
.bind(1_i64)
|
||
.bind(vec![0_u8; 4])
|
||
.bind("2026-09-12T12:00:00Z")
|
||
.execute(db.pool())
|
||
.await
|
||
.unwrap();
|
||
|
||
let ids = HashMap::from([("Rust".to_string(), id)]);
|
||
replace_matches(
|
||
&db,
|
||
Some(7),
|
||
&[(
|
||
1,
|
||
vec![
|
||
TopInterest {
|
||
name: "Rust".into(),
|
||
cos: 0.7,
|
||
z: 1.2,
|
||
},
|
||
TopInterest {
|
||
name: "Unknown".into(),
|
||
cos: 0.9,
|
||
z: 2.0,
|
||
},
|
||
],
|
||
)],
|
||
&ids,
|
||
)
|
||
.await
|
||
.unwrap();
|
||
replace_matches(
|
||
&db,
|
||
None,
|
||
&[(
|
||
1,
|
||
vec![TopInterest {
|
||
name: "Rust".into(),
|
||
cos: 0.8,
|
||
z: 1.5,
|
||
}],
|
||
)],
|
||
&ids,
|
||
)
|
||
.await
|
||
.unwrap();
|
||
|
||
assert_eq!(match_counts(&db).await.unwrap(), HashMap::from([(id, 1)]));
|
||
assert_eq!(
|
||
matches_for_articles(&db, &[1]).await.unwrap(),
|
||
[MatchRow {
|
||
article_id: 1,
|
||
interest_id: id,
|
||
name: "Rust".into(),
|
||
cos: 0.8,
|
||
z: 1.5
|
||
}]
|
||
);
|
||
let run_id: Option<i64> = sqlx::query_scalar(
|
||
"SELECT run_id FROM article_interests WHERE article_id = 1 AND interest_id = ?",
|
||
)
|
||
.bind(id)
|
||
.fetch_one(db.pool())
|
||
.await
|
||
.unwrap();
|
||
assert_eq!(run_id, None);
|
||
|
||
delete(&db, id).await.unwrap();
|
||
let match_rows: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM article_interests")
|
||
.fetch_one(db.pool())
|
||
.await
|
||
.unwrap();
|
||
let embedding_rows: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM interest_embeddings")
|
||
.fetch_one(db.pool())
|
||
.await
|
||
.unwrap();
|
||
assert_eq!(match_rows, 0);
|
||
assert_eq!(embedding_rows, 0);
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn grouped_sorts_categories_and_puts_uncategorized_last() {
|
||
let (_dir, db) = test_db().await;
|
||
let now = ts("2026-09-12T12:00:00Z");
|
||
add(&db, "zebra", Some("Animals"), now).await.unwrap();
|
||
add(&db, "Alpaca", Some("Animals"), now).await.unwrap();
|
||
let AddOutcome::Added(id) = add(&db, "rust", None, now).await.unwrap() else {
|
||
unreachable!();
|
||
};
|
||
add(&db, "Baking", Some("cooking"), now).await.unwrap();
|
||
|
||
assert_eq!(
|
||
grouped(&db).await.unwrap(),
|
||
[
|
||
("Animals".into(), vec!["Alpaca".into(), "zebra".into()]),
|
||
("cooking".into(), vec!["Baking".into()]),
|
||
("Other standing interests".into(), vec!["rust".into()]),
|
||
]
|
||
);
|
||
assert_eq!(uncategorized(&db).await.unwrap()[0].id, id);
|
||
|
||
set_category(&db, id, Some("Software"), now).await.unwrap();
|
||
assert!(uncategorized(&db).await.unwrap().is_empty());
|
||
set_category(&db, id, None, now).await.unwrap();
|
||
let rust = uncategorized(&db).await.unwrap().pop().unwrap();
|
||
assert_eq!(rust.categorized_at, None);
|
||
assert_eq!(
|
||
names(&db).await.unwrap(),
|
||
["Alpaca", "Baking", "rust", "zebra"]
|
||
);
|
||
}
|
||
|
||
#[tokio::test]
|
||
async fn empty_article_lookup_is_a_no_op() {
|
||
let (_dir, db) = test_db().await;
|
||
assert!(matches_for_articles(&db, &[]).await.unwrap().is_empty());
|
||
}
|
||
}
|