Files
at-container-registry/pkg/labeler/db.go
T

565 lines
16 KiB
Go

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()
}