mirror of
https://github.com/Sagit-chu/flvx.git
synced 2026-10-08 18:56:37 +08:00
feat: add nftables forwarding mode
This commit is contained in:
@@ -0,0 +1,54 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
)
|
||||
|
||||
type Manager struct {
|
||||
runner Runner
|
||||
}
|
||||
|
||||
func NewManager(runner Runner) *Manager {
|
||||
if runner == nil {
|
||||
runner = NewSSHRunner()
|
||||
}
|
||||
return &Manager{runner: runner}
|
||||
}
|
||||
|
||||
func (m *Manager) Test(ctx context.Context, cfg SSHConfig) error {
|
||||
if err := m.ensureInitialized(); err != nil {
|
||||
return err
|
||||
}
|
||||
return m.runner.Test(ctx, cfg)
|
||||
}
|
||||
|
||||
func (m *Manager) Reconcile(ctx context.Context, cfg SSHConfig, plan NodePlan) (ApplyResult, error) {
|
||||
if err := m.ensureInitialized(); err != nil {
|
||||
return ApplyResult{}, err
|
||||
}
|
||||
result := ApplyResult{
|
||||
NodeID: plan.NodeID,
|
||||
Script: RenderTable(plan),
|
||||
Hashes: PlanHashes(plan),
|
||||
}
|
||||
if err := m.runner.ApplyScript(ctx, cfg, result.Script); err != nil {
|
||||
return ApplyResult{}, err
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (m *Manager) Clear(ctx context.Context, cfg SSHConfig) error {
|
||||
if err := m.ensureInitialized(); err != nil {
|
||||
return err
|
||||
}
|
||||
script := RenderTable(NodePlan{})
|
||||
return m.runner.ApplyScript(ctx, cfg, script)
|
||||
}
|
||||
|
||||
func (m *Manager) ensureInitialized() error {
|
||||
if m == nil || m.runner == nil {
|
||||
return errors.New("nftables manager not initialized")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,113 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
type fakeRunner struct {
|
||||
scripts []string
|
||||
err error
|
||||
testErr error
|
||||
}
|
||||
|
||||
func (f *fakeRunner) ApplyScript(ctx context.Context, cfg SSHConfig, script string) error {
|
||||
f.scripts = append(f.scripts, script)
|
||||
return f.err
|
||||
}
|
||||
|
||||
func (f *fakeRunner) Test(ctx context.Context, cfg SSHConfig) error {
|
||||
return f.testErr
|
||||
}
|
||||
|
||||
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"}}},
|
||||
}
|
||||
|
||||
result, err := manager.Reconcile(context.Background(), SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"}, plan)
|
||||
if err != nil {
|
||||
t.Fatalf("Reconcile: %v", err)
|
||||
}
|
||||
if len(runner.scripts) != 1 {
|
||||
t.Fatalf("expected 1 script, got %d", len(runner.scripts))
|
||||
}
|
||||
if !strings.Contains(runner.scripts[0], "flvx forward:42 tcp") {
|
||||
t.Fatalf("script missing forward comment:\n%s", runner.scripts[0])
|
||||
}
|
||||
if result.NodeID != 7 || result.Hashes[42] == "" {
|
||||
t.Fatalf("unexpected result: %+v", result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerReconcileReturnsRunnerError(t *testing.T) {
|
||||
runner := &fakeRunner{err: errors.New("ssh failed")}
|
||||
manager := NewManager(runner)
|
||||
|
||||
_, err := manager.Reconcile(context.Background(), SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"}, NodePlan{NodeID: 7})
|
||||
if !errors.Is(err, runner.err) {
|
||||
t.Fatalf("expected original runner error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerClearAppliesEmptyTable(t *testing.T) {
|
||||
runner := &fakeRunner{}
|
||||
manager := NewManager(runner)
|
||||
|
||||
if err := manager.Clear(context.Background(), SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"}); err != nil {
|
||||
t.Fatalf("Clear: %v", err)
|
||||
}
|
||||
if len(runner.scripts) != 1 {
|
||||
t.Fatalf("expected 1 script, got %d", len(runner.scripts))
|
||||
}
|
||||
if strings.Contains(runner.scripts[0], "masquerade comment") {
|
||||
t.Fatalf("empty table should not include masquerade:\n%s", runner.scripts[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerTestPassesThroughRunnerError(t *testing.T) {
|
||||
runner := &fakeRunner{testErr: errors.New("probe failed")}
|
||||
manager := NewManager(runner)
|
||||
|
||||
err := manager.Test(context.Background(), SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"})
|
||||
if !errors.Is(err, runner.testErr) {
|
||||
t.Fatalf("expected original runner error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestManagerMethodsRequireInitializedRunner(t *testing.T) {
|
||||
cfg := SSHConfig{Host: "203.0.113.10", Port: 22, Username: "root"}
|
||||
plan := NodePlan{NodeID: 7}
|
||||
expected := errors.New("nftables manager not initialized")
|
||||
|
||||
var nilManager *Manager
|
||||
if err := nilManager.Test(context.Background(), cfg); err == nil || err.Error() != expected.Error() {
|
||||
t.Fatalf("expected not initialized error from nil manager Test, got %v", err)
|
||||
}
|
||||
|
||||
if _, err := nilManager.Reconcile(context.Background(), cfg, plan); err == nil || err.Error() != expected.Error() {
|
||||
t.Fatalf("expected not initialized error from nil manager Reconcile, got %v", err)
|
||||
}
|
||||
|
||||
if err := nilManager.Clear(context.Background(), cfg); err == nil || err.Error() != expected.Error() {
|
||||
t.Fatalf("expected not initialized error from nil manager Clear, 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)
|
||||
}
|
||||
|
||||
if _, err := manager.Reconcile(context.Background(), cfg, plan); err == nil || err.Error() != expected.Error() {
|
||||
t.Fatalf("expected not initialized error from Reconcile, got %v", err)
|
||||
}
|
||||
|
||||
if err := manager.Clear(context.Background(), cfg); err == nil || err.Error() != expected.Error() {
|
||||
t.Fatalf("expected not initialized error from Clear, got %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net"
|
||||
"net/url"
|
||||
"strconv"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func ParseSingleTarget(raw string) (Target, error) {
|
||||
value := strings.TrimSpace(raw)
|
||||
if value == "" {
|
||||
return Target{}, fmt.Errorf("目标地址不能为空")
|
||||
}
|
||||
if strings.Contains(value, ",") || strings.Contains(value, "\n") {
|
||||
return Target{}, fmt.Errorf("nftables 纯转发第一阶段仅支持单目标")
|
||||
}
|
||||
if hasScheme(value) {
|
||||
return Target{}, fmt.Errorf("目标地址必须是 host:port,不能包含 URL scheme")
|
||||
}
|
||||
host, portText, err := net.SplitHostPort(value)
|
||||
if err != nil {
|
||||
return Target{}, fmt.Errorf("目标地址必须是 host:port")
|
||||
}
|
||||
host = strings.TrimSpace(strings.Trim(host, "[]"))
|
||||
if host == "" {
|
||||
return Target{}, fmt.Errorf("目标主机不能为空")
|
||||
}
|
||||
port, err := strconv.Atoi(portText)
|
||||
if err != nil || port < 1 || port > 65535 {
|
||||
return Target{}, fmt.Errorf("目标端口必须在 1-65535 之间")
|
||||
}
|
||||
return Target{Host: host, Port: port}, nil
|
||||
}
|
||||
|
||||
func hasScheme(value string) bool {
|
||||
parsed, err := url.Parse(value)
|
||||
if err != nil || parsed.Scheme == "" {
|
||||
return false
|
||||
}
|
||||
if strings.Contains(value, "://") {
|
||||
return true
|
||||
}
|
||||
colon := strings.IndexByte(value, ':')
|
||||
if colon <= 0 || strings.Contains(parsed.Scheme, ".") {
|
||||
return false
|
||||
}
|
||||
suffix := value[colon+1:]
|
||||
return strings.IndexByte(suffix, ':') == -1
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
package nftables
|
||||
|
||||
import "testing"
|
||||
|
||||
func TestParseSingleTargetAcceptsHostPortAndIPv6(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
raw string
|
||||
host string
|
||||
port int
|
||||
}{
|
||||
{name: "hostname", raw: "example.com:443", host: "example.com", port: 443},
|
||||
{name: "ipv4", raw: "198.51.100.20:8443", host: "198.51.100.20", port: 8443},
|
||||
{name: "ipv6", raw: "[2001:db8::1]:443", host: "2001:db8::1", port: 443},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
target, err := ParseSingleTarget(tt.raw)
|
||||
if err != nil {
|
||||
t.Fatalf("ParseSingleTarget: %v", err)
|
||||
}
|
||||
if target.Host != tt.host || target.Port != tt.port {
|
||||
t.Fatalf("expected %s/%d, got %+v", tt.host, tt.port, target)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseSingleTargetRejectsUnsupportedValues(t *testing.T) {
|
||||
for _, raw := range []string{
|
||||
"",
|
||||
"example.com",
|
||||
"example.com:0",
|
||||
"example.com:65536",
|
||||
"a:1,b:2",
|
||||
"http://example.com:443",
|
||||
"https:443",
|
||||
"mailto:443",
|
||||
} {
|
||||
t.Run(raw, func(t *testing.T) {
|
||||
if _, err := ParseSingleTarget(raw); err == nil {
|
||||
t.Fatalf("expected error for %q", raw)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,113 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"net"
|
||||
"sort"
|
||||
"strings"
|
||||
)
|
||||
|
||||
func RenderTable(plan NodePlan) string {
|
||||
var b strings.Builder
|
||||
b.WriteString("table inet flvx {\n")
|
||||
b.WriteString(" chain prerouting {\n")
|
||||
b.WriteString(" type nat hook prerouting priority dstnat; policy accept;\n")
|
||||
for _, rule := range sortedRules(plan.Rules) {
|
||||
for _, protocol := range normalizedProtocols(rule.Protocols) {
|
||||
b.WriteString(fmt.Sprintf(" %s dport %d dnat %s to %s comment \"flvx forward:%d %s\"\n",
|
||||
protocol,
|
||||
rule.InPort,
|
||||
dnatFamilyPrefix(rule.TargetHost),
|
||||
formatDNATTarget(rule.TargetHost, rule.TargetPort),
|
||||
rule.ForwardID,
|
||||
protocol,
|
||||
))
|
||||
}
|
||||
}
|
||||
b.WriteString(" }\n\n")
|
||||
b.WriteString(" chain postrouting {\n")
|
||||
b.WriteString(" type nat hook postrouting priority srcnat; policy accept;\n")
|
||||
if len(plan.Rules) > 0 {
|
||||
b.WriteString(" masquerade comment \"flvx masquerade\"\n")
|
||||
}
|
||||
b.WriteString(" }\n\n")
|
||||
b.WriteString(" chain forward {\n")
|
||||
b.WriteString(" type filter hook forward priority filter; policy accept;\n")
|
||||
b.WriteString(" }\n")
|
||||
b.WriteString("}\n")
|
||||
return b.String()
|
||||
}
|
||||
|
||||
func RuleHash(rule Rule) string {
|
||||
protocols := normalizedProtocols(rule.Protocols)
|
||||
sum := sha256.Sum256([]byte(fmt.Sprintf("%d|%d|%s|%d|%s",
|
||||
rule.ForwardID,
|
||||
rule.InPort,
|
||||
strings.TrimSpace(rule.TargetHost),
|
||||
rule.TargetPort,
|
||||
strings.Join(protocols, ","),
|
||||
)))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func PlanHashes(plan NodePlan) map[int64]string {
|
||||
hashes := make(map[int64]string, len(plan.Rules))
|
||||
for _, rule := range plan.Rules {
|
||||
hashes[rule.ForwardID] = RuleHash(rule)
|
||||
}
|
||||
return hashes
|
||||
}
|
||||
|
||||
func sortedRules(rules []Rule) []Rule {
|
||||
out := append([]Rule(nil), rules...)
|
||||
sort.SliceStable(out, func(i, j int) bool {
|
||||
if out[i].InPort == out[j].InPort {
|
||||
return out[i].ForwardID < out[j].ForwardID
|
||||
}
|
||||
return out[i].InPort < out[j].InPort
|
||||
})
|
||||
return out
|
||||
}
|
||||
|
||||
func normalizedProtocols(protocols []string) []string {
|
||||
seen := map[string]struct{}{}
|
||||
out := make([]string, 0, 2)
|
||||
for _, protocol := range protocols {
|
||||
p := strings.ToLower(strings.TrimSpace(protocol))
|
||||
if p != "tcp" && p != "udp" {
|
||||
continue
|
||||
}
|
||||
if _, ok := seen[p]; ok {
|
||||
continue
|
||||
}
|
||||
seen[p] = struct{}{}
|
||||
out = append(out, p)
|
||||
}
|
||||
if len(out) == 0 {
|
||||
return []string{"tcp", "udp"}
|
||||
}
|
||||
sort.Strings(out)
|
||||
return out
|
||||
}
|
||||
|
||||
func formatDNATTarget(host string, port int) string {
|
||||
trimmed := strings.Trim(strings.TrimSpace(host), "[]")
|
||||
if ip := net.ParseIP(trimmed); ip != nil && ip.To4() == nil {
|
||||
return fmt.Sprintf("[%s]:%d", trimmed, port)
|
||||
}
|
||||
return fmt.Sprintf("%s:%d", trimmed, port)
|
||||
}
|
||||
|
||||
func dnatFamilyPrefix(host string) string {
|
||||
trimmed := strings.Trim(strings.TrimSpace(host), "[]")
|
||||
ip := net.ParseIP(trimmed)
|
||||
if ip == nil {
|
||||
return ""
|
||||
}
|
||||
if ip.To4() != nil {
|
||||
return "ip"
|
||||
}
|
||||
return "ip6"
|
||||
}
|
||||
@@ -0,0 +1,73 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestRenderTableIncludesDNATAndMasquerade(t *testing.T) {
|
||||
script := RenderTable(NodePlan{
|
||||
NodeID: 10,
|
||||
Rules: []Rule{
|
||||
{
|
||||
ForwardID: 42,
|
||||
InPort: 24000,
|
||||
TargetHost: "198.51.100.20",
|
||||
TargetPort: 443,
|
||||
Protocols: []string{"tcp", "udp"},
|
||||
},
|
||||
},
|
||||
})
|
||||
|
||||
expectedParts := []string{
|
||||
"table inet flvx",
|
||||
"type nat hook prerouting priority dstnat; policy accept;",
|
||||
"type nat hook postrouting priority srcnat; policy accept;",
|
||||
"tcp dport 24000 dnat ip to 198.51.100.20:443 comment \"flvx forward:42 tcp\"",
|
||||
"udp dport 24000 dnat ip to 198.51.100.20:443 comment \"flvx forward:42 udp\"",
|
||||
"masquerade comment \"flvx masquerade\"",
|
||||
}
|
||||
for _, part := range expectedParts {
|
||||
if !strings.Contains(script, part) {
|
||||
t.Fatalf("script missing %q:\n%s", part, script)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestRenderTableBracketsIPv6Target(t *testing.T) {
|
||||
script := RenderTable(NodePlan{
|
||||
NodeID: 10,
|
||||
Rules: []Rule{
|
||||
{ForwardID: 42, InPort: 24000, TargetHost: "2001:db8::1", TargetPort: 443, Protocols: []string{"tcp"}},
|
||||
},
|
||||
})
|
||||
if !strings.Contains(script, "dnat ip6 to [2001:db8::1]:443") {
|
||||
t.Fatalf("expected bracketed IPv6 dnat target, got:\n%s", script)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuleHashIsStable(t *testing.T) {
|
||||
rule := Rule{ForwardID: 42, InPort: 24000, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp", "udp"}}
|
||||
if RuleHash(rule) != RuleHash(rule) {
|
||||
t.Fatalf("expected stable rule hash")
|
||||
}
|
||||
if RuleHash(rule) == RuleHash(Rule{ForwardID: 42, InPort: 24001, TargetHost: "198.51.100.20", TargetPort: 443, Protocols: []string{"tcp", "udp"}}) {
|
||||
t.Fatalf("expected hash to change when port changes")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRuleHashIgnoresBindIPWhenRenderingDoesNotUseIt(t *testing.T) {
|
||||
base := Rule{
|
||||
ForwardID: 42,
|
||||
InPort: 24000,
|
||||
TargetHost: "198.51.100.20",
|
||||
TargetPort: 443,
|
||||
Protocols: []string{"tcp", "udp"},
|
||||
}
|
||||
withBind := base
|
||||
withBind.BindIP = "192.0.2.10"
|
||||
|
||||
if RuleHash(base) != RuleHash(withBind) {
|
||||
t.Fatalf("expected bind IP to be ignored by hash when it is not rendered")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,194 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"golang.org/x/crypto/ssh"
|
||||
)
|
||||
|
||||
type Runner interface {
|
||||
ApplyScript(ctx context.Context, cfg SSHConfig, script string) error
|
||||
Test(ctx context.Context, cfg SSHConfig) error
|
||||
}
|
||||
|
||||
type SSHRunner struct {
|
||||
Timeout time.Duration
|
||||
}
|
||||
|
||||
func NewSSHRunner() *SSHRunner {
|
||||
return &SSHRunner{Timeout: 15 * time.Second}
|
||||
}
|
||||
|
||||
func (r *SSHRunner) Test(ctx context.Context, cfg SSHConfig) error {
|
||||
return r.run(ctx, cfg, "command -v nft >/dev/null 2>&1 && nft --version >/dev/null 2>&1")
|
||||
}
|
||||
|
||||
func (r *SSHRunner) ApplyScript(ctx context.Context, cfg SSHConfig, script string) error {
|
||||
nft := nftBinary(cfg)
|
||||
command := "tmp=$(mktemp /tmp/flvx-nft-XXXXXX.nft) || exit 1\n" +
|
||||
"cleanup() {\n" +
|
||||
" rm -f \"$tmp\"\n" +
|
||||
"}\n" +
|
||||
"trap cleanup EXIT\n" +
|
||||
"cat > \"$tmp\" <<'EOF'\n" + script + "\nEOF\n" +
|
||||
nft + " -c -f \"$tmp\"\n" +
|
||||
"if " + nft + " list table inet flvx >/dev/null 2>&1; then\n" +
|
||||
" " + nft + " delete table inet flvx\n" +
|
||||
"fi\n" +
|
||||
nft + " -f \"$tmp\""
|
||||
return r.run(ctx, cfg, command)
|
||||
}
|
||||
|
||||
func (r *SSHRunner) run(ctx context.Context, cfg SSHConfig, command string) error {
|
||||
timeout := r.Timeout
|
||||
if timeout <= 0 {
|
||||
timeout = 15 * time.Second
|
||||
}
|
||||
runCtx, cancel := context.WithTimeout(ctx, timeout)
|
||||
defer cancel()
|
||||
|
||||
clientConfig, err := buildSSHClientConfig(cfg)
|
||||
if err != nil {
|
||||
return 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)
|
||||
}
|
||||
defer conn.Close()
|
||||
|
||||
sshConn, chans, reqs, err := ssh.NewClientConn(conn, addr, clientConfig)
|
||||
if err != nil {
|
||||
return 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)
|
||||
}
|
||||
defer session.Close()
|
||||
|
||||
var stderr bytes.Buffer
|
||||
session.Stderr = &stderr
|
||||
|
||||
done := make(chan error, 1)
|
||||
go func() {
|
||||
done <- session.Run(command)
|
||||
}()
|
||||
|
||||
select {
|
||||
case <-runCtx.Done():
|
||||
_ = session.Close()
|
||||
return 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 fmt.Errorf("远程执行失败: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
|
||||
func buildSSHClientConfig(cfg SSHConfig) (*ssh.ClientConfig, error) {
|
||||
if strings.TrimSpace(cfg.Host) == "" {
|
||||
return nil, fmt.Errorf("SSH 主机不能为空")
|
||||
}
|
||||
if strings.TrimSpace(cfg.Username) == "" {
|
||||
return nil, fmt.Errorf("SSH 用户名不能为空")
|
||||
}
|
||||
|
||||
auth, err := authMethods(cfg)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if len(auth) == 0 {
|
||||
return nil, fmt.Errorf("SSH 认证方式不能为空")
|
||||
}
|
||||
|
||||
return &ssh.ClientConfig{
|
||||
User: strings.TrimSpace(cfg.Username),
|
||||
Auth: auth,
|
||||
HostKeyCallback: ssh.InsecureIgnoreHostKey(),
|
||||
Timeout: 15 * time.Second,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func authMethods(cfg SSHConfig) ([]ssh.AuthMethod, error) {
|
||||
switch strings.ToLower(strings.TrimSpace(cfg.AuthType)) {
|
||||
case "":
|
||||
if strings.TrimSpace(cfg.PrivateKey) == "" {
|
||||
return nil, fmt.Errorf("SSH 私钥不能为空")
|
||||
}
|
||||
signer, err := parsePrivateKey(cfg.PrivateKey, cfg.Passphrase)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return []ssh.AuthMethod{ssh.PublicKeys(signer)}, nil
|
||||
case "password":
|
||||
if cfg.Password == "" {
|
||||
return nil, fmt.Errorf("SSH 密码不能为空")
|
||||
}
|
||||
return []ssh.AuthMethod{ssh.Password(cfg.Password)}, nil
|
||||
case "private_key":
|
||||
if strings.TrimSpace(cfg.PrivateKey) == "" {
|
||||
return nil, fmt.Errorf("SSH 私钥不能为空")
|
||||
}
|
||||
signer, err := parsePrivateKey(cfg.PrivateKey, cfg.Passphrase)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return []ssh.AuthMethod{ssh.PublicKeys(signer)}, nil
|
||||
default:
|
||||
return nil, fmt.Errorf("不支持的 SSH 认证方式: %s", cfg.AuthType)
|
||||
}
|
||||
}
|
||||
|
||||
func parsePrivateKey(privateKey, passphrase string) (ssh.Signer, error) {
|
||||
if passphrase != "" {
|
||||
signer, err := ssh.ParsePrivateKeyWithPassphrase([]byte(privateKey), []byte(passphrase))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("SSH 私钥解析失败: %w", err)
|
||||
}
|
||||
return signer, nil
|
||||
}
|
||||
signer, err := ssh.ParsePrivateKey([]byte(privateKey))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("SSH 私钥解析失败: %w", err)
|
||||
}
|
||||
return signer, nil
|
||||
}
|
||||
|
||||
func nftCommand(cfg SSHConfig, command string) string {
|
||||
return "sh -lc " + sshQuote(command)
|
||||
}
|
||||
|
||||
func nftBinary(cfg SSHConfig) string {
|
||||
if strings.EqualFold(strings.TrimSpace(cfg.SudoMode), "sudo") {
|
||||
return "sudo -n nft"
|
||||
}
|
||||
return "nft"
|
||||
}
|
||||
|
||||
func sshQuote(value string) string {
|
||||
return "'" + strings.ReplaceAll(value, "'", "'\"'\"'") + "'"
|
||||
}
|
||||
|
||||
func normalizedSSHPort(port int) int {
|
||||
if port <= 0 {
|
||||
return 22
|
||||
}
|
||||
return port
|
||||
}
|
||||
@@ -0,0 +1,42 @@
|
||||
package nftables
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"crypto/x509"
|
||||
"encoding/pem"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestAuthMethodsDefaultToPrivateKey(t *testing.T) {
|
||||
privateKey := mustGeneratePrivateKey(t)
|
||||
methods, err := authMethods(SSHConfig{PrivateKey: privateKey})
|
||||
if err != nil {
|
||||
t.Fatalf("authMethods: %v", err)
|
||||
}
|
||||
if len(methods) != 1 {
|
||||
t.Fatalf("expected 1 auth method, got %d", len(methods))
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthMethodsDefaultPrivateKeyRequiresKey(t *testing.T) {
|
||||
_, err := authMethods(SSHConfig{})
|
||||
if err == nil || !strings.Contains(err.Error(), "SSH 私钥不能为空") {
|
||||
t.Fatalf("expected private key required error, got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func mustGeneratePrivateKey(t *testing.T) string {
|
||||
t.Helper()
|
||||
|
||||
key, err := rsa.GenerateKey(rand.Reader, 1024)
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateKey: %v", err)
|
||||
}
|
||||
block := &pem.Block{
|
||||
Type: "RSA PRIVATE KEY",
|
||||
Bytes: x509.MarshalPKCS1PrivateKey(key),
|
||||
}
|
||||
return string(pem.EncodeToMemory(block))
|
||||
}
|
||||
@@ -0,0 +1,46 @@
|
||||
package nftables
|
||||
|
||||
const (
|
||||
ModeAgent = "agent"
|
||||
ModeNftables = "nftables"
|
||||
|
||||
StatusPending = "pending"
|
||||
StatusApplied = "applied"
|
||||
StatusError = "error"
|
||||
)
|
||||
|
||||
type Target struct {
|
||||
Host string
|
||||
Port int
|
||||
}
|
||||
|
||||
type Rule struct {
|
||||
ForwardID int64
|
||||
InPort int
|
||||
BindIP string
|
||||
TargetHost string
|
||||
TargetPort int
|
||||
Protocols []string
|
||||
}
|
||||
|
||||
type NodePlan struct {
|
||||
NodeID int64
|
||||
Rules []Rule
|
||||
}
|
||||
|
||||
type SSHConfig struct {
|
||||
Host string
|
||||
Port int
|
||||
Username string
|
||||
AuthType string
|
||||
Password string
|
||||
PrivateKey string
|
||||
Passphrase string
|
||||
SudoMode string
|
||||
}
|
||||
|
||||
type ApplyResult struct {
|
||||
NodeID int64
|
||||
Script string
|
||||
Hashes map[int64]string
|
||||
}
|
||||
Reference in New Issue
Block a user