feat(nftables): collect counters over ssh

This commit is contained in:
sagitchu
2026-06-06 18:36:22 +08:00
committed by sagit
parent 9e69e020ab
commit 03524f4a65
3 changed files with 84 additions and 12 deletions
@@ -46,6 +46,17 @@ func (m *Manager) Clear(ctx context.Context, cfg SSHConfig) error {
return m.runner.ApplyScript(ctx, cfg, script) return m.runner.ApplyScript(ctx, cfg, script)
} }
func (m *Manager) CollectCounters(ctx context.Context, cfg SSHConfig) ([]CounterSample, error) {
if err := m.ensureInitialized(); err != nil {
return nil, err
}
raw, err := m.runner.ListTableJSON(ctx, cfg)
if err != nil {
return nil, err
}
return ParseCounterSamples(raw)
}
func (m *Manager) ensureInitialized() error { func (m *Manager) ensureInitialized() error {
if m == nil || m.runner == nil { if m == nil || m.runner == nil {
return errors.New("nftables manager not initialized") return errors.New("nftables manager not initialized")
@@ -8,9 +8,11 @@ import (
) )
type fakeRunner struct { type fakeRunner struct {
scripts []string scripts []string
err error err error
testErr error testErr error
listJSON []byte
listJSONErr error
} }
func (f *fakeRunner) ApplyScript(ctx context.Context, cfg SSHConfig, script string) error { func (f *fakeRunner) ApplyScript(ctx context.Context, cfg SSHConfig, script string) error {
@@ -22,12 +24,16 @@ func (f *fakeRunner) Test(ctx context.Context, cfg SSHConfig) error {
return f.testErr return f.testErr
} }
func (f *fakeRunner) ListTableJSON(ctx context.Context, cfg SSHConfig) ([]byte, error) {
return f.listJSON, f.listJSONErr
}
func TestManagerReconcileAppliesRenderedScript(t *testing.T) { func TestManagerReconcileAppliesRenderedScript(t *testing.T) {
runner := &fakeRunner{} runner := &fakeRunner{}
manager := NewManager(runner) manager := NewManager(runner)
plan := NodePlan{ plan := NodePlan{
NodeID: 7, NodeID: 7,
Rules: []Rule{{ForwardID: 42, InPort: 24000, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp", "udp"}}}, Rules: []Rule{{ForwardID: 42, InPort: 24000, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp", "udp"}}},
} }
result, err := manager.Reconcile(context.Background(), SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"}, plan) result, err := manager.Reconcile(context.Background(), SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"}, plan)
@@ -80,6 +86,41 @@ func TestManagerTestPassesThroughRunnerError(t *testing.T) {
} }
} }
func TestManagerCollectCountersParsesRunnerTableJSON(t *testing.T) {
runner := &fakeRunner{listJSON: []byte(`{
"nftables": [
{"rule": {
"family": "inet",
"table": "flvx",
"chain": "forward",
"comment": "flvx forward:77 to-target tcp",
"expr": [
{"counter": {"packets": 3, "bytes": 2048}}
]
}}
]
}`)}
manager := NewManager(runner)
samples, err := manager.CollectCounters(context.Background(), SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"})
if err != nil {
t.Fatalf("CollectCounters: %v", err)
}
if len(samples) != 1 {
t.Fatalf("expected 1 sample, got %d: %+v", len(samples), samples)
}
want := CounterSample{
ForwardID: 77,
Direction: CounterDirectionToTarget,
Protocol: "tcp",
Bytes: 2048,
Packets: 3,
}
if samples[0] != want {
t.Fatalf("expected %+v, got %+v", want, samples[0])
}
}
func TestManagerMethodsRequireInitializedRunner(t *testing.T) { func TestManagerMethodsRequireInitializedRunner(t *testing.T) {
cfg := SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"} cfg := SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"}
plan := NodePlan{NodeID: 7} plan := NodePlan{NodeID: 7}
@@ -98,6 +139,10 @@ func TestManagerMethodsRequireInitializedRunner(t *testing.T) {
t.Fatalf("expected not initialized error from nil manager Clear, got %v", err) t.Fatalf("expected not initialized error from nil manager Clear, got %v", err)
} }
if _, err := nilManager.CollectCounters(context.Background(), cfg); err == nil || err.Error() != expected.Error() {
t.Fatalf("expected not initialized error from nil manager CollectCounters, got %v", err)
}
manager := &Manager{} manager := &Manager{}
if err := manager.Test(context.Background(), cfg); err == nil || err.Error() != expected.Error() { if err := manager.Test(context.Background(), cfg); err == nil || err.Error() != expected.Error() {
t.Fatalf("expected not initialized error from Test, got %v", err) t.Fatalf("expected not initialized error from Test, got %v", err)
@@ -110,4 +155,8 @@ func TestManagerMethodsRequireInitializedRunner(t *testing.T) {
if err := manager.Clear(context.Background(), cfg); err == nil || err.Error() != expected.Error() { if err := manager.Clear(context.Background(), cfg); err == nil || err.Error() != expected.Error() {
t.Fatalf("expected not initialized error from Clear, got %v", err) t.Fatalf("expected not initialized error from Clear, got %v", err)
} }
if _, err := manager.CollectCounters(context.Background(), cfg); err == nil || err.Error() != expected.Error() {
t.Fatalf("expected not initialized error from CollectCounters, got %v", err)
}
} }
+20 -8
View File
@@ -14,6 +14,7 @@ import (
type Runner interface { type Runner interface {
ApplyScript(ctx context.Context, cfg SSHConfig, script string) error ApplyScript(ctx context.Context, cfg SSHConfig, script string) error
Test(ctx context.Context, cfg SSHConfig) error Test(ctx context.Context, cfg SSHConfig) error
ListTableJSON(ctx context.Context, cfg SSHConfig) ([]byte, error)
} }
type SSHRunner struct { type SSHRunner struct {
@@ -44,7 +45,16 @@ func (r *SSHRunner) ApplyScript(ctx context.Context, cfg SSHConfig, script strin
return r.run(ctx, cfg, command) return r.run(ctx, cfg, command)
} }
func (r *SSHRunner) ListTableJSON(ctx context.Context, cfg SSHConfig) ([]byte, error) {
return r.runOutput(ctx, cfg, nftBinary(cfg)+" -j list table inet flvx")
}
func (r *SSHRunner) run(ctx context.Context, cfg SSHConfig, command string) error { func (r *SSHRunner) run(ctx context.Context, cfg SSHConfig, command string) error {
_, err := r.runOutput(ctx, cfg, command)
return err
}
func (r *SSHRunner) runOutput(ctx context.Context, cfg SSHConfig, command string) ([]byte, error) {
timeout := r.Timeout timeout := r.Timeout
if timeout <= 0 { if timeout <= 0 {
timeout = 15 * time.Second timeout = 15 * time.Second
@@ -54,31 +64,33 @@ func (r *SSHRunner) run(ctx context.Context, cfg SSHConfig, command string) erro
clientConfig, err := buildSSHClientConfig(cfg) clientConfig, err := buildSSHClientConfig(cfg)
if err != nil { if err != nil {
return err return nil, err
} }
addr := net.JoinHostPort(strings.TrimSpace(cfg.Host), fmt.Sprintf("%d", normalizedSSHPort(cfg.Port))) addr := net.JoinHostPort(strings.TrimSpace(cfg.Host), fmt.Sprintf("%d", normalizedSSHPort(cfg.Port)))
dialer := net.Dialer{Timeout: timeout} dialer := net.Dialer{Timeout: timeout}
conn, err := dialer.DialContext(runCtx, "tcp", addr) conn, err := dialer.DialContext(runCtx, "tcp", addr)
if err != nil { if err != nil {
return fmt.Errorf("SSH 连接失败: %w", err) return nil, fmt.Errorf("SSH 连接失败: %w", err)
} }
defer conn.Close() defer conn.Close()
sshConn, chans, reqs, err := ssh.NewClientConn(conn, addr, clientConfig) sshConn, chans, reqs, err := ssh.NewClientConn(conn, addr, clientConfig)
if err != nil { if err != nil {
return fmt.Errorf("SSH 认证失败: %w", err) return nil, fmt.Errorf("SSH 认证失败: %w", err)
} }
client := ssh.NewClient(sshConn, chans, reqs) client := ssh.NewClient(sshConn, chans, reqs)
defer client.Close() defer client.Close()
session, err := client.NewSession() session, err := client.NewSession()
if err != nil { if err != nil {
return fmt.Errorf("SSH 会话创建失败: %w", err) return nil, fmt.Errorf("SSH 会话创建失败: %w", err)
} }
defer session.Close() defer session.Close()
var stdout bytes.Buffer
var stderr bytes.Buffer var stderr bytes.Buffer
session.Stdout = &stdout
session.Stderr = &stderr session.Stderr = &stderr
done := make(chan error, 1) done := make(chan error, 1)
@@ -89,16 +101,16 @@ func (r *SSHRunner) run(ctx context.Context, cfg SSHConfig, command string) erro
select { select {
case <-runCtx.Done(): case <-runCtx.Done():
_ = session.Close() _ = session.Close()
return fmt.Errorf("SSH 命令超时: %w", runCtx.Err()) return nil, fmt.Errorf("SSH 命令超时: %w", runCtx.Err())
case err := <-done: case err := <-done:
if err != nil { if err != nil {
message := strings.TrimSpace(stderr.String()) message := strings.TrimSpace(stderr.String())
if message != "" { if message != "" {
return fmt.Errorf("远程执行失败: %s: %w", message, err) return nil, fmt.Errorf("远程执行失败: %s: %w", message, err)
} }
return fmt.Errorf("远程执行失败: %w", err) return nil, fmt.Errorf("远程执行失败: %w", err)
} }
return nil return stdout.Bytes(), nil
} }
} }