WireGuard server with an embedded admin console
Go backend that drives kernel WireGuard over netlink (wireguard-go as the fallback), nftables NAT with MSS clamping, forwarding and buffer sysctls, SQLite for peers, users, sessions, traffic history and the audit log. React console: dashboard with live rates and usage history, peer management with QR codes and .conf downloads, disconnect, session reset, key rotation, expiry, client-supplied keys, settings, users with admin and viewer roles, two-factor authentication with recovery codes, audit log. Docker image on Alpine with compose files for bridged and host networking, CI and GHCR publish workflows, performance notes.
This commit is contained in:
@@ -0,0 +1,300 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Peer is a client device.
|
||||
type Peer struct {
|
||||
ID string
|
||||
Name string
|
||||
PublicKey string
|
||||
PrivateKey string // empty when the client generated its own key pair
|
||||
PresharedKey string
|
||||
IPv4 string // tunnel address without prefix, e.g. 10.8.0.2
|
||||
IPv6 string // may be empty
|
||||
ClientRoutes string // AllowedIPs the *client* routes into the tunnel
|
||||
DNS string // override; empty means the server default
|
||||
Keepalive int // seconds; 0 means the server default
|
||||
MTU int // 0 means the server default
|
||||
Enabled bool
|
||||
ExpiresAt time.Time
|
||||
Notes string
|
||||
CreatedAt time.Time
|
||||
UpdatedAt time.Time
|
||||
RxTotal int64
|
||||
TxTotal int64
|
||||
LastHandshake time.Time
|
||||
LastEndpoint string
|
||||
}
|
||||
|
||||
const peerCols = `id, name, public_key, private_key, preshared_key, ipv4, ipv6, client_routes, dns, keepalive, mtu, enabled, expires_at, notes, created_at, updated_at, rx_total, tx_total, last_handshake, last_endpoint`
|
||||
|
||||
func scanPeer(row interface{ Scan(...any) error }) (*Peer, error) {
|
||||
var p Peer
|
||||
var priv, psk, v6 sql.NullString
|
||||
var enabled int
|
||||
var exp, created, updated int64
|
||||
var expN, hsN sql.NullInt64
|
||||
if err := row.Scan(&p.ID, &p.Name, &p.PublicKey, &priv, &psk, &p.IPv4, &v6, &p.ClientRoutes, &p.DNS, &p.Keepalive, &p.MTU, &enabled, &expN, &p.Notes, &created, &updated, &p.RxTotal, &p.TxTotal, &hsN, &p.LastEndpoint); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
_ = exp
|
||||
p.PrivateKey = priv.String
|
||||
p.PresharedKey = psk.String
|
||||
p.IPv6 = v6.String
|
||||
p.Enabled = enabled == 1
|
||||
if expN.Valid {
|
||||
p.ExpiresAt = time.Unix(expN.Int64, 0)
|
||||
}
|
||||
p.CreatedAt = time.Unix(created, 0)
|
||||
p.UpdatedAt = time.Unix(updated, 0)
|
||||
if hsN.Valid && hsN.Int64 > 0 {
|
||||
p.LastHandshake = time.Unix(hsN.Int64, 0)
|
||||
}
|
||||
return &p, nil
|
||||
}
|
||||
|
||||
func nullStr(s string) any {
|
||||
if s == "" {
|
||||
return nil
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
func nullTime(t time.Time) any {
|
||||
if t.IsZero() {
|
||||
return nil
|
||||
}
|
||||
return t.Unix()
|
||||
}
|
||||
|
||||
func boolInt(b bool) int {
|
||||
if b {
|
||||
return 1
|
||||
}
|
||||
return 0
|
||||
}
|
||||
|
||||
// CreatePeer inserts a peer.
|
||||
func (s *Store) CreatePeer(ctx context.Context, p *Peer) error {
|
||||
now := time.Now()
|
||||
p.CreatedAt, p.UpdatedAt = now, now
|
||||
_, err := s.db.ExecContext(ctx, `INSERT INTO peers(`+peerCols+`) VALUES(?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)`,
|
||||
p.ID, p.Name, p.PublicKey, nullStr(p.PrivateKey), nullStr(p.PresharedKey), p.IPv4, nullStr(p.IPv6), p.ClientRoutes, p.DNS, p.Keepalive, p.MTU, boolInt(p.Enabled), nullTime(p.ExpiresAt), p.Notes, now.Unix(), now.Unix(), p.RxTotal, p.TxTotal, nullTime(p.LastHandshake), p.LastEndpoint)
|
||||
return err
|
||||
}
|
||||
|
||||
// UpdatePeer writes every editable column of a peer.
|
||||
func (s *Store) UpdatePeer(ctx context.Context, p *Peer) error {
|
||||
p.UpdatedAt = time.Now()
|
||||
_, err := s.db.ExecContext(ctx, `UPDATE peers SET name=?, public_key=?, private_key=?, preshared_key=?, ipv4=?, ipv6=?, client_routes=?, dns=?, keepalive=?, mtu=?, enabled=?, expires_at=?, notes=?, updated_at=? WHERE id=?`,
|
||||
p.Name, p.PublicKey, nullStr(p.PrivateKey), nullStr(p.PresharedKey), p.IPv4, nullStr(p.IPv6), p.ClientRoutes, p.DNS, p.Keepalive, p.MTU, boolInt(p.Enabled), nullTime(p.ExpiresAt), p.Notes, p.UpdatedAt.Unix(), p.ID)
|
||||
return err
|
||||
}
|
||||
|
||||
// PeerByID loads one peer.
|
||||
func (s *Store) PeerByID(ctx context.Context, id string) (*Peer, error) {
|
||||
return scanPeer(s.db.QueryRowContext(ctx, `SELECT `+peerCols+` FROM peers WHERE id = ?`, id))
|
||||
}
|
||||
|
||||
// ListPeers returns every peer, newest first.
|
||||
func (s *Store) ListPeers(ctx context.Context) ([]*Peer, error) {
|
||||
rows, err := s.db.QueryContext(ctx, `SELECT `+peerCols+` FROM peers ORDER BY created_at DESC, id`)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []*Peer
|
||||
for rows.Next() {
|
||||
p, err := scanPeer(rows)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, p)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
// DeletePeer removes a peer and its traffic history.
|
||||
func (s *Store) DeletePeer(ctx context.Context, id string) error {
|
||||
_, err := s.db.ExecContext(ctx, `DELETE FROM peers WHERE id = ?`, id)
|
||||
return err
|
||||
}
|
||||
|
||||
// UsedAddresses returns every tunnel address in use, for allocation.
|
||||
func (s *Store) UsedAddresses(ctx context.Context) (v4, v6 []string, err error) {
|
||||
rows, err := s.db.QueryContext(ctx, `SELECT ipv4, ipv6 FROM peers`)
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
for rows.Next() {
|
||||
var a string
|
||||
var b sql.NullString
|
||||
if err := rows.Scan(&a, &b); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
v4 = append(v4, a)
|
||||
if b.Valid {
|
||||
v6 = append(v6, b.String)
|
||||
}
|
||||
}
|
||||
return v4, v6, rows.Err()
|
||||
}
|
||||
|
||||
// PeerCounters is the running total the collector flushes.
|
||||
type PeerCounters struct {
|
||||
ID string
|
||||
RxTotal int64
|
||||
TxTotal int64
|
||||
LastHandshake time.Time
|
||||
LastEndpoint string
|
||||
}
|
||||
|
||||
// TrafficSample is one bucket increment.
|
||||
type TrafficSample struct {
|
||||
PeerID string
|
||||
Bucket time.Time
|
||||
Rx, Tx int64
|
||||
}
|
||||
|
||||
// FlushCounters writes peer totals and traffic buckets in one transaction.
|
||||
func (s *Store) FlushCounters(ctx context.Context, counters []PeerCounters, samples []TrafficSample) error {
|
||||
if len(counters) == 0 && len(samples) == 0 {
|
||||
return nil
|
||||
}
|
||||
tx, err := s.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer tx.Rollback()
|
||||
for _, c := range counters {
|
||||
if _, err := tx.ExecContext(ctx, `UPDATE peers SET rx_total=?, tx_total=?, last_handshake=?, last_endpoint=? WHERE id=?`, c.RxTotal, c.TxTotal, nullTime(c.LastHandshake), c.LastEndpoint, c.ID); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
for _, t := range samples {
|
||||
if t.Rx == 0 && t.Tx == 0 {
|
||||
continue
|
||||
}
|
||||
if _, err := tx.ExecContext(ctx, `INSERT INTO traffic(peer_id, bucket_start, rx, tx) VALUES(?,?,?,?) ON CONFLICT(peer_id, bucket_start) DO UPDATE SET rx = rx + excluded.rx, tx = tx + excluded.tx`, t.PeerID, t.Bucket.Unix(), t.Rx, t.Tx); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return tx.Commit()
|
||||
}
|
||||
|
||||
// TrafficPoint is one row of a usage series.
|
||||
type TrafficPoint struct {
|
||||
Bucket time.Time `json:"t"`
|
||||
Rx int64 `json:"rx"`
|
||||
Tx int64 `json:"tx"`
|
||||
}
|
||||
|
||||
// TrafficSeries returns a peer's buckets since a time; peerID "" means all
|
||||
// peers summed.
|
||||
func (s *Store) TrafficSeries(ctx context.Context, peerID string, since time.Time) ([]TrafficPoint, error) {
|
||||
var rows *sql.Rows
|
||||
var err error
|
||||
if peerID == "" {
|
||||
rows, err = s.db.QueryContext(ctx, `SELECT bucket_start, SUM(rx), SUM(tx) FROM traffic WHERE bucket_start >= ? GROUP BY bucket_start ORDER BY bucket_start`, since.Unix())
|
||||
} else {
|
||||
rows, err = s.db.QueryContext(ctx, `SELECT bucket_start, rx, tx FROM traffic WHERE peer_id = ? AND bucket_start >= ? ORDER BY bucket_start`, peerID, since.Unix())
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []TrafficPoint
|
||||
for rows.Next() {
|
||||
var b, rx, tx int64
|
||||
if err := rows.Scan(&b, &rx, &tx); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, TrafficPoint{Bucket: time.Unix(b, 0), Rx: rx, Tx: tx})
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
// PeerUsage is a per-peer total over a window.
|
||||
type PeerUsage struct {
|
||||
PeerID string `json:"peerId"`
|
||||
Rx int64 `json:"rx"`
|
||||
Tx int64 `json:"tx"`
|
||||
}
|
||||
|
||||
// UsageSince sums traffic per peer since a time.
|
||||
func (s *Store) UsageSince(ctx context.Context, since time.Time) ([]PeerUsage, error) {
|
||||
rows, err := s.db.QueryContext(ctx, `SELECT peer_id, SUM(rx), SUM(tx) FROM traffic WHERE bucket_start >= ? GROUP BY peer_id`, since.Unix())
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []PeerUsage
|
||||
for rows.Next() {
|
||||
var u PeerUsage
|
||||
if err := rows.Scan(&u.PeerID, &u.Rx, &u.Tx); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, u)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
// PruneTraffic deletes buckets older than the cutoff.
|
||||
func (s *Store) PruneTraffic(ctx context.Context, before time.Time) error {
|
||||
_, err := s.db.ExecContext(ctx, `DELETE FROM traffic WHERE bucket_start < ?`, before.Unix())
|
||||
return err
|
||||
}
|
||||
|
||||
// AuditEntry is one administrative action.
|
||||
type AuditEntry struct {
|
||||
ID int64 `json:"id"`
|
||||
At time.Time `json:"at"`
|
||||
Actor string `json:"actor"`
|
||||
Action string `json:"action"`
|
||||
Target string `json:"target"`
|
||||
Detail string `json:"detail"`
|
||||
IP string `json:"ip"`
|
||||
}
|
||||
|
||||
// Audit appends an entry.
|
||||
func (s *Store) Audit(ctx context.Context, e AuditEntry) error {
|
||||
if e.At.IsZero() {
|
||||
e.At = time.Now()
|
||||
}
|
||||
_, err := s.db.ExecContext(ctx, `INSERT INTO audit(at, actor, action, target, detail, ip) VALUES(?,?,?,?,?,?)`, e.At.Unix(), e.Actor, e.Action, e.Target, e.Detail, e.IP)
|
||||
return err
|
||||
}
|
||||
|
||||
// ListAudit returns the newest entries.
|
||||
func (s *Store) ListAudit(ctx context.Context, limit int) ([]AuditEntry, error) {
|
||||
rows, err := s.db.QueryContext(ctx, `SELECT id, at, actor, action, target, detail, ip FROM audit ORDER BY id DESC LIMIT ?`, limit)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
out := []AuditEntry{}
|
||||
for rows.Next() {
|
||||
var e AuditEntry
|
||||
var at int64
|
||||
if err := rows.Scan(&e.ID, &at, &e.Actor, &e.Action, &e.Target, &e.Detail, &e.IP); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
e.At = time.Unix(at, 0)
|
||||
out = append(out, e)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
// PruneAudit keeps the newest n entries.
|
||||
func (s *Store) PruneAudit(ctx context.Context, keep int) error {
|
||||
_, err := s.db.ExecContext(ctx, `DELETE FROM audit WHERE id NOT IN (SELECT id FROM audit ORDER BY id DESC LIMIT ?)`, keep)
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,157 @@
|
||||
// Package store is the SQLite persistence layer. Everything WGX remembers --
|
||||
// peers and their keys, admin users, sessions, traffic history and the audit
|
||||
// log -- lives in one file under the data directory.
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
)
|
||||
|
||||
// Store wraps the database handle.
|
||||
type Store struct {
|
||||
db *sql.DB
|
||||
}
|
||||
|
||||
// Open opens (creating if needed) the database at path and applies the schema.
|
||||
func Open(path string) (*Store, error) {
|
||||
if path != ":memory:" {
|
||||
if err := os.MkdirAll(filepath.Dir(path), 0o700); err != nil {
|
||||
return nil, fmt.Errorf("create data directory: %w", err)
|
||||
}
|
||||
}
|
||||
dsn := path
|
||||
if path != ":memory:" {
|
||||
// The file holds private keys, so it is created unreadable to anyone
|
||||
// but the owner. `_pragma` options ride along in the DSN.
|
||||
f, err := os.OpenFile(path, os.O_RDWR|os.O_CREATE, 0o600)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("open database: %w", err)
|
||||
}
|
||||
f.Close()
|
||||
dsn = "file:" + path + "?_pragma=journal_mode(WAL)&_pragma=busy_timeout(5000)&_pragma=foreign_keys(ON)&_pragma=synchronous(NORMAL)"
|
||||
} else {
|
||||
dsn = "file::memory:?cache=shared&_pragma=foreign_keys(ON)"
|
||||
}
|
||||
db, err := sql.Open("sqlite", dsn)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// One connection: SQLite serialises writers anyway and a single handle
|
||||
// avoids "database is locked" surprises under WAL with the pure-Go driver.
|
||||
db.SetMaxOpenConns(1)
|
||||
s := &Store{db: db}
|
||||
if err := s.migrate(context.Background()); err != nil {
|
||||
db.Close()
|
||||
return nil, err
|
||||
}
|
||||
return s, nil
|
||||
}
|
||||
|
||||
// Close closes the database.
|
||||
func (s *Store) Close() error { return s.db.Close() }
|
||||
|
||||
// DB exposes the handle for the rare caller that needs raw SQL (tests).
|
||||
func (s *Store) DB() *sql.DB { return s.db }
|
||||
|
||||
const schema = `
|
||||
CREATE TABLE IF NOT EXISTS settings (
|
||||
key TEXT PRIMARY KEY,
|
||||
value TEXT NOT NULL
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS users (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
username TEXT NOT NULL UNIQUE COLLATE NOCASE,
|
||||
password_hash TEXT NOT NULL,
|
||||
role TEXT NOT NULL DEFAULT 'admin',
|
||||
totp_secret TEXT,
|
||||
totp_enabled INTEGER NOT NULL DEFAULT 0,
|
||||
created_at INTEGER NOT NULL,
|
||||
last_login_at INTEGER
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS recovery_codes (
|
||||
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||
code_hash TEXT NOT NULL,
|
||||
used_at INTEGER
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS recovery_codes_user ON recovery_codes(user_id);
|
||||
CREATE TABLE IF NOT EXISTS sessions (
|
||||
token_hash TEXT PRIMARY KEY,
|
||||
user_id INTEGER NOT NULL REFERENCES users(id) ON DELETE CASCADE,
|
||||
created_at INTEGER NOT NULL,
|
||||
last_seen_at INTEGER NOT NULL,
|
||||
expires_at INTEGER NOT NULL,
|
||||
ip TEXT NOT NULL DEFAULT '',
|
||||
user_agent TEXT NOT NULL DEFAULT '',
|
||||
totp_pending INTEGER NOT NULL DEFAULT 0
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS sessions_user ON sessions(user_id);
|
||||
CREATE TABLE IF NOT EXISTS peers (
|
||||
id TEXT PRIMARY KEY,
|
||||
name TEXT NOT NULL,
|
||||
public_key TEXT NOT NULL UNIQUE,
|
||||
private_key TEXT,
|
||||
preshared_key TEXT,
|
||||
ipv4 TEXT NOT NULL UNIQUE,
|
||||
ipv6 TEXT UNIQUE,
|
||||
client_routes TEXT NOT NULL,
|
||||
dns TEXT NOT NULL DEFAULT '',
|
||||
keepalive INTEGER NOT NULL DEFAULT 0,
|
||||
mtu INTEGER NOT NULL DEFAULT 0,
|
||||
enabled INTEGER NOT NULL DEFAULT 1,
|
||||
expires_at INTEGER,
|
||||
notes TEXT NOT NULL DEFAULT '',
|
||||
created_at INTEGER NOT NULL,
|
||||
updated_at INTEGER NOT NULL,
|
||||
rx_total INTEGER NOT NULL DEFAULT 0,
|
||||
tx_total INTEGER NOT NULL DEFAULT 0,
|
||||
last_handshake INTEGER,
|
||||
last_endpoint TEXT NOT NULL DEFAULT ''
|
||||
);
|
||||
CREATE TABLE IF NOT EXISTS traffic (
|
||||
peer_id TEXT NOT NULL REFERENCES peers(id) ON DELETE CASCADE,
|
||||
bucket_start INTEGER NOT NULL,
|
||||
rx INTEGER NOT NULL DEFAULT 0,
|
||||
tx INTEGER NOT NULL DEFAULT 0,
|
||||
PRIMARY KEY (peer_id, bucket_start)
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS traffic_bucket ON traffic(bucket_start);
|
||||
CREATE TABLE IF NOT EXISTS audit (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
at INTEGER NOT NULL,
|
||||
actor TEXT NOT NULL,
|
||||
action TEXT NOT NULL,
|
||||
target TEXT NOT NULL DEFAULT '',
|
||||
detail TEXT NOT NULL DEFAULT '',
|
||||
ip TEXT NOT NULL DEFAULT ''
|
||||
);
|
||||
CREATE INDEX IF NOT EXISTS audit_at ON audit(at);
|
||||
`
|
||||
|
||||
func (s *Store) migrate(ctx context.Context) error {
|
||||
if _, err := s.db.ExecContext(ctx, schema); err != nil {
|
||||
return fmt.Errorf("apply schema: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetSetting returns the raw value of a key, or "" when unset.
|
||||
func (s *Store) GetSetting(ctx context.Context, key string) (string, error) {
|
||||
var v string
|
||||
err := s.db.QueryRowContext(ctx, `SELECT value FROM settings WHERE key = ?`, key).Scan(&v)
|
||||
if err == sql.ErrNoRows {
|
||||
return "", nil
|
||||
}
|
||||
return v, err
|
||||
}
|
||||
|
||||
// SetSetting writes a key.
|
||||
func (s *Store) SetSetting(ctx context.Context, key, value string) error {
|
||||
_, err := s.db.ExecContext(ctx, `INSERT INTO settings(key, value) VALUES(?, ?) ON CONFLICT(key) DO UPDATE SET value = excluded.value`, key, value)
|
||||
return err
|
||||
}
|
||||
@@ -0,0 +1,239 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ErrNotFound is returned when a row does not exist.
|
||||
var ErrNotFound = errors.New("not found")
|
||||
|
||||
// User is an administrator account.
|
||||
type User struct {
|
||||
ID int64
|
||||
Username string
|
||||
PasswordHash string
|
||||
Role string
|
||||
TOTPSecret string
|
||||
TOTPEnabled bool
|
||||
CreatedAt time.Time
|
||||
LastLoginAt time.Time
|
||||
}
|
||||
|
||||
func scanUser(row interface{ Scan(...any) error }) (*User, error) {
|
||||
var u User
|
||||
var secret sql.NullString
|
||||
var totp int
|
||||
var created int64
|
||||
var last sql.NullInt64
|
||||
if err := row.Scan(&u.ID, &u.Username, &u.PasswordHash, &u.Role, &secret, &totp, &created, &last); err != nil {
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
return nil, err
|
||||
}
|
||||
u.TOTPSecret = secret.String
|
||||
u.TOTPEnabled = totp == 1
|
||||
u.CreatedAt = time.Unix(created, 0)
|
||||
if last.Valid {
|
||||
u.LastLoginAt = time.Unix(last.Int64, 0)
|
||||
}
|
||||
return &u, nil
|
||||
}
|
||||
|
||||
const userCols = `id, username, password_hash, role, totp_secret, totp_enabled, created_at, last_login_at`
|
||||
|
||||
// CountUsers returns how many users exist; zero means first-run setup is due.
|
||||
func (s *Store) CountUsers(ctx context.Context) (int, error) {
|
||||
var n int
|
||||
err := s.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM users`).Scan(&n)
|
||||
return n, err
|
||||
}
|
||||
|
||||
// CreateUser inserts a user and returns it.
|
||||
func (s *Store) CreateUser(ctx context.Context, username, passwordHash, role string) (*User, error) {
|
||||
now := time.Now().Unix()
|
||||
res, err := s.db.ExecContext(ctx, `INSERT INTO users(username, password_hash, role, created_at) VALUES(?, ?, ?, ?)`, username, passwordHash, role, now)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
id, _ := res.LastInsertId()
|
||||
return s.UserByID(ctx, id)
|
||||
}
|
||||
|
||||
// UserByID looks a user up by id.
|
||||
func (s *Store) UserByID(ctx context.Context, id int64) (*User, error) {
|
||||
return scanUser(s.db.QueryRowContext(ctx, `SELECT `+userCols+` FROM users WHERE id = ?`, id))
|
||||
}
|
||||
|
||||
// UserByName looks a user up by username (case-insensitive).
|
||||
func (s *Store) UserByName(ctx context.Context, name string) (*User, error) {
|
||||
return scanUser(s.db.QueryRowContext(ctx, `SELECT `+userCols+` FROM users WHERE username = ?`, name))
|
||||
}
|
||||
|
||||
// ListUsers returns every user ordered by username.
|
||||
func (s *Store) ListUsers(ctx context.Context) ([]*User, error) {
|
||||
rows, err := s.db.QueryContext(ctx, `SELECT `+userCols+` FROM users ORDER BY username`)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer rows.Close()
|
||||
var out []*User
|
||||
for rows.Next() {
|
||||
u, err := scanUser(rows)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
out = append(out, u)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
// SetPassword replaces a user's password hash.
|
||||
func (s *Store) SetPassword(ctx context.Context, id int64, hash string) error {
|
||||
_, err := s.db.ExecContext(ctx, `UPDATE users SET password_hash = ? WHERE id = ?`, hash, id)
|
||||
return err
|
||||
}
|
||||
|
||||
// SetRole changes a user's role.
|
||||
func (s *Store) SetRole(ctx context.Context, id int64, role string) error {
|
||||
_, err := s.db.ExecContext(ctx, `UPDATE users SET role = ? WHERE id = ?`, role, id)
|
||||
return err
|
||||
}
|
||||
|
||||
// SetTOTP stores a secret and whether it is active. An empty secret clears it.
|
||||
func (s *Store) SetTOTP(ctx context.Context, id int64, secret string, enabled bool) error {
|
||||
var sec any
|
||||
if secret != "" {
|
||||
sec = secret
|
||||
}
|
||||
en := 0
|
||||
if enabled {
|
||||
en = 1
|
||||
}
|
||||
_, err := s.db.ExecContext(ctx, `UPDATE users SET totp_secret = ?, totp_enabled = ? WHERE id = ?`, sec, en, id)
|
||||
return err
|
||||
}
|
||||
|
||||
// TouchLogin records a successful login.
|
||||
func (s *Store) TouchLogin(ctx context.Context, id int64) error {
|
||||
_, err := s.db.ExecContext(ctx, `UPDATE users SET last_login_at = ? WHERE id = ?`, time.Now().Unix(), id)
|
||||
return err
|
||||
}
|
||||
|
||||
// DeleteUser removes a user and, through cascades, their sessions and codes.
|
||||
func (s *Store) DeleteUser(ctx context.Context, id int64) error {
|
||||
_, err := s.db.ExecContext(ctx, `DELETE FROM users WHERE id = ?`, id)
|
||||
return err
|
||||
}
|
||||
|
||||
// ReplaceRecoveryCodes replaces a user's recovery codes with the given hashes.
|
||||
func (s *Store) ReplaceRecoveryCodes(ctx context.Context, id int64, hashes []string) error {
|
||||
tx, err := s.db.BeginTx(ctx, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer tx.Rollback()
|
||||
if _, err := tx.ExecContext(ctx, `DELETE FROM recovery_codes WHERE user_id = ?`, id); err != nil {
|
||||
return err
|
||||
}
|
||||
for _, h := range hashes {
|
||||
if _, err := tx.ExecContext(ctx, `INSERT INTO recovery_codes(user_id, code_hash) VALUES(?, ?)`, id, h); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return tx.Commit()
|
||||
}
|
||||
|
||||
// UseRecoveryCode marks a code used if it exists and is unused; it reports
|
||||
// whether it did.
|
||||
func (s *Store) UseRecoveryCode(ctx context.Context, id int64, hash string) (bool, error) {
|
||||
res, err := s.db.ExecContext(ctx, `UPDATE recovery_codes SET used_at = ? WHERE user_id = ? AND code_hash = ? AND used_at IS NULL`, time.Now().Unix(), id, hash)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
return n == 1, nil
|
||||
}
|
||||
|
||||
// RecoveryCodesLeft counts a user's unused recovery codes.
|
||||
func (s *Store) RecoveryCodesLeft(ctx context.Context, id int64) (int, error) {
|
||||
var n int
|
||||
err := s.db.QueryRowContext(ctx, `SELECT COUNT(*) FROM recovery_codes WHERE user_id = ? AND used_at IS NULL`, id).Scan(&n)
|
||||
return n, err
|
||||
}
|
||||
|
||||
// Session is one logged-in browser.
|
||||
type Session struct {
|
||||
TokenHash string
|
||||
UserID int64
|
||||
CreatedAt time.Time
|
||||
LastSeenAt time.Time
|
||||
ExpiresAt time.Time
|
||||
IP string
|
||||
UserAgent string
|
||||
TOTPPending bool
|
||||
}
|
||||
|
||||
// CreateSession stores a session.
|
||||
func (s *Store) CreateSession(ctx context.Context, sess Session) error {
|
||||
pending := 0
|
||||
if sess.TOTPPending {
|
||||
pending = 1
|
||||
}
|
||||
_, err := s.db.ExecContext(ctx, `INSERT INTO sessions(token_hash, user_id, created_at, last_seen_at, expires_at, ip, user_agent, totp_pending) VALUES(?, ?, ?, ?, ?, ?, ?, ?)`,
|
||||
sess.TokenHash, sess.UserID, sess.CreatedAt.Unix(), sess.LastSeenAt.Unix(), sess.ExpiresAt.Unix(), sess.IP, sess.UserAgent, pending)
|
||||
return err
|
||||
}
|
||||
|
||||
// SessionByHash loads a session.
|
||||
func (s *Store) SessionByHash(ctx context.Context, hash string) (*Session, error) {
|
||||
var sess Session
|
||||
var created, seen, exp int64
|
||||
var pending int
|
||||
err := s.db.QueryRowContext(ctx, `SELECT token_hash, user_id, created_at, last_seen_at, expires_at, ip, user_agent, totp_pending FROM sessions WHERE token_hash = ?`, hash).
|
||||
Scan(&sess.TokenHash, &sess.UserID, &created, &seen, &exp, &sess.IP, &sess.UserAgent, &pending)
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, ErrNotFound
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
sess.CreatedAt = time.Unix(created, 0)
|
||||
sess.LastSeenAt = time.Unix(seen, 0)
|
||||
sess.ExpiresAt = time.Unix(exp, 0)
|
||||
sess.TOTPPending = pending == 1
|
||||
return &sess, nil
|
||||
}
|
||||
|
||||
// TouchSession bumps last_seen and the sliding expiry.
|
||||
func (s *Store) TouchSession(ctx context.Context, hash string, expires time.Time) error {
|
||||
_, err := s.db.ExecContext(ctx, `UPDATE sessions SET last_seen_at = ?, expires_at = ? WHERE token_hash = ?`, time.Now().Unix(), expires.Unix(), hash)
|
||||
return err
|
||||
}
|
||||
|
||||
// ClearTOTPPending marks a session fully authenticated.
|
||||
func (s *Store) ClearTOTPPending(ctx context.Context, hash string) error {
|
||||
_, err := s.db.ExecContext(ctx, `UPDATE sessions SET totp_pending = 0 WHERE token_hash = ?`, hash)
|
||||
return err
|
||||
}
|
||||
|
||||
// DeleteSession logs one browser out.
|
||||
func (s *Store) DeleteSession(ctx context.Context, hash string) error {
|
||||
_, err := s.db.ExecContext(ctx, `DELETE FROM sessions WHERE token_hash = ?`, hash)
|
||||
return err
|
||||
}
|
||||
|
||||
// DeleteUserSessions logs a user out everywhere.
|
||||
func (s *Store) DeleteUserSessions(ctx context.Context, userID int64) error {
|
||||
_, err := s.db.ExecContext(ctx, `DELETE FROM sessions WHERE user_id = ?`, userID)
|
||||
return err
|
||||
}
|
||||
|
||||
// PruneSessions drops expired sessions.
|
||||
func (s *Store) PruneSessions(ctx context.Context) error {
|
||||
_, err := s.db.ExecContext(ctx, `DELETE FROM sessions WHERE expires_at < ?`, time.Now().Unix())
|
||||
return err
|
||||
}
|
||||
Reference in New Issue
Block a user