package labeler import ( "database/sql" "fmt" "io" "log/slog" "os" "path/filepath" "strings" "time" "github.com/bluesky-social/indigo/atproto/atcrypto" "github.com/bluesky-social/indigo/atproto/labeling" "github.com/tursodatabase/go-libsql" ) // LabelVersion is the ATProto label format version (currently 1). const LabelVersion int64 = labeling.ATPROTO_LABEL_VERSION const schema = ` CREATE TABLE IF NOT EXISTS takedowns ( id INTEGER PRIMARY KEY AUTOINCREMENT, input TEXT NOT NULL, subject_did TEXT NOT NULL, subject_repo TEXT NOT NULL DEFAULT '', subject_handle TEXT NOT NULL DEFAULT '', reason TEXT NOT NULL DEFAULT '', created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP, created_by TEXT NOT NULL DEFAULT '', reversed_at TIMESTAMP, reversed_by TEXT NOT NULL DEFAULT '' ); CREATE INDEX IF NOT EXISTS idx_takedowns_active ON takedowns(reversed_at, created_at DESC); CREATE INDEX IF NOT EXISTS idx_takedowns_subject ON takedowns(subject_did, subject_repo); 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, ver INTEGER NOT NULL DEFAULT 1, sig BLOB NOT NULL, subject_did TEXT NOT NULL, subject_repo TEXT NOT NULL DEFAULT '', takedown_id INTEGER REFERENCES takedowns(id) ); 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); CREATE INDEX IF NOT EXISTS idx_labels_uri ON labels(uri); CREATE INDEX IF NOT EXISTS idx_labels_takedown ON labels(takedown_id); ` // Label represents an ATProto label record stored locally. Its on-the-wire representation // is produced by ToLabeling() which round-trips through indigo's labeling package so the // signature stays valid byte-for-byte. // // TakedownID is a labeler-internal pointer to the takedown event that produced this // label (positive or negation). It's never serialized into ATProto wire format. type Label struct { ID int64 Src string URI string CID string Val string Neg bool Cts time.Time Exp *time.Time Ver int64 Sig []byte SubjectDID string SubjectRepo string TakedownID *int64 } // Takedown is a single operator-issued takedown action. Each Takedown owns one or more // Label rows linked by takedown_id. Reversal sets reversed_at / reversed_by in place. type Takedown struct { ID int64 Input string SubjectDID string SubjectRepo string SubjectHandle string Reason string CreatedAt time.Time CreatedBy string ReversedAt *time.Time ReversedBy string LabelCount int } // LibsqlSync configures optional embedded-replica sync to a remote libSQL database. // SyncURL empty means local-only mode. type LibsqlSync struct { SyncURL string AuthToken string SyncInterval time.Duration } // LabelerDB wraps the *sql.DB plus its libsql connector (when in embedded-replica mode) // so the caller can release file locks on shutdown. type LabelerDB struct { DB *sql.DB connector io.Closer } // Close closes the database and the libsql connector (if any). The connector close is // what releases file locks; without it a subsequent local-only open errors with // "database is locked" — the same gotcha the hold ran into. func (l *LabelerDB) Close() error { var dbErr, connErr error if l.DB != nil { dbErr = l.DB.Close() } if l.connector != nil { connErr = l.connector.Close() } if dbErr != nil { return dbErr } return connErr } // OpenDB opens or creates the labeler database. When sync.SyncURL is set, the DB runs // in embedded-replica mode (writes go to the remote, frames replicate to the local // file); otherwise it's a plain local libSQL file. Schema is applied either way. func OpenDB(dbPath string, sync LibsqlSync) (*LabelerDB, error) { if err := os.MkdirAll(filepath.Dir(dbPath), 0755); err != nil { return nil, fmt.Errorf("failed to create db directory: %w", err) } var ( db *sql.DB connector io.Closer ) if sync.SyncURL != "" { opts := []libsql.Option{libsql.WithAuthToken(sync.AuthToken)} if sync.SyncInterval > 0 { opts = append(opts, libsql.WithSyncInterval(sync.SyncInterval)) } conn, err := libsql.NewEmbeddedReplicaConnector(dbPath, sync.SyncURL, opts...) if err != nil { return nil, fmt.Errorf("failed to create libsql embedded replica connector: %w", err) } db = sql.OpenDB(conn) connector = conn slog.Info("Labeler database opened in embedded replica mode", "path", dbPath, "sync_url", sync.SyncURL) } else { dsn := dbPath if !strings.HasPrefix(dsn, "file:") && !strings.HasPrefix(dsn, ":memory:") { dsn = "file:" + dsn } var err error db, err = sql.Open("libsql", dsn) if err != nil { return nil, fmt.Errorf("failed to open database: %w", err) } slog.Info("Labeler database opened in local-only mode", "path", dbPath) } // Local PRAGMAs only — Bunny rejects PRAGMA assignments forwarded over the // replication protocol (same caveat as pkg/hold/db). if sync.SyncURL == "" { var journalMode string if err := db.QueryRow("PRAGMA journal_mode = WAL").Scan(&journalMode); err != nil { _ = closeIfNonNil(db, connector) return nil, fmt.Errorf("failed to set journal mode: %w", err) } var busyTimeout int if err := db.QueryRow("PRAGMA busy_timeout = 5000").Scan(&busyTimeout); err != nil { _ = closeIfNonNil(db, connector) return nil, fmt.Errorf("failed to set busy_timeout: %w", err) } } for _, stmt := range splitStatements(schema) { if _, err := db.Exec(stmt); err != nil { _ = closeIfNonNil(db, connector) return nil, fmt.Errorf("failed to apply schema: %w", err) } } return &LabelerDB{DB: db, connector: connector}, nil } // closeIfNonNil is the defensive cleanup for the failure path on OpenDB so we don't // leave file locks dangling if schema application fails. func closeIfNonNil(db *sql.DB, connector io.Closer) error { if db != nil { _ = db.Close() } if connector != nil { return connector.Close() } return nil } func splitStatements(sql string) []string { parts := strings.Split(sql, ";") out := make([]string, 0, len(parts)) for _, s := range parts { s = strings.TrimSpace(s) if s != "" { out = append(out, s) } } return out } // ToLabeling converts the row into indigo's label struct (deterministic CBOR shape). func (l *Label) ToLabeling() labeling.Label { out := labeling.Label{ CreatedAt: l.Cts.UTC().Format(time.RFC3339), SourceDID: l.Src, URI: l.URI, Val: l.Val, Version: l.Ver, } if l.CID != "" { s := l.CID out.CID = &s } if l.Exp != nil { s := l.Exp.UTC().Format(time.RFC3339) out.ExpiresAt = &s } if l.Neg { t := true out.Negated = &t } if len(l.Sig) > 0 { out.Sig = l.Sig } return out } // Sign computes a k256 signature over the deterministic CBOR encoding of the label // (without the sig field) and stores it on the row. func (l *Label) Sign(key *atcrypto.PrivateKeyK256) error { if l.Ver == 0 { l.Ver = LabelVersion } if l.Cts.IsZero() { l.Cts = time.Now().UTC() } pre := l.ToLabeling() pre.Sig = nil if err := pre.Sign(key); err != nil { return fmt.Errorf("failed to sign label: %w", err) } l.Sig = pre.Sig return nil } // CreateLabel inserts a freshly signed label and returns its sequence id. // Caller must Sign() first — CreateLabel rejects rows missing a signature. func CreateLabel(db *sql.DB, l *Label) (int64, error) { if len(l.Sig) == 0 { return 0, fmt.Errorf("refusing to insert unsigned label") } if l.Ver == 0 { l.Ver = LabelVersion } var expStr *string if l.Exp != nil { s := l.Exp.UTC().Format(time.RFC3339) expStr = &s } var takedownID any if l.TakedownID != nil { takedownID = *l.TakedownID } result, err := db.Exec( `INSERT INTO labels (src, uri, cid, val, neg, cts, exp, ver, sig, subject_did, subject_repo, takedown_id) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, l.Src, l.URI, nullableString(l.CID), l.Val, l.Neg, l.Cts.UTC().Format(time.RFC3339), expStr, l.Ver, l.Sig, l.SubjectDID, l.SubjectRepo, takedownID, ) if err != nil { return 0, fmt.Errorf("failed to insert label: %w", err) } id, err := result.LastInsertId() if err != nil { return 0, err } l.ID = id return id, nil } func nullableString(s string) any { if s == "" { return nil } return s } // 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, ver, sig, subject_did, subject_repo, takedown_id FROM labels WHERE id > ? ORDER BY id ASC LIMIT ?`, cursor, limit, ) if err != nil { return nil, err } defer rows.Close() return scanLabels(rows) } // LatestSeq returns the highest sequence id in the database, or 0 if empty. func LatestSeq(db *sql.DB) (int64, error) { var seq sql.NullInt64 if err := db.QueryRow(`SELECT MAX(id) FROM labels`).Scan(&seq); err != nil { return 0, err } if !seq.Valid { return 0, nil } return seq.Int64, nil } // CreateTakedown inserts a takedown event row and returns its id. The id should then // be stamped onto every label produced by this takedown (positive labels at issue time, // negation labels at reversal time) so the audit trail stays linked. func CreateTakedown(db *sql.DB, t *Takedown) (int64, error) { if t.CreatedAt.IsZero() { t.CreatedAt = time.Now().UTC() } result, err := db.Exec( `INSERT INTO takedowns (input, subject_did, subject_repo, subject_handle, reason, created_at, created_by) VALUES (?, ?, ?, ?, ?, ?, ?)`, t.Input, t.SubjectDID, t.SubjectRepo, t.SubjectHandle, t.Reason, t.CreatedAt.UTC().Format(time.RFC3339), t.CreatedBy, ) if err != nil { return 0, fmt.Errorf("failed to insert takedown: %w", err) } id, err := result.LastInsertId() if err != nil { return 0, err } t.ID = id return id, nil } // GetTakedown loads a single takedown row by id. Returns sql.ErrNoRows when missing. func GetTakedown(db *sql.DB, id int64) (*Takedown, error) { row := db.QueryRow( `SELECT t.id, t.input, t.subject_did, t.subject_repo, t.subject_handle, t.reason, t.created_at, t.created_by, t.reversed_at, t.reversed_by, (SELECT COUNT(*) FROM labels l WHERE l.takedown_id = t.id AND l.neg = 0) FROM takedowns t WHERE t.id = ?`, id, ) return scanTakedown(row.Scan) } // TakedownFilter scopes ListTakedowns to active, reversed, or all rows. type TakedownFilter int const ( TakedownAll TakedownFilter = iota // every takedown row, regardless of reversal state TakedownActive // only takedowns whose reversed_at is NULL TakedownReversed // only takedowns whose reversed_at is set ) // ListTakedowns returns takedown events ordered by created_at DESC, scoped by filter. // The total count reflects the same filter. func ListTakedowns(db *sql.DB, filter TakedownFilter, limit, offset int) ([]Takedown, int, error) { where := "" switch filter { case TakedownActive: where = "WHERE reversed_at IS NULL" case TakedownReversed: where = "WHERE reversed_at IS NOT NULL" } var total int if err := db.QueryRow(`SELECT COUNT(*) FROM takedowns ` + where).Scan(&total); err != nil { return nil, 0, err } rows, err := db.Query( `SELECT t.id, t.input, t.subject_did, t.subject_repo, t.subject_handle, t.reason, t.created_at, t.created_by, t.reversed_at, t.reversed_by, (SELECT COUNT(*) FROM labels l WHERE l.takedown_id = t.id AND l.neg = 0) FROM takedowns t `+where+` ORDER BY t.created_at DESC LIMIT ? OFFSET ?`, limit, offset, ) if err != nil { return nil, 0, err } defer rows.Close() var out []Takedown for rows.Next() { t, err := scanTakedown(rows.Scan) if err != nil { return nil, 0, err } out = append(out, *t) } return out, total, rows.Err() } // MarkTakedownReversed sets the reversed_at / reversed_by fields on the takedown row. // Refuses to overwrite an existing reversal. func MarkTakedownReversed(db *sql.DB, id int64, by string, at time.Time) error { if at.IsZero() { at = time.Now().UTC() } res, err := db.Exec( `UPDATE takedowns SET reversed_at = ?, reversed_by = ? WHERE id = ? AND reversed_at IS NULL`, at.UTC().Format(time.RFC3339), by, id, ) if err != nil { return fmt.Errorf("failed to mark takedown reversed: %w", err) } n, err := res.RowsAffected() if err != nil { return err } if n == 0 { return fmt.Errorf("takedown %d not found or already reversed", id) } return nil } // GetLabelsByTakedown returns all labels (positive + negations) linked to a takedown. func GetLabelsByTakedown(db *sql.DB, takedownID int64) ([]Label, error) { rows, err := db.Query( `SELECT id, src, uri, COALESCE(cid, ''), val, neg, cts, exp, ver, sig, subject_did, subject_repo, takedown_id FROM labels WHERE takedown_id = ? ORDER BY id ASC`, takedownID, ) if err != nil { return nil, err } defer rows.Close() return scanLabels(rows) } // NegateTakedownLabels signs+inserts negation labels for every active (non-negated) // label linked to the given takedown_id. Negations carry the same takedown_id so they // remain part of the takedown's audit trail. // // The NOT EXISTS subquery skips URIs that already have a later neg=1 row (from a prior // reversal call or from an external negation streamed in via subscribeLabels), so this // function is idempotent and won't emit duplicate negations. func NegateTakedownLabels(db *sql.DB, key *atcrypto.PrivateKeyK256, src string, takedownID int64) ([]Label, error) { rows, err := db.Query( `SELECT l1.uri, l1.subject_did, l1.subject_repo FROM labels l1 WHERE l1.takedown_id = ? AND 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 )`, takedownID, ) if err != nil { return nil, err } type entry struct { uri string did string repo string } var entries []entry for rows.Next() { var e entry if err := rows.Scan(&e.uri, &e.did, &e.repo); err != nil { rows.Close() return nil, err } entries = append(entries, e) } rows.Close() if err := rows.Err(); err != nil { return nil, err } id := takedownID out := make([]Label, 0, len(entries)) for _, e := range entries { neg := &Label{ Src: src, URI: e.uri, Val: "!takedown", Neg: true, Cts: time.Now().UTC(), SubjectDID: e.did, SubjectRepo: e.repo, TakedownID: &id, } if err := neg.Sign(key); err != nil { return out, err } if _, err := CreateLabel(db, neg); err != nil { return out, err } out = append(out, *neg) } return out, nil } func scanTakedown(scan func(...any) error) (*Takedown, error) { var ( t Takedown created string revAt *string ) if err := scan( &t.ID, &t.Input, &t.SubjectDID, &t.SubjectRepo, &t.SubjectHandle, &t.Reason, &created, &t.CreatedBy, &revAt, &t.ReversedBy, &t.LabelCount, ); err != nil { return nil, err } if ts, err := time.Parse(time.RFC3339, created); err == nil { t.CreatedAt = ts } if revAt != nil { if ts, err := time.Parse(time.RFC3339, *revAt); err == nil { t.ReversedAt = &ts } } return &t, nil } func scanLabels(rows *sql.Rows) ([]Label, error) { var labels []Label for rows.Next() { var ( l Label cts string exp *string tdID sql.NullInt64 ) if err := rows.Scan(&l.ID, &l.Src, &l.URI, &l.CID, &l.Val, &l.Neg, &cts, &exp, &l.Ver, &l.Sig, &l.SubjectDID, &l.SubjectRepo, &tdID); 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 } } if tdID.Valid { id := tdID.Int64 l.TakedownID = &id } labels = append(labels, l) } return labels, rows.Err() }