Files
routewatch/internal/database/database.go
T
clawbot df9e23d503
check / check (push) Successful in 2m36s
Serve /api/v1/stats counts from realtime in-memory counters (closes #27)
Realtime in-memory counters seeded at startup and adjusted on every insert, update and delete; no periodic recompute. Independent review passed: #29 (comment)

model: claude-opus-4-8 (implementation and review); merged by claude-fable-5
2026-09-22 09:41:16 +02:00

2165 lines
61 KiB
Go

// Package database provides SQLite storage for BGP routing data including ASNs, prefixes, announcements and peerings.
package database
import (
"context"
"database/sql"
_ "embed"
"encoding/json"
"errors"
"fmt"
"net"
"os"
"path/filepath"
"runtime"
"sync"
"time"
"git.eeqj.de/sneak/routewatch/internal/config"
"git.eeqj.de/sneak/routewatch/internal/logger"
"git.eeqj.de/sneak/routewatch/pkg/asinfo"
"github.com/google/uuid"
_ "github.com/mattn/go-sqlite3" // CGO SQLite driver
)
// IMPORTANT: NO schema changes are to be made outside of schema.sql
// We do NOT support migrations. All schema changes MUST be made in schema.sql only.
//
//go:embed schema.sql
var dbSchema string
const (
dirPermissions = 0750 // rwxr-x---
ipVersionV4 = 4
ipVersionV6 = 6
ipv6Length = 16
ipv4Offset = 12
ipv4Bits = 32
maxIPv4 = 0xFFFFFFFF
)
// SQLite memory tuning. cache_size and busy_timeout go in the DSN so every
// pooled connection gets them; the heap limits are process-wide and set once.
const (
// sqliteCacheSizeKiB is the per-connection page cache; negative means KiB.
// -65536 = 64 MiB, so at most 640 MiB across the 10-connection pool.
sqliteCacheSizeKiB = -65536
// sqliteBusyTimeoutMs is how long a connection waits on a locked database.
sqliteBusyTimeoutMs = 5000
// sqliteSoftHeapLimitBytes (1 GiB) makes SQLite recycle its cache rather
// than allocate once its C heap passes this size.
sqliteSoftHeapLimitBytes = 1073741824
// sqliteHardHeapLimitBytes (1.5 GiB) fails a statement with SQLITE_NOMEM
// instead of growing the C heap without bound.
sqliteHardHeapLimitBytes = 1610612736
)
// Common errors
var (
// ErrInvalidIP is returned when an IP address is malformed
ErrInvalidIP = errors.New("invalid IP address")
// ErrNoRoute is returned when no route is found for an IP
ErrNoRoute = errors.New("no route found")
// ErrNoStaleASN is returned when no ASN needs WHOIS refresh
ErrNoStaleASN = errors.New("no stale ASN found")
)
// Database manages the SQLite database connection and operations.
type Database struct {
db *sql.DB
logger *logger.Logger
path string
mu sync.Mutex
lockedAt time.Time
lockedBy string
counts *liveCounts
}
// New creates a new database connection and initializes the schema.
func New(cfg *config.Config, logger *logger.Logger) (*Database, error) {
dbPath := filepath.Join(cfg.GetStateDir(), "db.sqlite")
// Log database path
logger.Info("Opening database", "path", dbPath)
// Ensure directory exists
dir := filepath.Dir(dbPath)
if err := os.MkdirAll(dir, dirPermissions); err != nil {
return nil, fmt.Errorf("failed to create database directory: %w", err)
}
// Per-connection SQLite settings go in the DSN so every pooled connection
// gets them, not just the one that runs the Initialize pragmas. _txlock=
// immediate makes every transaction take the write lock at BEGIN. Without it
// a transaction that reads before writing starts as a reader and, when it
// then writes while another connection holds the write lock, fails at once
// with "database is locked" without waiting for _busy_timeout.
dsn := fmt.Sprintf(
"file:%s?_cache_size=%d&_synchronous=OFF&_busy_timeout=%d&_journal_mode=WAL&_txlock=immediate",
dbPath,
sqliteCacheSizeKiB,
sqliteBusyTimeoutMs,
)
db, err := sql.Open("sqlite3", dsn)
if err != nil {
return nil, fmt.Errorf("failed to open database: %w", err)
}
if err := db.Ping(); err != nil {
return nil, fmt.Errorf("failed to ping database: %w", err)
}
// Set connection pool parameters
// Multiple connections allow concurrent reads while writes are serialized
const maxConns = 10
db.SetMaxOpenConns(maxConns)
db.SetMaxIdleConns(maxConns)
db.SetConnMaxLifetime(0)
database := &Database{db: db, logger: logger, path: dbPath, counts: &liveCounts{}}
if err := database.Initialize(); err != nil {
return nil, fmt.Errorf("failed to initialize database: %w", err)
}
// Seed the in-memory statistics counters from the tables once, before the
// streamer starts writing. From here on every write keeps them current, so
// the stats endpoints never scan the tables to report counts.
if err := database.seedCounts(context.Background()); err != nil {
return nil, fmt.Errorf("failed to seed statistics counters: %w", err)
}
return database, nil
}
// Initialize creates the database schema if it doesn't exist.
func (d *Database) Initialize() error {
// Set SQLite pragmas for performance. Per-connection settings (cache_size,
// synchronous, busy_timeout, journal_mode) live in the DSN; temp_store is
// left at its default so DISTINCT temp B-trees spill to disk instead of C
// heap. The heap limits below are process-wide, so setting them once here is
// enough for the whole pool.
pragmas := []string{
"PRAGMA journal_mode=WAL", // Write-Ahead Logging
"PRAGMA analysis_limit=0", // Disable automatic ANALYZE
"PRAGMA auto_vacuum=INCREMENTAL", // Enable incremental vacuum
fmt.Sprintf("PRAGMA soft_heap_limit=%d", sqliteSoftHeapLimitBytes),
fmt.Sprintf("PRAGMA hard_heap_limit=%d", sqliteHardHeapLimitBytes),
}
for _, pragma := range pragmas {
if err := d.exec(pragma); err != nil {
d.logger.Warn("Failed to set pragma", "pragma", pragma, "error", err)
}
}
// Run WAL checkpoint on startup to consolidate any existing WAL
var walPages, checkpointedPages, movedPages int
err := d.db.QueryRow("PRAGMA wal_checkpoint(TRUNCATE)").Scan(&walPages, &checkpointedPages, &movedPages)
if err != nil {
d.logger.Warn("Failed to checkpoint WAL on startup", "error", err)
} else {
d.logger.Info("WAL checkpoint on startup",
"wal_pages", walPages,
"checkpointed", checkpointedPages,
"moved", movedPages,
)
}
err = d.exec(dbSchema)
if err != nil {
return err
}
return nil
}
// Close closes the database connection.
func (d *Database) Close() error {
return d.db.Close()
}
// lock acquires the database mutex and logs debug information
func (d *Database) lock(operation string) {
// Get caller information
_, file, line, _ := runtime.Caller(1)
caller := fmt.Sprintf("%s:%d", filepath.Base(file), line)
d.logger.Debug("Acquiring database lock", "operation", operation, "caller", caller)
d.mu.Lock()
d.lockedAt = time.Now()
d.lockedBy = fmt.Sprintf("%s (%s)", operation, caller)
d.logger.Debug("Database lock acquired", "operation", operation, "caller", caller)
}
// unlock releases the database mutex and logs debug information including hold duration
func (d *Database) unlock() {
holdDuration := time.Since(d.lockedAt)
lockedBy := d.lockedBy
d.lockedAt = time.Time{}
d.lockedBy = ""
d.mu.Unlock()
d.logger.Debug("Database lock released", "held_by", lockedBy, "duration_ms", holdDuration.Milliseconds())
}
// beginTx starts a new transaction with logging
func (d *Database) beginTx() (*loggingTx, error) {
tx, err := d.db.Begin()
if err != nil {
return nil, err
}
return &loggingTx{Tx: tx, logger: d.logger}, nil
}
// A live-route upsert is an UPDATE followed, only when no row matched, by an
// INSERT. The UPDATE's rows-affected count (1 for an existing key, 0 for a new
// one) is what lets the in-memory route counters stay exact without a COUNT(*).
// Callers hold the database write lock, so no other writer can insert the same
// key between the two statements. The id column is set only on INSERT, so an
// updated route keeps its original id, exactly as the previous ON CONFLICT
// upsert did.
const (
updateLiveRouteV4SQL = `UPDATE live_routes_v4 SET mask_length = ?, as_path = ?, next_hop = ?,
last_updated = ?, ip_start = ?, ip_end = ? WHERE prefix = ? AND origin_asn = ? AND peer_ip = ?`
insertLiveRouteV4SQL = `INSERT INTO live_routes_v4 (id, prefix, mask_length, origin_asn, peer_ip,
as_path, next_hop, last_updated, ip_start, ip_end) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`
updateLiveRouteV6SQL = `UPDATE live_routes_v6 SET mask_length = ?, as_path = ?, next_hop = ?,
last_updated = ? WHERE prefix = ? AND origin_asn = ? AND peer_ip = ?`
insertLiveRouteV6SQL = `INSERT INTO live_routes_v6 (id, prefix, mask_length, origin_asn, peer_ip,
as_path, next_hop, last_updated) VALUES (?, ?, ?, ?, ?, ?, ?, ?)`
)
// upsertRouteRowV4 updates an IPv4 live route, inserting it when no row matched,
// and reports whether a new row was inserted.
func upsertRouteRowV4(upd, ins *sql.Stmt, route *LiveRoute, pathJSON string) (inserted bool, err error) {
if route.V4IPStart == nil || route.V4IPEnd == nil {
return false, fmt.Errorf("IPv4 route %s missing range values", route.Prefix)
}
res, err := upd.Exec(route.MaskLength, pathJSON, route.NextHop, route.LastUpdated,
*route.V4IPStart, *route.V4IPEnd, route.Prefix, route.OriginASN, route.PeerIP)
if err != nil {
return false, err
}
affected, err := res.RowsAffected()
if err != nil {
return false, err
}
if affected > 0 {
return false, nil
}
_, err = ins.Exec(route.ID.String(), route.Prefix, route.MaskLength, route.OriginASN,
route.PeerIP, pathJSON, route.NextHop, route.LastUpdated, *route.V4IPStart, *route.V4IPEnd)
if err != nil {
return false, err
}
return true, nil
}
// upsertRouteRowV6 updates an IPv6 live route, inserting it when no row matched,
// and reports whether a new row was inserted.
func upsertRouteRowV6(upd, ins *sql.Stmt, route *LiveRoute, pathJSON string) (inserted bool, err error) {
res, err := upd.Exec(route.MaskLength, pathJSON, route.NextHop, route.LastUpdated,
route.Prefix, route.OriginASN, route.PeerIP)
if err != nil {
return false, err
}
affected, err := res.RowsAffected()
if err != nil {
return false, err
}
if affected > 0 {
return false, nil
}
_, err = ins.Exec(route.ID.String(), route.Prefix, route.MaskLength, route.OriginASN,
route.PeerIP, pathJSON, route.NextHop, route.LastUpdated)
if err != nil {
return false, err
}
return true, nil
}
// UpsertLiveRouteBatch inserts or updates multiple live routes in a single transaction
func (d *Database) UpsertLiveRouteBatch(routes []*LiveRoute) error {
if len(routes) == 0 {
return nil
}
d.lock("UpsertLiveRouteBatch")
defer d.unlock()
tx, err := d.beginTx()
if err != nil {
return fmt.Errorf("failed to begin transaction: %w", err)
}
defer func() {
if err := tx.Rollback(); err != nil && err != sql.ErrTxDone {
d.logger.Error("Failed to rollback transaction", "error", err)
}
}()
// Prepare the update and insert statements for both tables.
updV4, err := tx.Prepare(updateLiveRouteV4SQL)
if err != nil {
return fmt.Errorf("failed to prepare IPv4 update statement: %w", err)
}
defer func() { _ = updV4.Close() }()
insV4, err := tx.Prepare(insertLiveRouteV4SQL)
if err != nil {
return fmt.Errorf("failed to prepare IPv4 insert statement: %w", err)
}
defer func() { _ = insV4.Close() }()
updV6, err := tx.Prepare(updateLiveRouteV6SQL)
if err != nil {
return fmt.Errorf("failed to prepare IPv6 update statement: %w", err)
}
defer func() { _ = updV6.Close() }()
insV6, err := tx.Prepare(insertLiveRouteV6SQL)
if err != nil {
return fmt.Errorf("failed to prepare IPv6 insert statement: %w", err)
}
defer func() { _ = insV6.Close() }()
var newV4, newV6 int
for _, route := range routes {
pathJSON, err := json.Marshal(route.ASPath)
if err != nil {
return fmt.Errorf("failed to encode AS path: %w", err)
}
if route.IPVersion == ipVersionV4 {
inserted, err := upsertRouteRowV4(updV4, insV4, route, string(pathJSON))
if err != nil {
return fmt.Errorf("failed to upsert route %s: %w", route.Prefix, err)
}
if inserted {
newV4++
}
continue
}
inserted, err := upsertRouteRowV6(updV6, insV6, route, string(pathJSON))
if err != nil {
return fmt.Errorf("failed to upsert route %s: %w", route.Prefix, err)
}
if inserted {
newV6++
}
}
if err = tx.Commit(); err != nil {
return fmt.Errorf("failed to commit transaction: %w", err)
}
d.counts.addRoutes(newV4, newV6)
return nil
}
// DeleteLiveRouteBatch deletes multiple live routes in a single transaction
func (d *Database) DeleteLiveRouteBatch(deletions []LiveRouteDeletion) error {
if len(deletions) == 0 {
return nil
}
d.lock("DeleteLiveRouteBatch")
defer d.unlock()
tx, err := d.beginTx()
if err != nil {
return fmt.Errorf("failed to begin transaction: %w", err)
}
defer func() {
if err := tx.Rollback(); err != nil && err != sql.ErrTxDone {
d.logger.Error("Failed to rollback transaction", "error", err)
}
}()
// No longer need to separate deletions since we handle them in the loop below
// Prepare statements for both IPv4 and IPv6 tables
stmtV4WithOrigin, err := tx.Prepare(`DELETE FROM live_routes_v4 WHERE prefix = ? AND origin_asn = ? AND peer_ip = ?`)
if err != nil {
return fmt.Errorf("failed to prepare IPv4 delete with origin statement: %w", err)
}
defer func() { _ = stmtV4WithOrigin.Close() }()
stmtV4WithoutOrigin, err := tx.Prepare(`DELETE FROM live_routes_v4 WHERE prefix = ? AND peer_ip = ?`)
if err != nil {
return fmt.Errorf("failed to prepare IPv4 delete without origin statement: %w", err)
}
defer func() { _ = stmtV4WithoutOrigin.Close() }()
stmtV6WithOrigin, err := tx.Prepare(`DELETE FROM live_routes_v6 WHERE prefix = ? AND origin_asn = ? AND peer_ip = ?`)
if err != nil {
return fmt.Errorf("failed to prepare IPv6 delete with origin statement: %w", err)
}
defer func() { _ = stmtV6WithOrigin.Close() }()
stmtV6WithoutOrigin, err := tx.Prepare(`DELETE FROM live_routes_v6 WHERE prefix = ? AND peer_ip = ?`)
if err != nil {
return fmt.Errorf("failed to prepare IPv6 delete without origin statement: %w", err)
}
defer func() { _ = stmtV6WithoutOrigin.Close() }()
// Process deletions
var deletedV4, deletedV6 int64
for _, del := range deletions {
var stmt *sql.Stmt
// Select appropriate statement based on IP version and whether we have origin ASN
//nolint:nestif // Clear logic for selecting the right statement
if del.IPVersion == ipVersionV4 {
if del.OriginASN == 0 {
stmt = stmtV4WithoutOrigin
} else {
stmt = stmtV4WithOrigin
}
} else {
if del.OriginASN == 0 {
stmt = stmtV6WithoutOrigin
} else {
stmt = stmtV6WithOrigin
}
}
// Execute deletion
var res sql.Result
if del.OriginASN == 0 {
res, err = stmt.Exec(del.Prefix, del.PeerIP)
} else {
res, err = stmt.Exec(del.Prefix, del.OriginASN, del.PeerIP)
}
if err != nil {
return fmt.Errorf("failed to delete route %s: %w", del.Prefix, err)
}
// A deletion with no origin ASN can remove several rows, so use the
// exact rows-affected count to keep the in-memory route counters right.
affected, err := res.RowsAffected()
if err != nil {
return fmt.Errorf("failed to count deleted route %s: %w", del.Prefix, err)
}
if del.IPVersion == ipVersionV4 {
deletedV4 += affected
} else {
deletedV6 += affected
}
}
if err = tx.Commit(); err != nil {
return fmt.Errorf("failed to commit transaction: %w", err)
}
d.counts.addRoutes(-int(deletedV4), -int(deletedV6))
return nil
}
// UpdatePrefixesBatch updates the last_seen time for multiple prefixes in a single transaction
func (d *Database) UpdatePrefixesBatch(prefixes map[string]time.Time) error {
if len(prefixes) == 0 {
return nil
}
d.lock("UpdatePrefixesBatch")
defer d.unlock()
tx, err := d.beginTx()
if err != nil {
return fmt.Errorf("failed to begin transaction: %w", err)
}
defer func() {
if err := tx.Rollback(); err != nil && err != sql.ErrTxDone {
d.logger.Error("Failed to rollback transaction", "error", err)
}
}()
// Prepare statements for both IPv4 and IPv6 tables
selectV4Stmt, err := tx.Prepare("SELECT id FROM prefixes_v4 WHERE prefix = ?")
if err != nil {
return fmt.Errorf("failed to prepare IPv4 select statement: %w", err)
}
defer func() { _ = selectV4Stmt.Close() }()
updateV4Stmt, err := tx.Prepare("UPDATE prefixes_v4 SET last_seen = ? WHERE prefix = ?")
if err != nil {
return fmt.Errorf("failed to prepare IPv4 update statement: %w", err)
}
defer func() { _ = updateV4Stmt.Close() }()
insertV4Stmt, err := tx.Prepare("INSERT INTO prefixes_v4 (id, prefix, first_seen, last_seen) VALUES (?, ?, ?, ?)")
if err != nil {
return fmt.Errorf("failed to prepare IPv4 insert statement: %w", err)
}
defer func() { _ = insertV4Stmt.Close() }()
selectV6Stmt, err := tx.Prepare("SELECT id FROM prefixes_v6 WHERE prefix = ?")
if err != nil {
return fmt.Errorf("failed to prepare IPv6 select statement: %w", err)
}
defer func() { _ = selectV6Stmt.Close() }()
updateV6Stmt, err := tx.Prepare("UPDATE prefixes_v6 SET last_seen = ? WHERE prefix = ?")
if err != nil {
return fmt.Errorf("failed to prepare IPv6 update statement: %w", err)
}
defer func() { _ = updateV6Stmt.Close() }()
insertV6Stmt, err := tx.Prepare("INSERT INTO prefixes_v6 (id, prefix, first_seen, last_seen) VALUES (?, ?, ?, ?)")
if err != nil {
return fmt.Errorf("failed to prepare IPv6 insert statement: %w", err)
}
defer func() { _ = insertV6Stmt.Close() }()
var newV4, newV6 int
for prefix, timestamp := range prefixes {
ipVersion := detectIPVersion(prefix)
var selectStmt, updateStmt, insertStmt *sql.Stmt
if ipVersion == ipVersionV4 {
selectStmt, updateStmt, insertStmt = selectV4Stmt, updateV4Stmt, insertV4Stmt
} else {
selectStmt, updateStmt, insertStmt = selectV6Stmt, updateV6Stmt, insertV6Stmt
}
var id string
err = selectStmt.QueryRow(prefix).Scan(&id)
switch err {
case nil:
// Prefix exists, update last_seen
_, err = updateStmt.Exec(timestamp, prefix)
if err != nil {
return fmt.Errorf("failed to update prefix %s: %w", prefix, err)
}
case sql.ErrNoRows:
// Prefix doesn't exist, create it
_, err = insertStmt.Exec(generateUUID().String(), prefix, timestamp, timestamp)
if err != nil {
return fmt.Errorf("failed to insert prefix %s: %w", prefix, err)
}
if ipVersion == ipVersionV4 {
newV4++
} else {
newV6++
}
default:
return fmt.Errorf("failed to query prefix %s: %w", prefix, err)
}
}
if err = tx.Commit(); err != nil {
return fmt.Errorf("failed to commit transaction: %w", err)
}
d.counts.addPrefixes(newV4, newV6)
return nil
}
// GetOrCreateASNBatch creates or updates multiple ASNs in a single transaction
func (d *Database) GetOrCreateASNBatch(asns map[int]time.Time) error {
if len(asns) == 0 {
return nil
}
d.lock("GetOrCreateASNBatch")
defer d.unlock()
tx, err := d.beginTx()
if err != nil {
return fmt.Errorf("failed to begin transaction: %w", err)
}
defer func() {
if err := tx.Rollback(); err != nil && err != sql.ErrTxDone {
d.logger.Error("Failed to rollback transaction", "error", err)
}
}()
// Prepare statements
selectStmt, err := tx.Prepare(
"SELECT asn, handle, description, first_seen, last_seen FROM asns WHERE asn = ?")
if err != nil {
return fmt.Errorf("failed to prepare select statement: %w", err)
}
defer func() { _ = selectStmt.Close() }()
updateStmt, err := tx.Prepare("UPDATE asns SET last_seen = ? WHERE asn = ?")
if err != nil {
return fmt.Errorf("failed to prepare update statement: %w", err)
}
defer func() { _ = updateStmt.Close() }()
insertStmt, err := tx.Prepare(
"INSERT INTO asns (asn, handle, description, first_seen, last_seen) VALUES (?, ?, ?, ?, ?)")
if err != nil {
return fmt.Errorf("failed to prepare insert statement: %w", err)
}
defer func() { _ = insertStmt.Close() }()
var newASNs int
for number, timestamp := range asns {
var asn ASN
var handle, description sql.NullString
err = selectStmt.QueryRow(number).Scan(&asn.ASN, &handle, &description, &asn.FirstSeen, &asn.LastSeen)
if err == nil {
// ASN exists, update last_seen
_, err = updateStmt.Exec(timestamp, number)
if err != nil {
return fmt.Errorf("failed to update ASN %d: %w", number, err)
}
continue
}
if err == sql.ErrNoRows {
// ASN doesn't exist, create it
asn = ASN{
ASN: number,
FirstSeen: timestamp,
LastSeen: timestamp,
}
// Look up ASN info
if info, ok := asinfo.Get(number); ok {
asn.Handle = info.Handle
asn.Description = info.Description
}
_, err = insertStmt.Exec(asn.ASN, asn.Handle, asn.Description, asn.FirstSeen, asn.LastSeen)
if err != nil {
return fmt.Errorf("failed to insert ASN %d: %w", number, err)
}
newASNs++
continue
}
if err != nil {
return fmt.Errorf("failed to query ASN %d: %w", number, err)
}
}
if err = tx.Commit(); err != nil {
return fmt.Errorf("failed to commit transaction: %w", err)
}
d.counts.addASNs(newASNs)
return nil
}
// GetOrCreateASN retrieves an existing ASN or creates a new one if it doesn't exist.
func (d *Database) GetOrCreateASN(number int, timestamp time.Time) (*ASN, error) {
d.lock("GetOrCreateASN")
defer d.unlock()
tx, err := d.beginTx()
if err != nil {
return nil, err
}
defer func() {
if err := tx.Rollback(); err != nil && err != sql.ErrTxDone {
d.logger.Error("Failed to rollback transaction", "error", err)
}
}()
var asn ASN
var handle, description sql.NullString
err = tx.QueryRow("SELECT asn, handle, description, first_seen, last_seen FROM asns WHERE asn = ?", number).
Scan(&asn.ASN, &handle, &description, &asn.FirstSeen, &asn.LastSeen)
if err == nil {
// ASN exists, update last_seen
asn.Handle = handle.String
asn.Description = description.String
_, err = tx.Exec("UPDATE asns SET last_seen = ? WHERE asn = ?", timestamp, number)
if err != nil {
return nil, err
}
asn.LastSeen = timestamp
if err = tx.Commit(); err != nil {
d.logger.Error("Failed to commit transaction for ASN update", "asn", number, "error", err)
return nil, err
}
return &asn, nil
}
if err != sql.ErrNoRows {
return nil, err
}
// ASN doesn't exist, create it with ASN info lookup
asn = ASN{
ASN: number,
FirstSeen: timestamp,
LastSeen: timestamp,
}
// Look up ASN info
if info, ok := asinfo.Get(number); ok {
asn.Handle = info.Handle
asn.Description = info.Description
}
_, err = tx.Exec("INSERT INTO asns (asn, handle, description, first_seen, last_seen) VALUES (?, ?, ?, ?, ?)",
asn.ASN, asn.Handle, asn.Description, asn.FirstSeen, asn.LastSeen)
if err != nil {
return nil, err
}
if err = tx.Commit(); err != nil {
d.logger.Error("Failed to commit transaction for ASN creation", "asn", number, "error", err)
return nil, err
}
d.counts.addASNs(1)
return &asn, nil
}
// GetOrCreatePrefix retrieves an existing prefix or creates a new one if it doesn't exist.
func (d *Database) GetOrCreatePrefix(prefix string, timestamp time.Time) (*Prefix, error) {
d.lock("GetOrCreatePrefix")
defer d.unlock()
tx, err := d.beginTx()
if err != nil {
return nil, err
}
defer func() {
if err := tx.Rollback(); err != nil && err != sql.ErrTxDone {
d.logger.Error("Failed to rollback transaction", "error", err)
}
}()
// Determine table based on IP version
ipVersion := detectIPVersion(prefix)
tableName := "prefixes_v4"
if ipVersion == ipVersionV6 {
tableName = "prefixes_v6"
}
var p Prefix
var idStr string
query := fmt.Sprintf("SELECT id, prefix, first_seen, last_seen FROM %s WHERE prefix = ?", tableName)
err = tx.QueryRow(query, prefix).
Scan(&idStr, &p.Prefix, &p.FirstSeen, &p.LastSeen)
if err == nil {
// Prefix exists, update last_seen
p.ID, _ = uuid.Parse(idStr)
p.IPVersion = ipVersion
updateQuery := fmt.Sprintf("UPDATE %s SET last_seen = ? WHERE id = ?", tableName)
_, err = tx.Exec(updateQuery, timestamp, p.ID.String())
if err != nil {
return nil, err
}
p.LastSeen = timestamp
if err = tx.Commit(); err != nil {
d.logger.Error("Failed to commit transaction for prefix update", "prefix", prefix, "error", err)
return nil, err
}
return &p, nil
}
if err != sql.ErrNoRows {
return nil, err
}
// Prefix doesn't exist, create it
p = Prefix{
ID: generateUUID(),
Prefix: prefix,
IPVersion: ipVersion,
FirstSeen: timestamp,
LastSeen: timestamp,
}
insertQuery := fmt.Sprintf("INSERT INTO %s (id, prefix, first_seen, last_seen) VALUES (?, ?, ?, ?)", tableName)
_, err = tx.Exec(insertQuery, p.ID.String(), p.Prefix, p.FirstSeen, p.LastSeen)
if err != nil {
return nil, err
}
if err = tx.Commit(); err != nil {
d.logger.Error("Failed to commit transaction for prefix creation", "prefix", prefix, "error", err)
return nil, err
}
if ipVersion == ipVersionV4 {
d.counts.addPrefixes(1, 0)
} else {
d.counts.addPrefixes(0, 1)
}
return &p, nil
}
// RecordAnnouncement inserts a new BGP announcement or withdrawal into the database.
func (d *Database) RecordAnnouncement(announcement *Announcement) error {
d.lock("RecordAnnouncement")
defer d.unlock()
err := d.exec(`
INSERT INTO announcements (id, prefix_id, peer_asn, origin_asn, path, next_hop, timestamp, is_withdrawal)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)`,
announcement.ID.String(), announcement.PrefixID.String(),
announcement.PeerASN, announcement.OriginASN,
announcement.Path, announcement.NextHop, announcement.Timestamp, announcement.IsWithdrawal)
return err
}
// RecordPeering records a peering relationship between two ASNs.
func (d *Database) RecordPeering(asA, asB int, timestamp time.Time) error {
// Validate ASNs
if asA <= 0 || asB <= 0 {
return fmt.Errorf("invalid ASN: asA=%d, asB=%d", asA, asB)
}
if asA == asB {
return fmt.Errorf("cannot create peering with same ASN: %d", asA)
}
// Normalize: ensure asA < asB
if asA > asB {
asA, asB = asB, asA
}
d.lock("RecordPeering")
defer d.unlock()
tx, err := d.beginTx()
if err != nil {
return err
}
defer func() {
if err := tx.Rollback(); err != nil && err != sql.ErrTxDone {
d.logger.Error("Failed to rollback transaction", "error", err)
}
}()
var exists bool
err = tx.QueryRow("SELECT EXISTS(SELECT 1 FROM peerings WHERE as_a = ? AND as_b = ?)",
asA, asB).Scan(&exists)
if err != nil {
return err
}
if exists {
_, err = tx.Exec("UPDATE peerings SET last_seen = ? WHERE as_a = ? AND as_b = ?",
timestamp, asA, asB)
} else {
_, err = tx.Exec(`
INSERT INTO peerings (id, as_a, as_b, first_seen, last_seen)
VALUES (?, ?, ?, ?, ?)`,
generateUUID().String(), asA, asB, timestamp, timestamp)
}
if err != nil {
return err
}
if err = tx.Commit(); err != nil {
d.logger.Error("Failed to commit transaction for peering",
"as_a", asA,
"as_b", asB,
"error", err,
)
return err
}
if !exists {
d.counts.addPeerings(1)
}
return nil
}
// UpdatePeerBatch updates or creates multiple BGP peer records in a single transaction
func (d *Database) UpdatePeerBatch(peers map[string]PeerUpdate) error {
if len(peers) == 0 {
return nil
}
d.lock("UpdatePeerBatch")
defer d.unlock()
tx, err := d.beginTx()
if err != nil {
return fmt.Errorf("failed to begin transaction: %w", err)
}
defer func() {
if err := tx.Rollback(); err != nil && err != sql.ErrTxDone {
d.logger.Error("Failed to rollback transaction", "error", err)
}
}()
// Prepare statements
checkStmt, err := tx.Prepare("SELECT EXISTS(SELECT 1 FROM bgp_peers WHERE peer_ip = ?)")
if err != nil {
return fmt.Errorf("failed to prepare check statement: %w", err)
}
defer func() { _ = checkStmt.Close() }()
updateStmt, err := tx.Prepare(
"UPDATE bgp_peers SET peer_asn = ?, last_seen = ?, last_message_type = ? WHERE peer_ip = ?")
if err != nil {
return fmt.Errorf("failed to prepare update statement: %w", err)
}
defer func() { _ = updateStmt.Close() }()
insertStmt, err := tx.Prepare(
"INSERT INTO bgp_peers (id, peer_ip, peer_asn, first_seen, last_seen, last_message_type) VALUES (?, ?, ?, ?, ?, ?)")
if err != nil {
return fmt.Errorf("failed to prepare insert statement: %w", err)
}
defer func() { _ = insertStmt.Close() }()
var newPeers int
for _, update := range peers {
var exists bool
err = checkStmt.QueryRow(update.PeerIP).Scan(&exists)
if err != nil {
return fmt.Errorf("failed to check peer %s: %w", update.PeerIP, err)
}
if exists {
_, err = updateStmt.Exec(update.PeerASN, update.Timestamp, update.MessageType, update.PeerIP)
} else {
_, err = insertStmt.Exec(
generateUUID().String(), update.PeerIP, update.PeerASN,
update.Timestamp, update.Timestamp, update.MessageType)
}
if err != nil {
return fmt.Errorf("failed to update peer %s: %w", update.PeerIP, err)
}
if !exists {
newPeers++
}
}
if err = tx.Commit(); err != nil {
return fmt.Errorf("failed to commit transaction: %w", err)
}
d.counts.addPeers(newPeers)
return nil
}
// UpdatePeer updates or creates a BGP peer record
func (d *Database) UpdatePeer(peerIP string, peerASN int, messageType string, timestamp time.Time) error {
d.lock("UpdatePeer")
defer d.unlock()
tx, err := d.beginTx()
if err != nil {
return err
}
defer func() {
if err := tx.Rollback(); err != nil && err != sql.ErrTxDone {
d.logger.Error("Failed to rollback transaction", "error", err)
}
}()
var exists bool
err = tx.QueryRow("SELECT EXISTS(SELECT 1 FROM bgp_peers WHERE peer_ip = ?)", peerIP).Scan(&exists)
if err != nil {
return err
}
if exists {
_, err = tx.Exec(
"UPDATE bgp_peers SET peer_asn = ?, last_seen = ?, last_message_type = ? WHERE peer_ip = ?",
peerASN, timestamp, messageType, peerIP,
)
} else {
_, err = tx.Exec(
"INSERT INTO bgp_peers (id, peer_ip, peer_asn, first_seen, last_seen, last_message_type) VALUES (?, ?, ?, ?, ?, ?)",
generateUUID().String(), peerIP, peerASN, timestamp, timestamp, messageType,
)
}
if err != nil {
return err
}
if err = tx.Commit(); err != nil {
d.logger.Error("Failed to commit transaction for peer update",
"peer_ip", peerIP,
"peer_asn", peerASN,
"error", err,
)
return err
}
if !exists {
d.counts.addPeers(1)
}
return nil
}
// GetStats returns database statistics
func (d *Database) GetStats() (Stats, error) {
return d.GetStatsContext(context.Background())
}
// GetStatsContext returns database statistics with context support.
//
// The row counts (ASNs, prefixes, peerings, peers, live routes) come from the
// in-memory counters, seeded at startup and kept current on every write, so a
// read runs no COUNT(*) over the tables. The oldest/newest route timestamps are
// read from the ends of the last_updated index, and the file size from a
// stat(); neither is a table scan. The only remaining query is the prefix
// distribution.
func (d *Database) GetStatsContext(ctx context.Context) (Stats, error) {
var stats Stats
// Row counts from memory, as a single consistent snapshot.
d.counts.fill(&stats)
// Database file size is a cheap stat() on the file.
if fileInfo, err := os.Stat(d.path); err != nil {
d.logger.Warn("Failed to get database file size", "error", err)
} else {
stats.FileSizeBytes = fileInfo.Size()
}
// Oldest and newest route timestamps read one row from each end of the
// last_updated index (a log-time lookup, not a scan). Selecting the column
// directly lets the driver parse the DATETIME into time.Time; the old
// MIN/MAX union scan returned an untyped string that failed to scan and
// logged a warning on every call.
oldest, newest, err := d.routeTimestampRange(ctx)
if err != nil {
// Display-only fields; log but keep the rest of the stats.
d.logger.Warn("Failed to get route timestamps", "error", err)
} else {
stats.OldestRoute = oldest
stats.NewestRoute = newest
}
// Prefix distribution counts distinct prefixes per mask length. It stays a
// query over the covering (mask_length, prefix) index rather than an
// in-memory counter: maintaining distinct-prefix-per-mask in memory would
// need a per-prefix table of roughly a million entries, memory this service
// is tuned to avoid.
stats.IPv4PrefixDistribution, stats.IPv6PrefixDistribution, err = d.GetPrefixDistributionContext(ctx)
if err != nil {
// Log but don't fail.
d.logger.Warn("Failed to get prefix distribution", "error", err)
}
return stats, nil
}
// routeTimestampRange returns the earliest and latest last_updated across both
// live route tables, or nil values when both tables are empty. Each query reads
// one row from an end of the last_updated index rather than scanning the tables.
func (d *Database) routeTimestampRange(ctx context.Context) (oldest, newest *time.Time, err error) {
oldestV4, ok, err := d.scanRouteTimestamp(ctx,
"SELECT last_updated FROM live_routes_v4 ORDER BY last_updated ASC LIMIT 1")
if err != nil {
return nil, nil, err
}
if ok {
oldest = &oldestV4
}
oldestV6, ok, err := d.scanRouteTimestamp(ctx,
"SELECT last_updated FROM live_routes_v6 ORDER BY last_updated ASC LIMIT 1")
if err != nil {
return nil, nil, err
}
if ok && (oldest == nil || oldestV6.Before(*oldest)) {
oldest = &oldestV6
}
newestV4, ok, err := d.scanRouteTimestamp(ctx,
"SELECT last_updated FROM live_routes_v4 ORDER BY last_updated DESC LIMIT 1")
if err != nil {
return nil, nil, err
}
if ok {
newest = &newestV4
}
newestV6, ok, err := d.scanRouteTimestamp(ctx,
"SELECT last_updated FROM live_routes_v6 ORDER BY last_updated DESC LIMIT 1")
if err != nil {
return nil, nil, err
}
if ok && (newest == nil || newestV6.After(*newest)) {
newest = &newestV6
}
return oldest, newest, nil
}
// scanRouteTimestamp runs a single-row timestamp query. ok is false when the
// table is empty. The query selects the last_updated column directly so the
// driver parses the DATETIME value into a time.Time.
func (d *Database) scanRouteTimestamp(ctx context.Context, query string) (ts time.Time, ok bool, err error) {
err = d.db.QueryRowContext(ctx, query).Scan(&ts)
switch {
case errors.Is(err, sql.ErrNoRows):
return time.Time{}, false, nil
case err != nil:
return time.Time{}, false, err
default:
return ts, true, nil
}
}
// UpsertLiveRoute inserts or updates a live route
func (d *Database) UpsertLiveRoute(route *LiveRoute) error {
d.lock("UpsertLiveRoute")
defer d.unlock()
pathJSON, err := json.Marshal(route.ASPath)
if err != nil {
return fmt.Errorf("failed to encode AS path: %w", err)
}
updateSQL, insertSQL := updateLiveRouteV4SQL, insertLiveRouteV4SQL
if route.IPVersion == ipVersionV6 {
updateSQL, insertSQL = updateLiveRouteV6SQL, insertLiveRouteV6SQL
}
// The write lock is held, so no other writer can insert this key between the
// update and the insert even though they are separate autocommit statements.
upd, err := d.db.Prepare(updateSQL)
if err != nil {
return fmt.Errorf("failed to prepare update statement: %w", err)
}
defer func() { _ = upd.Close() }()
ins, err := d.db.Prepare(insertSQL)
if err != nil {
return fmt.Errorf("failed to prepare insert statement: %w", err)
}
defer func() { _ = ins.Close() }()
var inserted bool
if route.IPVersion == ipVersionV4 {
inserted, err = upsertRouteRowV4(upd, ins, route, string(pathJSON))
} else {
inserted, err = upsertRouteRowV6(upd, ins, route, string(pathJSON))
}
if err != nil {
return fmt.Errorf("failed to upsert route %s: %w", route.Prefix, err)
}
if inserted {
if route.IPVersion == ipVersionV4 {
d.counts.addRoutes(1, 0)
} else {
d.counts.addRoutes(0, 1)
}
}
return nil
}
// DeleteLiveRoute deletes a live route
// If originASN is 0, deletes all routes for the prefix/peer combination
func (d *Database) DeleteLiveRoute(prefix string, originASN int, peerIP string) error {
d.lock("DeleteLiveRoute")
defer d.unlock()
// Determine table based on prefix IP version
_, ipnet, err := net.ParseCIDR(prefix)
if err != nil {
return fmt.Errorf("invalid prefix format: %w", err)
}
isV4 := ipnet.IP.To4() != nil
// Literal per-table queries (rather than one formatted with the table name)
// so the delete carries no dynamically built SQL. A delete with no origin
// ASN can remove several rows.
var res sql.Result
switch {
case isV4 && originASN == 0:
res, err = d.db.Exec(`DELETE FROM live_routes_v4 WHERE prefix = ? AND peer_ip = ?`, prefix, peerIP)
case isV4:
res, err = d.db.Exec(
`DELETE FROM live_routes_v4 WHERE prefix = ? AND origin_asn = ? AND peer_ip = ?`,
prefix, originASN, peerIP)
case originASN == 0:
res, err = d.db.Exec(`DELETE FROM live_routes_v6 WHERE prefix = ? AND peer_ip = ?`, prefix, peerIP)
default:
res, err = d.db.Exec(
`DELETE FROM live_routes_v6 WHERE prefix = ? AND origin_asn = ? AND peer_ip = ?`,
prefix, originASN, peerIP)
}
if err != nil {
return err
}
affected, err := res.RowsAffected()
if err != nil {
return err
}
if isV4 {
d.counts.addRoutes(-int(affected), 0)
} else {
d.counts.addRoutes(0, -int(affected))
}
return nil
}
// GetPrefixDistribution returns the distribution of unique prefixes by mask length
func (d *Database) GetPrefixDistribution() (ipv4 []PrefixDistribution, ipv6 []PrefixDistribution, err error) {
return d.GetPrefixDistributionContext(context.Background())
}
// GetPrefixDistributionContext returns the distribution of unique prefixes by mask length with context support
func (d *Database) GetPrefixDistributionContext(ctx context.Context) (
ipv4 []PrefixDistribution, ipv6 []PrefixDistribution, err error) {
// IPv4 distribution - count unique prefixes from v4 table
query := `
SELECT mask_length, COUNT(DISTINCT prefix) as count
FROM live_routes_v4
GROUP BY mask_length
ORDER BY mask_length
`
rows4, err := d.db.QueryContext(ctx, query)
if err != nil {
return nil, nil, fmt.Errorf("failed to query IPv4 distribution: %w", err)
}
defer func() {
if rows4 != nil {
_ = rows4.Close()
}
}()
for rows4.Next() {
var dist PrefixDistribution
if err := rows4.Scan(&dist.MaskLength, &dist.Count); err != nil {
return nil, nil, fmt.Errorf("failed to scan IPv4 distribution: %w", err)
}
ipv4 = append(ipv4, dist)
}
// IPv6 distribution - count unique prefixes from v6 table
query = `
SELECT mask_length, COUNT(DISTINCT prefix) as count
FROM live_routes_v6
GROUP BY mask_length
ORDER BY mask_length
`
rows6, err := d.db.QueryContext(ctx, query)
if err != nil {
return nil, nil, fmt.Errorf("failed to query IPv6 distribution: %w", err)
}
defer func() {
if rows6 != nil {
_ = rows6.Close()
}
}()
for rows6.Next() {
var dist PrefixDistribution
if err := rows6.Scan(&dist.MaskLength, &dist.Count); err != nil {
return nil, nil, fmt.Errorf("failed to scan IPv6 distribution: %w", err)
}
ipv6 = append(ipv6, dist)
}
return ipv4, ipv6, nil
}
// GetLiveRouteCounts returns the count of IPv4 and IPv6 routes
func (d *Database) GetLiveRouteCounts() (ipv4Count, ipv6Count int, err error) {
return d.GetLiveRouteCountsContext(context.Background())
}
// GetLiveRouteCountsContext returns the count of IPv4 and IPv6 routes with context support
func (d *Database) GetLiveRouteCountsContext(ctx context.Context) (ipv4Count, ipv6Count int, err error) {
// Get IPv4 count from dedicated table
err = d.db.QueryRowContext(ctx, "SELECT COUNT(*) FROM live_routes_v4").Scan(&ipv4Count)
if err != nil {
return 0, 0, fmt.Errorf("failed to count IPv4 routes: %w", err)
}
// Get IPv6 count from dedicated table
err = d.db.QueryRowContext(ctx, "SELECT COUNT(*) FROM live_routes_v6").Scan(&ipv6Count)
if err != nil {
return 0, 0, fmt.Errorf("failed to count IPv6 routes: %w", err)
}
return ipv4Count, ipv6Count, nil
}
// GetASInfoForIP returns AS information for the given IP address
func (d *Database) GetASInfoForIP(ip string) (*ASInfo, error) {
return d.GetASInfoForIPContext(context.Background(), ip)
}
// GetASInfoForIPContext returns AS information for the given IP address with context support
func (d *Database) GetASInfoForIPContext(ctx context.Context, ip string) (*ASInfo, error) {
// Parse the IP to validate it
parsedIP := net.ParseIP(ip)
if parsedIP == nil {
return nil, fmt.Errorf("%w: %s", ErrInvalidIP, ip)
}
// Determine IP version
ipVersion := ipVersionV4
ipv4 := parsedIP.To4()
if ipv4 == nil {
ipVersion = ipVersionV6
}
// For IPv4, use optimized range query
if ipVersion == ipVersionV4 {
// Convert IP to 32-bit unsigned integer
ipUint := ipToUint32(ipv4)
query := `
SELECT DISTINCT lr.prefix, lr.mask_length, lr.origin_asn, lr.last_updated, a.handle, a.description
FROM live_routes_v4 lr
LEFT JOIN asns a ON a.asn = lr.origin_asn
WHERE lr.ip_start <= ? AND lr.ip_end >= ?
ORDER BY lr.mask_length DESC
LIMIT 1
`
var prefix string
var maskLength, originASN int
var lastUpdated time.Time
var handle, description sql.NullString
err := d.db.QueryRowContext(ctx, query, ipUint, ipUint).Scan(
&prefix, &maskLength, &originASN, &lastUpdated, &handle, &description)
if err != nil {
if err == sql.ErrNoRows {
return nil, fmt.Errorf("%w for IP %s", ErrNoRoute, ip)
}
return nil, fmt.Errorf("failed to query routes: %w", err)
}
age := time.Since(lastUpdated).Round(time.Second).String()
return &ASInfo{
ASN: originASN,
Handle: handle.String,
Description: description.String,
Prefix: prefix,
LastUpdated: lastUpdated,
Age: age,
}, nil
}
// For IPv6, use the original method since we don't have range optimization
query := `
SELECT DISTINCT lr.prefix, lr.mask_length, lr.origin_asn, lr.last_updated, a.handle, a.description
FROM live_routes_v6 lr
LEFT JOIN asns a ON a.asn = lr.origin_asn
ORDER BY lr.mask_length DESC
`
rows, err := d.db.QueryContext(ctx, query)
if err != nil {
return nil, fmt.Errorf("failed to query routes: %w", err)
}
defer func() { _ = rows.Close() }()
// Find the most specific matching prefix
var bestMatch struct {
prefix string
maskLength int
originASN int
lastUpdated time.Time
handle sql.NullString
description sql.NullString
}
bestMaskLength := -1
for rows.Next() {
var prefix string
var maskLength, originASN int
var lastUpdated time.Time
var handle, description sql.NullString
if err := rows.Scan(&prefix, &maskLength, &originASN, &lastUpdated, &handle, &description); err != nil {
continue
}
// Parse the prefix CIDR
_, ipNet, err := net.ParseCIDR(prefix)
if err != nil {
continue
}
// Check if the IP is in this prefix
if ipNet.Contains(parsedIP) && maskLength > bestMaskLength {
bestMatch.prefix = prefix
bestMatch.maskLength = maskLength
bestMatch.originASN = originASN
bestMatch.lastUpdated = lastUpdated
bestMatch.handle = handle
bestMatch.description = description
bestMaskLength = maskLength
}
}
if bestMaskLength == -1 {
return nil, fmt.Errorf("%w for IP %s", ErrNoRoute, ip)
}
age := time.Since(bestMatch.lastUpdated).Round(time.Second).String()
return &ASInfo{
ASN: bestMatch.originASN,
Handle: bestMatch.handle.String,
Description: bestMatch.description.String,
Prefix: bestMatch.prefix,
LastUpdated: bestMatch.lastUpdated,
Age: age,
}, nil
}
// ipToUint32 converts an IPv4 address to a 32-bit unsigned integer
func ipToUint32(ip net.IP) uint32 {
if len(ip) == ipv6Length {
// Convert to 4-byte representation
ip = ip[ipv4Offset:ipv6Length]
}
return uint32(ip[0])<<24 | uint32(ip[1])<<16 | uint32(ip[2])<<8 | uint32(ip[3])
}
// CalculateIPv4Range calculates the start and end IP addresses for an IPv4 CIDR block
func CalculateIPv4Range(cidr string) (start, end uint32, err error) {
_, ipNet, err := net.ParseCIDR(cidr)
if err != nil {
return 0, 0, err
}
// Get the network address (start of range)
ip := ipNet.IP.To4()
if ip == nil {
return 0, 0, fmt.Errorf("not an IPv4 address")
}
start = ipToUint32(ip)
// Calculate the end of the range
ones, bits := ipNet.Mask.Size()
hostBits := bits - ones
if hostBits >= ipv4Bits {
// Special case for /0 - entire IPv4 space
end = maxIPv4
} else {
// Safe to convert since we checked hostBits < 32
//nolint:gosec // hostBits is guaranteed to be < 32 from the check above
end = start | ((1 << uint(hostBits)) - 1)
}
return start, end, nil
}
// GetASDetails returns detailed information about an AS including prefixes
func (d *Database) GetASDetails(asn int) (*ASN, []LiveRoute, error) {
return d.GetASDetailsContext(context.Background(), asn)
}
// GetASDetailsContext returns detailed information about an AS including prefixes with context support
func (d *Database) GetASDetailsContext(ctx context.Context, asn int) (*ASN, []LiveRoute, error) {
// Get AS information
var asnInfo ASN
var handle, description sql.NullString
err := d.db.QueryRowContext(ctx,
"SELECT asn, handle, description, first_seen, last_seen FROM asns WHERE asn = ?",
asn,
).Scan(&asnInfo.ASN, &handle, &description, &asnInfo.FirstSeen, &asnInfo.LastSeen)
if err != nil {
if err == sql.ErrNoRows {
return nil, nil, fmt.Errorf("%w: AS%d", ErrNoRoute, asn)
}
return nil, nil, fmt.Errorf("failed to query AS: %w", err)
}
asnInfo.Handle = handle.String
asnInfo.Description = description.String
// Get prefixes announced by this AS from both tables
var allPrefixes []LiveRoute
// Query IPv4 prefixes
queryV4 := `
SELECT prefix, mask_length, MAX(last_updated) as last_updated
FROM live_routes_v4
WHERE origin_asn = ?
GROUP BY prefix, mask_length
`
rows4, err := d.db.QueryContext(ctx, queryV4, asn)
if err != nil {
return &asnInfo, nil, fmt.Errorf("failed to query IPv4 prefixes: %w", err)
}
defer func() { _ = rows4.Close() }()
for rows4.Next() {
var route LiveRoute
var lastUpdatedStr string
err := rows4.Scan(&route.Prefix, &route.MaskLength, &lastUpdatedStr)
if err != nil {
d.logger.Error("Failed to scan IPv4 prefix row", "error", err, "asn", asn)
continue
}
// Parse the timestamp string
route.LastUpdated, err = time.Parse("2006-01-02 15:04:05-07:00", lastUpdatedStr)
if err != nil {
// Try without timezone
route.LastUpdated, err = time.Parse("2006-01-02 15:04:05", lastUpdatedStr)
if err != nil {
d.logger.Error("Failed to parse timestamp", "error", err, "timestamp", lastUpdatedStr)
continue
}
}
route.OriginASN = asn
route.IPVersion = ipVersionV4
allPrefixes = append(allPrefixes, route)
}
// Query IPv6 prefixes
queryV6 := `
SELECT prefix, mask_length, MAX(last_updated) as last_updated
FROM live_routes_v6
WHERE origin_asn = ?
GROUP BY prefix, mask_length
`
rows6, err := d.db.QueryContext(ctx, queryV6, asn)
if err != nil {
return &asnInfo, allPrefixes, fmt.Errorf("failed to query IPv6 prefixes: %w", err)
}
defer func() { _ = rows6.Close() }()
for rows6.Next() {
var route LiveRoute
var lastUpdatedStr string
err := rows6.Scan(&route.Prefix, &route.MaskLength, &lastUpdatedStr)
if err != nil {
d.logger.Error("Failed to scan IPv6 prefix row", "error", err, "asn", asn)
continue
}
// Parse the timestamp string
route.LastUpdated, err = time.Parse("2006-01-02 15:04:05-07:00", lastUpdatedStr)
if err != nil {
// Try without timezone
route.LastUpdated, err = time.Parse("2006-01-02 15:04:05", lastUpdatedStr)
if err != nil {
d.logger.Error("Failed to parse timestamp", "error", err, "timestamp", lastUpdatedStr)
continue
}
}
route.OriginASN = asn
route.IPVersion = ipVersionV6
allPrefixes = append(allPrefixes, route)
}
return &asnInfo, allPrefixes, nil
}
// ASPeer represents a peering relationship with another AS including handle, description, and timestamps.
type ASPeer struct {
ASN int `json:"asn"`
Handle string `json:"handle"`
Description string `json:"description"`
FirstSeen time.Time `json:"first_seen"`
LastSeen time.Time `json:"last_seen"`
}
// GetASPeers returns all ASes that peer with the given AS
func (d *Database) GetASPeers(asn int) ([]ASPeer, error) {
return d.GetASPeersContext(context.Background(), asn)
}
// GetASPeersContext returns all ASes that peer with the given AS with context support
func (d *Database) GetASPeersContext(ctx context.Context, asn int) ([]ASPeer, error) {
query := `
SELECT
CASE
WHEN p.as_a = ? THEN p.as_b
ELSE p.as_a
END as peer_asn,
COALESCE(a.handle, '') as handle,
COALESCE(a.description, '') as description,
p.first_seen,
p.last_seen
FROM peerings p
LEFT JOIN asns a ON a.asn = CASE
WHEN p.as_a = ? THEN p.as_b
ELSE p.as_a
END
WHERE p.as_a = ? OR p.as_b = ?
ORDER BY peer_asn
`
rows, err := d.db.QueryContext(ctx, query, asn, asn, asn, asn)
if err != nil {
return nil, fmt.Errorf("failed to query AS peers: %w", err)
}
defer func() { _ = rows.Close() }()
var peers []ASPeer
for rows.Next() {
var peer ASPeer
err := rows.Scan(&peer.ASN, &peer.Handle, &peer.Description, &peer.FirstSeen, &peer.LastSeen)
if err != nil {
d.logger.Error("Failed to scan peer row", "error", err, "asn", asn)
continue
}
peers = append(peers, peer)
}
if err := rows.Err(); err != nil {
return nil, fmt.Errorf("error iterating AS peers: %w", err)
}
return peers, nil
}
// GetPrefixDetails returns detailed information about a prefix
func (d *Database) GetPrefixDetails(prefix string) ([]LiveRoute, error) {
return d.GetPrefixDetailsContext(context.Background(), prefix)
}
// GetPrefixDetailsContext returns detailed information about a prefix with context support
func (d *Database) GetPrefixDetailsContext(ctx context.Context, prefix string) ([]LiveRoute, error) {
// Determine if it's IPv4 or IPv6 by parsing the prefix
_, ipnet, err := net.ParseCIDR(prefix)
if err != nil {
return nil, fmt.Errorf("invalid prefix format: %w", err)
}
tableName := "live_routes_v4"
ipVersion := ipVersionV4
if ipnet.IP.To4() == nil {
tableName = "live_routes_v6"
ipVersion = ipVersionV6
}
//nolint:gosec // Table name is hardcoded based on IP version
query := fmt.Sprintf(`
SELECT lr.origin_asn, lr.peer_ip, lr.as_path, lr.next_hop, lr.last_updated,
a.handle, a.description
FROM %s lr
LEFT JOIN asns a ON a.asn = lr.origin_asn
WHERE lr.prefix = ?
ORDER BY lr.origin_asn, lr.peer_ip
`, tableName)
rows, err := d.db.QueryContext(ctx, query, prefix)
if err != nil {
return nil, fmt.Errorf("failed to query prefix details: %w", err)
}
defer func() { _ = rows.Close() }()
var routes []LiveRoute
for rows.Next() {
var route LiveRoute
var pathJSON string
var handle, description sql.NullString
err := rows.Scan(
&route.OriginASN, &route.PeerIP, &pathJSON, &route.NextHop,
&route.LastUpdated, &handle, &description,
)
if err != nil {
continue
}
// Decode AS path
if err := json.Unmarshal([]byte(pathJSON), &route.ASPath); err != nil {
route.ASPath = []int{}
}
route.Prefix = prefix
route.IPVersion = ipVersion
route.MaskLength, _ = ipnet.Mask.Size()
routes = append(routes, route)
}
if len(routes) == 0 {
return nil, fmt.Errorf("%w: %s", ErrNoRoute, prefix)
}
return routes, nil
}
// GetRandomPrefixesByLength returns a random sample of prefixes with the specified mask length
func (d *Database) GetRandomPrefixesByLength(maskLength, ipVersion, limit int) ([]LiveRoute, error) {
return d.GetRandomPrefixesByLengthContext(context.Background(), maskLength, ipVersion, limit)
}
// GetRandomPrefixesByLengthContext returns a random sample of prefixes with context support
func (d *Database) GetRandomPrefixesByLengthContext(
ctx context.Context, maskLength, ipVersion, limit int) ([]LiveRoute, error) {
// Select unique prefixes with their most recent route information
tableName := "live_routes_v4"
if ipVersion == ipVersionV6 {
tableName = "live_routes_v6"
}
//nolint:gosec // Table name is hardcoded based on IP version
query := fmt.Sprintf(`
WITH unique_prefixes AS (
SELECT prefix, MAX(last_updated) as max_updated
FROM %s
WHERE mask_length = ?
GROUP BY prefix
ORDER BY RANDOM()
LIMIT ?
)
SELECT lr.prefix, lr.mask_length, lr.origin_asn, lr.as_path,
lr.peer_ip, lr.last_updated
FROM %s lr
INNER JOIN unique_prefixes up ON lr.prefix = up.prefix AND lr.last_updated = up.max_updated
WHERE lr.mask_length = ?
`, tableName, tableName)
rows, err := d.db.QueryContext(ctx, query, maskLength, limit, maskLength)
if err != nil {
return nil, fmt.Errorf("failed to query random prefixes: %w", err)
}
defer func() {
_ = rows.Close()
}()
var routes []LiveRoute
for rows.Next() {
var route LiveRoute
var pathJSON string
err := rows.Scan(
&route.Prefix,
&route.MaskLength,
&route.OriginASN,
&pathJSON,
&route.PeerIP,
&route.LastUpdated,
)
// Set IP version based on which table we queried
route.IPVersion = ipVersion
if err != nil {
continue
}
// Decode AS path
if err := json.Unmarshal([]byte(pathJSON), &route.ASPath); err != nil {
route.ASPath = []int{}
}
routes = append(routes, route)
}
return routes, nil
}
// GetNextStaleASN returns a random ASN that needs WHOIS data refresh.
func (d *Database) GetNextStaleASN(ctx context.Context, staleThreshold time.Duration) (int, error) {
cutoff := time.Now().Add(-staleThreshold)
// Select a random stale ASN using ORDER BY RANDOM()
query := `
SELECT asn FROM asns
WHERE whois_updated_at IS NULL
OR whois_updated_at < ?
ORDER BY RANDOM()
LIMIT 1
`
var asn int
err := d.db.QueryRowContext(ctx, query, cutoff).Scan(&asn)
if err != nil {
if err == sql.ErrNoRows {
return 0, ErrNoStaleASN
}
return 0, fmt.Errorf("failed to get stale ASN: %w", err)
}
return asn, nil
}
// WHOISStats contains statistics about WHOIS data freshness.
type WHOISStats struct {
TotalASNs int `json:"total_asns"`
StaleASNs int `json:"stale_asns"`
FreshASNs int `json:"fresh_asns"`
NeverFetched int `json:"never_fetched"`
}
// GetWHOISStats returns statistics about WHOIS data freshness.
func (d *Database) GetWHOISStats(ctx context.Context, staleThreshold time.Duration) (*WHOISStats, error) {
cutoff := time.Now().Add(-staleThreshold)
query := `
SELECT
COUNT(*) as total,
COALESCE(SUM(CASE WHEN whois_updated_at IS NULL THEN 1 ELSE 0 END), 0) as never_fetched,
COALESCE(SUM(CASE WHEN whois_updated_at IS NOT NULL AND whois_updated_at < ? THEN 1 ELSE 0 END), 0) as stale,
COALESCE(SUM(CASE WHEN whois_updated_at IS NOT NULL AND whois_updated_at >= ? THEN 1 ELSE 0 END), 0) as fresh
FROM asns
`
var stats WHOISStats
err := d.db.QueryRowContext(ctx, query, cutoff, cutoff).Scan(
&stats.TotalASNs,
&stats.NeverFetched,
&stats.StaleASNs,
&stats.FreshASNs,
)
if err != nil {
return nil, fmt.Errorf("failed to get WHOIS stats: %w", err)
}
return &stats, nil
}
// UpdateASNWHOIS updates an ASN record with WHOIS data.
func (d *Database) UpdateASNWHOIS(ctx context.Context, update *ASNWHOISUpdate) error {
d.lock("UpdateASNWHOIS")
defer d.unlock()
query := `
UPDATE asns SET
as_name = ?,
org_name = ?,
org_id = ?,
address = ?,
country_code = ?,
abuse_email = ?,
abuse_phone = ?,
tech_email = ?,
tech_phone = ?,
rir = ?,
rir_registration_date = ?,
rir_last_modified = ?,
whois_raw = ?,
whois_updated_at = ?
WHERE asn = ?
`
_, err := d.db.ExecContext(ctx, query,
update.ASName,
update.OrgName,
update.OrgID,
update.Address,
update.CountryCode,
update.AbuseEmail,
update.AbusePhone,
update.TechEmail,
update.TechPhone,
update.RIR,
update.RIRRegDate,
update.RIRLastMod,
update.WHOISRaw,
time.Now(),
update.ASN,
)
if err != nil {
return fmt.Errorf("failed to update ASN WHOIS: %w", err)
}
return nil
}
// GetIPInfo returns comprehensive IP information for the /ip endpoint.
func (d *Database) GetIPInfo(ip string) (*IPInfo, error) {
return d.GetIPInfoContext(context.Background(), ip)
}
// GetIPInfoContext returns comprehensive IP information with context support.
func (d *Database) GetIPInfoContext(ctx context.Context, ip string) (*IPInfo, error) {
// Parse the IP to validate it
parsedIP := net.ParseIP(ip)
if parsedIP == nil {
return nil, fmt.Errorf("%w: %s", ErrInvalidIP, ip)
}
// Determine IP version
ipv4 := parsedIP.To4()
if ipv4 != nil {
return d.getIPv4Info(ctx, ip, ipv4)
}
return d.getIPv6Info(ctx, ip, parsedIP)
}
// getIPv4Info returns comprehensive IP information for an IPv4 address.
func (d *Database) getIPv4Info(ctx context.Context, ip string, ipv4 net.IP) (*IPInfo, error) {
info := &IPInfo{
IP: ip,
IPVersion: ipVersionV4,
}
ipUint := ipToUint32(ipv4)
// Get route info with peer count and prefix first_seen
query := `
SELECT
lr.prefix,
lr.mask_length,
lr.origin_asn,
lr.last_updated,
(SELECT COUNT(DISTINCT peer_ip) FROM live_routes_v4 WHERE prefix = lr.prefix) as num_peers,
p.first_seen,
a.handle,
a.description,
a.as_name,
a.org_name,
a.org_id,
a.address,
a.country_code,
a.abuse_email,
a.rir,
a.whois_updated_at
FROM live_routes_v4 lr
LEFT JOIN prefixes_v4 p ON p.prefix = lr.prefix
LEFT JOIN asns a ON a.asn = lr.origin_asn
WHERE lr.ip_start <= ? AND lr.ip_end >= ?
ORDER BY lr.mask_length DESC
LIMIT 1
`
var handle, description, asName, orgName, orgID, address, countryCode, abuseEmail, rir sql.NullString
var prefixFirstSeen sql.NullTime
var whoisUpdatedAt sql.NullTime
err := d.db.QueryRowContext(ctx, query, ipUint, ipUint).Scan(
&info.Netblock,
&info.MaskLength,
&info.ASN,
&info.LastSeen,
&info.NumPeers,
&prefixFirstSeen,
&handle,
&description,
&asName,
&orgName,
&orgID,
&address,
&countryCode,
&abuseEmail,
&rir,
&whoisUpdatedAt,
)
if err != nil {
if err == sql.ErrNoRows {
return nil, fmt.Errorf("%w for IP %s", ErrNoRoute, ip)
}
return nil, fmt.Errorf("failed to query routes: %w", err)
}
info.Handle = handle.String
info.Description = description.String
info.ASName = asName.String
info.OrgName = orgName.String
info.OrgID = orgID.String
info.Address = address.String
info.CountryCode = countryCode.String
info.AbuseEmail = abuseEmail.String
info.RIR = rir.String
if prefixFirstSeen.Valid {
info.FirstSeen = prefixFirstSeen.Time
}
// Check if WHOIS data needs refresh (never fetched or older than 30 days)
const staleThreshold = 30 * 24 * time.Hour
info.NeedsWHOISRefresh = !whoisUpdatedAt.Valid || time.Since(whoisUpdatedAt.Time) > staleThreshold
return info, nil
}
// getIPv6Info returns comprehensive IP information for an IPv6 address.
func (d *Database) getIPv6Info(ctx context.Context, ip string, parsedIP net.IP) (*IPInfo, error) {
info := &IPInfo{
IP: ip,
IPVersion: ipVersionV6,
}
// For IPv6, scan all routes and find best match
query := `
SELECT DISTINCT
lr.prefix,
lr.mask_length,
lr.origin_asn,
lr.last_updated,
a.handle,
a.description,
a.as_name,
a.org_name,
a.org_id,
a.address,
a.country_code,
a.abuse_email,
a.rir,
a.whois_updated_at
FROM live_routes_v6 lr
LEFT JOIN asns a ON a.asn = lr.origin_asn
ORDER BY lr.mask_length DESC
`
rows, err := d.db.QueryContext(ctx, query)
if err != nil {
return nil, fmt.Errorf("failed to query routes: %w", err)
}
defer func() { _ = rows.Close() }()
bestMaskLength := -1
for rows.Next() {
var prefix string
var maskLength, originASN int
var lastUpdated time.Time
var handle, description, asName, orgName, orgID, address, countryCode, abuseEmail, rir sql.NullString
var whoisUpdatedAt sql.NullTime
if err := rows.Scan(
&prefix, &maskLength, &originASN, &lastUpdated,
&handle, &description, &asName, &orgName, &orgID,
&address, &countryCode, &abuseEmail, &rir, &whoisUpdatedAt,
); err != nil {
continue
}
_, ipNet, err := net.ParseCIDR(prefix)
if err != nil {
continue
}
if ipNet.Contains(parsedIP) && maskLength > bestMaskLength {
info.Netblock = prefix
info.MaskLength = maskLength
info.ASN = originASN
info.LastSeen = lastUpdated
info.Handle = handle.String
info.Description = description.String
info.ASName = asName.String
info.OrgName = orgName.String
info.OrgID = orgID.String
info.Address = address.String
info.CountryCode = countryCode.String
info.AbuseEmail = abuseEmail.String
info.RIR = rir.String
bestMaskLength = maskLength
if !whoisUpdatedAt.Valid {
info.NeedsWHOISRefresh = true
} else {
const staleThreshold = 30 * 24 * time.Hour
info.NeedsWHOISRefresh = time.Since(whoisUpdatedAt.Time) > staleThreshold
}
}
}
if bestMaskLength == -1 {
return nil, fmt.Errorf("%w for IP %s", ErrNoRoute, ip)
}
// Get peer count and first_seen for IPv6
countQuery := `SELECT COUNT(DISTINCT peer_ip) FROM live_routes_v6 WHERE prefix = ?`
_ = d.db.QueryRowContext(ctx, countQuery, info.Netblock).Scan(&info.NumPeers)
firstSeenQuery := `SELECT first_seen FROM prefixes_v6 WHERE prefix = ?`
var prefixFirstSeen sql.NullTime
err = d.db.QueryRowContext(ctx, firstSeenQuery, info.Netblock).Scan(&prefixFirstSeen)
if err == nil && prefixFirstSeen.Valid {
info.FirstSeen = prefixFirstSeen.Time
}
return info, nil
}
// Vacuum runs incremental vacuum to reclaim unused pages without blocking writes.
// It frees up to the specified number of pages per call (0 = all freeable pages).
func (d *Database) Vacuum(ctx context.Context) error {
// Free up to 1000 pages per call (~4MB with default 4KB page size)
// This keeps each vacuum operation quick and non-blocking
const pagesToFree = 1000
_, err := d.db.ExecContext(ctx, fmt.Sprintf("PRAGMA incremental_vacuum(%d)", pagesToFree))
if err != nil {
return fmt.Errorf("failed to run incremental vacuum: %w", err)
}
return nil
}
// Analyze runs the SQLite ANALYZE command to update query planner statistics.
func (d *Database) Analyze(ctx context.Context) error {
_, err := d.db.ExecContext(ctx, "ANALYZE")
if err != nil {
return fmt.Errorf("failed to analyze database: %w", err)
}
return nil
}
// Checkpoint runs a WAL checkpoint to transfer data from the WAL to the main database.
// Uses TRUNCATE mode which blocks writers briefly but ensures complete checkpoint,
// keeping the WAL small for fast read performance.
func (d *Database) Checkpoint(ctx context.Context) error {
_, err := d.db.ExecContext(ctx, "PRAGMA wal_checkpoint(TRUNCATE)")
if err != nil {
return fmt.Errorf("failed to checkpoint WAL: %w", err)
}
return nil
}
// Ping checks if the database connection is alive with a lightweight query.
func (d *Database) Ping(ctx context.Context) error {
var result int
err := d.db.QueryRowContext(ctx, "SELECT 1").Scan(&result)
if err != nil {
return fmt.Errorf("database ping failed: %w", err)
}
return nil
}