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)
}
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 {
if m == nil || m.runner == nil {
return errors.New("nftables manager not initialized")
@@ -8,9 +8,11 @@ import (
)
type fakeRunner struct {
scripts []string
err error
testErr error
scripts []string
err error
testErr error
listJSON []byte
listJSONErr 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
}
func (f *fakeRunner) ListTableJSON(ctx context.Context, cfg SSHConfig) ([]byte, error) {
return f.listJSON, f.listJSONErr
}
func TestManagerReconcileAppliesRenderedScript(t *testing.T) {
runner := &fakeRunner{}
manager := NewManager(runner)
plan := NodePlan{
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)
@@ -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) {
cfg := SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"}
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)
}
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{}
if err := manager.Test(context.Background(), cfg); err == nil || err.Error() != expected.Error() {
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() {
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 {
ApplyScript(ctx context.Context, cfg SSHConfig, script string) error
Test(ctx context.Context, cfg SSHConfig) error
ListTableJSON(ctx context.Context, cfg SSHConfig) ([]byte, error)
}
type SSHRunner struct {
@@ -44,7 +45,16 @@ func (r *SSHRunner) ApplyScript(ctx context.Context, cfg SSHConfig, script strin
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 {
_, err := r.runOutput(ctx, cfg, command)
return err
}
func (r *SSHRunner) runOutput(ctx context.Context, cfg SSHConfig, command string) ([]byte, error) {
timeout := r.Timeout
if timeout <= 0 {
timeout = 15 * time.Second
@@ -54,31 +64,33 @@ func (r *SSHRunner) run(ctx context.Context, cfg SSHConfig, command string) erro
clientConfig, err := buildSSHClientConfig(cfg)
if err != nil {
return err
return nil, err
}
addr := net.JoinHostPort(strings.TrimSpace(cfg.Host), fmt.Sprintf("%d", normalizedSSHPort(cfg.Port)))
dialer := net.Dialer{Timeout: timeout}
conn, err := dialer.DialContext(runCtx, "tcp", addr)
if err != nil {
return fmt.Errorf("SSH 连接失败: %w", err)
return nil, fmt.Errorf("SSH 连接失败: %w", err)
}
defer conn.Close()
sshConn, chans, reqs, err := ssh.NewClientConn(conn, addr, clientConfig)
if err != nil {
return fmt.Errorf("SSH 认证失败: %w", err)
return nil, fmt.Errorf("SSH 认证失败: %w", err)
}
client := ssh.NewClient(sshConn, chans, reqs)
defer client.Close()
session, err := client.NewSession()
if err != nil {
return fmt.Errorf("SSH 会话创建失败: %w", err)
return nil, fmt.Errorf("SSH 会话创建失败: %w", err)
}
defer session.Close()
var stdout bytes.Buffer
var stderr bytes.Buffer
session.Stdout = &stdout
session.Stderr = &stderr
done := make(chan error, 1)
@@ -89,16 +101,16 @@ func (r *SSHRunner) run(ctx context.Context, cfg SSHConfig, command string) erro
select {
case <-runCtx.Done():
_ = session.Close()
return fmt.Errorf("SSH 命令超时: %w", runCtx.Err())
return nil, fmt.Errorf("SSH 命令超时: %w", runCtx.Err())
case err := <-done:
if err != nil {
message := strings.TrimSpace(stderr.String())
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
}
}