Files
RemLink/internal/database/nodes.go
T
qsc20001102 142e5dc7d6
ci / Go checks (ubuntu-latest) (push) Has been cancelled
ci / Go checks (windows-latest) (push) Has been cancelled
初版功能完成
2026-08-29 13:12:17 +08:00

365 lines
12 KiB
Go

package database
import (
"context"
"database/sql"
"errors"
"fmt"
"net/netip"
"time"
"remlink/internal/model"
)
// ErrNodeNotFound is returned when a requested node does not exist.
var ErrNodeNotFound = errors.New("node not found")
// Store groups repositories backed by a migrated RemLink database.
type Store struct {
db *sql.DB
}
// NewStore creates repositories over db. Open must have migrated db first.
func NewStore(db *sql.DB) *Store {
return &Store{db: db}
}
// DB exposes the underlying connection for transaction-oriented services.
func (s *Store) DB() *sql.DB { return s.db }
// CreateNode persists a newly registered node.
func (s *Store) CreateNode(ctx context.Context, node model.Node) error {
if err := validateNode(node); err != nil {
return err
}
now := time.Now().UTC()
if node.CreatedAt.IsZero() {
node.CreatedAt = now
}
if node.UpdatedAt.IsZero() {
node.UpdatedAt = node.CreatedAt
}
if node.Status == "" {
node.Status = model.NodeOffline
}
_, err := s.db.ExecContext(ctx, `
INSERT INTO nodes(
node_id, type, name, overlay_ip, wg_public_key, node_token_hash,
status, version, os_version, last_seen, created_at, updated_at
) VALUES(?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`, node.ID, node.Type, node.Name, node.OverlayIP.String(), node.WGPublicKey,
node.NodeTokenHash, node.Status, node.Version, node.OSVersion,
nullTime(node.LastSeen), formatTime(node.CreatedAt), formatTime(node.UpdatedAt))
if err != nil {
return fmt.Errorf("create node %q: %w", node.ID, err)
}
return nil
}
// GetNode returns a node by its stable UUID.
func (s *Store) GetNode(ctx context.Context, nodeID string) (model.Node, error) {
row := s.db.QueryRowContext(ctx, nodeSelect+` WHERE node_id = ?`, nodeID)
node, err := scanNode(row)
if errors.Is(err, sql.ErrNoRows) {
return model.Node{}, fmt.Errorf("%w: %s", ErrNodeNotFound, nodeID)
}
if err != nil {
return model.Node{}, fmt.Errorf("get node %q: %w", nodeID, err)
}
return node, nil
}
// ListNodes returns all nodes in stable creation order.
func (s *Store) ListNodes(ctx context.Context) ([]model.Node, error) {
rows, err := s.db.QueryContext(ctx, nodeSelect+` ORDER BY created_at, node_id`)
if err != nil {
return nil, fmt.Errorf("list nodes: %w", err)
}
defer rows.Close()
nodes := make([]model.Node, 0)
for rows.Next() {
node, err := scanNode(rows)
if err != nil {
return nil, fmt.Errorf("scan node: %w", err)
}
nodes = append(nodes, node)
}
if err := rows.Err(); err != nil {
return nil, fmt.Errorf("iterate nodes: %w", err)
}
return nodes, nil
}
// ListOverlayIPs returns every reserved node address.
func (s *Store) ListOverlayIPs(ctx context.Context) ([]netip.Addr, error) {
rows, err := s.db.QueryContext(ctx, `SELECT overlay_ip FROM nodes ORDER BY overlay_ip`)
if err != nil {
return nil, fmt.Errorf("list overlay addresses: %w", err)
}
defer rows.Close()
var addresses []netip.Addr
for rows.Next() {
var raw string
if err := rows.Scan(&raw); err != nil {
return nil, fmt.Errorf("scan overlay address: %w", err)
}
address, err := netip.ParseAddr(raw)
if err != nil || !address.Is4() {
return nil, fmt.Errorf("invalid stored overlay address %q", raw)
}
addresses = append(addresses, address)
}
if err := rows.Err(); err != nil {
return nil, fmt.Errorf("iterate overlay addresses: %w", err)
}
return addresses, nil
}
// UpdateNodeRegistration rotates the token and refreshes mutable client data.
func (s *Store) UpdateNodeRegistration(ctx context.Context, node model.Node) error {
if len(node.NodeTokenHash) == 0 {
return errors.New("node token hash must not be empty")
}
result, err := s.db.ExecContext(ctx, `
UPDATE nodes
SET name = ?, node_token_hash = ?, version = ?, os_version = ?, updated_at = ?
WHERE node_id = ?
`, node.Name, node.NodeTokenHash, node.Version, node.OSVersion,
formatTime(time.Now().UTC()), node.ID)
if err != nil {
return fmt.Errorf("update node registration %q: %w", node.ID, err)
}
return requireChanged(result, node.ID)
}
// UpdateNodeOverlayIP changes a node's authoritative allocation.
func (s *Store) UpdateNodeOverlayIP(ctx context.Context, nodeID string, address netip.Addr) error {
if !address.Is4() {
return errors.New("overlay address must be IPv4")
}
result, err := s.db.ExecContext(ctx, `
UPDATE nodes SET overlay_ip = ?, updated_at = ? WHERE node_id = ?
`, address.String(), formatTime(time.Now().UTC()), nodeID)
if err != nil {
return fmt.Errorf("update node %q overlay address: %w", nodeID, err)
}
return requireChanged(result, nodeID)
}
// UpdateNodeName applies the Admin API's bounded display-name edit.
func (s *Store) UpdateNodeName(ctx context.Context, nodeID, name string) error {
if name == "" || len(name) > 128 {
return errors.New("node name must contain 1 to 128 bytes")
}
result, err := s.db.ExecContext(ctx, `UPDATE nodes SET name = ?, updated_at = ? WHERE node_id = ?`,
name, formatTime(time.Now().UTC()), nodeID)
if err != nil {
return fmt.Errorf("update node %q name: %w", nodeID, err)
}
return requireChanged(result, nodeID)
}
// ReplaceNodeOverlayIPs atomically migrates all Node allocations. Temporary
// values avoid UNIQUE collisions when old and new pools overlap.
func (s *Store) ReplaceNodeOverlayIPs(ctx context.Context, assignments map[string]netip.Addr) error {
tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return err
}
defer tx.Rollback()
if err := replaceNodeOverlayIPs(ctx, tx, assignments); err != nil {
return err
}
return tx.Commit()
}
// ReplaceNodeOverlayIPsAndSetting commits the T18 Node allocation plan and
// authoritative network setting in one SQLite transaction. A Server crash can
// therefore observe either the complete old network or the complete new one,
// never new Node IPs paired with an old global Overlay configuration.
func (s *Store) ReplaceNodeOverlayIPsAndSetting(ctx context.Context, assignments map[string]netip.Addr, settingKey, settingValue string) error {
if settingKey == "" {
return errors.New("network setting key is required")
}
tx, err := s.db.BeginTx(ctx, nil)
if err != nil {
return err
}
defer tx.Rollback()
if err := replaceNodeOverlayIPs(ctx, tx, assignments); err != nil {
return err
}
if _, err := tx.ExecContext(ctx, `
INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)
ON CONFLICT(key) DO UPDATE SET value = excluded.value, updated_at = excluded.updated_at
`, settingKey, settingValue, formatTime(time.Now().UTC())); err != nil {
return fmt.Errorf("set migration setting %q: %w", settingKey, err)
}
return tx.Commit()
}
type nodeOverlayTransaction interface {
ExecContext(context.Context, string, ...any) (sql.Result, error)
}
func replaceNodeOverlayIPs(ctx context.Context, tx nodeOverlayTransaction, assignments map[string]netip.Addr) error {
for nodeID, address := range assignments {
if nodeID == "" || !address.Is4() {
return errors.New("Overlay migration requires Node IDs and IPv4 addresses")
}
result, err := tx.ExecContext(ctx, `UPDATE nodes SET overlay_ip = ?, updated_at = ? WHERE node_id = ?`,
"migrating:"+nodeID, formatTime(time.Now().UTC()), nodeID)
if err != nil {
return err
}
if count, _ := result.RowsAffected(); count != 1 {
return fmt.Errorf("%w: %s", ErrNodeNotFound, nodeID)
}
}
for nodeID, address := range assignments {
if _, err := tx.ExecContext(ctx, `UPDATE nodes SET overlay_ip = ?, updated_at = ? WHERE node_id = ?`,
address.String(), formatTime(time.Now().UTC()), nodeID); err != nil {
return err
}
}
return nil
}
// DeleteNode revokes a node and releases its unique address reservation.
func (s *Store) DeleteNode(ctx context.Context, nodeID string) error {
result, err := s.db.ExecContext(ctx, `DELETE FROM nodes WHERE node_id = ?`, nodeID)
if err != nil {
return fmt.Errorf("delete node %q: %w", nodeID, err)
}
return requireChanged(result, nodeID)
}
// UpdateNodeHeartbeat records liveness and current application metadata.
func (s *Store) UpdateNodeHeartbeat(ctx context.Context, nodeID string, status model.NodeStatus, at time.Time, version, osVersion string) error {
if !status.Valid() {
return fmt.Errorf("invalid node status %q", status)
}
result, err := s.db.ExecContext(ctx, `
UPDATE nodes SET status = ?, last_seen = ?, version = ?, os_version = ?, updated_at = ?
WHERE node_id = ?
`, status, formatTime(at), version, osVersion, formatTime(time.Now().UTC()), nodeID)
if err != nil {
return fmt.Errorf("update node %q heartbeat: %w", nodeID, err)
}
return requireChanged(result, nodeID)
}
// UpdateNodeStatus changes only the derived heartbeat status.
func (s *Store) UpdateNodeStatus(ctx context.Context, nodeID string, status model.NodeStatus) error {
if !status.Valid() {
return fmt.Errorf("invalid node status %q", status)
}
result, err := s.db.ExecContext(ctx, `
UPDATE nodes SET status = ?, updated_at = ? WHERE node_id = ?
`, status, formatTime(time.Now().UTC()), nodeID)
if err != nil {
return fmt.Errorf("update node %q status: %w", nodeID, err)
}
return requireChanged(result, nodeID)
}
const nodeSelect = `SELECT node_id, type, name, overlay_ip, wg_public_key,
node_token_hash, status, version, os_version, last_seen, created_at, updated_at FROM nodes`
type rowScanner interface {
Scan(dest ...any) error
}
func scanNode(scanner rowScanner) (model.Node, error) {
var (
node model.Node
nodeType string
status string
overlayIP string
lastSeen sql.NullString
createdAt string
updatedAt string
)
if err := scanner.Scan(&node.ID, &nodeType, &node.Name, &overlayIP,
&node.WGPublicKey, &node.NodeTokenHash, &status, &node.Version,
&node.OSVersion, &lastSeen, &createdAt, &updatedAt); err != nil {
return model.Node{}, err
}
parsedType, err := model.ParseNodeType(nodeType)
if err != nil {
return model.Node{}, err
}
node.Type = parsedType
node.Status = model.NodeStatus(status)
if !node.Status.Valid() {
return model.Node{}, fmt.Errorf("invalid stored node status %q", status)
}
node.OverlayIP, err = netip.ParseAddr(overlayIP)
if err != nil || !node.OverlayIP.Is4() {
return model.Node{}, fmt.Errorf("invalid stored overlay address %q", overlayIP)
}
node.CreatedAt, err = parseTime(createdAt)
if err != nil {
return model.Node{}, fmt.Errorf("parse created_at: %w", err)
}
node.UpdatedAt, err = parseTime(updatedAt)
if err != nil {
return model.Node{}, fmt.Errorf("parse updated_at: %w", err)
}
if lastSeen.Valid {
value, err := parseTime(lastSeen.String)
if err != nil {
return model.Node{}, fmt.Errorf("parse last_seen: %w", err)
}
node.LastSeen = &value
}
return node, nil
}
func validateNode(node model.Node) error {
if node.ID == "" {
return errors.New("node ID must not be empty")
}
if !node.Type.Valid() {
return fmt.Errorf("invalid node type %q", node.Type)
}
if node.Name == "" {
return errors.New("node name must not be empty")
}
if !node.OverlayIP.Is4() {
return errors.New("overlay address must be IPv4")
}
if node.WGPublicKey == "" {
return errors.New("WireGuard public key must not be empty")
}
if len(node.NodeTokenHash) == 0 {
return errors.New("node token hash must not be empty")
}
if node.Status != "" && !node.Status.Valid() {
return fmt.Errorf("invalid node status %q", node.Status)
}
return nil
}
func formatTime(value time.Time) string { return value.UTC().Format(time.RFC3339Nano) }
func parseTime(value string) (time.Time, error) { return time.Parse(time.RFC3339Nano, value) }
func nullTime(value *time.Time) any {
if value == nil {
return nil
}
return formatTime(*value)
}
func requireChanged(result sql.Result, nodeID string) error {
count, err := result.RowsAffected()
if err != nil {
return fmt.Errorf("read changed rows: %w", err)
}
if count == 0 {
return fmt.Errorf("%w: %s", ErrNodeNotFound, nodeID)
}
return nil
}