mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-09-28 07:36:38 +08:00
feat(nftables): collect counters over ssh
This commit is contained in:
@@ -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)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user