初版功能完成
This commit is contained in:
@@ -0,0 +1,364 @@
|
||||
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
|
||||
}
|
||||
Reference in New Issue
Block a user