Files

302 lines
7.5 KiB
Go

package labeler
import (
"database/sql"
"fmt"
"os"
"path/filepath"
"time"
_ "github.com/tursodatabase/go-libsql"
)
const schema = `
CREATE TABLE IF NOT EXISTS labels (
id INTEGER PRIMARY KEY AUTOINCREMENT,
src TEXT NOT NULL,
uri TEXT NOT NULL,
cid TEXT,
val TEXT NOT NULL,
neg BOOLEAN NOT NULL DEFAULT 0,
cts TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
exp TIMESTAMP,
subject_did TEXT NOT NULL,
subject_repo TEXT NOT NULL DEFAULT '',
UNIQUE(src, uri, val, neg)
);
CREATE INDEX IF NOT EXISTS idx_labels_subject ON labels(subject_did, subject_repo);
CREATE INDEX IF NOT EXISTS idx_labels_cts ON labels(cts DESC);
`
// Label represents an ATProto label (com.atproto.label.defs#label).
type Label struct {
ID int64
Src string
URI string
CID string
Val string
Neg bool
Cts time.Time
Exp *time.Time
SubjectDID string
SubjectRepo string
}
// OpenDB opens or creates the labeler database.
func OpenDB(dbPath string) (*sql.DB, error) {
if err := os.MkdirAll(filepath.Dir(dbPath), 0755); err != nil {
return nil, fmt.Errorf("failed to create db directory: %w", err)
}
db, err := sql.Open("libsql", "file:"+dbPath)
if err != nil {
return nil, fmt.Errorf("failed to open database: %w", err)
}
// Apply schema
for _, stmt := range splitStatements(schema) {
if _, err := db.Exec(stmt); err != nil {
return nil, fmt.Errorf("failed to apply schema: %w", err)
}
}
return db, nil
}
// splitStatements splits SQL by semicolons (go-libsql doesn't support multi-statement exec).
func splitStatements(sql string) []string {
var stmts []string
for _, s := range splitOnSemicolon(sql) {
s = trimSpace(s)
if s != "" {
stmts = append(stmts, s)
}
}
return stmts
}
func splitOnSemicolon(s string) []string {
var parts []string
start := 0
for i := 0; i < len(s); i++ {
if s[i] == ';' {
parts = append(parts, s[start:i])
start = i + 1
}
}
if start < len(s) {
parts = append(parts, s[start:])
}
return parts
}
func trimSpace(s string) string {
// Simple trim that handles newlines and spaces
i := 0
for i < len(s) && (s[i] == ' ' || s[i] == '\t' || s[i] == '\n' || s[i] == '\r') {
i++
}
j := len(s)
for j > i && (s[j-1] == ' ' || s[j-1] == '\t' || s[j-1] == '\n' || s[j-1] == '\r') {
j--
}
return s[i:j]
}
// CreateLabel inserts a new label into the database.
func CreateLabel(db *sql.DB, l *Label) (int64, error) {
result, err := db.Exec(
`INSERT INTO labels (src, uri, cid, val, neg, cts, exp, subject_did, subject_repo)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(src, uri, val, neg) DO UPDATE SET cts = excluded.cts`,
l.Src, l.URI, l.CID, l.Val, l.Neg, l.Cts.UTC().Format(time.RFC3339), l.Exp,
l.SubjectDID, l.SubjectRepo,
)
if err != nil {
return 0, fmt.Errorf("failed to create label: %w", err)
}
return result.LastInsertId()
}
// NegateLabel creates a negation label to reverse a previous label.
func NegateLabel(db *sql.DB, src, uri, val string, subjectDID, subjectRepo string) error {
_, err := db.Exec(
`INSERT INTO labels (src, uri, val, neg, cts, subject_did, subject_repo)
VALUES (?, ?, ?, 1, ?, ?, ?)`,
src, uri, val, time.Now().UTC().Format(time.RFC3339), subjectDID, subjectRepo,
)
return err
}
// GetLabelsSince returns labels with id > cursor, ordered by id ascending.
func GetLabelsSince(db *sql.DB, cursor int64, limit int) ([]Label, error) {
rows, err := db.Query(
`SELECT id, src, uri, COALESCE(cid, ''), val, neg, cts, exp, subject_did, subject_repo
FROM labels WHERE id > ? ORDER BY id ASC LIMIT ?`,
cursor, limit,
)
if err != nil {
return nil, err
}
defer rows.Close()
return scanLabels(rows)
}
// ListActiveTakedowns returns active (non-negated) takedown labels.
func ListActiveTakedowns(db *sql.DB, limit, offset int) ([]Label, int, error) {
var total int
err := db.QueryRow(
`SELECT COUNT(*) FROM labels l1
WHERE l1.val = '!takedown' AND l1.neg = 0
AND NOT EXISTS (
SELECT 1 FROM labels l2
WHERE l2.src = l1.src AND l2.uri = l1.uri AND l2.val = l1.val
AND l2.neg = 1 AND l2.id > l1.id
)
AND (l1.exp IS NULL OR l1.exp > CURRENT_TIMESTAMP)`,
).Scan(&total)
if err != nil {
return nil, 0, err
}
rows, err := db.Query(
`SELECT l1.id, l1.src, l1.uri, COALESCE(l1.cid, ''), l1.val, l1.neg, l1.cts, l1.exp, l1.subject_did, l1.subject_repo
FROM labels l1
WHERE l1.val = '!takedown' AND l1.neg = 0
AND NOT EXISTS (
SELECT 1 FROM labels l2
WHERE l2.src = l1.src AND l2.uri = l1.uri AND l2.val = l1.val
AND l2.neg = 1 AND l2.id > l1.id
)
AND (l1.exp IS NULL OR l1.exp > CURRENT_TIMESTAMP)
ORDER BY l1.cts DESC LIMIT ? OFFSET ?`,
limit, offset,
)
if err != nil {
return nil, 0, err
}
defer rows.Close()
labels, err := scanLabels(rows)
return labels, total, err
}
// GetLabelsForRepo returns all active labels for a specific DID + repository.
func GetLabelsForRepo(db *sql.DB, did, repo string) ([]Label, error) {
rows, err := db.Query(
`SELECT id, src, uri, COALESCE(cid, ''), val, neg, cts, exp, subject_did, subject_repo
FROM labels
WHERE subject_did = ? AND subject_repo = ?
ORDER BY cts DESC`,
did, repo,
)
if err != nil {
return nil, err
}
defer rows.Close()
return scanLabels(rows)
}
// NegateRepoLabels creates negation labels for all active takedown labels on a (DID, repo) pair.
func NegateRepoLabels(db *sql.DB, src, did, repo string) error {
rows, err := db.Query(
`SELECT uri FROM labels
WHERE subject_did = ? AND subject_repo = ? AND val = '!takedown' AND neg = 0`,
did, repo,
)
if err != nil {
return err
}
var uris []string
for rows.Next() {
var uri string
if err := rows.Scan(&uri); err != nil {
rows.Close()
return err
}
uris = append(uris, uri)
}
rows.Close()
if err := rows.Err(); err != nil {
return err
}
now := time.Now().UTC().Format(time.RFC3339)
for _, uri := range uris {
if _, err := db.Exec(
`INSERT INTO labels (src, uri, val, neg, cts, subject_did, subject_repo)
VALUES (?, ?, '!takedown', 1, ?, ?, ?)`,
src, uri, now, did, repo,
); err != nil {
return err
}
}
return nil
}
// NegateUserLabels creates negation labels for all active takedown labels on a DID (user-level).
func NegateUserLabels(db *sql.DB, src, did string) error {
rows, err := db.Query(
`SELECT uri, subject_repo FROM labels
WHERE subject_did = ? AND val = '!takedown' AND neg = 0`,
did,
)
if err != nil {
return err
}
type uriRepo struct {
uri string
repo string
}
var entries []uriRepo
for rows.Next() {
var e uriRepo
if err := rows.Scan(&e.uri, &e.repo); err != nil {
rows.Close()
return err
}
entries = append(entries, e)
}
rows.Close()
if err := rows.Err(); err != nil {
return err
}
now := time.Now().UTC().Format(time.RFC3339)
for _, e := range entries {
if _, err := db.Exec(
`INSERT INTO labels (src, uri, val, neg, cts, subject_did, subject_repo)
VALUES (?, ?, '!takedown', 1, ?, ?, ?)`,
src, e.uri, now, did, e.repo,
); err != nil {
return err
}
}
return nil
}
func scanLabels(rows *sql.Rows) ([]Label, error) {
var labels []Label
for rows.Next() {
var l Label
var cts string
var exp *string
if err := rows.Scan(&l.ID, &l.Src, &l.URI, &l.CID, &l.Val, &l.Neg, &cts, &exp, &l.SubjectDID, &l.SubjectRepo); err != nil {
return nil, err
}
if t, err := time.Parse(time.RFC3339, cts); err == nil {
l.Cts = t
}
if exp != nil {
if t, err := time.Parse(time.RFC3339, *exp); err == nil {
l.Exp = &t
}
}
labels = append(labels, l)
}
return labels, rows.Err()
}