247 lines
5.9 KiB
Go
247 lines
5.9 KiB
Go
package fleet
|
|
|
|
import (
|
|
"context"
|
|
"crypto/rand"
|
|
"encoding/hex"
|
|
"fmt"
|
|
"net"
|
|
"os/exec"
|
|
"strings"
|
|
"sync"
|
|
"time"
|
|
)
|
|
|
|
const immuneFailureThreshold = 5
|
|
const immunePauseDuration = 24 * time.Hour
|
|
|
|
// SubnetMapper sweeps operator-declared CIDRs with /24 immune pause after failures.
|
|
type SubnetMapper struct {
|
|
Store *Store
|
|
mu sync.Mutex
|
|
}
|
|
|
|
// SubnetCIDR is an operator-declared scan target.
|
|
type SubnetCIDR struct {
|
|
ID string `json:"id"`
|
|
CIDR string `json:"cidr"`
|
|
Enabled bool `json:"enabled"`
|
|
LastScanAt *time.Time `json:"last_scan_at,omitempty"`
|
|
}
|
|
|
|
// AddCIDR registers a CIDR for incremental discovery.
|
|
func (m *SubnetMapper) AddCIDR(ctx context.Context, cidr string) (*SubnetCIDR, error) {
|
|
if _, _, err := net.ParseCIDR(cidr); err != nil {
|
|
return nil, fmt.Errorf("invalid cidr: %w", err)
|
|
}
|
|
id := randomHex(16)
|
|
_, err := m.Store.DB().ExecContext(ctx, `
|
|
INSERT INTO subnet_cidrs (id, cidr, enabled) VALUES (?, ?, 1)`, id, cidr)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
return &SubnetCIDR{ID: id, CIDR: cidr, Enabled: true}, nil
|
|
}
|
|
|
|
// ListCIDRs returns declared subnets.
|
|
func (m *SubnetMapper) ListCIDRs(ctx context.Context) ([]SubnetCIDR, error) {
|
|
rows, err := m.Store.DB().QueryContext(ctx, `
|
|
SELECT id, cidr, enabled, last_scan_at FROM subnet_cidrs ORDER BY created_at`)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
defer rows.Close()
|
|
var out []SubnetCIDR
|
|
for rows.Next() {
|
|
var s SubnetCIDR
|
|
var enabled int
|
|
var lastScan sqlNullTime
|
|
if err := rows.Scan(&s.ID, &s.CIDR, &enabled, &lastScan); err != nil {
|
|
return nil, err
|
|
}
|
|
s.Enabled = enabled == 1
|
|
if lastScan.Valid {
|
|
s.LastScanAt = &lastScan.Time
|
|
}
|
|
out = append(out, s)
|
|
}
|
|
return out, rows.Err()
|
|
}
|
|
|
|
type sqlNullTime struct {
|
|
Valid bool
|
|
Time time.Time
|
|
}
|
|
|
|
func (n *sqlNullTime) Scan(src interface{}) error {
|
|
if src == nil {
|
|
n.Valid = false
|
|
return nil
|
|
}
|
|
switch v := src.(type) {
|
|
case string:
|
|
if v == "" {
|
|
n.Valid = false
|
|
return nil
|
|
}
|
|
t, err := time.Parse("2006-01-02 15:04:05", v)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
n.Time = t
|
|
n.Valid = true
|
|
case []byte:
|
|
return n.Scan(string(v))
|
|
}
|
|
return nil
|
|
}
|
|
|
|
// ScanResult holds hosts discovered in a sweep.
|
|
type ScanResult struct {
|
|
CIDR string `json:"cidr"`
|
|
Prefix string `json:"prefix_24"`
|
|
Alive []string `json:"alive"`
|
|
Skipped bool `json:"skipped"`
|
|
Reason string `json:"reason,omitempty"`
|
|
}
|
|
|
|
// SweepCIDR pings hosts in a CIDR unless /24 is immune-paused.
|
|
func (m *SubnetMapper) SweepCIDR(ctx context.Context, cidr string) (*ScanResult, error) {
|
|
ip, ipNet, err := net.ParseCIDR(cidr)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
|
|
prefix24 := prefixOf24(ip)
|
|
if paused, reason := m.isImmunePaused(ctx, prefix24); paused {
|
|
return &ScanResult{CIDR: cidr, Prefix: prefix24, Skipped: true, Reason: reason}, nil
|
|
}
|
|
|
|
alive, sweepErr := pingSweep(ctx, ipNet)
|
|
if sweepErr != nil {
|
|
_ = m.recordSubnetFailure(ctx, prefix24)
|
|
return &ScanResult{CIDR: cidr, Prefix: prefix24, Alive: alive}, sweepErr
|
|
}
|
|
|
|
_, _ = m.Store.DB().ExecContext(ctx, `
|
|
UPDATE subnet_cidrs SET last_scan_at = datetime('now') WHERE cidr = ?`, cidr)
|
|
|
|
return &ScanResult{CIDR: cidr, Prefix: prefix24, Alive: alive}, nil
|
|
}
|
|
|
|
// SweepAll runs enabled CIDRs sequentially.
|
|
func (m *SubnetMapper) SweepAll(ctx context.Context) ([]ScanResult, error) {
|
|
cidrs, err := m.ListCIDRs(ctx)
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
var results []ScanResult
|
|
for _, c := range cidrs {
|
|
if !c.Enabled {
|
|
continue
|
|
}
|
|
r, err := m.SweepCIDR(ctx, c.CIDR)
|
|
if err != nil {
|
|
return results, err
|
|
}
|
|
results = append(results, *r)
|
|
}
|
|
return results, nil
|
|
}
|
|
|
|
func prefixOf24(ip net.IP) string {
|
|
v4 := ip.To4()
|
|
if v4 == nil {
|
|
return ip.String() + "/64"
|
|
}
|
|
return fmt.Sprintf("%d.%d.%d.0/24", v4[0], v4[1], v4[2])
|
|
}
|
|
|
|
func (m *SubnetMapper) isImmunePaused(ctx context.Context, prefix string) (bool, string) {
|
|
var count int
|
|
var pausedUntil sqlNullTime
|
|
err := m.Store.DB().QueryRowContext(ctx, `
|
|
SELECT failure_count, paused_until FROM subnet_immune WHERE prefix = ?`, prefix).
|
|
Scan(&count, &pausedUntil)
|
|
if err != nil {
|
|
return false, ""
|
|
}
|
|
if pausedUntil.Valid && time.Now().Before(pausedUntil.Time) {
|
|
return true, fmt.Sprintf("/24 immune pause until %s", pausedUntil.Time.Format(time.RFC3339))
|
|
}
|
|
return false, ""
|
|
}
|
|
|
|
func (m *SubnetMapper) recordSubnetFailure(ctx context.Context, prefix string) error {
|
|
m.mu.Lock()
|
|
defer m.mu.Unlock()
|
|
|
|
var count int
|
|
_ = m.Store.DB().QueryRowContext(ctx, `SELECT failure_count FROM subnet_immune WHERE prefix = ?`, prefix).Scan(&count)
|
|
count++
|
|
|
|
pausedUntil := ""
|
|
if count >= immuneFailureThreshold {
|
|
until := time.Now().Add(immunePauseDuration)
|
|
pausedUntil = until.Format("2006-01-02 15:04:05")
|
|
}
|
|
|
|
_, err := m.Store.DB().ExecContext(ctx, `
|
|
INSERT INTO subnet_immune (prefix, failure_count, paused_until, updated_at)
|
|
VALUES (?, ?, ?, datetime('now'))
|
|
ON CONFLICT(prefix) DO UPDATE SET
|
|
failure_count = excluded.failure_count,
|
|
paused_until = excluded.paused_until,
|
|
updated_at = datetime('now')`,
|
|
prefix, count, nullStringOrNil(pausedUntil))
|
|
return err
|
|
}
|
|
|
|
func nullStringOrNil(s string) interface{} {
|
|
if s == "" {
|
|
return nil
|
|
}
|
|
return s
|
|
}
|
|
|
|
func pingSweep(ctx context.Context, ipNet *net.IPNet) ([]string, error) {
|
|
var alive []string
|
|
ip := ipNet.IP.Mask(ipNet.Mask)
|
|
for ip := incrementIP(ip); ipNet.Contains(ip); ip = incrementIP(ip) {
|
|
select {
|
|
case <-ctx.Done():
|
|
return alive, ctx.Err()
|
|
default:
|
|
}
|
|
if pingHost(ip.String()) {
|
|
alive = append(alive, ip.String())
|
|
}
|
|
}
|
|
return alive, nil
|
|
}
|
|
|
|
func pingHost(host string) bool {
|
|
out, err := exec.Command("ping", "-c", "1", "-W", "1", host).CombinedOutput()
|
|
if err != nil {
|
|
return false
|
|
}
|
|
return strings.Contains(string(out), "1 received") || strings.Contains(string(out), "1 packets received")
|
|
}
|
|
|
|
func incrementIP(ip net.IP) net.IP {
|
|
ip = ip.To16()
|
|
for i := len(ip) - 1; i >= 0; i-- {
|
|
ip[i]++
|
|
if ip[i] != 0 {
|
|
break
|
|
}
|
|
}
|
|
return ip
|
|
}
|
|
|
|
func randomHex(n int) string {
|
|
b := make([]byte, n)
|
|
_, _ = rand.Read(b)
|
|
return hex.EncodeToString(b)
|
|
}
|