179 lines
5.5 KiB
Go
179 lines
5.5 KiB
Go
package db
|
|||
|
|
|
||
|
|
import (
|
||
|
|
"context"
|
||
|
|
"database/sql"
|
||
|
|
"fmt"
|
||
|
|
"time"
|
||
|
|
)
|
||
|
|
|
||
|
|
// RegisterValidator inserts a new validator or updates an existing one's
|
||
|
|
// hostname/agent_version on re-registration. It never overwrites the state
|
||
|
|
// of a validator that's mid-assignment, so an agent restarting while it
|
||
|
|
// owns an IP doesn't silently lose that ownership.
|
||
|
|
func (d *DB) RegisterValidator(ctx context.Context, validatorID, hostname, osPortID, agentVersion string) error {
|
||
|
|
now := timeToDB(Now())
|
||
|
|
_, err := d.ExecContext(ctx, `
|
||
|
|
INSERT INTO validators (validator_id, hostname, os_port_id, agent_version, state, created_at, updated_at)
|
||
|
|
VALUES (?, ?, ?, ?, ?, ?, ?)
|
||
|
|
ON CONFLICT(validator_id) DO UPDATE SET
|
||
|
|
hostname=excluded.hostname,
|
||
|
|
os_port_id=excluded.os_port_id,
|
||
|
|
agent_version=excluded.agent_version,
|
||
|
|
updated_at=excluded.updated_at
|
||
|
|
`, validatorID, hostname, osPortID, agentVersion, ValidatorIdle, now, now)
|
||
|
|
if err != nil {
|
||
|
|
return fmt.Errorf("register validator: %w", err)
|
||
|
|
}
|
||
|
|
// A brand-new row already lands in ValidatorIdle via the INSERT branch;
|
||
|
|
// a re-registering validator that was 'unregistered' or 'unreachable'
|
||
|
|
// (but not mid-assignment) should also come back to idle.
|
||
|
|
_, err = d.ExecContext(ctx, `
|
||
|
|
UPDATE validators SET state=?, updated_at=?
|
||
|
|
WHERE validator_id=? AND state IN (?, ?)
|
||
|
|
`, ValidatorIdle, now, validatorID, ValidatorUnregistered, ValidatorUnreachable)
|
||
|
|
if err != nil {
|
||
|
|
return fmt.Errorf("register validator (reactivate): %w", err)
|
||
|
|
}
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
|
||
|
|
func (d *DB) Heartbeat(ctx context.Context, validatorID string) error {
|
||
|
|
now := timeToDB(Now())
|
||
|
|
res, err := d.ExecContext(ctx, `
|
||
|
|
UPDATE validators SET last_heartbeat_at=?, updated_at=?,
|
||
|
|
state = CASE WHEN state=? THEN ? ELSE state END
|
||
|
|
WHERE validator_id=?
|
||
|
|
`, now, now, ValidatorUnreachable, ValidatorIdle, validatorID)
|
||
|
|
if err != nil {
|
||
|
|
return fmt.Errorf("heartbeat: %w", err)
|
||
|
|
}
|
||
|
|
n, _ := res.RowsAffected()
|
||
|
|
if n == 0 {
|
||
|
|
return sql.ErrNoRows
|
||
|
|
}
|
||
|
|
return nil
|
||
|
|
}
|
||
|
|
|
||
|
|
func (d *DB) GetValidator(ctx context.Context, validatorID string) (*Validator, error) {
|
||
|
|
row := d.QueryRowContext(ctx, `
|
||
|
|
SELECT validator_id, hostname, os_port_id, state, current_ip_id, agent_version,
|
||
|
|
last_heartbeat_at, created_at, updated_at
|
||
|
|
FROM validators WHERE validator_id=?
|
||
|
|
`, validatorID)
|
||
|
|
return scanValidator(row)
|
||
|
|
}
|
||
|
|
|
||
|
|
func (d *DB) ListValidators(ctx context.Context) ([]Validator, error) {
|
||
|
|
rows, err := d.QueryContext(ctx, `
|
||
|
|
SELECT validator_id, hostname, os_port_id, state, current_ip_id, agent_version,
|
||
|
|
last_heartbeat_at, created_at, updated_at
|
||
|
|
FROM validators ORDER BY validator_id
|
||
|
|
`)
|
||
|
|
if err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
defer rows.Close()
|
||
|
|
var out []Validator
|
||
|
|
for rows.Next() {
|
||
|
|
v, err := scanValidator(rows)
|
||
|
|
if err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
out = append(out, *v)
|
||
|
|
}
|
||
|
|
return out, rows.Err()
|
||
|
|
}
|
||
|
|
|
||
|
|
// ListIdleValidators returns validators currently eligible to be handed a
|
||
|
|
// new IP to work on.
|
||
|
|
func (d *DB) ListIdleValidators(ctx context.Context) ([]Validator, error) {
|
||
|
|
rows, err := d.QueryContext(ctx, `
|
||
|
|
SELECT validator_id, hostname, os_port_id, state, current_ip_id, agent_version,
|
||
|
|
last_heartbeat_at, created_at, updated_at
|
||
|
|
FROM validators WHERE state=? ORDER BY validator_id
|
||
|
|
`, ValidatorIdle)
|
||
|
|
if err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
defer rows.Close()
|
||
|
|
var out []Validator
|
||
|
|
for rows.Next() {
|
||
|
|
v, err := scanValidator(rows)
|
||
|
|
if err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
out = append(out, *v)
|
||
|
|
}
|
||
|
|
return out, rows.Err()
|
||
|
|
}
|
||
|
|
|
||
|
|
// ListStaleHeartbeats returns validators whose last heartbeat predates the
|
||
|
|
// given cutoff and that aren't already marked unreachable.
|
||
|
|
func (d *DB) ListStaleHeartbeats(ctx context.Context, cutoff time.Time) ([]Validator, error) {
|
||
|
|
rows, err := d.QueryContext(ctx, `
|
||
|
|
SELECT validator_id, hostname, os_port_id, state, current_ip_id, agent_version,
|
||
|
|
last_heartbeat_at, created_at, updated_at
|
||
|
|
FROM validators
|
||
|
|
WHERE state != ? AND last_heartbeat_at IS NOT NULL AND last_heartbeat_at < ?
|
||
|
|
`, ValidatorUnreachable, timeToDB(cutoff))
|
||
|
|
if err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
defer rows.Close()
|
||
|
|
var out []Validator
|
||
|
|
for rows.Next() {
|
||
|
|
v, err := scanValidator(rows)
|
||
|
|
if err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
out = append(out, *v)
|
||
|
|
}
|
||
|
|
return out, rows.Err()
|
||
|
|
}
|
||
|
|
|
||
|
|
func (d *DB) MarkValidatorUnreachable(ctx context.Context, validatorID string) error {
|
||
|
|
_, err := d.ExecContext(ctx, `UPDATE validators SET state=?, updated_at=? WHERE validator_id=?`,
|
||
|
|
ValidatorUnreachable, timeToDB(Now()), validatorID)
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
|
||
|
|
// FreeValidator returns a validator to idle with no assigned IP. Used after
|
||
|
|
// an IP finishes (success or failure) or is reclaimed by the lease sweep.
|
||
|
|
func (d *DB) FreeValidator(ctx context.Context, validatorID string) error {
|
||
|
|
_, err := d.ExecContext(ctx, `
|
||
|
|
UPDATE validators SET state=?, current_ip_id=NULL, updated_at=?
|
||
|
|
WHERE validator_id=?
|
||
|
|
`, ValidatorIdle, timeToDB(Now()), validatorID)
|
||
|
|
return err
|
||
|
|
}
|
||
|
|
|
||
|
|
type rowScanner interface {
|
||
|
|
Scan(dest ...interface{}) error
|
||
|
|
}
|
||
|
|
|
||
|
|
func scanValidator(row rowScanner) (*Validator, error) {
|
||
|
|
var v Validator
|
||
|
|
var currentIPID sql.NullInt64
|
||
|
|
var lastHeartbeat sql.NullString
|
||
|
|
var createdAt, updatedAt string
|
||
|
|
if err := row.Scan(&v.ValidatorID, &v.Hostname, &v.OSPortID, &v.State, ¤tIPID,
|
||
|
|
&v.AgentVersion, &lastHeartbeat, &createdAt, &updatedAt); err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
if currentIPID.Valid {
|
||
|
|
v.CurrentIPID = ¤tIPID.Int64
|
||
|
|
}
|
||
|
|
hb, err := nullStringToTimePtr(lastHeartbeat)
|
||
|
|
if err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
v.LastHeartbeatAt = hb
|
||
|
|
if v.CreatedAt, err = dbToTime(createdAt); err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
if v.UpdatedAt, err = dbToTime(updatedAt); err != nil {
|
||
|
|
return nil, err
|
||
|
|
}
|
||
|
|
return &v, nil
|
||
|
|
}
|